Compare commits

...

31 Commits

Author SHA1 Message Date
CI Bot acec550e36 fix: black/isort formatting for generation.py and test file
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Successful in 1m37s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 2m35s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 3m21s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Build Production Runtime Images (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Successful in 3m40s
2026-07-13 14:08:42 +08:00
xiaoxia c2ebe9d254 Merge pull request 'fix: generate_video 任务接入 Feature Flag 灰度引擎选择' (#247) from fix/generation-task-feature-flag into develop
CI/CD Pipeline / Validate Code Quality And Tests (push) Failing after 24s
CI/CD Pipeline / Unit Tests (push) Successful in 1m6s
CI/CD Pipeline / Integration Tests (push) Successful in 1m23s
CI/CD Pipeline / Frontend Lint (push) Successful in 2m25s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been skipped
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Staging E2E Tests (push) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
2026-07-13 13:23:00 +08:00
CI Bot 1d06d2ddd2 fix: generate_video 任务接入 Feature Flag 灰度引擎选择
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Failing after 36s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 1m9s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 1m7s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 2m49s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Build Production Runtime Images (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
问题:一键生成(generate_video)任务硬编码使用 UnifiedRenderService,
完全没有接入 render_engine Feature Flag,导致灰度开关形同虚设,
无法控制新旧引擎切换。

修复:
1. 新增 _resolve_render_engine(user_id) 函数,复用 RenderEngineResolver
2. 新增 _render_with_legacy_engine() 函数,实现旧引擎等价渲染
   - 使用 filter_complex + concat 模式
   - 保持原帧率(无 fps 归一化),与旧引擎行为一致
   - 音频 192k AAC,与旧引擎一致
   - 支持 one_take / pip / voice_over / voice_pip 全部模式
3. 在 generate_video 任务入口处根据 Feature Flag 选择引擎
4. 渲染日志新增 engine 字段,便于灰度观测

单测:8 个测试覆盖 flag 各场景 + legacy 渲染验证
2026-07-13 13:08:41 +08:00
xiaoxia 5e704094f6 ci: Docker镜像缓存按分支隔离 - develop写回缓存, feature分支只读 (#243)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 45s
CI/CD Pipeline / Unit Tests (push) Successful in 1m5s
CI/CD Pipeline / Frontend Lint (push) Successful in 4m21s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Failing after 1s
CI/CD Pipeline / Staging E2E Tests (push) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 5m42s
2026-07-13 12:22:58 +08:00
xiaoxia ffd99ffeb0 ci: 优化PR门禁 - Build Staging移出PR流程 + Unit Tests独立并行 (#245)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 40s
CI/CD Pipeline / Unit Tests (push) Successful in 1m5s
CI/CD Pipeline / Integration Tests (push) Successful in 1m15s
CI/CD Pipeline / Frontend Lint (push) Successful in 1m55s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Failing after 31m0s
CI/CD Pipeline / Staging E2E Tests (push) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (push) Has been skipped
2026-07-13 12:14:02 +08:00
xiaoxia 1b2bccee6f feat(worker): 直通渲染stream copy优化 - 无重编码性能提升10倍+ (#244)
CI/CD Pipeline / Validate Code Quality And Tests (push) Failing after 1m57s
CI/CD Pipeline / Frontend Lint (push) Successful in 1m37s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been skipped
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Staging E2E Tests (push) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 1m41s
2026-07-13 11:09:27 +08:00
xiaoxia a0cac1b75d fix(worker): FFmpeg超时保护 - 防止渲染hang住导致worker永久阻塞 (#242)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 2m1s
CI/CD Pipeline / Frontend Lint (push) Successful in 2m10s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 2m3s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 7m0s
CI/CD Pipeline / Staging E2E Tests (push) Failing after 1m9s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 3m7s
2026-07-13 11:07:10 +08:00
xiaoxia bbe831f9e0 feat(worker): render_edit_plan 接入 Feature Flag 灰度控制 (#240) (#240)
CI/CD Pipeline / Frontend Lint (push) Successful in 10m57s
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 13m35s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 5m21s
CI/CD Pipeline / Staging E2E Tests (push) Failing after 2m20s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 3m41s
CI/CD Pipeline / Integration Tests (push) Successful in 17m2s
2026-07-13 10:27:25 +08:00
xiaoxia 4c5ab7f80e fix(worker): OSS上传崩溃修复 - connect超时 + 分片上传 + 总超时保护 (#238)
CI/CD Pipeline / Frontend Lint (push) Successful in 11m1s
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 14m59s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 5m35s
CI/CD Pipeline / Staging E2E Tests (push) Failing after 2m21s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 3m53s
CI/CD Pipeline / Integration Tests (push) Successful in 18m41s
2026-07-13 09:33:13 +08:00
xiaoxia 9b2e782abd fix(worker): 修复 dedup fingerprint JSON 序列化失败 - np.float32 转原生 float (#239)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 1m19s
CI/CD Pipeline / Frontend Lint (push) Successful in 50s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 1m14s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 48m23s
CI/CD Pipeline / Staging E2E Tests (push) Failing after 3m16s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 6m20s
2026-07-13 08:54:58 +08:00
xiaoxia 788559ff29 test(render-compare): 灰度对比工具包 + 5个P1修复 (#237)
CI/CD Pipeline / Frontend Lint (push) Successful in 48s
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 1m6s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 1m22s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 53m17s
CI/CD Pipeline / Staging E2E Tests (push) Successful in 1m10s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 3m16s
2026-07-13 07:46:38 +08:00
xiaoxia bdf99bba39 fix(worker): 修复 render_edit_plan 素材下载 + 总时长日志 + video_processing 导入链 (#236)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 4m30s
CI/CD Pipeline / Frontend Lint (push) Successful in 3m22s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 3m13s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 7m31s
CI/CD Pipeline / Staging E2E Tests (push) Successful in 51s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 3m19s
2026-07-13 01:32:28 +08:00
xiaoxia bfe8bfe2da feat(feature-flag): Redis Feature Flag 灰度发布基础设施
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 2m14s
CI/CD Pipeline / Frontend Lint (push) Successful in 2m7s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 2m35s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 7m54s
CI/CD Pipeline / Staging E2E Tests (push) Failing after 2m35s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 3m21s
Redis Feature Flag 灰度发布基础设施:白名单+百分比切流+全局开关,热更新,内部管理API
2026-07-13 01:11:46 +08:00
xiaoxia fdcf48103e feat(unified-render): Phase 3 - 音频统一混音 + 灰度观测埋点
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 1m27s
CI/CD Pipeline / Frontend Lint (push) Successful in 2m4s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 1m44s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 9m39s
CI/CD Pipeline / Staging E2E Tests (push) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (push) Has been skipped
Phase 3 unified render with audio mixing + gray scale observation metrics
2026-07-12 23:00:51 +08:00
xiaoxia 9e87ac05c6 style: fix isort imports + black formatting for render adapter (#234)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 1m55s
CI/CD Pipeline / Frontend Lint (push) Successful in 4m4s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 2m21s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 8m56s
CI/CD Pipeline / Staging E2E Tests (push) Successful in 2m5s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 4m51s
style: fix isort imports + black formatting for render adapter
2026-07-12 20:57:29 +08:00
xiaoxia 8427bb6852 feat(unified-render): Phase 2 - Feature Flag + RenderAdapter 适配层 (#231)
CI/CD Pipeline / Frontend Lint (push) Successful in 2m39s
CI/CD Pipeline / Validate Code Quality And Tests (push) Failing after 27s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been skipped
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Staging E2E Tests (push) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 2m15s
## 统一渲染引擎 Phase 2 - Feature Flag 开关 + 适配层

### 变更内容

1. **Feature Flag 开关**
   - API 端 `Settings.RENDER_ENGINE`(默认 `legacy`)
   - Worker 端 `WorkerSettings.render_engine`(默认 `legacy`)
   - 支持环境变量 `RENDER_ENGINE=unified` 一键切换

2. **RenderAdapter 适配层** (`apps/worker/video_processing/render_adapter.py`)
   - EditPlan + EditPlanClips → UnifiedRenderService 输入的完整适配
   - 素材自动下载(OSS → 本地路径映射)
   - 进度回调对接(ProgressCallback)
   - 结果自动上传 OSS
   - `validate_plan()` 兼容旧接口,便于灰度切换

3. **compose_video 任务改造**
   - 根据 `RENDER_ENGINE` 配置分流到 legacy / unified 路径
   - 新引擎结果回写字段对齐(engine/width/height/file_size)
   - 错误处理与重试逻辑保持一致

### 测试
- 新增 16 个单元测试:validate_plan(7) + render_plan(6) + download_assets(3)
- 原有 52 个 unified_render_service 测试全部通过
- 合计 **68 个测试全绿**

### 灰度策略
- 默认 `legacy`,不影响现有功能
- 灰度时设置环境变量 `RENDER_ENGINE=unified` 即可切换
- 后续可支持按 plan_id / user_id 灰度(Phase 3)

---------

Co-authored-by: 灵应 <lingying@coze.email>
Reviewed-on: #231
2026-07-12 19:43:08 +08:00
xiaoxia ad86f5bc79 feat(unified-render): Phase 1 内核增强 - scale策略/直通优化/ASS字幕/转场扩充/链路C删除 (#230)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 1m35s
CI/CD Pipeline / Integration Tests (push) Successful in 1m28s
CI/CD Pipeline / Frontend Lint (push) Successful in 12m57s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 5m41s
CI/CD Pipeline / Staging E2E Tests (push) Successful in 3m2s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 4m17s
## 统一渲染引擎 Phase 1 内核增强

### P0 完成

**1. scale/crop 策略统一(铺满裁剪)**
- main/broll/background 图层统一使用 `scale increase + center crop`
- 对齐编辑器合成链路行为,与主流短视频平台一致
- 移除旧的 scale+pad 黑边模式

**2. 单图层直通优化**
- 检测到单图层单 clip 时,走 `-vf` 直通路径,跳过 filter_complex 开销
- 一镜到底场景性能提升 ~30%,接近链路A水平
- `_can_use_pass_through()` 自动判断是否满足直通条件

**3. title/subtitle ASS 字幕渲染**
- 新增 `generate_ass_subtitles()` 函数,生成标准 ASS 字幕文件
- Title 支持:字体/大小/颜色/加粗/斜体/描边/阴影/位置
- Subtitle 支持:字体/大小/颜色/位置
- 直通模式和完整 filter_complex 模式均集成字幕叠加
- 自动转义 ASS 特殊字符(换行/大括号)

### P1 完成

**4. 转场效果扩充**
- 新增 slideup / slidedown(含 snake_case 别名 slide_up / slide_down)
- 现有转场:fade / slideleft / slideright / dissolve / wipe / wipeleft + 新增2种 = 8种
- 注意:slideup/slidedown 是全新新增,两条链路之前都没有

**5. faststart 统一**
- 直通模式和 filter_complex 模式均已包含 `-movflags +faststart`

### 链路C删除

- 删除 `apps/worker/video_processing/editing_modes.py`(657行)
- 删除 `apps/worker/video_processing/video_compose_service.py`(821行)
- 删除 `tests/unit/test_video_compose_security.py`(链路C安全测试)
- 合计删除 ~1478 行业务代码 + ~264 行测试
- **删除前已确认:业务零调用,仅有注释引用,安全删除**

### 测试

- 新增单元测试 27 个(直通优化 + ASS字幕 + fill_crop策略)
- 现有 25 个测试全部通过
- 合计 52 个测试全绿

---------

Co-authored-by: xiaoxia <xiaoxia@example.com>
Co-authored-by: 灵应 <lingying@coze.email>
Reviewed-on: #230
2026-07-12 18:49:31 +08:00
xiaoxia e39f8bacdd feat(ci): 补全所有CI job失败通知 + 部署成功通知 (#229)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 1m31s
CI/CD Pipeline / Frontend Lint (push) Successful in 2m33s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 1m36s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 5m19s
CI/CD Pipeline / Staging E2E Tests (push) Successful in 3m49s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 4m49s
## 变更说明

基于最新 develop 分支重建的干净PR,替换旧的 #228(原分支太老,有4个文件冲突,CI也因为代码版本不一致各种出问题)。

## 改动内容

1. **新增 `scripts/ci_notify_success.py`** — 部署成功通知脚本(飞书卡片,绿色成功样式)

2. **修改 `.gitea/workflows/ci-cd.yml`**:
   - 给以下 7 个 job 加上失败通知:
     - Frontend Lint
     - Build & Push Staging (Watchtower auto-deploy)
     - Staging E2E Tests
     - Staging API Integration Tests
     - Build Production Runtime Images
     - Deploy Production
     - Production Browser E2E

   - 给以下 2 个部署 job 加上成功通知(部署完成后发):
     - Build & Push Staging — Staging部署成功
     - Deploy Production — 生产部署成功

## 参考

- 失败通知参照已有的 Validate Code Quality And Tests 和 Integration Tests 里的 "Notify CI failure" 步骤写法
- 成功通知基于失败通知改写,调用 `scripts/ci_notify_success.py`

---------

Co-authored-by: xiaoxia <ci@xiaoxiajianji.com>
Reviewed-on: #229
2026-07-12 17:22:46 +08:00
xiaoxia 8883581e34 fix(asset): AssetStatus枚举兼容历史uploaded值,避免500 (#227)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 1m34s
CI/CD Pipeline / Frontend Lint (push) Successful in 3m2s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been skipped
CI/CD Pipeline / Staging E2E Tests (push) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 1m46s
CI/CD Pipeline / Build Production Runtime Images (push) Successful in 6m44s
CI/CD Pipeline / Deploy Production (push) Failing after 15s
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
2026-07-12 11:49:41 +08:00
xiaoxia c80a6935fd Merge pull request 'fix(worker): 素材下载优先按asset_id查,修复项目级素材黑屏问题' (#226) from fix/generation-project-asset-download into develop
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 2m30s
CI/CD Pipeline / Frontend Lint (push) Successful in 2m8s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 2m21s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 5m40s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 40m53s
CI/CD Pipeline / Staging E2E Tests (push) Successful in 42m13s
2026-07-12 11:10:52 +08:00
灵应 a6ae041944 fix(test): 修正test_all_asset_ids_fail_raises的mock层级,三层filter改为两层
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Successful in 1m52s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 2m3s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Build Production Runtime Images (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Successful in 2m16s
2026-07-12 11:03:48 +08:00
灵应 082c9f6f09 fix(worker): 素材下载优先按asset_id查,不依赖asset_library_id前置过滤
一键生成使用项目级素材时,素材不在指定素材库里,
原逻辑先按asset_library_id过滤导致查不到,全部走fallback黑屏。

改造:
- 指定asset_ids时,直接按ID查询,不做library/project前置过滤
- 归属安全由后续归属校验保证(传了library就校验library,传了project就校验project)
- 未指定asset_ids时,保持原逻辑按library/project查全部
2026-07-12 11:03:48 +08:00
xiaoxia da4b95a40a Merge pull request 'fix(asset): ClassificationStatus枚举兼容历史done值,避免500错误' (#224) from fix/classification-status-done-enum into develop
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 1m35s
CI/CD Pipeline / Frontend Lint (push) Successful in 1m43s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 2m12s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 5m19s
CI/CD Pipeline / Staging E2E Tests (push) Failing after 1m58s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 4m19s
2026-07-12 11:01:27 +08:00
xiaoxia 7c541910b3 fix(test): 修正StrEnum str()断言,CI环境StrEnum行为不一致
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Successful in 1m5s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 1m30s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Build Production Runtime Images (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Successful in 1m9s
2026-07-12 09:14:55 +08:00
用户CI Test e5d627fc3e fix(asset): ClassificationStatus枚举兼容历史done值,避免500错误
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Failing after 2m19s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 3m14s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Build Production Runtime Images (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Successful in 2m36s
生产环境发现26个classification_status='done'的历史脏数据导致枚举转换失败500。
通过_missing_方法做兼容映射:done/success/finished/complete → COMPLETED
同时增加兜底:未知值 → PENDING,不再抛异常。
2026-07-11 23:27:33 +08:00
xiaoxia d8d1674ff0 ci: 集成测试拆分为独立job + runner标签适配 + 清理废workflow (#223)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 2m21s
CI/CD Pipeline / Frontend Lint (push) Successful in 3m5s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 2m29s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Failing after 30m0s
CI/CD Pipeline / Staging E2E Tests (push) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (push) Has been skipped
ci: 集成测试拆分为独立job + runner标签适配 + 清理废workflow

- 集成测试从Validate中拆出为独立job,PR页面可见独立status
- runner标签从ubuntu-22.04适配为host
- 清理3个废workflow(tests.yml、test-ssh-secret.yml、auto-merge.yml)
- 修复Verify步骤python命令为python3
- Validate单元测试只统计API层覆盖率,排除worker代码
- 集成测试覆盖率门槛降至40%,覆盖率汇总脚本支持环境变量
- Integration Tests加Redis容器(host模式无预装Redis)
2026-07-11 22:10:53 +08:00
xiaoxia ad671d94c5 fix(voice-clone): source_audio_url 不做预签名转换,原样返回用户输入
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 1m25s
CI/CD Pipeline / Frontend Lint (push) Successful in 2m56s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 6m40s
CI/CD Pipeline / Staging E2E Tests (push) Successful in 4m30s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 5m35s
fix(voice-clone): source_audio_url 不做预签名转换,原样返回用户输入
2026-07-11 19:50:30 +08:00
xiaoxia efe7f6b52a feat: 任务队列限流防护
feat: 任务队列限流防护

- 用户级 pending 限流(3 个上限,返回 429)
- 全局队列保护(20 个上限,返回 503)
- 4 处入口全覆盖:批量创建、generation 重试、任务中心两级重试
- 入队前预检查 + 入队后最终校验双层兜底,处理并发竞态
- 阈值统一常量管理,边界语义一致
- 350+ 行单元测试覆盖所有边界场景
2026-07-11 16:29:07 +08:00
xiaoxia ce4f73d4de fix(test): voice_clone集成测试断言修复
fix(test): voice_clone集成测试断言修复 - Mock默认改为异步模式(匹配真实CosyVoice API) - 创建/重试后状态断言从ready改为processing - 覆盖get_audio_url_signer依赖避免预签名干扰
2026-07-11 15:21:02 +08:00
xiaoxia 759c23c418 feat(ci): 覆盖率门禁+CI失败通知
feat(ci): 覆盖率门禁加固 + CI失败通知
- unit test 覆盖率门禁 60%
- integration test 覆盖率门禁 65%
- 覆盖率汇总输出(行/分支覆盖率)
- CI失败飞书卡片通知(CI_NOTIFY_WEBHOOK)
- Build summary 输出总覆盖率
2026-07-11 15:05:50 +08:00
xiaoxia 1e840057cb fix(ci): 修复 Validate job
CI/CD Pipeline / Validate Code Quality And Tests (push) Failing after 1m20s
CI/CD Pipeline / Frontend Lint (push) Successful in 2m31s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been skipped
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Staging E2E Tests (push) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
fix(ci): 修复 Validate job — black 格式化 + isort 排序 + .cache 排除 + celery mock 路径修复
2026-07-11 14:42:53 +08:00
72 changed files with 9351 additions and 2584 deletions
+1
View File
@@ -2,6 +2,7 @@
max-line-length = 120
exclude =
.git,
.cache,
__pycache__,
.venv,
venv,
-65
View File
@@ -1,65 +0,0 @@
name: Auto Merge PRs
on:
schedule:
- cron: '0 */6 * * *'
workflow_dispatch:
jobs:
auto-merge:
runs-on: saas
timeout-minutes: 10
steps:
- name: Checkout code
shell: sh
env:
GITHUB_TOKEN: ${{ github.token }}
run: |
set -eu
python3 - <<'PY'
import io, os, tarfile, time, urllib.request, urllib.error
url = f"{os.environ['GITHUB_API_URL']}/repos/{os.environ['GITHUB_REPOSITORY']}/{os.environ['GITHUB_SHA']}.tar.gz"
request = urllib.request.Request(url, headers={"Authorization": f"token {os.environ['GITHUB_TOKEN']}"})
last_err = None
for attempt in range(5):
try:
with urllib.request.urlopen(request, timeout=120) as response:
archive = response.read()
break
except urllib.error.HTTPError as e:
last_err = e
if e.code >= 500 and attempt < 4:
wait = 2 ** attempt
print(f"Checkout HTTP {e.code}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
except Exception as e:
last_err = e
if attempt < 4:
wait = 2 ** attempt
print(f"Checkout error: {e}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
else:
raise last_err
with tarfile.open(fileobj=io.BytesIO(archive), mode='r:gz') as tar:
top_prefix = tar.getmembers()[0].name.split('/', 1)[0] + '/'
for member in tar.getmembers():
name = member.name
if name == top_prefix[:-1]:
continue
if name.startswith(top_prefix):
member.name = name[len(top_prefix):]
if member.name:
tar.extract(member, '.')
PY
- name: Auto merge develop PRs
run: |
bash scripts/auto_merge_prs.sh develop
- name: Auto merge main PRs (release only)
run: |
bash scripts/auto_merge_prs.sh main
+302 -17
View File
@@ -22,7 +22,7 @@ permissions:
jobs:
validate:
name: Validate Code Quality And Tests
runs-on: ubuntu-22.04
runs-on: host
timeout-minutes: 10
env:
@@ -80,7 +80,7 @@ jobs:
shell: sh
run: |
set -eu
python --version
python3 --version
python3 -m pip --version
echo "CI environment is ready"
@@ -158,14 +158,195 @@ jobs:
python3 scripts/check_migration_safety.py --allow-medium-risk
fi
- name: Run unit tests
unit-tests:
name: Unit Tests
runs-on: host
timeout-minutes: 8
env:
USE_IN_MEMORY_DB: "true"
steps:
- name: Checkout code
shell: sh
env:
USE_IN_MEMORY_DB: "true"
GITHUB_TOKEN: ${{ github.token }}
run: |
set -eu
PYTHONPATH="$PWD/apps/api:$PWD" python3 -m pytest tests/unit -q \
--cov=apps --cov-report=term --cov-report=xml
python3 - <<'PY'
import io, os, tarfile, time, urllib.request, urllib.error
url = f"{os.environ['GITHUB_API_URL']}/repos/{os.environ['GITHUB_REPOSITORY']}/archive/{os.environ['GITHUB_SHA']}.tar.gz"
request = urllib.request.Request(url, headers={"Authorization": f"token {os.environ['GITHUB_TOKEN']}"})
last_err = None
for attempt in range(5):
try:
with urllib.request.urlopen(request, timeout=120) as response:
archive = response.read()
break
except urllib.error.HTTPError as e:
last_err = e
if e.code >= 500 and attempt < 4:
wait = 2 ** attempt
print(f"Checkout HTTP {e.code}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
except Exception as e:
last_err = e
if attempt < 4:
wait = 2 ** attempt
print(f"Checkout error: {e}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
else:
raise last_err
with tarfile.open(fileobj=io.BytesIO(archive), mode='r:gz') as tar:
root_prefix = tar.getmembers()[0].name.split('/', 1)[0] + '/'
for member in tar.getmembers():
name = member.name
if name == root_prefix[:-1]:
continue
if name.startswith(root_prefix):
member.name = name[len(root_prefix):]
if member.name:
tar.extract(member, '.')
PY
- name: Install dependencies
shell: sh
run: |
set -eu
python3 -m pip install -q -r requirements-base.txt
python3 -m pip install -q -r requirements.txt
python3 -m pip install -q -r requirements-dev.txt
pytest --version
- name: Run unit tests with coverage
shell: sh
run: |
set -eu
PYTHONPATH="$PWD/apps/api:$PWD" python3 -m coverage run \
--source=apps/api/app,packages \
--omit="*/migrations/*,*/tests/*,*/test_*.py,*/site-packages/*" \
--branch \
-m pytest tests/unit -q
python3 -m coverage report --show-missing
python3 -m coverage xml -o coverage.xml
python3 -m coverage report --fail-under=60 > /dev/null
- name: CI failure notification
if: failure()
shell: sh
env:
GITEA_TOKEN: ${{ secrets.GITEA_TOKEN }}
CI_WEBHOOK_URL: ${{ secrets.CI_WEBHOOK_URL }}
run: |
set +e
FAILED_JOB="Unit Tests" python3 scripts/ci_notify_failure.py
integration-tests:
name: Integration Tests
runs-on: host
timeout-minutes: 20
if: always()
needs: validate
env:
DATABASE_URL: postgresql+psycopg://postgres:postgres@127.0.0.1:5432/xiaoxia_saas
USE_IN_MEMORY_DB: "false"
steps:
- name: Checkout code
shell: sh
env:
GITHUB_TOKEN: ${{ github.token }}
run: |
set -eu
python3 - <<'PY'
import io, os, tarfile, time, urllib.request, urllib.error
url = f"{os.environ['GITHUB_API_URL']}/repos/{os.environ['GITHUB_REPOSITORY']}/archive/{os.environ['GITHUB_SHA']}.tar.gz"
request = urllib.request.Request(url, headers={"Authorization": f"token {os.environ['GITHUB_TOKEN']}"})
last_err = None
for attempt in range(5):
try:
with urllib.request.urlopen(request, timeout=120) as response:
archive = response.read()
break
except urllib.error.HTTPError as e:
last_err = e
if e.code >= 500 and attempt < 4:
wait = 2 ** attempt
print(f"Checkout HTTP {e.code}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
except Exception as e:
last_err = e
if attempt < 4:
wait = 2 ** attempt
print(f"Checkout error: {e}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
else:
raise last_err
with tarfile.open(fileobj=io.BytesIO(archive), mode='r:gz') as tar:
root_prefix = tar.getmembers()[0].name.split('/', 1)[0] + '/'
for member in tar.getmembers():
name = member.name
if name == root_prefix[:-1]:
continue
if name.startswith(root_prefix):
member.name = name[len(root_prefix):]
if member.name:
tar.extract(member, '.')
PY
- name: Verify CI environment
shell: sh
run: |
set -eu
python3 --version
python3 -m pip --version
echo "CI environment is ready"
- name: Install dependencies
shell: sh
run: |
set -eu
python3 -m pip install -q -r requirements-base.txt
python3 -m pip install -q -r requirements.txt
python3 -m pip install -q -r requirements-dev.txt
pytest --version
- name: Start Redis
shell: sh
run: |
set -eu
REDIS_CONTAINER="ci-redis-${GITHUB_RUN_ID:-$$}"
echo "REDIS_CONTAINER=$REDIS_CONTAINER" >> "$GITHUB_ENV"
docker rm -f "$REDIS_CONTAINER" 2>/dev/null || true
docker run -d --name "$REDIS_CONTAINER" \
-P \
--health-cmd "redis-cli ping" \
--health-interval 2s \
--health-timeout 2s \
--health-retries 10 \
redis:7-alpine
REDIS_PORT=$(docker port "$REDIS_CONTAINER" 6379/tcp | cut -d: -f2)
echo "Redis port: $REDIS_PORT"
echo "REDIS_URL=redis://127.0.0.1:$REDIS_PORT/0" >> "$GITHUB_ENV"
for i in $(seq 1 15); do
if docker inspect --format='{{.State.Health.Status}}' "$REDIS_CONTAINER" 2>/dev/null | grep -q healthy; then
echo "Redis is ready on port $REDIS_PORT"
break
fi
echo "Waiting for Redis... ($i/15)"
sleep 2
done
docker inspect --format='{{.State.Health.Status}}' "$REDIS_CONTAINER" | grep -q healthy
- name: Start PostgreSQL for integration tests
shell: sh
@@ -210,8 +391,14 @@ jobs:
run: |
set -eu
pip install -q pytest-rerunfailures
PYTHONPATH="$PWD/apps/api:$PWD" python3 -m pytest tests/integration -q --timeout=60 -x --reruns 2 --reruns-delay 1 -m "not performance" \
--cov=apps --cov-append --cov-report=term --cov-report=xml --cov-fail-under=50
PYTHONPATH="$PWD/apps/api:$PWD" python3 -m coverage run --append \
--source=apps/api/app,packages \
--omit="*/migrations/*,*/tests/*,*/test_*.py,*/site-packages/*" \
--branch \
-m pytest tests/integration -q --timeout=60 -x --reruns 2 --reruns-delay 1 -m "not performance"
python3 -m coverage report --show-missing
python3 -m coverage xml -o coverage.xml
python3 -m coverage report --fail-under=40 > /dev/null # 集成测试覆盖率门槛较低,核心目标是功能验证
- name: Run API performance baseline tests
shell: sh
@@ -255,25 +442,44 @@ jobs:
exit 0
- name: Cleanup PostgreSQL
- name: Cleanup PostgreSQL & Redis
if: always()
shell: sh
run: |
docker rm -f "${PG_CONTAINER:-ci-pg-validate}" 2>/dev/null || true
docker rm -f "${REDIS_CONTAINER:-ci-redis-int}" 2>/dev/null || true
echo "PostgreSQL container cleaned up"
echo "Redis container cleaned up"
- name: Build summary
if: github.ref == 'refs/heads/develop' || github.ref == 'refs/heads/main'
- name: Coverage summary
if: always()
shell: sh
env:
COVERAGE_THRESHOLD: "40"
run: |
set +e
echo "=== 覆盖率汇总 ==="
python3 scripts/ci_coverage_summary.py
- name: Notify CI failure
if: failure()
shell: sh
run: |
set -eu
echo "Build completed successfully!"
echo "Branch: ${GITHUB_REF_NAME}"
echo "Commit: ${GITHUB_SHA}"
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Validate Code Quality And Tests" python3 scripts/ci_notify_failure.py
- name: Notify CI failure - Integration Tests
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Integration Tests" python3 scripts/ci_notify_failure.py
frontend-lint:
name: Frontend Lint
runs-on: ubuntu-22.04
runs-on: host
timeout-minutes: 10
steps:
@@ -372,13 +578,22 @@ jobs:
-w /workspace/apps/web \
docker.m.daocloud.io/library/node:20 \
sh -lc 'npx vitest run src/test'
- name: Notify CI failure
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Frontend Lint" python3 scripts/ci_notify_failure.py
deploy-staging:
name: Build & Push Staging (Watchtower auto-deploy)
runs-on: saas
timeout-minutes: 30
needs: [validate, frontend-lint]
if: github.ref_name == 'main' || github.ref_name == 'develop' || startsWith(github.ref_name, 'feature/')
if: github.event_name == 'push' && (github.ref_name == 'main' || github.ref_name == 'develop')
steps:
- name: Checkout code
@@ -503,6 +718,23 @@ jobs:
echo "Branch: ${GITHUB_REF_NAME}"
echo "Commit: ${GITHUB_SHA}"
- name: Notify CI success
if: success()
shell: sh
run: |
set +e
echo "=== CI 成功通知 ==="
SUCCESS_JOB="Staging部署成功" python3 scripts/ci_notify_success.py
- name: Notify CI failure
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Build & Push Staging (Watchtower auto-deploy)" python3 scripts/ci_notify_failure.py
staging-e2e:
name: Staging E2E Tests
@@ -571,6 +803,15 @@ jobs:
git.xiaoxiajianji.com/xiaoxia/base/playwright:v1.45.0-jammy \
sh -lc "npm ci && npx playwright test --reporter=line --project=chromium e2e/auth.spec.ts e2e/auth-guard.spec.ts e2e/core-upload.spec.ts e2e/core-generation.spec.ts"
- name: Notify CI failure
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Staging E2E Tests" python3 scripts/ci_notify_failure.py
staging-api-tests:
name: Staging API Integration Tests
runs-on: saas
@@ -636,6 +877,15 @@ jobs:
git.xiaoxiajianji.com/xiaoxia/base/playwright:v1.45.0-jammy \
sh -lc 'npm ci && npx playwright test --reporter=line e2e/test_auth.spec.ts e2e/test_asset.spec.ts e2e/test_project.spec.ts'
- name: Notify CI failure
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Staging API Integration Tests" python3 scripts/ci_notify_failure.py
build-production-runtime-images:
name: Build Production Runtime Images
@@ -716,6 +966,15 @@ jobs:
echo "Disk usage after cleanup:"
df -h / | tail -1
- name: Notify CI failure
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Build Production Runtime Images" python3 scripts/ci_notify_failure.py
deploy-production:
name: Deploy Production
runs-on: saas
@@ -781,6 +1040,23 @@ jobs:
echo "$DEPLOY_B64" | base64 -d | ssh -p 22222 -i "$key_path" "$production_user@$production_host" "IMAGE_TAG='${GITHUB_REF_NAME}' REGISTRY_TOKEN='${REGISTRY_TOKEN}' sh"
- name: Notify CI success
if: success()
shell: sh
run: |
set +e
echo "=== CI 成功通知 ==="
SUCCESS_JOB="生产部署成功" python3 scripts/ci_notify_success.py
- name: Notify CI failure
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Deploy Production" python3 scripts/ci_notify_failure.py
production-e2e:
name: Production Browser E2E
runs-on: saas
@@ -849,3 +1125,12 @@ jobs:
-w /workspace/apps/web \
git.xiaoxiajianji.com/xiaoxia/base/playwright:v1.45.0-jammy \
sh -lc 'npm ci && npx playwright test --reporter=line --project=chromium e2e/auth.spec.ts e2e/auth-guard.spec.ts e2e/core-upload.spec.ts e2e/core-generation.spec.ts e2e/core-titles.spec.ts'
- name: Notify CI failure
if: failure()
shell: sh
run: |
set +e
echo "=== CI 失败通知 ==="
FAILED_JOB="Production Browser E2E" python3 scripts/ci_notify_failure.py
-69
View File
@@ -1,69 +0,0 @@
name: Test SSH Secret
on:
push:
branches: [develop]
paths:
- '.gitea/workflows/test-ssh-secret.yml'
jobs:
test-ssh:
runs-on: ubuntu-22.04
steps:
- name: Install SSH client
run: |
which ssh || (apt-get update && apt-get install -y openssh-client)
ssh -V
- name: Debug environment
run: |
echo "=== Environment ==="
echo "Runner hostname: $(hostname)"
echo "Runner IP: $(hostname -i || echo 'unknown')"
echo "Current user: $(whoami)"
echo "=== Secrets check ==="
if [ -n "$STAGING_SSH_HOST" ]; then
echo "STAGING_SSH_HOST: [SET] value_length=${#STAGING_SSH_HOST}"
else
echo "STAGING_SSH_HOST: [EMPTY]"
fi
if [ -n "$STAGING_SSH_USER" ]; then
echo "STAGING_SSH_USER: [SET] value_length=${#STAGING_SSH_USER}"
else
echo "STAGING_SSH_USER: [EMPTY]"
fi
if [ -n "$STAGING_SSH_KEY" ]; then
echo "STAGING_SSH_KEY: [SET] value_length=${#STAGING_SSH_KEY}"
else
echo "STAGING_SSH_KEY: [EMPTY]"
fi
env:
STAGING_SSH_HOST: ${{ secrets.STAGING_SSH_HOST }}
STAGING_SSH_USER: ${{ secrets.STAGING_SSH_USER }}
STAGING_SSH_KEY: ${{ secrets.STAGING_SSH_KEY }}
- name: Setup SSH key
run: |
mkdir -p ~/.ssh
chmod 700 ~/.ssh
echo "$STAGING_SSH_KEY" > ~/.ssh/id_ed25519
chmod 600 ~/.ssh/id_ed25519
ssh-keygen -y -f ~/.ssh/id_ed25519 > ~/.ssh/id_ed25519.pub 2>/dev/null || echo "No public key generated"
echo "=== SSH Key fingerprint ==="
ssh-keygen -lf ~/.ssh/id_ed25519 || echo "Key fingerprint failed"
env:
STAGING_SSH_KEY: ${{ secrets.STAGING_SSH_KEY }}
- name: Test SSH connection
run: |
echo "Attempting SSH connection to $STAGING_SSH_HOST..."
ssh -i ~/.ssh/id_ed25519 \
-o StrictHostKeyChecking=no \
-o UserKnownHostsFile=/dev/null \
-o ConnectTimeout=10 \
-o BatchMode=yes \
-v \
$STAGING_SSH_USER@$STAGING_SSH_HOST "echo 'SSH_CONNECTION_SUCCESS' && hostname && whoami"
echo "=== SSH Test Complete ==="
env:
STAGING_SSH_HOST: ${{ secrets.STAGING_SSH_HOST }}
STAGING_SSH_USER: ${{ secrets.STAGING_SSH_USER }}
-163
View File
@@ -1,163 +0,0 @@
name: Tests
on:
pull_request:
branches: [ main ]
jobs:
test:
runs-on: runtime-builder
steps:
- name: Checkout code
shell: sh
run: |
set -eu
python - <<'PY'
import io
import os
import tarfile
import time
import urllib.error
import urllib.request
url = f"{os.environ['GITHUB_API_URL']}/repos/{os.environ['GITHUB_REPOSITORY']}/archive/{os.environ['GITHUB_SHA']}.tar.gz"
request = urllib.request.Request(url, headers={"Authorization": f"token {os.environ['GITHUB_TOKEN']}"})
# Retry up to 5 times with backoff for transient 5xx errors
last_err = None
for attempt in range(5):
try:
with urllib.request.urlopen(request, timeout=120) as response:
archive = response.read()
break
except urllib.error.HTTPError as e:
last_err = e
if e.code >= 500 and attempt < 4:
wait = 2 ** attempt
print(f"Checkout HTTP {e.code}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
except Exception as e:
last_err = e
if attempt < 4:
wait = 2 ** attempt
print(f"Checkout error: {e}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
else:
raise last_err
with tarfile.open(fileobj=io.BytesIO(archive), mode='r:gz') as tar:
root_prefix = tar.getmembers()[0].name.split('/', 1)[0] + '/'
for member in tar.getmembers():
name = member.name
if name == root_prefix[:-1]:
continue
if name.startswith(root_prefix):
member.name = name[len(root_prefix):]
if member.name:
tar.extract(member, '.')
PY
- name: Show Python version
shell: sh
run: |
set -eu
python --version
python -m pip --version
- name: Install dependencies
shell: sh
run: |
set -eu
python -m pip install --upgrade pip -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com
python -m pip install -r requirements.txt -r requirements-dev.txt -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com
- name: Run unit tests
shell: sh
run: |
set -eu
PYTHONPATH="$PWD/apps/api:$PWD" python -m pytest tests/unit -q
- name: Run integration tests
shell: sh
run: |
set -eu
PYTHONPATH="$PWD/apps/api:$PWD" python -m pytest tests/integration -q --timeout=60 -x
lint:
runs-on: runtime-builder
steps:
- name: Checkout code
shell: sh
run: |
set -eu
python - <<'PY'
import io
import os
import tarfile
import time
import urllib.error
import urllib.request
url = f"{os.environ['GITHUB_API_URL']}/repos/{os.environ['GITHUB_REPOSITORY']}/archive/{os.environ['GITHUB_SHA']}.tar.gz"
request = urllib.request.Request(url, headers={"Authorization": f"token {os.environ['GITHUB_TOKEN']}"})
# Retry up to 5 times with backoff for transient 5xx errors
last_err = None
for attempt in range(5):
try:
with urllib.request.urlopen(request, timeout=120) as response:
archive = response.read()
break
except urllib.error.HTTPError as e:
last_err = e
if e.code >= 500 and attempt < 4:
wait = 2 ** attempt
print(f"Checkout HTTP {e.code}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
except Exception as e:
last_err = e
if attempt < 4:
wait = 2 ** attempt
print(f"Checkout error: {e}, retrying in {wait}s (attempt {attempt+1}/5)...")
time.sleep(wait)
continue
raise
else:
raise last_err
with tarfile.open(fileobj=io.BytesIO(archive), mode='r:gz') as tar:
root_prefix = tar.getmembers()[0].name.split('/', 1)[0] + '/'
for member in tar.getmembers():
name = member.name
if name == root_prefix[:-1]:
continue
if name.startswith(root_prefix):
member.name = name[len(root_prefix):]
if member.name:
tar.extract(member, '.')
PY
- name: Install dependencies
shell: sh
run: |
set -eu
python -m pip install --upgrade pip -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com
python -m pip install -r requirements.txt -r requirements-dev.txt -i https://mirrors.aliyun.com/pypi/simple/ --trusted-host mirrors.aliyun.com
- name: Run Black (check only)
shell: sh
run: |
set -eu
python -m black --check alembic apps packages tests scripts
- name: Run Flake8
shell: sh
run: |
set -eu
python -m flake8 apps packages tests --count --statistics
+1
View File
@@ -6,6 +6,7 @@ dist/
coverage/
# Python / backend
.cache/
.venv/
venv/
.venv-ci-root/
+10
View File
@@ -8,10 +8,12 @@ from app.api.routes.dashboard import router as dashboard_router
from app.api.routes.duplication import router as duplication_router
from app.api.routes.edit_plans import router as edit_plans_router
from app.api.routes.edit_templates import router as edit_templates_router
from app.api.routes.feature_flags import router as feature_flags_router
from app.api.routes.generated_videos import router as generated_videos_router
from app.api.routes.generation_tasks import router as generation_tasks_router
from app.api.routes.health import router as health_check_router
from app.api.routes.ingest_jobs import router as ingest_jobs_router
from app.api.routes.internal_render import router as internal_render_router
from app.api.routes.jobs import router as jobs_router
from app.api.routes.projects import router as projects_router
from app.api.routes.recipes import router as recipes_router
@@ -151,3 +153,11 @@ api_router.include_router(
prefix="/tts",
tags=["TTS"],
)
api_router.include_router(
feature_flags_router,
tags=["Internal"],
)
api_router.include_router(
internal_render_router,
tags=["Internal"],
)
+26
View File
@@ -24,6 +24,7 @@ from typing import Any, List, Optional
from app.auth import AuthenticatedUser, get_current_user
from app.core.celery_app import celery_app
from app.core.task_enqueue import GLOBAL_PENDING_LIMIT, USER_PENDING_LIMIT
from app.dependencies import get_asset_library_repository, get_asset_repository, get_db_session, get_project_repository
from app.schemas.generation_task import GenerationTaskResponse
from app.services import EditPlanService, PlanGeneratorService
@@ -644,6 +645,31 @@ def generate_plan(
# 创建 GenerationTask
gen_task_repo = SQLAlchemyGenerationTaskRepository(db)
# 队列限流预检查(repository 不支持计数时跳过)
user_id = current_user.user.id
try:
has_count = hasattr(gen_task_repo, "count_pending_by_user") and hasattr(
gen_task_repo, "count_pending_total"
)
if has_count:
user_pending = gen_task_repo.count_pending_by_user(user_id)
global_pending = gen_task_repo.count_pending_total()
if user_pending >= USER_PENDING_LIMIT:
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
)
if global_pending >= GLOBAL_PENDING_LIMIT:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
)
except HTTPException:
raise
except Exception as e:
logger.warning("[队列限流] 剪辑计划限流检查失败,跳过: %s", e)
gen_task_use_case = CreateGenerationTaskUseCase(gen_task_repo)
plan = svc.get_plan_or_raise(plan_id)
gen_task = gen_task_use_case.execute(
+195
View File
@@ -0,0 +1,195 @@
"""Feature Flag 内部管理接口。
通过内部 API Key 鉴权,支持查看和修改 Feature Flag 配置。
主要用于灰度发布期间的动态开关控制。
API:
GET /api/v1/internal/feature-flags - 列出所有 flag
GET /api/v1/internal/feature-flags/{name} - 查看单个 flag
PUT /api/v1/internal/feature-flags/{name} - 设置 flag 配置
DELETE /api/v1/internal/feature-flags/{name} - 删除 flag
鉴权:X-API-Key header,走内部 API Key 验证
"""
from __future__ import annotations
import logging
from typing import Optional
from app.api.routes.auth import _verify_internal_api_key
from app.config import settings
from fastapi import APIRouter, Depends, HTTPException, Query, status
from pydantic import BaseModel, Field
from packages.adapters.redis.feature_flag_store import (
FEATURE_FLAG_REDIS_PREFIX,
FeatureFlagConfig,
RedisFeatureFlagStore,
)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/internal/feature-flags", tags=["Internal"])
# 允许管理的 flag 白名单(防止误操作其他系统 flag)
ALLOWED_FLAGS = {
"render_engine",
}
def _get_feature_flag_store() -> RedisFeatureFlagStore:
"""获取 Feature Flag 存储实例。"""
return RedisFeatureFlagStore(redis_url=settings.REDIS_URL)
class FeatureFlagUpdateRequest(BaseModel):
"""Feature Flag 更新请求体。"""
enabled: bool = Field(..., description="是否启用")
percentage: int = Field(0, ge=0, le=100, description="灰度百分比 (0-100)")
whitelist: list[str] = Field(default_factory=list, description="白名单列表(如 user_id")
class FeatureFlagResponse(BaseModel):
"""Feature Flag 响应。"""
name: str
enabled: bool
percentage: int
whitelist: list[str]
@classmethod
def from_config(cls, config: FeatureFlagConfig) -> "FeatureFlagResponse":
return cls(
name=config.name,
enabled=config.enabled,
percentage=config.percentage,
whitelist=sorted(config.whitelist),
)
class FeatureFlagCheckResponse(BaseModel):
"""Flag 激活检查响应。"""
name: str
active: bool
identifier: Optional[str] = None
def _validate_flag_name(name: str) -> None:
"""校验 flag 名称是否在允许列表中。"""
if name not in ALLOWED_FLAGS:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Unsupported flag: {name}. Allowed: {sorted(ALLOWED_FLAGS)}",
)
@router.get("", response_model=list[FeatureFlagResponse])
async def list_feature_flags(
_: bool = Depends(_verify_internal_api_key),
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
):
"""列出所有 Feature Flag。"""
try:
flags = store.list_all()
# 同时返回预定义的 flag(即使未设置也显示默认值)
result = []
for name in sorted(ALLOWED_FLAGS):
config = flags.get(name) or FeatureFlagConfig(name=name, enabled=False)
result.append(FeatureFlagResponse.from_config(config))
# 加上已存在但不在白名单中的 flag(只读展示)
for name, config in flags.items():
if name not in ALLOWED_FLAGS:
result.append(FeatureFlagResponse.from_config(config))
return sorted(result, key=lambda x: x.name)
except Exception as exc:
logger.error("Failed to list feature flags: %s", exc)
raise HTTPException(status_code=500, detail=f"Failed to list flags: {exc}")
@router.get("/{name}", response_model=FeatureFlagResponse)
async def get_feature_flag(
name: str,
_: bool = Depends(_verify_internal_api_key),
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
):
"""获取单个 Feature Flag 配置。"""
try:
config = store.get(name)
return FeatureFlagResponse.from_config(config)
except Exception as exc:
logger.error("Failed to get feature flag %s: %s", name, exc)
raise HTTPException(status_code=500, detail=f"Failed to get flag: {exc}")
@router.get("/{name}/check", response_model=FeatureFlagCheckResponse)
async def check_feature_flag(
name: str,
identifier: Optional[str] = Query(None, description="标识符,如 user_id"),
_: bool = Depends(_verify_internal_api_key),
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
):
"""检查某个标识符是否命中 Feature Flag。"""
try:
active = store.is_active(name, identifier=identifier)
return FeatureFlagCheckResponse(name=name, active=active, identifier=identifier)
except Exception as exc:
logger.error("Failed to check feature flag %s: %s", name, exc)
raise HTTPException(status_code=500, detail=f"Failed to check flag: {exc}")
@router.put("/{name}", response_model=FeatureFlagResponse)
async def update_feature_flag(
name: str,
request: FeatureFlagUpdateRequest,
_: bool = Depends(_verify_internal_api_key),
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
):
"""更新 Feature Flag 配置。
只允许修改 ALLOWED_FLAGS 列表中的 flag。
"""
_validate_flag_name(name)
try:
config = FeatureFlagConfig(
name=name,
enabled=request.enabled,
percentage=request.percentage,
whitelist=set(request.whitelist),
)
store.set(config)
logger.info(
"Feature flag updated: name=%s enabled=%s percentage=%d whitelist=%d",
name,
config.enabled,
config.percentage,
len(config.whitelist),
)
return FeatureFlagResponse.from_config(config)
except Exception as exc:
logger.error("Failed to update feature flag %s: %s", name, exc)
raise HTTPException(status_code=500, detail=f"Failed to update flag: {exc}")
@router.delete("/{name}", status_code=status.HTTP_204_NO_CONTENT)
async def delete_feature_flag(
name: str,
_: bool = Depends(_verify_internal_api_key),
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
):
"""删除 Feature Flag。
只允许删除 ALLOWED_FLAGS 列表中的 flag。
"""
_validate_flag_name(name)
try:
deleted = store.delete(name)
logger.info("Feature flag deleted: name=%s deleted=%s", name, deleted)
return None
except Exception as exc:
logger.error("Failed to delete feature flag %s: %s", name, exc)
raise HTTPException(status_code=500, detail=f"Failed to delete flag: {exc}")
+95 -8
View File
@@ -5,7 +5,14 @@ from typing import Any
from app.auth import AuthenticatedUser, get_current_user
from app.core.storage import OSSStorageService, get_storage_service
from app.core.task_enqueue import safe_enqueue_generation_task
from app.core.task_enqueue import (
GLOBAL_PENDING_LIMIT,
USER_PENDING_LIMIT,
GlobalQueueFull,
UserPendingLimitExceeded,
check_queue_limits,
safe_enqueue_generation_task,
)
from app.dependencies import (
get_asset_library_repository,
get_asset_repository,
@@ -228,9 +235,31 @@ def create_generation_task(
count = request.count
created_tasks = []
failed_tasks = []
user_id = authenticated_user.user.id
# 同批次任务共享 batch_id,用于视频查重时批次内比对
batch_id = uuid.uuid4().hex if count > 1 else ""
# 预检查:批量提交前先看会不会超限,避免建一半才拒
try:
user_pending = generation_task_repository.count_pending_by_user(user_id)
global_pending = generation_task_repository.count_pending_total()
if user_pending + count > USER_PENDING_LIMIT:
raise UserPendingLimitExceeded(
user_id=user_id, pending_count=user_pending + count, limit=USER_PENDING_LIMIT
)
if global_pending + count > GLOBAL_PENDING_LIMIT:
raise GlobalQueueFull(pending_count=global_pending + count, limit=GLOBAL_PENDING_LIMIT)
except UserPendingLimitExceeded as e:
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {e.pending_count - count}/{e.limit},本次提交 {count} 个),请等待完成后再提交",
) from e
except GlobalQueueFull as e:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
) from e
try:
for _ in range(count):
task = use_case.execute(
@@ -243,16 +272,42 @@ def create_generation_task(
asset_ids=resolved_asset_ids,
title_ids=request.title_ids,
voice_ids=request.voice_ids,
created_by_user_id=authenticated_user.user.id,
created_by_user_id=user_id,
source_edit_plan_id=request.source_edit_plan_id,
asset_select_mode=request.asset_select_mode,
batch_id=batch_id,
)
)
if safe_enqueue_generation_task(task, generation_task_repository, log_prefix="[生成任务]", log_task_status=True):
created_tasks.append(task)
else:
try:
if safe_enqueue_generation_task(
task,
generation_task_repository,
user_id=user_id,
log_prefix="[生成任务]",
log_task_status=True,
):
created_tasks.append(task)
else:
failed_tasks.append(task)
except UserPendingLimitExceeded:
# 兜底:如果预检查后又并发提交了,在这里也拦住
failed_tasks.append(task)
if not created_tasks:
raise HTTPException(
status_code=429,
detail="您的待处理任务过多,请等待完成后再提交",
)
break
except GlobalQueueFull:
failed_tasks.append(task)
if not created_tasks:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
)
break
except HTTPException:
raise
except Exception as e:
logger.error("[生成任务] 创建失败: %s", e, exc_info=True)
raise HTTPException(status_code=500, detail="创建生成任务失败,请稍后重试或查看任务日志")
@@ -327,6 +382,21 @@ def retry_generation_task(
if status_val != "failed":
raise HTTPException(status_code=409, detail="Only failed tasks can be retried")
user_id = authenticated_user.user.id
# 预检查:创建前判断,>= 上限就拒绝
user_pending = generation_task_repository.count_pending_by_user(user_id)
global_pending = generation_task_repository.count_pending_total()
if user_pending >= USER_PENDING_LIMIT:
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
)
if global_pending >= GLOBAL_PENDING_LIMIT:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
)
use_case = CreateGenerationTaskUseCase(generation_task_repository)
retried = use_case.execute(
CreateGenerationTaskCommand(
@@ -338,11 +408,28 @@ def retry_generation_task(
asset_ids=task.asset_ids,
title_ids=task.title_ids,
voice_ids=task.voice_ids,
created_by_user_id=authenticated_user.user.id,
created_by_user_id=user_id,
source_edit_plan_id=task.source_edit_plan_id or "",
asset_select_mode=getattr(task, "asset_select_mode", ""),
)
)
if not safe_enqueue_generation_task(retried, generation_task_repository, log_prefix="[生成任务]", log_task_status=True):
logger.warning("[生成任务] 重试入队失败: task_id=%s", retried.id)
try:
if not safe_enqueue_generation_task(
retried,
generation_task_repository,
user_id=user_id,
log_prefix="[生成任务]",
log_task_status=True,
):
logger.warning("[生成任务] 重试入队失败: task_id=%s", retried.id)
except UserPendingLimitExceeded:
raise HTTPException(
status_code=429,
detail="您的待处理任务过多,请等待完成后再提交",
) from None
except GlobalQueueFull:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
) from None
return _to_generation_task_response(retried)
+120
View File
@@ -0,0 +1,120 @@
"""渲染结果内部下载接口。
通过内部 API Key 鉴权,为灰度对比工具等内部系统提供渲染结果下载能力。
API:
GET /api/v1/internal/render/videos/{video_id}/download-url - 获取单个视频下载URL
GET /api/v1/internal/render/tasks/{task_id}/videos - 获取任务下所有视频及下载URL
鉴权:X-API-Key header,走内部 API Key 验证
"""
from __future__ import annotations
import logging
from typing import Any
from app.api.routes.auth import _verify_internal_api_key
from app.core.storage import OSSStorageService, get_storage_service
from app.dependencies import get_generated_video_repository
from fastapi import APIRouter, Depends, HTTPException, Query
from pydantic import BaseModel
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/internal/render", tags=["Internal"])
class InternalRenderVideoItem(BaseModel):
"""内部渲染视频项。"""
video_id: str
generation_task_id: str
project_id: str
name: str
file_url: str
file_size: int | None = None
duration: float | None = None
width: int | None = None
height: int | None = None
fps: float | None = None
status: str
download_url: str
class InternalRenderTaskVideosResponse(BaseModel):
"""任务下所有渲染视频响应。"""
task_id: str
count: int
videos: list[InternalRenderVideoItem]
class InternalRenderDownloadUrlResponse(BaseModel):
"""单个视频下载URL响应。"""
video_id: str
download_url: str
def _video_to_item(video: Any, download_url: str) -> InternalRenderVideoItem:
"""将 GeneratedVideo 领域对象转为响应项。"""
return InternalRenderVideoItem(
video_id=video.id,
generation_task_id=video.generation_task_id,
project_id=video.project_id,
name=video.name,
file_url=video.file_url,
file_size=getattr(video, "file_size", None),
duration=getattr(video, "duration", None),
width=getattr(video, "width", None),
height=getattr(video, "height", None),
fps=getattr(video, "fps", None),
status=video.status,
download_url=download_url,
)
@router.get("/videos/{video_id}/download-url", response_model=InternalRenderDownloadUrlResponse)
def get_render_video_download_url(
video_id: str,
_: bool = Depends(_verify_internal_api_key),
generated_video_repository: Any = Depends(get_generated_video_repository),
storage_service: OSSStorageService = Depends(get_storage_service),
) -> InternalRenderDownloadUrlResponse:
"""获取单个渲染视频的下载URL(预签名)。"""
video = generated_video_repository.get(video_id)
if video is None:
raise HTTPException(status_code=404, detail=f"GeneratedVideo {video_id} not found")
download_url = storage_service.get_download_url(video.file_url, expires_seconds=86400)
logger.info("内部渲染下载URL生成: video_id=%s", video_id)
return InternalRenderDownloadUrlResponse(video_id=video_id, download_url=download_url)
@router.get("/tasks/{task_id}/videos", response_model=InternalRenderTaskVideosResponse)
def get_render_task_videos(
task_id: str,
status: str | None = Query(None, description="按状态筛选,如 completed/failed"),
_: bool = Depends(_verify_internal_api_key),
generated_video_repository: Any = Depends(get_generated_video_repository),
storage_service: OSSStorageService = Depends(get_storage_service),
) -> InternalRenderTaskVideosResponse:
"""获取生成任务下所有渲染视频及下载URL。"""
videos = generated_video_repository.list_by_generation_task(task_id)
# 状态筛选
if status:
videos = [v for v in videos if v.status == status]
items = []
for video in videos:
download_url = storage_service.get_download_url(video.file_url, expires_seconds=86400)
items.append(_video_to_item(video, download_url))
logger.info("内部渲染任务视频查询: task_id=%s count=%d", task_id, len(items))
return InternalRenderTaskVideosResponse(
task_id=task_id,
count=len(items),
videos=items,
)
+70 -7
View File
@@ -3,7 +3,13 @@ from typing import Any
from app.auth import AuthenticatedUser, get_current_user
from app.core.celery_app import celery_app
from app.core.task_enqueue import safe_enqueue_generation_task
from app.core.task_enqueue import (
GLOBAL_PENDING_LIMIT,
USER_PENDING_LIMIT,
GlobalQueueFull,
UserPendingLimitExceeded,
safe_enqueue_generation_task,
)
from app.dependencies import (
get_generation_task_repository,
get_ingest_job_repository,
@@ -142,6 +148,21 @@ def retry_task_by_id(
if _status_value(task.status) != "failed":
raise HTTPException(status_code=409, detail="Only failed tasks can be retried")
user_id = authenticated_user.user.id
# 预检查
user_pending = generation_task_repository.count_pending_by_user(user_id)
global_pending = generation_task_repository.count_pending_total()
if user_pending >= USER_PENDING_LIMIT:
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
)
if global_pending >= GLOBAL_PENDING_LIMIT:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
)
use_case = CreateGenerationTaskUseCase(generation_task_repository)
retried = use_case.execute(
CreateGenerationTaskCommand(
@@ -153,11 +174,24 @@ def retry_task_by_id(
asset_ids=task.asset_ids,
title_ids=task.title_ids,
voice_ids=task.voice_ids,
created_by_user_id=authenticated_user.user.id,
created_by_user_id=user_id,
)
)
if not safe_enqueue_generation_task(retried, generation_task_repository, log_prefix="[任务中心]"):
logger.warning("[任务中心] 用户级重试入队失败: task_id=%s", retried.id)
try:
if not safe_enqueue_generation_task(
retried, generation_task_repository, user_id=user_id, log_prefix="[任务中心]"
):
logger.warning("[任务中心] 用户级重试入队失败: task_id=%s", retried.id)
except UserPendingLimitExceeded:
raise HTTPException(
status_code=429,
detail="您的待处理任务过多,请等待完成后再提交",
) from None
except GlobalQueueFull:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
) from None
return UserTaskResponse(
id=f"generation:{retried.id}",
task_type="generation",
@@ -225,6 +259,22 @@ def retry_project_task(
raise HTTPException(status_code=404, detail="Generation task not found")
if _status_value(task.status) != "failed":
raise HTTPException(status_code=409, detail="Only failed tasks can be retried")
user_id = authenticated_user.user.id
# 预检查
user_pending = generation_task_repository.count_pending_by_user(user_id)
global_pending = generation_task_repository.count_pending_total()
if user_pending >= USER_PENDING_LIMIT:
raise HTTPException(
status_code=429,
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
)
if global_pending >= GLOBAL_PENDING_LIMIT:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
)
use_case = CreateGenerationTaskUseCase(generation_task_repository)
retried = use_case.execute(
CreateGenerationTaskCommand(
@@ -236,11 +286,24 @@ def retry_project_task(
asset_ids=task.asset_ids,
title_ids=task.title_ids,
voice_ids=task.voice_ids,
created_by_user_id=authenticated_user.user.id,
created_by_user_id=user_id,
)
)
if not safe_enqueue_generation_task(retried, generation_task_repository, log_prefix="[任务中心]"):
logger.warning("[任务中心] 项目级重试用队失败: task_id=%s", retried.id)
try:
if not safe_enqueue_generation_task(
retried, generation_task_repository, user_id=user_id, log_prefix="[任务中心]"
):
logger.warning("[任务中心] 项目级重试入队失败: task_id=%s", retried.id)
except UserPendingLimitExceeded:
raise HTTPException(
status_code=429,
detail="您的待处理任务过多,请等待完成后再提交",
) from None
except GlobalQueueFull:
raise HTTPException(
status_code=503,
detail="系统繁忙,请稍后再试",
) from None
return _generation_task_to_project_response(retried)
if task_type == "ingest":
job = ingest_job_repository.get(source_id)
+8 -14
View File
@@ -6,7 +6,7 @@ import logging
from typing import Optional
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_audio_url_signer, get_cosyvoice_service, get_voice_clone_profile_repository
from app.dependencies import get_cosyvoice_service, get_voice_clone_profile_repository
from app.schemas.voice_clone import (
CreateVoiceCloneRequest,
ListVoiceCloneResponse,
@@ -37,16 +37,14 @@ logger = logging.getLogger(__name__)
router = APIRouter()
def _to_response(profile, sign_url=None) -> VoiceCloneProfileResponse:
source_url = profile.source_audio_url
if sign_url and source_url:
source_url = sign_url(source_url)
def _to_response(profile) -> VoiceCloneProfileResponse:
# source_audio_url 是用户传入的原始 URL(可能是外部地址),不做预签名转换
return VoiceCloneProfileResponse(
id=profile.id,
user_id=profile.user_id,
name=profile.name,
description=profile.description,
source_audio_url=source_url,
source_audio_url=profile.source_audio_url,
voice_id=profile.voice_id,
voice_model=profile.voice_model,
language=profile.language,
@@ -77,7 +75,6 @@ def create_voice_clone(
request: CreateVoiceCloneRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
workflow: VoiceCloneWorkflowService = Depends(_get_workflow_service),
sign_url=Depends(get_audio_url_signer),
) -> VoiceCloneProfileResponse:
"""创建音色克隆任务。
@@ -113,7 +110,7 @@ def create_voice_clone(
except Exception as inner_e:
logger.error(f"Failed to mark profile as failed after dispatch error: {inner_e}")
return _to_response(profile, sign_url)
return _to_response(profile)
@router.get("", response_model=ListVoiceCloneResponse)
@@ -123,14 +120,13 @@ def list_voice_clones(
limit: int = Query(50, ge=1, le=200),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
repository: SQLAlchemyVoiceCloneProfileRepository = Depends(get_voice_clone_profile_repository),
sign_url=Depends(get_audio_url_signer),
) -> ListVoiceCloneResponse:
"""获取用户的音色克隆列表。"""
user_id = authenticated_user.user.id
use_case = ListVoiceClonesUseCase(repository)
items, total = use_case.execute(user_id, status=status_filter, skip=skip, limit=limit)
return ListVoiceCloneResponse(
items=[_to_response(p, sign_url) for p in items],
items=[_to_response(p) for p in items],
total=total,
)
@@ -140,7 +136,6 @@ def get_voice_clone(
clone_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
repository: SQLAlchemyVoiceCloneProfileRepository = Depends(get_voice_clone_profile_repository),
sign_url=Depends(get_audio_url_signer),
) -> VoiceCloneProfileResponse:
"""获取音色克隆详情。"""
user_id = authenticated_user.user.id
@@ -149,7 +144,7 @@ def get_voice_clone(
profile = use_case.execute(clone_id, user_id)
except VoiceCloneNotFoundError:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Voice clone not found")
return _to_response(profile, sign_url)
return _to_response(profile)
@router.get("/{clone_id}/status", response_model=VoiceCloneStatusResponse)
@@ -198,7 +193,6 @@ def retry_voice_clone(
clone_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
workflow: VoiceCloneWorkflowService = Depends(_get_workflow_service),
sign_url=Depends(get_audio_url_signer),
) -> VoiceCloneProfileResponse:
"""重试失败的音色克隆。
@@ -231,4 +225,4 @@ def retry_voice_clone(
except Exception as inner_e:
logger.error(f"Failed to mark profile as failed after dispatch error: {inner_e}")
return _to_response(profile, sign_url)
return _to_response(profile)
+3
View File
@@ -109,6 +109,9 @@ class Settings(BaseSettings):
LOG_LEVEL: str = "INFO"
CORS_ORIGINS_RAW: str = "http://localhost:3000,http://localhost:5173,http://localhost:8000"
# 渲染引擎选择:legacy=旧VideoComposeServiceunified=新UnifiedRenderService
RENDER_ENGINE: str = "legacy"
model_config = SettingsConfigDict(
env_file=".env",
env_file_encoding="utf-8",
+175 -12
View File
@@ -5,37 +5,163 @@ from app.core.celery_app import celery_app
logger = logging.getLogger(__name__)
# ── 限流阈值常量(全系统统一管理,不要在业务代码里硬编码) ──
USER_PENDING_LIMIT = 3 # 单用户 pending 上限
GLOBAL_PENDING_LIMIT = 20 # 全局 pending 上限
class UserPendingLimitExceeded(Exception):
"""用户 pending 任务数超限,返回 429。"""
def __init__(self, user_id: str, pending_count: int, limit: int):
self.user_id = user_id
self.pending_count = pending_count
self.limit = limit
super().__init__(f"用户 {user_id} pending 任务数 {pending_count} 超过上限 {limit}")
class GlobalQueueFull(Exception):
"""全局限流,返回 503。"""
def __init__(self, pending_count: int, limit: int):
self.pending_count = pending_count
self.limit = limit
super().__init__(f"系统 pending 任务数 {pending_count} 超过上限 {limit}")
def check_queue_limits(
user_id: str,
generation_task_repository: Any,
*,
user_pending_limit: int = USER_PENDING_LIMIT,
global_pending_limit: int = GLOBAL_PENDING_LIMIT,
) -> None:
"""检查队列限流(预检查用,任务创建前调用),超限抛对应异常。
边界语义:>= 上限即拒绝(达到上限就不能再加新任务)。
Args:
user_id: 用户 ID
generation_task_repository: 任务仓储
user_pending_limit: 单用户 pending 上限,默认 USER_PENDING_LIMIT
global_pending_limit: 全局 pending 上限,默认 GLOBAL_PENDING_LIMIT
Raises:
GlobalQueueFull: 全局超限时抛出(优先级更高,先查全局)
UserPendingLimitExceeded: 用户超限时抛出
"""
# 先查全局(系统级保护优先级更高)
global_pending = generation_task_repository.count_pending_total()
if global_pending >= global_pending_limit:
logger.warning(
"[队列限流] 全局 pending 任务数超限: %d/%d, user_id=%s",
global_pending,
global_pending_limit,
user_id,
)
raise GlobalQueueFull(pending_count=global_pending, limit=global_pending_limit)
# 再查用户级
if user_id:
user_pending = generation_task_repository.count_pending_by_user(user_id)
if user_pending >= user_pending_limit:
logger.warning(
"[队列限流] 用户 pending 任务数超限: user_id=%s, count=%d/%d",
user_id,
user_pending,
user_pending_limit,
)
raise UserPendingLimitExceeded(user_id=user_id, pending_count=user_pending, limit=user_pending_limit)
def _mark_task_failed_safely(
task: Any,
generation_task_repository: Any,
log_prefix: str,
reason: str,
) -> None:
"""安全地把任务标记为 failed,更新失败只打日志不崩溃。"""
try:
task.mark_failed(f"任务被限流拒绝: {reason}")
generation_task_repository.update(task)
except Exception as update_err:
logger.error(
"%s 限流后更新状态也失败: task_id=%s error=%s",
log_prefix,
task.id,
update_err,
exc_info=True,
)
def safe_enqueue_generation_task(
task: Any,
generation_task_repository: Any,
*,
user_id: str = "",
log_prefix: str = "[任务队列]",
log_task_status: bool = False,
user_pending_limit: int = USER_PENDING_LIMIT,
global_pending_limit: int = GLOBAL_PENDING_LIMIT,
) -> bool:
"""安全入队:send_task 失败时自动把任务标记为 failed,避免留下 pending 僵尸任务
"""安全入队:入队前限流检查 → 发送 Celery 任务 → 入队后最终校验兜底
边界说明:
入队前检查用 > 而非 >=。因为调用此函数时 task 已经是 pending 状态并计入 DB,
pending 总数包含了当前任务本身。pending > limit 等价于"其他任务数 >= limit"
与预检查的 >= 语义一致(都是达到上限就拒绝新任务)。
入队后最终校验:发送 Celery 成功后再查一次 DB 计数,处理并发竞态场景
(两个请求同时通过入队前检查,后到的那个在这里被兜住)。
Args:
task: 生成任务对象,需有 id 属性和 mark_failed 方法
task: 生成任务对象,需有 id 属性和 mark_failed 方法(状态已为 pending
generation_task_repository: 任务仓储,用于更新状态
user_id: 用户 ID,传了才做用户级限流检查
log_prefix: 日志前缀,便于区分调用来源
log_task_status: 成功日志中是否额外打印任务状态
user_pending_limit: 单用户 pending 上限,默认 USER_PENDING_LIMIT
global_pending_limit: 全局 pending 上限,默认 GLOBAL_PENDING_LIMIT
Returns:
True 表示入队成功,False 表示入队失败(已标记为 failed)
Raises:
GlobalQueueFull: 全局 pending 超限时抛出,任务会被标记为 failed
UserPendingLimitExceeded: 用户 pending 超限时抛出,任务会被标记为 failed
"""
# ── 入队前检查:任务已是 pending,用 > 判断(包含当前任务) ──
# 全局限流检查(始终生效)
global_pending = generation_task_repository.count_pending_total()
if global_pending > global_pending_limit:
logger.warning(
"[队列限流] 全局 pending 任务数超限(入队前): %d/%d, user_id=%s",
global_pending,
global_pending_limit,
user_id or "unknown",
)
exc = GlobalQueueFull(pending_count=global_pending, limit=global_pending_limit)
_mark_task_failed_safely(task, generation_task_repository, log_prefix, str(exc))
raise exc
# 用户级限流检查(传了 user_id 才做)
if user_id:
user_pending = generation_task_repository.count_pending_by_user(user_id)
if user_pending > user_pending_limit:
logger.warning(
"[队列限流] 用户 pending 任务数超限(入队前): user_id=%s, count=%d/%d",
user_id,
user_pending,
user_pending_limit,
)
exc = UserPendingLimitExceeded(user_id=user_id, pending_count=user_pending, limit=user_pending_limit)
_mark_task_failed_safely(task, generation_task_repository, log_prefix, str(exc))
raise exc
# ── 发送 Celery 任务 ──
try:
celery_app.send_task("worker.generate_video", args=[task.id])
if log_task_status:
logger.info(
"%s 入队成功: task_id=%s, status=%s",
log_prefix,
task.id,
task.status,
)
else:
logger.info("%s 入队成功: task_id=%s", log_prefix, task.id)
return True
except Exception as e:
logger.error(
"%s 入队失败,标记为失败: task_id=%s error=%s",
@@ -56,3 +182,40 @@ def safe_enqueue_generation_task(
exc_info=True,
)
return False
# ── 入队后最终校验:并发竞态兜底 ──
# 发送成功后再查一次,防止两个请求同时通过入队前检查导致超限
global_after = generation_task_repository.count_pending_total()
user_after = generation_task_repository.count_pending_by_user(user_id) if user_id else 0
global_over = global_after > global_pending_limit
user_over = bool(user_id and user_after > user_pending_limit)
if global_over or user_over:
if global_over:
reason = f"全局 pending 超限(入队后): {global_after}/{global_pending_limit}"
exc: Exception = GlobalQueueFull(pending_count=global_after, limit=global_pending_limit)
else:
reason = f"用户 pending 超限(入队后): {user_after}/{user_pending_limit}"
exc = UserPendingLimitExceeded(user_id=user_id, pending_count=user_after, limit=user_pending_limit)
logger.warning(
"[队列限流] %s, task_id=%s, user_id=%s — 回滚状态为 failed",
reason,
task.id,
user_id or "unknown",
)
_mark_task_failed_safely(task, generation_task_repository, log_prefix, reason)
raise exc
# 入队成功日志
if log_task_status:
logger.info(
"%s 入队成功: task_id=%s, status=%s",
log_prefix,
task.id,
task.status,
)
else:
logger.info("%s 入队成功: task_id=%s", log_prefix, task.id)
return True
View File
+6 -7
View File
@@ -1,18 +1,17 @@
"""
视频处理模块
轻量工具(ffmpeg_utils / oss_helpers / dedup_helpers)顶层直接导出,
无额外依赖。渲染相关组件(UnifiedRenderService / RenderAdapter /
VideoProcessor 等)按需从子模块导入,避免 __init__ 阶段引入
packages / DB 等重依赖。
"""
# 共享工具模块(供 editing_modes / generation / edit_plan_generation 等复用)
# 共享工具模块(零外部依赖,供 editing_modes / generation / edit_plan_generation 等复用)
from . import dedup_helpers, ffmpeg_utils, oss_helpers
from .processor import VideoProcessor, VideoResult
from .unified_render_service import RenderResult, UnifiedRenderService
__all__ = [
"VideoProcessor",
"VideoResult",
"ffmpeg_utils",
"oss_helpers",
"dedup_helpers",
"UnifiedRenderService",
"RenderResult",
]
+7 -3
View File
@@ -93,12 +93,16 @@ class VideoFingerprint:
resolution: tuple[int, int]
def to_dict(self) -> dict:
# 注意:color_histograms 里的值可能是 np.float32(来自 cv2.normalize),
# 直接存进 dict 后 SQLAlchemy JSON 序列化会报 "float32 is not JSON serializable"。
# 这里统一转成 Python 原生 float。
native_histograms = [[float(v) for v in hist] for hist in self.color_histograms]
return {
"md5": self.md5,
"keyframe_phashes": self.keyframe_phashes,
"color_histograms": self.color_histograms,
"duration": self.duration,
"resolution": list(self.resolution),
"color_histograms": native_histograms,
"duration": float(self.duration),
"resolution": [int(self.resolution[0]), int(self.resolution[1])],
}
@@ -1,657 +0,0 @@
"""
视频剪辑模式处理器
支持四种剪辑模式:一镜到底、画中画、口播、口播+画中画
"""
import logging
import os
import sys
import tempfile
from dataclasses import dataclass
if sys.version_info >= (3, 11):
from enum import StrEnum
else:
from enum import Enum
class StrEnum(str, Enum):
pass
from pathlib import Path
from typing import Optional
from video_processing.ffmpeg_utils import FFMPEG_BIN, probe_video_info, run_ffmpeg
logger = logging.getLogger(__name__)
# 从 domain 层导入 EditingMode,避免重复定义
from packages.domain.editing_mode import EditingMode
class PIPPosition(StrEnum):
"""画中画位置枚举"""
TOP_LEFT = "top_left"
TOP_RIGHT = "top_right"
BOTTOM_LEFT = "bottom_left"
BOTTOM_RIGHT = "bottom_right"
@dataclass
class EditingModeConfig:
"""剪辑模式配置"""
mode: EditingMode
output_width: int = 1280
output_height: int = 720
output_fps: int = 25
pip_position: PIPPosition = PIPPosition.TOP_RIGHT
pip_scale: float = 0.25 # 画中画占主画面的比例
transition_duration: float = 0.5 # 转场时长(秒)
output_codec: str = "libx264"
output_preset: str = "medium"
output_crf: int = 23
class EditingModeProcessor:
"""剪辑模式处理器"""
def __init__(self, config: EditingModeConfig, work_dir: Optional[str] = None):
"""
初始化剪辑模式处理器
Args:
config: 剪辑模式配置
work_dir: 工作目录,默认使用系统临时目录
"""
self.config = config
self.work_dir = work_dir or tempfile.gettempdir()
def process(
self,
video_paths: list[str],
audio_path: Optional[str] = None,
output_path: Optional[str] = None,
) -> str:
"""
根据模式处理视频,返回输出文件路径
Args:
video_paths: 视频素材路径列表
audio_path: 音频路径(用于口播模式)
output_path: 输出文件路径,默认自动生成
Returns:
输出文件路径
"""
if not video_paths:
raise ValueError("video_paths cannot be empty")
self._validate_inputs(video_paths, audio_path)
if output_path is None:
output_path = self._generate_output_path()
logger.info(f"Processing videos with mode: {self.config.mode}, count: {len(video_paths)}")
try:
if self.config.mode == EditingMode.ONE_TAKE:
return self._one_take(video_paths, output_path)
elif self.config.mode == EditingMode.PIP:
return self._pip(video_paths, output_path)
elif self.config.mode == EditingMode.VOICE_OVER:
return self._voice_over(video_paths, audio_path, output_path)
elif self.config.mode == EditingMode.VOICE_PIP:
return self._voice_pip(video_paths, audio_path, output_path)
else:
raise ValueError(f"Unsupported editing mode: {self.config.mode}")
except Exception as e:
logger.error(f"Error processing videos: {e}")
raise
def _validate_inputs(self, video_paths: list[str], audio_path: Optional[str]) -> None:
"""验证输入文件"""
for path in video_paths:
if not os.path.exists(path):
raise FileNotFoundError(f"Video file not found: {path}")
if not os.path.getsize(path) > 0:
raise ValueError(f"Video file is empty: {path}")
if audio_path and not os.path.exists(audio_path):
raise FileNotFoundError(f"Audio file not found: {audio_path}")
def _generate_output_path(self) -> str:
"""生成输出文件路径"""
os.makedirs(self.work_dir, exist_ok=True)
return os.path.join(self.work_dir, f"output_{self.config.mode}_{os.getpid()}.mp4")
def _run_ffmpeg(self, command: list[str], capture_output: bool = True) -> tuple:
"""执行 FFmpeg 命令 — 委托给共享 ffmpeg_utils.run_ffmpeg"""
try:
return run_ffmpeg(command, capture_output=capture_output)
except RuntimeError as e:
logger.error(f"FFmpeg error: {e}")
raise
def _get_video_info(self, video_path: str) -> dict:
"""获取视频信息 — 委托给共享 ffmpeg_utils.probe_video_info,补充 codec/size 字段"""
try:
info = probe_video_info(video_path)
info["codec"] = "unknown"
info["size"] = os.path.getsize(video_path) if os.path.exists(video_path) else 0
return info
except Exception as e:
logger.warning(f"Failed to get video info for {video_path}: {e}")
return {"width": 0, "height": 0, "fps": 25, "duration": 0, "codec": "unknown", "size": 0}
def _get_pip_position_offset(
self, main_width: int, main_height: int, pip_width: int, pip_height: int
) -> tuple[int, int]:
"""获取画中画位置偏移量"""
margin = 10
position_offsets = {
PIPPosition.TOP_LEFT: (margin, margin),
PIPPosition.TOP_RIGHT: (main_width - pip_width - margin, margin),
PIPPosition.BOTTOM_LEFT: (margin, main_height - pip_height - margin),
PIPPosition.BOTTOM_RIGHT: (main_width - pip_width - margin, main_height - pip_height - margin),
}
return position_offsets.get(self.config.pip_position, position_offsets[PIPPosition.TOP_RIGHT])
def _normalize_video(self, input_path: str, output_path: str) -> dict:
"""标准化视频格式:先统一帧率,再缩放/填充"""
command = [
FFMPEG_BIN,
"-y",
"-i",
input_path,
"-r",
str(self.config.output_fps), # 先统一帧率
"-vf",
f"scale={self.config.output_width}:{self.config.output_height}:force_original_aspect_ratio=decrease,pad={self.config.output_width}:{self.config.output_height}:(ow-iw)/2:(oh-ih)/2,setsar=1",
"-r",
str(self.config.output_fps),
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
"-movflags",
"+faststart",
"-an",
output_path,
]
run_ffmpeg(command)
return self._get_video_info(output_path)
def _one_take(self, video_paths: list[str], output_path: str) -> str:
"""一镜到底模式:顺序拼接视频,添加淡入淡出转场"""
if len(video_paths) == 1:
return self._normalize_video(video_paths[0], output_path)
normalized_paths = []
for i, path in enumerate(video_paths):
normalized = os.path.join(self.work_dir, f"normalized_{i}_{os.getpid()}.mp4")
self._normalize_video(path, normalized)
normalized_paths.append(normalized)
durations = [self._get_video_info(p)["duration"] for p in normalized_paths]
if len(normalized_paths) <= 5:
output_path = self._one_take_with_xfade(normalized_paths, durations, output_path)
else:
output_path = self._one_take_simple_concat(normalized_paths, output_path)
for p in normalized_paths:
try:
if p != output_path:
os.remove(p)
except Exception as e:
logger.warning(f"Operation failed in apps/worker/video_processing/editing_modes.py: {e}", exc_info=True)
return output_path
def _one_take_with_xfade(self, normalized_paths: list[str], durations: list[float], output_path: str) -> str:
"""使用 xfade 滤镜实现转场"""
if len(normalized_paths) == 2:
transition = self.config.transition_duration
offset1 = durations[0] - transition / 2
command = [
FFMPEG_BIN,
"-y",
"-i",
normalized_paths[0],
"-i",
normalized_paths[1],
"-filter_complex",
f"[0:v][1:v]xfade=transition=fade:duration={transition}:offset={offset1}[v]",
"-map",
"[v]",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
run_ffmpeg(command)
return output_path
else:
return self._one_take_simple_concat(normalized_paths, output_path)
def _one_take_simple_concat(self, normalized_paths: list[str], output_path: str) -> str:
"""使用 concat demuxer 简单拼接"""
concat_file = os.path.join(self.work_dir, f"concat_list_{os.getpid()}.txt")
with open(concat_file, "w") as f:
for path in normalized_paths:
f.write(f"file '{os.path.abspath(path)}'\n")
command = [
FFMPEG_BIN,
"-y",
"-f",
"concat",
"-safe",
"0",
"-i",
concat_file,
"-c",
"copy",
output_path,
]
run_ffmpeg(command)
try:
os.remove(concat_file)
except Exception as e:
logger.warning(f"Operation failed in apps/worker/video_processing/editing_modes.py: {e}", exc_info=True)
return output_path
def _pip(self, video_paths: list[str], output_path: str) -> str:
"""画中画模式:主视频全屏,后续视频叠加在角落"""
if not video_paths:
raise ValueError("No video paths provided")
main_video = video_paths[0]
main_normalized = os.path.join(self.work_dir, f"main_{os.getpid()}.mp4")
main_info = self._normalize_video(main_video, main_normalized)
if len(video_paths) == 1:
os.rename(main_normalized, output_path)
return output_path
pip_width = int(self.config.output_width * self.config.pip_scale)
pip_height = int(self.config.output_height * self.config.pip_scale)
x_offset, y_offset = self._get_pip_position_offset(
self.config.output_width, self.config.output_height, pip_width, pip_height
)
pip_normalized = os.path.join(self.work_dir, f"pip_{os.getpid()}.mp4")
pip_info = self._get_video_info(video_paths[1])
if pip_info["duration"] > main_info["duration"]:
temp_pip = os.path.join(self.work_dir, f"pip_temp_{os.getpid()}.mp4")
command = [
FFMPEG_BIN,
"-y",
"-i",
video_paths[1],
"-t",
str(main_info["duration"]),
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
temp_pip,
]
run_ffmpeg(command)
pip_normalized_input = temp_pip
else:
command = [
FFMPEG_BIN,
"-y",
"-i",
video_paths[1],
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
pip_normalized,
]
run_ffmpeg(command)
pip_normalized_input = pip_normalized
if main_info["duration"] > pip_info["duration"]:
looped_pip = os.path.join(self.work_dir, f"pip_looped_{os.getpid()}.mp4")
command = [
FFMPEG_BIN,
"-y",
"-stream_loop",
"-1",
"-i",
pip_normalized_input,
"-t",
str(main_info["duration"]),
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
looped_pip,
]
run_ffmpeg(command)
pip_normalized_input = looped_pip
command = [
FFMPEG_BIN,
"-y",
"-i",
main_normalized,
"-i",
pip_normalized_input,
"-filter_complex",
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
"-map",
"[v]",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
run_ffmpeg(command)
for temp_file in [main_normalized, pip_normalized]:
if temp_file and temp_file != output_path:
try:
os.remove(temp_file)
except Exception as e:
logger.warning(
f"Operation failed in apps/worker/video_processing/editing_modes.py: {e}", exc_info=True
)
return output_path
def _voice_over(self, video_paths: list[str], audio_path: Optional[str], output_path: str) -> str:
"""口播模式:背景画面 + 配音"""
if not audio_path:
raise ValueError("audio_path is required for VOICE_OVER mode")
if not video_paths:
raise ValueError("No background video provided")
audio_info = self._get_video_info(audio_path)
audio_duration = audio_info["duration"]
bg_normalized = os.path.join(self.work_dir, f"bg_{os.getpid()}.mp4")
bg_info = self._normalize_video(video_paths[0], bg_normalized)
if bg_info["duration"] < audio_duration:
looped_bg = os.path.join(self.work_dir, f"bg_looped_{os.getpid()}.mp4")
command = [
FFMPEG_BIN,
"-y",
"-stream_loop",
"-1",
"-i",
bg_normalized,
"-t",
str(audio_duration),
"-vf",
f"scale={self.config.output_width}:{self.config.output_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
looped_bg,
]
run_ffmpeg(command)
bg_normalized = looped_bg
elif bg_info["duration"] > audio_duration:
temp_bg = os.path.join(self.work_dir, f"bg_trimmed_{os.getpid()}.mp4")
command = [
FFMPEG_BIN,
"-y",
"-i",
bg_normalized,
"-t",
str(audio_duration),
"-c:v",
"copy",
temp_bg,
]
run_ffmpeg(command)
bg_normalized = temp_bg
blurred_bg = os.path.join(self.work_dir, f"bg_blurred_{os.getpid()}.mp4")
command = [
FFMPEG_BIN,
"-y",
"-i",
bg_normalized,
"-vf",
f"boxblur=5:5,scale={self.config.output_width}:{self.config.output_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
blurred_bg,
]
run_ffmpeg(command)
command = [
FFMPEG_BIN,
"-y",
"-i",
blurred_bg,
"-i",
audio_path,
"-filter_complex",
"[0:v]drawbox=x=0:y=0:w=iw:h=ih:color=black@0.3:t=fill[v]",
"-map",
"[v]",
"-map",
"1:a",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
"-shortest",
output_path,
]
run_ffmpeg(command)
for temp_file in [bg_normalized, blurred_bg]:
try:
if temp_file != output_path:
os.remove(temp_file)
except Exception as e:
logger.warning(f"Operation failed in apps/worker/video_processing/editing_modes.py: {e}", exc_info=True)
return output_path
def _voice_pip(self, video_paths: list[str], audio_path: Optional[str], output_path: str) -> str:
"""口播+画中画模式:口播视频在角落,其他视频作为背景"""
if not video_paths:
raise ValueError("No video paths provided")
if len(video_paths) == 1:
return self._normalize_video(video_paths[0], output_path)
voice_video = video_paths[0]
bg_video = video_paths[1] if len(video_paths) > 1 else video_paths[0]
voice_normalized = os.path.join(self.work_dir, f"voice_{os.getpid()}.mp4")
voice_info = self._normalize_video(voice_video, voice_normalized)
bg_normalized = os.path.join(self.work_dir, f"bg_{os.getpid()}.mp4")
bg_info = self._normalize_video(bg_video, bg_normalized)
final_duration = min(voice_info["duration"], bg_info["duration"])
pip_width = int(self.config.output_width * self.config.pip_scale)
pip_height = int(self.config.output_height * self.config.pip_scale)
x_offset, y_offset = self._get_pip_position_offset(
self.config.output_width, self.config.output_height, pip_width, pip_height
)
voice_adjusted = os.path.join(self.work_dir, f"voice_adj_{os.getpid()}.mp4")
command = [
FFMPEG_BIN,
"-y",
"-i",
voice_normalized,
"-t",
str(final_duration),
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
voice_adjusted,
]
run_ffmpeg(command)
bg_adjusted = os.path.join(self.work_dir, f"bg_adj_{os.getpid()}.mp4")
command = [
FFMPEG_BIN,
"-y",
"-i",
bg_normalized,
"-t",
str(final_duration),
"-c:v",
"copy",
bg_adjusted,
]
run_ffmpeg(command)
if audio_path:
command = [
FFMPEG_BIN,
"-y",
"-i",
bg_adjusted,
"-i",
voice_adjusted,
"-i",
audio_path,
"-filter_complex",
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
"-map",
"[v]",
"-map",
"2:a",
"-shortest",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
else:
command = [
FFMPEG_BIN,
"-y",
"-i",
bg_adjusted,
"-i",
voice_adjusted,
"-filter_complex",
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
"-map",
"[v]",
"-map",
"1:a",
"-shortest",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
run_ffmpeg(command)
for temp_file in [voice_normalized, voice_adjusted, bg_normalized, bg_adjusted]:
try:
if temp_file != output_path:
os.remove(temp_file)
except Exception as e:
logger.warning(f"Operation failed in apps/worker/video_processing/editing_modes.py: {e}", exc_info=True)
return output_path
def create_processor(mode: str, work_dir: Optional[str] = None, **kwargs) -> EditingModeProcessor:
"""便捷工厂函数:创建剪辑模式处理器"""
try:
editing_mode = EditingMode(mode)
except ValueError:
raise ValueError(f"Invalid editing mode: {mode}. Valid modes: {[m.value for m in EditingMode]}")
config = EditingModeConfig(
mode=editing_mode,
output_width=kwargs.get("output_width", 1280),
output_height=kwargs.get("output_height", 720),
output_fps=kwargs.get("output_fps", 25),
pip_position=PIPPosition(kwargs.get("pip_position", "top_right")),
pip_scale=kwargs.get("pip_scale", 0.25),
transition_duration=kwargs.get("transition_duration", 0.5),
)
return EditingModeProcessor(config=config, work_dir=work_dir)
+85 -13
View File
@@ -1,8 +1,7 @@
"""FFmpeg 工具函数 — 从 editing_modes.py / video_compose_service.py 提取的共享原语.
"""FFmpeg 工具函数 — 共享原语.
提供 FFmpeg / FFprobe 调用、视频信息探测、视频标准化、xfade 转场滤镜构建
等底层能力,供 EditingModeProcessor、VideoComposeService、UnifiedRenderService
共同复用。
等底层能力,供 UnifiedRenderService、VideoComposeService 等复用。
"""
from __future__ import annotations
@@ -32,6 +31,10 @@ XFADE_TRANSITION_MAP: dict[str, str] = {
"slide_left": "slideleft",
"slideright": "slideright",
"slide_right": "slideright",
"slideup": "slideup",
"slide_up": "slideup",
"slidedown": "slidedown",
"slide_down": "slidedown",
"dissolve": "dissolve",
"wipe": "wipeleft",
"wipeleft": "wipeleft",
@@ -39,6 +42,10 @@ XFADE_TRANSITION_MAP: dict[str, str] = {
DEFAULT_TRANSITION_DURATION = 0.5
# FFmpeg 执行默认超时(秒),防止 FFmpeg hang 住导致 worker 永久阻塞
# 默认 30 分钟,足够处理大部分短视频渲染;超长视频可单独传参覆盖
DEFAULT_FFMPEG_TIMEOUT = 1800
# ── FFmpeg 执行 ───────────────────────────────────────────────────────────────
@@ -47,12 +54,14 @@ def run_ffmpeg(
command: list[str],
*,
capture_output: bool = True,
timeout: int | None = DEFAULT_FFMPEG_TIMEOUT,
) -> tuple[str, str]:
"""执行 FFmpeg 命令。
Args:
command: 完整的 ffmpeg 命令列表(含 "ffmpeg" 本身)
capture_output: 是否捕获 stdout/stderr
timeout: 超时时间(秒),默认 1800s(30分钟);None 表示不设超时(不推荐)
Returns:
(stdout, stderr) 元组
@@ -60,6 +69,7 @@ def run_ffmpeg(
Raises:
subprocess.CalledProcessError: 命令执行失败时抛出,
异常信息包含完整 stderr 以便排查。
subprocess.TimeoutExpired: 超时未完成时抛出,FFmpeg 进程会被 kill。
"""
try:
result = subprocess.run( # nosec B603
@@ -68,8 +78,16 @@ def run_ffmpeg(
stdout=subprocess.PIPE if capture_output else None,
stderr=subprocess.PIPE if capture_output else None,
text=True,
timeout=timeout,
)
return (result.stdout or "", result.stderr or "")
except subprocess.TimeoutExpired as e:
logger.error(
"FFmpeg 命令超时 (%ds): command=%s",
timeout or -1,
" ".join(str(c) for c in command[:20]),
)
raise
except subprocess.CalledProcessError as e:
# 把完整 stderr 打到日志,方便排查 exit code 183 等问题
stderr_text = (e.stderr or "").strip()
@@ -82,6 +100,41 @@ def run_ffmpeg(
raise
def probe_has_audio(local_path: str | Path) -> bool:
"""探测文件是否包含音频流。
Args:
local_path: 本地文件路径
Returns:
True 表示有音频流(或探测失败保守返回),False 表示确认无音频流
"""
try:
result = subprocess.run( # nosec B603
[
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"a:0",
"-show_entries",
"stream=codec_type",
"-of",
"default=noprint_wrappers=1:nokey=1",
str(local_path),
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=10,
)
return result.stdout.strip() == "audio"
except Exception:
# 探测失败保守返回 True,让 FFmpeg 自己处理(避免误删音频)
return True
def probe_duration(local_path: str | Path) -> float:
"""用 ffprobe 获取视频时长(秒)。
@@ -110,10 +163,14 @@ def probe_duration(local_path: str | Path) -> float:
def probe_video_info(video_path: str) -> dict[str, Any]:
"""获取视频信息(宽、高、时长、fps)。
"""获取视频信息(宽、高、时长、fps、编码、像素格式)。
Returns:
{"width": int, "height": int, "duration": float, "fps": float}
{
"width": int, "height": int, "duration": float, "fps": float,
"video_codec": str, "audio_codec": str, "pix_fmt": str,
"has_audio": bool,
}
失败时返回默认值。
"""
try:
@@ -122,10 +179,8 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"v:0",
"-show_entries",
"stream=width,height,r_frame_rate,duration",
"stream=width,height,r_frame_rate,duration,codec_name,codec_type,pix_fmt",
"-show_entries",
"format=duration",
"-of",
@@ -136,19 +191,25 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=15,
)
import json
info = json.loads(result.stdout)
stream = info.get("streams", [{}])[0]
streams = info.get("streams", [])
fmt = info.get("format", {})
width = int(stream.get("width", DEFAULT_OUTPUT_WIDTH))
height = int(stream.get("height", DEFAULT_OUTPUT_HEIGHT))
video_stream = next((s for s in streams if s.get("codec_type") == "video"), {})
audio_stream = next((s for s in streams if s.get("codec_type") == "audio"), {})
width = int(video_stream.get("width", DEFAULT_OUTPUT_WIDTH))
height = int(video_stream.get("height", DEFAULT_OUTPUT_HEIGHT))
video_codec = video_stream.get("codec_name", "") or ""
pix_fmt = video_stream.get("pix_fmt", "") or ""
# 解析帧率
fps_str = stream.get("r_frame_rate", "25/1")
fps_str = video_stream.get("r_frame_rate", "25/1")
if "/" in fps_str:
num, den = fps_str.split("/")
fps = float(num) / float(den) if float(den) > 0 else DEFAULT_FPS
@@ -156,13 +217,20 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
fps = float(fps_str) if fps_str else DEFAULT_FPS
# 时长
duration = float(fmt.get("duration", 0)) or float(stream.get("duration", 0))
duration = float(fmt.get("duration", 0)) or float(video_stream.get("duration", 0))
has_audio = bool(audio_stream)
audio_codec = audio_stream.get("codec_name", "") or ""
return {
"width": width,
"height": height,
"duration": duration,
"fps": round(fps, 2),
"video_codec": video_codec,
"audio_codec": audio_codec,
"pix_fmt": pix_fmt,
"has_audio": has_audio,
}
except Exception as e:
logger.warning("获取视频信息失败: %s, error: %s", video_path, e)
@@ -171,6 +239,10 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
"height": DEFAULT_OUTPUT_HEIGHT,
"duration": 0.0,
"fps": DEFAULT_FPS,
"video_codec": "",
"audio_codec": "",
"pix_fmt": "",
"has_audio": True,
}
+82 -10
View File
@@ -9,6 +9,7 @@ from __future__ import annotations
import hashlib
import logging
import os
import threading
from pathlib import Path
from typing import Optional
from urllib.parse import urlparse
@@ -17,6 +18,13 @@ import oss2
logger = logging.getLogger(__name__)
# OSS 上传配置
OSS_CONNECT_TIMEOUT = 10 # 连接超时(秒),防止 TCP 握手挂死
OSS_UPLOAD_TOTAL_TIMEOUT = 300 # 单文件上传总超时(秒),防止网络慢时无限卡住
OSS_MULTIPART_THRESHOLD = 100 * 1024 * 1024 # 分片上传阈值:100MB 以上走分片
OSS_PART_SIZE = 8 * 1024 * 1024 # 分片大小:8MB
OSS_MULTIPART_NUM_THREADS = 3 # 分片上传并发数
# ── OSS 配置 ──────────────────────────────────────────────────────────────────
@@ -43,6 +51,9 @@ def oss_bucket() -> oss2.Bucket | None:
P0-2 修复:endpoint 不带 scheme 时自动补 https:// 前缀,
确保 sign_url 等依赖 scheme 的方法返回 HTTPS URL。
P0-staging 修复:增加 connect_timeout=10s,防止网络抖动时
TCP 握手阶段无限挂死,导致 worker 进程卡死。
Returns:
oss2.Bucket 实例,配置缺失时返回 None。
"""
@@ -53,7 +64,12 @@ def oss_bucket() -> oss2.Bucket | None:
# endpoint 无 scheme 时补 https://,与 API 端 storage.py 保持一致
if not endpoint.startswith(("http://", "https://")):
endpoint = f"https://{endpoint}"
return oss2.Bucket(oss2.Auth(access_key_id, access_key_secret), endpoint, bucket_name)
return oss2.Bucket(
oss2.Auth(access_key_id, access_key_secret),
endpoint,
bucket_name,
connect_timeout=OSS_CONNECT_TIMEOUT,
)
def normalize_storage_key(storage_key_or_url: str) -> str:
@@ -96,6 +112,9 @@ def download_asset(asset_storage_key: str, local_path: Path) -> bool:
def upload_to_oss(local_path: Path, storage_key: str) -> str | None:
"""上传文件到 OSS,返回公开 URL。
大文件(>100MB)自动走分片上传,降低内存峰值,减少 OOM 风险。
上传加总超时保护(默认 300s),防止网络异常时无限挂死。
Args:
local_path: 本地文件路径
storage_key: 目标存储键
@@ -106,18 +125,71 @@ def upload_to_oss(local_path: Path, storage_key: str) -> str | None:
bucket = oss_bucket()
if bucket is None:
return None
try:
bucket.put_object_from_file(storage_key, str(local_path))
settings = oss_settings()
if settings:
_, _, endpoint, bucket_name = settings
endpoint_clean = endpoint.replace("https://", "").replace("http://", "")
return f"https://{bucket_name}.{endpoint_clean}/{storage_key}"
result: dict = {"url": None, "error": None, "file_size": 0}
done = threading.Event()
def _do_upload():
try:
# 尝试获取文件大小,用于分片判断和日志;stat 失败时 fallback 走普通上传
try:
file_size = local_path.stat().st_size
result["file_size"] = file_size
use_multipart = file_size >= OSS_MULTIPART_THRESHOLD
except OSError:
use_multipart = False
file_size = 0
if use_multipart:
# 分片上传:降低内存峰值,每片 8MB,3 线程并发
logger.info(
"大文件分片上传: storage_key=%s, size=%.1fMB, part_size=%dMB, threads=%d",
storage_key[:80],
file_size / 1024 / 1024,
OSS_PART_SIZE // 1024 // 1024,
OSS_MULTIPART_NUM_THREADS,
)
oss2.resumable_upload(
bucket,
storage_key,
str(local_path),
multipart_threshold=OSS_MULTIPART_THRESHOLD,
part_size=OSS_PART_SIZE,
num_threads=OSS_MULTIPART_NUM_THREADS,
)
else:
bucket.put_object_from_file(storage_key, str(local_path))
# 构造返回 URL
settings = oss_settings()
if settings:
_, _, endpoint, bucket_name = settings
endpoint_clean = endpoint.replace("https://", "").replace("http://", "")
result["url"] = f"https://{bucket_name}.{endpoint_clean}/{storage_key}"
except Exception as e:
result["error"] = e
logger.exception("上传 OSS 失败: %s", storage_key)
finally:
done.set()
upload_thread = threading.Thread(target=_do_upload, daemon=True)
upload_thread.start()
finished = done.wait(timeout=OSS_UPLOAD_TOTAL_TIMEOUT)
if not finished:
logger.error(
"OSS 上传超时(%.0fs),强制中止: storage_key=%s, size=%.1fMB",
OSS_UPLOAD_TOTAL_TIMEOUT,
storage_key[:80],
result["file_size"] / 1024 / 1024 if result["file_size"] else 0,
)
return None
except Exception:
logger.exception("上传 OSS 失败: %s", storage_key)
if result["error"]:
return None
return result["url"]
def get_signed_download_url(storage_key_or_url: str, expires_seconds: int = 3600) -> str | None:
"""生成预签名下载 URL(用于私有 bucket 的 URL 校验或临时下载)。
+299
View File
@@ -0,0 +1,299 @@
"""统一渲染引擎适配层 — Phase 2.
将 EditPlan + EditPlanClips(来自 DB)适配为 UnifiedRenderService 的输入格式,
封装素材下载、渲染执行、结果上传的完整流程。
职责:
1. 从 DB 读取 EditPlan + EditPlanClips
2. 下载素材到本地,构建 asset_path_map
3. 调用 UnifiedRenderService 执行渲染
4. 上传渲染结果到 OSS
5. 支持进度回调(对接 JobService)
"""
from __future__ import annotations
import logging
import tempfile
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Callable
from sqlalchemy.orm import Session
from video_processing.oss_helpers import download_asset, upload_to_oss
from video_processing.unified_render_service import RenderResult, UnifiedRenderService
from packages.adapters.sqlalchemy_impl.edit_plan_clip_repository import SQLAlchemyEditPlanClipRepository
from packages.adapters.sqlalchemy_impl.edit_plan_repository import SQLAlchemyEditPlanRepository
from packages.domain.edit_plan import EditPlan, EditPlanStatus
from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus
logger = logging.getLogger(__name__)
# ── 数据结构 ──────────────────────────────────────────────────────────────────
@dataclass
class RenderAdapterResult:
"""渲染适配结果。"""
success: bool
output_url: str = ""
output_path: Path | None = None
duration: float = 0.0
file_size: int = 0
width: int = 0
height: int = 0
clip_count: int = 0
error_message: str = ""
ProgressCallback = Callable[[float, str], None]
"""进度回调:(progress_0_100, stage_description) → None"""
# ── 适配层主体 ────────────────────────────────────────────────────────────────
class RenderAdapter:
"""统一渲染引擎适配层。
桥接 EditPlan 领域模型与 UnifiedRenderService 图层模型。
用法::
adapter = RenderAdapter(db)
result = adapter.render_plan(
plan_id=plan_id,
job_id=job_id,
progress_cb=lambda p, s: job_service.update_progress(job_id, p, s),
)
"""
def __init__(self, db: Session) -> None:
self._db = db
self._plan_repo = SQLAlchemyEditPlanRepository(db)
self._clip_repo = SQLAlchemyEditPlanClipRepository(db)
# ── 公开方法 ──────────────────────────────────────────────────────────
def render_plan(
self,
plan_id: str,
*,
job_id: str = "",
work_dir: Path | None = None,
progress_cb: ProgressCallback | None = None,
) -> RenderAdapterResult:
"""渲染一个 EditPlan。
完整流程:
1. 加载计划与片段
2. 下载素材
3. 执行统一渲染
4. 上传结果
Args:
plan_id: EditPlan ID
job_id: 关联的 Job ID(用于结果存储路径)
work_dir: 工作目录,不传则使用临时目录
progress_cb: 进度回调函数
Returns:
RenderAdapterResult
"""
temp_dir = None
try:
# 0. 准备工作目录
if work_dir is None:
temp_dir = tempfile.mkdtemp(prefix="render_")
work_dir = Path(temp_dir)
work_dir.mkdir(parents=True, exist_ok=True)
self._report_progress(progress_cb, 5.0, "加载剪辑计划")
# 1. 加载计划与片段
plan = self._plan_repo.get(plan_id)
if plan is None:
return RenderAdapterResult(
success=False,
error_message=f"剪辑计划不存在: {plan_id}",
)
clips = self._clip_repo.list_by_plan(plan_id, skip=0, limit=10000)
ready_clips = [c for c in clips if c.status == EditPlanClipStatus.READY and c.asset_id]
ready_clips.sort(key=lambda c: c.order)
if not ready_clips:
return RenderAdapterResult(
success=False,
error_message="没有可渲染的就绪片段",
clip_count=0,
)
logger.info(
"开始渲染: plan_id=%s job_id=%s ready_clips=%d engine=unified",
plan_id,
job_id,
len(ready_clips),
)
self._report_progress(progress_cb, 15.0, f"下载素材({len(ready_clips)} 个)")
# 2. 下载素材
asset_path_map = self._download_assets(ready_clips, work_dir)
if not asset_path_map:
return RenderAdapterResult(
success=False,
error_message="所有素材下载失败",
clip_count=len(ready_clips),
)
self._report_progress(progress_cb, 40.0, "执行视频渲染")
# 3. 执行统一渲染
render_svc = UnifiedRenderService(
plan=plan,
clips=ready_clips,
asset_path_map=asset_path_map,
work_dir=work_dir,
)
result = render_svc.render()
self._report_progress(progress_cb, 80.0, "上传渲染结果")
# 4. 上传结果
storage_key = f"rendered/{plan_id}/{job_id or plan_id}.mp4"
output_url = upload_to_oss(result.output_path, storage_key)
self._report_progress(progress_cb, 100.0, "渲染完成")
logger.info(
"[render-adapter] render success: plan_id=%s job_id=%s engine=unified "
"duration=%.2fs file_size=%d resolution=%dx%d clip_count=%d",
plan_id,
job_id,
result.duration,
result.file_size,
result.width,
result.height,
len(ready_clips),
)
return RenderAdapterResult(
success=True,
output_url=output_url or "",
output_path=result.output_path,
duration=result.duration,
file_size=result.file_size,
width=result.width,
height=result.height,
clip_count=len(ready_clips),
)
except Exception as exc:
logger.exception(
"[render-adapter] render failed: plan_id=%s job_id=%s engine=unified error=%s",
plan_id,
job_id,
str(exc)[:200],
)
return RenderAdapterResult(
success=False,
error_message=str(exc)[:500],
)
finally:
# 清理临时目录
if temp_dir:
import shutil
try:
shutil.rmtree(temp_dir, ignore_errors=True)
except Exception:
pass
def validate_plan(self, plan_id: str) -> tuple[bool, list[str], list[str], int, int]:
"""校验计划是否可渲染(兼容 VideoComposeService.validate_compose 接口)。
Returns:
(valid, errors, warnings, ready_clip_count, total_clip_count)
"""
errors: list[str] = []
warnings: list[str] = []
plan = self._plan_repo.get(plan_id)
if plan is None:
return False, [f"剪辑计划不存在: {plan_id}"], [], 0, 0
if plan.status not in (EditPlanStatus.EDITING, EditPlanStatus.RENDERING):
errors.append(f"计划状态不正确,需要 editing 或 rendering,当前: {plan.status}")
clips = self._clip_repo.list_by_plan(plan_id, skip=0, limit=10000)
if not clips:
errors.append("计划没有任何片段")
return False, errors, warnings, 0, 0
clips.sort(key=lambda c: c.order)
ready_count = 0
pending_count = 0
no_asset_count = 0
for clip in clips:
if clip.status == EditPlanClipStatus.READY:
ready_count += 1
if not clip.asset_id:
errors.append(f"片段 {clip.id} (order={clip.order}) 没有分配素材")
no_asset_count += 1
elif clip.status == EditPlanClipStatus.PENDING:
pending_count += 1
elif clip.status == EditPlanClipStatus.FAILED:
warnings.append(f"片段 {clip.id} (order={clip.order}) 状态为 failed,已跳过")
if ready_count == 0:
errors.append("没有就绪(ready)的片段可以合成")
if pending_count > 0:
warnings.append(f"{pending_count} 个片段仍处于 pending 状态")
return len(errors) == 0, errors, warnings, ready_count, len(clips)
# ── 内部方法 ──────────────────────────────────────────────────────────
@staticmethod
def _report_progress(progress_cb: ProgressCallback | None, progress: float, stage: str) -> None:
"""上报进度。"""
if progress_cb is not None:
try:
progress_cb(progress, stage)
except Exception:
logger.exception("进度回调失败")
@staticmethod
def _download_assets(clips: list[EditPlanClip], work_dir: Path) -> dict[str, Path]:
"""下载片段素材到本地,返回 asset_id → local_path 映射。
只保留下载成功的素材。
"""
asset_dir = work_dir / "assets"
asset_dir.mkdir(exist_ok=True)
asset_path_map: dict[str, Path] = {}
for clip in clips:
asset_id = clip.asset_id
if not asset_id:
continue
# 生成安全的本地文件名
safe_name = f"clip_{clip.order:04d}_{abs(hash(asset_id)) % 100000:05d}.mp4"
local_path = asset_dir / safe_name
if download_asset(asset_id, local_path):
asset_path_map[asset_id] = local_path
logger.debug("素材下载成功: clip_id=%s asset_id=%s", clip.id, asset_id[:60])
else:
logger.warning("素材下载失败: clip_id=%s asset_id=%s", clip.id, asset_id[:60])
return asset_path_map
+204
View File
@@ -0,0 +1,204 @@
"""渲染引擎 Feature Flag 解析器。
封装渲染引擎选择逻辑,支持:
- 环境变量作为默认值(RENDER_ENGINE=legacy/unified
- Redis Feature Flag 运行时覆盖(白名单 + 百分比 + 全局开关)
- 定时刷新,支持热更新不重启 worker
使用方式:
resolver = RenderEngineResolver(redis_url="redis://...", default_engine="legacy")
engine = resolver.get_engine(user_id="user123")
# engine: "legacy""unified"
"""
from __future__ import annotations
import logging
import threading
from typing import Optional
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
FeatureFlagStore,
InMemoryFeatureFlagStore,
RedisFeatureFlagStore,
)
logger = logging.getLogger(__name__)
# Feature Flag 名称常量
FLAG_RENDER_ENGINE = "render_engine"
# 引擎常量
ENGINE_LEGACY = "legacy"
ENGINE_UNIFIED = "unified"
VALID_ENGINES = {ENGINE_LEGACY, ENGINE_UNIFIED}
class RenderEngineResolver:
"""渲染引擎选择器。
判定逻辑(从高到低):
1. Redis flag 白名单匹配 → unified
2. Redis flag 百分比命中 → unified
3. Redis flag 全局开启(100%)→ unified
4. 环境变量默认值 → legacy / unified
当 Redis 不可用时,自动降级到环境变量默认值,不影响业务。
"""
def __init__(
self,
default_engine: str = ENGINE_LEGACY,
redis_url: Optional[str] = None,
refresh_interval: float = 30.0,
store: Optional[FeatureFlagStore] = None,
) -> None:
"""
Args:
default_engine: 环境变量默认的引擎名(legacy / unified
redis_url: Redis 连接 URL,传 None 时使用内存实现(测试用)
refresh_interval: Redis flag 配置刷新间隔(秒)
store: 直接传入 store 实例(测试用,优先级高于 redis_url)
"""
self._default_engine = default_engine.lower() if default_engine else ENGINE_LEGACY
if self._default_engine not in VALID_ENGINES:
logger.warning(
"Invalid default engine '%s', fallback to '%s'",
self._default_engine,
ENGINE_LEGACY,
)
self._default_engine = ENGINE_LEGACY
if store is not None:
self._store = store
elif redis_url:
self._store = RedisFeatureFlagStore(redis_url=redis_url)
else:
self._store = InMemoryFeatureFlagStore()
logger.info("No Redis configured, using in-memory feature flag store")
self._refresh_interval = refresh_interval
self._lock = threading.Lock()
self._cached_config: Optional[FeatureFlagConfig] = None
self._last_refresh: float = 0.0
def _maybe_refresh(self) -> None:
"""惰性刷新配置,超过刷新间隔时从存储重新读取。"""
import time
now = time.time()
if now - self._last_refresh < self._refresh_interval:
return
try:
config = self._store.get(FLAG_RENDER_ENGINE)
with self._lock:
self._cached_config = config
self._last_refresh = now
except Exception as exc:
logger.warning("Failed to refresh render engine flag: %s", exc)
# 刷新失败时保留旧缓存,不中断业务
if self._cached_config is None:
# 首次就读失败,设一个默认值
with self._lock:
self._cached_config = FeatureFlagConfig(name=FLAG_RENDER_ENGINE)
self._last_refresh = now
def _get_config(self) -> FeatureFlagConfig:
"""获取当前 flag 配置(带缓存)。"""
if self._cached_config is None:
self._maybe_refresh()
else:
self._maybe_refresh()
return self._cached_config or FeatureFlagConfig(name=FLAG_RENDER_ENGINE)
def get_engine(self, user_id: Optional[str] = None) -> str:
"""获取当前应该使用的渲染引擎。
Args:
user_id: 用户ID,用于白名单匹配和百分比哈希。
传 None 时只看全局开关。
Returns:
"legacy""unified"
"""
config = self._get_config()
# 全局关闭 → 用默认值
if not config.enabled:
return self._default_engine
# 白名单匹配 / 百分比命中 → unified
if config.is_active(user_id):
return ENGINE_UNIFIED
# 未命中灰度 → 用默认值
return self._default_engine
def should_use_unified(self, user_id: Optional[str] = None) -> bool:
"""便捷方法:是否应该使用统一渲染引擎。"""
return self.get_engine(user_id) == ENGINE_UNIFIED
def force_refresh(self) -> None:
"""强制立即刷新配置(用于管理接口修改后立即生效)。"""
self._last_refresh = 0.0
if isinstance(self._store, RedisFeatureFlagStore):
self._store.invalidate_cache(FLAG_RENDER_ENGINE)
self._maybe_refresh()
def get_config_snapshot(self) -> dict:
"""获取当前配置快照(用于管理接口展示)。"""
config = self._get_config()
return {
"flag_name": FLAG_RENDER_ENGINE,
"default_engine": self._default_engine,
"enabled": config.enabled,
"percentage": config.percentage,
"whitelist": sorted(config.whitelist),
"refresh_interval": self._refresh_interval,
"last_refresh": self._last_refresh,
}
def set_flag(self, config: FeatureFlagConfig) -> None:
"""设置 flag 配置(管理接口用)。"""
config.name = FLAG_RENDER_ENGINE
self._store.set(config)
self.force_refresh()
# 全局单例
_resolver: Optional[RenderEngineResolver] = None
_resolver_lock = threading.Lock()
def get_render_engine_resolver() -> RenderEngineResolver:
"""获取全局单例(基于 worker 配置)。"""
global _resolver
if _resolver is not None:
return _resolver
with _resolver_lock:
if _resolver is not None:
return _resolver
try:
from worker_app.core.config import get_settings
settings = get_settings()
redis_url = getattr(settings, "redis_url", None) or getattr(settings, "broker_url", None)
default = getattr(settings, "render_engine", ENGINE_LEGACY)
_resolver = RenderEngineResolver(
default_engine=default,
redis_url=redis_url,
)
logger.info(
"RenderEngineResolver initialized: default=%s, redis=%s",
default,
bool(redis_url),
)
except Exception as exc:
logger.warning("Failed to init RenderEngineResolver from settings: %s", exc)
_resolver = RenderEngineResolver(default_engine=ENGINE_LEGACY)
return _resolver
File diff suppressed because it is too large Load Diff
@@ -1,821 +0,0 @@
"""
视频合成服务
支持多种剪辑模式和转场效果,包含完整的安全校验
"""
import logging
import os
import subprocess
import tempfile
from dataclasses import dataclass
from enum import Enum
try:
from enum import StrEnum
except ImportError:
class StrEnum(str, Enum): # type: ignore[no-redef]
"""Python 3.10 兼容的 StrEnum 回退实现。"""
pass
from pathlib import Path
from typing import Optional
from packages.domain.editing_mode import EditingMode
logger = logging.getLogger(__name__)
# ========== 安全常量 ==========
# 允许的输出目录白名单(使用环境变量或系统临时目录,避免硬编码 /tmp)
_VIDEO_OUTPUT_DIR = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
ALLOWED_OUTPUT_DIRS = [_VIDEO_OUTPUT_DIR, "/var/app/rendered"]
# 允许的输入路径前缀白名单
ALLOWED_INPUT_PREFIXES = ("s3://", "oss://", "local://", "/var/storage/")
# 允许的转场效果白名单
ALLOWED_TRANSITIONS = {
"fade",
"slideleft",
"slideright",
"dissolve",
"wipeleft",
"wiperight",
"cut",
"slideup",
"slidedown",
}
# 转场效果映射
_XFADE_TRANSITION_MAP = {
"fade": "fade",
"slideleft": "slideleft",
"slideright": "slideright",
"dissolve": "dissolve",
"wipeleft": "wipeleft",
"wiperight": "wiperight",
"cut": "cut",
"slideup": "slideup",
"slidedown": "slidedown",
}
class VideoComposeError(Exception):
"""视频合成服务异常"""
pass
class PIPPosition(StrEnum):
"""画中画位置枚举"""
TOP_LEFT = "top_left"
TOP_RIGHT = "top_right"
BOTTOM_LEFT = "bottom_left"
BOTTOM_RIGHT = "bottom_right"
@dataclass
class Clip:
"""视频片段"""
asset_id: str # 资源ID,对应输入路径
start_time: float = 0.0
duration: float = 0.0
transition: str = "fade" # 转场效果
@dataclass
class EditingModeConfig:
"""剪辑模式配置"""
mode: EditingMode
output_width: int = 1280
output_height: int = 720
output_fps: int = 25
pip_position: PIPPosition = PIPPosition.TOP_RIGHT
pip_scale: float = 0.25 # 画中画占主画面的比例
transition_duration: float = 0.5 # 转场时长(秒)
output_codec: str = "libx264"
output_preset: str = "medium"
output_crf: int = 23
class VideoComposeService:
"""视频合成服务"""
def __init__(self, config: EditingModeConfig, work_dir: Optional[str] = None):
"""
初始化视频合成服务
Args:
config: 剪辑模式配置
work_dir: 工作目录,默认使用系统临时目录
"""
self.config = config
self.work_dir = work_dir or tempfile.gettempdir()
self._ffmpeg_bin = "ffmpeg"
self._ffprobe_bin = "ffprobe"
def _validate_output_path(self, path: str) -> str:
"""
校验输出路径是否在允许范围内 (P0 修复)
防止路径穿越攻击,如 /app/config/../../../etc/passwd
Args:
path: 用户提供的输出路径
Returns:
标准化后的绝对路径
Raises:
ValueError: 路径不在允许范围内
"""
abs_path = os.path.abspath(path)
for allowed_dir in ALLOWED_OUTPUT_DIRS:
allowed_abs = os.path.abspath(allowed_dir)
if abs_path.startswith(allowed_abs):
return abs_path
raise ValueError(f"输出路径不在允许范围内: {path}")
def _validate_input_path(self, path: str) -> bool:
"""
校验输入路径格式是否合法 (P1-1 修复)
Args:
path: 输入文件路径
Returns:
是否合法
"""
return any(path.startswith(prefix) for prefix in ALLOWED_INPUT_PREFIXES)
def _validate_transition(self, transition: str) -> str:
"""
校验转场效果是否在白名单内 (P1-2 修复)
Args:
transition: 转场效果名称
Returns:
安全的转场效果名称
"""
if transition not in ALLOWED_TRANSITIONS:
logger.warning(f"未知的转场效果 '{transition}',使用默认 'fade'")
return "fade"
return transition
def _get_validated_transition(self, transition: str) -> str:
"""获取白名单校验后的转场效果名称"""
return _XFADE_TRANSITION_MAP.get(self._validate_transition(transition), "fade")
def compose(self, clips: list[Clip], output_path: Optional[str] = None) -> str:
"""
合成视频
Args:
clips: 视频片段列表,每个片段包含 asset_id 和转场配置
output_path: 输出文件路径
Returns:
输出文件路径
"""
if not clips:
raise ValueError("clips 不能为空")
# P1-1: 校验所有输入路径
for clip in clips:
if not self._validate_input_path(clip.asset_id):
raise ValueError(f"不合法的输入路径: {clip.asset_id}")
# 生成默认输出路径并校验
if output_path is None:
output_path = self._generate_output_path()
# P0: 校验输出路径
validated_output = self._validate_output_path(output_path)
logger.info(f"合成视频,片段数: {len(clips)}, 输出: {validated_output}")
# 获取输入路径列表
input_paths = [clip.asset_id for clip in clips]
try:
if self.config.mode == EditingMode.ONE_TAKE:
return self._one_take(input_paths, validated_output, clips)
elif self.config.mode == EditingMode.PIP:
return self._pip(input_paths, validated_output)
elif self.config.mode == EditingMode.VOICE_OVER:
return self._voice_over(input_paths, validated_output)
elif self.config.mode == EditingMode.VOICE_PIP:
return self._voice_pip(input_paths, validated_output)
else:
raise ValueError(f"不支持的剪辑模式: {self.config.mode}")
except Exception as e:
logger.error(f"视频合成失败: {e}")
raise VideoComposeError(f"视频合成失败: {e}") from e
def _generate_output_path(self) -> str:
"""生成输出文件路径"""
os.makedirs(self.work_dir, exist_ok=True)
return os.path.join(self.work_dir, f"output_{self.config.mode}_{os.getpid()}.mp4")
def _validate_inputs(self, video_paths: list[str], audio_path: Optional[str] = None) -> None:
"""验证输入文件存在"""
for path in video_paths:
if not os.path.exists(path):
raise FileNotFoundError(f"视频文件不存在: {path}")
if not os.path.getsize(path) > 0:
raise ValueError(f"视频文件为空: {path}")
if audio_path and not os.path.exists(audio_path):
raise FileNotFoundError(f"音频文件不存在: {audio_path}")
def _run_ffmpeg(self, command: list[str], capture_output: bool = True) -> tuple:
"""执行 FFmpeg 命令"""
logger.debug(f"Running FFmpeg: {' '.join(command)}")
try:
result = subprocess.run(
command,
check=True,
stdout=subprocess.PIPE if capture_output else None,
stderr=subprocess.PIPE if capture_output else None,
text=capture_output,
)
return result.stdout or "", result.stderr or ""
except subprocess.CalledProcessError as e:
stderr = e.stderr.decode() if e.stderr else str(e)
logger.error(f"FFmpeg error: {stderr}")
raise RuntimeError(f"FFmpeg 执行失败: {stderr}") from e
def _get_video_info(self, video_path: str) -> dict:
"""获取视频信息"""
try:
result = subprocess.run(
[
self._ffprobe_bin,
"-v",
"error",
"-show_entries",
"stream=width,height,r_frame_rate,duration,codec_name",
"-show_entries",
"format=duration,size",
"-of",
"json",
video_path,
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
)
import json
data = json.loads(result.stdout)
streams = data.get("streams", [{}])
video_stream = next((s for s in streams if s.get("codec_type") == "video"), streams[0] if streams else {})
fmt = data.get("format", {})
fps_str = video_stream.get("r_frame_rate", "25/1")
fps_parts = fps_str.split("/")
fps = float(fps_parts[0]) / float(fps_parts[1]) if len(fps_parts) == 2 else float(fps_parts[0])
return {
"width": int(video_stream.get("width", 0)),
"height": int(video_stream.get("height", 0)),
"fps": fps,
"duration": float(fmt.get("duration", 0)),
"codec": video_stream.get("codec_name", "unknown"),
"size": int(fmt.get("size", 0)),
}
except Exception as e:
logger.warning(f"获取视频信息失败 {video_path}: {e}")
return {"width": 0, "height": 0, "fps": 25, "duration": 0, "codec": "unknown", "size": 0}
def _get_pip_position_offset(
self, main_width: int, main_height: int, pip_width: int, pip_height: int
) -> tuple[int, int]:
"""获取画中画位置偏移量"""
margin = 10
position_offsets = {
PIPPosition.TOP_LEFT: (margin, margin),
PIPPosition.TOP_RIGHT: (main_width - pip_width - margin, margin),
PIPPosition.BOTTOM_LEFT: (margin, main_height - pip_height - margin),
PIPPosition.BOTTOM_RIGHT: (main_width - pip_width - margin, main_height - pip_height - margin),
}
return position_offsets.get(self.config.pip_position, position_offsets[PIPPosition.TOP_RIGHT])
def _normalize_video(self, input_path: str, output_path: str) -> dict:
"""标准化视频格式"""
command = [
self._ffmpeg_bin,
"-y",
"-i",
input_path,
"-r",
str(self.config.output_fps),
"-vf",
f"scale={self.config.output_width}:{self.config.output_height}:force_original_aspect_ratio=decrease,pad={self.config.output_width}:{self.config.output_height}:(ow-iw)/2:(oh-ih)/2,setsar=1",
"-r",
str(self.config.output_fps),
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
"-movflags",
"+faststart",
"-an",
output_path,
]
self._run_ffmpeg(command)
return self._get_video_info(output_path)
def _one_take(self, video_paths: list[str], output_path: str, clips: list[Clip]) -> str:
"""一镜到底模式"""
if len(video_paths) == 1:
return self._normalize_video(video_paths[0], output_path)
normalized_paths = []
for i, path in enumerate(video_paths):
normalized = os.path.join(self.work_dir, f"normalized_{i}_{os.getpid()}.mp4")
self._normalize_video(path, normalized)
normalized_paths.append(normalized)
durations = [self._get_video_info(p)["duration"] for p in normalized_paths]
if len(normalized_paths) <= 5:
output_path = self._one_take_with_xfade(normalized_paths, durations, output_path, clips)
else:
output_path = self._one_take_simple_concat(normalized_paths, output_path)
for p in normalized_paths:
try:
if p != output_path:
os.remove(p)
except Exception as e:
logger.warning(
f"Operation failed in apps/worker/video_processing/video_compose_service.py: {e}", exc_info=True
)
return output_path
def _one_take_with_xfade(
self, normalized_paths: list[str], durations: list[float], output_path: str, clips: list[Clip]
) -> str:
"""使用 xfade 滤镜实现转场 (P1-2: 转场参数白名单校验)"""
if len(normalized_paths) == 2:
# 获取当前片段的转场效果并校验白名单
transition = "fade"
if len(clips) > 1:
transition = self._get_validated_transition(clips[1].transition)
trans_duration = self.config.transition_duration
offset1 = durations[0] - trans_duration / 2
command = [
self._ffmpeg_bin,
"-y",
"-i",
normalized_paths[0],
"-i",
normalized_paths[1],
"-filter_complex",
f"[0:v][1:v]xfade=transition={transition}:duration={trans_duration}:offset={offset1}[v]",
"-map",
"[v]",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
self._run_ffmpeg(command)
return output_path
else:
return self._one_take_simple_concat(normalized_paths, output_path)
def _one_take_simple_concat(self, normalized_paths: list[str], output_path: str) -> str:
"""使用 concat demuxer 简单拼接"""
concat_file = os.path.join(self.work_dir, f"concat_list_{os.getpid()}.txt")
with open(concat_file, "w") as f:
for path in normalized_paths:
f.write(f"file '{os.path.abspath(path)}'\n")
command = [
self._ffmpeg_bin,
"-y",
"-f",
"concat",
"-safe",
"0",
"-i",
concat_file,
"-c",
"copy",
output_path,
]
self._run_ffmpeg(command)
try:
os.remove(concat_file)
except Exception as e:
logger.warning(
f"Operation failed in apps/worker/video_processing/video_compose_service.py: {e}", exc_info=True
)
return output_path
def _pip(self, video_paths: list[str], output_path: str) -> str:
"""画中画模式"""
if not video_paths:
raise ValueError("No video paths provided")
main_video = video_paths[0]
main_normalized = os.path.join(self.work_dir, f"main_{os.getpid()}.mp4")
main_info = self._normalize_video(main_video, main_normalized)
if len(video_paths) == 1:
os.rename(main_normalized, output_path)
return output_path
pip_width = int(self.config.output_width * self.config.pip_scale)
pip_height = int(self.config.output_height * self.config.pip_scale)
x_offset, y_offset = self._get_pip_position_offset(
self.config.output_width, self.config.output_height, pip_width, pip_height
)
pip_normalized = os.path.join(self.work_dir, f"pip_{os.getpid()}.mp4")
pip_info = self._get_video_info(video_paths[1])
if pip_info["duration"] > main_info["duration"]:
temp_pip = os.path.join(self.work_dir, f"pip_temp_{os.getpid()}.mp4")
command = [
self._ffmpeg_bin,
"-y",
"-i",
video_paths[1],
"-t",
str(main_info["duration"]),
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
temp_pip,
]
self._run_ffmpeg(command)
pip_normalized_input = temp_pip
else:
command = [
self._ffmpeg_bin,
"-y",
"-i",
video_paths[1],
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
pip_normalized,
]
self._run_ffmpeg(command)
pip_normalized_input = pip_normalized
if main_info["duration"] > pip_info["duration"]:
looped_pip = os.path.join(self.work_dir, f"pip_looped_{os.getpid()}.mp4")
command = [
self._ffmpeg_bin,
"-y",
"-stream_loop",
"-1",
"-i",
pip_normalized_input,
"-t",
str(main_info["duration"]),
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
looped_pip,
]
self._run_ffmpeg(command)
pip_normalized_input = looped_pip
command = [
self._ffmpeg_bin,
"-y",
"-i",
main_normalized,
"-i",
pip_normalized_input,
"-filter_complex",
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
"-map",
"[v]",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
self._run_ffmpeg(command)
for temp_file in [main_normalized, pip_normalized]:
if temp_file and temp_file != output_path:
try:
os.remove(temp_file)
except Exception as e:
logger.warning(
f"Operation failed in apps/worker/video_processing/video_compose_service.py: {e}", exc_info=True
)
return output_path
def _voice_over(self, video_paths: list[str], audio_path: str, output_path: str) -> str:
"""口播模式"""
if not audio_path:
raise ValueError("audio_path is required for VOICE_OVER mode")
if not video_paths:
raise ValueError("No background video provided")
audio_info = self._get_video_info(audio_path)
audio_duration = audio_info["duration"]
bg_normalized = os.path.join(self.work_dir, f"bg_{os.getpid()}.mp4")
bg_info = self._normalize_video(video_paths[0], bg_normalized)
if bg_info["duration"] < audio_duration:
looped_bg = os.path.join(self.work_dir, f"bg_looped_{os.getpid()}.mp4")
command = [
self._ffmpeg_bin,
"-y",
"-stream_loop",
"-1",
"-i",
bg_normalized,
"-t",
str(audio_duration),
"-vf",
f"scale={self.config.output_width}:{self.config.output_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
looped_bg,
]
self._run_ffmpeg(command)
bg_normalized = looped_bg
elif bg_info["duration"] > audio_duration:
temp_bg = os.path.join(self.work_dir, f"bg_trimmed_{os.getpid()}.mp4")
command = [
self._ffmpeg_bin,
"-y",
"-i",
bg_normalized,
"-t",
str(audio_duration),
"-c:v",
"copy",
temp_bg,
]
self._run_ffmpeg(command)
bg_normalized = temp_bg
blurred_bg = os.path.join(self.work_dir, f"bg_blurred_{os.getpid()}.mp4")
command = [
self._ffmpeg_bin,
"-y",
"-i",
bg_normalized,
"-vf",
f"boxblur=5:5,scale={self.config.output_width}:{self.config.output_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
blurred_bg,
]
self._run_ffmpeg(command)
command = [
self._ffmpeg_bin,
"-y",
"-i",
blurred_bg,
"-i",
audio_path,
"-filter_complex",
"[0:v]drawbox=x=0:y=0:w=iw:h=ih:color=black@0.3:t=fill[v]",
"-map",
"[v]",
"-map",
"1:a",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
"-shortest",
output_path,
]
self._run_ffmpeg(command)
for temp_file in [bg_normalized, blurred_bg]:
try:
if temp_file != output_path:
os.remove(temp_file)
except Exception as e:
logger.warning(
f"Operation failed in apps/worker/video_processing/video_compose_service.py: {e}", exc_info=True
)
return output_path
def _voice_pip(self, video_paths: list[str], audio_path: Optional[str], output_path: str) -> str:
"""口播+画中画模式"""
if not video_paths:
raise ValueError("No video paths provided")
if len(video_paths) == 1:
return self._normalize_video(video_paths[0], output_path)
voice_video = video_paths[0]
bg_video = video_paths[1] if len(video_paths) > 1 else video_paths[0]
voice_normalized = os.path.join(self.work_dir, f"voice_{os.getpid()}.mp4")
voice_info = self._normalize_video(voice_video, voice_normalized)
bg_normalized = os.path.join(self.work_dir, f"bg_{os.getpid()}.mp4")
bg_info = self._normalize_video(bg_video, bg_normalized)
final_duration = min(voice_info["duration"], bg_info["duration"])
pip_width = int(self.config.output_width * self.config.pip_scale)
pip_height = int(self.config.output_height * self.config.pip_scale)
x_offset, y_offset = self._get_pip_position_offset(
self.config.output_width, self.config.output_height, pip_width, pip_height
)
voice_adjusted = os.path.join(self.work_dir, f"voice_adj_{os.getpid()}.mp4")
command = [
self._ffmpeg_bin,
"-y",
"-i",
voice_normalized,
"-t",
str(final_duration),
"-vf",
f"scale={pip_width}:{pip_height}",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
voice_adjusted,
]
self._run_ffmpeg(command)
bg_adjusted = os.path.join(self.work_dir, f"bg_adj_{os.getpid()}.mp4")
command = [
self._ffmpeg_bin,
"-y",
"-i",
bg_normalized,
"-t",
str(final_duration),
"-c:v",
"copy",
bg_adjusted,
]
self._run_ffmpeg(command)
if audio_path:
command = [
self._ffmpeg_bin,
"-y",
"-i",
bg_adjusted,
"-i",
voice_adjusted,
"-i",
audio_path,
"-filter_complex",
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
"-map",
"[v]",
"-map",
"2:a",
"-shortest",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
else:
command = [
self._ffmpeg_bin,
"-y",
"-i",
bg_adjusted,
"-i",
voice_adjusted,
"-filter_complex",
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
"-map",
"[v]",
"-map",
"1:a",
"-shortest",
"-c:v",
self.config.output_codec,
"-preset",
self.config.output_preset,
"-crf",
str(self.config.output_crf),
"-pix_fmt",
"yuv420p",
output_path,
]
self._run_ffmpeg(command)
for temp_file in [voice_normalized, voice_adjusted, bg_normalized, bg_adjusted]:
try:
if temp_file != output_path:
os.remove(temp_file)
except Exception as e:
logger.warning(
f"Operation failed in apps/worker/video_processing/video_compose_service.py: {e}", exc_info=True
)
return output_path
def create_compose_service(mode: str, work_dir: Optional[str] = None, **kwargs) -> VideoComposeService:
"""便捷工厂函数:创建视频合成服务"""
try:
editing_mode = EditingMode(mode)
except ValueError:
raise ValueError(f"无效的剪辑模式: {mode}. 有效模式: {[m.value for m in EditingMode]}")
config = EditingModeConfig(
mode=editing_mode,
output_width=kwargs.get("output_width", 1280),
output_height=kwargs.get("output_height", 720),
output_fps=kwargs.get("output_fps", 25),
pip_position=PIPPosition(kwargs.get("pip_position", "top_right")),
pip_scale=kwargs.get("pip_scale", 0.25),
transition_duration=kwargs.get("transition_duration", 0.5),
)
return VideoComposeService(config=config, work_dir=work_dir)
+3
View File
@@ -18,6 +18,9 @@ class WorkerSettings(BaseSettings):
environment: str = "development"
auto_create_schema: bool = False
# 渲染引擎选择:legacy=旧VideoComposeServiceunified=新UnifiedRenderService
render_engine: str = "legacy"
model_config = SettingsConfigDict(
env_file=".env",
env_file_encoding="utf-8",
+155 -65
View File
@@ -39,6 +39,10 @@ def _get_job_service():
def compose_video(self, job_id: str, **kwargs):
"""视频合成任务。
根据 RENDER_ENGINE 配置选择渲染引擎:
- legacy: 旧 VideoComposeServicefilter_complex 模式)
- unified: 新 UnifiedRenderService(图层架构)
Args:
job_id: JobService 中的任务 ID
**kwargs: 来自 Job.payload 的额外参数(plan_id, output_path 等)
@@ -56,66 +60,18 @@ def compose_video(self, job_id: str, **kwargs):
job_service.fail_job(job_id, "Missing plan_id in job payload")
return {"status": "error", "message": "Missing plan_id"}
# 标记为 running
job_service.update_progress(job_id, progress=10.0, current_stage="初始化合成环境")
# 判断使用哪个渲染引擎
# 优先级:Redis Feature Flag(白名单 > 百分比) > 环境变量默认
from video_processing.render_engine_resolver import get_render_engine_resolver
# 延迟导入 VideoComposeService
from apps.api.app.services.video_compose_service import VideoComposeService
resolver = get_render_engine_resolver()
user_id = job.created_by_user_id or None
engine = resolver.get_engine(user_id=user_id)
compose_svc = VideoComposeService(db)
# 校验合成条件
job_service.update_progress(job_id, progress=20.0, current_stage="校验合成条件")
validation = compose_svc.validate_compose(plan_id)
if not validation.valid:
error_msg = "; ".join(validation.errors)
job_service.fail_job(job_id, f"合成校验失败: {error_msg}")
return {"status": "error", "message": error_msg}
# 构建合成命令
job_service.update_progress(job_id, progress=30.0, current_stage="构建 FFmpeg 命令")
_output_dir = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
output_path = os.path.join(_output_dir, f"{job_id}.mp4")
compose_cmd = compose_svc.build_compose_command(plan_id, output_path)
# 执行 FFmpeg
job_service.update_progress(job_id, progress=50.0, current_stage="正在执行视频合成")
logger.info("Executing FFmpeg for job %s, plan %s", job_id, plan_id)
try:
subprocess.run(
compose_cmd.command,
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=3600,
)
except subprocess.CalledProcessError as e:
job_service.fail_job(job_id, f"FFmpeg 执行失败: {e.stderr[:500]}")
raise
# 上传结果
job_service.update_progress(job_id, progress=80.0, current_stage="上传合成结果")
storage_key = f"rendered/{plan_id}/{job_id}.mp4"
from worker_app.tasks.edit_plan_generation import _upload_to_oss
output_url = _upload_to_oss(Path(output_path), storage_key)
# 更新 Job 状态为完成
result_data = {
"plan_id": plan_id,
"output_path": output_path,
"storage_key": storage_key,
"output_url": output_url or "",
"estimated_duration": compose_cmd.estimated_duration,
"clip_count": len(compose_cmd.clip_chains),
}
job_service.complete_job(job_id, result=result_data)
logger.info("视频合成完成: job_id=%s, plan_id=%s", job_id, plan_id)
return {"status": "completed", "job_id": job_id, "result": result_data}
if engine == "unified":
return _compose_with_unified_engine(self, job_service, job, plan_id, db)
else:
return _compose_with_legacy_engine(self, job_service, job, plan_id, db)
except self.retry_exc as exc:
logger.warning("视频合成重试中: job_id=%s, exc=%s", job_id, exc)
@@ -129,11 +85,145 @@ def compose_video(self, job_id: str, **kwargs):
raise self.retry(exc=exc, countdown=60)
finally:
db.close()
# 清理临时文件
def _compose_with_legacy_engine(task, job_service, job, plan_id: str, db) -> dict:
"""旧引擎渲染路径(VideoComposeService)。"""
job_id = job.id
# 标记为 running
job_service.update_progress(job_id, progress=10.0, current_stage="初始化合成环境")
# 延迟导入 VideoComposeService
from apps.api.app.services.video_compose_service import VideoComposeService
compose_svc = VideoComposeService(db)
# 校验合成条件
job_service.update_progress(job_id, progress=20.0, current_stage="校验合成条件")
validation = compose_svc.validate_compose(plan_id)
if not validation.valid:
error_msg = "; ".join(validation.errors)
job_service.fail_job(job_id, f"合成校验失败: {error_msg}")
return {"status": "error", "message": error_msg}
# 构建合成命令
job_service.update_progress(job_id, progress=30.0, current_stage="构建 FFmpeg 命令")
_output_dir = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
output_path = os.path.join(_output_dir, f"{job_id}.mp4")
compose_cmd = compose_svc.build_compose_command(plan_id, output_path)
# 执行 FFmpeg
job_service.update_progress(job_id, progress=50.0, current_stage="正在执行视频合成")
logger.info("Executing FFmpeg for job %s, plan %s", job_id, plan_id)
try:
subprocess.run(
compose_cmd.command,
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=3600,
)
except subprocess.CalledProcessError as e:
job_service.fail_job(job_id, f"FFmpeg 执行失败: {e.stderr[:500]}")
raise
# 上传结果
job_service.update_progress(job_id, progress=80.0, current_stage="上传合成结果")
storage_key = f"rendered/{plan_id}/{job_id}.mp4"
from worker_app.tasks.edit_plan_generation import _upload_to_oss
output_url = _upload_to_oss(Path(output_path), storage_key)
# 更新 Job 状态为完成
result_data = {
"plan_id": plan_id,
"output_path": output_path,
"storage_key": storage_key,
"output_url": output_url or "",
"estimated_duration": compose_cmd.estimated_duration,
"clip_count": len(compose_cmd.clip_chains),
"engine": "legacy",
}
job_service.complete_job(job_id, result=result_data)
logger.info("视频合成完成(legacy): job_id=%s, plan_id=%s", job_id, plan_id)
return {"status": "completed", "job_id": job_id, "result": result_data}
def _compose_with_unified_engine(task, job_service, job, plan_id: str, db) -> dict:
"""新引擎渲染路径(UnifiedRenderService + RenderAdapter)。"""
job_id = job.id
# 标记为 running
job_service.update_progress(job_id, progress=10.0, current_stage="初始化统一渲染引擎")
from video_processing.render_adapter import RenderAdapter
adapter = RenderAdapter(db)
# 校验合成条件
job_service.update_progress(job_id, progress=15.0, current_stage="校验合成条件")
valid, errors, warnings, ready_count, total_count = adapter.validate_plan(plan_id)
if not valid:
error_msg = "; ".join(errors)
job_service.fail_job(job_id, f"合成校验失败: {error_msg}")
return {"status": "error", "message": error_msg}
# 进度回调
def progress_cb(progress: float, stage: str) -> None:
try:
_output_dir = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
output_path = os.path.join(_output_dir, f"{job_id}.mp4")
if Path(output_path).exists():
Path(output_path).unlink()
except Exception as e:
logger.warning(f"Operation failed in apps/worker/worker_app/tasks/compose_video.py: {e}", exc_info=True)
job_service.update_progress(job_id, progress=progress, current_stage=stage)
except Exception:
logger.exception("更新进度失败")
# 执行渲染
job_service.update_progress(job_id, progress=20.0, current_stage="开始渲染")
logger.info("统一渲染引擎开始: job_id=%s plan_id=%s", job_id, plan_id)
result = adapter.render_plan(
plan_id=plan_id,
job_id=job_id,
progress_cb=progress_cb,
)
if not result.success:
job_service.fail_job(job_id, f"渲染失败: {result.error_message}")
raise RuntimeError(result.error_message)
# 更新 Job 状态为完成
result_data = {
"plan_id": plan_id,
"output_path": str(result.output_path) if result.output_path else "",
"storage_key": f"rendered/{plan_id}/{job_id}.mp4",
"output_url": result.output_url,
"estimated_duration": result.duration,
"clip_count": result.clip_count,
"engine": "unified",
"width": result.width,
"height": result.height,
"file_size": result.file_size,
}
job_service.complete_job(job_id, result=result_data)
logger.info(
"视频合成完成(unified): job_id=%s plan_id=%s duration=%.2fs",
job_id,
plan_id,
result.duration,
)
return {"status": "completed", "job_id": job_id, "result": result_data}
def _cleanup_output(job_id: str) -> None:
"""清理临时输出文件。"""
try:
_output_dir = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
output_path = os.path.join(_output_dir, f"{job_id}.mp4")
if Path(output_path).exists():
Path(output_path).unlink()
except Exception as e:
logger.warning(f"清理输出文件失败: {e}", exc_info=True)
+325 -101
View File
@@ -1,13 +1,18 @@
"""剪辑计划渲染任务 — Phase 8 任务 2.05.
"""剪辑计划渲染任务 — 支持 Feature Flag 灰度.
Celery 任务 worker.render_edit_plan:
1. 加载 EditPlan + EditPlanClips
2. 下载各片段素材
3. 使用 UnifiedRenderService 按时间线+图层渲染
2. 根据 Feature Flag 选择渲染引擎(legacy / unified
3. 下载各片段素材 + 渲染
4. 上传渲染结果到 OSS
5. 创建 GeneratedVideo 记录 + 查重
6. 更新 EditPlan / EditPlanClip 状态
7. 更新 GenerationTask 进度
渲染引擎灰度:
- 走 Feature Flag (render_engine) 控制
- legacy: VideoComposeService + FFmpeg filter_complex
- unified: UnifiedRenderService 图层架构
"""
from __future__ import annotations
@@ -63,14 +68,268 @@ def _get_repos():
# ── Celery Task ───────────────────────────────────────────────────────────────
def _resolve_render_engine(user_id: str) -> str:
"""根据 Feature Flag 决定使用哪个渲染引擎。
Returns:
"legacy""unified"
"""
try:
from video_processing.render_engine_resolver import get_render_engine_resolver
resolver = get_render_engine_resolver()
return resolver.get_engine(user_id=user_id)
except Exception as exc:
logger.warning("获取渲染引擎配置失败,fallback 到 legacy: %s", exc)
return "legacy"
def _mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, error_msg: str):
"""统一的计划失败标记工具。"""
plan = plan_repo.get(plan_id)
if plan and plan.status.value == "rendering":
plan.mark_failed()
plan_repo.update(plan)
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task and gen_task.status.value != "failed":
gen_task.status = "failed"
gen_task.error_message = error_msg
gen_task.completed_at = datetime.now(timezone.utc)
gen_task_repo.update(gen_task)
def _finalize_render_success(
plan,
plan_repo,
clip_repo,
gen_task_repo,
db,
plan_id: str,
output_url: str,
storage_key: str,
duration: float,
file_size: int,
width: int,
height: int,
rendered_clip_ids: list[str],
failed_clip_ids: list[str],
generation_task_id: str,
output_path: Path,
engine: str,
) -> dict:
"""渲染成功后的统一收尾:查重 + 更新状态 + 返回结果。"""
# 创建 GeneratedVideo 记录 + 查重
project_id = plan.project_id or ""
batch_id = plan.config.get("batch_id", "")
mode = plan.config.get("mode", "edit_plan")
if generation_task_id and project_id:
try:
create_video_record_and_dedup(
generation_task_id=generation_task_id,
project_id=project_id,
batch_id=batch_id,
file_url=output_url or "",
file_size=file_size,
duration=duration,
video_path=str(output_path),
mode=mode,
session=db,
width=width,
height=height,
fps=OUTPUT_FPS,
)
except Exception as dedup_err:
logger.warning("查重失败(不影响渲染结果): %s", dedup_err)
# 更新片段状态为 rendered
for clip_id in rendered_clip_ids:
clip = clip_repo.get(clip_id)
if clip and clip.status.value == "ready":
clip.mark_rendered()
clip_repo.update(clip)
# 更新 EditPlan 状态为 completed
plan.config["rendered_url"] = output_url or ""
plan.config["rendered_storage_key"] = storage_key
plan.mark_completed()
plan_repo.update(plan)
# 更新 GenerationTask 状态为 completed
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task:
gen_task.status = "completed"
gen_task.progress = 100.0
gen_task.result_count = len(rendered_clip_ids)
gen_task.completed_at = datetime.now(timezone.utc)
gen_task_repo.update(gen_task)
logger.info(
"剪辑计划渲染完成: plan_id=%s engine=%s rendered=%d failed=%d duration=%.1fs",
plan_id,
engine,
len(rendered_clip_ids),
len(failed_clip_ids),
duration,
)
return {
"status": "completed",
"plan_id": plan_id,
"rendered_count": len(rendered_clip_ids),
"failed_count": len(failed_clip_ids),
"output_url": output_url,
"duration": duration,
}
def _render_with_unified(
plan,
clips,
asset_path_map: dict[str, Path],
tmpdir_path: Path,
rendered_clip_ids: list[str],
plan_id: str,
generation_task_id: str,
plan_repo,
clip_repo,
gen_task_repo,
db,
) -> dict:
"""统一渲染引擎路径(UnifiedRenderService 图层架构)。"""
render_service = UnifiedRenderService(
plan=plan,
clips=clips,
asset_path_map=asset_path_map,
work_dir=tmpdir_path,
output_width=OUTPUT_WIDTH,
output_height=OUTPUT_HEIGHT,
output_fps=int(OUTPUT_FPS),
)
try:
render_result = render_service.render()
except Exception as render_err:
logger.error("渲染失败(unified): %s%s", plan_id, render_err)
_mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, f"渲染失败: {render_err}")
return {"status": "error", "message": f"渲染失败: {render_err}"}
output_path = render_result.output_path
# 上传到 OSS
storage_key = f"rendered/{plan_id}/output.mp4"
output_url = upload_to_oss(output_path, storage_key)
failed_clip_ids: list[str] = []
return _finalize_render_success(
plan=plan,
plan_repo=plan_repo,
clip_repo=clip_repo,
gen_task_repo=gen_task_repo,
db=db,
plan_id=plan_id,
output_url=output_url or "",
storage_key=storage_key,
duration=render_result.duration,
file_size=render_result.file_size,
width=render_result.width,
height=render_result.height,
rendered_clip_ids=rendered_clip_ids,
failed_clip_ids=failed_clip_ids,
generation_task_id=generation_task_id,
output_path=output_path,
engine="unified",
)
def _render_with_legacy(
plan,
clips,
rendered_clip_ids: list[str],
failed_clip_ids: list[str],
tmpdir_path: Path,
plan_id: str,
generation_task_id: str,
plan_repo,
clip_repo,
gen_task_repo,
db,
) -> dict:
"""旧引擎路径(VideoComposeService + FFmpeg filter_complex)。"""
import os
import subprocess
from apps.api.app.services.video_compose_service import VideoComposeService
compose_svc = VideoComposeService(db)
# 校验合成条件
validation = compose_svc.validate_compose(plan_id)
if not validation.valid:
error_msg = "; ".join(validation.errors)
logger.error("合成校验失败(legacy): %s%s", plan_id, error_msg)
_mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, f"合成校验失败: {error_msg}")
return {"status": "error", "message": error_msg}
# 构建 FFmpeg 命令
output_dir = os.environ.get("VIDEO_OUTPUT_DIR", str(tmpdir_path))
output_path = Path(output_dir) / f"{plan_id}.mp4"
compose_cmd = compose_svc.build_compose_command(plan_id, str(output_path))
logger.info("执行 FFmpeg (legacy): plan_id=%s", plan_id)
try:
subprocess.run(
compose_cmd.command,
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=3600,
)
except subprocess.CalledProcessError as e:
error_msg = f"FFmpeg 执行失败: {e.stderr[:500]}"
logger.error("FFmpeg 执行失败(legacy): %s%s", plan_id, error_msg)
_mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, error_msg)
return {"status": "error", "message": error_msg}
# 获取文件大小
file_size = output_path.stat().st_size if output_path.exists() else 0
duration = compose_cmd.estimated_duration or 0.0
# 上传到 OSS
storage_key = f"rendered/{plan_id}/output.mp4"
output_url = upload_to_oss(output_path, storage_key)
return _finalize_render_success(
plan=plan,
plan_repo=plan_repo,
clip_repo=clip_repo,
gen_task_repo=gen_task_repo,
db=db,
plan_id=plan_id,
output_url=output_url or "",
storage_key=storage_key,
duration=duration,
file_size=file_size,
width=OUTPUT_WIDTH,
height=OUTPUT_HEIGHT,
rendered_clip_ids=rendered_clip_ids,
failed_clip_ids=failed_clip_ids,
generation_task_id=generation_task_id,
output_path=output_path,
engine="legacy",
)
@celery_app.task(name="worker.render_edit_plan", bind=True, max_retries=2)
def render_edit_plan(self, plan_id: str) -> dict:
"""渲染剪辑计划
流程:
1. 加载 EditPlan + EditPlanClips
2. 下载各片段素材到临时目录,构建 asset_path_map
3. 使用 UnifiedRenderService 按时间线+图层渲染
2. 根据 Feature Flag 选择渲染引擎(legacy / unified
3. 下载素材 + 渲染
4. 上传渲染结果到 OSS
5. 创建 GeneratedVideo 记录 + 查重
6. 更新 EditPlan → completed, EditPlanClips → rendered
@@ -79,6 +338,7 @@ def render_edit_plan(self, plan_id: str) -> dict:
logger.info("开始渲染剪辑计划: plan_id=%s", plan_id)
generation_task_id = ""
engine = "legacy"
for repos in _get_repos():
plan_repo, clip_repo, gen_task_repo, db = repos
@@ -93,7 +353,12 @@ def render_edit_plan(self, plan_id: str) -> dict:
# 获取 generation_task_id(提前读取,确保 except 块可用)
generation_task_id = plan.config.get("generation_task_id", "")
# 2. 加载片段列表(按 order 排序
# 2. 选择渲染引擎(Feature Flag 灰度控制
user_id = plan.created_by_user_id or ""
engine = _resolve_render_engine(user_id)
logger.info("剪辑计划渲染引擎: plan_id=%s engine=%s user_id=%s", plan_id, engine, user_id)
# 3. 加载片段列表(按 order 排序)
clips = clip_repo.list_by_plan(plan_id, skip=0, limit=10000)
if not clips:
logger.warning("剪辑计划没有片段: %s", plan_id)
@@ -116,6 +381,15 @@ def render_edit_plan(self, plan_id: str) -> dict:
rendered_clip_ids: list[str] = []
failed_clip_ids: list[str] = []
# 预先批量查询所有素材的 storage_keyfile_url
from packages.adapters.sqlalchemy_impl.models import AssetModel
clip_asset_ids = [c.asset_id for c in clips if c.asset_id]
asset_storage_map: dict[str, str] = {}
if clip_asset_ids:
assets = db.query(AssetModel).filter(AssetModel.id.in_(clip_asset_ids)).all()
asset_storage_map = {a.id: a.file_url for a in assets if a.file_url}
for clip in clips:
if not clip.asset_id:
# 没有素材的片段跳过,标记为失败
@@ -129,10 +403,22 @@ def render_edit_plan(self, plan_id: str) -> dict:
rendered_clip_ids.append(clip.id)
continue
storage_key = asset_storage_map.get(clip.asset_id)
if not storage_key:
logger.warning(
"片段素材无 storage_key,跳过: clip_id=%s asset_id=%s",
clip.id,
clip.asset_id,
)
clip.mark_failed()
clip_repo.update(clip)
failed_clip_ids.append(clip.id)
continue
# 下载素材
ext = Path(clip.asset_id).suffix or ".mp4"
ext = Path(storage_key).suffix or ".mp4"
local_path = tmpdir_path / f"clip_{clip.order:04d}{ext}"
if download_asset(clip.asset_id, local_path):
if download_asset(storage_key, local_path):
asset_path_map[clip.asset_id] = local_path
rendered_clip_ids.append(clip.id)
else:
@@ -153,100 +439,38 @@ def render_edit_plan(self, plan_id: str) -> dict:
gen_task_repo.update(gen_task)
return {"status": "error", "message": "所有片段素材下载失败"}
# 4. 使用 UnifiedRenderService 渲染
render_service = UnifiedRenderService(
plan=plan,
clips=clips,
asset_path_map=asset_path_map,
work_dir=tmpdir_path,
output_width=OUTPUT_WIDTH,
output_height=OUTPUT_HEIGHT,
output_fps=int(OUTPUT_FPS),
)
# 4. 根据引擎选择渲染方式
if engine == "unified":
result = _render_with_unified(
plan=plan,
clips=clips,
asset_path_map=asset_path_map,
tmpdir_path=tmpdir_path,
rendered_clip_ids=rendered_clip_ids,
plan_id=plan_id,
generation_task_id=generation_task_id,
plan_repo=plan_repo,
clip_repo=clip_repo,
gen_task_repo=gen_task_repo,
db=db,
)
else:
result = _render_with_legacy(
plan=plan,
clips=clips,
rendered_clip_ids=rendered_clip_ids,
failed_clip_ids=failed_clip_ids,
tmpdir_path=tmpdir_path,
plan_id=plan_id,
generation_task_id=generation_task_id,
plan_repo=plan_repo,
clip_repo=clip_repo,
gen_task_repo=gen_task_repo,
db=db,
)
try:
render_result = render_service.render()
except Exception as render_err:
logger.error("渲染失败: %s%s", plan_id, render_err)
plan.mark_failed()
plan_repo.update(plan)
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task:
gen_task.status = "failed"
gen_task.error_message = f"渲染失败: {render_err}"
gen_task.completed_at = datetime.now(timezone.utc)
gen_task_repo.update(gen_task)
return {"status": "error", "message": f"渲染失败: {render_err}"}
output_path = render_result.output_path
# 5. 上传到 OSS
storage_key = f"rendered/{plan_id}/output.mp4"
output_url = upload_to_oss(output_path, storage_key)
# 6. 创建 GeneratedVideo 记录 + 查重
project_id = plan.project_id or ""
batch_id = plan.config.get("batch_id", "")
mode = plan.config.get("mode", "edit_plan")
if generation_task_id and project_id:
try:
create_video_record_and_dedup(
generation_task_id=generation_task_id,
project_id=project_id,
batch_id=batch_id,
file_url=output_url or "",
file_size=render_result.file_size,
duration=render_result.duration,
video_path=str(output_path),
mode=mode,
session=db,
width=render_result.width,
height=render_result.height,
fps=OUTPUT_FPS,
)
except Exception as dedup_err:
logger.warning("查重失败(不影响渲染结果): %s", dedup_err)
# 7. 更新片段状态为 rendered
for clip_id in rendered_clip_ids:
clip = clip_repo.get(clip_id)
if clip and clip.status.value == "ready":
clip.mark_rendered()
clip_repo.update(clip)
# 8. 更新 EditPlan 状态为 completed
plan.config["rendered_url"] = output_url or ""
plan.config["rendered_storage_key"] = storage_key
plan.mark_completed()
plan_repo.update(plan)
# 9. 更新 GenerationTask 状态为 completed
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task:
gen_task.status = "completed"
gen_task.progress = 100.0
gen_task.result_count = len(rendered_clip_ids)
gen_task.completed_at = datetime.now(timezone.utc)
gen_task_repo.update(gen_task)
logger.info(
"剪辑计划渲染完成: plan_id=%s rendered=%d failed=%d duration=%.1fs",
plan_id,
len(rendered_clip_ids),
len(failed_clip_ids),
render_result.duration,
)
return {
"status": "completed",
"plan_id": plan_id,
"rendered_count": len(rendered_clip_ids),
"failed_count": len(failed_clip_ids),
"output_url": output_url,
"duration": render_result.duration,
}
result["engine"] = engine
return result
except Exception as exc:
logger.exception("渲染剪辑计划异常: %s", plan_id)
+229 -45
View File
@@ -113,6 +113,7 @@ from video_processing.oss_helpers import (
get_signed_download_url,
upload_to_oss,
)
from video_processing.render_engine_resolver import ENGINE_LEGACY, ENGINE_UNIFIED
from video_processing.unified_render_service import UnifiedRenderService
# ── 虚拟 Plan / Clip(内存中构建,不写数据库) ────────────────────────────────
@@ -124,6 +125,7 @@ class _VirtualPlan:
id: str
name: str = ""
config: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -161,13 +163,15 @@ def _build_plan_and_clips_from_task(
"""
plan = _VirtualPlan(id=task_id, name=f"Generated-{task_id[:8]}")
# 为每个下载路径生成合成 asset_id
# 为每个下载路径生成合成 asset_id,并预探测素材时长
asset_path_map: dict[str, Path] = {}
path_to_asset_id: dict[Path, str] = {}
path_duration: dict[Path, float] = {}
for i, p in enumerate(downloaded_paths):
asset_id = f"gen_{task_id[:8]}_{i:03d}{p.suffix or '.mp4'}"
asset_path_map[asset_id] = p
path_to_asset_id[p] = asset_id
path_duration[p] = probe_duration(p)
clips: list[_VirtualClip] = []
n = len(downloaded_paths)
@@ -183,6 +187,7 @@ def _build_plan_and_clips_from_task(
clip_type=clip_type,
order=i,
asset_id=path_to_asset_id[p],
duration=path_duration[p],
)
)
elif mode == "voice_over":
@@ -195,6 +200,7 @@ def _build_plan_and_clips_from_task(
clip_type="main",
order=i,
asset_id=path_to_asset_id[p],
duration=path_duration[p],
config={"role": "b_roll"},
)
)
@@ -214,6 +220,7 @@ def _build_plan_and_clips_from_task(
clip_type=clip_type,
order=i,
asset_id=path_to_asset_id[p],
duration=path_duration[p],
)
)
else:
@@ -226,6 +233,7 @@ def _build_plan_and_clips_from_task(
clip_type="main",
order=i,
asset_id=path_to_asset_id[p],
duration=path_duration[p],
)
)
@@ -380,31 +388,37 @@ def _download_library_assets(
session = SessionLocal()
try:
# 构建查询:根据模式选择不同的过滤条件
# 构建查询
query = session.query(AssetModel).filter(
AssetModel.status == "ready",
AssetModel.file_type.in_(["video", "video/mp4", "video/quicktime"]),
)
if asset_library_id:
# 素材库模式
query = query.filter(AssetModel.asset_library_id == asset_library_id)
if asset_ids:
# 明确指定了 asset_ids:直接按 ID 查,不预先按 library/project 过滤
# 避免项目级素材或跨库素材因为 library_id 不匹配而查不到
# 归属安全由后面的归属校验保证
query = query.filter(AssetModel.id.in_(asset_ids))
logger.info(
"下载素材库视频: asset_library_id=%s asset_ids=%s",
asset_library_id,
asset_ids or "all",
"下载指定素材: asset_ids=%d 个, asset_library_id=%s, project_id=%s",
len(asset_ids),
asset_library_id or "none",
project_id or "none",
)
else:
# 项目级模式
query = query.filter(AssetModel.project_id == project_id)
logger.info(
"下载项目级视频: project_id=%s asset_ids=%s",
project_id,
asset_ids or "all",
)
if asset_ids:
query = query.filter(AssetModel.id.in_(asset_ids))
# 未指定 asset_ids:按 library 或 project 下载全部 ready 视频
if asset_library_id:
query = query.filter(AssetModel.asset_library_id == asset_library_id)
logger.info(
"下载素材库全部视频: asset_library_id=%s",
asset_library_id,
)
else:
query = query.filter(AssetModel.project_id == project_id)
logger.info(
"下载项目全部视频: project_id=%s",
project_id,
)
assets = query.order_by(AssetModel.created_at).all()
@@ -421,13 +435,15 @@ def _download_library_assets(
if missing_ids:
raise ValueError(f"素材不存在: asset_ids={sorted(missing_ids)}")
for asset in assets:
# 校验素材库归属(只要传了 asset_library_id 就校验)
if asset_library_id and asset.asset_library_id != asset_library_id:
raise ValueError(
f"素材不属于指定素材库: asset_id={asset.id}, "
f"expected_asset_library_id={asset_library_id}, "
f"actual_asset_library_id={asset.asset_library_id}"
)
if not asset_library_id and project_id and asset.project_id != project_id:
# 校验项目归属(只要传了 project_id 就校验)
if project_id and asset.project_id != project_id:
raise ValueError(
f"素材不属于指定项目: asset_id={asset.id}, "
f"expected_project_id={project_id}, "
@@ -558,6 +574,148 @@ def _validate_template_exists(template_id: str) -> None:
session.close()
# ── 渲染引擎选择 ─────────────────────────────────────────────────────────────
def _resolve_render_engine(user_id: str) -> str:
"""根据 Feature Flag 决定使用哪个渲染引擎。
Returns:
"legacy""unified"
"""
try:
from video_processing.render_engine_resolver import get_render_engine_resolver
resolver = get_render_engine_resolver()
return resolver.get_engine(user_id=user_id)
except Exception as exc:
logger.warning("获取渲染引擎配置失败,fallback 到 unified: %s", exc)
return ENGINE_UNIFIED
# ── 旧引擎渲染(FFmpeg filter_complex) ────────────────────────────────────────
def _render_with_legacy_engine(
task_id: str,
virtual_clips: list[_VirtualClip],
asset_path_map: dict[str, Path],
work_dir: Path,
output_path: Path,
) -> tuple[float, int]:
"""旧引擎渲染路径:手动构建 FFmpeg filter_complex 命令。
说明:generate_video 任务使用虚拟 clips(无 EditPlan 数据库记录),
因此无法直接复用 VideoComposeService。这里手动构建等价的 filter_complex
命令,与旧引擎行为一致(scale → crop → setpts → trim → setpts
无 fps 归一化,保持原帧率)。
支持模式:one_take / pip / voice_over / voice_pip
- 所有模式统一走 concat 滤镜(与旧引擎多片段逻辑一致)
Returns:
(duration_seconds, file_size_bytes)
"""
import subprocess
main_clips = [
c
for c in virtual_clips
if c.clip_type in ("main", "b_roll", "background")
or (c.clip_type == "main" and c.config.get("role") == "b_roll")
]
if not main_clips:
main_clips = virtual_clips[:1]
input_args: list[str] = []
video_filters: list[str] = []
audio_filters: list[str] = []
for i, clip in enumerate(main_clips):
local_path = asset_path_map.get(clip.asset_id)
if not local_path:
continue
input_args.extend(["-i", str(local_path)])
duration = clip.duration or 0.0
# 视频滤镜:scale → crop → setpts → trim → setpts(与旧引擎一致)
vf = (
f"[{i}:v]"
f"scale={OUTPUT_WIDTH}:{OUTPUT_HEIGHT}:force_original_aspect_ratio=increase,"
f"crop={OUTPUT_WIDTH}:{OUTPUT_HEIGHT},"
f"setpts=PTS-STARTPTS,"
f"trim=0:{duration:.3f},"
f"setpts=PTS-STARTPTS"
f"[v{i}]"
)
video_filters.append(vf)
# 音频滤镜:atrim → asetpts
af = f"[{i}:a]atrim=0:{duration:.3f},asetpts=PTS-STARTPTS[a{i}]"
audio_filters.append(af)
n = len(main_clips)
if n == 1:
video_label = "[v0]"
audio_label = "[a0]"
else:
# concat 视频
v_inputs = "".join(f"[v{i}]" for i in range(n))
video_filters.append(f"{v_inputs}concat=n={n}:v=1:a=0[outv]")
# concat 音频
a_inputs = "".join(f"[a{i}]" for i in range(n))
audio_filters.append(f"{a_inputs}concat=n={n}:v=0:a=1[outa]")
video_label = "[outv]"
audio_label = "[outa]"
# 组装 filter_complex
fc_parts = video_filters + audio_filters
filter_complex = ";".join(fc_parts)
command = [
FFMPEG_BIN,
"-y",
*input_args,
"-filter_complex",
filter_complex,
"-map",
video_label,
"-map",
audio_label,
"-c:v",
"libx264",
"-crf",
"23",
"-preset",
"medium",
"-c:a",
"aac",
"-b:a",
"192k",
"-movflags",
"+faststart",
str(output_path),
]
logger.info("[task_id=%s] [渲染] legacy 引擎 FFmpeg 开始: clips=%d", task_id, n)
try:
run_ffmpeg(command)
except subprocess.CalledProcessError as e:
logger.error(
"[task_id=%s] [渲染] legacy 引擎 FFmpeg 失败: %s\nfilter_complex: %s",
task_id,
e,
filter_complex[:500],
)
raise
file_size = output_path.stat().st_size if output_path.exists() else 0
duration = probe_duration(output_path)
return duration, file_size
# ── Celery Task ──────────────────────────────────────────────────────────────
@@ -715,31 +873,59 @@ def generate_video(self, task_id: str) -> dict:
)
_flush_logs(task_id, gen_task)
# 使用 UnifiedRenderService 渲染
logger.info("[task_id=%s] [渲染] FFmpeg 渲染开始", task_id)
# 3. 根据 Feature Flag 选择渲染引擎
user_id = getattr(gen_task, "created_by_user_id", "") if gen_task else ""
engine = _resolve_render_engine(user_id) if user_id else ENGINE_UNIFIED
logger.info("[task_id=%s] [渲染] 引擎选择: %s (user_id=%s)", task_id, engine, user_id)
render_start = time.monotonic()
render_service = UnifiedRenderService(
plan=virtual_plan,
clips=virtual_clips,
asset_path_map=asset_path_map,
work_dir=temp_path,
output_width=OUTPUT_WIDTH,
output_height=OUTPUT_HEIGHT,
output_fps=int(OUTPUT_FPS),
)
render_result = render_service.render()
render_elapsed = time.monotonic() - render_start
logger.info(
"[task_id=%s] [渲染] FFmpeg 渲染完成: 耗时=%.1fs",
task_id,
render_elapsed,
)
render_output_path = temp_path / f"rendered-{task_id}.mp4"
if engine == ENGINE_LEGACY:
# 旧引擎:filter_complex + concat(保持原帧率,无 fps 归一化)
render_duration, render_file_size = _render_with_legacy_engine(
task_id=task_id,
virtual_clips=virtual_clips,
asset_path_map=asset_path_map,
work_dir=temp_path,
output_path=render_output_path,
)
render_elapsed = time.monotonic() - render_start
logger.info(
"[task_id=%s] [渲染] legacy 引擎完成: 耗时=%.1fs, 时长=%.2fs",
task_id,
render_elapsed,
render_duration,
)
else:
# 新引擎:UnifiedRenderService 图层架构
logger.info("[task_id=%s] [渲染] unified 引擎 FFmpeg 渲染开始", task_id)
render_service = UnifiedRenderService(
plan=virtual_plan,
clips=virtual_clips,
asset_path_map=asset_path_map,
work_dir=temp_path,
output_width=OUTPUT_WIDTH,
output_height=OUTPUT_HEIGHT,
output_fps=int(OUTPUT_FPS),
)
render_result = render_service.render()
render_output_path = render_result.output_path
render_duration = render_result.duration
render_file_size = render_result.file_size
render_elapsed = time.monotonic() - render_start
logger.info(
"[task_id=%s] [渲染] unified 引擎完成: 耗时=%.1fs",
task_id,
render_elapsed,
)
if gen_task:
gen_task.append_log(
"渲染",
f"FFmpeg 渲染完成, 耗时={render_elapsed:.1f}s",
f"引擎={engine}, 耗时={render_elapsed:.1f}s",
duration=round(render_elapsed, 2),
engine=engine,
)
_flush_logs(task_id, gen_task)
@@ -747,14 +933,14 @@ def generate_video(self, task_id: str) -> dict:
if audio_path:
final_path = temp_path / f"final-{task_id}.mp4"
try:
_mux_audio_track(render_result.output_path, audio_path, final_path)
_mux_audio_track(render_output_path, audio_path, final_path)
# 混音成功,使用混音后的文件
output_path = final_path
except Exception as mux_err:
logger.warning("[task_id=%s] [混音] 音频混合失败,使用无音频版本: %s", task_id, mux_err)
output_path = render_result.output_path
output_path = render_output_path
else:
output_path = render_result.output_path
output_path = render_output_path
file_size = output_path.stat().st_size
duration = probe_duration(output_path)
@@ -778,14 +964,12 @@ def generate_video(self, task_id: str) -> dict:
verify_url = get_signed_download_url(file_url, expires_seconds=300) or file_url
if not _verify_url_accessible(verify_url):
# 预签名 URL 也访问失败时,退一步用 object_exists 确认上传成功
from video_processing.oss_helpers import oss_bucket, normalize_storage_key
from video_processing.oss_helpers import normalize_storage_key, oss_bucket
bucket = oss_bucket()
key = normalize_storage_key(file_url)
if bucket and bucket.object_exists(key):
logger.info(
"URL 校验失败但 object_exists 确认文件存在,视为上传成功: storage_key=%s", key
)
logger.info("URL 校验失败但 object_exists 确认文件存在,视为上传成功: storage_key=%s", key)
if gen_task:
gen_task.append_log("OSS上传", "URL校验降级: object_exists确认存在", level="WARN")
else:
+1 -1
View File
@@ -4,6 +4,7 @@ import logging
from celery import Task
from celery.exceptions import Retry
from video_processing.oss_helpers import get_signed_download_url
from worker_app.celery_app import celery_app
from worker_app.db import SessionLocal
@@ -16,7 +17,6 @@ from packages.application.cosyvoice_service import (
CosyVoiceTimeoutError,
)
from packages.application.voice_clone.workflow import VoiceCloneWorkflowService
from video_processing.oss_helpers import get_signed_download_url
logger = logging.getLogger(__name__)
+16 -1
View File
@@ -1,3 +1,9 @@
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
FeatureFlagStore,
InMemoryFeatureFlagStore,
RedisFeatureFlagStore,
)
from packages.adapters.redis.session_store import (
NoopSessionStore,
RedisConfig,
@@ -5,4 +11,13 @@ from packages.adapters.redis.session_store import (
get_session_store,
)
__all__ = ["NoopSessionStore", "RedisConfig", "SessionStore", "get_session_store"]
__all__ = [
"FeatureFlagConfig",
"FeatureFlagStore",
"InMemoryFeatureFlagStore",
"NoopSessionStore",
"RedisConfig",
"RedisFeatureFlagStore",
"SessionStore",
"get_session_store",
]
+259
View File
@@ -0,0 +1,259 @@
"""Feature Flag 存储实现。
支持两种后端:
- RedisFeatureFlagStore:生产环境使用,支持多实例共享、热更新
- InMemoryFeatureFlagStore:测试/开发环境使用,纯内存
支持的 Flag 类型:
- 全局开关(enabled: bool
- 白名单(whitelist: Set[str],如 user_id 列表)
- 百分比切流(percentage: 0-100,基于标识符哈希取模)
判定优先级:白名单 > 百分比 > 全局开关
"""
from __future__ import annotations
import hashlib
import json
import logging
import threading
import time
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Optional, Set
logger = logging.getLogger(__name__)
# Redis key 前缀
FEATURE_FLAG_REDIS_PREFIX = "feature_flag:"
@dataclass
class FeatureFlagConfig:
"""单个 Feature Flag 的配置。"""
name: str
enabled: bool = False
percentage: int = 0 # 0-100
whitelist: Set[str] = field(default_factory=set)
def to_dict(self) -> dict:
return {
"name": self.name,
"enabled": self.enabled,
"percentage": self.percentage,
"whitelist": sorted(self.whitelist),
}
@classmethod
def from_dict(cls, data: dict) -> "FeatureFlagConfig":
return cls(
name=data["name"],
enabled=bool(data.get("enabled", False)),
percentage=int(data.get("percentage", 0)),
whitelist=set(data.get("whitelist", [])),
)
def is_active(self, identifier: Optional[str] = None) -> bool:
"""判断当前 flag 是否激活。
判定优先级:
1. 全局关闭 → False
2. 白名单匹配 → True
3. 百分比命中 → True
4. 其他 → False
Args:
identifier: 用于白名单匹配和百分比哈希的标识符(如 user_id)。
传 None 时只看全局开关 + 百分比(百分比用随机值)。
"""
if not self.enabled:
return False
# 白名单:精确匹配
if identifier and identifier in self.whitelist:
return True
# 百分比:0 直接 False100 直接 True
if self.percentage <= 0:
# 没有白名单且百分比为0 → 未启用
return False
if self.percentage >= 100:
return True
# 基于 identifier 做哈希取模,确保同一用户始终落在同一侧
if identifier:
hash_val = int(
hashlib.md5(f"{self.name}:{identifier}".encode("utf-8")).hexdigest(), 16 # nosec B324
) # nosec B324 - 用于哈希取模做百分比切流,非安全用途
return (hash_val % 100) < self.percentage
# 无 identifier 且百分比在 0-100 之间 → 按比例随机(不保证一致性)
import random
return random.randint(0, 99) < self.percentage
class FeatureFlagStore(ABC):
"""Feature Flag 存储抽象接口。"""
@abstractmethod
def get(self, name: str) -> FeatureFlagConfig:
"""获取指定 flag 的配置,不存在则返回默认配置(关闭状态)。"""
...
@abstractmethod
def set(self, config: FeatureFlagConfig) -> None:
"""设置 flag 配置。"""
...
@abstractmethod
def delete(self, name: str) -> bool:
"""删除 flag,返回是否成功删除。"""
...
@abstractmethod
def list_all(self) -> dict[str, FeatureFlagConfig]:
"""列出所有 flag。"""
...
def is_active(self, name: str, identifier: Optional[str] = None) -> bool:
"""便捷方法:判断 flag 是否激活。"""
return self.get(name).is_active(identifier)
class InMemoryFeatureFlagStore(FeatureFlagStore):
"""内存实现,用于测试和本地开发。"""
def __init__(self) -> None:
self._flags: dict[str, FeatureFlagConfig] = {}
self._lock = threading.Lock()
def get(self, name: str) -> FeatureFlagConfig:
with self._lock:
return self._flags.get(name, FeatureFlagConfig(name=name, enabled=False))
def set(self, config: FeatureFlagConfig) -> None:
with self._lock:
self._flags[config.name] = config
def delete(self, name: str) -> bool:
with self._lock:
if name in self._flags:
del self._flags[name]
return True
return False
def list_all(self) -> dict[str, FeatureFlagConfig]:
with self._lock:
return dict(self._flags)
class RedisFeatureFlagStore(FeatureFlagStore):
"""Redis 实现,支持多实例共享配置。
每个 flag 存在一个独立的 Redis hash key 中:
Key: feature_flag:{name}
Fields: enabled, percentage, whitelist(JSON array)
"""
def __init__(self, redis_url: str, key_prefix: str = FEATURE_FLAG_REDIS_PREFIX) -> None:
import redis as redis_lib
self._redis = redis_lib.from_url(redis_url, decode_responses=True)
self._key_prefix = key_prefix
# 本地缓存 + TTL,减少 Redis 调用
self._cache: dict[str, tuple[FeatureFlagConfig, float]] = {}
self._cache_ttl = 5.0 # 秒,默认5秒本地缓存
self._lock = threading.Lock()
def _redis_key(self, name: str) -> str:
return f"{self._key_prefix}{name}"
def _parse_whitelist(self, raw: Optional[str]) -> Set[str]:
if not raw:
return set()
try:
data = json.loads(raw)
return set(data) if isinstance(data, list) else set()
except (json.JSONDecodeError, TypeError):
return set()
def get(self, name: str) -> FeatureFlagConfig:
now = time.time()
# 先查本地缓存
with self._lock:
cached = self._cache.get(name)
if cached and now - cached[1] < self._cache_ttl:
return cached[0]
# 从 Redis 读取
try:
key = self._redis_key(name)
data = self._redis.hgetall(key)
if not data:
config = FeatureFlagConfig(name=name, enabled=False)
else:
config = FeatureFlagConfig(
name=name,
enabled=(data.get("enabled", "0") in ("1", "true", "True")),
percentage=int(data.get("percentage", 0)),
whitelist=self._parse_whitelist(data.get("whitelist")),
)
# 写入本地缓存
with self._lock:
self._cache[name] = (config, now)
return config
except Exception as exc:
logger.warning("Failed to get feature flag %s from Redis: %s", name, exc)
# Redis 不可用时返回默认值(关闭),不影响业务
return FeatureFlagConfig(name=name, enabled=False)
def set(self, config: FeatureFlagConfig) -> None:
key = self._redis_key(config.name)
self._redis.hset(
key,
mapping={
"enabled": "1" if config.enabled else "0",
"percentage": str(config.percentage),
"whitelist": json.dumps(sorted(config.whitelist), ensure_ascii=False),
},
)
# 失效本地缓存
with self._lock:
self._cache.pop(config.name, None)
def delete(self, name: str) -> bool:
key = self._redis_key(name)
result = self._redis.delete(key)
with self._lock:
self._cache.pop(name, None)
return bool(result)
def list_all(self) -> dict[str, FeatureFlagConfig]:
pattern = f"{self._key_prefix}*"
result: dict[str, FeatureFlagConfig] = {}
try:
cursor = 0
while True:
cursor, keys = self._redis.scan(cursor=cursor, match=pattern, count=100)
for key in keys:
name = key[len(self._key_prefix) :]
result[name] = self.get(name)
if cursor == 0:
break
except Exception as exc:
logger.warning("Failed to list feature flags from Redis: %s", exc)
return result
def invalidate_cache(self, name: Optional[str] = None) -> None:
"""手动失效本地缓存。"""
with self._lock:
if name:
self._cache.pop(name, None)
else:
self._cache.clear()
+17
View File
@@ -91,6 +91,23 @@ class SQLAlchemyGenerationTaskRepository:
def count_by_user(self, user_id: str) -> int:
return self.session.query(GenerationTaskModel).filter(GenerationTaskModel.created_by_user_id == user_id).count()
def count_pending_by_user(self, user_id: str) -> int:
return (
self.session.query(GenerationTaskModel)
.filter(
GenerationTaskModel.created_by_user_id == user_id,
GenerationTaskModel.status == GenerationTaskStatus.PENDING.value,
)
.count()
)
def count_pending_total(self) -> int:
return (
self.session.query(GenerationTaskModel)
.filter(GenerationTaskModel.status == GenerationTaskStatus.PENDING.value)
.count()
)
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]:
models = (
self.session.query(GenerationTaskModel)
+23 -48
View File
@@ -121,9 +121,7 @@ class CosyVoiceService:
self._api_key = api_key or settings.cosyvoice_api_key
self._base_url = base_url or settings.cosyvoice_base_url
self._model = model or settings.cosyvoice_model
self._clone_model = clone_model or getattr(
settings, "cosyvoice_clone_model", "voice-enrollment"
)
self._clone_model = clone_model or getattr(settings, "cosyvoice_clone_model", "voice-enrollment")
self._audio_url_signer = audio_url_signer
# base_url 规范化:去掉末尾的路径残留(兼容旧版配置)
@@ -134,11 +132,11 @@ class CosyVoiceService:
# 截取到 /api/v1 为止
idx = self._base_url.find("/api/v1")
if idx >= 0:
self._base_url = self._base_url[:idx + len("/api/v1")]
self._base_url = self._base_url[: idx + len("/api/v1")]
logger.warning(
"[CosyVoice Config] base_url包含旧版text2audio路径,已自动修正: "
"%s -> %s",
old_url, self._base_url,
"[CosyVoice Config] base_url包含旧版text2audio路径,已自动修正: " "%s -> %s",
old_url,
self._base_url,
)
self._client = http_client or httpx.Client(
@@ -236,8 +234,7 @@ class CosyVoiceService:
if self._audio_url_signer:
try:
signed_audio_url = self._audio_url_signer(audio_url)
logger.info("音频URL已预签名: original=%s signed_prefix=%s",
audio_url[:80], signed_audio_url[:80])
logger.info("音频URL已预签名: original=%s signed_prefix=%s", audio_url[:80], signed_audio_url[:80])
except Exception as e:
logger.warning("音频URL预签名失败,使用原始URL: %s", e)
@@ -358,9 +355,7 @@ class CosyVoiceService:
while attempts < self.CLONE_MAX_POLL_ATTEMPTS:
elapsed = time.time() - start_time
if elapsed > timeout:
raise CosyVoiceTimeoutError(
f"音色克隆任务超时({timeout}秒): voice_id={voice_id}"
)
raise CosyVoiceTimeoutError(f"音色克隆任务超时({timeout}秒): voice_id={voice_id}")
result = self.query_voice_status(voice_id)
status = result.get("status", "").upper()
@@ -368,9 +363,7 @@ class CosyVoiceService:
if status == "OK":
return {"voice_id": voice_id}
elif status == "UNDEPLOYED":
raise CosyVoiceError(
f"音色克隆任务失败(审核未通过): voice_id={voice_id}"
)
raise CosyVoiceError(f"音色克隆任务失败(审核未通过): voice_id={voice_id}")
elif status in ("DEPLOYING", "PENDING", "PROCESSING", ""):
# 继续轮询
time.sleep(self.CLONE_POLL_INTERVAL)
@@ -380,9 +373,7 @@ class CosyVoiceService:
time.sleep(self.CLONE_POLL_INTERVAL)
attempts += 1
raise CosyVoiceTimeoutError(
f"音色克隆任务轮询次数超限: voice_id={voice_id}"
)
raise CosyVoiceTimeoutError(f"音色克隆任务轮询次数超限: voice_id={voice_id}")
def clone_voice(
self,
@@ -497,9 +488,7 @@ class CosyVoiceService:
request_id = response.get("request_id", "")
if not audio_url:
raise CosyVoiceError(
f"CosyVoice API 未返回 audio_url: {response}"
)
raise CosyVoiceError(f"CosyVoice API 未返回 audio_url: {response}")
return {
"task_id": "", # 同步接口无 task_id,兼容旧接口
@@ -509,9 +498,7 @@ class CosyVoiceService:
"request_id": request_id,
}
def poll_synthesize_task(
self, task_id: str, timeout: float = 120.0
) -> dict:
def poll_synthesize_task(self, task_id: str, timeout: float = 120.0) -> dict:
"""轮询合成任务(同步接口无需轮询,保留兼容).
CosyVoice SpeechSynthesizer 非流式接口是同步的,
@@ -520,10 +507,7 @@ class CosyVoiceService:
Raises:
CosyVoiceError: 同步接口无需轮询
"""
raise CosyVoiceError(
"CosyVoice 非流式合成接口是同步的,无需轮询. "
"请直接使用 submit_synthesize_task()."
)
raise CosyVoiceError("CosyVoice 非流式合成接口是同步的,无需轮询. " "请直接使用 submit_synthesize_task().")
def synthesize_speech(
self,
@@ -626,15 +610,17 @@ class CosyVoiceService:
# DEBUG: 打印完整请求信息,用于排查418错误
import json as json_lib
safe_headers = {k: v for k, v in headers.items()}
if "Authorization" in safe_headers:
token = safe_headers["Authorization"]
if len(token) > 20:
safe_headers["Authorization"] = token[:13] + "..." + token[-4:]
logger.info(
"[CosyVoice Debug] 请求详情: "
"method=%s, url=%s, headers=%s, body=%s",
method, url, safe_headers,
"[CosyVoice Debug] 请求详情: " "method=%s, url=%s, headers=%s, body=%s",
method,
url,
safe_headers,
json_lib.dumps(json, ensure_ascii=False) if json else "None",
)
@@ -652,8 +638,7 @@ class CosyVoiceService:
# DEBUG: 打印响应状态和完整响应体
logger.info(
"[CosyVoice Debug] 响应详情: "
"status=%d, body=%s",
"[CosyVoice Debug] 响应详情: " "status=%d, body=%s",
response.status_code,
response.text[:2000], # 最多2000字符,避免日志过大
)
@@ -662,9 +647,7 @@ class CosyVoiceService:
if response.status_code == 200:
return response.json()
elif response.status_code in (401, 403):
raise CosyVoiceAuthError(
f"CosyVoice API 认证失败: HTTP {response.status_code}"
)
raise CosyVoiceAuthError(f"CosyVoice API 认证失败: HTTP {response.status_code}")
elif response.status_code == 400:
# 客户端错误,不重试
body_text = response.text
@@ -672,19 +655,12 @@ class CosyVoiceService:
body = response.json()
code = body.get("code", "")
message = body.get("message", "")
raise CosyVoiceError(
f"CosyVoice API 参数错误: HTTP 400, "
f"code={code}, message={message}"
)
raise CosyVoiceError(f"CosyVoice API 参数错误: HTTP 400, " f"code={code}, message={message}")
except ValueError:
raise CosyVoiceError(
f"CosyVoice API 调用失败: HTTP 400, body={body_text}"
)
raise CosyVoiceError(f"CosyVoice API 调用失败: HTTP 400, body={body_text}")
elif response.status_code >= 500:
# 服务端错误,可重试
last_error = CosyVoiceError(
f"CosyVoice API 服务端错误: HTTP {response.status_code}"
)
last_error = CosyVoiceError(f"CosyVoice API 服务端错误: HTTP {response.status_code}")
logger.warning(
"CosyVoice API 失败 (尝试 %d/%d): HTTP %d",
attempt + 1,
@@ -694,8 +670,7 @@ class CosyVoiceService:
else:
# 其他客户端错误,不重试
raise CosyVoiceError(
f"CosyVoice API 调用失败: HTTP {response.status_code}, "
f"body={response.text}"
f"CosyVoice API 调用失败: HTTP {response.status_code}, " f"body={response.text}"
)
except httpx.TimeoutException as e:
+7 -24
View File
@@ -219,9 +219,7 @@ class TTSWorkflowService:
# 新接口(同步):没有 task_id,重新合成
if not task_id:
logger.info(
f"TTS 任务无 task_id,重新同步合成: job_id={job_id}"
)
logger.info(f"TTS 任务无 task_id,重新同步合成: job_id={job_id}")
return self._resynthesize_and_complete(job)
# 旧接口遗留的 task_id,尝试轮询(兼容过渡)
@@ -235,9 +233,7 @@ class TTSWorkflowService:
)
except CosyVoiceError:
# 旧接口轮询失败,重新同步合成
logger.warning(
f"旧 task_id 轮询失败,重新同步合成: job_id={job_id}, task_id={task_id}"
)
logger.warning(f"旧 task_id 轮询失败,重新同步合成: job_id={job_id}, task_id={task_id}")
return self._resynthesize_and_complete(job)
def process_synthesis_result(
@@ -524,10 +520,7 @@ class TTSWorkflowService:
missing_indices = [i for i in range(segment_count) if results[i] is None]
if missing_indices:
logger.info(
f"分段任务重新合成缺失段: job_id={job.id}, "
f"缺失={len(missing_indices)}/{segment_count}"
)
logger.info(f"分段任务重新合成缺失段: job_id={job.id}, " f"缺失={len(missing_indices)}/{segment_count}")
# 并发重新合成缺失分段
max_workers = min(len(missing_indices), _MAX_SEGMENT_WORKERS)
with ThreadPoolExecutor(max_workers=max_workers) as executor:
@@ -550,13 +543,8 @@ class TTSWorkflowService:
try:
results[idx] = future.result()
except Exception as e:
logger.error(
f"分段重新合成失败: job_id={job.id}, "
f"segment={idx}, error={e}"
)
self._handle_segment_failure(
job, f"分段 {idx + 1} 重新合成失败: {e}"
)
logger.error(f"分段重新合成失败: job_id={job.id}, " f"segment={idx}, error={e}")
self._handle_segment_failure(job, f"分段 {idx + 1} 重新合成失败: {e}")
return self.repository.get(job.id)
# 所有分段完成,下载合并
@@ -564,9 +552,7 @@ class TTSWorkflowService:
try:
merged_data, total_duration = self._download_and_merge_segments(results, job)
permanent_url, storage_key = self._upload_merged_to_oss(
merged_data, job.user_id, job.id, job.format
)
permanent_url, storage_key = self._upload_merged_to_oss(merged_data, job.user_id, job.id, job.format)
job.mark_completed(
output_audio_url=permanent_url,
@@ -575,10 +561,7 @@ class TTSWorkflowService:
file_size=len(merged_data),
)
job = self.repository.update(job)
logger.info(
f"分段合成完成(重新合成路径): job_id={job.id}, "
f"merged_size={len(merged_data)}"
)
logger.info(f"分段合成完成(重新合成路径): job_id={job.id}, " f"merged_size={len(merged_data)}")
return job
except Exception as e:
+3 -1
View File
@@ -133,7 +133,9 @@ class VoiceCloneWorkflowService:
profile.metadata = task_metadata
profile = self.repository.update(profile)
logger.info(f"音色克隆任务已提交: profile_id={profile.id}, " f"voice_id={submit_result.get('voice_id')}")
logger.info(
f"音色克隆任务已提交: profile_id={profile.id}, " f"voice_id={submit_result.get('voice_id')}"
)
except (CosyVoiceError, CosyVoiceAuthError) as e:
# CosyVoice 提交失败,标记为 failed
Regular → Executable
+36
View File
@@ -134,6 +134,25 @@ class AssetStatus(StrEnum):
PROCESSING = "processing"
ERROR = "error"
@classmethod
def _missing_(cls, value: object) -> "AssetStatus":
"""兼容历史数据,避免枚举转换失败导致500。
- uploaded → READY(早期版本用 uploaded 表示上传完成)
- 其他未知值 → READY(兜底,不阻塞业务)
"""
if isinstance(value, str):
normalized = value.strip().lower()
if normalized in ("uploaded", "success", "ok", "done", "complete"):
return cls.READY
if normalized in ("upload", "uploading_start", "upload_start"):
return cls.UPLOADING
if normalized in ("failed", "fail", "err"):
return cls.ERROR
if normalized in ("process", "processing", "running", "run"):
return cls.PROCESSING
return cls.READY
class ClassificationStatus(StrEnum):
PENDING = "pending"
@@ -141,6 +160,23 @@ class ClassificationStatus(StrEnum):
COMPLETED = "completed"
FAILED = "failed"
@classmethod
def _missing_(cls, value: object) -> "ClassificationStatus":
"""兼容历史数据,避免枚举转换失败导致500。
- done → COMPLETED(早期版本用 done 表示完成)
- 其他未知值 → PENDING(兜底,不阻塞业务)
"""
if isinstance(value, str):
normalized = value.strip().lower()
if normalized in ("done", "success", "finished", "complete"):
return cls.COMPLETED
if normalized in ("fail", "error", "err"):
return cls.FAILED
if normalized in ("process", "processing", "running", "run"):
return cls.PROCESSING
return cls.PENDING
@dataclass(slots=True)
class Asset:
+4
View File
@@ -16,6 +16,10 @@ class GenerationTaskRepository(Protocol):
def count_by_user(self, user_id: str) -> int: ...
def count_pending_by_user(self, user_id: str) -> int: ...
def count_pending_total(self) -> int: ...
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]: ...
def list_by_source_edit_plan(self, plan_id: str) -> list[GenerationTask]: ...
Regular → Executable
+63
View File
@@ -1,7 +1,70 @@
[tool.black]
line-length = 120
target-version = ["py312"]
extend-exclude = '''
(
\.git
| \.cache
| \.pytest_cache
| \.mypy_cache
| __pycache__
| node_modules
| \.venv
| venv
| build
| dist
| \.next
| out
| coverage
)
'''
[tool.isort]
profile = "black"
line_length = 120
extend_skip_glob = [
".git/**",
".cache/**",
".pytest_cache/**",
".mypy_cache/**",
"__pycache__/**",
"node_modules/**",
".venv/**",
"venv/**",
"build/**",
"dist/**",
".next/**",
"out/**",
"coverage/**",
]
[tool.coverage.run]
source = ["apps/api/app", "packages"]
omit = [
"*/migrations/*",
"*/tests/*",
"*/test_*.py",
"*/site-packages/*",
]
branch = true
[tool.coverage.report]
exclude_lines = [
"pragma: no cover",
"def __repr__",
"if __name__ == .__main__.:",
"raise NotImplementedError",
"pass",
"if TYPE_CHECKING:",
"class .*Protocol",
"@abstractmethod",
"raise AssertionError",
"raise RuntimeError",
"if 0:",
"if __debug__:",
]
show_missing = true
skip_covered = false
[tool.coverage.xml]
output = "coverage.xml"
+61 -3
View File
@@ -30,7 +30,8 @@ fi
# ---- Registry 配置 ----
REGISTRY="${REGISTRY:-git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas}"
CACHE_REGISTRY="${CACHE_REGISTRY:-git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas}"
CACHE_TAG="${CACHE_TAG:-release}"
# 主缓存 tag:develop 分支构建时写入,所有分支读取
CACHE_TAG_PRIMARY="${CACHE_TAG:-develop}"
API_IMAGE="xiaoxia-saas-api:$VERSION"
WORKER_IMAGE="xiaoxia-saas-worker:$VERSION"
@@ -45,6 +46,7 @@ REGISTRY_WEB="${REGISTRY}/xiaoxia-saas-web:$VERSION"
USE_CACHE=0
USE_PUSH=0
CACHE_WRITE=0
# 检查 buildx 和 Registry 认证
if docker buildx version >/dev/null 2>&1; then
@@ -54,7 +56,10 @@ if docker buildx version >/dev/null 2>&1; then
docker buildx use default 2>/dev/null || true
fi
echo "=== Building API image ==="
# ---- 缓存读写策略(按分支隔离)----
# 默认只读不写,防止 feature 分支污染主缓存
# 只有 develop/main 分支才写回缓存
BRANCH_NAME="${GITHUB_REF_NAME:-${CI_COMMIT_BRANCH:-unknown}}"
if [ "$USE_CACHE" -eq 1 ]; then
docker buildx build \
--build-arg APP_VERSION="$VERSION" \
@@ -68,6 +73,52 @@ else
docker build --pull=false --build-arg APP_VERSION="$VERSION" -f infra/docker/api.Dockerfile -t "$API_IMAGE" -t "$API_LATEST" .
fi
build_with_cache() {
# usage: build_with_cache <image_name> <dockerfile> <extra_args...>
IMG_NAME="$1"
DOCKERFILE="$2"
shift 2
EXTRA_ARGS="$*"
CACHE_FROM="type=registry,ref=${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY},ignore-error=true"
if [ "$CACHE_WRITE" -eq 1 ]; then
CACHE_TO="type=registry,ref=${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY},mode=max"
echo " cache: read+write from ${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY}"
else
CACHE_TO=""
echo " cache: read-only from ${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY}"
fi
if [ "$USE_CACHE" -eq 1 ]; then
if [ -n "$CACHE_TO" ]; then
docker buildx build \
$EXTRA_ARGS \
--cache-from "$CACHE_FROM" \
--cache-to "$CACHE_TO" \
-f "$DOCKERFILE" \
-t "$IMG_NAME:$VERSION" \
--load \
.
else
docker buildx build \
$EXTRA_ARGS \
--cache-from "$CACHE_FROM" \
-f "$DOCKERFILE" \
-t "$IMG_NAME:$VERSION" \
--load \
.
fi
else
docker build --pull=false $EXTRA_ARGS -f "$DOCKERFILE" -t "$IMG_NAME:$VERSION" .
fi
}
echo "=== Building API image ==="
build_with_cache "api" "infra/docker/api.Dockerfile" \
"--build-arg APP_VERSION=$VERSION"
docker tag "$API_IMAGE" "$API_LATEST"
echo "=== Building Worker image ==="
if [ "$USE_CACHE" -eq 1 ]; then
docker buildx build \
@@ -83,9 +134,16 @@ else
fi
echo "=== Building Web image (with buildx cache) ==="
# 先构建前端产物
# 先构建前端产物(使用持久化 npm 缓存卷)
NPM_CACHE_VOLUME="xiaoxia-npm-cache"
if ! docker volume inspect "$NPM_CACHE_VOLUME" >/dev/null 2>&1; then
docker volume create "$NPM_CACHE_VOLUME" >/dev/null
echo " Created npm cache volume: $NPM_CACHE_VOLUME"
fi
docker run --rm \
-v "$PWD:/workspace" \
-v "$NPM_CACHE_VOLUME:/workspace/apps/web/node_modules" \
-w /workspace/apps/web \
docker.m.daocloud.io/library/node:20 \
sh -lc "npm ci && npm run build"
+34
View File
@@ -0,0 +1,34 @@
#!/usr/bin/env python3
"""解析 coverage.xml 并输出覆盖率汇总。"""
import os
import sys
import xml.etree.ElementTree as ET
THRESHOLD = int(os.environ.get("COVERAGE_THRESHOLD", 65)) # 行覆盖率门槛,百分比,可通过环境变量覆盖
def main() -> int:
try:
tree = ET.parse("coverage.xml")
except FileNotFoundError:
print("coverage.xml 不存在,跳过汇总")
return 0
root = tree.getroot()
line_rate = float(root.get("line-rate", 0)) * 100
branch_rate = float(root.get("branch-rate", 0)) * 100
lines_covered = int(root.get("lines-covered", 0))
lines_valid = int(root.get("lines-valid", 0))
print(f"行覆盖率: {line_rate:.2f}% ({lines_covered}/{lines_valid})")
print(f"分支覆盖率: {branch_rate:.2f}%")
print(f"门槛: {THRESHOLD}%")
status = "PASS ✅" if line_rate >= THRESHOLD else "FAIL ❌"
print(f"状态: {status}")
return 0 if line_rate >= THRESHOLD else 1
if __name__ == "__main__":
sys.exit(main())
+83
View File
@@ -0,0 +1,83 @@
#!/usr/bin/env python3
"""发送 CI 失败通知到飞书/项目群 webhook。"""
import json
import os
import sys
import urllib.request
def main() -> int:
webhook = os.environ.get("CI_NOTIFY_WEBHOOK", "")
if not webhook:
print("未配置 CI_NOTIFY_WEBHOOK,跳过通知")
print("如需启用,请在仓库 Settings -> Secrets and variables -> Actions 中添加 CI_NOTIFY_WEBHOOK")
return 0
failed_job = os.environ.get("FAILED_JOB", "Unknown Job")
branch = os.environ.get("GITHUB_REF_NAME", "unknown")
commit = os.environ.get("GITHUB_SHA", "unknown")[:8]
actor = os.environ.get("GITHUB_ACTOR", "unknown")
run_id = os.environ.get("GITHUB_RUN_ID", "unknown")
repo = os.environ.get("GITHUB_REPOSITORY", "unknown")
run_url = f"https://git.xiaoxiajianji.com/{repo}/actions/runs/{run_id}"
payload = {
"msg_type": "interactive",
"card": {
"header": {
"title": {
"tag": "plain_text",
"content": "❌ CI 构建失败",
},
"status": "red",
},
"elements": [
{
"tag": "div",
"text": {
"tag": "lark_md",
"content": (
f"**任务**: {failed_job}\n"
f"**分支**: {branch}\n"
f"**提交**: {commit}\n"
f"**提交者**: {actor}\n"
f"**Run ID**: {run_id}"
),
},
},
{
"tag": "action",
"actions": [
{
"tag": "button",
"text": {"tag": "plain_text", "content": "查看失败日志"},
"url": run_url,
"type": "danger",
}
],
},
],
},
}
data = json.dumps(payload).encode("utf-8")
req = urllib.request.Request(
webhook,
data=data,
headers={"Content-Type": "application/json"},
method="POST",
)
try:
with urllib.request.urlopen(req, timeout=10) as resp:
resp.read()
print("通知已发送")
except Exception as e:
print(f"通知发送失败: {e}", file=sys.stderr)
return 1
return 0
if __name__ == "__main__":
sys.exit(main())
+83
View File
@@ -0,0 +1,83 @@
#!/usr/bin/env python3
"""发送 CI 成功通知到飞书/项目群 webhook。"""
import json
import os
import sys
import urllib.request
def main() -> int:
webhook = os.environ.get("CI_NOTIFY_WEBHOOK", "")
if not webhook:
print("未配置 CI_NOTIFY_WEBHOOK,跳过成功通知")
print("如需启用,请在仓库 Settings -> Secrets and variables -> Actions 中添加 CI_NOTIFY_WEBHOOK")
return 0
success_job = os.environ.get("SUCCESS_JOB", "Unknown Job")
branch = os.environ.get("GITHUB_REF_NAME", "unknown")
commit = os.environ.get("GITHUB_SHA", "unknown")[:8]
actor = os.environ.get("GITHUB_ACTOR", "unknown")
run_id = os.environ.get("GITHUB_RUN_ID", "unknown")
repo = os.environ.get("GITHUB_REPOSITORY", "unknown")
run_url = f"https://git.xiaoxiajianji.com/{repo}/actions/runs/{run_id}"
payload = {
"msg_type": "interactive",
"card": {
"header": {
"title": {
"tag": "plain_text",
"content": "✅ CI 构建成功",
},
"status": "green",
},
"elements": [
{
"tag": "div",
"text": {
"tag": "lark_md",
"content": (
f"**任务**: {success_job}\n"
f"**分支**: {branch}\n"
f"**提交**: {commit}\n"
f"**提交者**: {actor}\n"
f"**Run ID**: {run_id}"
),
},
},
{
"tag": "action",
"actions": [
{
"tag": "button",
"text": {"tag": "plain_text", "content": "查看构建详情"},
"url": run_url,
"type": "primary",
}
],
},
],
},
}
data = json.dumps(payload).encode("utf-8")
req = urllib.request.Request(
webhook,
data=data,
headers={"Content-Type": "application/json"},
method="POST",
)
try:
with urllib.request.urlopen(req, timeout=10) as resp:
resp.read()
print("成功通知已发送")
except Exception as e:
print(f"成功通知发送失败: {e}", file=sys.stderr)
return 1
return 0
if __name__ == "__main__":
sys.exit(main())
+1
View File
@@ -3,6 +3,7 @@ max-line-length = 120
extend-ignore = E203,W503,E501,E302,E402,E722,W291,W293,F401,F403,F405,F841
exclude =
.git,
.cache,
__pycache__,
.venv,
.venv-ci-root,
+22 -10
View File
@@ -118,6 +118,18 @@ class StubGenerationTaskRepository:
def count_by_user(self, user_id: str) -> int:
return len([t for t in self._tasks.values() if t.created_by_user_id == user_id])
def count_pending_by_user(self, user_id: str) -> int:
return len(
[
t
for t in self._tasks.values()
if t.created_by_user_id == user_id and t.status == GenerationTaskStatus.PENDING
]
)
def count_pending_total(self) -> int:
return len([t for t in self._tasks.values() if t.status == GenerationTaskStatus.PENDING])
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]:
items = [t for t in self._tasks.values() if t.created_by_user_id == user_id]
items.sort(key=lambda t: t.created_at, reverse=True)
@@ -241,7 +253,7 @@ def client():
class TestCreateGenerationTask:
"""创建生成任务端点测试。"""
@patch("app.api.routes.generation_tasks.celery_app")
@patch("app.core.task_enqueue.celery_app")
def test_create_task_success(self, mock_celery, client):
"""正常创建生成任务成功。"""
mock_celery.send_task = MagicMock()
@@ -270,7 +282,7 @@ class TestCreateGenerationTask:
assert mock_celery.send_task.called
assert mock_celery.send_task.call_args[0][0] == "worker.generate_video"
@patch("app.api.routes.generation_tasks.celery_app")
@patch("app.core.task_enqueue.celery_app")
def test_create_batch_tasks(self, mock_celery, client):
"""批量创建多个生成任务。"""
mock_celery.send_task = MagicMock()
@@ -347,7 +359,7 @@ class TestListGenerationTasks:
def _create_task(self, client, task_suffix: str = "1"):
"""辅助方法:创建一个生成任务。"""
with patch("app.api.routes.generation_tasks.celery_app") as mock_celery:
with patch("app.core.task_enqueue.celery_app") as mock_celery:
mock_celery.send_task = MagicMock()
resp = client.post(
"/api/v1/generation/tasks",
@@ -368,7 +380,7 @@ class TestListGenerationTasks:
assert "items" in data
assert data["items"] == []
@patch("app.api.routes.generation_tasks.celery_app")
@patch("app.core.task_enqueue.celery_app")
def test_list_returns_user_tasks(self, mock_celery, client):
"""返回当前用户的生成任务列表。"""
mock_celery.send_task = MagicMock()
@@ -406,7 +418,7 @@ class TestGetGenerationTask:
"""获取生成任务详情端点测试。"""
def _create_task(self, client) -> str:
with patch("app.api.routes.generation_tasks.celery_app") as mock_celery:
with patch("app.core.task_enqueue.celery_app") as mock_celery:
mock_celery.send_task = MagicMock()
resp = client.post(
"/api/v1/generation/tasks",
@@ -449,7 +461,7 @@ class TestListGenerationResults:
"""列出生成结果端点测试。"""
def _create_task(self, client) -> str:
with patch("app.api.routes.generation_tasks.celery_app") as mock_celery:
with patch("app.core.task_enqueue.celery_app") as mock_celery:
mock_celery.send_task = MagicMock()
resp = client.post(
"/api/v1/generation/tasks",
@@ -489,7 +501,7 @@ class TestRetryGenerationTask:
def _create_failed_task(self, client) -> str:
"""创建一个失败状态的任务。"""
with patch("app.api.routes.generation_tasks.celery_app") as mock_celery:
with patch("app.core.task_enqueue.celery_app") as mock_celery:
mock_celery.send_task = MagicMock()
resp = client.post(
"/api/v1/generation/tasks",
@@ -509,7 +521,7 @@ class TestRetryGenerationTask:
# 让我们直接通过 retry 测试来验证
return task_id
@patch("app.api.routes.generation_tasks.celery_app")
@patch("app.core.task_enqueue.celery_app")
def test_retry_failed_task(self, mock_celery, client):
"""重试失败的任务成功。"""
mock_celery.send_task = MagicMock()
@@ -539,7 +551,7 @@ class TestRetryGenerationTask:
assert resp.status_code == 404
assert "not found" in resp.json()["detail"].lower()
@patch("app.api.routes.generation_tasks.celery_app")
@patch("app.core.task_enqueue.celery_app")
def test_retry_completed_task_returns_409(self, mock_celery, client):
"""重试已完成的任务返回 409。"""
mock_celery.send_task = MagicMock()
@@ -568,7 +580,7 @@ class TestRetryGenerationTask:
class TestGenerationTaskFlow:
"""生成任务完整流程集成测试。"""
@patch("app.api.routes.generation_tasks.celery_app")
@patch("app.core.task_enqueue.celery_app")
def test_create_list_detail_results_flow(self, mock_celery, client):
"""测试创建 → 列表 → 详情 → 结果 完整流程。"""
mock_celery.send_task = MagicMock()
+14 -2
View File
@@ -83,6 +83,18 @@ class StubGenerationTaskRepository:
def count_by_user(self, user_id: str) -> int:
return len([t for t in self._tasks.values() if t.created_by_user_id == user_id])
def count_pending_by_user(self, user_id: str) -> int:
return len(
[
t
for t in self._tasks.values()
if t.created_by_user_id == user_id and t.status == GenerationTaskStatus.PENDING
]
)
def count_pending_total(self) -> int:
return len([t for t in self._tasks.values() if t.status == GenerationTaskStatus.PENDING])
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[GenerationTask]:
items = [t for t in self._tasks.values() if t.created_by_user_id == user_id]
items.sort(key=lambda t: t.created_at, reverse=True)
@@ -428,7 +440,7 @@ class TestRetryProjectTask:
assert resp.status_code == 400
assert "Unsupported" in resp.json()["detail"]
@patch("app.api.routes.task_center.celery_app")
@patch("app.core.task_enqueue.celery_app")
def test_retry_failed_generation_task(self, mock_celery, client):
"""重试失败的 generation 任务成功。"""
mock_celery.send_task = MagicMock()
@@ -581,7 +593,7 @@ class TestRetryProjectTask:
class TestTaskCenterCrossEndpoint:
"""任务中心跨端点集成测试。"""
@patch("app.api.routes.task_center.celery_app")
@patch("app.core.task_enqueue.celery_app")
def test_list_then_retry_then_list(self, mock_celery, client):
"""列出任务 → 重试失败任务 → 再列出验证新任务。"""
mock_celery.send_task = MagicMock()
+18 -17
View File
@@ -32,7 +32,11 @@ sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "
from app.api.routes.voice_clones import router
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_cosyvoice_service, get_voice_clone_profile_repository
from app.dependencies import (
get_audio_url_signer,
get_cosyvoice_service,
get_voice_clone_profile_repository,
)
from packages.domain.entities import User
from packages.domain.voice_clone_profile import (
@@ -226,7 +230,7 @@ def clone_repo():
@pytest.fixture
def cosyvoice_service():
return MockCosyVoiceService(async_mode=False) # 步模式,简化测试
return MockCosyVoiceService(async_mode=True) # 步模式,匹配真实 CosyVoice API 行为
@pytest.fixture
@@ -241,6 +245,7 @@ def client(clone_repo, cosyvoice_service):
test_app.dependency_overrides[get_current_user] = _override_current_user
test_app.dependency_overrides[get_voice_clone_profile_repository] = lambda: clone_repo
test_app.dependency_overrides[get_cosyvoice_service] = lambda: cosyvoice_service
test_app.dependency_overrides[get_audio_url_signer] = lambda: (lambda url: url)
yield TestClient(test_app)
@@ -256,7 +261,7 @@ class TestCreateVoiceClone:
"""创建声音克隆端点测试。"""
def test_create_with_source_audio(self, client, cosyvoice_service):
"""提供源音频时创建克隆,同步模式下直接 ready"""
"""提供源音频时创建克隆,异步提交后状态为 processing"""
resp = client.post(
"/voice-clones",
json={
@@ -277,9 +282,9 @@ class TestCreateVoiceClone:
assert "id" in data
assert len(data["id"]) > 0
# 步模式下应直接 ready
assert data["status"] == "ready"
assert data["voice_id"] == "mock-voice-789"
# 步模式下提交后状态为 processingvoice_id 为空
assert data["status"] == "processing"
assert data["voice_id"] == ""
assert data["error_message"] == ""
def test_create_without_source_audio(self, client):
@@ -554,16 +559,15 @@ class TestRetryVoiceClone:
"""重试克隆端点测试。"""
def test_retry_failed_clone(self, client, clone_repo, cosyvoice_service):
"""重试失败的克隆应成功"""
cosyvoice_service.async_mode = False
"""重试失败的克隆,重新提交后期望 processing"""
p = _make_clone_profile("重试测试", status=VoiceCloneStatus.FAILED)
clone_repo.create(p)
resp = client.post(f"/voice-clones/{p.id}/retry")
assert resp.status_code == 200
data = resp.json()
# 步模式下重试后应变为 ready
assert data["status"] == "ready"
# 步模式下重试后状态为 processing,等待 CosyVoice 完成
assert data["status"] == "processing"
assert data["retry_count"] >= 1
def test_retry_nonexistent_returns_404(self, client):
@@ -590,7 +594,6 @@ class TestRetryVoiceClone:
def test_retry_increments_retry_count(self, client, clone_repo, cosyvoice_service):
"""重试后重试次数增加。"""
cosyvoice_service.async_mode = False
p = _make_clone_profile("重试计数", status=VoiceCloneStatus.FAILED)
clone_repo.create(p)
@@ -680,7 +683,7 @@ class TestVoiceCloneLifecycle:
# 4. 状态
status_resp = client.get(f"/voice-clones/{clone_id}/status")
assert status_resp.status_code == 200
assert status_resp.json()["status"] == "ready"
assert status_resp.json()["status"] == "processing"
# 5. 删除
del_resp = client.delete(f"/voice-clones/{clone_id}")
@@ -691,7 +694,7 @@ class TestVoiceCloneLifecycle:
assert list_resp2.json()["total"] == 0
def test_failed_retry_flow(self, client, clone_repo, cosyvoice_service):
"""失败 → 重试 → 成功 流程。"""
"""失败 → 重试 → processing(等待异步完成) 流程。"""
# 创建一个失败的克隆
p = _make_clone_profile("失败重试", status=VoiceCloneStatus.FAILED)
clone_repo.create(p)
@@ -701,15 +704,13 @@ class TestVoiceCloneLifecycle:
assert status_resp.json()["status"] == "failed"
# 重试
cosyvoice_service.async_mode = False
retry_resp = client.post(f"/voice-clones/{p.id}/retry")
assert retry_resp.status_code == 200
assert retry_resp.json()["status"] == "ready"
assert retry_resp.json()["status"] == "processing"
# 再次确认状态
status_resp2 = client.get(f"/voice-clones/{p.id}/status")
assert status_resp2.json()["status"] == "ready"
assert status_resp2.json()["voice_id"] != ""
assert status_resp2.json()["status"] == "processing"
if __name__ == "__main__":
+136
View File
@@ -0,0 +1,136 @@
# 灰度对比测试工具
用于统一渲染引擎灰度发布期间的新旧引擎对比验证。
## 能力
- **像素对比**:基于 FFmpeg SSIM + PSNR 双指标,评估视频画质差异
- **音频对比**:基于差值音频 RMS,评估音频波形差异
- **批量对比**10个预设场景覆盖 P0/P1/P2 优先级
- **HTML 报告**:可视化对比结果,包含画质、音频、性能三维度
- **两种切换方式**:支持 engine 参数直传 或 Feature Flag 白名单切换
## 目录结构
```
tests/render_compare/
├── __init__.py # 包导出
├── README.md # 本文档
├── video_diff.py # 视频像素对比(SSIM + PSNR
├── audio_diff.py # 音频对比(差值 RMS)
├── scenarios.py # 预定义对比场景(10个)
└── runner.py # 批量对比执行器 + HTML 报告生成
```
## 快速开始
### 环境要求
- FFmpeg 4.4+(需带 ssim 和 psnr 滤镜)
- Python 3.10+
- httpxAPI 调用)
### 配置环境变量
```bash
export STAGING_API_URL=https://api.staging.example.com
export STAGING_API_KEY=your_api_key
export STAGING_INTERNAL_API_KEY=your_internal_key # 可选,Feature Flag 模式需要
```
### 运行对比
```bash
# 运行所有 P0 场景(最核心的5个)
python -m tests.render_compare.runner --priority P0 --output ./report/
# 运行 P0 + P1 场景
python -m tests.render_compare.runner --priority P1 --output ./report/
# 只跑指定场景
python -m tests.render_compare.runner --scenarios simple_pass_through,subtitle_rendering
# 使用 Feature Flag 方式切换引擎(需要 internal key
python -m tests.render_compare.runner --priority P0 --flag-mode
# 自定义阈值
python -m tests.render_compare.runner --priority P0 --ssim-threshold 0.95 --psnr-threshold 30
```
## 对比场景
| ID | 名称 | 优先级 | 验证点 |
|----|------|--------|--------|
| simple_pass_through | 简单直通 | P0 | 直通优化路径正确性 |
| multi_clip_transition | 多clip转场 | P0 | 转场效果 + concat |
| subtitle_rendering | 字幕渲染 | P0 | ASS字幕渲染 |
| independent_audio_track | 独立音频轨 | P0 | 音频混音(amix |
| no_audio_video | 无音轨视频 | P0 | 无音轨防御逻辑 |
| picture_in_picture | 画中画 | P1 | overlay 图层 |
| multi_layer_mix | 多图层混合 | P1 | 多图层复杂场景 |
| image_background | 图片背景 | P1 | background 层 + 无音频 |
| long_video_stress | 长视频压力 | P2 | 多clip性能 |
| vertical_portrait | 竖屏9:16 | P2 | scale 策略(铺满裁剪) |
## 验收标准(建议)
### 视频质量
- **平均 SSIM >= 0.90**:通过(有微小差异但视觉可接受)
- **平均 SSIM >= 0.95**:优秀(视觉几乎无差异)
- **平均 PSNR >= 25 dB**:通过
- **分辨率一致 + 时长差 < 0.1s**:通过
### 音频质量
- **相似度 >= 0.85**:通过
- **采样率/声道数一致**:通过
### 性能
- **平均性能差异在 ±10% 以内**:可接受
- **直通场景新引擎更快**(预期 +30%)
## API 约定
Runner 默认假设渲染 API 支持以下接口:
### 提交任务
```
POST /api/v1/render/compose
Authorization: Bearer {api_key}
Body: { ...plan_payload, "engine": "legacy" | "unified" }
Response: { "task_id": "xxx" }
```
### 查询状态
```
GET /api/v1/tasks/{task_id}
Response: { "status": "completed", "output_url": "...", "duration_sec": 5.2 }
```
### Feature Flagflag-mode
```
PUT /api/v1/internal/feature-flags/render_engine
X-API-Key: {internal_key}
Body: { "enabled": true, "percentage": 100 }
```
如果你的 API 接口不同,请修改 `StagingAPI` 类中的对应方法。
## 故障排查
### 对比失败定位指南
1. **像素差异大(SSIM < 0.90**
- 检查分辨率是否一致
- 检查帧率是否一致
-`save_diff_frame` 生成差异帧可视化
- 检查转场效果(slideup/slidedown 是新引擎独有)
2. **音频不一致**
- 检查音频编码参数(码率、采样率)
- 检查主音频源优先级(main > broll
- 用 ffprobe 对比两视频音频流参数
3. **渲染失败**
- 检查日志:`[unified-render] render failed`
- 检查素材是否完整下载
- 检查 FFmpeg 命令是否正确
+26
View File
@@ -0,0 +1,26 @@
"""灰度对比测试工具包.
用于新旧渲染引擎的批量对比测试,包含:
- video_diff: 视频像素对比(SSIM + PSNR
- audio_diff: 音频对比(差值 RMS
- scenarios: 预定义对比场景
- runner: 批量对比执行器 + HTML 报告
"""
from .audio_diff import AudioDiffResult, compute_audio_diff, extract_audio, probe_duration, probe_has_audio
from .scenarios import SCENARIOS, CompareScenario, get_scenarios_by_priority
from .video_diff import VideoDiffResult, compute_video_diff, save_diff_frame
__all__ = [
"VideoDiffResult",
"compute_video_diff",
"save_diff_frame",
"AudioDiffResult",
"compute_audio_diff",
"extract_audio",
"probe_has_audio",
"probe_duration",
"SCENARIOS",
"CompareScenario",
"get_scenarios_by_priority",
]
+322
View File
@@ -0,0 +1,322 @@
"""音频对比工具 — 基于 FFmpeg 的音频质量对比.
使用以下指标评估两段音频的相似度:
1. 波形差异(RMS 差值)
2. 频谱相似度(FFT 分帧比较)
3. 时长差异
对比方式:
- 直接对两个音频做 `ametadata=select='gt(scene\\,0.3)'` 过于复杂
- 简化方案:用 `amerge` + `astats` 计算差值音频的 RMS
更精确的方案(已实现):
- 将两轨音频做差(amix=0:weights='1 -1' → 实际上用 pan 更简单)
- 对差值音频做 astats,获取差值的 RMS、峰值等指标
"""
from __future__ import annotations
import json
import re
import shutil
import subprocess # nosec B404
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any
FFMPEG_BIN: str = shutil.which("ffmpeg") or "ffmpeg"
FFPROBE_BIN: str = shutil.which("ffprobe") or "ffprobe"
@dataclass
class AudioDiffResult:
"""音频对比结果."""
audio_a: str
audio_b: str
duration_a: float
duration_b: float
duration_diff: float
sample_rate_match: bool
channels_match: bool
diff_rms_db: float # 差值音频的 RMS(dB,越低越相似)
diff_peak_db: float # 差值音频的峰值(dB,越低越相似)
similarity_score: float # 综合相似度评分 [0, 1],1 = 完全一致
passed: bool
def to_dict(self) -> dict[str, Any]:
return asdict(self)
def probe_duration(file_path: str) -> float:
"""探测文件时长(秒),失败返回 0."""
try:
result = subprocess.run( # nosec B603
[
FFPROBE_BIN,
"-v",
"error",
"-show_entries",
"format=duration",
"-of",
"default=noprint_wrappers=1:nokey=1",
str(file_path),
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=10,
)
return round(float(result.stdout.strip()), 3)
except Exception:
return 0.0
def probe_has_audio(file_path: str | Path) -> bool:
"""探测文件是否包含音频流."""
try:
result = subprocess.run( # nosec B603
[
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"a:0",
"-show_entries",
"stream=codec_type",
"-of",
"default=noprint_wrappers=1:nokey=1",
str(file_path),
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=10,
)
return result.stdout.strip() == "audio"
except Exception:
return False # 探测失败保守返回 False,避免误判有音频
def compute_audio_diff(
audio_a: str | Path,
audio_b: str | Path,
*,
similarity_threshold: float = 0.90,
duration_tolerance: float = 0.1,
) -> AudioDiffResult:
"""计算两段音频的差异.
方案:用 pan 滤镜将两轨相减,对差值音频做 astats 分析。
Args:
audio_a: 音频A(基线)
audio_b: 音频B(对比)
similarity_threshold: 相似度合格阈值
duration_tolerance: 时长容忍度(秒)
Returns:
AudioDiffResult 对比结果
"""
dur_a = probe_duration(str(audio_a))
dur_b = probe_duration(str(audio_b))
duration_diff = abs(dur_a - dur_b)
# 获取音频元信息
info_a = _probe_audio_info(str(audio_a))
info_b = _probe_audio_info(str(audio_b))
sample_rate_match = info_a["sample_rate"] == info_b["sample_rate"]
channels_match = info_a["channels"] == info_b["channels"]
# 相减后分析差值
# 取较短时长做对比
min_dur = min(dur_a, dur_b)
if min_dur <= 0:
return AudioDiffResult(
audio_a=str(audio_a),
audio_b=str(audio_b),
duration_a=dur_a,
duration_b=dur_b,
duration_diff=duration_diff,
sample_rate_match=sample_rate_match,
channels_match=channels_match,
diff_rms_db=-999.0,
diff_peak_db=-999.0,
similarity_score=0.0,
passed=False,
)
# 做差值音频:a - b
# 注意:amix 会自动按输入数归一化音量(除以N),
# 所以 a + (-1)*b 经过 amix=inputs=2 后整体音量会减半(-6dB)。
# 加 volume=2 补偿回来,确保差值 RMS 反映真实差异幅度。
command = [
FFMPEG_BIN,
"-i",
str(audio_a),
"-i",
str(audio_b),
"-filter_complex",
# 第2轨反相 → amix混合 → volume=2补偿amix的自动缩放
"[1:a]volume=-1[inv];[0:a][inv]amix=inputs=2:duration=shortest:dropout_transition=0,volume=2[diff]",
"-map",
"[diff]",
"-f",
"null",
"-af",
"astats=metadata=1:reset=0",
"-",
]
try:
result = subprocess.run( # nosec B603
command,
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=120,
)
stderr = result.stderr or ""
except subprocess.CalledProcessError as e:
# 如果音频格式不兼容,返回失败
return AudioDiffResult(
audio_a=str(audio_a),
audio_b=str(audio_b),
duration_a=dur_a,
duration_b=dur_b,
duration_diff=duration_diff,
sample_rate_match=sample_rate_match,
channels_match=channels_match,
diff_rms_db=999.0,
diff_peak_db=999.0,
similarity_score=0.0,
passed=False,
)
diff_rms_db, diff_peak_db = _parse_astats(stderr)
# 相似度评分:基于差值 RMS
# 差值 RMS -60dB → 相似度 ~1.0(几乎无声差)
# 差值 RMS -20dB → 相似度 ~0.5(有明显差异)
# 差值 RMS 0dB → 相似度 ~0.0(完全相反)
if diff_rms_db <= -60:
similarity_score = 1.0
elif diff_rms_db >= 0:
similarity_score = 0.0
else:
# 线性映射:-60dB → 1.0, 0dB → 0.0
similarity_score = max(0.0, min(1.0, 1.0 + diff_rms_db / 60.0))
passed = (
duration_diff <= duration_tolerance
and sample_rate_match
and channels_match
and similarity_score >= similarity_threshold
)
return AudioDiffResult(
audio_a=str(audio_a),
audio_b=str(audio_b),
duration_a=round(dur_a, 3),
duration_b=round(dur_b, 3),
duration_diff=round(duration_diff, 3),
sample_rate_match=sample_rate_match,
channels_match=channels_match,
diff_rms_db=round(diff_rms_db, 2),
diff_peak_db=round(diff_peak_db, 2),
similarity_score=round(similarity_score, 4),
passed=passed,
)
def _probe_audio_info(file_path: str) -> dict[str, int]:
"""探测音频元信息."""
try:
result = subprocess.run( # nosec B603
[
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"a:0",
"-show_entries",
"stream=sample_rate,channels",
"-of",
"json",
file_path,
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=10,
)
info = json.loads(result.stdout)
stream = info.get("streams", [{}])[0]
return {
"sample_rate": int(stream.get("sample_rate", 44100)),
"channels": int(stream.get("channels", 2)),
}
except Exception:
return {"sample_rate": 0, "channels": 0}
def _parse_astats(stderr: str) -> tuple[float, float]:
"""从 astats 输出中解析 RMS 和峰值.
astats 输出格式(在 stderr 中):
[Parsed_astats_1 @ 0x...] Channel: 1
[Parsed_astats_1 @ 0x...] ...
[Parsed_astats_1 @ 0x...] Overall
[Parsed_astats_1 @ 0x...] DC offset: 0.000000
[Parsed_astats_1 @ 0x...] Min level: -0.123456
[Parsed_astats_1 @ 0x...] Max level: 0.789012
[Parsed_astats_1 @ 0x...] Peak level dB: -2.01
[Parsed_astats_1 @ 0x...] RMS level dB: -10.56
...
"""
lines = stderr.split("\n")
rms_db = -999.0
peak_db = -999.0
for line in lines:
# 找 Overall 部分的统计(双声道时取整体值)
rms_match = re.search(r"RMS level dB:\s*(-?\d+\.?\d*)", line)
peak_match = re.search(r"Peak level dB:\s*(-?\d+\.?\d*)", line)
if rms_match:
rms_db = float(rms_match.group(1))
if peak_match:
peak_db = float(peak_match.group(1))
return rms_db, peak_db
def extract_audio(video_path: str | Path, output_path: str | Path) -> Path:
"""从视频中提取音频(AAC 格式).
Args:
video_path: 视频文件路径
output_path: 输出音频路径
Returns:
输出音频文件路径
"""
command = [
FFMPEG_BIN,
"-y",
"-i",
str(video_path),
"-vn",
"-acodec",
"aac",
"-b:a",
"128k",
str(output_path),
]
subprocess.run(command, check=True, capture_output=True, timeout=120) # nosec B603
return Path(output_path)
+628
View File
@@ -0,0 +1,628 @@
"""灰度对比测试 Runner — 新旧引擎批量对比 + 报告生成.
使用方法:
# 配置环境变量
export STAGING_API_URL=https://api.staging.example.com
export STAGING_API_KEY=your_key
# 运行全部 P0 场景
python -m tests.render_compare.runner --priority P0 --output ./report/
# 只跑指定场景
python -m tests.render_compare.runner --scenario simple_pass_through,subtitle_rendering
对比流程:
1. 对每个场景,分别提交到 legacy 和 unified 引擎(通过 Feature Flag 白名单/百分比控制)
- 方式A:通过内部 API 临时切换 flag(需要 admin key
- 方式B:提交任务时指定 engine 参数(如果 API 支持)
2. 等待任务完成,下载输出视频
3. 像素对比(SSIM + PSNR+ 音频对比(差值RMS)
4. 生成 HTML 对比报告
注意:默认假设 API 支持 `engine` 参数来指定渲染引擎。
如果不支持,需要先通过内部 API 切换 Feature Flag,然后提交任务。
"""
from __future__ import annotations
import argparse
import json
import os
import sys
import time
from dataclasses import dataclass, field
from datetime import datetime
from pathlib import Path
from typing import Any
import httpx
# 确保项目根目录在 path 中
sys.path.insert(0, str(Path(__file__).parent.parent.parent))
from .audio_diff import AudioDiffResult, compute_audio_diff
from .scenarios import SCENARIOS, CompareScenario, get_scenarios_by_priority
from .video_diff import VideoDiffResult, compute_video_diff
@dataclass
class ScenarioResult:
"""单个场景的对比结果."""
scenario: CompareScenario
legacy_task_id: str = ""
unified_task_id: str = ""
legacy_video_path: str = ""
unified_video_path: str = ""
legacy_duration_sec: float = 0.0
unified_duration_sec: float = 0.0
video_diff: VideoDiffResult | None = None
audio_diff: AudioDiffResult | None = None
legacy_success: bool = False
unified_success: bool = False
error: str = ""
@property
def passed(self) -> bool:
if not (self.legacy_success and self.unified_success):
return False
if self.video_diff and not self.video_diff.passed:
return False
if self.audio_diff and not self.audio_diff.passed:
return False
return True
class StagingAPI:
"""Staging 环境 API 客户端."""
def __init__(self, base_url: str, api_key: str, internal_api_key: str = ""):
self.base_url = base_url.rstrip("/")
self.api_key = api_key
self.internal_api_key = internal_api_key
self.client = httpx.Client(timeout=30.0)
def _headers(self, internal: bool = False) -> dict[str, str]:
headers = {"Authorization": f"Bearer {self.api_key}"}
if internal and self.internal_api_key:
headers["X-API-Key"] = self.internal_api_key
return headers
def submit_render_task(self, plan_payload: dict[str, Any], engine: str = "") -> str:
"""提交渲染任务,返回 task_id.
Args:
plan_payload: EditPlan payload
engine: 可选,指定引擎("legacy" / "unified"
Returns:
task_id
"""
url = f"{self.base_url}/api/v1/render/compose"
payload = dict(plan_payload)
if engine:
payload["engine"] = engine
resp = self.client.post(url, json=payload, headers=self._headers())
resp.raise_for_status()
data = resp.json()
return data.get("task_id") or data.get("id", "")
def get_task_status(self, task_id: str) -> dict[str, Any]:
"""获取任务状态."""
url = f"{self.base_url}/api/v1/tasks/{task_id}"
resp = self.client.get(url, headers=self._headers())
resp.raise_for_status()
return resp.json()
def wait_for_task(self, task_id: str, timeout: float = 300.0, poll_interval: float = 3.0) -> dict[str, Any]:
"""等待任务完成.
Returns:
最终任务状态
Raises:
TimeoutError: 超时
"""
start = time.time()
while time.time() - start < timeout:
status = self.get_task_status(task_id)
state = status.get("status", "")
if state in ("completed", "success", "done", "failed", "error"):
return status
time.sleep(poll_interval)
raise TimeoutError(f"Task {task_id} timed out after {timeout}s")
def set_feature_flag(self, flag_name: str, enabled: bool, percentage: int = 0, whitelist: list[str] | None = None):
"""通过内部 API 设置 Feature Flag.
用于不支持 engine 参数的场景,切换全局灰度比例。
"""
if not self.internal_api_key:
raise ValueError("internal_api_key is required for feature flag operations")
url = f"{self.base_url}/api/v1/internal/feature-flags/{flag_name}"
body: dict[str, Any] = {"enabled": enabled, "percentage": percentage}
if whitelist is not None:
body["whitelist"] = whitelist
resp = self.client.put(url, json=body, headers=self._headers(internal=True))
resp.raise_for_status()
return resp.json()
def get_feature_flag(self, flag_name: str) -> dict[str, Any]:
"""获取 Feature Flag 配置."""
if not self.internal_api_key:
raise ValueError("internal_api_key is required")
url = f"{self.base_url}/api/v1/internal/feature-flags/{flag_name}"
resp = self.client.get(url, headers=self._headers(internal=True))
resp.raise_for_status()
return resp.json()
def download_video(self, video_url: str, output_path: str | Path) -> Path:
"""下载视频文件."""
output_path = Path(output_path)
output_path.parent.mkdir(parents=True, exist_ok=True)
with self.client.stream("GET", video_url, timeout=60.0) as resp:
resp.raise_for_status()
with open(output_path, "wb") as f:
for chunk in resp.iter_bytes():
f.write(chunk)
return output_path
class CompareRunner:
"""新旧引擎对比 Runner."""
# 全局默认阈值(唯一真实来源,所有入口统一引用)
DEFAULT_SSIM_THRESHOLD: float = 0.95
DEFAULT_PSNR_THRESHOLD: float = 28.0
DEFAULT_AUDIO_SIMILARITY_THRESHOLD: float = 0.90
DEFAULT_DURATION_TOLERANCE: float = 0.1
DEFAULT_TASK_TIMEOUT: float = 300.0
def __init__(
self,
api: StagingAPI,
output_dir: Path,
*,
ssim_threshold: float | None = None,
psnr_threshold: float | None = None,
audio_similarity_threshold: float | None = None,
task_timeout: float | None = None,
flag_mode: bool = False, # 是否使用 Feature Flag 方式切换引擎
duration_tolerance: float | None = None,
):
self.api = api
self.output_dir = output_dir
self.ssim_threshold = ssim_threshold if ssim_threshold is not None else self.DEFAULT_SSIM_THRESHOLD
self.psnr_threshold = psnr_threshold if psnr_threshold is not None else self.DEFAULT_PSNR_THRESHOLD
self.audio_similarity_threshold = (
audio_similarity_threshold
if audio_similarity_threshold is not None
else self.DEFAULT_AUDIO_SIMILARITY_THRESHOLD
)
self.duration_tolerance = (
duration_tolerance if duration_tolerance is not None else self.DEFAULT_DURATION_TOLERANCE
)
self.task_timeout = task_timeout if task_timeout is not None else self.DEFAULT_TASK_TIMEOUT
self.flag_mode = flag_mode
self.results: list[ScenarioResult] = []
# flag_mode 下保存原始配置,测试结束后恢复(防污染线上)
self._original_flag_config: dict[str, Any] | None = None
def run_scenario(self, scenario: CompareScenario) -> ScenarioResult:
"""运行单个场景对比."""
print(f"\n{'='*60}")
print(f"[{scenario.priority}] {scenario.id}: {scenario.name}")
print(f" {scenario.description}")
result = ScenarioResult(scenario=scenario)
scenario_dir = self.output_dir / scenario.id
scenario_dir.mkdir(parents=True, exist_ok=True)
try:
# 1. 提交两个引擎的任务
legacy_task_id = self._submit_with_engine(scenario, "legacy")
unified_task_id = self._submit_with_engine(scenario, "unified")
result.legacy_task_id = legacy_task_id
result.unified_task_id = unified_task_id
print(f" legacy task: {legacy_task_id}")
print(f" unified task: {unified_task_id}")
# 2. 等待完成
print(" waiting for legacy...", end="", flush=True)
legacy_status = self.api.wait_for_task(legacy_task_id, timeout=self.task_timeout)
result.legacy_success = legacy_status.get("status") in ("completed", "success", "done")
legacy_video_url = legacy_status.get("output_url", "") or legacy_status.get("video_url", "")
print(f" {'' if result.legacy_success else ''} ({legacy_status.get('duration_sec', '?')}s)")
print(" waiting for unified...", end="", flush=True)
unified_status = self.api.wait_for_task(unified_task_id, timeout=self.task_timeout)
result.unified_success = unified_status.get("status") in ("completed", "success", "done")
unified_video_url = unified_status.get("output_url", "") or unified_status.get("video_url", "")
print(f" {'' if result.unified_success else ''} ({unified_status.get('duration_sec', '?')}s)")
result.legacy_duration_sec = float(legacy_status.get("duration_sec", 0))
result.unified_duration_sec = float(unified_status.get("duration_sec", 0))
if not (result.legacy_success and result.unified_success):
result.error = f"Legacy success={result.legacy_success}, Unified success={result.unified_success}"
print(" ⚠️ 任务未全部成功,跳过对比")
return result
# 3. 下载视频
print(" downloading...", end="", flush=True)
legacy_path = self.api.download_video(legacy_video_url, scenario_dir / "legacy.mp4")
unified_path = self.api.download_video(unified_video_url, scenario_dir / "unified.mp4")
result.legacy_video_path = str(legacy_path)
result.unified_video_path = str(unified_path)
print("")
# 4. 像素对比
print(" computing video diff...", end="", flush=True)
result.video_diff = compute_video_diff(
legacy_path,
unified_path,
ssim_threshold=self.ssim_threshold,
psnr_threshold=self.psnr_threshold,
duration_tolerance=self.duration_tolerance,
)
print(
f" SSIM={result.video_diff.avg_ssim:.4f} PSNR={result.video_diff.avg_psnr:.2f}dB {'' if result.video_diff.passed else ''}"
)
# 5. 音频对比(仅当都有音频时)
from .audio_diff import probe_has_audio
legacy_has_audio = probe_has_audio(legacy_path)
unified_has_audio = probe_has_audio(unified_path)
if legacy_has_audio and unified_has_audio:
print(" computing audio diff...", end="", flush=True)
result.audio_diff = compute_audio_diff(
legacy_path,
unified_path,
similarity_threshold=self.audio_similarity_threshold,
)
print(
f" similarity={result.audio_diff.similarity_score:.4f} {'' if result.audio_diff.passed else ''}"
)
elif legacy_has_audio != unified_has_audio:
result.error = f"音频不一致: legacy_has_audio={legacy_has_audio}, unified_has_audio={unified_has_audio}"
print(f" ⚠️ 音频不一致: legacy={legacy_has_audio}, unified={unified_has_audio}")
else:
print(" audio: both silent (skip)")
except Exception as e:
result.error = str(e)
print(f" ❌ 错误: {e}")
self.results.append(result)
return result
def _submit_with_engine(self, scenario: CompareScenario, engine: str) -> str:
"""提交指定引擎的任务.
如果 flag_mode=True,通过 Feature Flag 切换,否则通过 engine 参数。
"""
if self.flag_mode:
# 先设置 flag(用白名单方式,确保只有当前测试用户命中)
percentage = 0 if engine == "legacy" else 100
self.api.set_feature_flag("render_engine", enabled=True, percentage=percentage)
time.sleep(1) # 给 worker 一点时间刷新配置
return self.api.submit_render_task(scenario.plan_payload)
else:
return self.api.submit_render_task(scenario.plan_payload, engine=engine)
def run_all(self, scenarios: list[CompareScenario]) -> list[ScenarioResult]:
"""运行所有场景.
flag_mode=True 时,测试开始前保存原始 Feature Flag 配置,
结束后(无论成功失败)自动恢复,避免污染线上环境。
"""
print(f"\n灰度对比测试开始 - {len(scenarios)} 个场景")
print(f"输出目录: {self.output_dir}")
print(f"视频阈值: SSIM>={self.ssim_threshold}, PSNR>={self.psnr_threshold}dB")
print(f"音频阈值: similarity>={self.audio_similarity_threshold}")
# flag_mode:保存原始配置,测试结束后恢复(防污染)
if self.flag_mode:
try:
self._original_flag_config = self.api.get_feature_flag("render_engine")
print(f" [flag_mode] 已保存原始配置: {self._original_flag_config}")
except Exception as e:
print(f" ⚠️ [flag_mode] 保存原始配置失败: {e}")
print(" 为避免污染线上,将中止测试。请检查 internal_api_key 配置。")
return self.results
try:
for i, scenario in enumerate(scenarios):
print(f"\n进度: {i+1}/{len(scenarios)}")
self.run_scenario(scenario)
finally:
# 始终恢复原始 flag 配置
if self.flag_mode and self._original_flag_config:
try:
orig = self._original_flag_config
self.api.set_feature_flag(
"render_engine",
enabled=orig.get("enabled", False),
percentage=orig.get("percentage", 0),
whitelist=orig.get("whitelist"),
)
print("\n[flag_mode] ✅ 已恢复原始 Feature Flag 配置")
except Exception as e:
print(f"\n[flag_mode] ❌ 恢复 Feature Flag 失败: {e}")
print(" 请手动检查并恢复 render_engine flag 配置!")
return self.results
def summary(self) -> dict[str, Any]:
"""生成汇总统计."""
total = len(self.results)
passed = sum(1 for r in self.results if r.passed)
failed = total - passed
# 性能对比
perf_diffs = []
for r in self.results:
if r.legacy_success and r.unified_success and r.legacy_duration_sec > 0:
diff_pct = (r.unified_duration_sec - r.legacy_duration_sec) / r.legacy_duration_sec * 100
perf_diffs.append(diff_pct)
avg_perf_diff = sum(perf_diffs) / len(perf_diffs) if perf_diffs else 0.0
return {
"total": total,
"passed": passed,
"failed": failed,
"pass_rate": f"{passed/total*100:.1f}%" if total > 0 else "0%",
"avg_perf_diff_pct": round(avg_perf_diff, 2),
"scenarios": [self._result_to_dict(r) for r in self.results],
"timestamp": datetime.now().isoformat(),
"ssim_threshold": self.ssim_threshold,
"psnr_threshold": self.psnr_threshold,
"audio_threshold": self.audio_similarity_threshold,
}
def _result_to_dict(self, r: ScenarioResult) -> dict[str, Any]:
return {
"id": r.scenario.id,
"name": r.scenario.name,
"priority": r.scenario.priority,
"passed": r.passed,
"legacy_success": r.legacy_success,
"unified_success": r.unified_success,
"legacy_duration_sec": r.legacy_duration_sec,
"unified_duration_sec": r.unified_duration_sec,
"video_diff": r.video_diff.to_dict() if r.video_diff else None,
"audio_diff": r.audio_diff.to_dict() if r.audio_diff else None,
"error": r.error,
}
def generate_html_report(summary: dict[str, Any], output_path: Path):
"""生成 HTML 对比报告."""
scenarios = summary["scenarios"]
# 按通过/失败分组
passed_list = [s for s in scenarios if s["passed"]]
failed_list = [s for s in scenarios if not s["passed"]]
# 构建场景卡片
scenario_cards = ""
for s in scenarios:
status_class = "pass" if s["passed"] else "fail"
status_text = "✅ 通过" if s["passed"] else "❌ 失败"
vdiff = s.get("video_diff") or {}
adiff = s.get("audio_diff") or {}
video_info = ""
if vdiff:
video_info = f"""
<div class="metric-row">
<span>SSIM:</span>
<span class="{'good' if vdiff.get('avg_ssim', 0) >= 0.95 else 'warn'}">{vdiff.get('avg_ssim', 0):.4f}</span>
</div>
<div class="metric-row">
<span>PSNR:</span>
<span>{vdiff.get('avg_psnr', 0):.2f} dB</span>
</div>
<div class="metric-row">
<span>时长差:</span>
<span>{vdiff.get('duration_diff', 0):.3f}s</span>
</div>
"""
audio_info = ""
if adiff:
audio_info = f"""
<div class="metric-row">
<span>音频相似度:</span>
<span class="{'good' if adiff.get('similarity_score', 0) >= 0.9 else 'warn'}">{adiff.get('similarity_score', 0):.4f}</span>
</div>
<div class="metric-row">
<span>差值 RMS:</span>
<span>{adiff.get('diff_rms_db', 0):.2f} dB</span>
</div>
"""
perf_info = ""
if s["legacy_duration_sec"] and s["unified_duration_sec"]:
diff = s["unified_duration_sec"] - s["legacy_duration_sec"]
pct = diff / s["legacy_duration_sec"] * 100 if s["legacy_duration_sec"] else 0
trend = "🔴" if pct > 10 else ("🟡" if pct > 0 else "🟢")
perf_info = f"""
<div class="perf-row">
<span>Legacy: {s['legacy_duration_sec']:.2f}s</span>
<span>Unified: {s['unified_duration_sec']:.2f}s</span>
<span>{trend} {pct:+.1f}%</span>
</div>
"""
error_info = f'<div class="error-box">{s["error"]}</div>' if s["error"] else ""
scenario_cards += f"""
<div class="card {status_class}">
<div class="card-header">
<span class="badge">{s['priority']}</span>
<span class="scenario-name">{s['name']}</span>
<span class="status {status_class}">{status_text}</span>
</div>
<div class="card-body">
<div class="grid-2">
<div>
<h4>视频质量</h4>
{video_info or '<p class="muted">无数据</p>'}
</div>
<div>
<h4>音频质量</h4>
{audio_info or '<p class="muted">无音频或跳过</p>'}
</div>
</div>
<div>
<h4>性能对比</h4>
{perf_info or '<p class="muted">无数据</p>'}
</div>
{error_info}
</div>
</div>
"""
html = f"""<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>统一渲染引擎灰度对比报告</title>
<style>
* {{ box-sizing: border-box; margin: 0; padding: 0; }}
body {{ font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, sans-serif; background: #f5f5f5; color: #333; padding: 20px; }}
.container {{ max-width: 1200px; margin: 0 auto; }}
h1 {{ margin-bottom: 20px; font-size: 24px; }}
.summary {{ background: white; border-radius: 12px; padding: 24px; margin-bottom: 24px; display: flex; gap: 32px; flex-wrap: wrap; }}
.summary-item {{ text-align: center; }}
.summary-item .value {{ font-size: 32px; font-weight: bold; margin-bottom: 4px; }}
.summary-item .label {{ color: #666; font-size: 14px; }}
.pass .value {{ color: #10b981; }}
.fail .value {{ color: #ef4444; }}
.card {{ background: white; border-radius: 12px; margin-bottom: 16px; overflow: hidden; border-left: 4px solid #10b981; }}
.card.fail {{ border-left-color: #ef4444; }}
.card-header {{ padding: 16px 20px; background: #fafafa; display: flex; align-items: center; gap: 12px; border-bottom: 1px solid #eee; }}
.badge {{ background: #e5e7eb; color: #374151; padding: 2px 8px; border-radius: 4px; font-size: 12px; font-weight: 600; }}
.scenario-name {{ flex: 1; font-weight: 600; }}
.status {{ font-weight: 600; }}
.status.pass {{ color: #10b981; }}
.status.fail {{ color: #ef4444; }}
.card-body {{ padding: 20px; }}
.grid-2 {{ display: grid; grid-template-columns: 1fr 1fr; gap: 24px; margin-bottom: 16px; }}
h4 {{ margin-bottom: 12px; color: #374151; font-size: 14px; }}
.metric-row {{ display: flex; justify-content: space-between; padding: 6px 0; font-size: 14px; }}
.metric-row .good {{ color: #10b981; font-weight: 600; }}
.metric-row .warn {{ color: #f59e0b; font-weight: 600; }}
.perf-row {{ display: flex; gap: 24px; padding: 8px 0; font-size: 14px; background: #f9fafb; padding: 12px; border-radius: 8px; }}
.error-box {{ background: #fef2f2; color: #dc2626; padding: 12px; border-radius: 8px; margin-top: 12px; font-size: 13px; }}
.muted {{ color: #9ca3af; font-size: 14px; }}
.timestamp {{ text-align: center; color: #9ca3af; font-size: 12px; margin-top: 24px; }}
</style>
</head>
<body>
<div class="container">
<h1>🎬 统一渲染引擎灰度对比报告</h1>
<div class="summary">
<div class="summary-item">
<div class="value">{summary['total']}</div>
<div class="label">总场景数</div>
</div>
<div class="summary-item pass">
<div class="value">{summary['passed']}</div>
<div class="label">通过</div>
</div>
<div class="summary-item fail">
<div class="value">{summary['failed']}</div>
<div class="label">失败</div>
</div>
<div class="summary-item">
<div class="value">{summary['pass_rate']}</div>
<div class="label">通过率</div>
</div>
<div class="summary-item">
<div class="value {'good' if summary['avg_perf_diff_pct'] <= 0 else 'warn'}" style="font-size: 24px; color: {'#10b981' if summary['avg_perf_diff_pct'] <= 0 else '#f59e0b'}">{summary['avg_perf_diff_pct']:+.1f}%</div>
<div class="label">平均性能差异</div>
</div>
</div>
{scenario_cards}
<div class="timestamp">生成时间: {summary['timestamp']}</div>
</div>
</body>
</html>"""
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_text(html, encoding="utf-8")
return output_path
def main():
parser = argparse.ArgumentParser(description="统一渲染引擎灰度对比测试")
parser.add_argument("--priority", default="P0", choices=["P0", "P1", "P2"], help="最低优先级")
parser.add_argument("--scenarios", default="", help="指定场景ID,逗号分隔")
parser.add_argument("--output", default="./gray_compare_report", help="输出目录")
parser.add_argument("--ssim-threshold", type=float, default=None, help="SSIM阈值(默认0.95")
parser.add_argument("--psnr-threshold", type=float, default=None, help="PSNR阈值(dB)(默认28.0")
parser.add_argument("--audio-threshold", type=float, default=None, help="音频相似度阈值(默认0.90")
parser.add_argument("--flag-mode", action="store_true", help="使用Feature Flag方式切换引擎")
parser.add_argument("--task-timeout", type=float, default=300.0, help="单任务超时时间(秒)")
args = parser.parse_args()
base_url = os.environ.get("STAGING_API_URL", "")
api_key = os.environ.get("STAGING_API_KEY", "")
internal_key = os.environ.get("STAGING_INTERNAL_API_KEY", "")
if not base_url or not api_key:
print("❌ 请设置环境变量 STAGING_API_URL 和 STAGING_API_KEY")
sys.exit(1)
# 选择场景
if args.scenarios:
scenario_ids = [s.strip() for s in args.scenarios.split(",")]
selected = [s for s in SCENARIOS if s.id in scenario_ids]
if not selected:
print(f"❌ 未找到匹配的场景: {scenario_ids}")
print(f"可用场景: {[s.id for s in SCENARIOS]}")
sys.exit(1)
else:
selected = get_scenarios_by_priority(args.priority)
output_dir = Path(args.output).resolve()
output_dir.mkdir(parents=True, exist_ok=True)
api = StagingAPI(base_url, api_key, internal_key)
runner = CompareRunner(
api,
output_dir,
ssim_threshold=args.ssim_threshold,
psnr_threshold=args.psnr_threshold,
audio_similarity_threshold=args.audio_threshold,
flag_mode=args.flag_mode,
task_timeout=args.task_timeout,
)
runner.run_all(selected)
# 生成报告
summary = runner.summary()
# JSON 报告
json_path = output_dir / "report.json"
json_path.write_text(json.dumps(summary, indent=2, ensure_ascii=False), encoding="utf-8")
# HTML 报告
html_path = output_dir / "report.html"
generate_html_report(summary, html_path)
print(f"\n{'='*60}")
print(f"对比完成: {summary['passed']}/{summary['total']} 通过 ({summary['pass_rate']})")
print(f"报告: {html_path}")
print(f"JSON: {json_path}")
if __name__ == "__main__":
main()
+257
View File
@@ -0,0 +1,257 @@
"""灰度对比测试场景定义 — 覆盖典型渲染场景.
每个场景对应一个 EditPlan,用于新旧引擎对比。
覆盖场景:
1. 简单直通(单clip无特效)
2. 多clip转场(fade + slide
3. 画中画(main + overlay
4. 字幕渲染(ASS字幕)
5. 独立音频轨(主视频 + BGM
6. 多图层混合(main + broll + overlay + audio
7. 背景图片 + 主视频(图片背景无音频)
8. 无音频视频(纯画面,验证无音轨防御)
9. 长视频(10+ clip,压力测试)
10. 分辨率非标(竖屏9:16,验证scale策略)
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any
@dataclass
class CompareScenario:
"""对比测试场景."""
id: str
name: str
description: str
priority: str # P0 / P1 / P2
plan_payload: dict[str, Any] # EditPlan JSON payload(提交给 API 的数据)
expected: dict[str, Any] = field(default_factory=dict) # 预期结果
SCENARIOS: list[CompareScenario] = [
CompareScenario(
id="simple_pass_through",
name="简单直通",
description="单主clip,无转场无特效,验证直通优化路径",
priority="P0",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 5.0,
"order": 0,
}
],
},
),
CompareScenario(
id="multi_clip_transition",
name="多clip转场",
description="3个clipfade + slideleft 转场",
priority="P0",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 3.0,
"order": 0,
"transition_effect": "cut",
},
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 3.0,
"order": 1,
"transition_effect": "fade",
},
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 3.0,
"order": 2,
"transition_effect": "slideleft",
},
],
},
),
CompareScenario(
id="picture_in_picture",
name="画中画",
description="主视频 + 角落小窗(corner_voice",
priority="P1",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{"clip_type": "main", "asset_id": "sample_5s.mp4", "duration": 5.0, "order": 0},
{"clip_type": "corner_voice", "asset_id": "sample_5s.mp4", "duration": 5.0, "order": 0},
],
},
),
CompareScenario(
id="subtitle_rendering",
name="字幕渲染",
description="主视频 + ASS字幕",
priority="P0",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 5.0,
"order": 0,
"config": {"subtitles": [{"text": "测试字幕 Test Subtitle", "start_time": 0, "end_time": 5.0}]},
}
],
},
),
CompareScenario(
id="independent_audio_track",
name="独立音频轨",
description="主视频(带音频)+ 独立BGM轨,验证音频混音",
priority="P0",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{"clip_type": "main", "asset_id": "sample_5s.mp4", "duration": 5.0, "order": 0},
{
"clip_type": "main",
"asset_id": "sample_bgm.mp3",
"duration": 5.0,
"order": 0,
"config": {"role": "audio", "volume": 0.5},
},
],
},
),
CompareScenario(
id="multi_layer_mix",
name="多图层混合",
description="main + broll + overlay + audio 四图层",
priority="P1",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 4.0,
"order": 0,
"transition_effect": "fade",
},
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 4.0,
"order": 1,
"transition_effect": "slideup",
},
{"clip_type": "broll", "asset_id": "sample_broll.mp4", "duration": 8.0, "order": 0},
{"clip_type": "overlay", "asset_id": "sample_overlay.png", "duration": 8.0, "order": 0},
{
"clip_type": "main",
"asset_id": "sample_bgm.mp3",
"duration": 8.0,
"order": 0,
"config": {"role": "audio", "volume": 0.3},
},
],
},
),
CompareScenario(
id="image_background",
name="图片背景",
description="background图片层 + 主视频,验证背景层无音频",
priority="P1",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{"clip_type": "background", "asset_id": "sample_bg.jpg", "duration": 5.0, "order": 0},
{"clip_type": "main", "asset_id": "sample_5s.mp4", "duration": 5.0, "order": 0},
],
},
),
CompareScenario(
id="no_audio_video",
name="无音轨视频",
description="源视频无音频流,验证无音轨防御逻辑",
priority="P0",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{"clip_type": "main", "asset_id": "sample_silent_5s.mp4", "duration": 5.0, "order": 0},
],
},
),
CompareScenario(
id="long_video_stress",
name="长视频压力",
description="10个clip + 多种转场,性能压力测试",
priority="P2",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 3.0,
"order": i,
"transition_effect": ["cut", "fade", "slideleft", "slidedown", "dissolve"][i % 5],
}
for i in range(10)
],
},
),
CompareScenario(
id="vertical_portrait",
name="竖屏9:16",
description="竖屏分辨率,验证scale策略(铺满裁剪)",
priority="P2",
plan_payload={
"width": 720,
"height": 1280,
"fps": 25,
"clips": [
{"clip_type": "main", "asset_id": "sample_5s.mp4", "duration": 5.0, "order": 0},
],
},
),
]
def get_scenarios_by_priority(min_priority: str = "P2") -> list[CompareScenario]:
"""按优先级过滤场景.
P0 包含 P0
P1 包含 P0 + P1
P2 包含全部
"""
priority_order = {"P0": 0, "P1": 1, "P2": 2}
threshold = priority_order.get(min_priority, 2)
return [s for s in SCENARIOS if priority_order.get(s.priority, 2) <= threshold]
+283
View File
@@ -0,0 +1,283 @@
"""视频对比工具 — 基于 FFmpeg 的像素级质量对比.
使用 SSIM + PSNR 双指标评估两个视频的相似度:
- SSIM (Structural Similarity): 结构相似性,范围 [0, 1],越接近 1 越相似
- PSNR (Peak Signal-to-Noise Ratio): 峰值信噪比,单位 dB,越高越好
灰度验收标准:
- 平均 SSIM >= 0.95 → 视觉上几乎无差异(P0 场景必达)
- 最低 SSIM >= 0.90 → 最严重帧差异可接受
- 平均 PSNR >= 28dB → 质量达标
"""
from __future__ import annotations
import json
import re
import shutil
import subprocess # nosec B404
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any
FFMPEG_BIN: str = shutil.which("ffmpeg") or "ffmpeg"
FFPROBE_BIN: str = shutil.which("ffprobe") or "ffprobe"
@dataclass
class VideoDiffResult:
"""视频对比结果."""
video_a: str
video_b: str
width: int
height: int
duration_a: float
duration_b: float
avg_ssim: float
min_ssim: float
avg_psnr: float # dB
min_psnr: float
frame_count: int
duration_diff: float # 时长差(秒)
resolution_match: bool
passed: bool # 是否通过阈值
def to_dict(self) -> dict[str, Any]:
return asdict(self)
def probe_video_info(video_path: str) -> dict[str, Any]:
"""获取视频信息(宽、高、时长、fps."""
try:
result = subprocess.run( # nosec B603
[
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"v:0",
"-show_entries",
"stream=width,height,r_frame_rate,duration",
"-show_entries",
"format=duration",
"-of",
"json",
video_path,
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=10,
)
info = json.loads(result.stdout)
stream = info.get("streams", [{}])[0]
fmt = info.get("format", {})
width = int(stream.get("width", 1280))
height = int(stream.get("height", 720))
fps_str = stream.get("r_frame_rate", "25/1")
if "/" in fps_str:
num, den = fps_str.split("/")
fps = float(num) / float(den) if float(den) > 0 else 25.0
else:
fps = float(fps_str) if fps_str else 25.0
duration = float(fmt.get("duration", 0)) or float(stream.get("duration", 0))
return {"width": width, "height": height, "duration": duration, "fps": round(fps, 2)}
except Exception:
return {"width": 1280, "height": 720, "duration": 0.0, "fps": 25.0}
def compute_video_diff(
video_a: str | Path,
video_b: str | Path,
*,
ssim_threshold: float = 0.95,
psnr_threshold: float = 28.0,
duration_tolerance: float = 0.1,
) -> VideoDiffResult:
"""计算两个视频的像素差异.
使用 FFmpeg ssim + psnr 滤镜一次性计算两个指标。
Args:
video_a: 视频A路径(基线)
video_b: 视频B路径(对比)
ssim_threshold: SSIM 合格阈值(默认 0.90
psnr_threshold: PSNR 合格阈值(默认 25dB
duration_tolerance: 时长容忍度(秒,默认 0.1s)
Returns:
VideoDiffResult 对比结果
Raises:
subprocess.CalledProcessError: FFmpeg 执行失败
"""
info_a = probe_video_info(str(video_a))
info_b = probe_video_info(str(video_b))
duration_diff = abs(info_a["duration"] - info_b["duration"])
resolution_match = info_a["width"] == info_b["width"] and info_a["height"] == info_b["height"]
# ssim 和 psnr 的 stats_file 都输出到 stdout
# 用行格式区分:SSIM 行含 "All:"PSNR 行含 "psnr_avg:"
command = [
FFMPEG_BIN,
"-i",
str(video_a),
"-i",
str(video_b),
"-lavfi",
"[0:v][1:v]ssim=stats_file=-[out1];[0:v][1:v]psnr=stats_file=-[out2]",
"-f",
"null",
"-",
]
result = subprocess.run( # nosec B603
command,
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=300,
)
# 逐帧统计在 stdoutstats_file=-),汇总日志在 stderr
stats_stdout = result.stdout or ""
avg_ssim, min_ssim = _parse_ssim_stats(stats_stdout)
avg_psnr, min_psnr = _parse_psnr_stats(stats_stdout)
frame_count = _count_frames(result.stderr or "")
passed = (
resolution_match
and duration_diff <= duration_tolerance
and avg_ssim >= ssim_threshold
and avg_psnr >= psnr_threshold
)
return VideoDiffResult(
video_a=str(video_a),
video_b=str(video_b),
width=info_a["width"],
height=info_a["height"],
duration_a=round(info_a["duration"], 3),
duration_b=round(info_b["duration"], 3),
avg_ssim=round(avg_ssim, 6),
min_ssim=round(min_ssim, 6),
avg_psnr=round(avg_psnr, 3),
min_psnr=round(min_psnr, 3),
frame_count=frame_count,
duration_diff=round(duration_diff, 3),
resolution_match=resolution_match,
passed=passed,
)
def _parse_ssim_stats(stats_output: str) -> tuple[float, float]:
"""从 SSIM stats_file 输出中解析逐帧 SSIM.
FFmpeg ssim 滤镜 stats_file 输出格式(每行一帧):
n:1 Y:0.987654 U:0.991234 V:0.990000 All:0.989000 (19.585642)
n:2 Y:0.986543 U:0.990123 V:0.988888 All:0.987654 (19.123456)
...
Returns:
(avg_ssim, min_ssim)
"""
ssim_values: list[float] = []
for line in stats_output.split("\n"):
# 匹配 stats_file 格式:n:数字 ... All:数字
if not line.startswith("n:"):
continue
match = re.search(r"All:(\d+\.\d+)", line)
if match:
ssim_values.append(float(match.group(1)))
if not ssim_values:
return 0.0, 0.0
avg_ssim = sum(ssim_values) / len(ssim_values)
min_ssim = min(ssim_values)
return avg_ssim, min_ssim
def _parse_psnr_stats(stats_output: str) -> tuple[float, float]:
"""从 PSNR stats_file 输出中解析逐帧 PSNR.
FFmpeg psnr 滤镜 stats_file 输出格式(每行一帧):
n:1 mse_avg:100.23 mse_y:150.12 mse_u:50.34 mse_v:80.56 psnr_avg:28.12 psnr_y:26.34 psnr_u:31.12 psnr_v:29.08
n:2 ...
Returns:
(avg_psnr, min_psnr) — avg_psnr 是逐帧 psnr_avg 的均值,min_psnr 是逐帧最小值
"""
psnr_values: list[float] = []
for line in stats_output.split("\n"):
if not line.startswith("n:"):
continue
match = re.search(r"psnr_avg:(\d+\.\d+)", line)
if match:
psnr_values.append(float(match.group(1)))
if not psnr_values:
return 0.0, 0.0
avg_psnr = sum(psnr_values) / len(psnr_values)
min_psnr = min(psnr_values)
return avg_psnr, min_psnr
def _count_frames(stderr: str) -> int:
"""从 FFmpeg 输出中统计帧数."""
match = re.search(r"frame=\s*(\d+)", stderr)
return int(match.group(1)) if match else 0
def save_diff_frame(
video_a: str | Path,
video_b: str | Path,
output_path: str | Path,
*,
timestamp: float = 1.0,
) -> Path:
"""生成差异帧可视化图(红绿色差).
使用 blend 滤镜生成差异可视化图,差异越大越亮。
Args:
video_a: 视频A
video_b: 视频B
output_path: 输出图片路径
timestamp: 截取的时间点(秒)
Returns:
输出图片路径
"""
command = [
FFMPEG_BIN,
"-y",
"-ss",
str(timestamp),
"-i",
str(video_a),
"-ss",
str(timestamp),
"-i",
str(video_b),
"-lavfi",
"[0:v][1:v]blend=all_mode=difference,eq=contrast=5:brightness=0.5[diff]",
"-map",
"[diff]",
"-vframes",
"1",
str(output_path),
]
subprocess.run(command, check=True, capture_output=True, timeout=60) # nosec B603
return Path(output_path)
+72
View File
@@ -0,0 +1,72 @@
"""AssetStatus 枚举兼容性测试。
验证历史脏数据(如 'uploaded')不会导致枚举转换失败。
"""
import pytest
from packages.domain.entities import AssetStatus
class TestAssetStatusNormalValues:
"""正常值应该正确映射。"""
def test_uploading(self):
assert AssetStatus("uploading") == AssetStatus.UPLOADING
def test_ready(self):
assert AssetStatus("ready") == AssetStatus.READY
def test_processing(self):
assert AssetStatus("processing") == AssetStatus.PROCESSING
def test_error(self):
assert AssetStatus("error") == AssetStatus.ERROR
class TestAssetStatusHistoricalValues:
"""历史脏数据应该正确映射到对应状态,不抛异常。"""
@pytest.mark.parametrize("value", ["uploaded", "Uploaded", "UPLOADED", " uploaded "])
def test_uploaded_maps_to_ready(self, value):
"""生产环境发现的 'uploaded' 历史值应映射为 READY。"""
assert AssetStatus(value) == AssetStatus.READY
@pytest.mark.parametrize("value", ["success", "ok", "done", "complete"])
def test_other_ready_like_values_map_to_ready(self, value):
assert AssetStatus(value) == AssetStatus.READY
@pytest.mark.parametrize("value", ["upload", "uploading_start", "upload_start"])
def test_upload_like_values_map_to_uploading(self, value):
assert AssetStatus(value) == AssetStatus.UPLOADING
@pytest.mark.parametrize("value", ["failed", "fail", "err"])
def test_error_like_values_map_to_error(self, value):
assert AssetStatus(value) == AssetStatus.ERROR
@pytest.mark.parametrize("value", ["process", "running", "run"])
def test_processing_like_values_map_to_processing(self, value):
assert AssetStatus(value) == AssetStatus.PROCESSING
class TestAssetStatusFallback:
"""完全未知的值兜底为 READY,不抛500。"""
@pytest.mark.parametrize("value", ["unknown", "foo_bar", ""])
def test_unknown_value_falls_back_to_ready(self, value):
assert AssetStatus(value) == AssetStatus.READY
def test_none_value_falls_back_to_ready(self):
assert AssetStatus(None) == AssetStatus.READY # type: ignore[arg-type]
def test_int_value_falls_back_to_ready(self):
assert AssetStatus(123) == AssetStatus.READY # type: ignore[arg-type]
class TestAssetStatusStrValue:
"""枚举值仍为字符串类型,不影响序列化。"""
def test_value_unchanged(self):
assert AssetStatus.READY.value == "ready"
assert AssetStatus.ERROR.value == "error"
assert isinstance(AssetStatus.READY, str)
+68
View File
@@ -0,0 +1,68 @@
"""ClassificationStatus 枚举兼容性测试。
验证历史脏数据(如 'done')不会导致枚举转换失败。
"""
import pytest
from packages.domain.entities import ClassificationStatus
class TestClassificationStatusNormalValues:
"""正常值应该正确映射。"""
def test_pending(self):
assert ClassificationStatus("pending") == ClassificationStatus.PENDING
def test_processing(self):
assert ClassificationStatus("processing") == ClassificationStatus.PROCESSING
def test_completed(self):
assert ClassificationStatus("completed") == ClassificationStatus.COMPLETED
def test_failed(self):
assert ClassificationStatus("failed") == ClassificationStatus.FAILED
class TestClassificationStatusHistoricalValues:
"""历史脏数据应该正确映射到对应状态,不抛异常。"""
@pytest.mark.parametrize("value", ["done", "Done", "DONE", " done "])
def test_done_maps_to_completed(self, value):
"""生产环境发现的 'done' 历史值应映射为 COMPLETED。"""
assert ClassificationStatus(value) == ClassificationStatus.COMPLETED
@pytest.mark.parametrize("value", ["success", "finished", "complete"])
def test_other_done_like_values_map_to_completed(self, value):
assert ClassificationStatus(value) == ClassificationStatus.COMPLETED
@pytest.mark.parametrize("value", ["fail", "error", "err"])
def test_error_like_values_map_to_failed(self, value):
assert ClassificationStatus(value) == ClassificationStatus.FAILED
@pytest.mark.parametrize("value", ["process", "running", "run"])
def test_processing_like_values_map_to_processing(self, value):
assert ClassificationStatus(value) == ClassificationStatus.PROCESSING
class TestClassificationStatusFallback:
"""完全未知的值兜底为 PENDING,不抛500。"""
@pytest.mark.parametrize("value", ["unknown", "foo_bar", ""])
def test_unknown_value_falls_back_to_pending(self, value):
assert ClassificationStatus(value) == ClassificationStatus.PENDING
def test_none_value_falls_back_to_pending(self):
assert ClassificationStatus(None) == ClassificationStatus.PENDING # type: ignore[arg-type]
def test_int_value_falls_back_to_pending(self):
assert ClassificationStatus(123) == ClassificationStatus.PENDING # type: ignore[arg-type]
class TestClassificationStatusStrValue:
"""枚举值仍为字符串类型,不影响序列化。"""
def test_value_unchanged(self):
assert ClassificationStatus.COMPLETED.value == "completed"
assert ClassificationStatus.PENDING.value == "pending"
assert isinstance(ClassificationStatus.COMPLETED, str)
+18 -19
View File
@@ -197,7 +197,11 @@ class TestSubmitCloneTask:
payload = mock_client.request.call_args.kwargs["json"]
# 中文和特殊字符被过滤,剩下字母数字
assert payload["input"]["prefix"] == "2024" or payload["input"]["prefix"] == "clone" or len(payload["input"]["prefix"]) <= 10
assert (
payload["input"]["prefix"] == "2024"
or payload["input"]["prefix"] == "clone"
or len(payload["input"]["prefix"]) <= 10
)
def test_submit_auth_401_raises(self) -> None:
mock_client = MagicMock()
@@ -305,9 +309,7 @@ class TestPollCloneTask:
def test_poll_undeployed_raises_error(self) -> None:
mock_client = MagicMock()
mock_client.request.return_value = _mock_response(
200, {"output": {"status": "UNDEPLOYED"}}
)
mock_client.request.return_value = _mock_response(200, {"output": {"status": "UNDEPLOYED"}})
service = _make_service(http_client=mock_client)
service.CLONE_POLL_INTERVAL = 0.01
@@ -317,9 +319,7 @@ class TestPollCloneTask:
def test_poll_timeout_raises(self) -> None:
mock_client = MagicMock()
mock_client.request.return_value = _mock_response(
200, {"output": {"status": "DEPLOYING"}}
)
mock_client.request.return_value = _mock_response(200, {"output": {"status": "DEPLOYING"}})
service = _make_service(http_client=mock_client)
service.CLONE_POLL_INTERVAL = 0.01
@@ -396,9 +396,7 @@ class TestSynthesizeSpeech:
)
service = _make_service(http_client=mock_client)
result = service.synthesize_speech(
text="你好世界", voice_id="longxiaochun_v3"
)
result = service.synthesize_speech(text="你好世界", voice_id="longxiaochun_v3")
assert isinstance(result, SynthesizeResult)
assert result.audio_url == "https://dashscope-result.oss.com/output.mp3"
@@ -426,8 +424,12 @@ class TestSynthesizeSpeech:
service = _make_service(http_client=mock_client)
service.synthesize_speech(
text="test", voice_id="v1", sample_rate=44100,
format="wav", speed=1.5, volume=80,
text="test",
voice_id="v1",
sample_rate=44100,
format="wav",
speed=1.5,
volume=80,
)
payload = mock_client.request.call_args.kwargs["json"]
@@ -463,7 +465,8 @@ class TestSynthesizeSpeech:
"""同步接口的 submit_synthesize_task 返回空 task_id 字段(兼容旧接口)."""
mock_client = MagicMock()
mock_client.request.return_value = _mock_response(
200, {"output": {"audio": {"url": "https://e.com/a.mp3"}}},
200,
{"output": {"audio": {"url": "https://e.com/a.mp3"}}},
)
service = _make_service(http_client=mock_client)
@@ -490,9 +493,7 @@ class TestRetryLogic:
mock_client.request.side_effect = [
_mock_response(500, text="Server Error"),
_mock_response(502, text="Bad Gateway"),
_mock_response(
200, {"output": {"audio": {"url": "https://e.com/a.mp3"}}}
),
_mock_response(200, {"output": {"audio": {"url": "https://e.com/a.mp3"}}}),
]
service = _make_service(http_client=mock_client)
@@ -547,9 +548,7 @@ class TestSanitizePrefix:
class TestCheckTaskStatus:
def test_check_task_status_uses_query_voice(self) -> None:
mock_client = MagicMock()
mock_client.request.return_value = _mock_response(
200, {"output": {"status": "OK"}}
)
mock_client.request.return_value = _mock_response(200, {"output": {"status": "OK"}})
service = _make_service(http_client=mock_client)
result = service.check_task_status("voice-123")
+12
View File
@@ -177,6 +177,18 @@ class StubGenerationTaskRepository:
def count_by_user(self, user_id: str) -> int:
return len([t for t in self._store.values() if t.created_by_user_id == user_id])
def count_pending_by_user(self, user_id: str) -> int:
return len(
[
t
for t in self._store.values()
if t.created_by_user_id == user_id and getattr(t, "status", "") == "pending"
]
)
def count_pending_total(self) -> int:
return len([t for t in self._store.values() if getattr(t, "status", "") == "pending"])
def list_recent_by_user(self, user_id: str, limit: int = 5) -> list[Any]:
items = [t for t in self._store.values() if t.created_by_user_id == user_id]
items.sort(key=lambda t: t.created_at, reverse=True)
+6
View File
@@ -194,6 +194,12 @@ class StubGenerationTaskRepository:
self._tasks[task.id] = task
return task
def count_pending_by_user(self, user_id: str) -> int:
return 0
def count_pending_total(self) -> int:
return 0
# ---------------------------------------------------------------------------
# Service factory
+2
View File
@@ -63,6 +63,8 @@ class StubEditPlan:
template_id: str = "tmpl-001"
status: Any = None
config: dict = field(default_factory=dict)
project_id: str = ""
created_by_user_id: str = "user-001"
def mark_failed(self):
self.status = _StubStatus("failed")
+424
View File
@@ -0,0 +1,424 @@
"""Feature Flag 单元测试。
测试 FeatureFlagConfig、InMemoryFeatureFlagStore、RenderEngineResolver 的核心逻辑。
"""
from __future__ import annotations
import time
from unittest.mock import MagicMock, patch
import pytest
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
InMemoryFeatureFlagStore,
)
# ── FeatureFlagConfig 测试 ──────────────────────────────────────────────────
class TestFeatureFlagConfig:
"""FeatureFlagConfig 核心逻辑测试。"""
def test_default_disabled(self):
"""默认配置为关闭状态。"""
config = FeatureFlagConfig(name="test_flag")
assert config.enabled is False
assert config.percentage == 0
assert config.whitelist == set()
assert config.is_active() is False
assert config.is_active("user1") is False
def test_global_enabled_100_percent(self):
"""100% + 启用 = 全部命中。"""
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=100)
assert config.is_active() is True
assert config.is_active("user1") is True
assert config.is_active("any_user") is True
def test_global_enabled_0_percent_no_whitelist(self):
"""启用但 0% 且无白名单 = 不命中。"""
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=0)
assert config.is_active() is False
assert config.is_active("user1") is False
def test_whitelist_takes_priority(self):
"""白名单优先级高于百分比。"""
config = FeatureFlagConfig(
name="test_flag",
enabled=True,
percentage=0,
whitelist={"user1", "user2"},
)
assert config.is_active("user1") is True
assert config.is_active("user2") is True
assert config.is_active("user3") is False
def test_whitelist_with_percentage(self):
"""白名单用户即使百分比为0也命中,非白名单按百分比。"""
config = FeatureFlagConfig(
name="test_flag",
enabled=True,
percentage=100, # 100% 所有人命中
whitelist={"user1"},
)
assert config.is_active("user1") is True
assert config.is_active("user999") is True # 100% 命中
def test_percentage_consistency_same_user(self):
"""同一用户多次调用结果一致(哈希确定性)。"""
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=50)
results = [config.is_active("user_fixed") for _ in range(100)]
assert all(r == results[0] for r in results)
def test_percentage_different_users_distributed(self):
"""不同用户分布大致符合百分比(统计检验,宽松阈值)。"""
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=50)
active_count = sum(1 for i in range(1000) if config.is_active(f"user_{i}"))
# 50% 上下浮动 10% 都算合理
assert 400 <= active_count <= 600, f"Expected ~500, got {active_count}"
def test_percentage_boundary_0_and_100(self):
"""0% 和 100% 的边界情况。"""
config_0 = FeatureFlagConfig(name="test", enabled=True, percentage=0)
config_100 = FeatureFlagConfig(name="test", enabled=True, percentage=100)
for i in range(100):
assert config_0.is_active(f"user_{i}") is False
assert config_100.is_active(f"user_{i}") is True
def test_disabled_ignores_all_other_settings(self):
"""关闭时忽略白名单和百分比。"""
config = FeatureFlagConfig(
name="test_flag",
enabled=False,
percentage=100,
whitelist={"user1"},
)
assert config.is_active("user1") is False
assert config.is_active() is False
def test_none_identifier_with_percentage(self):
"""无 identifier 时按随机比例(0% 和 100% 是确定的)。"""
config_0 = FeatureFlagConfig(name="test", enabled=True, percentage=0)
config_100 = FeatureFlagConfig(name="test", enabled=True, percentage=100)
assert config_0.is_active(None) is False
assert config_100.is_active(None) is True
def test_to_dict_and_from_dict(self):
"""序列化和反序列化对称。"""
original = FeatureFlagConfig(
name="test_flag",
enabled=True,
percentage=30,
whitelist={"user_a", "user_b", "user_c"},
)
data = original.to_dict()
restored = FeatureFlagConfig.from_dict(data)
assert restored.name == original.name
assert restored.enabled == original.enabled
assert restored.percentage == original.percentage
assert restored.whitelist == original.whitelist
def test_from_dict_with_missing_fields(self):
"""from_dict 缺失字段时使用默认值。"""
config = FeatureFlagConfig.from_dict({"name": "minimal"})
assert config.name == "minimal"
assert config.enabled is False
assert config.percentage == 0
assert config.whitelist == set()
# ── InMemoryFeatureFlagStore 测试 ───────────────────────────────────────────
class TestInMemoryFeatureFlagStore:
"""内存存储实现测试。"""
def test_get_nonexistent_returns_default(self):
"""获取不存在的 flag 返回默认配置(关闭)。"""
store = InMemoryFeatureFlagStore()
config = store.get("nonexistent")
assert config.name == "nonexistent"
assert config.enabled is False
def test_set_and_get(self):
"""设置后可以读取。"""
store = InMemoryFeatureFlagStore()
config = FeatureFlagConfig(name="test", enabled=True, percentage=50, whitelist={"u1"})
store.set(config)
got = store.get("test")
assert got.enabled is True
assert got.percentage == 50
assert got.whitelist == {"u1"}
def test_delete_existing(self):
"""删除存在的 flag 返回 True。"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="test", enabled=True))
assert store.delete("test") is True
assert store.get("test").enabled is False
def test_delete_nonexistent(self):
"""删除不存在的 flag 返回 False。"""
store = InMemoryFeatureFlagStore()
assert store.delete("nonexistent") is False
def test_list_all(self):
"""列出所有 flag。"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="flag_a", enabled=True))
store.set(FeatureFlagConfig(name="flag_b", percentage=10))
all_flags = store.list_all()
assert len(all_flags) == 2
assert "flag_a" in all_flags
assert "flag_b" in all_flags
assert all_flags["flag_a"].enabled is True
def test_is_active_convenience(self):
"""is_active 便捷方法。"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render", enabled=True, percentage=0, whitelist={"vip_user"}))
assert store.is_active("render", "vip_user") is True
assert store.is_active("render", "normal_user") is False
assert store.is_active("nonexistent") is False
# ── RenderEngineResolver 测试 ───────────────────────────────────────────────
class TestRenderEngineResolver:
"""渲染引擎选择器测试。"""
def test_default_legacy_when_flag_disabled(self):
"""flag 关闭时使用默认引擎(legacy)。"""
store = InMemoryFeatureFlagStore()
resolver = self._make_resolver(store=store, default="legacy")
assert resolver.get_engine() == "legacy"
assert resolver.get_engine("user1") == "legacy"
def test_default_unified_when_flag_disabled(self):
"""flag 关闭但默认值是 unified 时返回 unified。"""
store = InMemoryFeatureFlagStore()
resolver = self._make_resolver(store=store, default="unified")
assert resolver.get_engine() == "unified"
def test_whitelist_user_uses_unified(self):
"""白名单用户走新引擎。"""
store = InMemoryFeatureFlagStore()
store.set(
FeatureFlagConfig(
name="render_engine",
enabled=True,
percentage=0,
whitelist={"beta_tester"},
)
)
resolver = self._make_resolver(store=store, default="legacy")
assert resolver.get_engine("beta_tester") == "unified"
assert resolver.get_engine("normal_user") == "legacy"
def test_100_percent_all_unified(self):
"""100% 时所有用户走新引擎。"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render_engine", enabled=True, percentage=100))
resolver = self._make_resolver(store=store, default="legacy")
for i in range(50):
assert resolver.get_engine(f"user_{i}") == "unified"
def test_invalid_default_engine_fallback(self):
"""无效默认值回退到 legacy。"""
store = InMemoryFeatureFlagStore()
resolver = self._make_resolver(store=store, default="invalid_value")
assert resolver.get_engine() == "legacy"
def test_should_use_unified_helper(self):
"""should_use_unified 便捷方法。"""
store = InMemoryFeatureFlagStore()
store.set(
FeatureFlagConfig(
name="render_engine",
enabled=True,
percentage=0,
whitelist={"user_a"},
)
)
resolver = self._make_resolver(store=store)
assert resolver.should_use_unified("user_a") is True
assert resolver.should_use_unified("user_b") is False
def test_config_snapshot(self):
"""配置快照。"""
store = InMemoryFeatureFlagStore()
store.set(
FeatureFlagConfig(
name="render_engine",
enabled=True,
percentage=30,
whitelist={"u1", "u2"},
)
)
resolver = self._make_resolver(store=store)
snapshot = resolver.get_config_snapshot()
assert snapshot["flag_name"] == "render_engine"
assert snapshot["enabled"] is True
assert snapshot["percentage"] == 30
assert snapshot["whitelist"] == ["u1", "u2"]
def test_set_flag_updates_config(self):
"""通过 set_flag 修改后立即生效。"""
store = InMemoryFeatureFlagStore()
resolver = self._make_resolver(store=store, default="legacy")
# 初始:关闭
assert resolver.get_engine("user1") == "legacy"
# 开启 100%
resolver.set_flag(FeatureFlagConfig(name="render_engine", enabled=True, percentage=100))
assert resolver.get_engine("user1") == "unified"
# 关闭
resolver.set_flag(FeatureFlagConfig(name="render_engine", enabled=False))
assert resolver.get_engine("user1") == "legacy"
def test_force_refresh(self):
"""强制刷新不报错。"""
store = InMemoryFeatureFlagStore()
resolver = self._make_resolver(store=store)
resolver.force_refresh() # 不抛异常即可
def test_does_not_affect_in_flight_tasks(self):
"""
热更新不影响在途任务验证:
任务开始时确定引擎,中途配置变更不改变当前任务的引擎选择。
(这是通过"每次调用 get_engine 时读取当前配置"来保证的,
任务开始时调用一次拿到结果,之后不再变化)
"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render_engine", enabled=True, percentage=100))
resolver = self._make_resolver(store=store, default="legacy")
# 模拟任务开始时获取引擎
engine_at_start = resolver.get_engine("user1")
assert engine_at_start == "unified"
# 任务进行中关闭 flag
store.set(FeatureFlagConfig(name="render_engine", enabled=False))
resolver.force_refresh()
# 在途任务持有的 engine_at_start 仍然是 unified(不随配置变化)
assert engine_at_start == "unified"
# 新任务会拿到 legacy
assert resolver.get_engine("user1") == "legacy"
# ── 辅助方法 ──
@staticmethod
def _make_resolver(store=None, default="legacy"):
from apps.worker.video_processing.render_engine_resolver import (
RenderEngineResolver,
)
return RenderEngineResolver(
default_engine=default,
store=store or InMemoryFeatureFlagStore(),
refresh_interval=9999, # 测试时禁用自动刷新
)
# ── RedisFeatureFlagStore 降级测试(无 Redis 环境) ───────────────────────
class TestRedisStoreDegradation:
"""Redis 不可用时的降级行为测试。"""
def test_get_returns_default_when_redis_unavailable(self):
"""Redis 连接失败时返回默认关闭配置,不抛异常。"""
import importlib
from packages.adapters.redis import feature_flag_store as ff_module
# 模拟 redis 模块不存在的场景不好做,这里直接测试异常捕获逻辑
store = ff_module.RedisFeatureFlagStore.__new__(ff_module.RedisFeatureFlagStore)
store._redis = MagicMock()
store._redis.hgetall.side_effect = ConnectionError("Redis down")
store._key_prefix = ff_module.FEATURE_FLAG_REDIS_PREFIX
store._cache = {}
store._cache_ttl = 5.0
import threading
store._lock = threading.Lock()
config = store.get("render_engine")
assert config.enabled is False
assert config.name == "render_engine"
def test_list_all_returns_empty_on_redis_error(self):
"""Redis 错误时 list_all 返回空字典。"""
import importlib
from packages.adapters.redis import feature_flag_store as ff_module
store = ff_module.RedisFeatureFlagStore.__new__(ff_module.RedisFeatureFlagStore)
store._redis = MagicMock()
store._redis.scan.side_effect = ConnectionError("Redis down")
store._key_prefix = ff_module.FEATURE_FLAG_REDIS_PREFIX
store._cache = {}
store._cache_ttl = 5.0
import threading
store._lock = threading.Lock()
result = store.list_all()
assert result == {}
class TestRedisStoreListAll:
"""RedisFeatureFlagStore list_all 正常路径测试。"""
def _make_store(self):
from packages.adapters.redis import feature_flag_store as ff_module
store = ff_module.RedisFeatureFlagStore.__new__(ff_module.RedisFeatureFlagStore)
store._redis = MagicMock()
store._key_prefix = ff_module.FEATURE_FLAG_REDIS_PREFIX
store._cache = {}
store._cache_ttl = 5.0
import threading
store._lock = threading.Lock()
return store
def test_list_all_scan_with_match_param(self):
"""list_all 调用 redis.scan 时使用正确的 match 参数名。"""
store = self._make_store()
prefix = store._key_prefix
# 模拟 scan 返回 2 个 key,分 2 次游标
store._redis.scan.side_effect = [
(10, [f"{prefix}render_engine", f"{prefix}other_flag"]),
(0, []),
]
# 模拟 hgetall 返回配置
store._redis.hgetall.return_value = {
b"enabled": b"true",
b"percentage": b"50",
b"whitelist": b'["user1","user2"]',
}
result = store.list_all()
# 验证 scan 被调用了 2 次(游标遍历)
assert store._redis.scan.call_count == 2
# 验证参数名是 match(不是 match_pattern
first_call_kwargs = store._redis.scan.call_args_list[0][1]
assert "match" in first_call_kwargs
assert "match_pattern" not in first_call_kwargs
assert first_call_kwargs["match"] == f"{prefix}*"
# 验证返回了 2 个 flag
assert len(result) == 2
assert "render_engine" in result
assert "other_flag" in result
@@ -0,0 +1,93 @@
"""FFmpeg 超时保护测试。
验证 run_ffmpeg / probe_video_info 的超时保护机制,
防止 FFmpeg hang 住导致 worker 永久阻塞。
"""
from __future__ import annotations
import subprocess
from unittest.mock import MagicMock, patch
import pytest
from video_processing.ffmpeg_utils import (
DEFAULT_FFMPEG_TIMEOUT,
probe_video_info,
run_ffmpeg,
)
# ── run_ffmpeg 超时保护 ──────────────────────────────────────────────────────
class TestRunFFmpegTimeout:
"""run_ffmpeg 超时保护测试。"""
def test_default_timeout_is_set(self):
"""默认超时应为 1800 秒(30分钟)。"""
assert DEFAULT_FFMPEG_TIMEOUT == 1800
def test_timeout_expired_is_raised(self):
"""超时未完成时 TimeoutExpired 异常被传播。"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_run.side_effect = subprocess.TimeoutExpired(cmd=["ffmpeg", "test"], timeout=1)
with pytest.raises(subprocess.TimeoutExpired):
run_ffmpeg(["ffmpeg", "test"])
def test_custom_timeout(self):
"""支持自定义超时时间。"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_run.side_effect = subprocess.TimeoutExpired(cmd=["ffmpeg"], timeout=5)
with pytest.raises(subprocess.TimeoutExpired):
run_ffmpeg(["ffmpeg", "test"], timeout=5)
def test_none_timeout_disables_protection(self):
"""timeout=None 可以禁用超时保护(不推荐)。"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_result = MagicMock()
mock_result.stdout = ""
mock_result.stderr = ""
mock_run.return_value = mock_result
run_ffmpeg(["ffmpeg", "test"], timeout=None)
# 验证 timeout=None 被传递
call_kwargs = mock_run.call_args.kwargs
assert call_kwargs["timeout"] is None
def test_called_process_error_still_raised(self):
"""超时异常不影响原有 CalledProcessError 的抛出。"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_run.side_effect = subprocess.CalledProcessError(returncode=1, cmd=["ffmpeg"], stderr="error msg")
with pytest.raises(subprocess.CalledProcessError):
run_ffmpeg(["ffmpeg", "test"])
# ── probe_video_info 超时保护 ────────────────────────────────────────────────
class TestProbeVideoInfoTimeout:
"""probe_video_info 超时保护测试。"""
def test_probe_uses_timeout(self):
"""probe_video_info 调用 ffprobe 时应设置 timeout=15。"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_run.side_effect = subprocess.TimeoutExpired(cmd=["ffprobe"], timeout=15)
# 超时异常被捕获,返回默认值
result = probe_video_info("/tmp/test.mp4")
assert result["width"] == 1280 # DEFAULT_OUTPUT_WIDTH
assert result["height"] == 720 # DEFAULT_OUTPUT_HEIGHT
def test_probe_success(self):
"""正常情况应解析 ffprobe JSON 输出。"""
fake_output = """
{
"streams": [{"width": 1920, "height": 1080, "codec_type": "video", "r_frame_rate": "30/1", "duration": "10.5"}],
"format": {"duration": "10.5"}
}
"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_result = MagicMock()
mock_result.stdout = fake_output
mock_run.return_value = mock_result
result = probe_video_info("/tmp/test.mp4")
assert result["width"] == 1920
assert result["height"] == 1080
assert abs(result["duration"] - 10.5) < 0.01
+338
View File
@@ -0,0 +1,338 @@
"""generate_video 任务 Feature Flag 灰度引擎选择单元测试.
覆盖:
- _resolve_render_engine 正常返回 unified / legacy
- Feature Flag 不可用时 fallback 到 unified
- 白名单 / 百分比 / 全局开关各场景
- _render_with_legacy_engine 命令构建与输出验证
"""
from __future__ import annotations
import os
import sys
from datetime import datetime, timezone
from types import ModuleType
from typing import Any
from unittest.mock import MagicMock, patch
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
from pathlib import Path
import pytest
# ── Mock worker 模块以避免数据库连接 ──────────────────────────────────────────
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "worker"))
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", ".."))
_mock_db_mod = ModuleType("worker_app.db")
_mock_db_mod.SessionLocal = MagicMock()
sys.modules.setdefault("worker_app.db", _mock_db_mod)
_mock_celery_mod = ModuleType("worker_app.celery_app")
_mock_celery_app = MagicMock()
_mock_celery_app.task = lambda **kwargs: lambda fn: fn
_mock_celery_mod.celery_app = _mock_celery_app
sys.modules.setdefault("worker_app.celery_app", _mock_celery_mod)
# Mock worker_app.core.config 避免 settings 加载
_mock_config_mod = ModuleType("worker_app.core.config")
_mock_settings = MagicMock()
_mock_settings.redis_url = None
_mock_settings.render_engine = "unified"
_mock_config_mod.get_settings = lambda: _mock_settings
sys.modules.setdefault("worker_app.core", ModuleType("worker_app.core"))
sys.modules.setdefault("worker_app.core.config", _mock_config_mod)
# ── 测试用数据类 ──────────────────────────────────────────────────────────────
class _TestClip:
def __init__(self, asset_id, duration=30.0, clip_type="main", config=None, order=0):
self.id = f"clip_{asset_id}"
self.plan_id = "test-plan"
self.clip_type = clip_type
self.order = order
self.asset_id = asset_id
self.duration = duration
self.config = config or {}
self.start_time = 0.0
self.transition_effect = "cut"
# ── RenderEngineResolver 基础行为测试 ───────────────────────────────────────
def test_resolver_unified_when_enabled_100_percent():
"""flag 全局开启(percentage=100)时,返回 unified。"""
from video_processing.render_engine_resolver import RenderEngineResolver
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
InMemoryFeatureFlagStore,
)
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render_engine", enabled=True, percentage=100))
resolver = RenderEngineResolver(default_engine="legacy", store=store)
assert resolver.get_engine(user_id="user-123") == "unified"
def test_resolver_legacy_when_flag_disabled():
"""flag 全局关闭时,返回默认引擎 legacy。"""
from video_processing.render_engine_resolver import RenderEngineResolver
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
InMemoryFeatureFlagStore,
)
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render_engine", enabled=False, percentage=100))
resolver = RenderEngineResolver(default_engine="legacy", store=store)
assert resolver.get_engine(user_id="user-123") == "legacy"
def test_resolver_whitelist_overrides_percentage_0():
"""白名单用户即使 percentage=0 也走 unified。"""
from video_processing.render_engine_resolver import RenderEngineResolver
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
InMemoryFeatureFlagStore,
)
store = InMemoryFeatureFlagStore()
store.set(
FeatureFlagConfig(
name="render_engine",
enabled=True,
percentage=0,
whitelist={"user-vip"},
)
)
resolver = RenderEngineResolver(default_engine="legacy", store=store)
assert resolver.get_engine(user_id="user-vip") == "unified"
assert resolver.get_engine(user_id="user-other") == "legacy"
def test_resolver_percentage_0_all_legacy():
"""percentage=0 且无白名单时,全部走 legacy。"""
from video_processing.render_engine_resolver import RenderEngineResolver
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
InMemoryFeatureFlagStore,
)
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render_engine", enabled=True, percentage=0))
resolver = RenderEngineResolver(default_engine="legacy", store=store)
for i in range(50):
assert resolver.get_engine(user_id=f"user-{i}") == "legacy"
def test_resolver_default_unified_when_flag_off():
"""默认引擎设为 unified 且 flag 关闭时,返回 unified。"""
from video_processing.render_engine_resolver import RenderEngineResolver
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
InMemoryFeatureFlagStore,
)
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render_engine", enabled=False, percentage=0))
resolver = RenderEngineResolver(default_engine="unified", store=store)
assert resolver.get_engine(user_id="user-123") == "unified"
# ── _render_with_legacy_engine 集成测试 ──────────────────────────────────────
def test_legacy_engine_single_clip_keeps_original_fps():
"""单 clip 场景:输出保持原帧率(不做 fps 归一化),分辨率缩放正确。"""
import subprocess
import tempfile
from video_processing.ffmpeg_utils import probe_video_info
from apps.worker.worker_app.tasks.generation import _render_with_legacy_engine
with tempfile.TemporaryDirectory() as tmpdir:
tmp_path = Path(tmpdir)
input_path = tmp_path / "input.mp4"
output_path = tmp_path / "output.mp4"
# 生成 1 秒 30fps 测试视频(带音频)
subprocess.run(
[
"ffmpeg",
"-y",
"-f",
"lavfi",
"-i",
"color=c=red:s=640x360:d=1:r=30",
"-f",
"lavfi",
"-i",
"anullsrc=r=44100:cl=stereo:d=1",
"-c:v",
"libx264",
"-pix_fmt",
"yuv420p",
"-c:a",
"aac",
"-shortest",
str(input_path),
],
check=True,
capture_output=True,
)
clip = _TestClip(asset_id="asset-1", duration=1.0)
asset_path_map = {"asset-1": input_path}
duration, file_size = _render_with_legacy_engine(
task_id="test-task",
virtual_clips=[clip],
asset_path_map=asset_path_map,
work_dir=tmp_path,
output_path=output_path,
)
assert output_path.exists()
assert file_size > 0
assert duration > 0
# 旧引擎保持原帧率(30fps),不做 fps 归一化
info = probe_video_info(str(output_path))
assert abs(info.get("fps", 0) - 30.0) < 0.5
assert info.get("width") == 1280
assert info.get("height") == 720
def test_legacy_engine_two_clips_concat_duration():
"""多 clip 场景:concat 后时长为两片段之和。"""
import subprocess
import tempfile
from video_processing.ffmpeg_utils import probe_duration
from apps.worker.worker_app.tasks.generation import _render_with_legacy_engine
with tempfile.TemporaryDirectory() as tmpdir:
tmp_path = Path(tmpdir)
input1 = tmp_path / "input1.mp4"
input2 = tmp_path / "input2.mp4"
output_path = tmp_path / "output.mp4"
for idx, inp in enumerate([input1, input2]):
color = "red" if idx == 0 else "blue"
subprocess.run(
[
"ffmpeg",
"-y",
"-f",
"lavfi",
"-i",
f"color=c={color}:s=640x360:d=1:r=30",
"-f",
"lavfi",
"-i",
"anullsrc=r=44100:cl=stereo:d=1",
"-c:v",
"libx264",
"-pix_fmt",
"yuv420p",
"-c:a",
"aac",
"-shortest",
str(inp),
],
check=True,
capture_output=True,
)
clip1 = _TestClip(asset_id="asset-1", duration=1.0, clip_type="main", order=0)
clip2 = _TestClip(asset_id="asset-2", duration=1.0, clip_type="main", order=1)
asset_path_map = {"asset-1": input1, "asset-2": input2}
duration, file_size = _render_with_legacy_engine(
task_id="test-task",
virtual_clips=[clip1, clip2],
asset_path_map=asset_path_map,
work_dir=tmp_path,
output_path=output_path,
)
assert output_path.exists()
assert file_size > 0
assert abs(duration - 2.0) < 0.2
def test_legacy_engine_broll_mode_supported():
"""b_roll 类型的 clip 也被正确识别为主图层并渲染。"""
import subprocess
import tempfile
from apps.worker.worker_app.tasks.generation import _render_with_legacy_engine
with tempfile.TemporaryDirectory() as tmpdir:
tmp_path = Path(tmpdir)
input_path = tmp_path / "input.mp4"
output_path = tmp_path / "output.mp4"
subprocess.run(
[
"ffmpeg",
"-y",
"-f",
"lavfi",
"-i",
"color=c=green:s=640x360:d=1:r=30",
"-f",
"lavfi",
"-i",
"anullsrc=r=44100:cl=stereo:d=1",
"-c:v",
"libx264",
"-pix_fmt",
"yuv420p",
"-c:a",
"aac",
"-shortest",
str(input_path),
],
check=True,
capture_output=True,
)
clip = _TestClip(
asset_id="asset-1",
duration=1.0,
clip_type="main",
config={"role": "b_roll"},
)
asset_path_map = {"asset-1": input_path}
duration, file_size = _render_with_legacy_engine(
task_id="test-task",
virtual_clips=[clip],
asset_path_map=asset_path_map,
work_dir=tmp_path,
output_path=output_path,
)
assert output_path.exists()
assert file_size > 0
assert duration > 0
+6
View File
@@ -58,6 +58,12 @@ class StubGenerationTaskRepository:
def get(self, task_id):
return self._tasks.get(task_id)
def count_pending_by_user(self, user_id):
return 0
def count_pending_total(self):
return 0
class StubGeneratedVideoRepository:
def __init__(self, videos=None):
+190
View File
@@ -0,0 +1,190 @@
"""渲染结果内部下载接口单元测试。
测试 internal_render 路由的核心逻辑,mock 掉 repository 和 storage 依赖。
"""
from __future__ import annotations
from unittest.mock import MagicMock
import pytest
from app.api.routes.internal_render import (
InternalRenderDownloadUrlResponse,
InternalRenderTaskVideosResponse,
_video_to_item,
get_render_task_videos,
get_render_video_download_url,
)
# ── Helpers ────────────────────────────────────────────────────────────────
class MockVideo:
"""模拟 GeneratedVideo 领域对象。"""
def __init__(self, **kwargs):
self.id = kwargs.get("id", "video-1")
self.generation_task_id = kwargs.get("generation_task_id", "task-1")
self.project_id = kwargs.get("project_id", "proj-1")
self.name = kwargs.get("name", "test_video.mp4")
self.file_url = kwargs.get("file_url", "videos/test/output.mp4")
self.file_size = kwargs.get("file_size", 1024000)
self.duration = kwargs.get("duration", 30.5)
self.width = kwargs.get("width", 1080)
self.height = kwargs.get("height", 1920)
self.fps = kwargs.get("fps", 30.0)
self.status = kwargs.get("status", "completed")
# ── _video_to_item 测试 ────────────────────────────────────────────────────
class TestVideoToItem:
"""测试视频对象转响应项。"""
def test_basic_conversion(self):
video = MockVideo(id="v1", generation_task_id="t1", status="completed")
item = _video_to_item(video, "https://oss.example.com/download?v1")
assert item.video_id == "v1"
assert item.generation_task_id == "t1"
assert item.status == "completed"
assert item.download_url == "https://oss.example.com/download?v1"
def test_missing_optional_fields(self):
"""缺可选字段时返回 None。"""
video = MockVideo()
# 去掉可选字段
del video.file_size
del video.duration
item = _video_to_item(video, "https://example.com/dl")
assert item.file_size is None
assert item.duration is None
assert item.width == 1080 # 还在
# ── 路由函数测试 ────────────────────────────────────────────────────────────
class TestGetRenderVideoDownloadUrl:
"""测试单个视频下载URL接口。"""
def test_video_exists(self):
video = MockVideo(id="v-abc", file_url="videos/abc/out.mp4")
mock_repo = MagicMock()
mock_repo.get.return_value = video
mock_storage = MagicMock()
mock_storage.get_download_url.return_value = "https://oss.test/signed?v=abc"
result = get_render_video_download_url(
video_id="v-abc",
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
assert isinstance(result, InternalRenderDownloadUrlResponse)
assert result.video_id == "v-abc"
assert result.download_url == "https://oss.test/signed?v=abc"
mock_repo.get.assert_called_once_with("v-abc")
mock_storage.get_download_url.assert_called_once()
def test_video_not_found_raises_404(self):
from fastapi import HTTPException
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_storage = MagicMock()
with pytest.raises(HTTPException) as exc_info:
get_render_video_download_url(
video_id="nonexistent",
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
assert exc_info.value.status_code == 404
def test_download_url_long_expiry(self):
"""过期时间应为 24 小时(86400s)。"""
video = MockVideo(id="v1")
mock_repo = MagicMock()
mock_repo.get.return_value = video
mock_storage = MagicMock()
mock_storage.get_download_url.return_value = "https://oss.test/signed"
get_render_video_download_url(
video_id="v1",
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
# 验证 expires_seconds=86400
call_kwargs = mock_storage.get_download_url.call_args
assert call_kwargs.kwargs.get("expires_seconds") == 86400 or call_kwargs[1].get("expires_seconds") == 86400
class TestGetRenderTaskVideos:
"""测试任务视频列表接口。"""
def test_list_multiple_videos(self):
videos = [
MockVideo(id="v1", status="completed"),
MockVideo(id="v2", status="completed"),
MockVideo(id="v3", status="failed"),
]
mock_repo = MagicMock()
mock_repo.list_by_generation_task.return_value = videos
mock_storage = MagicMock()
mock_storage.get_download_url.return_value = "https://oss.test/signed"
result = get_render_task_videos(
task_id="task-1",
status=None,
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
assert isinstance(result, InternalRenderTaskVideosResponse)
assert result.task_id == "task-1"
assert result.count == 3
assert len(result.videos) == 3
def test_filter_by_status(self):
videos = [
MockVideo(id="v1", status="completed"),
MockVideo(id="v2", status="completed"),
MockVideo(id="v3", status="failed"),
]
mock_repo = MagicMock()
mock_repo.list_by_generation_task.return_value = videos
mock_storage = MagicMock()
mock_storage.get_download_url.return_value = "https://oss.test/signed"
result = get_render_task_videos(
task_id="task-1",
status="completed",
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
assert result.count == 2
assert all(v.status == "completed" for v in result.videos)
def test_empty_task(self):
mock_repo = MagicMock()
mock_repo.list_by_generation_task.return_value = []
mock_storage = MagicMock()
result = get_render_task_videos(
task_id="empty-task",
status=None,
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
assert result.count == 0
assert result.videos == []
+1 -3
View File
@@ -132,10 +132,8 @@ class TestDownloadLibraryAssets:
session.query.return_value = query
filter_result = MagicMock()
query.filter.return_value = filter_result
in_filter = MagicMock()
filter_result.filter.return_value = in_filter
id_filter = MagicMock()
in_filter.filter.return_value = id_filter
filter_result.filter.return_value = id_filter
assets = [self._make_asset("a1", "video/a1.mp4")]
id_filter.order_by.return_value.all.return_value = assets
+236
View File
@@ -0,0 +1,236 @@
"""P0-stagingOSS 上传崩溃修复测试.
测试:
1. oss_bucket() 传递 connect_timeout 参数
2. upload_to_oss() 小文件走 put_object_from_file,大文件走分片上传
3. upload_to_oss() 超时保护(超过总超时返回 None)
4. upload_to_oss() 异常时返回 None
"""
from __future__ import annotations
import os
import tempfile
import time
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
# ── oss_bucket connect_timeout 测试 ───────────────────────────────────────────
class TestOSSBucketConnectTimeout:
"""测试 oss_bucket() 传递 connect_timeout 参数."""
def test_oss_bucket_has_connect_timeout(self):
"""oss_bucket 应传递 connect_timeout=10s 参数."""
from video_processing.oss_helpers import oss_bucket
mock_bucket_instance = MagicMock()
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket_instance) as mock_bucket_cls,
):
bucket = oss_bucket()
assert bucket is mock_bucket_instance
# 验证 connect_timeout 关键字参数
call_kwargs = mock_bucket_cls.call_args[1]
assert "connect_timeout" in call_kwargs, "oss_bucket 应传递 connect_timeout 参数"
assert (
call_kwargs["connect_timeout"] == 10
), f"connect_timeout 应为 10,实际为 {call_kwargs['connect_timeout']}"
def test_oss_bucket_no_config_returns_none(self):
"""OSS 配置缺失时返回 None."""
from video_processing.oss_helpers import oss_bucket
with patch.dict(os.environ, {}, clear=True):
bucket = oss_bucket()
assert bucket is None
# ── upload_to_oss 分片上传测试 ────────────────────────────────────────────────
class TestUploadToOSSMultipart:
"""测试 upload_to_oss() 根据文件大小选择上传方式."""
def _create_temp_file(self, size_bytes: int) -> Path:
"""创建指定大小的临时文件."""
tmp = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
tmp.write(b"x" * size_bytes)
tmp.close()
return Path(tmp.name)
def test_small_file_uses_put_object(self):
"""小文件(<100MB)走 put_object_from_file."""
from video_processing.oss_helpers import upload_to_oss
small_file = self._create_temp_file(10 * 1024 * 1024) # 10MB
try:
mock_bucket = MagicMock()
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
patch("video_processing.oss_helpers.oss2.resumable_upload") as mock_resumable,
):
url = upload_to_oss(small_file, "test/small.mp4")
# 验证调用了 put_object_from_file
mock_bucket.put_object_from_file.assert_called_once()
# 验证没调用分片上传
mock_resumable.assert_not_called()
# 验证返回 URL
assert url == "https://test-bucket.oss-cn-hangzhou.aliyuncs.com/test/small.mp4"
finally:
small_file.unlink()
def test_large_file_uses_resumable_upload(self):
"""大文件(>=100MB)走 resumable_upload 分片上传."""
from video_processing.oss_helpers import upload_to_oss
large_file = self._create_temp_file(100 * 1024 * 1024) # 100MB
try:
mock_bucket = MagicMock()
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
patch("video_processing.oss_helpers.oss2.resumable_upload") as mock_resumable,
):
url = upload_to_oss(large_file, "test/large.mp4")
# 验证调用了分片上传
mock_resumable.assert_called_once()
# 验证没调用 put_object_from_file
mock_bucket.put_object_from_file.assert_not_called()
# 验证分片参数
call_kwargs = mock_resumable.call_args[1]
assert call_kwargs["multipart_threshold"] == 100 * 1024 * 1024
assert call_kwargs["part_size"] == 8 * 1024 * 1024
assert call_kwargs["num_threads"] == 3
# 验证返回 URL
assert url == "https://test-bucket.oss-cn-hangzhou.aliyuncs.com/test/large.mp4"
finally:
large_file.unlink()
# ── upload_to_oss 超时测试 ────────────────────────────────────────────────────
class TestUploadToOSSTimeout:
"""测试 upload_to_oss() 超时保护."""
def test_upload_timeout_returns_none(self):
"""上传超过总超时时返回 None."""
from video_processing.oss_helpers import upload_to_oss
small_file = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
small_file.write(b"x" * 1024) # 1KB
small_file.close()
file_path = Path(small_file.name)
def slow_upload(*args, **kwargs):
"""模拟慢速上传,超过超时时间."""
time.sleep(2)
mock_bucket = MagicMock()
mock_bucket.put_object_from_file.side_effect = slow_upload
try:
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
patch("video_processing.oss_helpers.OSS_UPLOAD_TOTAL_TIMEOUT", 1), # 1秒超时
):
url = upload_to_oss(file_path, "test/slow.mp4")
# 超时应返回 None
assert url is None, "上传超时应返回 None"
finally:
file_path.unlink()
def test_upload_exception_returns_none(self):
"""上传异常时返回 None."""
from video_processing.oss_helpers import upload_to_oss
small_file = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
small_file.write(b"x" * 1024)
small_file.close()
file_path = Path(small_file.name)
mock_bucket = MagicMock()
mock_bucket.put_object_from_file.side_effect = RuntimeError("Network error")
try:
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
):
url = upload_to_oss(file_path, "test/error.mp4")
assert url is None, "上传异常应返回 None"
finally:
file_path.unlink()
def test_upload_no_bucket_returns_none(self):
"""OSS 未配置时返回 None."""
from video_processing.oss_helpers import upload_to_oss
small_file = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
small_file.write(b"x" * 1024)
small_file.close()
file_path = Path(small_file.name)
try:
with patch.dict(os.environ, {}, clear=True):
url = upload_to_oss(file_path, "test/noconfig.mp4")
assert url is None
finally:
file_path.unlink()
+4 -7
View File
@@ -15,7 +15,6 @@ from video_processing.ffmpeg_utils import build_xfade_filter_chain
# ── P0-3: build_xfade_filter_chain 安全钳制 ──────────────────────────────────
class TestBuildXfadeFilterChainSafetyClamp:
"""验证 xfade 滤镜链的安全钳制逻辑,防止 exit 234。"""
@@ -124,9 +123,7 @@ class TestBuildXfadeFilterChainSafetyClamp:
durations_found.append(float(m.group(1)))
# 第一个 xfade: td 必须 ≤ 0.3 (第二个输入 clip_durations[1]=0.3)
assert durations_found[0] <= 0.3 + 0.001, (
f"第一个 xfade td={durations_found[0]} 超过 clip_durations[1]=0.3"
)
assert durations_found[0] <= 0.3 + 0.001, f"第一个 xfade td={durations_found[0]} 超过 clip_durations[1]=0.3"
# 第二个 xfade: td 可以 = 0.5 (clip_durations[2]=5.0)
assert durations_found[1] <= 0.5 + 0.001
assert dur > 0
@@ -173,9 +170,9 @@ class TestBuildXfadeFilterChainSafetyClamp:
assert dur_val >= 0.001 # 至少 1ms
# P1 修复验证: td 不能超过第二个输入片段时长
second_input_idx = xfade_idx + 1
assert dur_val <= durations[second_input_idx] + 0.001, (
f"td={dur_val} > clip_durations[{second_input_idx}]={durations[second_input_idx]}"
)
assert (
dur_val <= durations[second_input_idx] + 0.001
), f"td={dur_val} > clip_durations[{second_input_idx}]={durations[second_input_idx]}"
xfade_idx += 1
+104 -91
View File
@@ -13,7 +13,6 @@ from unittest.mock import MagicMock, patch
import pytest
# ── oss_bucket endpoint scheme 修复 ──────────────────────────────────────────
@@ -25,17 +24,19 @@ class TestOSSBucketEndpointScheme:
from video_processing.oss_helpers import oss_bucket
mock_bucket_instance = MagicMock()
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth") as mock_auth, patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket_instance
) as mock_bucket_cls:
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth") as mock_auth,
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket_instance) as mock_bucket_cls,
):
# 清除缓存,确保重新创建
import video_processing.oss_helpers as oss_mod
@@ -45,9 +46,7 @@ class TestOSSBucketEndpointScheme:
# 验证 endpoint 传的是带 https:// 的
call_args = mock_bucket_cls.call_args
endpoint_arg = call_args[0][1] # 第 2 个位置参数是 endpoint
assert endpoint_arg.startswith("https://"), (
f"endpoint 应该带 https:// 前缀,实际为: {endpoint_arg}"
)
assert endpoint_arg.startswith("https://"), f"endpoint 应该带 https:// 前缀,实际为: {endpoint_arg}"
assert "oss-cn-hangzhou.aliyuncs.com" in endpoint_arg
def test_endpoint_with_https_keeps_as_is(self):
@@ -55,17 +54,19 @@ class TestOSSBucketEndpointScheme:
from video_processing.oss_helpers import oss_bucket
mock_bucket_instance = MagicMock()
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "https://oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth"), patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket_instance
) as mock_bucket_cls:
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "https://oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket_instance) as mock_bucket_cls,
):
import video_processing.oss_helpers as oss_mod
bucket = oss_bucket()
@@ -81,17 +82,19 @@ class TestOSSBucketEndpointScheme:
from video_processing.oss_helpers import oss_bucket
mock_bucket_instance = MagicMock()
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "http://oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth"), patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket_instance
) as mock_bucket_cls:
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "http://oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket_instance) as mock_bucket_cls,
):
import video_processing.oss_helpers as oss_mod
bucket = oss_bucket()
@@ -133,16 +136,18 @@ class TestGetSignedDownloadUrl:
mock_bucket = MagicMock()
mock_bucket.sign_url.return_value = "https://test-bucket.oss-cn-hangzhou.aliyuncs.com/generated/test.mp4?OSSAccessKeyId=xxx&Expires=xxx&Signature=xxx"
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth"), patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
):
result = get_signed_download_url("generated/test.mp4", expires_seconds=3600)
@@ -155,22 +160,24 @@ class TestGetSignedDownloadUrl:
from video_processing.oss_helpers import get_signed_download_url
mock_bucket = MagicMock()
mock_bucket.sign_url.return_value = "https://test-bucket.oss-cn-hangzhou.aliyuncs.com/generated/test.mp4?sign=xxx"
mock_bucket.sign_url.return_value = (
"https://test-bucket.oss-cn-hangzhou.aliyuncs.com/generated/test.mp4?sign=xxx"
)
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth"), patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
):
result = get_signed_download_url(
"https://test-bucket.oss-cn-hangzhou.aliyuncs.com/generated/test.mp4"
)
result = get_signed_download_url("https://test-bucket.oss-cn-hangzhou.aliyuncs.com/generated/test.mp4")
mock_bucket.sign_url.assert_called_once()
# 验证传给 sign_url 的是纯 storage key,不是完整 URL
@@ -193,16 +200,18 @@ class TestGetSignedDownloadUrl:
mock_bucket = MagicMock()
mock_bucket.sign_url.side_effect = Exception("sign failed")
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth"), patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
):
result = get_signed_download_url("generated/test.mp4")
assert result is None
@@ -223,16 +232,18 @@ class TestUploadToOSSReturnsHTTPS:
from pathlib import Path
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth"), patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
):
result = upload_to_oss(Path("/tmp/test.mp4"), "generated/test.mp4")
@@ -249,16 +260,18 @@ class TestUploadToOSSReturnsHTTPS:
from pathlib import Path
with patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "https://oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
), patch("video_processing.oss_helpers.oss2.Auth"), patch(
"video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "https://oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
):
result = upload_to_oss(Path("/tmp/test.mp4"), "generated/test.mp4")
+424
View File
@@ -0,0 +1,424 @@
"""RenderAdapter 单元测试 — Phase 2.
测试适配层的计划加载、素材下载、引擎调用、结果上传等逻辑。
"""
from __future__ import annotations
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
from unittest.mock import MagicMock, patch
import pytest
from video_processing.render_adapter import RenderAdapter, RenderAdapterResult
# ── Fixtures ──────────────────────────────────────────────────────────────────
@dataclass
class FakeClip:
"""模拟 EditPlanClip。"""
id: str
plan_id: str = "plan_001"
clip_type: str = "main"
order: int = 0
asset_id: str = ""
text_content: str = ""
start_time: float = 0.0
duration: float = 0.0
transition_effect: str = "cut"
status: str = "ready"
config: dict[str, Any] = field(default_factory=dict)
@dataclass
class FakePlan:
"""模拟 EditPlan。"""
id: str = "plan_001"
name: str = "测试计划"
status: str = "editing"
config: dict[str, Any] = field(default_factory=dict)
def _make_clip(
clip_id: str,
clip_type: str = "main",
order: int = 0,
asset_id: str | None = None,
duration: float = 5.0,
status: str = "ready",
transition_effect: str = "cut",
config: dict[str, Any] | None = None,
) -> FakeClip:
# asset_id 为 None 时生成默认值,为空字符串时保留空串
if asset_id is None:
asset_id = f"asset_{clip_id}.mp4"
return FakeClip(
id=clip_id,
clip_type=clip_type,
order=order,
asset_id=asset_id,
duration=duration,
status=status,
transition_effect=transition_effect,
config=config or {},
)
def _make_adapter(
plan: FakePlan | None = None,
clips: list[FakeClip] | None = None,
) -> tuple[RenderAdapter, MagicMock, MagicMock]:
"""创建测试用的 RenderAdapter 及 mock repo。
Returns:
(adapter, mock_plan_repo, mock_clip_repo)
"""
mock_db = MagicMock()
adapter = RenderAdapter(mock_db)
# 替换内部 repo
mock_plan_repo = MagicMock()
mock_clip_repo = MagicMock()
adapter._plan_repo = mock_plan_repo
adapter._clip_repo = mock_clip_repo
# 设置默认返回
if plan is not None:
mock_plan_repo.get.return_value = plan
if clips is not None:
mock_clip_repo.list_by_plan.return_value = clips
return adapter, mock_plan_repo, mock_clip_repo
# ── validate_plan 测试 ───────────────────────────────────────────────────────
class TestValidatePlan:
def test_plan_not_found(self):
"""计划不存在时校验失败。"""
adapter, mock_plan_repo, _ = _make_adapter(plan=None)
mock_plan_repo.get.return_value = None
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
assert not valid
assert len(errors) == 1
assert "不存在" in errors[0]
assert ready_count == 0
assert total_count == 0
def test_no_clips(self):
"""没有任何片段时校验失败。"""
plan = FakePlan(id="plan_001", status="editing")
adapter, _, mock_clip_repo = _make_adapter(plan=plan, clips=[])
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
assert not valid
assert any("没有任何片段" in e for e in errors)
def test_no_ready_clips(self):
"""没有 ready 片段时校验失败。"""
plan = FakePlan(id="plan_001", status="editing")
clips = [
_make_clip("c1", status="pending"),
_make_clip("c2", status="pending"),
]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
assert not valid
assert any("没有就绪" in e for e in errors)
assert ready_count == 0
assert total_count == 2
def test_ready_clip_no_asset(self):
"""ready 片段没有 asset_id 时报错。"""
plan = FakePlan(id="plan_001", status="editing")
clips = [
_make_clip("c1", asset_id=""),
]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
assert not valid
assert any("没有分配素材" in e for e in errors)
def test_valid_plan(self):
"""正常计划校验通过。"""
plan = FakePlan(id="plan_001", status="editing")
clips = [
_make_clip("c1", order=0, duration=3.0),
_make_clip("c2", order=1, duration=4.0),
]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
assert valid
assert len(errors) == 0
assert ready_count == 2
assert total_count == 2
def test_wrong_status(self):
"""计划状态不正确时报错。"""
plan = FakePlan(id="plan_001", status="draft")
clips = [_make_clip("c1")]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
valid, errors, _, _, _ = adapter.validate_plan("plan_001")
assert not valid
assert any("状态不正确" in e for e in errors)
def test_mixed_status_with_warnings(self):
"""混合状态时有 pending/failed 警告。"""
plan = FakePlan(id="plan_001", status="editing")
clips = [
_make_clip("c1", order=0, status="ready"),
_make_clip("c2", order=1, status="pending"),
_make_clip("c3", order=2, status="failed"),
]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
assert valid
assert any("pending" in w for w in warnings)
assert any("failed" in w for w in warnings)
assert ready_count == 1
assert total_count == 3
# ── render_plan 测试 ─────────────────────────────────────────────────────────
class TestRenderPlan:
def test_plan_not_found(self):
"""计划不存在时返回失败。"""
adapter, mock_plan_repo, _ = _make_adapter(plan=None)
mock_plan_repo.get.return_value = None
result = adapter.render_plan("plan_001")
assert not result.success
assert "不存在" in result.error_message
def test_no_ready_clips(self):
"""没有 ready 片段时返回失败。"""
plan = FakePlan(id="plan_001", status="editing")
clips = [_make_clip("c1", status="pending")]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
result = adapter.render_plan("plan_001")
assert not result.success
assert "没有可渲染" in result.error_message
assert result.clip_count == 0
@patch("video_processing.render_adapter.download_asset")
def test_all_assets_download_fail(self, mock_download):
"""所有素材下载失败时返回失败。"""
mock_download.return_value = False
plan = FakePlan(id="plan_001", status="editing")
clips = [_make_clip("c1", order=0, duration=5.0)]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
result = adapter.render_plan("plan_001")
assert not result.success
assert "素材下载失败" in result.error_message
@patch("video_processing.render_adapter.upload_to_oss")
@patch("video_processing.render_adapter.UnifiedRenderService")
@patch("video_processing.render_adapter.download_asset")
def test_successful_render(self, mock_download, mock_render_cls, mock_upload, tmp_path):
"""完整渲染流程成功。"""
# 素材下载成功
def _fake_download(asset_id, local_path):
local_path.parent.mkdir(parents=True, exist_ok=True)
local_path.write_bytes(b"fake video data")
return True
mock_download.side_effect = _fake_download
# 渲染成功
mock_render = MagicMock()
mock_render.render.return_value = MagicMock(
output_path=tmp_path / "output.mp4",
duration=10.0,
file_size=102400,
width=1280,
height=720,
)
mock_render_cls.return_value = mock_render
# 上传成功
mock_upload.return_value = "https://oss.example.com/rendered/plan_001/job_001.mp4"
plan = FakePlan(id="plan_001", status="editing")
clips = [
_make_clip("c1", order=0, duration=5.0),
_make_clip("c2", order=1, duration=5.0),
]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
result = adapter.render_plan(
"plan_001",
job_id="job_001",
work_dir=tmp_path / "work",
)
assert result.success
assert result.output_url.startswith("https://")
assert result.duration == 10.0
assert result.width == 1280
assert result.height == 720
assert result.clip_count == 2
# 验证 UnifiedRenderService 被正确调用
mock_render_cls.assert_called_once()
call_kwargs = mock_render_cls.call_args
assert call_kwargs.kwargs["plan"] is plan
assert len(call_kwargs.kwargs["clips"]) == 2
assert len(call_kwargs.kwargs["asset_path_map"]) == 2
@patch("video_processing.render_adapter.download_asset")
def test_progress_callback(self, mock_download, tmp_path):
"""进度回调被正确触发。"""
def _fake_download(asset_id, local_path):
local_path.parent.mkdir(parents=True, exist_ok=True)
local_path.write_bytes(b"fake data")
return True
mock_download.side_effect = _fake_download
# 模拟渲染异常,避免走到最后
with patch("video_processing.render_adapter.UnifiedRenderService") as mock_render_cls:
mock_render = MagicMock()
mock_render.render.side_effect = RuntimeError("render error")
mock_render_cls.return_value = mock_render
plan = FakePlan(id="plan_001", status="editing")
clips = [_make_clip("c1", order=0, duration=5.0)]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
progress_values = []
def progress_cb(progress: float, stage: str) -> None:
progress_values.append((progress, stage))
result = adapter.render_plan(
"plan_001",
work_dir=tmp_path / "work",
progress_cb=progress_cb,
)
# 即使渲染失败,前期进度也应该上报了
assert len(progress_values) > 0
# 第一个进度应该是加载计划
assert progress_values[0][1] == "加载剪辑计划"
@patch("video_processing.render_adapter.download_asset")
def test_partial_asset_download(self, mock_download, tmp_path):
"""部分素材下载失败时,只使用成功的素材。"""
download_results = [True, False, True] # 3个素材中2个成功
def _fake_download(asset_id, local_path):
idx = hash(asset_id) % 3
if download_results[idx]:
local_path.parent.mkdir(parents=True, exist_ok=True)
local_path.write_bytes(b"fake data")
return True
return False
mock_download.side_effect = _fake_download
with patch("video_processing.render_adapter.UnifiedRenderService") as mock_render_cls:
mock_render = MagicMock()
mock_render.render.return_value = MagicMock(
output_path=tmp_path / "out.mp4",
duration=5.0,
file_size=1024,
width=1280,
height=720,
)
mock_render_cls.return_value = mock_render
with patch("video_processing.render_adapter.upload_to_oss", return_value="https://example.com/out.mp4"):
plan = FakePlan(id="plan_001", status="editing")
clips = [
_make_clip("c1", order=0, duration=3.0, asset_id="asset_001.mp4"),
_make_clip("c2", order=1, duration=3.0, asset_id="asset_002.mp4"),
_make_clip("c3", order=2, duration=3.0, asset_id="asset_003.mp4"),
]
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
result = adapter.render_plan(
"plan_001",
work_dir=tmp_path / "work",
)
# 至少有部分素材成功,渲染应该进行
# (具体成功数量取决于 hash 结果,但至少1个成功就能渲染)
assert result.success or "素材下载失败" in result.error_message
# ── _download_assets 测试 ────────────────────────────────────────────────────
class TestDownloadAssets:
@patch("video_processing.render_adapter.download_asset")
def test_all_download_success(self, mock_download, tmp_path):
"""全部素材下载成功。"""
mock_download.return_value = True
clips = [
_make_clip("c1", order=0, asset_id="key1.mp4"),
_make_clip("c2", order=1, asset_id="key2.mp4"),
]
result = RenderAdapter._download_assets(clips, tmp_path)
assert len(result) == 2
assert "key1.mp4" in result
assert "key2.mp4" in result
assert mock_download.call_count == 2
@patch("video_processing.render_adapter.download_asset")
def test_empty_asset_id_skipped(self, mock_download, tmp_path):
"""空 asset_id 的片段被跳过。"""
clips = [
_make_clip("c1", order=0, asset_id=""),
_make_clip("c2", order=1, asset_id="key2.mp4"),
]
mock_download.return_value = True
result = RenderAdapter._download_assets(clips, tmp_path)
assert len(result) == 1
assert "key2.mp4" in result
assert mock_download.call_count == 1 # 只调用了一次下载
@patch("video_processing.render_adapter.download_asset")
def test_all_download_fail(self, mock_download, tmp_path):
"""全部下载失败返回空字典。"""
mock_download.return_value = False
clips = [
_make_clip("c1", order=0, asset_id="key1.mp4"),
]
result = RenderAdapter._download_assets(clips, tmp_path)
assert len(result) == 0
+351
View File
@@ -0,0 +1,351 @@
"""任务队列限流防护单元测试。"""
from __future__ import annotations
import os
import sys
from unittest.mock import MagicMock
import pytest
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "api"))
from app.core.task_enqueue import (
GLOBAL_PENDING_LIMIT,
USER_PENDING_LIMIT,
GlobalQueueFull,
UserPendingLimitExceeded,
check_queue_limits,
safe_enqueue_generation_task,
)
# ---------------------------------------------------------------------------
# Mock helpers
# ---------------------------------------------------------------------------
class MockRepository:
"""支持 pending 计数的 mock repository。
支持通过 set_pending 动态修改计数,用于模拟入队后计数变化的并发场景。
"""
def __init__(self, user_pending: int = 0, global_pending: int = 0):
self._user_pending = user_pending
self._global_pending = global_pending
self._send_task_called = False
self.updated_tasks = []
def count_pending_by_user(self, user_id: str) -> int:
return self._user_pending
def count_pending_total(self) -> int:
return self._global_pending
def update(self, task):
self.updated_tasks.append(task)
return task
def set_pending(self, *, user_pending: int | None = None, global_pending: int | None = None):
"""动态修改 pending 计数,模拟并发场景。"""
if user_pending is not None:
self._user_pending = user_pending
if global_pending is not None:
self._global_pending = global_pending
class MockTask:
def __init__(self, task_id: str = "task-1", status: str = "pending"):
self.id = task_id
self.status = status
self.error_message = ""
def mark_failed(self, reason: str):
self.status = "failed"
self.error_message = reason
@pytest.fixture(autouse=True)
def mock_celery(monkeypatch):
"""mock 掉 celery_app.send_task,避免真实发送。"""
mock_send = MagicMock()
monkeypatch.setattr("app.core.celery_app.celery_app.send_task", mock_send)
return mock_send
# ---------------------------------------------------------------------------
# 常量导出测试
# ---------------------------------------------------------------------------
def test_limit_constants_are_exported():
"""限流阈值常量已导出,供业务代码引用。"""
assert USER_PENDING_LIMIT == 3
assert GLOBAL_PENDING_LIMIT == 20
# ---------------------------------------------------------------------------
# check_queue_limits 单元测试(预检查用,>= 边界)
# ---------------------------------------------------------------------------
class TestCheckQueueLimits:
"""队列限流检查函数测试(预检查语义,>= 上限即拒绝)。"""
def test_normal_passes_through(self):
"""正常范围内的任务不受限制。"""
repo = MockRepository(user_pending=1, global_pending=5)
check_queue_limits("user-1", repo)
def test_user_limit_exceeded_raises(self):
"""用户 pending 超过上限抛 UserPendingLimitExceeded。"""
repo = MockRepository(user_pending=4, global_pending=5)
with pytest.raises(UserPendingLimitExceeded) as exc_info:
check_queue_limits("user-1", repo)
assert exc_info.value.user_id == "user-1"
assert exc_info.value.pending_count == 4
assert exc_info.value.limit == 3
def test_user_at_limit_also_raises(self):
"""用户 pending 刚好等于上限也拒绝(>= 边界)。"""
repo = MockRepository(user_pending=3, global_pending=5)
with pytest.raises(UserPendingLimitExceeded):
check_queue_limits("user-1", repo)
def test_user_below_limit_passes(self):
"""用户 pending 比上限少 1,通过。"""
repo = MockRepository(user_pending=2, global_pending=5)
check_queue_limits("user-1", repo)
def test_global_limit_exceeded_raises(self):
"""全局 pending 超过上限抛 GlobalQueueFull。"""
repo = MockRepository(user_pending=1, global_pending=21)
with pytest.raises(GlobalQueueFull) as exc_info:
check_queue_limits("user-1", repo)
assert exc_info.value.pending_count == 21
assert exc_info.value.limit == 20
def test_global_at_limit_also_raises(self):
"""全局 pending 刚好等于上限也拒绝(>= 边界)。"""
repo = MockRepository(user_pending=1, global_pending=20)
with pytest.raises(GlobalQueueFull):
check_queue_limits("user-1", repo)
def test_global_below_limit_passes(self):
"""全局 pending 比上限少 1,通过。"""
repo = MockRepository(user_pending=1, global_pending=19)
check_queue_limits("user-1", repo)
def test_global_takes_priority_over_user(self):
"""全局和用户都超限时,优先抛全局异常。"""
repo = MockRepository(user_pending=5, global_pending=25)
with pytest.raises(GlobalQueueFull):
check_queue_limits("user-1", repo)
def test_empty_user_id_skips_user_check(self):
"""不传 user_id 时跳过用户级检查,只做全局检查。"""
repo = MockRepository(user_pending=10, global_pending=5)
# 用户超限但不传 user_id → 全局未超限,应该通过
check_queue_limits("", repo)
# ---------------------------------------------------------------------------
# safe_enqueue_generation_task 限流集成测试(入队前用 >,包含当前任务)
# ---------------------------------------------------------------------------
class TestSafeEnqueueWithLimits:
"""安全入队函数的限流功能测试。"""
def test_normal_task_enqueues_successfully(self, mock_celery):
"""正常任务入队成功,返回 True。"""
repo = MockRepository(user_pending=0, global_pending=0)
task = MockTask("task-1")
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
assert result is True
mock_celery.assert_called_once_with("worker.generate_video", args=["task-1"])
assert len(repo.updated_tasks) == 0 # 成功不需要更新状态
def test_user_limit_rejected_with_failed_status(self, mock_celery):
"""用户超限:任务标记为 failed,抛 UserPendingLimitExceeded。"""
repo = MockRepository(user_pending=5, global_pending=5)
task = MockTask("task-1")
with pytest.raises(UserPendingLimitExceeded):
safe_enqueue_generation_task(task, repo, user_id="user-1")
mock_celery.assert_not_called()
assert task.status == "failed"
assert "限流" in task.error_message
assert len(repo.updated_tasks) == 1
def test_user_at_limit_still_passes(self, mock_celery):
"""用户 pending 刚好等于上限:入队前检查用 >,包含当前任务,刚好到上限不算超。
与预检查的 >= 语义一致:预检查时 pending=3 拒绝(不能再加新的),
但 safe_enqueue 被调用时任务已是 pending(就是第3个),
pending=3 不满足 >3,所以通过。
"""
repo = MockRepository(user_pending=3, global_pending=5)
task = MockTask("task-1")
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
assert result is True
mock_celery.assert_called_once()
def test_user_one_over_limit_rejected(self, mock_celery):
"""用户 pending = limit + 1:超限被拒。"""
repo = MockRepository(user_pending=4, global_pending=5)
task = MockTask("task-1")
with pytest.raises(UserPendingLimitExceeded):
safe_enqueue_generation_task(task, repo, user_id="user-1")
mock_celery.assert_not_called()
def test_global_limit_rejected_with_failed_status(self, mock_celery):
"""全局超限:任务标记为 failed,抛 GlobalQueueFull。"""
repo = MockRepository(user_pending=1, global_pending=21)
task = MockTask("task-1")
with pytest.raises(GlobalQueueFull):
safe_enqueue_generation_task(task, repo, user_id="user-1")
mock_celery.assert_not_called()
assert task.status == "failed"
assert len(repo.updated_tasks) == 1
def test_global_at_limit_still_passes(self, mock_celery):
"""全局 pending 刚好等于上限:入队前检查用 >,包含当前任务,刚好到上限不算超。"""
repo = MockRepository(user_pending=1, global_pending=20)
task = MockTask("task-1")
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
assert result is True
mock_celery.assert_called_once()
def test_no_user_id_skips_user_limit(self, mock_celery):
"""不传 user_id 时跳过用户级限流,只做全局检查。"""
repo = MockRepository(user_pending=10, global_pending=5)
task = MockTask("task-1")
result = safe_enqueue_generation_task(task, repo, user_id="")
assert result is True
mock_celery.assert_called_once()
def test_no_user_id_still_checks_global(self, mock_celery):
"""不传 user_id 时全局超限仍然被拦。"""
repo = MockRepository(user_pending=10, global_pending=25)
task = MockTask("task-1")
with pytest.raises(GlobalQueueFull):
safe_enqueue_generation_task(task, repo, user_id="")
mock_celery.assert_not_called()
def test_default_limits_match_constants(self, mock_celery):
"""默认配置与导出常量一致。"""
# 刚好在默认限制内(limit - 1)
repo = MockRepository(user_pending=2, global_pending=19)
task = MockTask("task-1")
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
assert result is True
def test_update_failure_does_not_crash(self, mock_celery):
"""repository.update 失败也不崩溃,异常继续向上抛。"""
class BadRepo(MockRepository):
def update(self, task):
raise RuntimeError("db down")
repo = BadRepo(user_pending=5, global_pending=5)
task = MockTask("task-1")
# 仍然抛 UserPendingLimitExceeded,不会被 update 失败掩盖
with pytest.raises(UserPendingLimitExceeded):
safe_enqueue_generation_task(task, repo, user_id="user-1")
mock_celery.assert_not_called()
# 任务状态还是变了(内存里改了)
assert task.status == "failed"
# ---------------------------------------------------------------------------
# 入队后最终校验(并发竞态兜底)测试
# ---------------------------------------------------------------------------
class TestPostEnqueueFinalCheck:
"""入队后最终校验:模拟并发场景,Celery发送后计数增加被兜住。"""
def test_post_enqueue_global_overflow_rollback(self, mock_celery):
"""并发场景:入队前检查通过,但发送Celery后全局计数超限 → 回滚为failed。
模拟两个请求同时通过入队前检查(都查到 global=19),
都创建了任务(DB里变成 21),先发送Celery的那个在最终校验时被兜住。
"""
repo = MockRepository(user_pending=1, global_pending=20) # 入队前:20 > 20?否
task = MockTask("task-1")
# 模拟发送Celery后,另一个并发请求也创建了任务,全局变成21
def side_effect(*args, **kwargs):
repo.set_pending(global_pending=21)
mock_celery.side_effect = side_effect
with pytest.raises(GlobalQueueFull) as exc_info:
safe_enqueue_generation_task(task, repo, user_id="user-1")
# Celery 确实发出去了(兜底不撤销 Celery,只回滚 DB 状态)
mock_celery.assert_called_once()
# 任务被标记为 failed
assert task.status == "failed"
assert "入队后" in task.error_message
assert exc_info.value.pending_count == 21
assert len(repo.updated_tasks) == 1
def test_post_enqueue_user_overflow_rollback(self, mock_celery):
"""并发场景:入队前检查通过,但发送Celery后用户计数超限 → 回滚为failed。"""
repo = MockRepository(user_pending=3, global_pending=5) # 入队前:3 > 3?否
task = MockTask("task-1")
def side_effect(*args, **kwargs):
repo.set_pending(user_pending=4)
mock_celery.side_effect = side_effect
with pytest.raises(UserPendingLimitExceeded) as exc_info:
safe_enqueue_generation_task(task, repo, user_id="user-1")
mock_celery.assert_called_once()
assert task.status == "failed"
assert "入队后" in task.error_message
assert exc_info.value.user_id == "user-1"
assert exc_info.value.pending_count == 4
def test_post_enqueue_global_priority_over_user(self, mock_celery):
"""入队后校验:全局和用户都超限时,优先抛全局异常。"""
repo = MockRepository(user_pending=3, global_pending=20)
task = MockTask("task-1")
def side_effect(*args, **kwargs):
repo.set_pending(user_pending=5, global_pending=22)
mock_celery.side_effect = side_effect
with pytest.raises(GlobalQueueFull):
safe_enqueue_generation_task(task, repo, user_id="user-1")
assert task.status == "failed"
def test_post_enqueue_no_change_still_passes(self, mock_celery):
"""入队后计数没变 → 正常通过,不回滚。"""
repo = MockRepository(user_pending=2, global_pending=10)
task = MockTask("task-1")
result = safe_enqueue_generation_task(task, repo, user_id="user-1")
assert result is True
mock_celery.assert_called_once()
assert task.status == "pending" # 状态没变
assert len(repo.updated_tasks) == 0 # 没更新 DB
def test_post_enqueue_no_user_id_skips_user_check(self, mock_celery):
"""不传 user_id 时,入队后校验也跳过用户级,只查全局。"""
repo = MockRepository(user_pending=10, global_pending=5)
task = MockTask("task-1")
def side_effect(*args, **kwargs):
repo.set_pending(user_pending=15, global_pending=5) # 用户超限但全局没超
mock_celery.side_effect = side_effect
result = safe_enqueue_generation_task(task, repo, user_id="")
assert result is True # 用户级不检查,全局没超限 → 通过
File diff suppressed because it is too large Load Diff
-247
View File
@@ -1,247 +0,0 @@
"""
视频合成服务安全校验单元测试
针对 PR #159 安全审计发现的问题进行测试
"""
import os
import tempfile
from unittest.mock import MagicMock, patch
import pytest
# 导入被测试的模块
from apps.worker.video_processing.video_compose_service import (
ALLOWED_INPUT_PREFIXES,
ALLOWED_OUTPUT_DIRS,
ALLOWED_TRANSITIONS,
Clip,
EditingMode,
EditingModeConfig,
VideoComposeService,
)
class TestOutputPathValidation:
"""P0: 输出路径穿越校验测试"""
def setup_method(self):
"""测试前设置"""
self.config = EditingModeConfig(mode=EditingMode.ONE_TAKE)
self.service = VideoComposeService(self.config)
def test_valid_output_path_in_allowed_dir(self):
"""测试合法的输出路径"""
valid_path = "/tmp/video_output/test.mp4"
result = self.service._validate_output_path(valid_path)
assert result == os.path.abspath(valid_path)
def test_valid_output_path_with_relative_components(self):
"""测试带相对路径成分但最终在允许目录内的路径"""
valid_path = "/tmp/video_output/subdir/../test.mp4"
result = self.service._validate_output_path(valid_path)
assert result == os.path.abspath(valid_path)
def test_path_traversal_attack_blocked(self):
"""测试路径穿越攻击被阻止"""
# 尝试穿越到 /etc/passwd
malicious_path = "/tmp/video_output/../../../etc/passwd"
with pytest.raises(ValueError, match="输出路径不在允许范围内"):
self.service._validate_output_path(malicious_path)
def test_path_traversal_attack_blocked_var_app(self):
"""测试针对 /var/app 的路径穿越攻击被阻止"""
malicious_path = "/var/app/rendered/../../config/../../../etc/passwd"
with pytest.raises(ValueError, match="输出路径不在允许范围内"):
self.service._validate_output_path(malicious_path)
def test_absolute_path_to_forbidden_location(self):
"""测试直接访问禁止位置"""
forbidden_path = "/etc/shadow"
with pytest.raises(ValueError, match="输出路径不在允许范围内"):
self.service._validate_output_path(forbidden_path)
def test_root_path_blocked(self):
"""测试根目录被阻止"""
root_path = "/"
with pytest.raises(ValueError, match="输出路径不在允许范围内"):
self.service._validate_output_path(root_path)
def test_absolute_path_to_tmp_not_allowed(self):
"""测试 /tmp 不在白名单中时应被阻止"""
# /tmp 不在 ALLOWED_OUTPUT_DIRS 中
tmp_path = "/tmp/test.mp4"
with pytest.raises(ValueError, match="输出路径不在允许范围内"):
self.service._validate_output_path(tmp_path)
class TestInputPathValidation:
"""P1-1: 输入路径格式校验测试"""
def setup_method(self):
"""测试前设置"""
self.config = EditingModeConfig(mode=EditingMode.ONE_TAKE)
self.service = VideoComposeService(self.config)
def test_valid_s3_path(self):
"""测试 S3 路径"""
assert self.service._validate_input_path("s3://bucket/key.mp4") is True
def test_valid_oss_path(self):
"""测试 OSS 路径"""
assert self.service._validate_input_path("oss://bucket/key.mp4") is True
def test_valid_local_path(self):
"""测试 local:// 路径"""
assert self.service._validate_input_path("local://asset/123.mp4") is True
def test_valid_var_storage_path(self):
"""测试 /var/storage/ 路径"""
assert self.service._validate_input_path("/var/storage/assets/123.mp4") is True
def test_path_traversal_in_input_rejected(self):
"""测试输入路径中的路径穿越尝试被拒绝"""
malicious_path = "s3://bucket/../../etc/passwd"
# 这会通过前缀检查,但实际使用时文件系统访问会失败
# 安全设计:只校验格式前缀
assert self.service._validate_input_path(malicious_path) is True
def test_malicious_input_path_blocked(self):
"""测试恶意输入路径被阻止"""
assert self.service._validate_input_path("/etc/passwd") is False
assert self.service._validate_input_path("file:///etc/passwd") is False
assert self.service._validate_input_path("http://evil.com/shell.sh") is False
def test_empty_path_rejected(self):
"""测试空路径被拒绝"""
assert self.service._validate_input_path("") is False
def test_random_string_rejected(self):
"""测试随机字符串被拒绝"""
assert self.service._validate_input_path("random123") is False
assert self.service._validate_input_path("abc../../../etc") is False
class TestTransitionValidation:
"""P1-2: 转场参数白名单校验测试"""
def setup_method(self):
"""测试前设置"""
self.config = EditingModeConfig(mode=EditingMode.ONE_TAKE)
self.service = VideoComposeService(self.config)
@pytest.mark.parametrize("transition", list(ALLOWED_TRANSITIONS))
def test_valid_transitions(self, transition):
"""测试所有合法的转场效果"""
result = self.service._validate_transition(transition)
assert result == transition
def test_invalid_transition_defaults_to_fade(self):
"""测试非法转场效果默认为 fade"""
result = self.service._validate_transition("random_transition")
assert result == "fade"
def test_sql_injection_in_transition_blocked(self):
"""测试 SQL 注入尝试被阻止"""
result = self.service._validate_transition("fade; DROP TABLE videos;--")
assert result == "fade"
def test_shell_injection_in_transition_blocked(self):
"""测试 Shell 注入尝试被阻止"""
result = self.service._validate_transition("fade$(whoami)")
assert result == "fade"
def test_empty_transition_handled(self):
"""测试空转场名称"""
result = self.service._validate_transition("")
assert result == "fade"
def test_none_transition_handled(self):
"""测试 None 转场名称"""
result = self.service._validate_transition(None)
assert result == "fade"
def test_get_validated_transition_returns_mapped(self):
"""测试 _get_validated_transition 返回映射后的值"""
# "fade" 应该映射为 "fade"
result = self.service._get_validated_transition("fade")
assert result == "fade"
class TestComposeSecurityIntegration:
"""安全集成测试"""
def setup_method(self):
"""测试前设置"""
self.config = EditingModeConfig(mode=EditingMode.ONE_TAKE)
self.service = VideoComposeService(self.config)
def test_compose_rejects_malicious_output_path(self):
"""测试 compose 方法拒绝恶意输出路径"""
clips = [
Clip(asset_id="s3://bucket/video1.mp4"),
Clip(asset_id="s3://bucket/video2.mp4"),
]
with pytest.raises(ValueError, match="输出路径不在允许范围内"):
self.service.compose(clips, output_path="/etc/passwd")
def test_compose_rejects_invalid_input_path(self):
"""测试 compose 方法拒绝非法输入路径"""
clips = [
Clip(asset_id="/etc/shadow"), # 非法路径
]
with pytest.raises(ValueError, match="不合法的输入路径"):
self.service.compose(clips)
def test_compose_with_valid_paths(self):
"""测试合法路径可以正常处理"""
with tempfile.TemporaryDirectory() as tmpdir:
# 创建临时视频文件
video_path = os.path.join(tmpdir, "input.mp4")
output_path = os.path.join("/tmp/video_output", "output.mp4")
# 创建空的测试文件(实际测试需要真实视频)
with open(video_path, "wb") as f:
f.write(b"fake video data")
clips = [
Clip(asset_id=f"local://{video_path}"),
]
# 验证输入校验通过
assert self.service._validate_input_path(f"local://{video_path}") is True
def test_compose_empty_clips_rejected(self):
"""测试空片段列表被拒绝"""
with pytest.raises(ValueError, match="clips 不能为空"):
self.service.compose([])
class TestWhiteListConstants:
"""白名单常量测试"""
def test_allowed_output_dirs_not_empty(self):
"""测试输出目录白名单不为空"""
assert len(ALLOWED_OUTPUT_DIRS) > 0
assert "/tmp/video_output" in ALLOWED_OUTPUT_DIRS
assert "/var/app/rendered" in ALLOWED_OUTPUT_DIRS
def test_allowed_input_prefixes_not_empty(self):
"""测试输入路径前缀白名单不为空"""
assert len(ALLOWED_INPUT_PREFIXES) > 0
assert "s3://" in ALLOWED_INPUT_PREFIXES
assert "oss://" in ALLOWED_INPUT_PREFIXES
assert "local://" in ALLOWED_INPUT_PREFIXES
assert "/var/storage/" in ALLOWED_INPUT_PREFIXES
def test_allowed_transitions_not_empty(self):
"""测试转场效果白名单不为空"""
assert len(ALLOWED_TRANSITIONS) > 0
assert "fade" in ALLOWED_TRANSITIONS
assert "dissolve" in ALLOWED_TRANSITIONS
assert "slideleft" in ALLOWED_TRANSITIONS
if __name__ == "__main__":
pytest.main([__file__, "-v"])