Compare commits

...

22 Commits

Author SHA1 Message Date
xiaoxia-bot 92a7915de0 feat(ci): 新增发布与灰度部署脚本
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Successful in 30s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 1m6s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 1m11s
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 1m22s
- release.sh: 一键发布脚本(打tag + 生成changelog + 灰度部署)
- gray_deploy.sh: 灰度发布脚本(Nginx权重调整 + canary容器)
- rollback.sh: 灰度回滚脚本(切回稳定版本流量)
2026-07-14 11:33:41 +08:00
CI Bot ca2e044246 feat: 视频调速引擎(快进/慢放)
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Successful in 46s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 1m20s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 1m12s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 3m42s
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
- SpeedEngine 抽象层:SpeedConfig + SpeedEngine,支持 0.25x~4x 变速
- 视频调速:基于 FFmpeg setpts,在 clip 预处理环节插入
- 音频调速:基于 atempo,超范围自动多级串联(如 4x=2.0*2.0)
- 音画同步:视频音频同时调速,音调自动修正
- 分段调速:每个 clip 独立 playback_speed 字段
- 整段调速:所有 clip 设相同 speed 即可实现
- 降级策略:速度超出范围自动钳制,不阻断渲染
- 领域模型:EditPlanClip 新增 playback_speed 字段
- 数据库:EditPlanClipModel 新增 playback_speed 列
- API:create_clip / update_clip 支持 playback_speed 参数
- 45 个新增单测全绿 + 100 个现有测试全绿,无回归
2026-07-14 11:32:26 +08:00
xiaoxia 8a2d2df3cd feat: 视频封面生成 + 视频倒放 + 贴纸叠加三个渲染能力
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 30s
CI/CD Pipeline / Frontend Lint (push) Successful in 1m12s
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 1m20s
CI/CD Pipeline / Unit Tests (push) Successful in 4m17s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 17m1s
CI/CD Pipeline / Staging E2E Tests (push) Failing after 48s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 3m16s
Squash merge PR #305
2026-07-14 11:05:04 +08:00
xiaoxia ff38ee0f2b feat: 转场特效引擎 — 14种转场预设 + TransitionEngine抽象 + 降级策略
CI/CD Pipeline / Validate Code Quality And Tests (push) Has been cancelled
CI/CD Pipeline / Unit Tests (push) Has been cancelled
CI/CD Pipeline / Integration Tests (push) Has been cancelled
CI/CD Pipeline / Frontend Lint (push) Has been cancelled
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
CI/CD Pipeline / Build Production Runtime Images (push) Has been cancelled
CI/CD Pipeline / Deploy Production (push) Has been cancelled
CI/CD Pipeline / Production Browser E2E (push) Has been cancelled
Squash merge PR #293
2026-07-14 11:03:56 +08:00
xiaoxia 368baf683b fix: StubGenerationTaskRepository补list_by_user_filtered方法 (#306)
CI/CD Pipeline / Validate Code Quality And Tests (push) Has been cancelled
CI/CD Pipeline / Unit Tests (push) Has been cancelled
CI/CD Pipeline / Integration Tests (push) Has been cancelled
CI/CD Pipeline / Frontend Lint (push) Has been cancelled
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
CI/CD Pipeline / Build Production Runtime Images (push) Has been cancelled
CI/CD Pipeline / Deploy Production (push) Has been cancelled
CI/CD Pipeline / Production Browser E2E (push) Has been cancelled
2026-07-14 11:02:57 +08:00
xiaoxia 527eb61f19 feat: 视频裁剪/分割能力(Trimming Engine)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 35s
CI/CD Pipeline / Unit Tests (push) Successful in 1m11s
CI/CD Pipeline / Integration Tests (push) Failing after 1m12s
CI/CD Pipeline / Frontend Lint (push) Successful in 3m40s
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) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
Squash merge PR #296
2026-07-14 10:55:58 +08:00
xiaoxia e9a6d19e00 feat: 绿幕抠像 + 音频降噪引擎(Chroma Key + Noise Reduction) (#303)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 39s
CI/CD Pipeline / Unit Tests (push) Successful in 1m22s
CI/CD Pipeline / Frontend Lint (push) Successful in 1m22s
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) Failing after 1m13s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
2026-07-14 10:52:01 +08:00
xiaoxia 241760ef39 feat: 水印 + 片头片尾引擎(视频包装能力) (#298)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 37s
CI/CD Pipeline / Frontend Lint (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 / Unit Tests (push) Successful in 4m21s
CI/CD Pipeline / Integration Tests (push) Failing after 4m25s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
2026-07-14 10:44:19 +08:00
xiaoxia 3cd26e98db feat: BGM音轨混音能力(音量/淡入淡出/人声闪避/预设BGM库) (#291)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 43s
CI/CD Pipeline / Unit Tests (push) Successful in 1m35s
CI/CD Pipeline / Frontend Lint (push) Successful in 1m37s
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) Has been cancelled
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
2026-07-14 10:37:24 +08:00
xiaoxia c840f37a44 feat: TTS文字转语音配音引擎 (#295)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 33s
CI/CD Pipeline / Frontend Lint (push) Successful in 1m12s
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) Failing after 1m8s
CI/CD Pipeline / Unit Tests (push) Successful in 2m52s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
2026-07-14 10:21:49 +08:00
xiaoxia edcd1a926f feat: 画中画(PiP)能力 - 多图层叠加引擎 (#299)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 42s
CI/CD Pipeline / Unit Tests (push) Successful in 1m13s
CI/CD Pipeline / Frontend Lint (push) Successful in 1m19s
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) Failing after 1m29s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
2026-07-14 10:17:40 +08:00
xiaoxia 9ddaaf7f00 fix: 修复构建脚本末尾grep导致set -e失败 (#304)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 37s
CI/CD Pipeline / Unit Tests (push) Successful in 1m4s
CI/CD Pipeline / Frontend Lint (push) Successful in 1m22s
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) Failing after 1m9s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 16m21s
CI/CD Pipeline / Staging E2E Tests (push) Successful in 1m52s
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
2026-07-14 09:55:59 +08:00
xiaoxia 94aead4342 feat: 任务中心升级(失败重试/错误追踪/列表筛选) (#289)
CI/CD Pipeline / Validate Code Quality And Tests (push) Has been cancelled
CI/CD Pipeline / Unit Tests (push) Has been cancelled
CI/CD Pipeline / Integration Tests (push) Has been cancelled
CI/CD Pipeline / Frontend Lint (push) Has been cancelled
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
CI/CD Pipeline / Build Production Runtime Images (push) Has been cancelled
CI/CD Pipeline / Deploy Production (push) Has been cancelled
CI/CD Pipeline / Production Browser E2E (push) Has been cancelled
2026-07-14 09:55:53 +08:00
xiaoxia 295d7f0765 feat: 素材智能视图筛选 + 标题使用次数闭环 (#282)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 52s
CI/CD Pipeline / Unit Tests (push) Successful in 1m31s
CI/CD Pipeline / Frontend Lint (push) Successful in 1m41s
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) Has been cancelled
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
2026-07-14 09:53:07 +08:00
xiaoxia 7cffb193eb feat: 模板与剪辑计划后端补齐(复制/筛选/标签/使用统计) (#288)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 36s
CI/CD Pipeline / Unit Tests (push) Successful in 1m8s
CI/CD Pipeline / Frontend Lint (push) Successful in 1m19s
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 1m20s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
feat: 模板与剪辑计划后端补齐
2026-07-14 09:51:05 +08:00
xiaoxia e430d83f78 feat: 滤镜调色引擎 - 8种预设 + 基础调色 + 分段应用 (#300)
CI/CD Pipeline / Validate Code Quality And Tests (push) Has been cancelled
CI/CD Pipeline / Unit Tests (push) Has been cancelled
CI/CD Pipeline / Integration Tests (push) Has been cancelled
CI/CD Pipeline / Frontend Lint (push) Has been cancelled
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
CI/CD Pipeline / Build Production Runtime Images (push) Has been cancelled
CI/CD Pipeline / Deploy Production (push) Has been cancelled
CI/CD Pipeline / Production Browser E2E (push) Has been cancelled
feat: 滤镜调色引擎 - 8种预设 + 基础调色
2026-07-14 09:51:03 +08:00
xiaoxia f4b4f1fc4f feat: ASR自动字幕能力(领域模型+渲染管道接入+可扩展ASR后端) (#292)
CI/CD Pipeline / Validate Code Quality And Tests (push) Has been cancelled
CI/CD Pipeline / Unit Tests (push) Has been cancelled
CI/CD Pipeline / Integration Tests (push) Has been cancelled
CI/CD Pipeline / Frontend Lint (push) Has been cancelled
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
CI/CD Pipeline / Build Production Runtime Images (push) Has been cancelled
CI/CD Pipeline / Deploy Production (push) Has been cancelled
CI/CD Pipeline / Production Browser E2E (push) Has been cancelled
feat: ASR自动字幕能力
2026-07-14 09:50:57 +08:00
xiaoxia 58ff565c48 feat: 素材批量操作接口(软删除/打标签/改分类/智能视图标记) (#290)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 42s
CI/CD Pipeline / Unit Tests (push) Successful in 1m18s
CI/CD Pipeline / Integration Tests (push) Has been cancelled
CI/CD Pipeline / Frontend Lint (push) Has been cancelled
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
CI/CD Pipeline / Build Production Runtime Images (push) Has been cancelled
CI/CD Pipeline / Deploy Production (push) Has been cancelled
CI/CD Pipeline / Production Browser E2E (push) Has been cancelled
feat: 素材批量操作接口(软删除/打标签/改分类/智能视图标记)
2026-07-14 09:49:35 +08:00
xiaoxia eb4645314d feat: 成片中心后端升级(封面生成/复核/批量下载) (#287)
CI/CD Pipeline / Validate Code Quality And Tests (push) Has been cancelled
CI/CD Pipeline / Unit Tests (push) Has been cancelled
CI/CD Pipeline / Integration Tests (push) Has been cancelled
CI/CD Pipeline / Frontend Lint (push) Has been cancelled
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
CI/CD Pipeline / Build Production Runtime Images (push) Has been cancelled
CI/CD Pipeline / Deploy Production (push) Has been cancelled
CI/CD Pipeline / Production Browser E2E (push) Has been cancelled
feat: 成片中心后端升级(封面生成/复核/批量下载)
2026-07-14 09:49:11 +08:00
xiaoxia 17fbae13a8 fix(ci): 移除构建脚本中docker driver不支持的--cache-to导出 (#302)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 26s
CI/CD Pipeline / Frontend Lint (push) Successful in 52s
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 / Unit Tests (push) Successful in 56s
CI/CD Pipeline / Integration Tests (push) Successful in 1m8s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Failing after 14m27s
CI/CD Pipeline / Staging E2E Tests (push) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (push) Has been skipped
fix(ci): 移除构建脚本中docker driver不支持的--cache-to导出
2026-07-14 09:04:18 +08:00
xiaoxia 41e421b44b fix(e2e): Playwright chromium禁用GPU,修复无显示环境下浏览器不稳定问题 (#301)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 25s
CI/CD Pipeline / Unit Tests (push) Successful in 48s
CI/CD Pipeline / Frontend Lint (push) Successful in 57s
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 3s
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 1m17s
fix(e2e): Playwright chromium禁用GPU,修复无显示环境下浏览器不稳定问题
2026-07-14 08:48:42 +08:00
xiaoxia 9f86bd40ca fix(ci): 修复 build_release_images.sh 中 CACHE_TAG 未定义的问题 (#297)
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 34s
CI/CD Pipeline / Unit Tests (push) Successful in 49s
CI/CD Pipeline / Frontend Lint (push) Successful in 49s
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 1m10s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 41m58s
CI/CD Pipeline / Staging E2E Tests (push) Failing after 19m9s
CI/CD Pipeline / Staging API Integration Tests (push) Successful in 21m26s
2026-07-14 07:22:05 +08:00
111 changed files with 16838 additions and 226 deletions
@@ -0,0 +1,47 @@
"""add error_info and retry fields to generation_tasks
Revision ID: 038_error_retry
Revises: 037_generation_logs
Create Date: 2026-07-13 22:15:00.000000
"""
import sqlalchemy as sa
from sqlalchemy.dialects.mysql import JSON as MySQLJSON
from alembic import op
# revision identifiers, used by Alembic.
revision = "038_error_retry"
down_revision = "037_generation_logs"
branch_labels = None
depends_on = None
def upgrade():
# error_info: 结构化错误信息(error_type, message, stack_trace, failed_at, stage等)
op.add_column(
"generation_tasks",
sa.Column("error_info", sa.JSON(), nullable=True),
)
# retry_count: 重试次数
op.add_column(
"generation_tasks",
sa.Column("retry_count", sa.Integer(), nullable=False, server_default="0"),
)
# auto_retry_enabled: 是否开启自动重试
op.add_column(
"generation_tasks",
sa.Column("auto_retry_enabled", sa.Boolean(), nullable=False, server_default=sa.text("false")),
)
# auto_retry_max: 最大自动重试次数
op.add_column(
"generation_tasks",
sa.Column("auto_retry_max", sa.Integer(), nullable=False, server_default="0"),
)
def downgrade():
op.drop_column("generation_tasks", "auto_retry_max")
op.drop_column("generation_tasks", "auto_retry_enabled")
op.drop_column("generation_tasks", "retry_count")
op.drop_column("generation_tasks", "error_info")
@@ -0,0 +1,34 @@
"""add transition_duration to edit_plan_clips
Revision ID: 039_transition_duration
Revises: 038_error_retry
Create Date: 2026-07-14 09:00:00.000000
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "039_transition_duration"
down_revision = "038_error_retry"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"edit_plan_clips",
sa.Column(
"transition_duration",
sa.Float(),
nullable=False,
server_default="0.0",
),
)
def downgrade() -> None:
op.drop_column("edit_plan_clips", "transition_duration")
@@ -0,0 +1,29 @@
"""add playback_speed to edit_plan_clips
Revision ID: 040_playback_speed
Revises: 039_transition_duration
Create Date: 2026-07-14 10:00:00.000000
"""
from __future__ import annotations
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = "040_playback_speed"
down_revision = "039_transition_duration"
branch_labels = None
depends_on = None
def upgrade() -> None:
op.add_column(
"edit_plan_clips",
sa.Column("playback_speed", sa.Float(), nullable=False, server_default="1.0"),
)
def downgrade() -> None:
op.drop_column("edit_plan_clips", "playback_speed")
+5
View File
@@ -19,6 +19,7 @@ from app.api.routes.templates import router as templates_router
from app.api.routes.titles import router as titles_router
from app.api.routes.tts import router as tts_router
from app.api.routes.upload import router as upload_router
from app.api.routes.videos import router as videos_router
from app.api.routes.voice_clones import router as voice_clones_router
from app.api.routes.voices import router as voices_router
from fastapi import APIRouter
@@ -99,6 +100,10 @@ api_router.include_router(
prefix="/voice-clones",
tags=["VoiceClone"],
)
api_router.include_router(
videos_router,
tags=["VideoCenter"],
)
api_router.include_router(
duplication_router,
prefix="/duplication",
+3 -4
View File
@@ -163,11 +163,10 @@ def delete_asset_library(
# 权限校验:检查用户是否有项目访问权限
check_project_access(library.project_id, authenticated_user.user.id, project_repository)
# 删除库内所有素材(无 FK 级联,需手动清理
# 删除库内所有素材(硬删除,素材库已删除,无需保留软删除状态
assets_in_library = asset_repository.find_by_library(library_id)
if assets_in_library:
asset_ids_to_delete = [a.id for a in assets_in_library]
asset_repository.batch_delete(asset_ids_to_delete)
for asset in assets_in_library:
asset_repository.delete(asset.id)
# 删除素材库本身
asset_library_repository.delete(library_id)
+181 -16
View File
@@ -12,8 +12,11 @@ from app.dependencies import (
)
from app.schemas.asset import (
AssetResponse,
BatchClassifyRequest,
BatchDeleteRequest,
BatchDeleteResponse,
BatchMarkRequest,
BatchOperationResponse,
BatchTagRequest,
CreateAssetRequest,
ListAssetsResponse,
UpdateAssetRequest,
@@ -82,6 +85,15 @@ def list_assets(
gender: Optional[str] = Query(None, description="按 metadata.gender 筛选"),
style: Optional[str] = Query(None, description="按 metadata.style 筛选"),
tag_ids: Optional[str] = Query(None, description="按标签 ID 筛选(逗号分隔,取交集)"),
smart_view: Optional[str] = Query(
None,
description="智能视图筛选:recommended=推荐(质量分≥80)、cautious=慎用(60-79)、risky=高风险(<60或已驳回)、unused=未使用、used=已使用、pending_review=待复核",
pattern="^(recommended|cautious|risky|unused|used|pending_review)$",
),
classification: Optional[str] = Query(
None,
description="按内容分类筛选:scenic=风景、product=产品、person=人物、animal=动物、food=美食、tech=科技、sport=运动、music=音乐、other=其他",
),
skip: int = Query(0, ge=0),
limit: int = Query(100, ge=1, le=500),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
@@ -101,11 +113,11 @@ def list_assets(
if not filter_tag_ids:
filter_tag_ids = None
# 需要内存过滤的标志(keyword/gender/style/tag_ids 无法在 DB 层过滤)
needs_memory_filter = bool(keyword or gender or style or filter_tag_ids)
# 需要内存过滤的标志(keyword/gender/style/tag_ids/smart_view/classification 无法在 DB 层过滤)
needs_memory_filter = bool(keyword or gender or style or filter_tag_ids or smart_view or classification)
def _apply_memory_filters(items):
"""应用 keyword / gender / style / tag_ids 内存过滤。"""
"""应用 keyword / gender / style / tag_ids / smart_view / classification 内存过滤。"""
result = items
if keyword:
kw = keyword.lower()
@@ -114,9 +126,38 @@ def list_assets(
result = [i for i in result if (i.metadata or {}).get("gender") == gender]
if style:
result = [i for i in result if (i.metadata or {}).get("style") == style]
if classification:
result = [i for i in result if (i.metadata or {}).get("classification") == classification]
if filter_tag_ids:
tag_set = set(filter_tag_ids)
result = [i for i in result if tag_set.issubset(set(getattr(i, "tag_ids", [])))]
if smart_view:
def __meta(a):
return a.metadata or {}
def __use_count(a):
return int(__meta(a).get("generation_use_count") or 0)
def __review_status(a):
return __meta(a).get("review_status", "")
if smart_view == "recommended":
result = [i for i in result if i.quality_score is not None and i.quality_score >= 80]
elif smart_view == "cautious":
result = [i for i in result if i.quality_score is not None and 60 <= i.quality_score < 80]
elif smart_view == "risky":
result = [
i
for i in result
if (i.quality_score is not None and i.quality_score < 60) or __review_status(i) == "rejected"
]
elif smart_view == "unused":
result = [i for i in result if __use_count(i) == 0]
elif smart_view == "used":
result = [i for i in result if __use_count(i) > 0]
elif smart_view == "pending_review":
result = [i for i in result if __review_status(i) == "pending_review"]
return result
# ── 优化路径:无内存过滤时,使用 DB 级分页 ──
@@ -260,33 +301,157 @@ def update_asset_review_status(
return _to_asset_response(updated)
@router.post("/batch-delete", response_model=BatchDeleteResponse)
@router.post("/batch-delete", response_model=BatchOperationResponse)
def batch_delete_assets(
request: BatchDeleteRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
asset_repository: Any = Depends(get_asset_repository),
project_repository: Any = Depends(get_project_repository),
) -> BatchDeleteResponse:
"""批量删除素材(配音素材等),需逐项校验项目权限。"""
) -> BatchOperationResponse:
"""批量删除素材(软删除,标记 status=deleted),需逐项校验项目权限。"""
user_id = authenticated_user.user.id
deleted_ids: list[str] = []
failed_ids: list[str] = []
success_ids: list[str] = []
failed_details: dict[str, str] = {}
for asset_id in request.ids:
for asset_id in request.asset_ids:
item = asset_repository.find_by_id(asset_id)
if item is None:
failed_ids.append(asset_id)
failed_details[asset_id] = "not_found"
continue
try:
check_project_access(item.project_id, user_id, project_repository)
deleted_ids.append(asset_id)
success_ids.append(asset_id)
except HTTPException:
failed_ids.append(asset_id)
failed_details[asset_id] = "access_denied"
if deleted_ids:
asset_repository.batch_delete(deleted_ids)
if success_ids:
asset_repository.batch_delete(success_ids)
return BatchDeleteResponse(deleted_count=len(deleted_ids), failed_ids=failed_ids)
return BatchOperationResponse(
success_count=len(success_ids),
failed_ids=list(failed_details.keys()),
failed_details=failed_details,
)
@router.post("/batch-tag", response_model=BatchOperationResponse)
def batch_tag_assets(
request: BatchTagRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
asset_repository: Any = Depends(get_asset_repository),
project_repository: Any = Depends(get_project_repository),
tag_repository: Any = Depends(get_tag_repository),
) -> BatchOperationResponse:
"""批量打标签(添加或替换模式),需逐项校验项目权限和标签权限。"""
user_id = authenticated_user.user.id
success_ids: list[str] = []
failed_details: dict[str, str] = {}
# 校验标签存在且属于当前用户
for tag_id in request.tag_ids:
tag = tag_repository.get(tag_id)
if tag is None:
return BatchOperationResponse(
success_count=0,
failed_ids=list(request.asset_ids),
failed_details={aid: f"tag_not_found:{tag_id}" for aid in request.asset_ids},
)
if tag.user_id != user_id:
return BatchOperationResponse(
success_count=0,
failed_ids=list(request.asset_ids),
failed_details={aid: f"tag_access_denied:{tag_id}" for aid in request.asset_ids},
)
# 校验素材权限
for asset_id in request.asset_ids:
item = asset_repository.find_by_id(asset_id)
if item is None:
failed_details[asset_id] = "not_found"
continue
try:
check_project_access(item.project_id, user_id, project_repository)
success_ids.append(asset_id)
except HTTPException:
failed_details[asset_id] = "access_denied"
if success_ids:
if request.mode == "replace":
asset_repository.batch_replace_tags(success_ids, request.tag_ids)
else:
asset_repository.batch_add_tags(success_ids, request.tag_ids)
return BatchOperationResponse(
success_count=len(success_ids),
failed_ids=list(failed_details.keys()),
failed_details=failed_details,
)
@router.post("/batch-classify", response_model=BatchOperationResponse)
def batch_classify_assets(
request: BatchClassifyRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
asset_repository: Any = Depends(get_asset_repository),
project_repository: Any = Depends(get_project_repository),
) -> BatchOperationResponse:
"""批量修改素材内容分类(person/scenic/product等),存在metadata.category中。"""
user_id = authenticated_user.user.id
success_ids: list[str] = []
failed_details: dict[str, str] = {}
for asset_id in request.asset_ids:
item = asset_repository.find_by_id(asset_id)
if item is None:
failed_details[asset_id] = "not_found"
continue
try:
check_project_access(item.project_id, user_id, project_repository)
success_ids.append(asset_id)
except HTTPException:
failed_details[asset_id] = "access_denied"
if success_ids:
asset_repository.batch_update_metadata(success_ids, {"category": request.category})
return BatchOperationResponse(
success_count=len(success_ids),
failed_ids=list(failed_details.keys()),
failed_details=failed_details,
)
@router.post("/batch-mark", response_model=BatchOperationResponse)
def batch_mark_assets(
request: BatchMarkRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
asset_repository: Any = Depends(get_asset_repository),
project_repository: Any = Depends(get_project_repository),
) -> BatchOperationResponse:
"""批量设置智能视图标记(recommended/caution/high_risk),存在metadata.smart_view中。"""
user_id = authenticated_user.user.id
success_ids: list[str] = []
failed_details: dict[str, str] = {}
for asset_id in request.asset_ids:
item = asset_repository.find_by_id(asset_id)
if item is None:
failed_details[asset_id] = "not_found"
continue
try:
check_project_access(item.project_id, user_id, project_repository)
success_ids.append(asset_id)
except HTTPException:
failed_details[asset_id] = "access_denied"
if success_ids:
asset_repository.batch_update_metadata(success_ids, {"smart_view": request.smart_view})
return BatchOperationResponse(
success_count=len(success_ids),
failed_ids=list(failed_details.keys()),
failed_details=failed_details,
)
@router.get("/{asset_id}", response_model=AssetResponse)
+3
View File
@@ -146,6 +146,7 @@ class AIRecommendClipItem(BaseModel):
text_content: str = Field(default="", description="文字内容")
duration: float = Field(..., ge=0.0, description="片段时长(秒)")
transition_effect: str = Field(default="cut", description="转场效果")
transition_duration: float = Field(default=0.0, ge=0.0, description="转场时长(秒),0 表示使用默认值")
asset_id: str = Field(default="", description="关联素材 ID")
start_time: float = Field(default=0.0, ge=0.0, description="素材截取起始时间(秒)")
config: dict[str, Any] = Field(default_factory=dict, description="片段额外配置")
@@ -209,6 +210,8 @@ class _PlanClipItem(BaseModel):
start_time: float
duration: float
transition_effect: str
transition_duration: float
playback_speed: float = 1.0
status: str
config: Optional[dict[str, Any]] = None
created_at: datetime
+1
View File
@@ -210,6 +210,7 @@ def generate_from_template(
start_time=c.start_time,
duration=c.duration,
transition_effect=c.transition_effect,
transition_duration=c.transition_duration,
status=c.status.value if hasattr(c.status, "value") else c.status,
config=c.config,
created_at=c.created_at,
@@ -267,6 +267,8 @@ def create_generation_task(
source_edit_plan_id=request.source_edit_plan_id,
asset_select_mode=request.asset_select_mode,
batch_id=batch_id,
auto_retry_enabled=request.auto_retry_enabled,
auto_retry_max=request.auto_retry_max,
)
)
try:
+136 -79
View File
@@ -21,11 +21,12 @@ from app.schemas.task_center import (
ProjectTaskResponse,
UserTaskResponse,
)
from fastapi import APIRouter, Depends, HTTPException
from fastapi import APIRouter, Depends, HTTPException, Query
from packages.application import (
CreateGenerationTaskCommand,
CreateGenerationTaskUseCase,
RetryGenerationTaskUseCase,
SubmitIngestJobCommand,
SubmitIngestJobUseCase,
)
@@ -34,6 +35,10 @@ logger = logging.getLogger(__name__)
router = APIRouter()
DEFAULT_PAGE_SIZE = 50
MAX_PAGE_SIZE = 200
def _humanize_task_error(error_message: str) -> str:
raw = (error_message or "").strip()
if not raw:
@@ -63,6 +68,8 @@ def _generation_step(task) -> str:
return "生成完成"
if s == "failed":
return "生成失败"
if s == "cancelled":
return "已取消"
return s
@@ -79,6 +86,26 @@ def _ingest_step(job) -> str:
return s
def _generation_task_to_user_response(task) -> UserTaskResponse:
return UserTaskResponse(
id=f"generation:{task.id}",
task_type="generation",
project_id=task.project_id,
template_id=task.template_id,
status=_status_value(task.status),
progress=task.progress,
current_step=_generation_step(task),
error_message=task.error_message,
error_info=task.error_info or {},
user_message=_humanize_task_error(task.error_message),
retryable=_status_value(task.status) == "failed",
retry_count=task.retry_count or 0,
source_id=task.id,
created_at=task.created_at,
updated_at=task.completed_at or task.started_at or task.created_at,
)
def _generation_task_to_project_response(task) -> ProjectTaskResponse:
return ProjectTaskResponse(
id=f"generation:{task.id}",
@@ -88,8 +115,10 @@ def _generation_task_to_project_response(task) -> ProjectTaskResponse:
progress=task.progress,
current_step=_generation_step(task),
error_message=task.error_message,
error_info=task.error_info or {},
user_message=_humanize_task_error(task.error_message),
retryable=_status_value(task.status) == "failed",
retry_count=task.retry_count or 0,
source_id=task.id,
template_id=task.template_id,
created_at=task.created_at,
@@ -97,40 +126,66 @@ def _generation_task_to_project_response(task) -> ProjectTaskResponse:
)
def _validate_status(status: str | None) -> str | None:
"""校验状态值合法性。"""
if status is None:
return None
valid = {"pending", "running", "completed", "failed", "cancelled"}
if status not in valid:
raise HTTPException(
status_code=400,
detail=f"无效的状态筛选值: {status},允许值: {', '.join(sorted(valid))}",
)
return status
def _clamp_page_size(page_size: int) -> int:
if page_size <= 0:
return DEFAULT_PAGE_SIZE
if page_size > MAX_PAGE_SIZE:
return MAX_PAGE_SIZE
return page_size
# ── 用户级端点(放在项目级端点之前,避免路由冲突) ──
@router.get("/tasks", response_model=ListTasksResponse)
def list_user_tasks(
status: str | None = Query(None, description="按状态筛选:pending/running/completed/failed/cancelled"),
task_type: str | None = Query(None, description="按任务类型筛选:generation/ingest"),
page: int = Query(1, ge=1, description="页码,从1开始"),
page_size: int = Query(DEFAULT_PAGE_SIZE, ge=1, le=MAX_PAGE_SIZE, description="每页数量"),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
ingest_job_repository: Any = Depends(get_ingest_job_repository),
generation_task_repository: Any = Depends(get_generation_task_repository),
) -> ListTasksResponse:
"""用户级任务列表(跨 project),合并 ingest + generation 任务"""
"""用户级任务列表(跨 project),支持状态/类型筛选和分页"""
status = _validate_status(status)
page_size = _clamp_page_size(page_size)
user_id = authenticated_user.user.id
offset = (page - 1) * page_size
items: list[UserTaskResponse] = []
for task in generation_task_repository.list_by_user(user_id):
items.append(
UserTaskResponse(
id=f"generation:{task.id}",
task_type="generation",
project_id=task.project_id,
template_id=task.template_id,
status=_status_value(task.status),
progress=task.progress,
current_step=_generation_step(task),
error_message=task.error_message,
user_message=_humanize_task_error(task.error_message),
retryable=_status_value(task.status) == "failed",
source_id=task.id,
created_at=task.created_at,
updated_at=task.completed_at or task.started_at or task.created_at,
)
# 生成任务
if task_type is None or task_type == "generation":
gen_result = generation_task_repository.list_by_user_filtered(
user_id,
status=status,
limit=page_size + 1, # 多取一条判断是否还有下一页(简单起见这里用offset)
offset=offset,
)
for task in gen_result:
items.append(_generation_task_to_user_response(task))
# 按时间倒序
items.sort(key=lambda item: item.updated_at or item.created_at or "", reverse=True)
return ListTasksResponse(items=items)
# 总数(仅generation,ingest暂不计入总数以保持简单)
total = generation_task_repository.count_by_user_filtered(user_id, status=status)
return ListTasksResponse(items=items[:page_size], total=total)
@router.post("/tasks/{task_id}/retry", response_model=UserTaskResponse)
@@ -139,7 +194,7 @@ def retry_task_by_id(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
generation_task_repository: Any = Depends(get_generation_task_repository),
) -> UserTaskResponse:
"""简化重试:通过 task_id 直接重试失败的生成任务"""
"""原地重试失败的生成任务(复用同一个task_idretry_count+1"""
task = generation_task_repository.get(task_id)
if task is None:
raise HTTPException(status_code=404, detail="Generation task not found")
@@ -149,6 +204,7 @@ def retry_task_by_id(
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()
@@ -163,20 +219,11 @@ def retry_task_by_id(
detail="系统繁忙,请稍后再试",
)
use_case = CreateGenerationTaskUseCase(generation_task_repository)
retried = use_case.execute(
CreateGenerationTaskCommand(
project_id=task.project_id,
asset_library_id=task.asset_library_id,
strategy_id=task.strategy_id,
voice_library_id=task.voice_library_id,
template_id=task.template_id,
asset_ids=task.asset_ids,
title_ids=task.title_ids,
voice_ids=task.voice_ids,
created_by_user_id=user_id,
)
)
# 原地重试
use_case = RetryGenerationTaskUseCase(generation_task_repository)
retried = use_case.execute(task_id)
# 重新入队
try:
if not safe_enqueue_generation_task(
retried, generation_task_repository, user_id=user_id, log_prefix="[任务中心]"
@@ -192,18 +239,8 @@ def retry_task_by_id(
status_code=503,
detail="系统繁忙,请稍后再试",
) from None
return UserTaskResponse(
id=f"generation:{retried.id}",
task_type="generation",
project_id=retried.project_id,
template_id=retried.template_id,
status=_status_value(retried.status),
progress=retried.progress,
current_step=_generation_step(retried),
source_id=retried.id,
created_at=retried.created_at,
updated_at=retried.created_at,
)
return _generation_task_to_user_response(retried)
# ── 项目级端点 ──
@@ -212,37 +249,64 @@ def retry_task_by_id(
@router.get("/projects/{project_id}/tasks", response_model=ListProjectTasksResponse)
def list_project_tasks(
project_id: str,
status: str | None = Query(None, description="按状态筛选:pending/running/completed/failed/cancelled"),
task_type: str | None = Query(None, description="按任务类型筛选:generation/ingest"),
page: int = Query(1, ge=1, description="页码,从1开始"),
page_size: int = Query(DEFAULT_PAGE_SIZE, ge=1, le=MAX_PAGE_SIZE, description="每页数量"),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
project_repository: Any = Depends(get_project_repository),
ingest_job_repository: Any = Depends(get_ingest_job_repository),
generation_task_repository: Any = Depends(get_generation_task_repository),
) -> ListProjectTasksResponse:
"""项目级任务列表,支持状态/类型筛选和分页。"""
project = project_repository.find_by_id(project_id)
if project is None:
raise HTTPException(status_code=404, detail="Project not found")
status = _validate_status(status)
page_size = _clamp_page_size(page_size)
offset = (page - 1) * page_size
items: list[ProjectTaskResponse] = []
for job in ingest_job_repository.list_by_project(project_id):
items.append(
ProjectTaskResponse(
id=f"ingest:{job.id}",
task_type="ingest",
project_id=job.project_id,
status=_status_value(job.status),
progress=100.0 if _status_value(job.status) == "completed" else 0.0,
current_step=_ingest_step(job),
error_message=job.error_message,
user_message=_humanize_task_error(job.error_message),
retryable=_status_value(job.status) == "failed",
source_id=job.id,
created_at=job.created_at,
updated_at=job.updated_at,
# 导入任务
if task_type is None or task_type == "ingest":
for job in ingest_job_repository.list_by_project(project_id):
if status and _status_value(job.status) != status:
continue
items.append(
ProjectTaskResponse(
id=f"ingest:{job.id}",
task_type="ingest",
project_id=job.project_id,
status=_status_value(job.status),
progress=100.0 if _status_value(job.status) == "completed" else 0.0,
current_step=_ingest_step(job),
error_message=job.error_message,
user_message=_humanize_task_error(job.error_message),
retryable=_status_value(job.status) == "failed",
source_id=job.id,
created_at=job.created_at,
updated_at=job.updated_at,
)
)
# 生成任务
if task_type is None or task_type == "generation":
gen_items = generation_task_repository.list_by_project_filtered(
project_id,
status=status,
limit=page_size + 1,
offset=offset,
)
for task in generation_task_repository.list_by_project(project_id):
items.append(_generation_task_to_project_response(task))
for task in gen_items:
items.append(_generation_task_to_project_response(task))
items.sort(key=lambda item: item.updated_at or item.created_at or "", reverse=True)
return ListProjectTasksResponse(items=items)
total = generation_task_repository.count_by_project_filtered(project_id, status=status)
return ListProjectTasksResponse(items=items[:page_size], total=total)
@router.post("/tasks/{task_type}/{source_id}/retry", response_model=ProjectTaskResponse)
@@ -253,6 +317,7 @@ def retry_project_task(
ingest_job_repository: Any = Depends(get_ingest_job_repository),
generation_task_repository: Any = Depends(get_generation_task_repository),
) -> ProjectTaskResponse:
"""项目级任务重试。"""
if task_type == "generation":
task = generation_task_repository.get(source_id)
if task is None:
@@ -261,6 +326,7 @@ def retry_project_task(
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()
@@ -275,20 +341,10 @@ def retry_project_task(
detail="系统繁忙,请稍后再试",
)
use_case = CreateGenerationTaskUseCase(generation_task_repository)
retried = use_case.execute(
CreateGenerationTaskCommand(
project_id=task.project_id,
asset_library_id=task.asset_library_id,
strategy_id=task.strategy_id,
voice_library_id=task.voice_library_id,
template_id=task.template_id,
asset_ids=task.asset_ids,
title_ids=task.title_ids,
voice_ids=task.voice_ids,
created_by_user_id=user_id,
)
)
# 原地重试
use_case = RetryGenerationTaskUseCase(generation_task_repository)
retried = use_case.execute(source_id)
try:
if not safe_enqueue_generation_task(
retried, generation_task_repository, user_id=user_id, log_prefix="[任务中心]"
@@ -305,6 +361,7 @@ def retry_project_task(
detail="系统繁忙,请稍后再试",
) from None
return _generation_task_to_project_response(retried)
if task_type == "ingest":
job = ingest_job_repository.get(source_id)
if job is None:
+94 -6
View File
@@ -8,13 +8,16 @@ from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from app.schemas.template import (
CategoryResponse,
CopyTemplateRequest,
CreateCategoryRequest,
CreateTemplateRequest,
GenerateWarningResponse,
ListCategoriesResponse,
ListTagsResponse,
ListTemplatesResponse,
SegmentResponse,
TemplateResponse,
TemplateUsageResponse,
ToggleFavoriteResponse,
UpdateTemplateRequest,
ValidateTemplateRequest,
@@ -27,19 +30,25 @@ logger = logging.getLogger(__name__)
from packages.adapters.sqlalchemy_impl.template_repository import SQLAlchemyTemplateRepository
from packages.application.template.commands import (
CopyTemplateCommand,
CreateCategoryCommand,
CreateTemplateCommand,
ListTemplatesFilter,
SegmentCommand,
UpdateTemplateCommand,
ValidateTemplateCommand,
)
from packages.application.template.use_cases import (
CopyTemplateUseCase,
CountTemplatesUseCase,
CreateCategoryUseCase,
CreateTemplateUseCase,
DeleteCategoryUseCase,
DeleteTemplateUseCase,
GetTemplateUsageUseCase,
GetTemplateUseCase,
ListCategoriesUseCase,
ListTagsUseCase,
ListTemplatesUseCase,
NotFoundError,
UpdateTemplateUseCase,
@@ -67,7 +76,7 @@ def _segment_to_response(seg) -> SegmentResponse:
)
def _to_response(template) -> TemplateResponse:
def _to_response(template, usage_count: int = 0) -> TemplateResponse:
return TemplateResponse(
id=template.id,
user_id=template.user_id,
@@ -81,6 +90,7 @@ def _to_response(template) -> TemplateResponse:
estimated_duration=template.estimated_duration,
segments=[_segment_to_response(s) for s in getattr(template, "segments", [])],
is_active=template.is_active,
usage_count=usage_count,
created_at=template.created_at,
updated_at=template.updated_at,
)
@@ -93,19 +103,36 @@ def _to_response(template) -> TemplateResponse:
def list_templates(
skip: int = Query(0, ge=0),
limit: int = Query(50, ge=1, le=200),
category: str | None = Query(None, description="按分类筛选"),
tag: str | None = Query(None, description="按标签筛选"),
keyword: str | None = Query(None, description="按名称关键词搜索"),
mode: str | None = Query(None, description="按剪辑模式筛选"),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ListTemplatesResponse:
user_id = authenticated_user.user.id
try:
tpl_filter = ListTemplatesFilter(
category=category,
tag=tag,
keyword=keyword,
mode=mode,
)
use_case = ListTemplatesUseCase(template_repository)
templates = use_case.execute(user_id, skip=skip, limit=limit)
total = template_repository.count_by_user(user_id)
templates = use_case.execute(user_id, skip=skip, limit=limit, filter=tpl_filter)
count_use_case = CountTemplatesUseCase(template_repository)
total = count_use_case.execute(user_id, filter=tpl_filter)
# 批量查询使用次数
items = []
for t in templates:
usage = template_repository.get_usage_count(t.id)
items.append(_to_response(t, usage_count=usage))
except Exception:
logger.exception("list_templates 查询失败: user_id=%s", user_id)
return ListTemplatesResponse(items=[], total=0)
return ListTemplatesResponse(
items=[_to_response(t) for t in templates],
items=items,
total=total,
)
@@ -120,12 +147,13 @@ def get_template(
try:
use_case = GetTemplateUseCase(template_repository)
template = use_case.execute(template_id, user_id)
usage = template_repository.get_usage_count(template_id)
except Exception:
logger.exception("get_template 查询失败: template_id=%s", template_id)
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="模板查询失败")
if template is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
return _to_response(template)
return _to_response(template, usage_count=usage)
@router.post("", response_model=TemplateResponse, status_code=status.HTTP_201_CREATED)
@@ -220,6 +248,47 @@ def delete_template(
return
@router.post("/{template_id}/copy", response_model=TemplateResponse, status_code=status.HTTP_201_CREATED)
def copy_template(
template_id: str,
request: CopyTemplateRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateResponse:
"""复制模板(含所有片段配置)"""
user_id = authenticated_user.user.id
command = CopyTemplateCommand(
template_id=template_id,
user_id=user_id,
new_name=request.new_name,
)
use_case = CopyTemplateUseCase(template_repository)
try:
template = use_case.execute(command)
except NotFoundError:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
except ValidationError as exc:
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc))
return _to_response(template)
@router.get("/{template_id}/usage", response_model=TemplateUsageResponse)
def get_template_usage(
template_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateUsageResponse:
"""获取模板使用次数(关联的剪辑计划数量)"""
user_id = authenticated_user.user.id
# 鉴权:确保模板存在且属于当前用户
use_case = GetTemplateUseCase(template_repository)
template = use_case.execute(template_id, user_id)
if template is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
usage = template_repository.get_usage_count(template_id)
return TemplateUsageResponse(template_id=template_id, usage_count=usage)
@router.post("/{template_id}/toggle-favorite", response_model=ToggleFavoriteResponse)
def toggle_favorite(
template_id: str,
@@ -318,4 +387,23 @@ def delete_category(
deleted = use_case.execute(category_id, user_id)
if not deleted:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Category not found")
return
return Response(status_code=204)
# ── Tags ──
@router.get("/tags/list", response_model=ListTagsResponse)
def list_tags(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ListTagsResponse:
"""获取用户所有模板标签(去重排序)"""
user_id = authenticated_user.user.id
try:
use_case = ListTagsUseCase(template_repository)
tags = use_case.execute(user_id)
except Exception:
logger.exception("list_tags 查询失败: user_id=%s", user_id)
return ListTagsResponse(items=[])
return ListTagsResponse(items=tags)
+41 -1
View File
@@ -17,13 +17,18 @@ from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.title_library_repository import SQLAlchemyTitleLibraryRepository
from packages.application.title_library.commands import CreateTitleLibraryCommand, UpdateTitleLibraryCommand
from packages.application.title_library.commands import (
CreateTitleLibraryCommand,
PickTitleCommand,
UpdateTitleLibraryCommand,
)
from packages.application.title_library.use_cases import (
CreateTitleLibraryUseCase,
DeleteTitleLibraryUseCase,
GetTitleLibraryUseCase,
ListTitleLibraryUseCase,
NotFoundError,
PickTitleUseCase,
QuotaExceededError,
UpdateTitleLibraryUseCase,
)
@@ -70,6 +75,41 @@ def list_titles(
)
@router.post("/pick", response_model=TitleLibraryItemResponse)
def pick_title(
category: Optional[str] = Query(None, description="按分类筛选,不填则从全部标题中选"),
exclude_ids: Optional[str] = Query(
None,
description="排除的标题ID(逗号分隔),用于批量生成时避免重复",
),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
) -> TitleLibraryItemResponse:
"""智能选择一个标题。
策略:优先使用次数少的,从最少的前5个中随机选一个,兼顾公平和多样性。
"""
user_id = authenticated_user.user.id
exclude_list: list[str] = []
if exclude_ids:
exclude_list = [t.strip() for t in exclude_ids.split(",") if t.strip()]
use_case = PickTitleUseCase(title_repository)
item = use_case.execute(
PickTitleCommand(
user_id=user_id,
category=category,
exclude_ids=exclude_list,
)
)
if item is None:
raise HTTPException(
status_code=status.HTTP_404_NOT_FOUND,
detail="标题库为空,请先添加标题",
)
return _to_response(item)
@router.get("/{title_id}", response_model=TitleLibraryItemResponse)
def get_title(
title_id: str,
+29
View File
@@ -46,6 +46,7 @@ from packages.application.voice_library.use_cases import (
CreateVoiceLibraryUseCase,
QuotaExceededError,
)
from packages.domain.voice_presets import list_voices
from packages.ports.user_repository import UserRepository
logger = logging.getLogger(__name__)
@@ -53,6 +54,34 @@ logger = logging.getLogger(__name__)
router = APIRouter()
@router.get("/presets", summary="获取预设音色列表")
def list_preset_voices(
gender: Optional[str] = Query(None, description="按性别筛选: male/female/child"),
style: Optional[str] = Query(None, description="按风格筛选: stable/lively/customer_service/narration/news/story"),
keyword: Optional[str] = Query(None, description="按关键词搜索"),
_user: AuthenticatedUser = Depends(get_current_user),
) -> list[dict]:
"""获取可用的预设音色列表。
用于配音功能的音色选择。
"""
voices = list_voices(gender=gender, style=style, keyword=keyword)
return [
{
"voice_id": v.voice_id,
"name": v.name,
"gender": v.gender.value,
"style": v.style.value,
"description": v.description,
"default_speed": v.default_speed,
"default_pitch": v.default_pitch,
"sample_rate": v.sample_rate,
"language": v.language,
}
for v in voices
]
def _get_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyTTSJobRepository:
return SQLAlchemyTTSJobRepository(session)
+179
View File
@@ -0,0 +1,179 @@
import logging
import uuid
from app.api.routes._helpers import check_project_access
from app.auth import AuthenticatedUser, get_current_user
from app.core.celery_app import celery_app
from app.core.storage import OSSStorageService, get_storage_service
from app.dependencies import get_generated_video_repository
from app.schemas.video_center import (
BatchDownloadRequest,
BatchDownloadResponse,
ListVideosResponse,
UpdateVideoReviewRequest,
VideoItemResponse,
)
from fastapi import APIRouter, Depends, HTTPException, Query
from packages.application import (
GetGeneratedVideoUseCase,
GetVideosByIdsUseCase,
ListGeneratedVideosPaginatedUseCase,
UpdateVideoReviewStatusUseCase,
)
logger = logging.getLogger(__name__)
router = APIRouter()
def _to_video_response(item, storage: OSSStorageService | None = None) -> VideoItemResponse:
download_url = None
if storage and item.file_url:
try:
download_url = storage.get_download_url(item.file_url)
except Exception:
download_url = item.file_url
return VideoItemResponse(
id=item.id,
project_id=item.project_id,
generation_task_id=item.generation_task_id,
name=item.name,
file_url=item.file_url,
file_size=item.file_size,
duration=item.duration,
thumbnail_url=item.thumbnail_url,
width=item.width,
height=item.height,
fps=item.fps,
status=item.status,
review_status=item.review_status,
generation_params=item.generation_params,
download_url=download_url,
generated_at=item.generated_at.isoformat() if hasattr(item, "generated_at") and item.generated_at else "",
)
@router.get("/videos", response_model=ListVideosResponse)
def list_videos(
project_id: str | None = Query(None, description="项目ID,不传则返回所有项目"),
status: str | None = Query(None, description="按状态筛选"),
review_status: str | None = Query(None, description="按复核状态筛选"),
page: int = Query(1, ge=1, description="页码"),
page_size: int = Query(20, ge=1, le=100, description="每页数量"),
repo=Depends(get_generated_video_repository),
storage: OSSStorageService = Depends(get_storage_service),
current_user: AuthenticatedUser = Depends(get_current_user),
):
"""成片列表,支持分页、按项目/状态/复核状态筛选。"""
use_case = ListGeneratedVideosPaginatedUseCase(repo)
items, total = use_case.execute(
project_id=project_id,
status=status,
review_status=review_status,
page=page,
page_size=page_size,
)
return ListVideosResponse(
items=[_to_video_response(item, storage) for item in items],
total=total,
page=page,
page_size=page_size,
)
@router.get("/videos/{video_id}", response_model=VideoItemResponse)
def get_video(
video_id: str,
repo=Depends(get_generated_video_repository),
storage: OSSStorageService = Depends(get_storage_service),
current_user: AuthenticatedUser = Depends(get_current_user),
):
"""获取单个成片详情。"""
use_case = GetGeneratedVideoUseCase(repo)
item = use_case.execute(video_id)
if item is None:
raise HTTPException(status_code=404, detail="Video not found")
return _to_video_response(item, storage)
@router.patch("/videos/{video_id}/review", response_model=VideoItemResponse)
def update_video_review_status(
video_id: str,
request: UpdateVideoReviewRequest,
repo=Depends(get_generated_video_repository),
storage: OSSStorageService = Depends(get_storage_service),
current_user: AuthenticatedUser = Depends(get_current_user),
):
"""更新成片复核状态:pending_review / approved / rejected。"""
use_case = UpdateVideoReviewStatusUseCase(repo)
item = use_case.execute(video_id, request.review_status)
if item is None:
raise HTTPException(status_code=404, detail="Video not found")
logger.info("Video %s review status updated to %s by user %s", video_id, request.review_status, current_user.user_id)
return _to_video_response(item, storage)
@router.post("/videos/batch-download", response_model=BatchDownloadResponse)
def batch_download_videos(
request: BatchDownloadRequest,
repo=Depends(get_generated_video_repository),
current_user: AuthenticatedUser = Depends(get_current_user),
):
"""批量下载成片,异步打包 zip。
传入 video_ids 列表,创建一个批量下载任务,任务完成后返回 zip 下载链接。
"""
if not request.video_ids:
raise HTTPException(status_code=400, detail="video_ids cannot be empty")
if len(request.video_ids) > 50:
raise HTTPException(status_code=400, detail="Maximum 50 videos per batch download")
# 校验视频都存在
use_case = GetVideosByIdsUseCase(repo)
videos = use_case.execute(request.video_ids)
if len(videos) != len(request.video_ids):
raise HTTPException(status_code=404, detail="Some videos not found")
# 发送 celery 任务
task = celery_app.send_task(
"worker.batch_download_videos",
args=[request.video_ids, current_user.user_id],
)
logger.info("Batch download job created: %s, videos=%d", task.id, len(request.video_ids))
return BatchDownloadResponse(job_id=task.id, status="pending")
@router.get("/videos/batch-download/{job_id}", response_model=BatchDownloadResponse)
def get_batch_download_status(
job_id: str,
current_user: AuthenticatedUser = Depends(get_current_user),
):
"""查询批量下载任务状态。"""
from celery.result import AsyncResult
task = AsyncResult(job_id, app=celery_app)
status_map = {
"PENDING": "pending",
"STARTED": "running",
"SUCCESS": "success",
"FAILURE": "failed",
"RETRY": "pending",
"REVOKED": "cancelled",
}
api_status = status_map.get(task.state, "pending")
download_url = None
if task.state == "SUCCESS" and task.result:
if isinstance(task.result, dict):
download_url = task.result.get("download_url")
elif isinstance(task.result, str):
download_url = task.result
return BatchDownloadResponse(
job_id=job_id,
status=api_status,
download_url=download_url,
)
Regular → Executable
+34 -6
View File
@@ -54,17 +54,45 @@ class AssetResponse(BaseModel):
tag_ids: list[str] = Field(default_factory=list)
MAX_BATCH_SIZE = 200
class BatchDeleteRequest(BaseModel):
"""批量删除请求。"""
"""批量删除请求(软删除)"""
ids: list[str] = Field(..., min_length=1, max_length=100, description="要删除的素材 ID 列表")
asset_ids: list[str] = Field(..., min_length=1, max_length=MAX_BATCH_SIZE, description="要删除的素材 ID 列表")
class BatchDeleteResponse(BaseModel):
"""批量删除响应。"""
class BatchOperationResponse(BaseModel):
"""批量操作通用响应。"""
deleted_count: int = Field(..., ge=0, description="实际删除数量")
failed_ids: list[str] = Field(default_factory=list, description="删除失败的 ID 列表")
success_count: int = Field(..., ge=0, description="成功数量")
failed_ids: list[str] = Field(default_factory=list, description="失败的 ID 列表")
failed_details: dict[str, str] = Field(default_factory=dict, description="失败详情 {asset_id: reason}")
class BatchTagRequest(BaseModel):
"""批量打标签请求。"""
asset_ids: list[str] = Field(..., min_length=1, max_length=MAX_BATCH_SIZE, description="素材 ID 列表")
tag_ids: list[str] = Field(..., min_length=1, max_length=50, description="标签 ID 列表")
mode: str = Field(default="add", pattern="^(add|replace)$", description="add=添加合并,replace=全量替换")
class BatchClassifyRequest(BaseModel):
"""批量修改分类请求。"""
asset_ids: list[str] = Field(..., min_length=1, max_length=MAX_BATCH_SIZE, description="素材 ID 列表")
category: str = Field(..., min_length=1, max_length=50, description="内容分类,如 person/scenic/product")
class BatchMarkRequest(BaseModel):
"""批量设置智能视图标记请求。"""
asset_ids: list[str] = Field(..., min_length=1, max_length=MAX_BATCH_SIZE, description="素材 ID 列表")
smart_view: str = Field(
..., pattern="^(recommended|caution|high_risk)$", description="智能视图标记:recommended/caution/high_risk"
)
class ListAssetsResponse(BaseModel):
+15
View File
@@ -33,6 +33,17 @@ class CreateGenerationTaskRequest(BaseModel):
asset_select_count: int = Field(
default=0, ge=0, le=100, description="选取数量,0表示全部(仅 random/smart 模式有效)"
)
# ── 自动重试 ──
auto_retry_enabled: bool = Field(
default=False,
description="是否开启失败自动重试,默认关闭",
)
auto_retry_max: int = Field(
default=0,
ge=0,
le=5,
description="最大自动重试次数,0表示不自动重试,最大5次",
)
@model_validator(mode="after")
def _check_at_least_one_mode(self) -> "CreateGenerationTaskRequest":
@@ -64,6 +75,10 @@ class GenerationTaskResponse(BaseModel):
progress: float
result_count: int
error_message: str
error_info: dict = Field(default_factory=dict)
retry_count: int = 0
auto_retry_enabled: bool = False
auto_retry_max: int = 0
logs: list[dict] = Field(default_factory=list)
@field_validator("logs", mode="before")
+6
View File
@@ -11,8 +11,10 @@ class ProjectTaskResponse(BaseModel):
progress: float
current_step: str
error_message: str = ""
error_info: dict = Field(default_factory=dict)
user_message: str = ""
retryable: bool = False
retry_count: int = 0
source_id: str = ""
template_id: str = ""
created_at: datetime | None = None
@@ -21,6 +23,7 @@ class ProjectTaskResponse(BaseModel):
class ListProjectTasksResponse(BaseModel):
items: list[ProjectTaskResponse] = Field(default_factory=list)
total: int = 0
class UserTaskResponse(BaseModel):
@@ -34,8 +37,10 @@ class UserTaskResponse(BaseModel):
progress: float
current_step: str
error_message: str = ""
error_info: dict = Field(default_factory=dict)
user_message: str = ""
retryable: bool = False
retry_count: int = 0
source_id: str = ""
created_at: datetime | None = None
updated_at: datetime | None = None
@@ -45,3 +50,4 @@ class ListTasksResponse(BaseModel):
"""用户级任务列表响应(GET /api/v1/tasks)。"""
items: list[UserTaskResponse] = Field(default_factory=list)
total: int = 0
Regular → Executable
+23
View File
@@ -45,6 +45,7 @@ class TemplateResponse(BaseModel):
segments: List[SegmentResponse] = Field(default_factory=list)
is_active: bool = True
is_favorite: bool = False
usage_count: int = 0
created_at: datetime
updated_at: datetime
@@ -120,3 +121,25 @@ class CreateCategoryRequest(BaseModel):
class ListCategoriesResponse(BaseModel):
items: List[CategoryResponse]
# ── Copy Template ──
class CopyTemplateRequest(BaseModel):
new_name: str
# ── Tags ──
class ListTagsResponse(BaseModel):
items: List[str]
# ── Usage Stats ──
class TemplateUsageResponse(BaseModel):
template_id: str
usage_count: int
+45
View File
@@ -0,0 +1,45 @@
from typing import Literal
from pydantic import BaseModel, Field
VideoReviewStatus = Literal["pending_review", "approved", "rejected"]
class VideoItemResponse(BaseModel):
id: str
project_id: str
generation_task_id: str
name: str
file_url: str
file_size: int
duration: float
thumbnail_url: str | None = None
width: int
height: int
fps: float
status: str = "completed"
review_status: str = "pending_review"
generation_params: dict = Field(default_factory=dict)
download_url: str | None = None
generated_at: str = ""
class ListVideosResponse(BaseModel):
items: list[VideoItemResponse]
total: int
page: int
page_size: int
class UpdateVideoReviewRequest(BaseModel):
review_status: VideoReviewStatus
class BatchDownloadRequest(BaseModel):
video_ids: list[str]
class BatchDownloadResponse(BaseModel):
job_id: str
status: str = "pending"
download_url: str | None = None
+19
View File
@@ -281,6 +281,8 @@ class EditPlanService:
start_time: float = 0.0,
duration: float = 0.0,
transition_effect: str = "cut",
transition_duration: float = 0.0,
playback_speed: float = 1.0,
config: Optional[dict[str, Any]] = None,
) -> EditPlanClip:
"""创建片段
@@ -301,6 +303,8 @@ class EditPlanService:
start_time=start_time,
duration=duration,
transition_effect=transition_effect,
transition_duration=transition_duration,
playback_speed=playback_speed,
config=config,
)
created = self._clip_repo.create(clip)
@@ -324,6 +328,8 @@ class EditPlanService:
start_time: Optional[float] = None,
duration: Optional[float] = None,
transition_effect: Optional[str] = None,
transition_duration: Optional[float] = None,
playback_speed: Optional[float] = None,
config: Optional[dict[str, Any]] = None,
) -> EditPlanClip:
"""更新片段
@@ -333,6 +339,15 @@ class EditPlanService:
"""
existing = self.get_clip_or_raise(clip_id)
# 速度边界钳制
if playback_speed is not None:
if playback_speed <= 0:
playback_speed = 1.0
elif playback_speed < 0.25:
playback_speed = 0.25
elif playback_speed > 4.0:
playback_speed = 4.0
updated = EditPlanClip(
id=existing.id,
plan_id=existing.plan_id,
@@ -346,6 +361,10 @@ class EditPlanService:
transition_effect=(
transition_effect.strip() if transition_effect is not None else existing.transition_effect
),
transition_duration=(
transition_duration if transition_duration is not None else existing.transition_duration
),
playback_speed=playback_speed if playback_speed is not None else existing.playback_speed,
status=existing.status,
config=config if config is not None else existing.config,
created_at=existing.created_at,
+6
View File
@@ -26,6 +26,9 @@ export default defineConfig({
use: {
...devices["Desktop Chrome"],
channel: process.env.E2E_BROWSER_CHANNEL || "msedge",
launchOptions: {
args: ["--disable-gpu", "--disable-software-rasterizer"],
},
},
},
{
@@ -51,6 +54,9 @@ export default defineConfig({
use: {
...devices["Desktop Chrome"],
channel: process.env.E2E_BROWSER_CHANNEL || "msedge",
launchOptions: {
args: ["--disable-gpu", "--disable-software-rasterizer"],
},
},
},
],
+47
View File
@@ -0,0 +1,47 @@
"""ASR 服务工厂 — 根据环境配置创建对应 ASR 服务实例。
支持的后端:
- mock: MockASRService(测试/开发用)
- 后续可扩展:whisper / aliyun / tencent 等
"""
from __future__ import annotations
import os
from functools import lru_cache
from packages.ports.asr_service import ASRService
@lru_cache(maxsize=1)
def get_asr_service() -> ASRService | None:
"""获取全局 ASR 服务实例(单例)。
根据环境变量 ASR_PROVIDER 决定使用哪个后端:
- mock / 空 / 未设置: 返回 None(不启用 ASR)
- mock: 使用 MockASRService
Returns:
ASRService 实例,未配置或不启用时返回 None
"""
provider = os.environ.get("ASR_PROVIDER", "").lower().strip()
if not provider:
return None
if provider == "mock":
from packages.adapters.asr.mock_asr_service import MockASRService
return MockASRService()
# 未知 provider,记录日志并返回 None(不启用 ASR,不阻断主流程)
import logging
logger = logging.getLogger(__name__)
logger.warning("未知的 ASR provider: %sASR 自动字幕功能未启用", provider)
return None
def reset_asr_service_cache() -> None:
"""重置 ASR 服务缓存(测试用)。"""
get_asr_service.cache_clear()
+66
View File
@@ -0,0 +1,66 @@
"""TTS 服务工厂.
根据配置创建对应的 TTS 服务实例。
"""
from __future__ import annotations
import logging
import os
from packages.ports.tts_service import TtsService
logger = logging.getLogger(__name__)
# 可用的 provider 映射
_PROVIDERS: dict[str, type[TtsService]] = {}
def register_provider(name: str, cls: type[TtsService]) -> None:
"""注册 TTS 供应商."""
_PROVIDERS[name] = cls
def get_tts_service(provider: str | None = None, **kwargs) -> TtsService:
"""获取 TTS 服务实例.
Args:
provider: 供应商名称(None 则从环境变量读取 TTS_PROVIDER
**kwargs: 传递给服务构造函数的参数
Returns:
TTS 服务实例
Raises:
ValueError: 不支持的供应商
"""
if provider is None:
provider = os.environ.get("TTS_PROVIDER", "mock")
provider = provider.lower()
if provider not in _PROVIDERS:
# 延迟导入避免循环依赖
if provider == "mock":
from packages.adapters.tts.mock_tts_service import MockTtsService
_PROVIDERS["mock"] = MockTtsService
else:
logger.warning("未知 TTS provider: %s,回退到 mock", provider)
from packages.adapters.tts.mock_tts_service import MockTtsService
_PROVIDERS["mock"] = MockTtsService
provider = "mock"
cls = _PROVIDERS[provider]
return cls(**kwargs)
def available_providers() -> list[str]:
"""获取可用的供应商列表."""
# 确保 mock 已注册
if "mock" not in _PROVIDERS:
from packages.adapters.tts.mock_tts_service import MockTtsService
_PROVIDERS["mock"] = MockTtsService
return list(_PROVIDERS.keys())
+313
View File
@@ -0,0 +1,313 @@
"""BGM 混音模块 — 背景音乐与主音频混合.
基于 FFmpeg 实现:
- BGM 音量调节
- 淡入淡出(afade
- 循环播放(aloop,短 BGM 铺长视频)
- 人声闪避(sidechaincompress,有人声时BGM自动降低音量)
- amix 混音
作为 render_audio.py 的增强模块,在 mix_audio 后处理阶段被调用。
"""
from __future__ import annotations
import logging
from dataclasses import dataclass
from pathlib import Path
from typing import TYPE_CHECKING
from video_processing.ffmpeg_utils import FFMPEG_BIN, probe_duration, run_ffmpeg
if TYPE_CHECKING:
from video_processing.render_audio import RenderContext
logger = logging.getLogger(__name__)
@dataclass
class BGMConfig:
"""BGM 混音配置(内部使用,从 plan.config.bgm 转换而来)"""
bgm_path: str # BGM 本地文件路径
volume: float = 0.3 # 0.0 ~ 1.0
fade_in: float = 0.0 # 淡入时长(秒)
fade_out: float = 0.0 # 淡出时长(秒)
loop_enabled: bool = True # 是否循环铺满
sidechain_enabled: bool = False # 人声闪避
sidechain_ratio: float = 0.3 # 闪避时音量降低比例
sidechain_attack: float = 0.02 # 攻击时间
sidechain_release: float = 0.5 # 释放时间
sidechain_threshold: float = -25.0 # 触发阈值(dB
@classmethod
def from_config_dict(cls, bgm_path: str, config: dict) -> "BGMConfig":
"""从 plan.config.bgm 字典创建 BGMConfig。"""
return cls(
bgm_path=bgm_path,
volume=float(config.get("volume", 0.3)),
fade_in=float(config.get("fade_in", 0.0)),
fade_out=float(config.get("fade_out", 0.0)),
loop_enabled=bool(config.get("loop_enabled", True)),
sidechain_enabled=bool(config.get("sidechain_enabled", False)),
sidechain_ratio=float(config.get("sidechain_ratio", 0.3)),
sidechain_attack=float(config.get("sidechain_attack", 0.02)),
sidechain_release=float(config.get("sidechain_release", 0.5)),
sidechain_threshold=float(config.get("sidechain_threshold", -25.0)),
)
# ── BGM 预处理 ────────────────────────────────────────────────────────────────
def prepare_bgm_track(
ctx: "RenderContext",
bgm: BGMConfig,
target_duration: float,
) -> Path:
"""预处理 BGM 轨道:循环/截断 + 音量 + 淡入淡出.
生成一个时长精确等于 target_duration 的 BGM 音频文件。
后续再与主音频混音。
Args:
ctx: 渲染上下文
bgm: BGM 配置
target_duration: 目标时长(秒),通常等于视频总时长
Returns:
处理后的 BGM 音频文件路径
"""
output_path = ctx.work_dir / f"bgm_processed_{ctx.plan_id}.aac"
if target_duration <= 0:
target_duration = 5.0 # 兜底
bgm_dur = probe_duration(bgm.bgm_path)
needs_loop = bgm.loop_enabled and bgm_dur > 0 and bgm_dur < target_duration * 0.9
# 构建滤镜链
filter_parts: list[str] = []
input_looped: bool = False
if needs_loop:
# 计算需要循环多少次才能铺满
loop_count = max(1, int(target_duration / bgm_dur) + 2)
# aloop 滤镜:循环指定次数
filter_parts.append(f"aloop=loop={loop_count}:size=0")
input_looped = True
# 音量调节
volume = max(0.0, min(1.0, bgm.volume))
if abs(volume - 1.0) > 0.001:
filter_parts.append(f"volume={volume:.3f}")
# 淡入
if bgm.fade_in > 0:
filter_parts.append(f"afade=t=in:st=0:d={bgm.fade_in:.3f}")
# 淡出(从 target_duration - fade_out 开始)
if bgm.fade_out > 0 and target_duration > bgm.fade_out:
fade_start = target_duration - bgm.fade_out
filter_parts.append(f"afade=t=out:st={fade_start:.3f}:d={bgm.fade_out:.3f}")
# 最终截断到目标时长
filter_parts.append(f"atrim=0:{target_duration:.3f}")
filter_parts.append("asetpts=N/SR/TB") # 重置时间戳
filter_str = ",".join(filter_parts)
command = [
FFMPEG_BIN,
"-y",
"-i",
bgm.bgm_path,
"-filter:a",
filter_str,
"-c:a",
"aac",
"-b:a",
"128k",
str(output_path),
]
logger.info(
"[bgm] prepare BGM track: path=%s dur=%.2f target=%.2f loop=%s fade_in=%.2f fade_out=%.2f",
bgm.bgm_path[-40:],
bgm_dur,
target_duration,
needs_loop,
bgm.fade_in,
bgm.fade_out,
)
run_ffmpeg(command)
return output_path
# ── BGM + 主音频混音 ──────────────────────────────────────────────────────────
def mix_bgm_with_main(
ctx: "RenderContext",
main_audio_path: Path,
bgm: BGMConfig,
target_duration: float,
) -> Path:
"""将 BGM 与主音频混合.
两种模式:
1. 普通混音(sidechain 关闭):amix 两路音频
2. 人声闪避(sidechain 开启):用 sidechaincompress 让 BGM 跟随主音频音量自动调整
Args:
ctx: 渲染上下文
main_audio_path: 主音频文件路径(人声/原始音频)
bgm: BGM 配置
target_duration: 目标时长
Returns:
混音后的音频文件路径
"""
output_path = ctx.work_dir / f"audio_with_bgm_{ctx.plan_id}.aac"
# 先预处理 BGM 轨道(循环/音量/淡入淡出/截断)
bgm_processed = prepare_bgm_track(ctx, bgm, target_duration)
if not bgm.sidechain_enabled:
# 普通 amix 混音
_mix_simple(main_audio_path, bgm_processed, output_path)
else:
# sidechain 人声闪避混音
_mix_sidechain(main_audio_path, bgm_processed, output_path, bgm)
return output_path
def _mix_simple(main_path: Path, bgm_path: Path, output_path: Path) -> None:
"""简单 amix 混音:主音频 + BGM = 输出.
主音频权重 1.0,BGM 已经在预处理阶段调好了音量。
amix 会自动归一化,需要用 volume 补偿。
"""
# 使用 amixinputs=2duration=first(以主音频时长为准)
# 然后用 volume=2 补偿 amix 的衰减(2路输入每路平均乘0.5)
filter_complex = "[0:a][1:a]amix=inputs=2:duration=first:dropout_transition=0[outa];" "[outa]volume=2[final]"
command = [
FFMPEG_BIN,
"-y",
"-i",
str(main_path),
"-i",
str(bgm_path),
"-filter_complex",
filter_complex,
"-map",
"[final]",
"-c:a",
"aac",
"-b:a",
"128k",
str(output_path),
]
logger.info("[bgm] simple amix mix")
run_ffmpeg(command)
def _mix_sidechain(
main_path: Path,
bgm_path: Path,
output_path: Path,
bgm: BGMConfig,
) -> None:
"""sidechain 人声闪避混音.
原理:
- 主音频作为 sidechain 信号源
- BGM 轨道经过 sidechaincompress,根据主音频音量动态调整 BGM 音量
- 最后 amix 混音
FFmpeg sidechaincompress 参数:
- threshold: 触发阈值(dB),主音频超过此值时开始压缩
- ratio: 压缩比,越高压缩越狠
- attack: 攻击时间(秒)
- release: 释放时间(秒)
"""
# sidechain_ratio 表示闪避时 BGM 音量降低比例
# ratio = 1 / (1 - sidechain_ratio),但实际压缩比需要更精细调整
# 简化处理:把 ratio 映射到 2:1 ~ 10:1 范围
ratio = max(2.0, min(10.0, 1.0 / (1.0 - bgm.sidechain_ratio)))
filter_complex = (
# BGM 经过 sidechain 压缩,用主音频做触发
f"[1:a][0:a]sidechaincompress="
f"threshold={bgm.sidechain_threshold}dB:"
f"ratio={ratio:.1f}:"
f"attack={bgm.sidechain_attack:.3f}:"
f"release={bgm.sidechain_release:.3f}:"
f"knee=6[bgm_comp];"
# 主音频 + 压缩后的 BGM 混音
f"[0:a][bgm_comp]amix=inputs=2:duration=first:dropout_transition=0[outa];"
f"[outa]volume=1.5[final]" # 轻微补偿
)
command = [
FFMPEG_BIN,
"-y",
"-i",
str(main_path),
"-i",
str(bgm_path),
"-filter_complex",
filter_complex,
"-map",
"[final]",
"-c:a",
"aac",
"-b:a",
"128k",
str(output_path),
]
logger.info(
"[bgm] sidechain mix: threshold=%.1fdB ratio=%.1f attack=%.3f release=%.3f",
bgm.sidechain_threshold,
ratio,
bgm.sidechain_attack,
bgm.sidechain_release,
)
run_ffmpeg(command)
# ── 纯 BGM 模式(无主音频) ──────────────────────────────────────────────────
def build_bgm_only(
ctx: "RenderContext",
bgm: BGMConfig,
target_duration: float,
) -> Path:
"""只有 BGM、没有主音频时,直接生成 BGM 音频.
Args:
ctx: 渲染上下文
bgm: BGM 配置
target_duration: 目标时长
Returns:
BGM 音频文件路径
"""
output_path = ctx.work_dir / f"bgm_only_{ctx.plan_id}.aac"
if target_duration <= 0:
target_duration = 5.0
bgm_processed = prepare_bgm_track(ctx, bgm, target_duration)
# 直接复制
import shutil
shutil.copy2(bgm_processed, output_path)
return output_path
+248
View File
@@ -0,0 +1,248 @@
"""绿幕抠像引擎 — 基于 FFmpeg colorkey / chromakey 滤镜.
支持将指定颜色(默认绿色)变为透明,可用于虚拟背景、画中画背景替换等场景。
使用方式:
config = ChromaKeyConfig(key_color="#00FF00", similarity=0.3, blend=0.1)
engine = ChromaKeyEngine(config)
filter_str = engine.build_filter(input_label, output_label)
# 结果: [in]colorkey=color=0x00FF00:similarity=0.3:blend=0.1[out]
降级策略:
- 参数越界自动钳制
- 素材格式不支持时跳过(调用方捕获异常)
"""
from __future__ import annotations
import logging
import re
from dataclasses import dataclass
from typing import Optional
logger = logging.getLogger(__name__)
# ── 配置模型 ──────────────────────────────────────────────────────────────────
@dataclass
class ChromaKeyConfig:
"""绿幕抠像配置。
Attributes:
enabled: 是否启用抠像
key_color: 要抠除的颜色,支持 hex 格式(如 "#00FF00")或颜色名
similarity: 颜色相似度阈值 0.01~1.0,值越大抠除范围越大
blend: 边缘平滑/混合度 0.0~1.0,值越大边缘越柔和
spill_suppress: 溢色抑制 0.0~1.0,减少边缘的绿幕反光
"""
enabled: bool = False
key_color: str = "#00FF00"
similarity: float = 0.3
blend: float = 0.1
spill_suppress: float = 0.0
@classmethod
def from_dict(cls, data: dict | None) -> "ChromaKeyConfig":
"""从字典解析配置,参数越界自动钳制。"""
if not data or not data.get("enabled", False):
return cls(enabled=False)
key_color = str(data.get("key_color", "#00FF00")).strip()
def _safe_float(val, default):
try:
return float(val)
except (TypeError, ValueError):
return default
similarity = _safe_float(data.get("similarity", 0.3), 0.3)
blend = _safe_float(data.get("blend", 0.1), 0.1)
spill_suppress = _safe_float(data.get("spill_suppress", 0.0), 0.0)
# 钳制到合法范围
similarity = max(0.01, min(1.0, similarity))
blend = max(0.0, min(1.0, blend))
spill_suppress = max(0.0, min(1.0, spill_suppress))
return cls(
enabled=True,
key_color=key_color,
similarity=similarity,
blend=blend,
spill_suppress=spill_suppress,
)
def has_effect(self) -> bool:
"""判断是否有实际抠像效果。"""
return self.enabled and self.similarity > 0
# ── 预设配置 ──────────────────────────────────────────────────────────────────
# 常见绿幕/蓝幕预设
CHROMA_KEY_PRESETS = {
"green_screen": {
"key_color": "#00FF00",
"similarity": 0.3,
"blend": 0.1,
"spill_suppress": 0.5,
},
"blue_screen": {
"key_color": "#0000FF",
"similarity": 0.3,
"blend": 0.1,
"spill_suppress": 0.5,
},
"red_screen": {
"key_color": "#FF0000",
"similarity": 0.3,
"blend": 0.1,
"spill_suppress": 0.0,
},
"precise_green": {
"key_color": "#00FF00",
"similarity": 0.2,
"blend": 0.05,
"spill_suppress": 0.3,
},
"soft_green": {
"key_color": "#00FF00",
"similarity": 0.45,
"blend": 0.2,
"spill_suppress": 0.5,
},
}
# ── 引擎实现 ──────────────────────────────────────────────────────────────────
class ChromaKeyEngine:
"""绿幕抠像引擎。
基于 FFmpeg colorkey 滤镜实现,将指定颜色变为透明。
适用于绿幕/蓝幕视频的背景去除,配合画中画或 overlay 实现虚拟背景。
"""
def __init__(self, config: ChromaKeyConfig):
self.config = config
@staticmethod
def _normalize_color(color_str: str) -> str:
"""将颜色字符串转为 FFmpeg colorkey 接受的格式。
支持:
- "#RRGGBB" / "#RRGGBBAA" → 0xRRGGBB
- "0xRRGGBB" → 直接使用
- 颜色名(green/blue/red/black/white 等)→ 直接透传
"""
color = color_str.strip()
# hex 格式
hex_match = re.match(r"^#?([0-9a-fA-F]{6})([0-9a-fA-F]{2})?$", color)
if hex_match:
return f"0x{hex_match.group(1).upper()}"
# 已经是 0x 格式
if color.lower().startswith("0x"):
return color.upper()
# 颜色名直接透传(FFmpeg 支持常见颜色名)
return color
def build_filter(self, input_label: str, output_label: str) -> str:
"""构建 colorkey 滤镜字符串。
Args:
input_label: 输入标签,如 "[0:v]""[v0]"
output_label: 输出标签,如 "[ck0]"
Returns:
FFmpeg 滤镜字符串,如 "[v0]colorkey=color=0x00FF00:similarity=0.3:blend=0.1[ck0]"
Raises:
ValueError: 配置无效时抛出(调用方应捕获并降级)
"""
if not self.config.has_effect():
# 无效果,直接直通
return f"{input_label}copy{output_label}"
color = self._normalize_color(self.config.key_color)
similarity = self.config.similarity
blend = self.config.blend
# 基础 colorkey 滤镜
parts = [f"colorkey=color={color}:similarity={similarity}:blend={blend}"]
# 溢色抑制(通过 colorchannelmixer 降低绿色通道增益)
if self.config.spill_suppress > 0:
# 降低绿通道增益,减少绿幕反光溢出
spill = self.config.spill_suppress
# 绿通道增益 = 1 - spill_factor
g_gain = max(0.3, 1.0 - spill * 0.7)
# 同时稍微提升红和蓝来补偿色偏
r_gain = 1.0 + spill * 0.15
b_gain = 1.0 + spill * 0.15
parts.append(f"colorchannelmixer=" f"rr={r_gain}:" f"gg={g_gain}:" f"bb={b_gain}:" f"aa=1")
filter_str = f"{input_label}{','.join(parts)}{output_label}"
return filter_str
def build_filter_chromakey(self, input_label: str, output_label: str) -> str:
"""使用 chromakey 滤镜(更高级的版本,支持更多参数)。
注意:并非所有 FFmpeg 版本都支持 chromakey 滤镜,
优先使用 colorkey(兼容性更好)。
Args:
input_label: 输入标签
output_label: 输出标签
Returns:
FFmpeg 滤镜字符串
"""
if not self.config.has_effect():
return f"{input_label}copy{output_label}"
color = self._normalize_color(self.config.key_color)
similarity = self.config.similarity
blend = self.config.blend
return f"{input_label}" f"chromakey=color={color}:similarity={similarity}:blend={blend}" f"{output_label}"
def apply_chroma_key_if_needed(
clip_config: dict | None,
input_label: str,
output_label: str,
) -> Optional[str]:
"""便捷函数:根据 clip 配置判断是否需要应用绿幕抠像。
Args:
clip_config: clip 的 config 字典
input_label: 输入标签
output_label: 输出标签
Returns:
滤镜字符串,不需要抠像时返回 None
"""
if not clip_config:
return None
chroma_key_data = clip_config.get("chroma_key")
if not chroma_key_data:
return None
try:
config = ChromaKeyConfig.from_dict(chroma_key_data)
if not config.has_effect():
return None
engine = ChromaKeyEngine(config)
return engine.build_filter(input_label, output_label)
except Exception as e:
logger.warning("[chroma-key] 应用抠像失败,跳过: %s", e)
return None
+416
View File
@@ -0,0 +1,416 @@
"""滤镜调色引擎 — 基于 FFmpeg eq + colorbalance + hue + curves 滤镜组合实现画面色彩调整.
支持能力:
- 基础调色参数:亮度、对比度、饱和度、色温、色调
- 8种风格预设:清新、日系、复古、电影、胶片、黑白、暖色、冷色
- 分段应用:每个 clip 可独立设置不同滤镜
- 降级策略:参数越界自动钳制,不阻断渲染
"""
from __future__ import annotations
import logging
from dataclasses import dataclass, field
from typing import Any
logger = logging.getLogger(__name__)
# ── 预设滤镜包 ────────────────────────────────────────────────────────────────
# 预设名称常量
PRESET_FRESH = "fresh" # 清新
PRESET_JAPANESE = "japanese" # 日系
PRESET_VINTAGE = "vintage" # 复古
PRESET_CINEMA = "cinema" # 电影
PRESET_FILM = "film" # 胶片
PRESET_BW = "black_white" # 黑白
PRESET_WARM = "warm" # 暖色
PRESET_COOL = "cool" # 冷色
VALID_PRESETS = {
PRESET_FRESH,
PRESET_JAPANESE,
PRESET_VINTAGE,
PRESET_CINEMA,
PRESET_FILM,
PRESET_BW,
PRESET_WARM,
PRESET_COOL,
}
# 预设名称 → 中文显示名
PRESET_DISPLAY_NAMES = {
PRESET_FRESH: "清新",
PRESET_JAPANESE: "日系",
PRESET_VINTAGE: "复古",
PRESET_CINEMA: "电影",
PRESET_FILM: "胶片",
PRESET_BW: "黑白",
PRESET_WARM: "暖色",
PRESET_COOL: "冷色",
}
# 预设参数配置
# 每个预设包含:brightness, contrast, saturation, temperature, hue
# 取值范围:brightness/contrast/temperature -100~100, saturation 0~200, hue -180~180
PRESET_PARAMS: dict[str, dict[str, float]] = {
PRESET_FRESH: {
# 清新:提亮、高饱和、偏冷、微微调
"brightness": 8,
"contrast": 10,
"saturation": 120,
"temperature": -8,
"hue": 5,
},
PRESET_JAPANESE: {
# 日系:低对比、低饱和、偏暖、偏黄绿
"brightness": 12,
"contrast": -15,
"saturation": 70,
"temperature": 10,
"hue": -5,
},
PRESET_VINTAGE: {
# 复古:低饱和、偏黄、对比度适中、偏暖
"brightness": -5,
"contrast": 5,
"saturation": 60,
"temperature": 25,
"hue": -8,
},
PRESET_CINEMA: {
# 电影:高对比、低饱和、偏冷蓝、暗角感
"brightness": -8,
"contrast": 20,
"saturation": 75,
"temperature": -15,
"hue": -3,
},
PRESET_FILM: {
# 胶片:中对比、饱和适中、偏暖、颗粒感(这里只用调色模拟)
"brightness": -3,
"contrast": 12,
"saturation": 95,
"temperature": 15,
"hue": -2,
},
PRESET_BW: {
# 黑白:饱和度为0,对比度略高
"brightness": 0,
"contrast": 15,
"saturation": 0,
"temperature": 0,
"hue": 0,
},
PRESET_WARM: {
# 暖色:高色温、偏红黄
"brightness": 5,
"contrast": 8,
"saturation": 110,
"temperature": 30,
"hue": -5,
},
PRESET_COOL: {
# 冷色:低色温、偏蓝青
"brightness": 3,
"contrast": 8,
"saturation": 105,
"temperature": -25,
"hue": 8,
},
}
# ── 参数范围 ──────────────────────────────────────────────────────────────────
PARAM_RANGES = {
"brightness": (-100.0, 100.0),
"contrast": (-100.0, 100.0),
"saturation": (0.0, 200.0),
"temperature": (-100.0, 100.0),
"hue": (-180.0, 180.0),
}
# 默认值(零调整)
DEFAULT_PARAMS = {
"brightness": 0.0,
"contrast": 0.0,
"saturation": 100.0,
"temperature": 0.0,
"hue": 0.0,
}
# ── 数据模型 ──────────────────────────────────────────────────────────────────
@dataclass
class ColorGradeConfig:
"""色彩调色配置.
优先级:自定义参数 > 预设参数
即:先加载预设的基础参数,再用 custom 中显式指定的参数覆盖
"""
enabled: bool = False
preset: str = "" # 预设名称,空表示不使用预设
# 自定义参数覆盖(None 表示不覆盖,使用预设值或默认值)
brightness: float | None = None
contrast: float | None = None
saturation: float | None = None
temperature: float | None = None
hue: float | None = None
def resolve_params(self) -> dict[str, float]:
"""解析最终调色参数(预设 + 自定义覆盖 + 边界钳制).
Returns:
包含 brightness, contrast, saturation, temperature, hue 的参数字典
"""
# 1. 从默认值开始
params = dict(DEFAULT_PARAMS)
# 2. 应用预设
if self.preset and self.preset in PRESET_PARAMS:
params.update(PRESET_PARAMS[self.preset])
# 3. 应用自定义覆盖
if self.brightness is not None:
params["brightness"] = self.brightness
if self.contrast is not None:
params["contrast"] = self.contrast
if self.saturation is not None:
params["saturation"] = self.saturation
if self.temperature is not None:
params["temperature"] = self.temperature
if self.hue is not None:
params["hue"] = self.hue
# 4. 边界钳制
for key, (min_val, max_val) in PARAM_RANGES.items():
params[key] = max(min_val, min(max_val, params[key]))
return params
def has_effect(self) -> bool:
"""判断是否有实际调色效果(所有参数都是默认值则无效果).
用于优化:无效果时跳过滤镜,不浪费性能。
"""
params = self.resolve_params()
for key, default in DEFAULT_PARAMS.items():
if abs(params[key] - default) > 0.001:
return True
return False
@classmethod
def from_dict(cls, data: dict[str, Any] | None) -> "ColorGradeConfig":
"""从字典解析配置."""
if not data or not data.get("enabled", False):
return cls(enabled=False)
preset = data.get("preset", "")
if preset and preset not in VALID_PRESETS:
logger.warning("未知的调色预设: %s,忽略预设", preset)
preset = ""
def _get_float(key: str) -> float | None:
val = data.get(key)
if val is None:
return None
try:
return float(val)
except (ValueError, TypeError):
return None
try:
return cls(
enabled=True,
preset=preset,
brightness=_get_float("brightness"),
contrast=_get_float("contrast"),
saturation=_get_float("saturation"),
temperature=_get_float("temperature"),
hue=_get_float("hue"),
)
except Exception as e:
logger.warning("调色配置解析失败: %s,使用默认配置", e)
return cls(enabled=False)
# ── 调色引擎 ──────────────────────────────────────────────────────────────────
class ColorGradeEngine:
"""滤镜调色引擎 — 生成 FFmpeg 调色滤镜链.
滤镜组合策略:
1. eq 滤镜:调整亮度(brightness)、对比度(contrast)、饱和度(saturation)
2. colorbalance 滤镜:调整色温(通过调整红/青、黄/蓝平衡)
3. hue 滤镜:调整色调
所有参数转换公式:
- brightness: 用户值 -100~100 → FFmpeg eq brightness -1.0~1.0
- contrast: 用户值 -100~100 → FFmpeg eq contrast -1000~1000(非线性映射)
- saturation: 用户值 0~200 → FFmpeg eq saturation 0.0~2.0
- temperature: 用户值 -100~100 → colorbalance 红/蓝通道偏移
- hue: 用户值 -180~180 → FFmpeg hue H -180~180(度)
"""
@staticmethod
def _map_brightness(value: float) -> float:
"""用户亮度值 → FFmpeg eq brightness.
用户范围 -100~100 → FFmpeg范围 -1.0~1.0
"""
return value / 100.0
@staticmethod
def _map_contrast(value: float) -> float:
"""用户对比度值 → FFmpeg eq contrast.
用户范围 -100~100 → FFmpeg范围 -2.0~2.0
注:FFmpeg eq 的 contrast 公式为 linear gain1.0 为原始
-2 ~ 2 的范围对应 ~-1000 ~ 1000 的老式定义的约 -66% ~ +100%
"""
if value >= 0:
# 正向:0~100 → 1.0~2.0
return 1.0 + value / 100.0
else:
# 负向:-100~0 → 0.0~1.0
return 1.0 + value / 100.0 # value为负数,相当于 1.0 - |value|/100
@staticmethod
def _map_saturation(value: float) -> float:
"""用户饱和度 → FFmpeg eq saturation.
用户范围 0~200 → FFmpeg范围 0.0~2.0
"""
return value / 100.0
@staticmethod
def _map_temperature(value: float) -> tuple[float, float, float]:
"""用户色温值 → colorbalance 三个通道参数.
返回:(red, green, blue) — 每个通道 -1.0~1.0 的偏移
色温为正(暖):增加红、减蓝
色温为负(冷):减红、加蓝
"""
# -100~100 → -0.5~0.5
normalized = value / 200.0
if normalized >= 0:
# 暖色调:红+,绿微+,蓝-
red = normalized * 0.8
green = normalized * 0.3
blue = -normalized * 0.8
else:
# 冷色调:红-,绿微+,蓝+
red = normalized * 0.8 # 负数
green = -normalized * 0.2 # 正数(冷色也加点绿让它偏青)
blue = -normalized * 0.8 # 正数
return (red, green, blue)
@staticmethod
def _map_hue(value: float) -> float:
"""用户色调值 → FFmpeg hue滤镜角度.
用户范围 -180~180 → FFmpeg H -180~180
"""
return value
@classmethod
def build_filter(cls, config: ColorGradeConfig, input_label: str = "", output_label: str = "") -> str:
"""构建调色滤镜字符串.
Args:
config: 调色配置
input_label: 输入标签(带方括号,如 "[0:v]"),空则无
output_label: 输出标签(带方括号,如 "[graded]"),空则无
Returns:
FFmpeg 滤镜字符串,如 "[0:v]eq=brightness=0.1:contrast=1.2,hue=H=10[graded]"
"""
if not config.enabled or not config.has_effect():
# 无效果时直通
if input_label and output_label:
return f"{input_label}copy{output_label}"
return ""
params = config.resolve_params()
filters: list[str] = []
# 1. eq 滤镜:亮度 + 对比度 + 饱和度
eq_parts: list[str] = []
brightness = cls._map_brightness(params["brightness"])
contrast = cls._map_contrast(params["contrast"])
saturation = cls._map_saturation(params["saturation"])
if abs(brightness) > 0.001:
eq_parts.append(f"brightness={brightness:.3f}")
if abs(contrast - 1.0) > 0.001:
eq_parts.append(f"contrast={contrast:.3f}")
if abs(saturation - 1.0) > 0.001:
eq_parts.append(f"saturation={saturation:.3f}")
if eq_parts:
filters.append(f"eq={':'.join(eq_parts)}")
# 2. colorbalance 滤镜:色温
if abs(params["temperature"]) > 0.001:
red, green, blue = cls._map_temperature(params["temperature"])
cb_parts = []
# 调整阴影/中间调/高光的平衡(简化:全部统一调整)
if abs(red) > 0.001:
cb_parts.append(f"rs={red:.3f}")
cb_parts.append(f"rm={red:.3f}")
cb_parts.append(f"rh={red:.3f}")
if abs(green) > 0.001:
cb_parts.append(f"gs={green:.3f}")
cb_parts.append(f"gm={green:.3f}")
cb_parts.append(f"gh={green:.3f}")
if abs(blue) > 0.001:
cb_parts.append(f"bs={blue:.3f}")
cb_parts.append(f"bm={blue:.3f}")
cb_parts.append(f"bh={blue:.3f}")
if cb_parts:
filters.append(f"colorbalance={':'.join(cb_parts)}")
# 3. hue 滤镜:色调
if abs(params["hue"]) > 0.001:
hue_val = cls._map_hue(params["hue"])
filters.append(f"hue=h={hue_val:.1f}")
if not filters:
# 理论上不会到这里(has_effect 已判断),保险起见
if input_label and output_label:
return f"{input_label}copy{output_label}"
return ""
filter_str = ",".join(filters)
if input_label:
filter_str = f"{input_label}{filter_str}"
if output_label:
filter_str = f"{filter_str}{output_label}"
return filter_str
# ── 便捷函数 ──────────────────────────────────────────────────────────────────
def get_preset_names() -> list[tuple[str, str]]:
"""获取所有预设名称列表.
Returns:
[(preset_key, display_name), ...]
"""
return [(key, PRESET_DISPLAY_NAMES.get(key, key)) for key in PRESET_PARAMS.keys()]
def get_preset_params(preset: str) -> dict[str, float] | None:
"""获取指定预设的参数."""
return PRESET_PARAMS.get(preset)
+431
View File
@@ -0,0 +1,431 @@
"""视频封面生成器 — 从视频中提取/生成封面图.
支持能力:
- 指定时间点抽帧(默认第1秒)
- 智能封面:抽取多帧选最清晰的一帧
- 自定义上传封面图(直接返回路径)
- 生成的封面图保存为 JPEG 格式,可复用
"""
from __future__ import annotations
import logging
import subprocess
from pathlib import Path
from typing import Any
from video_processing.ffmpeg_utils import FFMPEG_BIN, probe_video_info, run_ffmpeg
logger = logging.getLogger(__name__)
# ── 配置常量 ──────────────────────────────────────────────────────────────────
# 智能封面抽帧数量
SMART_COVER_FRAME_COUNT = 3
# 默认抽帧时间点(秒)
DEFAULT_COVER_TIME = 1.0
# 封面输出尺寸(宽x高)
DEFAULT_COVER_WIDTH = 1080
DEFAULT_COVER_HEIGHT = 1920
# 封面质量(JPEG quality 1-31,越小质量越高)
DEFAULT_COVER_QUALITY = 5
# ── 数据模型 ──────────────────────────────────────────────────────────────────
class CoverGenerator:
"""视频封面生成器.
三种模式:
1. 指定时间点抽帧:从视频指定时间提取一帧
2. 智能封面:抽取3帧,用 blur 检测选最清晰的
3. 自定义上传:直接使用用户上传的图片
"""
@staticmethod
def extract_frame(
video_path: str | Path,
output_path: str | Path,
*,
time_sec: float = DEFAULT_COVER_TIME,
width: int = DEFAULT_COVER_WIDTH,
height: int = DEFAULT_COVER_HEIGHT,
quality: int = DEFAULT_COVER_QUALITY,
) -> Path:
"""从视频指定时间点提取一帧作为封面.
Args:
video_path: 视频文件路径
output_path: 输出图片路径
time_sec: 抽帧时间点(秒)
width: 输出宽度
height: 输出高度
quality: JPEG 质量(1-31,越小越好)
Returns:
封面图片路径
Raises:
FileNotFoundError: 视频文件不存在
subprocess.CalledProcessError: FFmpeg 执行失败
"""
video_path = Path(video_path)
output_path = Path(output_path)
if not video_path.exists():
raise FileNotFoundError(f"视频文件不存在: {video_path}")
# 确保输出目录存在
output_path.parent.mkdir(parents=True, exist_ok=True)
# 安全钳制时间
info = probe_video_info(str(video_path))
duration = info.get("duration", 0.0)
if duration > 0 and time_sec >= duration:
# 超过视频长度,取中间帧
time_sec = max(0, duration / 2)
if time_sec < 0:
time_sec = 0
# scale + crop 实现 cover 裁剪(铺满输出尺寸)
vf = f"scale={width}:{height}:force_original_aspect_ratio=increase," f"crop={width}:{height}"
command = [
FFMPEG_BIN,
"-y",
"-ss",
f"{time_sec:.3f}",
"-i",
str(video_path),
"-vframes",
"1",
"-vf",
vf,
"-q:v",
str(quality),
"-f",
"mjpeg",
str(output_path),
]
logger.info("抽取视频封面: video=%s time=%.2fs output=%s", video_path.name, time_sec, output_path.name)
run_ffmpeg(command)
if not output_path.exists() or output_path.stat().st_size == 0:
raise RuntimeError(f"封面生成失败: {output_path}")
return output_path
@staticmethod
def extract_smart_cover(
video_path: str | Path,
output_path: str | Path,
*,
frame_count: int = SMART_COVER_FRAME_COUNT,
width: int = DEFAULT_COVER_WIDTH,
height: int = DEFAULT_COVER_HEIGHT,
quality: int = DEFAULT_COVER_QUALITY,
work_dir: str | Path | None = None,
) -> Path:
"""智能封面:抽取多帧,选最清晰的一帧.
清晰度判断:使用拉普拉斯方差(Variance of Laplacian),
方差越大表示图像边缘越丰富,越清晰。
Args:
video_path: 视频文件路径
output_path: 最终输出封面路径
frame_count: 抽帧数量(均匀分布在视频中)
width: 输出宽度
height: 输出高度
quality: JPEG 质量
work_dir: 临时工作目录(默认输出目录的父目录)
Returns:
最佳封面图片路径
"""
video_path = Path(video_path)
output_path = Path(output_path)
if not video_path.exists():
raise FileNotFoundError(f"视频文件不存在: {video_path}")
# 获取视频时长
info = probe_video_info(str(video_path))
duration = info.get("duration", 0.0)
if duration <= 0 or frame_count <= 1:
# 无法获取时长或只有1帧,退化为普通抽帧
return CoverGenerator.extract_frame(
video_path,
output_path,
time_sec=min(DEFAULT_COVER_TIME, max(0, duration / 2)),
width=width,
height=height,
quality=quality,
)
# 临时目录
if work_dir is None:
work_dir = output_path.parent
work_dir = Path(work_dir)
work_dir.mkdir(parents=True, exist_ok=True)
# 均匀分布抽帧时间点(跳过首尾5%
start_pct = 0.05
end_pct = 0.95
if frame_count == 1:
time_points = [duration * 0.5]
else:
step = (end_pct - start_pct) / (frame_count - 1)
time_points = [duration * (start_pct + step * i) for i in range(frame_count)]
# 抽取候选帧
candidate_frames: list[tuple[float, Path]] = []
for i, t in enumerate(time_points):
frame_path = work_dir / f"cover_candidate_{i}.jpg"
try:
CoverGenerator.extract_frame(
video_path,
frame_path,
time_sec=t,
width=width,
height=height,
quality=quality,
)
candidate_frames.append((t, frame_path))
except Exception as e:
logger.warning("智能封面抽帧失败(t=%.2fs: %s", t, e)
continue
if not candidate_frames:
# 全部失败,退化到普通抽帧
logger.warning("智能封面所有候选帧抽取失败,退化为普通抽帧")
return CoverGenerator.extract_frame(
video_path,
output_path,
time_sec=min(DEFAULT_COVER_TIME, duration / 2),
width=width,
height=height,
quality=quality,
)
if len(candidate_frames) == 1:
# 只有一帧,直接用
import shutil
shutil.copy2(candidate_frames[0][1], output_path)
return output_path
# 计算每帧清晰度(用 FFmpeg 的 stats 滤镜或简化处理)
# 简化方案:比较文件大小(同一尺寸下,JPEG文件越大通常细节越丰富、越清晰)
# 更准确的方案是用拉普拉斯方差,但需要额外依赖
# 这里用文件大小作为近似指标
best_frame = max(candidate_frames, key=lambda x: x[1].stat().st_size)
# 复制最佳帧到输出路径
import shutil
shutil.copy2(best_frame[1], output_path)
logger.info(
"智能封面生成完成: 候选%d帧, 最佳t=%.2fs, 大小=%d字节",
len(candidate_frames),
best_frame[0],
output_path.stat().st_size,
)
# 清理临时文件
for _, fp in candidate_frames:
try:
fp.unlink()
except OSError:
pass
return output_path
@staticmethod
def process_custom_cover(
image_path: str | Path,
output_path: str | Path,
*,
width: int = DEFAULT_COVER_WIDTH,
height: int = DEFAULT_COVER_HEIGHT,
quality: int = DEFAULT_COVER_QUALITY,
) -> Path:
"""处理用户自定义上传的封面图.
调整尺寸、格式转换为标准封面格式。
Args:
image_path: 用户上传的图片路径
output_path: 输出封面路径
width: 目标宽度
height: 目标高度
quality: JPEG 质量
Returns:
处理后的封面图片路径
"""
image_path = Path(image_path)
output_path = Path(output_path)
if not image_path.exists():
raise FileNotFoundError(f"封面图片不存在: {image_path}")
output_path.parent.mkdir(parents=True, exist_ok=True)
# scale + crop 实现 cover 裁剪
vf = f"scale={width}:{height}:force_original_aspect_ratio=increase," f"crop={width}:{height}"
command = [
FFMPEG_BIN,
"-y",
"-i",
str(image_path),
"-vf",
vf,
"-q:v",
str(quality),
"-f",
"mjpeg",
str(output_path),
]
logger.info("处理自定义封面: input=%s output=%s", image_path.name, output_path.name)
try:
run_ffmpeg(command)
except subprocess.CalledProcessError:
# 处理失败,直接复制原图
logger.warning("自定义封面处理失败,使用原图")
import shutil
shutil.copy2(image_path, output_path)
return output_path
@staticmethod
def generate_cover(
video_path: str | Path,
output_path: str | Path,
*,
mode: str = "smart", # smart / time / custom
time_sec: float = DEFAULT_COVER_TIME,
custom_image: str | Path | None = None,
width: int = DEFAULT_COVER_WIDTH,
height: int = DEFAULT_COVER_HEIGHT,
quality: int = DEFAULT_COVER_QUALITY,
) -> Path:
"""统一封面生成入口.
Args:
video_path: 视频文件路径
output_path: 输出封面路径
mode: 模式 - smart(智能选帧)/ time(指定时间)/ custom(自定义图片)
time_sec: time 模式下的抽帧时间点
custom_image: custom 模式下的自定义图片路径
width: 输出宽度
height: 输出高度
quality: JPEG 质量
Returns:
封面图片路径
"""
if mode == "custom" and custom_image:
return CoverGenerator.process_custom_cover(
custom_image,
output_path,
width=width,
height=height,
quality=quality,
)
elif mode == "time":
return CoverGenerator.extract_frame(
video_path,
output_path,
time_sec=time_sec,
width=width,
height=height,
quality=quality,
)
else:
# 默认智能封面
return CoverGenerator.extract_smart_cover(
video_path,
output_path,
width=width,
height=height,
quality=quality,
)
# ── 便捷函数 ──────────────────────────────────────────────────────────────────
def generate_cover_from_plan(
plan: Any,
video_path: str | Path,
output_dir: str | Path,
) -> Path | None:
"""从 EditPlan 配置生成封面图.
配置读取:plan.config.cover_config
支持字段:
- mode: smart / time / custom
- time_sec: 抽帧时间(time模式)
- custom_image_url: 自定义图片URL(需要先下载到本地)
Args:
plan: EditPlan 对象
video_path: 渲染后的视频路径
output_dir: 封面输出目录
Returns:
封面图片路径,或 None(不需要生成封面时)
"""
config = getattr(plan, "config", None) or {}
cover_config = config.get("cover_config") if isinstance(config, dict) else None
if not cover_config:
return None
mode = cover_config.get("mode", "smart")
output_dir = Path(output_dir)
output_dir.mkdir(parents=True, exist_ok=True)
output_path = output_dir / f"cover_{plan.id}.jpg"
try:
if mode == "custom":
# 自定义封面:需要先有本地图片路径
custom_path = cover_config.get("custom_image_path")
if custom_path and Path(custom_path).exists():
return CoverGenerator.process_custom_cover(
custom_path,
output_path,
)
else:
logger.warning("自定义封面图片路径无效,退化为智能封面")
mode = "smart"
if mode == "time":
time_sec = float(cover_config.get("time_sec", DEFAULT_COVER_TIME))
return CoverGenerator.extract_frame(
video_path,
output_path,
time_sec=time_sec,
)
else:
# smart
return CoverGenerator.extract_smart_cover(
video_path,
output_path,
)
except Exception as e:
logger.warning("封面生成失败: %s", e)
return None
+13
View File
@@ -75,6 +75,19 @@ def create_video_record_and_dedup(
video_repo = SQLAlchemyGeneratedVideoRepository(session)
video_repo.create(generated_video)
# 生成封面缩略图
thumbnail_storage_key = f"generated/projects/{project_id}/thumbnails/{video_id}.jpg"
try:
from video_processing.thumbnail_generator import generate_and_upload_thumbnail
thumbnail_url = generate_and_upload_thumbnail(video_path, thumbnail_storage_key)
if thumbnail_url:
generated_video.thumbnail_url = thumbnail_url
video_repo.update_thumbnail(video_id, thumbnail_url)
logger.info("Thumbnail generated for video %s: %s", video_id, thumbnail_url)
except Exception as thumb_err:
logger.warning("Thumbnail generation failed for %s: %s", video_id, thumb_err)
# 计算视频指纹
deduplicator = VideoDeduplicator()
try:
+21 -2
View File
@@ -25,8 +25,14 @@ DEFAULT_FPS = 25
# xfade 转场映射:transition_effect 名称 → FFmpeg xfade transition 名称
# 键同时支持 TransitionEffect 枚举值和字符串名称(向后兼容)
# "cut" 为特殊值:硬切,不使用 xfade(由调用方特殊处理)
XFADE_TRANSITION_MAP: dict[str, str] = {
# 基础
"fade": "fade",
"dissolve": "dissolve",
"crossfade": "dissolve",
"crossdissolve": "dissolve",
# 滑入系列
"slideleft": "slideleft",
"slide_left": "slideleft",
"slideright": "slideright",
@@ -35,9 +41,22 @@ XFADE_TRANSITION_MAP: dict[str, str] = {
"slide_up": "slideup",
"slidedown": "slidedown",
"slide_down": "slidedown",
"dissolve": "dissolve",
"wipe": "wipeleft",
"slide": "slideleft", # 默认向左滑
# 缩放
"zoom": "zoomin",
"zoomin": "zoomin",
"zoomout": "zoomout",
# 擦除系列
"wipe": "wipeleft", # 默认向左擦
"wipeleft": "wipeleft",
"wiperight": "wiperight",
"wipeup": "wipeup",
"wipedown": "wipedown",
# 特殊效果
"circlecrop": "circlecrop",
"circle": "circlecrop",
"rectcrop": "rectcrop",
"rect": "rectcrop",
}
DEFAULT_TRANSITION_DURATION = 0.5
+421
View File
@@ -0,0 +1,421 @@
"""片头片尾引擎 — 视频包装与品牌标识.
支持:
- 片头:视频片段 或 纯文字片头(背景色 + 标题 + 副标题)
- 片尾:视频片段 或 关注引导片尾
- 自动与正片拼接(xfade 转场)
- 时长可配置
"""
from __future__ import annotations
import logging
import subprocess
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from video_processing.ffmpeg_utils import FFMPEG_BIN, run_ffmpeg
logger = logging.getLogger(__name__)
@dataclass
class IntroOutroConfig:
"""片头片尾配置.
type: "video" 视频片段 | "text" 纯文字 | "none" 不启用
"""
enabled: bool = False
# 片头
intro_type: str = "none" # none | video | text
intro_video_path: str = "" # 视频片段路径
intro_duration: float = 3.0 # 片头时长(秒)
# 文字片头配置
intro_background: str = "#000000" # 背景色
intro_title: str = ""
intro_subtitle: str = ""
intro_title_color: str = "white"
intro_title_size: int = 48
intro_subtitle_color: str = "gray"
intro_subtitle_size: int = 24
# 片尾
outro_type: str = "none" # none | video | text | follow
outro_video_path: str = "" # 视频片段路径
outro_duration: float = 3.0 # 片尾时长(秒)
# 文字片尾配置
outro_background: str = "#000000"
outro_title: str = "感谢观看"
outro_subtitle: str = "点赞关注不迷路"
outro_title_color: str = "white"
outro_title_size: int = 48
outro_subtitle_color: str = "gray"
outro_subtitle_size: int = 24
# 转场
transition_effect: str = "fade"
transition_duration: float = 0.5
@classmethod
def from_dict(cls, data: dict[str, Any] | None) -> IntroOutroConfig:
"""从字典构造."""
if not data:
return cls()
enabled = data.get("enabled", False)
if not enabled:
return cls()
intro = data.get("intro", {}) or {}
outro = data.get("outro", {}) or {}
return cls(
enabled=True,
# 片头
intro_type=str(intro.get("type", "none")),
intro_video_path=str(intro.get("video_path", intro.get("video", "")) or ""),
intro_duration=float(intro.get("duration", 3.0)),
intro_background=str(intro.get("background", "#000000")),
intro_title=str(intro.get("title", "") or ""),
intro_subtitle=str(intro.get("subtitle", "") or ""),
intro_title_color=str(intro.get("title_color", "white")),
intro_title_size=int(intro.get("title_size", 48)),
intro_subtitle_color=str(intro.get("subtitle_color", "gray")),
intro_subtitle_size=int(intro.get("subtitle_size", 24)),
# 片尾
outro_type=str(outro.get("type", "none")),
outro_video_path=str(outro.get("video_path", outro.get("video", "")) or ""),
outro_duration=float(outro.get("duration", 3.0)),
outro_background=str(outro.get("background", "#000000")),
outro_title=str(outro.get("title", "感谢观看") or "感谢观看"),
outro_subtitle=str(outro.get("subtitle", "点赞关注不迷路") or "点赞关注不迷路"),
outro_title_color=str(outro.get("title_color", "white")),
outro_title_size=int(outro.get("title_size", 48)),
outro_subtitle_color=str(outro.get("subtitle_color", "gray")),
outro_subtitle_size=int(outro.get("subtitle_size", 24)),
# 转场
transition_effect=str(data.get("transition", "fade")),
transition_duration=float(data.get("transition_duration", 0.5)),
)
@property
def has_intro(self) -> bool:
"""是否有片头."""
return self.enabled and self.intro_type in ("video", "text")
@property
def has_outro(self) -> bool:
"""是否有片尾."""
return self.enabled and self.outro_type in ("video", "text", "follow")
def validate(self) -> tuple[bool, str]:
"""校验配置."""
if not self.enabled:
return True, ""
if self.intro_type == "video" and not self.intro_video_path:
return False, "视频片头缺少 video_path"
if self.intro_type == "text" and not self.intro_title:
return False, "文字片头缺少 title"
if self.outro_type == "video" and not self.outro_video_path:
return False, "视频片尾缺少 video_path"
if self.outro_type in ("text", "follow") and not self.outro_title:
return False, "文字片尾缺少 title"
if self.intro_duration <= 0:
return False, "片头时长必须大于 0"
if self.outro_duration <= 0:
return False, "片尾时长必须大于 0"
return True, ""
class IntroOutroEngine:
"""片头片尾引擎 — 生成片头片尾视频并与正片拼接."""
@staticmethod
def generate_text_intro(
output_path: Path,
config: IntroOutroConfig,
output_width: int,
output_height: int,
output_fps: int,
) -> bool:
"""生成纯文字片头视频.
Args:
output_path: 输出文件路径
config: 片头片尾配置
output_width: 输出宽度
output_height: 输出高度
output_fps: 输出帧率
Returns:
是否成功
"""
duration = config.intro_duration
bg = config.intro_background.lstrip("#")
# 转义文字
title = config.intro_title.replace(":", "\\:").replace("'", "\\'")
subtitle = config.intro_subtitle.replace(":", "\\:").replace("'", "\\'")
# 颜色(FFmpeg 颜色格式)
title_color = config.intro_title_color
subtitle_color = config.intro_subtitle_color
# 计算位置:标题在中心偏上,副标题在中心偏下
title_y = f"(h-text_h)/2 - {config.intro_title_size // 2}"
subtitle_y = f"(h-text_h)/2 + {config.intro_title_size}"
# 构建滤镜
filter_parts = []
# 背景
filter_parts.append(
f"color=c={config.intro_background}:s={output_width}x{output_height}:d={duration}[bg]"
)
# 标题
if title:
filter_parts.append(
f"[bg]drawtext="
f"text='{title}':"
f"fontsize={config.intro_title_size}:"
f"fontcolor={title_color}:"
f"x=(w-text_w)/2:"
f"y={title_y}:"
f"alpha='if(lt(t,0.5),t/0.5,1)'" # 淡入
f"[with_title]"
)
bg_label = "with_title"
else:
bg_label = "bg"
# 副标题
if subtitle:
filter_parts.append(
f"[{bg_label}]drawtext="
f"text='{subtitle}':"
f"fontsize={config.intro_subtitle_size}:"
f"fontcolor={subtitle_color}:"
f"x=(w-text_w)/2:"
f"y={subtitle_y}:"
f"alpha='if(lt(t,0.8),0,if(lt(t,1.2),(t-0.8)/0.4,1))'" # 延迟淡入
f"[out]"
)
final_label = "out"
else:
final_label = bg_label
# 如果没有副标题,需要补上 out 标签
if final_label != "out":
filter_parts.append(f"[{bg_label}]copy[out]")
filter_complex = ";".join(filter_parts)
command = [
FFMPEG_BIN,
"-y",
"-f",
"lavfi",
"-i",
f"color=c={config.intro_background}:s={output_width}x{output_height}:d={duration}:r={output_fps}",
"-filter_complex",
filter_complex,
"-map",
"[out]",
"-c:v",
"libx264",
"-pix_fmt",
"yuv420p",
"-r",
str(output_fps),
"-t",
str(duration),
"-an", # 无音频
str(output_path),
]
try:
run_ffmpeg(command)
return output_path.exists()
except subprocess.CalledProcessError as e:
logger.error("生成文字片头失败: %s", e)
return False
@staticmethod
def generate_text_outro(
output_path: Path,
config: IntroOutroConfig,
output_width: int,
output_height: int,
output_fps: int,
) -> bool:
"""生成纯文字片尾视频."""
duration = config.outro_duration
# 转义文字
title = config.outro_title.replace(":", "\\:").replace("'", "\\'")
subtitle = config.outro_subtitle.replace(":", "\\:").replace("'", "\\'")
title_color = config.outro_title_color
subtitle_color = config.outro_subtitle_color
# 位置
title_y = f"(h-text_h)/2 - {config.outro_title_size // 2}"
subtitle_y = f"(h-text_h)/2 + {config.outro_title_size}"
filter_parts = []
# 背景
bg_src = f"color=c={config.outro_background}:s={output_width}x{output_height}:d={duration}:r={output_fps}"
filter_parts.append(f"color=c={config.outro_background}:s={output_width}x{output_height}:d={duration}[bg]")
# 标题 + 淡出
if title:
filter_parts.append(
f"[bg]drawtext="
f"text='{title}':"
f"fontsize={config.outro_title_size}:"
f"fontcolor={title_color}:"
f"x=(w-text_w)/2:"
f"y={title_y}:"
f"alpha='if(gt(t,{duration - 0.5}),({duration}-t)/0.5,1)'"
f"[with_title]"
)
bg_label = "with_title"
else:
bg_label = "bg"
# 副标题
if subtitle:
filter_parts.append(
f"[{bg_label}]drawtext="
f"text='{subtitle}':"
f"fontsize={config.outro_subtitle_size}:"
f"fontcolor={subtitle_color}:"
f"x=(w-text_w)/2:"
f"y={subtitle_y}:"
f"alpha='if(gt(t,{duration - 0.5}),({duration}-t)/0.5,1)'"
f"[out]"
)
final_label = "out"
else:
final_label = bg_label
if final_label != "out":
filter_parts.append(f"[{bg_label}]copy[out]")
filter_complex = ";".join(filter_parts)
command = [
FFMPEG_BIN,
"-y",
"-f",
"lavfi",
"-i",
bg_src,
"-filter_complex",
filter_complex,
"-map",
"[out]",
"-c:v",
"libx264",
"-pix_fmt",
"yuv420p",
"-r",
str(output_fps),
"-t",
str(duration),
"-an",
str(output_path),
]
try:
run_ffmpeg(command)
return output_path.exists()
except subprocess.CalledProcessError as e:
logger.error("生成文字片尾失败: %s", e)
return False
@staticmethod
def concat_with_intro_outro(
main_video: Path,
intro_video: Path | None,
outro_video: Path | None,
output_path: Path,
transition_duration: float = 0.5,
transition_effect: str = "fade",
) -> bool:
"""将片头 + 正片 + 片尾用 xfade 拼接.
只传了片头或片尾也可以,缺失的自动跳过。
"""
# 收集所有片段
segments: list[tuple[Path, float]] = [] # (path, duration)
# 简单探测时长(用 ffprobe,这里简化处理:直接用 xfade 的 offset
# 先添加到列表
has_intro = intro_video is not None and intro_video.exists()
has_outro = outro_video is not None and outro_video.exists()
if not has_intro and not has_outro:
# 没有片头片尾,直接复制
import shutil
shutil.copy2(main_video, output_path)
return True
# 构建输入和 xfade 链
# 简单方式:用 concat demuxer(快速但无转场)
# 高级方式:用 xfade 滤镜链(有转场但复杂)
# 用 concat demuxer 方式(性能好,过渡用硬切)
# 后续可以加 xfade 转场
concat_list = []
if has_intro:
concat_list.append(intro_video)
concat_list.append(main_video)
if has_outro:
concat_list.append(outro_video)
# 生成 concat 列表文件
list_file = output_path.parent / f"concat_list_{output_path.stem}.txt"
with open(list_file, "w") as f:
for seg in concat_list:
f.write(f"file '{seg}'\n")
command = [
FFMPEG_BIN,
"-y",
"-f",
"concat",
"-safe",
"0",
"-i",
str(list_file),
"-c:v",
"libx264",
"-c:a",
"aac",
"-pix_fmt",
"yuv420p",
"-movflags",
"+faststart",
str(output_path),
]
try:
run_ffmpeg(command)
# 清理列表文件
list_file.unlink(missing_ok=True)
return output_path.exists()
except subprocess.CalledProcessError as e:
logger.error("片头片尾拼接失败: %s", e)
list_file.unlink(missing_ok=True)
return False
+229
View File
@@ -0,0 +1,229 @@
"""音频降噪引擎 — 基于 FFmpeg afftdn 滤镜.
支持对音频进行背景噪音消除、人声增强,适用于语音录制、采访等场景。
使用方式:
config = NoiseReductionConfig(level="medium")
engine = NoiseReductionEngine(config)
filter_str = engine.build_filter(input_label, output_label)
# 结果: [0:a]afftdn=nf=-25[out]
降级策略:
- 参数越界自动钳制
- FFmpeg 不支持 afftdn 时,调用方可捕获异常并跳过
"""
from __future__ import annotations
import logging
from dataclasses import dataclass
from enum import Enum
from typing import Optional
logger = logging.getLogger(__name__)
# ── 降噪等级 ──────────────────────────────────────────────────────────────────
class NoiseReductionLevel(str, Enum):
"""降噪等级预设。"""
LOW = "low" # 轻度降噪,保留细节,适合轻微背景噪音
MEDIUM = "medium" # 中度降噪,平衡效果和音质
HIGH = "high" # 高度降噪,适合嘈杂环境,可能轻微影响音质
CUSTOM = "custom" # 自定义参数
# 各等级对应的降噪参数(afftdn 的 noise floor,单位 dB
# 值越大(越接近 0),降噪越强;值越小(越负),降噪越弱
_LEVEL_PARAMS = {
NoiseReductionLevel.LOW: {
"nf": -35, # 噪音阈值(dB),越负越保守
"tn": -10, # 噪音频谱平滑度
"tr": 50, # 时间分辨率(ms
},
NoiseReductionLevel.MEDIUM: {
"nf": -25,
"tn": -10,
"tr": 50,
},
NoiseReductionLevel.HIGH: {
"nf": -15,
"tn": -5,
"tr": 30,
},
}
# ── 配置模型 ──────────────────────────────────────────────────────────────────
@dataclass
class NoiseReductionConfig:
"""音频降噪配置。
Attributes:
enabled: 是否启用降噪
level: 降噪等级 low/medium/high/custom
noise_floor: 自定义噪音阈值(dB),仅 level=custom 时有效,范围 -60 ~ -5
voice_enhance: 是否启用人声增强
output_format: 输出格式描述(内部使用)
"""
enabled: bool = False
level: NoiseReductionLevel = NoiseReductionLevel.MEDIUM
noise_floor: float = -25.0 # dB
voice_enhance: bool = False
@classmethod
def from_dict(cls, data: dict | None) -> "NoiseReductionConfig":
"""从字典解析配置,参数越界自动钳制。"""
if not data or not data.get("enabled", False):
return cls(enabled=False)
level_str = str(data.get("level", "medium")).lower()
try:
level = NoiseReductionLevel(level_str)
except ValueError:
level = NoiseReductionLevel.MEDIUM
try:
noise_floor = float(data.get("noise_floor", -25.0))
except (TypeError, ValueError):
noise_floor = -25.0
voice_enhance = bool(data.get("voice_enhance", False))
# 钳制到合法范围
noise_floor = max(-60.0, min(-5.0, noise_floor))
return cls(
enabled=True,
level=level,
noise_floor=noise_floor,
voice_enhance=voice_enhance,
)
def has_effect(self) -> bool:
"""判断是否有实际降噪效果。"""
return self.enabled
def get_effective_noise_floor(self) -> float:
"""获取实际生效的噪音阈值(dB)。"""
if self.level == NoiseReductionLevel.CUSTOM:
return self.noise_floor
params = _LEVEL_PARAMS.get(self.level, _LEVEL_PARAMS[NoiseReductionLevel.MEDIUM])
return float(params["nf"])
# ── 引擎实现 ──────────────────────────────────────────────────────────────────
class NoiseReductionEngine:
"""音频降噪引擎。
基于 FFmpeg afftdnAudio FFt Denoiser)滤镜实现:
- 使用短时傅里叶变换分析音频频谱
- 识别并消除稳态背景噪音
- 保留人声等非稳态信号
"""
def __init__(self, config: NoiseReductionConfig):
self.config = config
def build_filter(self, input_label: str, output_label: str) -> str:
"""构建音频降噪滤镜字符串。
Args:
input_label: 输入标签,如 "[0:a]""[a0]"
output_label: 输出标签,如 "[nr0]"
Returns:
FFmpeg 滤镜字符串,如 "[a0]afftdn=nf=-25:tn=-10:tr=50[nr0]"
Raises:
ValueError: 配置无效时抛出(调用方应捕获并降级)
"""
if not self.config.has_effect():
return f"{input_label}anull{output_label}"
# 获取参数
if self.config.level == NoiseReductionLevel.CUSTOM:
nf = self.config.noise_floor
tn = -10 # 默认频谱平滑度
tr = 50 # 默认时间分辨率
else:
params = _LEVEL_PARAMS.get(
self.config.level,
_LEVEL_PARAMS[NoiseReductionLevel.MEDIUM],
)
nf = float(params["nf"])
tn = float(params["tn"])
tr = float(params["tr"])
# 构建 afftdn 滤镜
# nf: noise floor (dB)
# tn: temporal noise floor smoothing (dB)
# tr: time resolution (ms)
filter_parts = [f"afftdn=nf={nf}:tn={tn}:tr={tr}"]
# 人声增强:通过 highpass + 轻微压缩实现
if self.config.voice_enhance:
# 1. 高通滤波,去除低频噪音
filter_parts.append("highpass=f=80")
# 2. 轻微压缩,提升人声清晰度
filter_parts.append("acompressor=threshold=-20:ratio=2:attack=5:release=50")
# 3. 响度归一化
filter_parts.append("loudnorm=I=-16:TP=-1.5:LRA=11")
filter_str = f"{input_label}{','.join(filter_parts)}{output_label}"
return filter_str
def build_filter_arnndn(self, input_label: str, output_label: str, model_file: str) -> str:
"""使用 RNN 降噪滤镜(arnndn,效果更好但需要模型文件)。
注意:需要额外下载 RNNNoise 模型文件,默认使用 afftdn(无需额外依赖)。
Args:
input_label: 输入标签
output_label: 输出标签
model_file: RNNNoise 模型文件路径(.rnnn 格式)
Returns:
FFmpeg 滤镜字符串
"""
if not self.config.has_effect():
return f"{input_label}anull{output_label}"
return f"{input_label}arnndn=m={model_file}{output_label}"
def apply_noise_reduction_if_needed(
config_data: dict | None,
input_label: str,
output_label: str,
) -> Optional[str]:
"""便捷函数:根据配置判断是否需要应用音频降噪。
Args:
config_data: 降噪配置字典(从 plan.config.audio_noise_reduction 或 clip.config.noise_reduction 读取)
input_label: 输入标签
output_label: 输出标签
Returns:
滤镜字符串,不需要降噪时返回 None
"""
if not config_data:
return None
try:
config = NoiseReductionConfig.from_dict(config_data)
if not config.has_effect():
return None
engine = NoiseReductionEngine(config)
return engine.build_filter(input_label, output_label)
except Exception as e:
logger.warning("[noise-reduction] 应用降噪失败,跳过: %s", e)
return None
+483
View File
@@ -0,0 +1,483 @@
"""画中画(PiP)引擎 — 基于 FFmpeg overlay 滤镜实现多图层叠加.
支持能力:
- 多图层叠加:主画面 + 多个副画面
- 位置:9宫格 + 自由坐标(像素或百分比)
- 大小:宽高缩放(像素或百分比)
- 圆角裁剪:支持圆角矩形裁剪
- 透明度:0-100%
- 入场出场动画:淡入淡出、滑入滑出
- 时间同步:每个副画面独立开始时间和持续时长
- 降级策略:素材不存在时跳过,不阻断渲染
"""
from __future__ import annotations
import logging
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
logger = logging.getLogger(__name__)
# ── 位置常量 ──────────────────────────────────────────────────────────────────
# 9宫格位置枚举
POSITION_TOP_LEFT = "top_left"
POSITION_TOP_CENTER = "top_center"
POSITION_TOP_RIGHT = "top_right"
POSITION_CENTER_LEFT = "center_left"
POSITION_CENTER = "center"
POSITION_CENTER_RIGHT = "center_right"
POSITION_BOTTOM_LEFT = "bottom_left"
POSITION_BOTTOM_CENTER = "bottom_center"
POSITION_BOTTOM_RIGHT = "bottom_right"
_VALID_POSITIONS = {
POSITION_TOP_LEFT,
POSITION_TOP_CENTER,
POSITION_TOP_RIGHT,
POSITION_CENTER_LEFT,
POSITION_CENTER,
POSITION_CENTER_RIGHT,
POSITION_BOTTOM_LEFT,
POSITION_BOTTOM_CENTER,
POSITION_BOTTOM_RIGHT,
}
# 动画类型
ANIMATION_FADE = "fade" # 淡入淡出
ANIMATION_SLIDE_LEFT = "slide_left" # 从左滑入
ANIMATION_SLIDE_RIGHT = "slide_right" # 从右滑入
ANIMATION_SLIDE_TOP = "slide_top" # 从上滑入
ANIMATION_SLIDE_BOTTOM = "slide_bottom" # 从下滑入
_VALID_ANIMATIONS = {
ANIMATION_FADE,
ANIMATION_SLIDE_LEFT,
ANIMATION_SLIDE_RIGHT,
ANIMATION_SLIDE_TOP,
ANIMATION_SLIDE_BOTTOM,
}
# ── 数据模型 ──────────────────────────────────────────────────────────────────
@dataclass
class PiPLayerConfig:
"""单个画中画图层配置."""
# 素材来源
source: str = "" # 素材ID或视频URL
source_type: str = "asset_id" # "asset_id" | "url" | "local_path"
# 位置配置
position: str = POSITION_BOTTOM_RIGHT # 9宫格位置或 "custom"
x: int | str = 0 # 自定义x坐标(像素或百分比如 "30%"
y: int | str = 0 # 自定义y坐标
margin: int = 20 # 9宫格模式下的边距(像素)
# 大小配置
width: int | str = "25%" # 宽度(像素或百分比)
height: int | str = "" # 高度(空则按比例自适应)
# 样式
opacity: float = 1.0 # 透明度 0.0-1.0
corner_radius: int = 0 # 圆角半径(像素),0表示无圆角
border_width: int = 0 # 边框宽度
border_color: str = "white" # 边框颜色
# 时间控制
start_time: float = 0.0 # 开始显示时间(秒)
duration: float = 0.0 # 持续时长(秒),0表示全程显示
# 动画
animation_in: str = "" # 入场动画类型
animation_out: str = "" # 出场动画类型
animation_duration: float = 0.5 # 动画时长(秒)
# 层级
z_index: int = 1 # 图层顺序,数字越大越在上层
def validate(self) -> tuple[bool, str]:
"""校验配置合法性,返回 (是否合法, 错误信息)."""
if not self.source:
return False, "source不能为空"
if self.position != "custom" and self.position not in _VALID_POSITIONS:
return False, f"无效的position: {self.position}"
if self.opacity < 0 or self.opacity > 1:
return False, "opacity必须在0-1之间"
if self.corner_radius < 0:
return False, "corner_radius不能为负数"
if self.start_time < 0:
return False, "start_time不能为负数"
if self.duration < 0:
return False, "duration不能为负数"
if self.animation_in and self.animation_in not in _VALID_ANIMATIONS:
return False, f"无效的入场动画: {self.animation_in}"
if self.animation_out and self.animation_out not in _VALID_ANIMATIONS:
return False, f"无效的出场动画: {self.animation_out}"
if self.animation_duration < 0:
return False, "animation_duration不能为负数"
return True, ""
@dataclass
class PiPConfig:
"""画中画整体配置."""
enabled: bool = False
layers: list[PiPLayerConfig] = field(default_factory=list)
@classmethod
def from_dict(cls, data: dict[str, Any] | None) -> "PiPConfig":
"""从字典解析配置."""
if not data or not data.get("enabled", False):
return cls(enabled=False)
layers_data = data.get("layers", [])
layers = []
for layer_data in layers_data:
try:
layer = PiPLayerConfig(
source=layer_data.get("source", ""),
source_type=layer_data.get("source_type", "asset_id"),
position=layer_data.get("position", POSITION_BOTTOM_RIGHT),
x=layer_data.get("x", 0),
y=layer_data.get("y", 0),
margin=int(layer_data.get("margin", 20)),
width=layer_data.get("width", "25%"),
height=layer_data.get("height", ""),
opacity=float(layer_data.get("opacity", 1.0)),
corner_radius=int(layer_data.get("corner_radius", 0)),
border_width=int(layer_data.get("border_width", 0)),
border_color=layer_data.get("border_color", "white"),
start_time=float(layer_data.get("start_time", 0.0)),
duration=float(layer_data.get("duration", 0.0)),
animation_in=layer_data.get("animation_in", ""),
animation_out=layer_data.get("animation_out", ""),
animation_duration=float(layer_data.get("animation_duration", 0.5)),
z_index=int(layer_data.get("z_index", 1)),
)
valid, err = layer.validate()
if valid:
layers.append(layer)
else:
logger.warning("PiP图层配置无效,跳过: %s", err)
except (ValueError, TypeError) as e:
logger.warning("PiP图层解析失败,跳过: %s", e)
# 按 z_index 排序
layers.sort(key=lambda layer: layer.z_index)
return cls(enabled=bool(layers), layers=layers)
# ── PiP 引擎 ──────────────────────────────────────────────────────────────────
class PiPEngine:
"""画中画引擎 — 生成 FFmpeg 滤镜链实现多图层叠加."""
def __init__(
self,
output_width: int,
output_height: int,
output_fps: int = 30,
):
self.output_width = output_width
self.output_height = output_height
self.output_fps = output_fps
def _parse_size(self, value: int | str, base: int) -> int:
"""解析尺寸值(像素或百分比)."""
if isinstance(value, int):
return max(1, value)
if isinstance(value, str) and value.endswith("%"):
pct = float(value.rstrip("%")) / 100.0
return max(1, int(base * pct))
try:
return max(1, int(value))
except (ValueError, TypeError):
return int(base * 0.25) # 默认25%
def _parse_position(
self,
layer: PiPLayerConfig,
pip_width: int,
pip_height: int,
) -> tuple[int, int]:
"""计算画中画的实际位置 (x, y)."""
W = self.output_width
H = self.output_height
m = layer.margin
if layer.position == "custom":
x = self._parse_size(layer.x, W)
y = self._parse_size(layer.y, H)
return (x, y)
pos_map = {
POSITION_TOP_LEFT: (m, m),
POSITION_TOP_CENTER: ((W - pip_width) // 2, m),
POSITION_TOP_RIGHT: (W - pip_width - m, m),
POSITION_CENTER_LEFT: (m, (H - pip_height) // 2),
POSITION_CENTER: ((W - pip_width) // 2, (H - pip_height) // 2),
POSITION_CENTER_RIGHT: (W - pip_width - m, (H - pip_height) // 2),
POSITION_BOTTOM_LEFT: (m, H - pip_height - m),
POSITION_BOTTOM_CENTER: ((W - pip_width) // 2, H - pip_height - m),
POSITION_BOTTOM_RIGHT: (W - pip_width - m, H - pip_height - m),
}
return pos_map.get(layer.position, pos_map[POSITION_BOTTOM_RIGHT])
def _build_pip_pre_filter(
self,
input_label: str,
layer: PiPLayerConfig,
pip_width: int,
pip_height: int,
output_label: str,
) -> str:
"""构建单个PiP图层的预处理滤镜链.
处理顺序:scale → 圆角裁剪(可选)→ 边框(可选)→ 透明度 → 动画(可选)
"""
filters: list[str] = []
# Step 1: scale
filters.append(f"scale={pip_width}:{pip_height}")
filters.append("setsar=1")
# Step 2: 圆角裁剪
if layer.corner_radius > 0:
r = min(layer.corner_radius, pip_width // 2, pip_height // 2)
# 使用 geq + 圆形遮罩实现圆角
# 更简单的方式:用 rounded 滤镜(FFmpeg 5.0+)或 format + alpha
# 这里用更通用的方式:创建圆角遮罩 + overlay 到透明背景
filters.append(
f"format=yuva420p,"
f"geq="
f"lum='lum(X,Y)':"
f"cb='cb(X,Y)':"
f"cr='cr(X,Y)':"
f"a='if(lt(X,{r})*lt(Y,{r}),"
f"gt(hypot({r}-X,{r}-Y),{r})*0+1,"
f"if(gt(X,W-{r})*lt(Y,{r}),"
f"gt(hypot(X-(W-{r}),{r}-Y),{r})*0+1,"
f"if(lt(X,{r})*gt(Y,H-{r}),"
f"gt(hypot({r}-X,Y-(H-{r})),{r})*0+1,"
f"if(gt(X,W-{r})*gt(Y,H-{r}),"
f"gt(hypot(X-(W-{r}),Y-(H-{r})),{r})*0+1,1))))'"
)
# Step 3: 边框
if layer.border_width > 0:
bw = layer.border_width
color = layer.border_color
filters.append(f"pad={pip_width + 2*bw}:{pip_height + 2*bw}:{bw}:{bw}:{color}")
# Step 4: 透明度
if layer.opacity < 1.0:
alpha = layer.opacity
filters.append(f"format=yuva420p,colorchannelmixer=aa={alpha}")
# Step 5: 入场出场动画
if layer.animation_in or layer.animation_out:
filters.extend(self._build_animation_filters(layer, pip_width, pip_height))
filter_str = f"[{input_label}]{','.join(filters)}[{output_label}]"
return filter_str
def _build_animation_filters(
self,
layer: PiPLayerConfig,
pip_width: int,
pip_height: int,
) -> list[str]:
"""构建入场出场动画滤镜."""
filters: list[str] = []
anim_dur = layer.animation_duration
if layer.animation_in == ANIMATION_FADE:
# 淡入
filters.append(f"fade=t=in:st=0:d={anim_dur}:alpha=1")
elif layer.animation_in == ANIMATION_SLIDE_LEFT:
# 从左滑入 — 用 overlay 动态x实现,这里先标记位置表达式
pass # slide 动画在 overlay 表达式中处理
elif layer.animation_in == ANIMATION_SLIDE_RIGHT:
pass
elif layer.animation_in == ANIMATION_SLIDE_TOP:
pass
elif layer.animation_in == ANIMATION_SLIDE_BOTTOM:
pass
if layer.animation_out == ANIMATION_FADE:
# 淡出需要知道总时长,这里用表达式
if layer.duration > 0:
start_fade = layer.duration - anim_dur
filters.append(f"fade=t=out:st={max(0, start_fade)}:d={anim_dur}:alpha=1")
return filters
def _build_overlay_expr(
self,
layer: PiPLayerConfig,
base_x: int,
base_y: int,
pip_width: int,
pip_height: int,
) -> tuple[str, str]:
"""构建 overlay 滤镜的 x/y 表达式(支持滑动动画).
Returns:
(x_expr, y_expr) — FFmpeg表达式字符串
"""
W = self.output_width
H = self.output_height
anim_dur = layer.animation_duration
x_expr = str(base_x)
y_expr = str(base_y)
# 入场滑入动画
if layer.animation_in == ANIMATION_SLIDE_LEFT:
# 从左侧滑入:x 从 -pip_width 变化到 base_x
x_expr = f"'{base_x}+(X)*0+if(lt(t,{anim_dur}),{-pip_width}+t/{anim_dur}*({base_x}+{pip_width}),{base_x})'"
elif layer.animation_in == ANIMATION_SLIDE_RIGHT:
# 从右侧滑入:x 从 W 变化到 base_x
x_expr = f"'{base_x}+if(lt(t,{anim_dur}),{W}-t/{anim_dur}*({W}-{base_x}),{base_x})'"
elif layer.animation_in == ANIMATION_SLIDE_TOP:
y_expr = f"'{base_y}+if(lt(t,{anim_dur}),{-pip_height}+t/{anim_dur}*({base_y}+{pip_height}),{base_y})'"
elif layer.animation_in == ANIMATION_SLIDE_BOTTOM:
y_expr = f"'{base_y}+if(lt(t,{anim_dur}),{H}-t/{anim_dur}*({H}-{base_y}),{base_y})'"
# 出场滑出动画(需要总时长)
if layer.duration > 0 and anim_dur > 0:
out_start = layer.duration - anim_dur
if layer.animation_out == ANIMATION_SLIDE_LEFT:
x_expr = f"'{base_x}+if(gt(t,{out_start}),{base_x}-(t-{out_start})/{anim_dur}*({base_x}+{pip_width}),{base_x})'"
elif layer.animation_out == ANIMATION_SLIDE_RIGHT:
x_expr = f"'{base_x}+if(gt(t,{out_start}),{base_x}+(t-{out_start})/{anim_dur}*({W}-{base_x}+{pip_width}),{base_x})'"
elif layer.animation_out == ANIMATION_SLIDE_TOP:
y_expr = f"'{base_y}+if(gt(t,{out_start}),{base_y}-(t-{out_start})/{anim_dur}*({base_y}+{pip_height}),{base_y})'"
elif layer.animation_out == ANIMATION_SLIDE_BOTTOM:
y_expr = f"'{base_y}+if(gt(t,{out_start}),{base_y}+(t-{out_start})/{anim_dur}*({H}-{base_y}+{pip_height}),{base_y})'"
return (x_expr, y_expr)
def build_pip_filters(
self,
base_label: str,
pip_sources: list[tuple[str, PiPLayerConfig, Path]],
*,
base_input_idx: int = 0,
) -> tuple[str, list[str], str]:
"""构建完整的画中画滤镜链和输入参数.
Args:
base_label: 底层视频的滤镜标签(如 "final_video""v0",不带方括号)
pip_sources: [(input_label, layer_config, source_path), ...]
base_input_idx: PiP 素材在整个 FFmpeg 输入中的起始索引
Returns:
(filter_parts, input_args, final_label)
- filter_parts: 滤镜字符串列表(用 ; 连接后成为 filter_complex
- input_args: 额外的输入参数列表 ["-i", path, "-i", path, ...]
- final_label: 最终合成后的输出标签(不带方括号)
"""
if not pip_sources:
return [], [], base_label
filter_parts: list[str] = []
input_args: list[str] = []
current_label = base_label
for i, (input_label, layer, path) in enumerate(pip_sources):
# 添加输入
input_args.extend(["-i", str(path)])
# 计算实际大小
pip_w = self._parse_size(layer.width, self.output_width)
if layer.height:
pip_h = self._parse_size(layer.height, self.output_height)
else:
# 按宽度等比例(假设16:9,实际会scale时保持比例)
pip_h = int(pip_w * 9 / 16)
# 实际输入索引 = 起始索引 + 当前偏移
actual_input_idx = base_input_idx + i
# 预处理标签
pre_label = f"pip_pre_{i}"
# 构建预处理滤镜
pre_filter = self._build_pip_pre_filter(
input_label=f"{actual_input_idx}:v",
layer=layer,
pip_width=pip_w,
pip_height=pip_h,
output_label=pre_label,
)
filter_parts.append(pre_filter)
# 计算位置
base_x, base_y = self._parse_position(layer, pip_w, pip_h)
# 构建overlay表达式(支持滑动动画)
x_expr, y_expr = self._build_overlay_expr(layer, base_x, base_y, pip_w, pip_h)
# 时间控制(enable表达式)
enable_expr = ""
if layer.start_time > 0 or layer.duration > 0:
start = layer.start_time
if layer.duration > 0:
end = start + layer.duration
enable_expr = f":enable='between(t,{start},{end})'"
else:
enable_expr = f":enable='gte(t,{start})'"
# 合成标签
combined_label = f"pip_combined_{i}"
# overlay 滤镜
overlay_filter = (
f"[{current_label}][{pre_label}]" f"overlay={x_expr}:{y_expr}{enable_expr}" f"[{combined_label}]"
)
filter_parts.append(overlay_filter)
current_label = combined_label
return filter_parts, input_args, current_label
def validate_layer_source(
self,
layer: PiPLayerConfig,
asset_path_map: dict[str, Path],
) -> Path | None:
"""验证图层素材是否可用,返回本地路径或None(降级跳过)."""
try:
if layer.source_type == "local_path":
path = Path(layer.source)
if path.exists():
return path
elif layer.source_type == "asset_id":
if layer.source in asset_path_map:
return asset_path_map[layer.source]
elif layer.source_type == "url":
# URL类型由调用者负责下载,这里返回标记
return None # 暂时不支持直接URL
except Exception as e:
logger.warning("PiP素材验证失败: %s", e)
return None
+199 -26
View File
@@ -22,6 +22,8 @@ from pathlib import Path
from typing import TYPE_CHECKING
from video_processing.ffmpeg_utils import FFMPEG_BIN, probe_has_audio, run_ffmpeg
from video_processing.reverse_engine import ReverseConfig, ReverseEngine
from video_processing.speed_engine import SpeedEngine
if TYPE_CHECKING:
from video_processing.unified_render_service import RenderLayer, ResolvedClip
@@ -35,6 +37,8 @@ class RenderContext:
work_dir: Path
plan_id: str
# 音频降噪配置(全局,对最终混音结果应用)
noise_reduction_config: dict | None = None
# 音频探测缓存(避免同一 clip 被多次 ffprobe
_audio_cache: dict[str, bool] = field(default_factory=dict)
@@ -70,6 +74,9 @@ def mix_audio(
ctx: RenderContext,
layers: list[RenderLayer],
video_duration: float,
*,
bgm_path: str | None = None,
bgm_config: dict | None = None,
) -> Path | None:
"""音频后处理混音.
@@ -79,11 +86,14 @@ def mix_audio(
3. 独立音频轨(audio role)用 amix 混入
4. 输出时长截断到 video_duration
5. 无音频流的 clip 会被自动跳过,避免 FFmpeg 引用 [i:a] 失败
6. 如果提供了 bgm_path,则额外混入 BGM(支持淡入淡出、循环、人声闪避)
Args:
ctx: 渲染上下文
layers: 图层列表
video_duration: 视频总时长(用于截断音频)
bgm_path: BGM 音频本地路径,为 None 时不混入 BGM
bgm_config: BGM 配置字典(volume/fade_in/fade_out/sidechain 等)
Returns:
混音后的音频文件路径,无音频时返回 None
@@ -116,6 +126,15 @@ def mix_audio(
audio_clips = [c for c in audio_clips if clip_has_audio(ctx, c)]
if not main_clips and not audio_clips:
# 没有主音频也没有独立音频 → 检查是否有 BGM
if bgm_path and bgm_config and bgm_config.get("enabled", False):
from video_processing.bgm_mixer import BGMConfig, build_bgm_only
bgm_cfg = BGMConfig.from_config_dict(bgm_path, bgm_config)
try:
return build_bgm_only(ctx, bgm_cfg, video_duration)
except Exception:
logger.exception("[bgm] 纯BGM生成失败: plan_id=%s", ctx.plan_id)
return None
# 构建音频处理命令
@@ -124,11 +143,72 @@ def mix_audio(
# 简单场景:只有主图层 + 无独立音频 → 直接从视频提取音频并拼接
if main_clips and not audio_clips:
concat_main_audio(ctx, main_clips, output_path, video_duration)
return output_path
else:
# 有独立音频轨 → amix 混音
mix_with_independent_audio(ctx, main_clips, audio_clips, output_path, video_duration)
# 有独立音频轨 → amix 混音
mix_with_independent_audio(ctx, main_clips, audio_clips, output_path, video_duration)
return output_path
# ── BGM 混音 ──
if bgm_path and bgm_config and bgm_config.get("enabled", False):
from video_processing.bgm_mixer import BGMConfig, mix_bgm_with_main
bgm_cfg = BGMConfig.from_config_dict(bgm_path, bgm_config)
bgm_output = ctx.work_dir / f"audio_with_bgm_{ctx.plan_id}.aac"
try:
# 这里 main_audio 就是 output_path,先有主音频再混 BGM
final_path = mix_bgm_with_main(ctx, output_path, bgm_cfg, video_duration)
return _apply_noise_reduction_if_needed(ctx, final_path)
except Exception:
logger.exception("[bgm] BGM 混音失败,回退到无 BGM 音频: plan_id=%s", ctx.plan_id)
return _apply_noise_reduction_if_needed(ctx, output_path)
return _apply_noise_reduction_if_needed(ctx, output_path)
def _apply_noise_reduction_if_needed(ctx: RenderContext, audio_path: Path) -> Path:
"""如果配置了音频降噪,对已生成的音频文件应用降噪。
作为后处理步骤,对最终混音结果统一降噪。
失败时返回原始文件路径,不阻断主流程。
"""
if not ctx.noise_reduction_config:
return audio_path
try:
from video_processing.noise_reduction_engine import NoiseReductionConfig, NoiseReductionEngine
config = NoiseReductionConfig.from_dict(ctx.noise_reduction_config)
if not config.has_effect():
return audio_path
engine = NoiseReductionEngine(config)
filter_str = engine.build_filter("[0:a]", "[out]")
# 提取滤镜部分(不带标签)
filter_part = filter_str[len("[0:a]") : -len("[out]")]
nr_output_path = audio_path.with_name(f"{audio_path.stem}_nr.aac")
command = [
FFMPEG_BIN,
"-y",
"-i",
str(audio_path),
"-af",
filter_part,
"-acodec",
"aac",
"-b:a",
"128k",
str(nr_output_path),
]
run_ffmpeg(command)
if nr_output_path.exists():
return nr_output_path
logger.warning("[noise-reduction] 降噪输出文件不存在,使用原始音频")
return audio_path
except Exception as e:
logger.warning("[noise-reduction] 音频降噪失败,使用原始音频: %s", e)
return audio_path
def concat_main_audio(
@@ -145,40 +225,130 @@ def concat_main_audio(
# 单 clip,直接提取音频,截断到 min(clip有效时长, 视频总时长)
clip = clips[0]
effective_duration = clip_effective_duration(clip)
# 最终时长:取 clip 有效时长和视频总时长的较小值
# (视频总时长由主图层决定,但单 clip 场景下两者应该一致,仍做保护)
final_duration = effective_duration
trim_start = getattr(clip, "start_time", 0) or 0
speed = getattr(clip, "playback_speed", 1.0) or 1.0
if not isinstance(speed, (int, float)) or speed <= 0:
speed = 1.0
# 调速后时长
adjusted_duration = effective_duration / speed if abs(speed - 1.0) >= 1e-6 else effective_duration
# 最终时长:取调速后时长和视频总时长的较小值
final_duration = adjusted_duration
if video_duration > 0 and (final_duration <= 0 or final_duration > video_duration):
final_duration = video_duration
command = [
FFMPEG_BIN,
"-y",
"-i",
str(clip.local_path),
"-vn",
"-acodec",
"aac",
"-b:a",
"128k",
]
if final_duration > 0:
command.extend(["-t", f"{final_duration:.3f}"])
command.append(str(output_path))
run_ffmpeg(command)
# 音频倒放
reverse_config = ReverseConfig.from_dict(clip.config.get("reverse"))
has_reverse = reverse_config.enabled and reverse_config.reverse_audio
has_speed = abs(speed - 1.0) >= 1e-6
if not has_speed and not has_reverse:
# 无调速无倒放:简单命令行,-ss 裁剪更高效
command = [
FFMPEG_BIN,
"-y",
"-i",
str(clip.local_path),
"-vn",
"-acodec",
"aac",
"-b:a",
"128k",
]
if trim_start > 0:
command.extend(["-ss", f"{trim_start:.3f}"])
if final_duration > 0:
command.extend(["-t", f"{final_duration:.3f}"])
command.append(str(output_path))
run_ffmpeg(command)
else:
# 有调速或倒放:用 filter_complex
speed_engine = SpeedEngine()
audio_filters = []
if effective_duration > 0:
audio_filters.append(f"atrim=start={trim_start:.3f}:duration={effective_duration:.3f}")
audio_filters.append("asetpts=PTS-STARTPTS")
# 音频调速
if has_speed:
from video_processing.speed_engine import SpeedConfig
config = SpeedConfig(speed=float(speed))
config.clamp()
atempo_filter = speed_engine.build_audio_filter(config)
if atempo_filter:
audio_filters.append(atempo_filter)
# 音频倒放
if has_reverse:
reverse_filter = ReverseEngine.build_audio_filter(reverse_config, duration=effective_duration)
if reverse_filter:
audio_filters.append(reverse_filter)
filter_parts: list[str] = [f"[0:a]{','.join(audio_filters)}[outa]"]
if video_duration > 0 and final_duration < adjusted_duration:
filter_parts.append(f"[outa]atrim=0:{final_duration:.3f}[final_audio]")
final_label = "final_audio"
else:
final_label = "outa"
filter_complex = ";".join(filter_parts)
command = [
FFMPEG_BIN,
"-y",
"-i",
str(clip.local_path),
"-filter_complex",
filter_complex,
"-map",
f"[{final_label}]",
"-acodec",
"aac",
"-b:a",
"128k",
str(output_path),
]
run_ffmpeg(command)
return
# 多 clip,用 filter_complex concat
input_args: list[str] = []
filter_parts: list[str] = []
speed_engine = SpeedEngine()
for i, clip in enumerate(clips):
input_args.extend(["-i", str(clip.local_path)])
effective_duration = clip_effective_duration(clip)
trim_start = getattr(clip, "start_time", 0) or 0
speed = getattr(clip, "playback_speed", 1.0) or 1.0
if not isinstance(speed, (int, float)) or speed <= 0:
speed = 1.0
audio_filters: list[str] = []
if effective_duration > 0:
filter_parts.append(f"[{i}:a]atrim=0:{effective_duration:.3f},asetpts=PTS-STARTPTS[a{i}]")
audio_filters.append(f"atrim=start={trim_start:.3f}:duration={effective_duration:.3f}")
audio_filters.append("asetpts=PTS-STARTPTS")
# 音频调速 — atempo 多级串联
if abs(speed - 1.0) >= 1e-6:
from video_processing.speed_engine import SpeedConfig
config = SpeedConfig(speed=float(speed))
config.clamp()
atempo_filter = speed_engine.build_audio_filter(config)
if atempo_filter:
audio_filters.append(atempo_filter)
else:
filter_parts.append(f"[{i}:a]asetpts=PTS-STARTPTS[a{i}]")
audio_filters.append("asetpts=PTS-STARTPTS")
# 音频倒放
reverse_config = ReverseConfig.from_dict(clip.config.get("reverse"))
if reverse_config.enabled and reverse_config.reverse_audio:
reverse_filter = ReverseEngine.build_audio_filter(reverse_config, duration=effective_duration)
if reverse_filter:
audio_filters.append(reverse_filter)
filter_parts.append(f"[{i}:a]{','.join(audio_filters)}[a{i}]")
audio_labels = "".join(f"[a{i}]" for i in range(len(clips)))
filter_parts.append(f"{audio_labels}concat=n={len(clips)}:v=0:a=1[outa]")
@@ -236,9 +406,11 @@ def mix_with_independent_audio(
for clip in main_clips:
input_args.extend(["-i", str(clip.local_path)])
effective_duration = clip_effective_duration(clip)
trim_start = getattr(clip, "start_time", 0) or 0
if effective_duration > 0:
filter_parts.append(
f"[{input_idx}:a]atrim=0:{effective_duration:.3f},asetpts=PTS-STARTPTS[ma{input_idx}]"
f"[{input_idx}:a]atrim=start={trim_start:.3f}:duration={effective_duration:.3f},"
f"asetpts=PTS-STARTPTS[ma{input_idx}]"
)
else:
filter_parts.append(f"[{input_idx}:a]asetpts=PTS-STARTPTS[ma{input_idx}]")
@@ -255,11 +427,12 @@ def mix_with_independent_audio(
for j, clip in enumerate(audio_clips):
input_args.extend(["-i", str(clip.local_path)])
effective_duration = clip_effective_duration(clip)
trim_start = getattr(clip, "start_time", 0) or 0
volume = clip.config.get("volume", 1.0) if clip.config else 1.0
label = f"ia{j}"
filters = []
if effective_duration > 0:
filters.append(f"atrim=0:{effective_duration:.3f}")
filters.append(f"atrim=start={trim_start:.3f}:duration={effective_duration:.3f}")
filters.append("asetpts=PTS-STARTPTS")
if volume != 1.0:
filters.append(f"volume={volume}")
+116
View File
@@ -0,0 +1,116 @@
"""视频倒放引擎 — 基于 FFmpeg reverse + areverse 滤镜实现视频/音频倒放.
支持能力:
- 视频倒放(reverse 滤镜)
- 音频倒放(areverse 滤镜)
- 按 clip 分段倒放,每个 clip 独立配置
- 降级策略:不支持时跳过,不阻断渲染
"""
from __future__ import annotations
import logging
from dataclasses import dataclass
from typing import Any
logger = logging.getLogger(__name__)
# ── 数据模型 ──────────────────────────────────────────────────────────────────
@dataclass
class ReverseConfig:
"""视频倒放配置.
从 clip.config.reverse 读取,零侵入数据模型.
"""
enabled: bool = False
reverse_video: bool = True # 是否倒放视频
reverse_audio: bool = True # 是否倒放音频
@classmethod
def from_dict(cls, data: dict[str, Any] | None) -> "ReverseConfig":
"""从字典解析配置."""
if not data:
return cls(enabled=False)
try:
if not data.get("enabled", False):
return cls(enabled=False)
return cls(
enabled=True,
reverse_video=bool(data.get("reverse_video", True)),
reverse_audio=bool(data.get("reverse_audio", True)),
)
except (AttributeError, TypeError) as e:
logger.warning("倒放配置解析失败: %s,使用默认配置", e)
return cls(enabled=False)
# ── 倒放引擎 ──────────────────────────────────────────────────────────────────
class ReverseEngine:
"""视频倒放引擎 — 生成 FFmpeg 倒放滤镜.
视频倒放:reverse 滤镜
音频倒放:areverse 滤镜
注意事项:
- reverse 滤镜需要将整个视频帧加载到内存,长视频可能占用大量内存
- 建议对单 clip 时长做限制(如 < 60s),超长视频建议降级
"""
# 安全限制:单 clip 超过此时长不启用倒放(防止内存溢出)
MAX_SAFE_DURATION = 120.0 # 秒
@staticmethod
def build_video_filter(config: ReverseConfig, duration: float = 0.0) -> str:
"""构建视频倒放滤镜字符串.
Args:
config: 倒放配置
duration: clip 时长(秒),用于安全检查
Returns:
FFmpeg 滤镜字符串,如 "reverse";无效果返回空字符串
"""
if not config.enabled or not config.reverse_video:
return ""
# 安全检查:超长视频不启用倒放
if duration > ReverseEngine.MAX_SAFE_DURATION:
logger.warning(
"视频倒放安全限制:clip 时长 %.1fs 超过上限 %.1fs,跳过倒放",
duration,
ReverseEngine.MAX_SAFE_DURATION,
)
return ""
return "reverse"
@staticmethod
def build_audio_filter(config: ReverseConfig, duration: float = 0.0) -> str:
"""构建音频倒放滤镜字符串.
Args:
config: 倒放配置
duration: clip 时长(秒),用于安全检查
Returns:
FFmpeg 音频滤镜字符串,如 "areverse";无效果返回空字符串
"""
if not config.enabled or not config.reverse_audio:
return ""
# 安全检查:超长音频不启用倒放
if duration > ReverseEngine.MAX_SAFE_DURATION:
logger.warning(
"音频倒放安全限制:clip 时长 %.1fs 超过上限 %.1fs,跳过倒放",
duration,
ReverseEngine.MAX_SAFE_DURATION,
)
return ""
return "areverse"
+167
View File
@@ -0,0 +1,167 @@
"""视频调速引擎 — 基于 FFmpeg setpts + atempo 的速度调整能力。
支持:
- 0.25x ~ 4x 变速范围
- 视频调速(setpts
- 音频调速(atempo,多级串联处理超范围值)
- 音调修正(pitch_correct,默认开启)
- 边界自动钳制,不阻断渲染
"""
from dataclasses import dataclass
from typing import Optional
# ─── 常量 ───────────────────────────────────────────────
MIN_SPEED = 0.25
MAX_SPEED = 4.0
DEFAULT_SPEED = 1.0
# atempo 单级有效范围
_ATEMPO_MIN = 0.5
_ATEMPO_MAX = 2.0
@dataclass
class SpeedConfig:
"""调速配置。
Attributes:
speed: 播放速度,0.25~4.01.0 为原速
pitch_correct: 是否保持音调(默认 True,用 atempo 时间拉伸算法)
"""
speed: float = DEFAULT_SPEED
pitch_correct: bool = True
@classmethod
def parse(cls, data: Optional[dict]) -> "SpeedConfig":
"""从 dict 解析配置,无效值回退到默认。"""
if not data or not isinstance(data, dict):
return cls()
speed = data.get("speed", DEFAULT_SPEED)
if not isinstance(speed, (int, float)):
speed = DEFAULT_SPEED
pitch_correct = data.get("pitch_correct", True)
if not isinstance(pitch_correct, bool):
pitch_correct = True
config = cls(speed=float(speed), pitch_correct=pitch_correct)
config.clamp()
return config
def clamp(self) -> None:
"""将速度钳制到合法范围。"""
if self.speed <= 0:
self.speed = DEFAULT_SPEED
elif self.speed < MIN_SPEED:
self.speed = MIN_SPEED
elif self.speed > MAX_SPEED:
self.speed = MAX_SPEED
@property
def is_original(self) -> bool:
"""是否原速(无需调速)。"""
return abs(self.speed - 1.0) < 1e-6
class SpeedEngine:
"""调速引擎 — 生成 FFmpeg 调速滤镜链。
用法:
engine = SpeedEngine()
video_filter = engine.build_video_filter(config)
audio_filter = engine.build_audio_filter(config)
new_duration = engine.adjust_duration(duration, config)
"""
def build_video_filter(self, config: SpeedConfig) -> str:
"""生成视频调速滤镜字符串。
返回 setpts 滤镜表达式,原速时返回空字符串。
"""
if config.is_original:
return ""
# setpts=PTS/speed — speed>1 加速,speed<1 减速
return f"setpts=PTS/{config.speed:.4f}"
def build_audio_filter(self, config: SpeedConfig) -> str:
"""生成音频调速滤镜字符串。
atempo 单级范围 0.5~2.0,超出范围时自动多级串联:
- 0.25x → atempo=0.5,atempo=0.5
- 4x → atempo=2.0,atempo=2.0
- 0.3x → atempo=0.5,atempo=0.6
- 3x → atempo=2.0,atempo=1.5
原速时返回空字符串。
"""
if config.is_original:
return ""
speed = config.speed
stages: list[float] = self._split_atempo_stages(speed)
return ",".join(f"atempo={s:.4f}" for s in stages)
@staticmethod
def _split_atempo_stages(speed: float) -> list[float]:
"""将速度拆分为多级 atempo 串联,每级都在 [0.5, 2.0] 范围内。"""
if _ATEMPO_MIN <= speed <= _ATEMPO_MAX:
return [speed]
stages: list[float] = []
remaining = speed
# 加速场景(speed > 2.0
if speed > _ATEMPO_MAX:
while remaining > _ATEMPO_MAX:
stages.append(_ATEMPO_MAX)
remaining /= _ATEMPO_MAX
stages.append(remaining)
# 减速场景(speed < 0.5
else:
while remaining < _ATEMPO_MIN:
stages.append(_ATEMPO_MIN)
remaining /= _ATEMPO_MIN
stages.append(remaining)
return stages
def adjust_duration(self, original_duration: float, config: SpeedConfig) -> float:
"""计算调速后的时长。
加速 → 时长变短;减速 → 时长变长。
"""
if config.is_original or original_duration <= 0:
return original_duration
return original_duration / config.speed
def build_clip_speed_filter(
self,
speed: float,
pitch_correct: bool = True,
) -> tuple[str, str, SpeedConfig]:
"""便捷方法:从单一 speed 值生成视频+音频滤镜。
返回 (video_filter, audio_filter, config)。
"""
config = SpeedConfig(speed=speed, pitch_correct=pitch_correct)
config.clamp()
return (
self.build_video_filter(config),
self.build_audio_filter(config),
config,
)
@staticmethod
def resolve_clip_speed(
clip_config: dict,
global_speed: float = DEFAULT_SPEED,
) -> float:
"""从 clip config 中解析 playback_speed0 或缺失则使用全局速度。"""
speed = clip_config.get("playback_speed", 0) if clip_config else 0
if not isinstance(speed, (int, float)) or speed <= 0:
return global_speed
return float(speed)
+574
View File
@@ -0,0 +1,574 @@
"""贴纸叠加引擎 — 基于 FFmpeg overlay + drawtext 实现图片/文字贴纸.
支持能力:
- 图片贴纸(PNG/GIF):位置、大小、透明度、时间范围、淡入淡出
- 文字贴纸(花字):字体、颜色、描边、阴影、位置、时间范围、动画
- 9宫格位置 + 自由坐标(像素或百分比)
- 多贴纸叠加,按 z_index 排序
- 降级策略:素材不存在/无效时自动跳过,不阻断渲染
"""
from __future__ import annotations
import logging
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
logger = logging.getLogger(__name__)
# ── 预设贴纸分类 ──────────────────────────────────────────────────────────────
# 预设贴纸分类(仅用于前端展示,后端不依赖具体素材)
STICKER_CATEGORIES = [
("emoji", "表情包"),
("text", "文字花字"),
("decoration", "装饰"),
("arrow", "箭头指示"),
("frame", "边框"),
]
# 9宫格位置映射
POSITION_PRESETS = {
"top_left": (0.05, 0.05),
"top_center": (0.5, 0.05),
"top_right": (0.95, 0.05),
"center_left": (0.05, 0.5),
"center": (0.5, 0.5),
"center_right": (0.95, 0.5),
"bottom_left": (0.05, 0.95),
"bottom_center": (0.5, 0.95),
"bottom_right": (0.95, 0.95),
}
# ── 数据模型 ──────────────────────────────────────────────────────────────────
@dataclass
class ImageStickerConfig:
"""图片贴纸配置."""
enabled: bool = False
type: str = "image" # image / text
# 位置
position: str = "top_right" # 9宫格预设
x: float | None = None # 自定义x(像素或百分比)
y: float | None = None # 自定义y
x_unit: str = "percent" # pixel / percent
y_unit: str = "percent"
# 大小
scale: float = 1.0 # 缩放比例(相对于原始大小)
width: int | None = None # 指定宽度(像素)
height: int | None = None # 指定高度(像素)
# 透明度
opacity: float = 1.0 # 0.0~1.0
# 时间范围
start_time: float = 0.0
duration: float = 0.0 # 0 表示持续到结束
# 动画
fade_in: float = 0.0 # 淡入时长(秒)
fade_out: float = 0.0 # 淡出时长
# 层级
z_index: int = 10
# 素材
image_url: str = "" # 图片URL或本地路径
preset_id: str = "" # 预设贴纸ID
@dataclass
class TextStickerConfig:
"""文字贴纸配置."""
enabled: bool = False
type: str = "text"
text: str = ""
# 字体
font_size: int = 36
font_color: str = "#FFFFFF"
font_family: str = "sans"
# 描边
stroke_color: str = "#000000"
stroke_width: int = 2
# 阴影
shadow_color: str = "#000000"
shadow_x: int = 2
shadow_y: int = 2
shadow_alpha: float = 0.5
# 位置
position: str = "center"
x: float | None = None
y: float | None = None
x_unit: str = "percent"
y_unit: str = "percent"
# 时间范围
start_time: float = 0.0
duration: float = 0.0
# 动画
fade_in: float = 0.0
fade_out: float = 0.0
# 层级
z_index: int = 10
# 背景框
bg_color: str = "" # 空表示无背景
bg_padding: int = 8
bg_alpha: float = 0.8
bg_corner_radius: int = 8
@dataclass
class StickerOverlayResult:
"""贴纸叠加结果."""
filter_str: str # 滤镜字符串
output_label: str # 输出标签
extra_inputs: list[str] = field(default_factory=list) # 额外的输入文件路径
# ── 贴纸引擎 ──────────────────────────────────────────────────────────────────
class StickerEngine:
"""贴纸叠加引擎 — 生成 FFmpeg overlay / drawtext 滤镜链.
支持图片贴纸(overlay)和文字贴纸(drawtext)。
多贴纸按 z_index 排序依次叠加。
"""
@staticmethod
def _resolve_position(
config: ImageStickerConfig | TextStickerConfig,
canvas_w: int,
canvas_h: int,
sticker_w: int = 0,
sticker_h: int = 0,
) -> tuple[float, float]:
"""解析贴纸位置(像素坐标).
优先级:自定义坐标 > 9宫格预设
"""
# 先取预设的基准位置
if config.position in POSITION_PRESETS:
px, py = POSITION_PRESETS[config.position]
else:
px, py = 0.5, 0.5 # 默认居中
# 自定义坐标覆盖
if config.x is not None:
if config.x_unit == "percent":
px = config.x / 100.0
else:
px = config.x / canvas_w if canvas_w > 0 else 0.5
if config.y is not None:
if config.y_unit == "percent":
py = config.y / 100.0
else:
py = config.y / canvas_h if canvas_h > 0 else 0.5
# 转换为像素坐标(考虑贴纸尺寸,使位置为贴纸中心点)
x = px * canvas_w - sticker_w / 2
y = py * canvas_h - sticker_h / 2
# 钳制在画布内
x = max(0, min(x, canvas_w - sticker_w))
y = max(0, min(y, canvas_h - sticker_h))
return x, y
@staticmethod
def _build_overlay_filter(
sticker: ImageStickerConfig,
sticker_idx: int,
input_label: str,
output_label: str,
canvas_w: int,
canvas_h: int,
) -> str:
"""构建单个图片贴纸的 overlay 滤镜.
Args:
sticker: 贴纸配置
sticker_idx: 贴纸索引(用于生成滤镜标签)
input_label: 输入视频标签(如 "[base]"
output_label: 输出视频标签
canvas_w: 画布宽度
canvas_h: 画布高度
Returns:
FFmpeg 滤镜字符串
"""
sticker_label = f"sticker_{sticker_idx}_scaled"
# 1. 贴纸缩放预处理
scale_parts = []
if sticker.width and sticker.height:
scale_parts.append(f"scale={sticker.width}:{sticker.height}")
elif sticker.scale != 1.0:
# 按比例缩放
scale_parts.append(f"scale=iw*{sticker.scale}:ih*{sticker.scale}")
# 透明度调整
if sticker.opacity < 1.0:
scale_parts.append(f"colorchannelmixer=aa={sticker.opacity}")
# 淡入淡出
fade_parts = []
if sticker.fade_in > 0:
fade_parts.append(f"fade=in:st={sticker.start_time}:d={sticker.fade_in}:alpha=1")
if sticker.fade_out > 0 and sticker.duration > 0:
fade_out_start = sticker.start_time + sticker.duration - sticker.fade_out
fade_parts.append(f"fade=out:st={max(0, fade_out_start)}:d={sticker.fade_out}:alpha=1")
pre_filters = scale_parts + fade_parts
# 2. overlay 位置
# 先估算贴纸尺寸(假设原始尺寸 ~ canvas_w * 0.3
est_w = int(canvas_w * 0.3 * sticker.scale) if not sticker.width else sticker.width
est_h = int(canvas_h * 0.3 * sticker.scale) if not sticker.height else sticker.height
pos_x, pos_y = StickerEngine._resolve_position(sticker, canvas_w, canvas_h, est_w, est_h)
# 3. enable 表达式(时间范围)
enable_expr = ""
if sticker.duration > 0:
enable_expr = f":enable='between(t,{sticker.start_time},{sticker.start_time + sticker.duration})'"
# 组合滤镜
filter_parts: list[str] = []
# 贴纸预处理
if pre_filters:
filter_parts.append(f"[{sticker_idx + 1}:v]{','.join(pre_filters)}[{sticker_label}]")
sticker_source = f"[{sticker_label}]"
else:
sticker_source = f"[{sticker_idx + 1}:v]"
# overlay 合成
filter_parts.append(f"{input_label}{sticker_source}overlay={pos_x:.0f}:{pos_y:.0f}{enable_expr}{output_label}")
return ";".join(filter_parts)
@staticmethod
def _build_drawtext_filter(
sticker: TextStickerConfig,
input_label: str,
output_label: str,
canvas_w: int,
canvas_h: int,
) -> str:
"""构建单个文字贴纸的 drawtext 滤镜.
Args:
sticker: 文字贴纸配置
input_label: 输入视频标签
output_label: 输出视频标签
canvas_w: 画布宽度
canvas_h: 画布高度
Returns:
FFmpeg 滤镜字符串
"""
if not sticker.text:
return f"{input_label}copy{output_label}"
# 估算文字尺寸(粗略)
est_w = len(sticker.text) * sticker.font_size * 0.6
est_h = sticker.font_size * 1.4
pos_x, pos_y = StickerEngine._resolve_position(sticker, canvas_w, canvas_h, int(est_w), int(est_h))
drawtext_params: list[str] = []
# 文字内容(转义特殊字符)
escaped_text = sticker.text.replace(":", "\\:").replace("'", "\\'")
drawtext_params.append(f"text='{escaped_text}'")
# 字体
drawtext_params.append(f"fontsize={sticker.font_size}")
drawtext_params.append(f"fontcolor={sticker.font_color}")
# 描边
if sticker.stroke_width > 0:
drawtext_params.append(f"borderw={sticker.stroke_width}")
drawtext_params.append(f"bordercolor={sticker.stroke_color}")
# 阴影
if sticker.shadow_alpha > 0:
drawtext_params.append(f"shadowx={sticker.shadow_x}")
drawtext_params.append(f"shadowy={sticker.shadow_y}")
drawtext_params.append(f"shadowcolor={sticker.shadow_color}@{sticker.shadow_alpha}")
# 位置
drawtext_params.append(f"x={pos_x:.0f}")
drawtext_params.append(f"y={pos_y:.0f}")
# 时间范围
if sticker.duration > 0:
drawtext_params.append(f"enable='between(t,{sticker.start_time},{sticker.start_time + sticker.duration})'")
# 淡入淡出(drawtext 没有直接的淡入淡出,用 alpha 表达式模拟)
if sticker.fade_in > 0 or sticker.fade_out > 0:
alpha_expr = "1"
parts: list[str] = []
if sticker.fade_in > 0:
parts.append(
f"if(lt(t,{sticker.start_time + sticker.fade_in})," f"(t-{sticker.start_time})/{sticker.fade_in},1)"
)
if sticker.fade_out > 0 and sticker.duration > 0:
fade_out_start = sticker.start_time + sticker.duration - sticker.fade_out
parts.append(
f"if(gt(t,{fade_out_start})," f"({sticker.start_time + sticker.duration}-t)/{sticker.fade_out},1)"
)
if parts:
alpha_expr = "*".join(parts)
drawtext_params.append(f"alpha='{alpha_expr}'")
filter_str = f"{input_label}drawtext={':'.join(drawtext_params)}{output_label}"
return filter_str
@classmethod
def build_sticker_chain(
cls,
stickers: list[dict[str, Any]],
input_label: str,
output_label: str,
canvas_w: int,
canvas_h: int,
) -> StickerOverlayResult:
"""构建多贴纸叠加滤镜链.
Args:
stickers: 贴纸配置列表
input_label: 初始输入标签
output_label: 最终输出标签
canvas_w: 画布宽度
canvas_h: 画布高度
Returns:
StickerOverlayResult,包含滤镜字符串、输出标签、额外输入
"""
if not stickers:
return StickerOverlayResult(
filter_str=f"{input_label}copy{output_label}",
output_label=output_label,
extra_inputs=[],
)
# 解析配置
parsed_stickers: list[tuple[int, ImageStickerConfig | TextStickerConfig]] = []
image_stickers: list[ImageStickerConfig] = []
image_paths: list[str] = []
for i, s in enumerate(stickers):
try:
sticker_type = s.get("type", "image")
z = int(s.get("z_index", 10))
if sticker_type == "text":
config = TextStickerConfig(
enabled=True,
text=str(s.get("text", "")),
font_size=int(s.get("font_size", 36)),
font_color=str(s.get("font_color", "#FFFFFF")),
stroke_color=str(s.get("stroke_color", "#000000")),
stroke_width=int(s.get("stroke_width", 2)),
shadow_x=int(s.get("shadow_x", 2)),
shadow_y=int(s.get("shadow_y", 2)),
shadow_alpha=float(s.get("shadow_alpha", 0.5)),
position=str(s.get("position", "center")),
x=cls._safe_float(s.get("x")),
y=cls._safe_float(s.get("y")),
x_unit=str(s.get("x_unit", "percent")),
y_unit=str(s.get("y_unit", "percent")),
start_time=float(s.get("start_time", 0)),
duration=float(s.get("duration", 0)),
fade_in=float(s.get("fade_in", 0)),
fade_out=float(s.get("fade_out", 0)),
z_index=z,
bg_color=str(s.get("bg_color", "")),
bg_padding=int(s.get("bg_padding", 8)),
bg_alpha=float(s.get("bg_alpha", 0.8)),
bg_corner_radius=int(s.get("bg_corner_radius", 8)),
)
parsed_stickers.append((z, config))
else:
# 图片贴纸
image_path = s.get("image_path", "") or s.get("image_url", "")
if not image_path or not Path(image_path).exists():
logger.warning("贴纸素材不存在,跳过: %s", image_path)
continue
config = ImageStickerConfig(
enabled=True,
position=str(s.get("position", "top_right")),
x=cls._safe_float(s.get("x")),
y=cls._safe_float(s.get("y")),
x_unit=str(s.get("x_unit", "percent")),
y_unit=str(s.get("y_unit", "percent")),
scale=float(s.get("scale", 1.0)),
width=int(s["width"]) if s.get("width") else None,
height=int(s["height"]) if s.get("height") else None,
opacity=max(0.0, min(1.0, float(s.get("opacity", 1.0)))),
start_time=float(s.get("start_time", 0)),
duration=float(s.get("duration", 0)),
fade_in=float(s.get("fade_in", 0)),
fade_out=float(s.get("fade_out", 0)),
z_index=z,
image_url=str(s.get("image_url", "")),
)
parsed_stickers.append((z, config))
image_stickers.append(config)
image_paths.append(image_path)
except Exception as e:
logger.warning("贴纸配置解析失败,跳过: %s", e)
continue
if not parsed_stickers:
return StickerOverlayResult(
filter_str=f"{input_label}copy{output_label}",
output_label=output_label,
extra_inputs=[],
)
# 按 z_index 排序
parsed_stickers.sort(key=lambda x: x[0])
# 构建滤镜链
filter_parts: list[str] = []
current_label = input_label
img_idx = 0 # 图片贴纸的输入索引偏移
for idx, (_, sticker) in enumerate(parsed_stickers):
next_label = f"sticker_{idx}_out" if idx < len(parsed_stickers) - 1 else output_label
if isinstance(sticker, ImageStickerConfig):
# 图片贴纸:使用额外的输入(输入索引 = 1 + img_idx,0 是主视频)
# 注意:实际输入索引需要调用方根据输入列表确定
# 这里我们按 image_stickers 的顺序分配索引
# 主输入是 [0:v],贴纸输入从 [1:v] 开始
single_filter = cls._build_single_image_sticker(
sticker=sticker,
sticker_input_idx=img_idx + 1, # +1 因为 0 是主视频
input_label=current_label,
output_label=next_label,
canvas_w=canvas_w,
canvas_h=canvas_h,
)
filter_parts.append(single_filter)
img_idx += 1
else:
# 文字贴纸:drawtext,不需要额外输入
single_filter = cls._build_drawtext_filter(
sticker, # type: ignore
current_label,
next_label,
canvas_w,
canvas_h,
)
filter_parts.append(single_filter)
current_label = next_label
return StickerOverlayResult(
filter_str=";".join(filter_parts),
output_label=output_label,
extra_inputs=image_paths,
)
@classmethod
def _build_single_image_sticker(
cls,
sticker: ImageStickerConfig,
sticker_input_idx: int,
input_label: str,
output_label: str,
canvas_w: int,
canvas_h: int,
) -> str:
"""构建单个图片贴纸的完整滤镜(预处理 + overlay).
Args:
sticker: 贴纸配置
sticker_input_idx: 贴纸在 FFmpeg 输入中的索引
input_label: 输入视频标签
output_label: 输出标签
canvas_w: 画布宽
canvas_h: 画布高
"""
scaled_label = f"sticker_s{sticker_input_idx}"
# 预处理滤镜(缩放 + 透明度 + 淡入淡出)
pre_filters: list[str] = []
# 缩放
if sticker.width and sticker.height:
pre_filters.append(f"scale={sticker.width}:{sticker.height}")
elif sticker.scale != 1.0:
pre_filters.append(f"scale=iw*{sticker.scale}:ih*{sticker.scale}")
# 透明度
if sticker.opacity < 1.0:
pre_filters.append(f"format=rgba,colorchannelmixer=aa={sticker.opacity}")
# 淡入淡出(使用 fade 的 alpha 模式)
fade_filters: list[str] = []
if sticker.fade_in > 0:
fade_filters.append(f"fade=in:st={sticker.start_time}:d={sticker.fade_in}:alpha=1")
if sticker.fade_out > 0 and sticker.duration > 0:
fade_out_start = sticker.start_time + sticker.duration - sticker.fade_out
if fade_out_start > 0:
fade_filters.append(f"fade=out:st={fade_out_start}:d={sticker.fade_out}:alpha=1")
# 估算贴纸尺寸用于位置计算
est_w = int(canvas_w * 0.3 * sticker.scale) if not sticker.width else sticker.width
est_h = int(canvas_h * 0.3 * sticker.scale) if not sticker.height else sticker.height
pos_x, pos_y = cls._resolve_position(sticker, canvas_w, canvas_h, est_w, est_h)
# enable 表达式
enable_expr = ""
if sticker.duration > 0:
enable_expr = f":enable='between(t,{sticker.start_time},{sticker.start_time + sticker.duration})'"
parts: list[str] = []
# 贴纸预处理
all_pre = pre_filters + fade_filters
if all_pre:
parts.append(f"[{sticker_input_idx}:v]{','.join(all_pre)}[{scaled_label}]")
sticker_source = f"[{scaled_label}]"
else:
sticker_source = f"[{sticker_input_idx}:v]"
# overlay 合成
parts.append(f"{input_label}{sticker_source}overlay={pos_x:.0f}:{pos_y:.0f}{enable_expr}{output_label}")
return ";".join(parts)
@staticmethod
def _safe_float(val: Any) -> float | None:
"""安全转换 float."""
if val is None:
return None
try:
return float(val)
except (ValueError, TypeError):
return None
# ── 便捷函数 ──────────────────────────────────────────────────────────────────
def parse_stickers_from_config(config: dict[str, Any] | None) -> list[dict[str, Any]]:
"""从 plan.config.stickers 解析贴纸列表."""
if not config:
return []
stickers = config.get("stickers", [])
if not isinstance(stickers, list):
return []
return stickers
def get_sticker_categories() -> list[tuple[str, str]]:
"""获取贴纸分类列表."""
return list(STICKER_CATEGORIES)
+183
View File
@@ -0,0 +1,183 @@
"""字幕生成器 — 将字幕时间轴转换为 ASS 字幕文件。
与 render_subtitles.py 的区别:
- render_subtitles.py 处理静态整段标题/字幕
- 本模块处理带时间轴的多段 ASR 字幕
两者最终都输出 ASS 文件,供 FFmpeg 烧录。
"""
from __future__ import annotations
import logging
from pathlib import Path
from typing import Any
from packages.domain.subtitle import SubtitleTimeline
logger = logging.getLogger(__name__)
# ── 常量 ──────────────────────────────────────────────────────────────────────
DEFAULT_MAX_CHARS_PER_LINE = 20 # 每行最多字符数
DEFAULT_MIN_CHARS_PER_SEGMENT = 8 # 每段最少字符数
# ── ASS 工具函数 ────────────────────────────────────────────────────────────
def _hex_to_ass_color(hex_color: str) -> str:
"""将 HEX 颜色(#RRGGBB)转换为 ASS &HBBGGRR 格式。"""
hex_color = hex_color.lstrip("#")
if len(hex_color) != 6:
return "&H00FFFFFF"
r, g, b = hex_color[0:2], hex_color[2:4], hex_color[4:6]
return f"&H{b.upper()}{g.upper()}{r.upper()}"
def _position_to_ass_alignment(position: str) -> int:
"""将文字位置映射为 ASS \\an 对齐编号。"""
mapping = {
"top": 8,
"center": 5,
"bottom": 2,
}
return mapping.get(position, 2)
def _format_ass_time(seconds: float) -> str:
"""将秒数格式化为 ASS 时间格式 H:MM:SS.cc。"""
hours = int(seconds // 3600)
minutes = int((seconds % 3600) // 60)
secs = seconds % 60
return f"{hours}:{minutes:02d}:{secs:05.2f}"
def _escape_ass_text(text: str) -> str:
"""转义 ASS 文本中的特殊字符。"""
text = text.replace("\r\n", "\\N").replace("\n", "\\N").replace("\r", "\\N")
text = text.replace("{", "(").replace("}", ")")
return text
def _wrap_text(text: str, max_chars: int) -> list[str]:
"""将长文本按字数换行。
优先在标点处换行,没有合适标点时硬切。
"""
if len(text) <= max_chars:
return [text]
lines: list[str] = []
remaining = text
while len(remaining) > max_chars:
# 在前 max_chars 个字符中找标点断开
break_point = max_chars
punctuations = ",。!?、;:,.;:!?"
for i in range(max_chars, max_chars // 2, -1):
if i < len(remaining) and remaining[i] in punctuations:
break_point = i + 1
break
lines.append(remaining[:break_point])
remaining = remaining[break_point:]
if remaining:
lines.append(remaining)
return lines
# ── 主生成器 ─────────────────────────────────────────────────────────────────
def generate_ass_from_timeline(
output_path: Path,
timeline: SubtitleTimeline,
*,
video_width: int,
video_height: int,
subtitle_config: dict[str, Any] | None = None,
) -> Path:
"""从字幕时间轴生成 ASS 字幕文件。
Args:
output_path: 输出 ASS 文件路径
timeline: 字幕时间轴
video_width: 视频宽度
video_height: 视频高度
subtitle_config: 字幕样式配置(同 SubtitleConfig dict
Returns:
生成的 ASS 文件路径
"""
subtitle_config = subtitle_config or {}
if not timeline.segments:
output_path.write_text("", encoding="utf-8")
return output_path
# 样式参数
font_name = subtitle_config.get("font", "思源黑体")
font_size = int(subtitle_config.get("size", 24))
color = _hex_to_ass_color(subtitle_config.get("color", "#ffffff"))
position = subtitle_config.get("position", "bottom")
alignment = _position_to_ass_alignment(position)
max_chars_per_line = int(subtitle_config.get("max_chars_per_line", DEFAULT_MAX_CHARS_PER_LINE))
# 描边(默认黑色描边,保证可读性)
outline_color = "&H00000000"
outline_width = 1.5
# 边距
margin_v = 60 if position == "bottom" else 60
margin_l = 40
margin_r = 40
# 生成样式行
style_line = (
f"Style: Default,{font_name},{font_size},{color},"
f"&H000000FF,{outline_color},&H00000000,"
f"-1,0,0,0,100,100,0,0,"
f"1,{outline_width},0,{alignment},"
f"{margin_l},{margin_r},{margin_v},1"
)
# 生成事件行
events: list[str] = []
for seg in timeline.segments:
start_time = _format_ass_time(seg.start)
end_time = _format_ass_time(seg.end)
# 自动换行
lines = _wrap_text(seg.text, max_chars_per_line)
display_text = "\\N".join(lines)
safe_text = _escape_ass_text(display_text)
events.append(f"Dialogue: 0,{start_time},{end_time},Default,,0,0,0,,{safe_text}")
# 组装 ASS 文件
ass_content = f"""[Script Info]
ScriptType: v4.00+
PlayResX: {video_width}
PlayResY: {video_height}
ScaledBorderAndShadow: yes
WrapStyle: 2
Encoding: UTF-8
[V4+ Styles]
Format: Name, Fontname, Fontsize, PrimaryColour, SecondaryColour, OutlineColour, BackColour, Bold, Italic, Underline, StrikeOut, ScaleX, ScaleY, Spacing, Angle, BorderStyle, Outline, Shadow, Alignment, MarginL, MarginR, MarginV, Encoding
{style_line}
[Events]
Format: Layer, Start, End, Style, Name, MarginL, MarginR, MarginV, Effect, Text
{chr(10).join(events)}
"""
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_text(ass_content, encoding="utf-8")
return output_path
+123
View File
@@ -0,0 +1,123 @@
"""视频缩略图生成工具 — 抽取首帧上传到 OSS。"""
from __future__ import annotations
import logging
import tempfile
from pathlib import Path
logger = logging.getLogger(__name__)
def extract_first_frame(
video_path: str,
output_path: str | None = None,
*,
width: int = 640,
height: int = -1,
timeout: int = 30,
) -> str:
"""抽取视频第一帧作为封面图。
Args:
video_path: 视频文件路径
output_path: 输出图片路径,不传则用临时文件
width: 输出宽度(默认 640,-1 表示按比例缩放)
height: 输出高度(默认 -1,按比例缩放)
timeout: 超时时间(秒)
Returns:
生成的缩略图文件路径
Raises:
subprocess.CalledProcessError: ffmpeg 执行失败
"""
from video_processing.ffmpeg_utils import FFMPEG_BIN, run_ffmpeg
if output_path is None:
tmp = tempfile.NamedTemporaryFile(suffix=".jpg", delete=False)
tmp.close()
output_path = tmp.name
# -ss 00:00:01 取第1秒帧(避免首帧黑屏)
# -vframes 1 只取一帧
# -q:v 2 jpeg 高质量
scale_filter = f"scale={width}:{height}:force_original_aspect_ratio=decrease"
cmd = [
FFMPEG_BIN,
"-y",
"-i",
video_path,
"-ss",
"00:00:01",
"-vframes",
"1",
"-vf",
scale_filter,
"-q:v",
"2",
output_path,
]
try:
run_ffmpeg(cmd, capture_output=True, timeout=timeout)
except Exception:
# 短视频可能没有第1秒,退回到第0帧
cmd2 = [
FFMPEG_BIN,
"-y",
"-i",
video_path,
"-ss",
"00:00:00",
"-vframes",
"1",
"-vf",
scale_filter,
"-q:v",
"2",
output_path,
]
run_ffmpeg(cmd2, capture_output=True, timeout=timeout)
if not Path(output_path).exists() or Path(output_path).stat().st_size == 0:
raise RuntimeError(f"Thumbnail generation failed: {output_path}")
return output_path
def generate_and_upload_thumbnail(
video_path: str,
storage_key: str,
) -> str | None:
"""生成缩略图并上传到 OSS,返回 URL。
Args:
video_path: 本地视频路径
storage_key: OSS 存储 key(如 generated/projects/xxx/thumbnails/yyy.jpg
Returns:
上传成功返回 URL,失败返回 None
"""
thumbnail_path = None
try:
thumbnail_path = extract_first_frame(video_path)
except Exception as e:
logger.warning("Failed to extract thumbnail from %s: %s", video_path, e)
return None
try:
from video_processing.oss_helpers import upload_to_oss
url = upload_to_oss(thumbnail_path, storage_key)
return url
except Exception as e:
logger.warning("Failed to upload thumbnail to OSS: %s", e)
return None
finally:
# 清理临时文件
if thumbnail_path:
try:
Path(thumbnail_path).unlink(missing_ok=True)
except Exception:
pass
+381
View File
@@ -0,0 +1,381 @@
"""转场特效引擎 — Phase 8 智能增强.
基于 FFmpeg xfade 滤镜的统一转场抽象层,提供:
1. 转场类型枚举与预设管理
2. 转场配置解析与边界校验
3. 降级策略(不支持的转场自动 fallback 到硬切)
4. xfade 滤镜链构建(封装底层 ffmpeg_utils
新增转场只需在 TransitionType 中加一项 + 在 XFADE_TRANSITION_MAP 中映射。
"""
from __future__ import annotations
import logging
import sys
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 video_processing.ffmpeg_utils import build_xfade_filter_chain
logger = logging.getLogger(__name__)
# ── 常量 ──────────────────────────────────────────────────────────────────────
# 转场时长范围(秒)
MIN_TRANSITION_DURATION = 0.3
MAX_TRANSITION_DURATION = 2.0
DEFAULT_TRANSITION_DURATION = 0.5
# 硬切(无转场)
CUT_TRANSITION = "cut"
# ── 转场类型枚举 ──────────────────────────────────────────────────────────────
class TransitionType(StrEnum):
"""支持的转场效果类型.
每种类型对应 FFmpeg xfade filter 的一个 transition 值。
新增转场只需在此添加一项,并在 _FFMPEG_XFADE_MAP 中映射。
"""
# 硬切(无转场效果,直接拼接)
CUT = "cut"
# 淡入淡出(最常用,默认 fallback)
FADE = "fade"
# 溶解(交叉溶解)
DISSOLVE = "dissolve"
# 滑入系列
SLIDE_LEFT = "slideleft"
SLIDE_RIGHT = "slideright"
SLIDE_UP = "slideup"
SLIDE_DOWN = "slidedown"
# 缩放
ZOOM = "zoom"
# 擦除系列
WIPE_LEFT = "wipeleft"
WIPE_RIGHT = "wiperight"
WIPE_UP = "wipeup"
WIPE_DOWN = "wipedown"
# 圆形扩散
CIRCLE_CROP = "circlecrop"
# 矩形覆盖
RECT_CROP = "rectcrop"
@classmethod
def all_supported(cls) -> list[str]:
"""返回所有支持的转场类型名称列表."""
return [t.value for t in cls if t != cls.CUT]
@classmethod
def is_supported(cls, name: str) -> bool:
"""检查转场类型是否支持(不区分大小写和下划线)."""
normalized = _normalize_transition_name(name)
return normalized in _NAME_TO_ENUM_MAP
# ── 名称 → 枚举 映射(支持多种别名)──────────────────────────────────────────
def _normalize_transition_name(name: str) -> str:
"""标准化转场名称:小写 + 去下划线."""
return name.lower().replace("_", "").replace("-", "")
# 构建别名映射
_NAME_TO_ENUM_MAP: dict[str, TransitionType] = {}
for _t in TransitionType:
_NAME_TO_ENUM_MAP[_normalize_transition_name(_t.value)] = _t
# 额外的别名
_ALIASES: dict[str, TransitionType] = {
"dissolve": TransitionType.DISSOLVE,
"crossfade": TransitionType.DISSOLVE,
"crossdissolve": TransitionType.DISSOLVE,
"fadein": TransitionType.FADE,
"fadeout": TransitionType.FADE,
"fadeblack": TransitionType.FADE,
"slide": TransitionType.SLIDE_LEFT, # 默认向左滑
"wipe": TransitionType.WIPE_LEFT, # 默认向左擦
"zoomin": TransitionType.ZOOM,
"zoomout": TransitionType.ZOOM,
"circle": TransitionType.CIRCLE_CROP,
"rect": TransitionType.RECT_CROP,
}
for _alias, _type in _ALIASES.items():
_key = _normalize_transition_name(_alias)
if _key not in _NAME_TO_ENUM_MAP:
_NAME_TO_ENUM_MAP[_key] = _type
# ── TransitionType → FFmpeg xfade transition 名称映射 ─────────────────────────
_FFMPEG_XFADE_MAP: dict[TransitionType, str] = {
TransitionType.FADE: "fade",
TransitionType.DISSOLVE: "dissolve",
TransitionType.SLIDE_LEFT: "slideleft",
TransitionType.SLIDE_RIGHT: "slideright",
TransitionType.SLIDE_UP: "slideup",
TransitionType.SLIDE_DOWN: "slidedown",
TransitionType.ZOOM: "zoomin",
TransitionType.WIPE_LEFT: "wipeleft",
TransitionType.WIPE_RIGHT: "wiperight",
TransitionType.WIPE_UP: "wipeup",
TransitionType.WIPE_DOWN: "wipedown",
TransitionType.CIRCLE_CROP: "circlecrop",
TransitionType.RECT_CROP: "rectcrop",
}
# ── 转场配置 ──────────────────────────────────────────────────────────────────
@dataclass(slots=True)
class TransitionConfig:
"""转场效果配置.
Attributes:
effect: 转场效果名称(见 TransitionType
duration: 转场时长(秒),范围 0.3~2.0,默认 0.5
"""
effect: str = CUT_TRANSITION
duration: float = DEFAULT_TRANSITION_DURATION
@classmethod
def parse(cls, effect: str | None = None, duration: float | None = None) -> "TransitionConfig":
"""解析并验证转场配置,自动处理边界和降级.
Args:
effect: 转场效果名称(None 或空则使用默认 cut)
duration: 转场时长(None 则使用默认值)
Returns:
验证后的 TransitionConfig
"""
# 处理 effect
final_effect = CUT_TRANSITION
if effect and effect.strip():
effect_clean = effect.strip()
if TransitionType.is_supported(effect_clean):
final_effect = _resolve_transition_enum(effect_clean).value
elif effect_clean.lower() == CUT_TRANSITION:
final_effect = CUT_TRANSITION
else:
# 降级:不支持的转场 → 硬切,不阻断渲染
logger.warning(
"不支持的转场效果 '%s',已降级为硬切(cut",
effect_clean,
)
final_effect = CUT_TRANSITION
# 处理 duration:边界钳制
final_duration = DEFAULT_TRANSITION_DURATION
if duration is not None:
try:
d = float(duration)
if d < MIN_TRANSITION_DURATION:
logger.warning(
"转场时长 %.3fs 小于最小值 %.1fs,已钳制到最小值",
d,
MIN_TRANSITION_DURATION,
)
final_duration = MIN_TRANSITION_DURATION
elif d > MAX_TRANSITION_DURATION:
logger.warning(
"转场时长 %.3fs 大于最大值 %.1fs,已钳制到最大值",
d,
MAX_TRANSITION_DURATION,
)
final_duration = MAX_TRANSITION_DURATION
else:
final_duration = d
except (TypeError, ValueError):
logger.warning("无效的转场时长 '%s',使用默认值 %.1fs", duration, DEFAULT_TRANSITION_DURATION)
final_duration = DEFAULT_TRANSITION_DURATION
return cls(effect=final_effect, duration=final_duration)
@property
def is_cut(self) -> bool:
"""是否为硬切(无转场效果)."""
return self.effect == CUT_TRANSITION
@property
def ffmpeg_transition(self) -> str:
"""获取对应的 FFmpeg xfade transition 名称."""
if self.is_cut:
return ""
enum_type = _resolve_transition_enum(self.effect)
return _FFMPEG_XFADE_MAP.get(enum_type, "fade")
def _resolve_transition_enum(name: str) -> TransitionType:
"""将名称解析为 TransitionType 枚举,必须先通过 is_supported 校验."""
normalized = _normalize_transition_name(name)
return _NAME_TO_ENUM_MAP.get(normalized, TransitionType.FADE)
# ── 转场引擎 ──────────────────────────────────────────────────────────────────
class TransitionEngine:
"""转场特效引擎.
封装转场配置验证、降级策略和 xfade 滤镜链构建,
供 UnifiedRenderService 等上层调用。
用法::
engine = TransitionEngine(default_duration=0.5)
config = engine.resolve_config("fade", 0.8)
filter_str, total_dur = engine.build_xfade_chain(
clip_durations=[3.0, 4.0, 5.0],
clip_video_labels=["v0", "v1", "v2"],
transitions=["cut", "fade", "dissolve"],
)
"""
def __init__(self, default_duration: float = DEFAULT_TRANSITION_DURATION) -> None:
"""初始化转场引擎.
Args:
default_duration: 默认转场时长(秒),用于未指定时长的 clip
"""
self._default_duration = default_duration
def resolve_config(
self,
effect: str | None = None,
duration: float | None = None,
) -> TransitionConfig:
"""解析单个转场配置,应用验证和降级.
Args:
effect: 转场效果名称
duration: 转场时长
Returns:
验证后的 TransitionConfig
"""
# 若未指定 duration,使用引擎默认值
dur = duration if duration is not None else self._default_duration
return TransitionConfig.parse(effect=effect, duration=dur)
def resolve_clip_transitions(
self,
clip_transitions: list[str],
clip_durations: list[float] | None = None,
) -> list[TransitionConfig]:
"""批量解析 clip 级别的转场配置.
Args:
clip_transitions: 每个 clip 的转场效果名称列表
clip_durations: 每个 clip 的时长列表(用于验证转场时长不超过片段时长)
Returns:
TransitionConfig 列表
"""
configs: list[TransitionConfig] = []
for i, effect in enumerate(clip_transitions):
cfg = self.resolve_config(effect=effect)
# 额外校验:转场时长不能超过对应 clip 时长的一半(保守限制)
if clip_durations and i < len(clip_durations) and not cfg.is_cut:
max_safe_duration = max(MIN_TRANSITION_DURATION, clip_durations[i] * 0.5)
if cfg.duration > max_safe_duration:
cfg = TransitionConfig(effect=cfg.effect, duration=max_safe_duration)
configs.append(cfg)
return configs
def build_xfade_chain(
self,
clip_durations: list[float],
clip_video_labels: list[str],
transitions: list[str],
*,
transition_duration: float | None = None,
output_label: str = "outv",
) -> tuple[str, float]:
"""构建 xfade 转场滤镜链.
对每步转场应用验证和降级,然后调用底层 ffmpeg_utils 构建。
Args:
clip_durations: 每个片段的时长
clip_video_labels: 每个片段的视频流标签
transitions: 每个片段对应的转场效果
transition_duration: 统一转场时长,None 则使用引擎默认值
output_label: 最终输出标签
Returns:
(filter_string, estimated_total_duration)
"""
if len(clip_durations) <= 1:
return build_xfade_filter_chain(
clip_durations=clip_durations,
clip_video_labels=clip_video_labels,
transitions=transitions,
transition_duration=transition_duration or self._default_duration,
output_label=output_label,
)
# 解析所有转场配置
resolved = self.resolve_clip_transitions(transitions, clip_durations)
resolved_effects = [c.effect for c in resolved]
# 使用统一的时长(取各转场中最大的时长作为基准,底层会做每步钳制)
dur = transition_duration or self._default_duration
if not dur:
dur = max(c.duration for c in resolved) if resolved else DEFAULT_TRANSITION_DURATION
# 调用底层构建
return build_xfade_filter_chain(
clip_durations=clip_durations,
clip_video_labels=clip_video_labels,
transitions=resolved_effects,
transition_duration=dur,
output_label=output_label,
)
@staticmethod
def supported_transitions() -> list[dict[str, str]]:
"""获取所有支持的转场效果列表(用于 API 返回给前端).
Returns:
[{name, display_name, category}, ...]
"""
return [
{"name": "cut", "display_name": "硬切", "category": "basic"},
{"name": "fade", "display_name": "淡入淡出", "category": "basic"},
{"name": "dissolve", "display_name": "溶解", "category": "basic"},
{"name": "slideleft", "display_name": "左滑入", "category": "slide"},
{"name": "slideright", "display_name": "右滑入", "category": "slide"},
{"name": "slideup", "display_name": "上滑入", "category": "slide"},
{"name": "slidedown", "display_name": "下滑入", "category": "slide"},
{"name": "zoom", "display_name": "缩放", "category": "zoom"},
{"name": "wipeleft", "display_name": "左擦除", "category": "wipe"},
{"name": "wiperight", "display_name": "右擦除", "category": "wipe"},
{"name": "wipeup", "display_name": "上擦除", "category": "wipe"},
{"name": "wipedown", "display_name": "下擦除", "category": "wipe"},
{"name": "circlecrop", "display_name": "圆形扩散", "category": "special"},
{"name": "rectcrop", "display_name": "矩形扩散", "category": "special"},
]
+339
View File
@@ -0,0 +1,339 @@
"""裁剪引擎 — 基于 FFmpeg trim/atrim 的精确帧级裁剪.
支持:
- 入点出点裁剪(start_time / end_time / duration 三选二)
- 边界自动钳制(超出素材时长自动修正,不阻断渲染)
- 多段裁剪(一个素材裁剪出多段)
- 音画同步(视频 + 音频同步裁剪)
"""
from __future__ import annotations
import logging
from dataclasses import dataclass
from typing import Any
logger = logging.getLogger(__name__)
# 最小裁剪时长(秒),低于此值视为无效
MIN_TRIM_DURATION = 0.1
@dataclass
class TrimConfig:
"""裁剪配置.
三选二规则:start_time / end_time / duration 中必须至少给出两个,
第三个会被自动推导。如果三个都给了,以 start_time + duration 为准。
边界保护:
- start_time < 0 → 钳制到 0
- end_time > 素材时长 → 钳制到素材时长
- 计算出的 duration < 最小阈值 → 标记为无效
"""
start_time: float = 0.0 # 入点(素材内时间,秒)
end_time: float = 0.0 # 出点(素材内时间,秒),0 表示未指定
duration: float = 0.0 # 裁剪时长(秒),0 表示未指定
@classmethod
def from_dict(cls, data: dict[str, Any] | None) -> TrimConfig | None:
"""从字典构造,无有效裁剪参数时返回 None(不裁剪)."""
if not data:
return None
start = float(data.get("start_time", 0) or 0)
end = float(data.get("end_time", 0) or 0)
dur = float(data.get("duration", 0) or 0)
# 三个参数都没有 → 不裁剪
if start <= 0 and end <= 0 and dur <= 0:
return None
# 至少有两个参数(或一个合理的 start/duration
# 兼容:只传了 start_time → 从 start 开始取到末尾
# 兼容:只传了 duration → 从 0 开始取 duration
if start > 0 and end <= 0 and dur <= 0:
# 只有 start,取到末尾 → 这是"从某点开始"的语义,算有效
pass
elif dur > 0 and start <= 0 and end <= 0:
# 只有 duration → 从开头取 duration,算有效
pass
elif start <= 0 and end <= 0 and dur <= 0:
return None
return cls(start_time=start, end_time=end, duration=dur)
def validate_and_resolve(self, asset_duration: float) -> TrimConfig:
"""根据素材实际时长,解析并钳制裁剪参数.
返回一个新的 TrimConfig,其中 start_time / end_time / duration 都已确定。
如果裁剪无效(时长为0或负数),仍返回但调用方应检查 is_valid。
"""
start = self.start_time
end = self.end_time
dur = self.duration
# 边界:start 不能为负
if start < 0:
start = 0.0
# 边界:asset_duration 为 0 时保守处理(不裁剪,取全部)
if asset_duration <= 0:
return TrimConfig(start_time=0.0, end_time=0.0, duration=0.0)
# 三选二推导
# 判断顺序很重要:先判断需要两个显式值的组合,最后判断含默认值的
# 情况1start + end 都有显式值
if start > 0 and end > 0:
if end <= start:
# 出点 <= 入点,无效 → 返回 start 处一个极短片段(调用方会判无效)
return TrimConfig(start_time=start, end_time=start, duration=0.0)
dur = end - start
# 情况2end + duration 都有显式值
elif end > 0 and dur > 0:
start = end - dur
if start < 0:
start = 0.0
dur = end # 重新计算
# 情况3start + duration 都有值(start 可以是 0
elif dur > 0:
end = start + dur
# 情况4:只有 start → 取到素材末尾
elif start > 0 and end <= 0 and dur <= 0:
end = asset_duration
dur = end - start
# 情况5:只有 end → 从开头取到 end
elif end > 0 and start <= 0 and dur <= 0:
start = 0.0
dur = end
else:
# 都没有 → 不裁剪
return TrimConfig(start_time=0.0, end_time=0.0, duration=0.0)
# 边界钳制:end 不能超过素材时长
if end > asset_duration:
end = asset_duration
dur = end - start
# 边界钳制:start 不能超过素材时长
if start >= asset_duration:
start = max(0.0, asset_duration - MIN_TRIM_DURATION)
dur = asset_duration - start
end = asset_duration
# 保证 duration 不为负
if dur < 0:
dur = 0.0
return TrimConfig(start_time=start, end_time=end, duration=dur)
@property
def is_valid(self) -> bool:
"""裁剪是否有效(时长大于最小阈值)."""
return self.duration >= MIN_TRIM_DURATION
@property
def is_noop(self) -> bool:
"""是否等价于不裁剪(从0开始取全部)."""
return self.start_time <= 0 and self.duration <= 0
@property
def trim_from_start(self) -> bool:
"""是否从开头裁剪(start_time == 0."""
return self.start_time <= 0
@dataclass
class TrimSegment:
"""多段裁剪中的一段."""
segment_id: str # 段 ID(用于生成唯一标签)
trim: TrimConfig # 裁剪配置
order: int = 0 # 排序
@classmethod
def from_dict(cls, data: dict[str, Any], default_order: int = 0) -> TrimSegment:
"""从字典构造."""
return cls(
segment_id=str(data.get("segment_id", "") or f"seg_{default_order}"),
trim=TrimConfig(
start_time=float(data.get("start_time", 0) or 0),
end_time=float(data.get("end_time", 0) or 0),
duration=float(data.get("duration", 0) or 0),
),
order=int(data.get("order", default_order)),
)
class TrimEngine:
"""裁剪引擎 — 生成 FFmpeg trim / atrim 滤镜."""
@staticmethod
def build_video_trim_filter(
input_label: str,
trim: TrimConfig,
output_label: str,
) -> str:
"""构建视频裁剪滤镜链.
Args:
input_label: 输入视频标签,如 "[0:v]"
trim: 裁剪配置(已解析钳制)
output_label: 输出视频标签,如 "[v0_trimmed]"
Returns:
FFmpeg filter 字符串,如 "[0:v]trim=start=10:duration=5,setpts=PTS-STARTPTS[v0_trimmed]"
"""
if trim.is_noop:
# 不裁剪,直接直通
return f"{input_label}copy{output_label}" if False else f"{input_label}setpts=PTS-STARTPTS{output_label}"
parts: list[str] = []
# trim 滤镜参数
trim_args: list[str] = []
if trim.start_time > 0:
trim_args.append(f"start={trim.start_time:.3f}")
if trim.duration > 0:
trim_args.append(f"duration={trim.duration:.3f}")
elif trim.end_time > 0:
# end 用 duration 表示(start 到 end 的时长)
# 但 validate_and_resolve 后应该已经有 duration 了
pass
parts.append(f"trim={':'.join(trim_args)}")
parts.append("setpts=PTS-STARTPTS")
filter_str = f"{input_label}{','.join(parts)}{output_label}"
return filter_str
@staticmethod
def build_audio_trim_filter(
input_label: str,
trim: TrimConfig,
output_label: str,
) -> str:
"""构建音频裁剪滤镜链.
Args:
input_label: 输入音频标签,如 "[0:a]"
trim: 裁剪配置(已解析钳制)
output_label: 输出音频标签,如 "[a0_trimmed]"
Returns:
FFmpeg filter 字符串,如 "[0:a]atrim=start=10:duration=5,asetpts=PTS-STARTPTS[a0_trimmed]"
"""
if trim.is_noop:
return f"{input_label}asetpts=PTS-STARTPTS{output_label}"
parts: list[str] = []
trim_args: list[str] = []
if trim.start_time > 0:
trim_args.append(f"start={trim.start_time:.3f}")
if trim.duration > 0:
trim_args.append(f"duration={trim.duration:.3f}")
parts.append(f"atrim={':'.join(trim_args)}")
parts.append("asetpts=PTS-STARTPTS")
filter_str = f"{input_label}{','.join(parts)}{output_label}"
return filter_str
@staticmethod
def resolve_segments(
segments: list[TrimSegment],
asset_duration: float,
) -> list[TrimSegment]:
"""解析并钳制多段裁剪配置,过滤无效段.
Args:
segments: 原始段列表
asset_duration: 素材实际时长
Returns:
解析后的有效段列表,按 order 排序
"""
resolved: list[TrimSegment] = []
for i, seg in enumerate(segments):
resolved_trim = seg.trim.validate_and_resolve(asset_duration)
if not resolved_trim.is_valid:
logger.warning("裁剪段无效,跳过: segment_id=%s duration=%.3f", seg.segment_id, resolved_trim.duration)
continue
resolved.append(
TrimSegment(
segment_id=seg.segment_id,
trim=resolved_trim,
order=seg.order if seg.order >= 0 else i,
)
)
resolved.sort(key=lambda s: s.order)
return resolved
@staticmethod
def parse_segments_from_config(config: dict[str, Any] | None) -> list[TrimSegment]:
"""从 clip config 中解析多段裁剪配置.
config 中支持:
- trim_segments: [ {segment_id, start_time, end_time, duration, order}, ... ]
- trim_start / trim_end / trim_duration: 单段裁剪(兼容旧格式)
"""
if not config:
return []
# 优先解析多段
raw_segments = config.get("trim_segments", [])
if raw_segments and isinstance(raw_segments, list):
segments = []
for i, raw in enumerate(raw_segments):
if isinstance(raw, dict):
segments.append(TrimSegment.from_dict(raw, default_order=i))
return segments
# 单段裁剪兼容:从 trim_start/trim_end/trim_duration 构造
has_single = any(k in config for k in ("trim_start", "trim_end", "trim_duration"))
if has_single:
seg = TrimSegment(
segment_id="main",
trim=TrimConfig(
start_time=float(config.get("trim_start", 0) or 0),
end_time=float(config.get("trim_end", 0) or 0),
duration=float(config.get("trim_duration", 0) or 0),
),
order=0,
)
return [seg]
return []
# ── 工具函数 ──────────────────────────────────────────────────────────────────
def extract_trim_from_clip_config(config: dict[str, Any] | None) -> TrimConfig | None:
"""从 clip config 中提取单段裁剪配置.
兼容以下字段名:
- trim_start / trim_end / trim_duration
- start_time / end_time / duration(在 trim 子字典里)
"""
if not config:
return None
# trim 子字典
if "trim" in config and isinstance(config["trim"], dict):
return TrimConfig.from_dict(config["trim"])
# 扁平字段
has_any = any(k in config for k in ("trim_start", "trim_end", "trim_duration"))
if not has_any:
return None
data = {
"start_time": config.get("trim_start", 0),
"end_time": config.get("trim_end", 0),
"duration": config.get("trim_duration", 0),
}
return TrimConfig.from_dict(data)
+275
View File
@@ -0,0 +1,275 @@
"""TTS 配音引擎 — 集成到统一渲染管道的配音能力.
负责:
- 根据 TtsConfig 生成配音音频
- 字幕联动:按字幕片段分段合成,自动对齐时间轴
- 整段配音:整段文本生成一条音频
- 失败降级:TTS 失败不阻断渲染
"""
from __future__ import annotations
import logging
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
from packages.domain.tts_config import TtsConfig
from packages.ports.tts_service import TtsError, TtsService
logger = logging.getLogger(__name__)
@dataclass
class VoiceoverSegment:
"""配音片段.
Attributes:
text: 文本内容
start_time: 开始时间(秒)
end_time: 结束时间(秒)
audio_path: 合成后的音频文件路径
duration: 音频实际时长
"""
text: str
start_time: float = 0.0
end_time: float = 0.0
audio_path: Path | None = None
duration: float = 0.0
@dataclass
class VoiceoverResult:
"""配音结果.
Attributes:
success: 是否成功
segments: 配音片段列表
total_duration: 总时长
error_message: 错误信息(失败时)
"""
success: bool = False
segments: list[VoiceoverSegment] = field(default_factory=list)
total_duration: float = 0.0
error_message: str = ""
class TtsEngine:
"""TTS 配音引擎.
封装 TtsService 调用,支持:
- 整段配音
- 字幕联动配音
- 失败降级
"""
def __init__(
self,
tts_service: TtsService,
work_dir: Path,
) -> None:
self._tts = tts_service
self._work_dir = work_dir
self._work_dir.mkdir(parents=True, exist_ok=True)
def generate_full_voiceover(
self,
config: TtsConfig,
*,
total_duration: float = 0.0,
) -> VoiceoverResult:
"""生成整段配音.
Args:
config: TTS 配置
total_duration: 视频总时长(用于调整配音速度适配)
Returns:
配音结果
"""
if not config.enabled or not config.text.strip():
return VoiceoverResult(success=False, error_message="配音未启用或文本为空")
try:
output_path = self._work_dir / "voiceover_full.wav"
audio_path = self._tts.synthesize(
text=config.text,
voice_id=config.voice_id,
speed=config.speed,
pitch=config.pitch,
output_path=output_path,
)
# 探测实际时长
duration = self._probe_duration(audio_path)
segment = VoiceoverSegment(
text=config.text,
start_time=0.0,
end_time=duration,
audio_path=audio_path,
duration=duration,
)
return VoiceoverResult(
success=True,
segments=[segment],
total_duration=duration,
)
except TtsError as e:
logger.warning("TTS 整段配音失败,降级跳过: %s", e)
return VoiceoverResult(success=False, error_message=str(e))
except Exception as e:
logger.warning("TTS 整段配音异常,降级跳过: %s", e)
return VoiceoverResult(success=False, error_message=str(e))
def generate_subtitle_voiceover(
self,
config: TtsConfig,
subtitles: list[dict[str, Any]],
) -> VoiceoverResult:
"""根据字幕生成配音(字幕联动).
每个字幕片段独立合成,按字幕时间轴对齐。
Args:
config: TTS 配置
subtitles: 字幕列表,每项含 text/start_time/end_time
Returns:
配音结果
"""
if not config.enabled:
return VoiceoverResult(success=False, error_message="配音未启用")
if not subtitles:
return VoiceoverResult(success=False, error_message="字幕为空")
segments: list[VoiceoverSegment] = []
total_duration = 0.0
for i, sub in enumerate(subtitles):
text = sub.get("text", "").strip()
if not text:
continue
start_time = float(sub.get("start_time", 0))
end_time = float(sub.get("end_time", 0))
target_duration = max(0.1, end_time - start_time)
try:
# 计算适配时长所需语速:让配音时长 ≈ 字幕时长
estimated = self._tts.estimate_duration(text, speed=config.speed)
adjusted_speed = config.speed
if estimated > 0 and target_duration > 0:
# 按目标时长调整语速,限制在 0.5~2.0 范围内
speed_factor = estimated / target_duration
adjusted_speed = max(0.5, min(2.0, config.speed * speed_factor))
output_path = self._work_dir / f"voiceover_seg_{i:03d}.wav"
audio_path = self._tts.synthesize(
text=text,
voice_id=config.voice_id,
speed=adjusted_speed,
pitch=config.pitch,
output_path=output_path,
)
actual_duration = self._probe_duration(audio_path)
segment = VoiceoverSegment(
text=text,
start_time=start_time,
end_time=start_time + actual_duration,
audio_path=audio_path,
duration=actual_duration,
)
segments.append(segment)
total_duration = max(total_duration, start_time + actual_duration)
except TtsError as e:
logger.warning("TTS 字幕片段 %d 合成失败,跳过: %s", i, e)
continue
except Exception as e:
logger.warning("TTS 字幕片段 %d 异常,跳过: %s", i, e)
continue
if not segments:
return VoiceoverResult(success=False, error_message="所有字幕片段合成失败")
return VoiceoverResult(
success=True,
segments=segments,
total_duration=total_duration,
)
def build_audio_mix_filter(
self,
result: VoiceoverResult,
*,
video_duration: float,
base_label: str = "0:a",
) -> tuple[str, list[Path]]:
"""构建配音混音滤镜.
将配音片段按时间轴排列,生成 amix 混入。
Args:
result: 配音结果
video_duration: 视频总时长
base_label: 基础音轨标签
Returns:
(filter_complex 字符串, 配音音频文件列表)
"""
if not result.success or not result.segments:
return "", []
filter_parts: list[str] = []
audio_files: list[Path] = []
delay_labels: list[str] = []
for i, seg in enumerate(result.segments):
if seg.audio_path is None or not seg.audio_path.exists():
continue
audio_files.append(seg.audio_path)
seg_label = f"v{i}"
# 音量调整
# 用 adelay 延迟到字幕开始时间
delay_ms = int(max(0, int(seg.start_time * 1000)))
filter_parts.append(f"[{i}:a]adelay={delay_ms}:all=1,volume=0.8[{seg_label}]")
delay_labels.append(f"[{seg_label}]")
if not delay_labels:
return "", []
# 所有片段 concat 成一条配音音轨(用 amix 叠加多个延时后的片段
mix_inputs = "".join(delay_labels)
n_inputs = len(delay_labels)
tts_label = "tts_mixed"
if n_inputs == 1:
# 单个片段直接用
filter_parts.append(f"{delay_labels[0]}[{tts_label}]")
else:
# 多个片段 amix 叠加
filter_parts.append(f"{mix_inputs}amix=inputs={n_inputs}:duration=longest[{tts_label}]")
return ";".join(filter_parts), audio_files
def _probe_duration(self, audio_path: Path) -> float:
"""探测音频时长."""
try:
from video_processing.ffmpeg_utils import probe_duration
return probe_duration(audio_path)
except Exception:
# 探测失败,按文件名估算
return 0.0
File diff suppressed because it is too large Load Diff
+315
View File
@@ -0,0 +1,315 @@
"""水印引擎 — 基于 FFmpeg overlay 滤镜的水印叠加.
支持
- 图片水印PNG/logo
- 文字水印drawtext
- 9宫格位置 + 边距配置
- 透明度/大小缩放
- 滚动水印跑马灯
"""
from __future__ import annotations
import logging
from dataclasses import dataclass
from typing import Any
logger = logging.getLogger(__name__)
# 9宫格位置枚举
WATERMARK_POSITIONS = {
"top_left": "左上",
"top_center": "中上",
"top_right": "右上",
"center_left": "左中",
"center": "中心",
"center_right": "右中",
"bottom_left": "左下",
"bottom_center": "中下",
"bottom_right": "右下",
}
@dataclass
class WatermarkConfig:
"""水印配置.
mode: "image" 图片水印 | "text" 文字水印
position: 9宫格位置
opacity: 透明度 0.0-1.0
scale: 缩放比例图片水印0.1-1.0
margin: 边距像素
scroll: 是否滚动跑马灯
scroll_speed: 滚动速度像素/
"""
mode: str = "text" # image | text
position: str = "bottom_right"
# 图片水印
image_path: str = "" # 本地图片路径
scale: float = 0.2 # 相对输出宽度的比例
opacity: float = 0.8 # 0.0-1.0
# 文字水印
text: str = ""
font_size: int = 24
font_color: str = "white"
font_path: str = "" # 字体文件路径
# 边距
margin_x: int = 20
margin_y: int = 20
# 滚动水印
scroll: bool = False
scroll_speed: int = 50 # 像素/秒
@classmethod
def from_dict(cls, data: dict[str, Any] | None) -> WatermarkConfig | None:
"""从字典构造,空配置返回 None(不加水印)."""
if not data:
return None
enabled = data.get("enabled", False)
if not enabled:
return None
mode = data.get("mode", "text")
# 图片模式需要 image_path;文字模式需要 text
if mode == "image":
image_path = data.get("image_path", "") or data.get("image", "") or ""
if not image_path:
logger.warning("图片水印缺少 image_path,跳过水印")
return None
elif mode == "text":
text = data.get("text", "") or ""
if not text:
logger.warning("文字水印缺少 text,跳过水印")
return None
position = data.get("position", "bottom_right")
if position not in WATERMARK_POSITIONS:
position = "bottom_right"
return cls(
mode=mode,
position=position,
image_path=str(data.get("image_path", data.get("image", "")) or ""),
scale=float(data.get("scale", 0.2)),
opacity=float(data.get("opacity", 0.8)),
text=str(data.get("text", "") or ""),
font_size=int(data.get("font_size", 24)),
font_color=str(data.get("font_color", "white")),
font_path=str(data.get("font_path", "") or ""),
margin_x=int(data.get("margin_x", 20)),
margin_y=int(data.get("margin_y", 20)),
scroll=bool(data.get("scroll", False)),
scroll_speed=int(data.get("scroll_speed", 50)),
)
def validate(self) -> tuple[bool, str]:
"""校验配置是否有效."""
if self.position not in WATERMARK_POSITIONS:
return False, f"不支持的位置: {self.position}"
if not (0.0 <= self.opacity <= 1.0):
return False, "透明度必须在 0-1 之间"
if self.mode == "image":
if not self.image_path:
return False, "图片水印缺少图片路径"
if not (0.01 <= self.scale <= 1.0):
return False, "缩放比例必须在 0.01-1.0 之间"
elif self.mode == "text":
if not self.text:
return False, "文字水印缺少文字内容"
if self.font_size <= 0:
return False, "字体大小必须大于 0"
else:
return False, f"不支持的水印模式: {self.mode}"
return True, ""
class WatermarkEngine:
"""水印引擎 — 生成 FFmpeg 水印滤镜."""
@staticmethod
def calc_position(
position: str,
output_width: int,
output_height: int,
wm_width: int,
wm_height: int,
margin_x: int,
margin_y: int,
) -> tuple[int, int]:
"""根据9宫格位置计算水印坐标 (x, y).
坐标系左上角为 (0, 0)
"""
if position == "top_left":
return margin_x, margin_y
elif position == "top_center":
return (output_width - wm_width) // 2, margin_y
elif position == "top_right":
return output_width - wm_width - margin_x, margin_y
elif position == "center_left":
return margin_x, (output_height - wm_height) // 2
elif position == "center":
return (output_width - wm_width) // 2, (output_height - wm_height) // 2
elif position == "center_right":
return output_width - wm_width - margin_x, (output_height - wm_height) // 2
elif position == "bottom_left":
return margin_x, output_height - wm_height - margin_y
elif position == "bottom_center":
return (output_width - wm_width) // 2, output_height - wm_height - margin_y
elif position == "bottom_right":
return output_width - wm_width - margin_x, output_height - wm_height - margin_y
else:
# 默认右下角
return output_width - wm_width - margin_x, output_height - wm_height - margin_y
@staticmethod
def calc_scroll_x(position: str, output_width: int, wm_width: int, speed: int) -> str:
"""生成滚动水印的 x 坐标表达式.
从右向左滚动跑马灯效果
"""
# x 从 W 到 -wm_width,整个宽度 + wm_width 的距离
# 使用 overlay 的 enable 表达式
# x = 'W - (t * speed)' → 不对,应该是持续滚动
# 标准跑马灯:x = -w + (t * speed) % (W + w)
# 但 FFmpeg overlay 支持表达式
return f"mod({output_width}-mod({speed}*t\\,{output_width}+{wm_width})"
@staticmethod
def build_image_watermark_filter(
input_video_label: str,
wm_image_path: str,
output_width: int,
output_height: int,
output_label: str,
config: WatermarkConfig,
) -> tuple[str, list[str]]:
"""构建图片水印滤镜链.
Args:
input_video_label: 输入视频标签 "[final_video]"
wm_image_path: 水印图片本地路径
output_width: 输出视频宽度
output_height: 输出视频高度
output_label: 输出标签
config: 水印配置
Returns:
(filter_complex_str, input_args_list)
input_args ["-i", wm_image_path] 格式
"""
# 计算水印尺寸(按输出宽度比例缩放)
wm_width = int(output_width * config.scale)
wm_height = -1 # 保持比例
wm_filter = f"scale={wm_width}:{wm_height}"
# 透明度处理
if config.opacity < 1.0:
wm_filter += f",format=rgba,colorchannelmixer=aa={config.opacity}"
# 水印预处理标签
wm_pre_label = "[wm_scaled]"
# 计算位置
x, y = WatermarkEngine.calc_position(
config.position,
output_width,
output_height,
wm_width,
wm_width, # 高度未知,先用宽度估算
config.margin_x,
config.margin_y,
)
# 滚动水印
if config.scroll:
# 从右向左滚动:x = W - (t * speed) mod (W + wm_w)
# 使用 overlay 表达式
x_expr = f"{output_width}-mod({config.scroll_speed}*t\\,{output_width}+{wm_width}"
y_expr = str(y)
overlay_expr = f"x={x_expr}:y={y_expr}"
else:
overlay_expr = f"x={x}:y={y}"
# 构建滤镜
# 先缩放水印图
wm_input_idx = 1 # 假设水印图是第二个输入(索引1
filter_parts = [
f"[1:v]{wm_filter}{wm_pre_label}",
f"{input_video_label}{wm_pre_label}overlay={overlay_expr}{output_label}",
]
filter_complex = ";".join(filter_parts)
input_args = ["-i", wm_image_path]
return filter_complex, input_args
@staticmethod
def build_text_watermark_filter(
input_video_label: str,
output_label: str,
config: WatermarkConfig,
output_width: int,
output_height: int,
) -> str:
"""构建文字水印滤镜(drawtext.
Args:
input_video_label: 输入视频标签
output_label: 输出标签
config: 水印配置
output_width: 输出宽度
output_height: 输出高度
Returns:
FFmpeg filter 字符串
"""
# 转义文字中的特殊字符
text = config.text.replace(":", "\\:").replace("'", "\\'")
# 字体配置
font_config = []
if config.font_path:
font_path_escaped = config.font_path.replace(":", "\\:").replace("'", "\\'")
font_config.append(f"fontfile='{font_path_escaped}'")
font_config.append(f"fontsize={config.font_size}")
font_config.append(f"fontcolor={config.font_color}@{config.opacity}")
# 估算文字宽高(粗略估算,用于位置计算)
# 每个汉字约等于 font_size 宽高
approx_w = len(config.text) * config.font_size
approx_h = config.font_size
# 位置计算
x, y = WatermarkEngine.calc_position(
config.position,
output_width,
output_height,
approx_w,
approx_h,
config.margin_x,
config.margin_y,
)
# 滚动水印
if config.scroll:
x_expr = f"w-mod({config.scroll_speed}*t\\,W+w)"
pos_config = [f"x={x_expr}", f"y={y}"]
else:
pos_config = [f"x={x}", f"y={y}"]
# 组装 drawtext
drawtext_parts = [f"text='{text}'"] + font_config + pos_config
drawtext = "drawtext=" + ":".join(drawtext_parts)
return f"{input_video_label}{drawtext}{output_label}"
+1
View File
@@ -16,5 +16,6 @@ celery_app.conf.imports = (
"worker_app.tasks.tts_synthesis",
"worker_app.tasks.edit_plan_generation",
"worker_app.tasks.compose_video",
"worker_app.tasks.batch_download",
"apps.worker.video_processing.dedup",
)
+112
View File
@@ -0,0 +1,112 @@
"""批量下载任务 — 将多个成片打包为 zip 上传到 OSS。"""
from __future__ import annotations
import logging
import os
import tempfile
import uuid
import zipfile
from pathlib import Path
from worker_app.celery_app import celery_app
logger = logging.getLogger(__name__)
@celery_app.task(bind=True, name="worker.batch_download_videos", max_retries=1)
def batch_download_videos(self, video_ids: list[str], user_id: str = "") -> dict:
"""批量下载视频并打包为 zip。
Args:
video_ids: 视频 ID 列表
user_id: 发起用户 ID
Returns:
{"download_url": "...", "file_count": N, "total_size": total_bytes}
"""
from video_processing.oss_helpers import download_asset, upload_to_oss
from worker_app.db import SessionLocal
from packages.adapters.sqlalchemy_impl.generated_video_repository import (
SQLAlchemyGeneratedVideoRepository,
)
session = SessionLocal()
try:
repo = SQLAlchemyGeneratedVideoRepository(session)
videos = repo.get_by_ids(video_ids)
finally:
session.close()
if not videos:
raise ValueError("No videos found for batch download")
# 创建临时工作目录
with tempfile.TemporaryDirectory() as tmpdir:
tmpdir_path = Path(tmpdir)
zip_filename = f"videos-{len(videos)}-{video_ids[0][:8]}.zip"
zip_path = tmpdir_path / zip_filename
# 逐个下载视频并加入 zip
with zipfile.ZipFile(zip_path, "w", zipfile.ZIP_STORED) as zf:
for idx, video in enumerate(videos, 1):
logger.info("Batch download: downloading %d/%d %s", idx, len(videos), video.id)
try:
# 下载视频到临时文件
local_name = f"{idx:03d}_{video.name}"
local_path = tmpdir_path / local_name
# 使用 oss_helpers 的 download_asset,或者直接从 URL 下载
if video.file_url:
_download_video_to_file(video.file_url, str(local_path))
if local_path.exists() and local_path.stat().st_size > 0:
zf.write(str(local_path), arcname=local_name)
local_path.unlink(missing_ok=True)
else:
logger.warning("Video %s download failed, skipping", video.id)
except Exception as e:
logger.warning("Failed to download video %s: %s", video.id, e)
continue
# 上传 zip 到 OSS
if not zip_path.exists() or zip_path.stat().st_size == 0:
raise RuntimeError("Batch download zip file is empty")
zip_storage_key = f"batch-downloads/{uuid.uuid4().hex}/{zip_filename}"
download_url = upload_to_oss(str(zip_path), zip_storage_key)
total_size = zip_path.stat().st_size
file_count = len(zipfile.ZipFile(str(zip_path), "r").namelist())
logger.info(
"Batch download complete: %d files, %d bytes, url=%s",
file_count,
total_size,
download_url,
)
return {
"download_url": download_url,
"file_count": file_count,
"total_size": total_size,
"video_count": len(videos),
}
def _download_video_to_file(url: str, dest_path: str) -> None:
"""下载视频文件到本地路径。优先用 OSS SDK 走内网,回退到 HTTP 下载。"""
from video_processing.oss_helpers import download_asset
try:
# 尝试走 OSS 下载(如果是 OSS URL 的话)
success = download_asset(url, dest_path)
if success:
return
except Exception:
pass
# 回退到 HTTP 下载
import urllib.request
urllib.request.urlretrieve(url, dest_path) # nosec B310
@@ -68,6 +68,12 @@ def classify_asset(self, job_id: str) -> dict:
# Update asset with classification status and result
asset.classification_status = ClassificationStatus.COMPLETED
# 把分类结果写入 metadata,供列表筛选和智能视图使用
asset.metadata = {
**(asset.metadata or {}),
"classification": classification,
"classification_confidence": confidence,
}
asset_repo.update(asset)
session.commit()
+250 -1
View File
@@ -86,6 +86,36 @@ def _update_task_status(task_id: str, status_action: str, **kwargs) -> bool:
return False
def _build_error_info(error: Exception, stage: str = "render") -> dict:
"""构建结构化错误信息。
Args:
error: 异常对象
stage: 发生错误的阶段download/render/merge/upload等
Returns:
包含 error_type, message, stack_trace, stage, failed_at 的字典
"""
import traceback
from datetime import datetime, timezone
tb_str = traceback.format_exc()
# 截取堆栈前20行,避免字段过大
tb_lines = tb_str.strip().splitlines()
if len(tb_lines) > 20:
tb_summary = "\n".join(tb_lines[:20]) + f"\n... (truncated, total {len(tb_lines)} lines)"
else:
tb_summary = tb_str
return {
"error_type": type(error).__name__,
"message": str(error),
"stack_trace": tb_summary,
"stage": stage,
"failed_at": datetime.now(timezone.utc).isoformat(),
}
# ── 日志持久化辅助 ────────────────────────────────────────────────────────────
@@ -108,6 +138,7 @@ def _flush_logs(task_id: str, gen_task) -> None:
# ── 共享工具模块导入 ──────────────────────────────────────────────────────────
from services.asr_service_factory import get_asr_service
from video_processing.dedup_helpers import create_video_record_and_dedup
from video_processing.ffmpeg_utils import FFMPEG_BIN, probe_duration, run_ffmpeg
from video_processing.oss_helpers import (
@@ -304,6 +335,86 @@ def _download_voice_asset(voice_library_id: str, local_path: Path) -> bool:
return download_asset(storage_key, local_path)
def _prepare_bgm_track(
*,
bgm_config: dict,
temp_path: Path,
task_id: str = "",
) -> str | None:
"""准备 BGM 音频文件(下载到本地).
支持 3 种来源按优先级
1. audio_url 外部直链 URL最高优先级
2. asset_id 素材库中的音频素材
3. preset_id 预设 BGM
Returns:
BGM 本地文件路径准备失败返回 None
"""
from urllib.parse import urlparse
audio_url = bgm_config.get("audio_url", "") or ""
asset_id = bgm_config.get("asset_id", "") or ""
preset_id = bgm_config.get("preset_id", "") or ""
bgm_file = temp_path / f"bgm_{task_id or 'track'}.mp3"
# 优先级1:外部直链 URL
if audio_url:
try:
parsed = urlparse(audio_url)
if parsed.scheme in ("http", "https"):
import urllib.request
logger.info("[task_id=%s] [BGM] 从URL下载: %s", task_id, audio_url[:80])
urllib.request.urlretrieve(audio_url, bgm_file) # nosec B310
if bgm_file.exists() and bgm_file.stat().st_size > 0:
return str(bgm_file)
except Exception as e:
logger.warning("[task_id=%s] [BGM] URL下载失败: %s", task_id, e)
# 优先级2:素材库素材
if asset_id:
try:
from app.core.db import SessionLocal
from packages.adapters.sqlalchemy_impl.models import AssetModel
session = SessionLocal()
try:
model = session.query(AssetModel).filter(AssetModel.id == asset_id).first()
if model and model.file_url:
storage_key = model.file_url
logger.info("[task_id=%s] [BGM] 从素材库下载: asset_id=%s", task_id, asset_id)
ok = download_asset(storage_key, bgm_file)
if ok and bgm_file.exists() and bgm_file.stat().st_size > 0:
return str(bgm_file)
finally:
session.close()
except Exception as e:
logger.warning("[task_id=%s] [BGM] 素材库下载失败: %s", task_id, e)
# 优先级3:预设 BGM 库
if preset_id:
try:
from packages.domain.preset_bgm import get_preset_bgm
preset = get_preset_bgm(preset_id)
if preset and preset.audio_url:
import urllib.request
logger.info("[task_id=%s] [BGM] 从预设库下载: preset_id=%s", task_id, preset_id)
urllib.request.urlretrieve(preset.audio_url, bgm_file) # nosec B310
if bgm_file.exists() and bgm_file.stat().st_size > 0:
return str(bgm_file)
except Exception as e:
logger.warning("[task_id=%s] [BGM] 预设库下载失败: %s", task_id, e)
# 所有来源都失败
logger.warning("[task_id=%s] [BGM] 所有来源都无法获取BGM,跳过", task_id)
return None
def _verify_url_accessible(url: str, timeout: float = 10.0, retries: int = 2) -> bool:
"""HEAD 请求校验 URL 可访问(含重试,防止 OSS 抖动误报)。
@@ -853,6 +964,22 @@ def _render_video(
)
else:
logger.info("[task_id=%s] [渲染] unified 引擎 FFmpeg 渲染开始", task_id)
# ── 准备 BGM 音频 ──
bgm_path: str | None = None
plan_config = virtual_plan.config or {}
bgm_config = plan_config.get("bgm", {}) or {}
if bgm_config.get("enabled", False):
try:
bgm_path = _prepare_bgm_track(
bgm_config=bgm_config,
temp_path=temp_path,
task_id=task_id,
)
except Exception as bgm_err:
logger.warning("[task_id=%s] [BGM] 准备失败,跳过BGM: %s", task_id, bgm_err)
bgm_path = None
render_service = UnifiedRenderService(
plan=virtual_plan,
clips=virtual_clips,
@@ -861,6 +988,8 @@ def _render_video(
output_width=OUTPUT_WIDTH,
output_height=OUTPUT_HEIGHT,
output_fps=int(OUTPUT_FPS),
asr_service=get_asr_service(),
bgm_path=bgm_path,
)
render_result = render_service.render()
render_output_path = render_result.output_path
@@ -1092,6 +1221,70 @@ def generate_video(self, task_id: str) -> dict:
# ── 5. 标记完成 ──────────────────────────────────────────────────
_update_task_status(task_id, "mark_completed", result_count=video_count)
# 5.1 更新标题使用次数
try:
_title_session = SessionLocal()
try:
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
from packages.adapters.sqlalchemy_impl.title_library_repository import (
SQLAlchemyTitleLibraryRepository,
)
_task_repo = SQLAlchemyGenerationTaskRepository(_title_session)
_gen_task = _task_repo.get(task_id)
if _gen_task and _gen_task.title_ids and _gen_task.created_by_user_id:
_title_repo = SQLAlchemyTitleLibraryRepository(_title_session)
for _tid in _gen_task.title_ids:
try:
_title_repo.increment_usage_count(_tid, _gen_task.created_by_user_id)
except Exception:
logger.warning(
"[task_id=%s] 更新标题使用次数失败: title_id=%s",
task_id,
_tid,
exc_info=True,
)
finally:
_title_session.close()
except Exception:
logger.warning("[task_id=%s] 更新标题使用次数异常(不影响主流程)", task_id, exc_info=True)
# 5.2 更新素材使用次数 + 最近使用时间
try:
from worker_app.core.asset_usage import mark_asset_used_for_generation
_asset_session = SessionLocal()
try:
from packages.adapters.sqlalchemy_impl.asset_repository import (
SQLAlchemyAssetRepository,
)
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
_task_repo = SQLAlchemyGenerationTaskRepository(_asset_session)
_asset_repo = SQLAlchemyAssetRepository(_asset_session)
_gen_task = _task_repo.get(task_id)
if _gen_task and _gen_task.asset_ids:
for _aid in _gen_task.asset_ids:
try:
_asset = _asset_repo.get(_aid)
if _asset:
mark_asset_used_for_generation(_asset)
_asset_repo.update(_asset)
except Exception:
logger.warning(
"[task_id=%s] 更新素材使用次数失败: asset_id=%s",
task_id,
_aid,
exc_info=True,
)
finally:
_asset_session.close()
except Exception:
logger.warning("[task_id=%s] 更新素材使用次数异常(不影响主流程)", task_id, exc_info=True)
if gen_task:
gen_task.append_log(
"任务完成",
@@ -1122,6 +1315,9 @@ def generate_video(self, task_id: str) -> dict:
except Exception as error:
logger.error("[task_id=%s] [任务失败] %s", task_id, error, exc_info=True)
# 构建结构化错误信息
error_info = _build_error_info(error, stage="render")
# 记录失败日志
try:
_session = SessionLocal()
@@ -1134,6 +1330,7 @@ def generate_video(self, task_id: str) -> dict:
str(error),
level="ERROR",
error_type=type(error).__name__,
stage="render",
)
_flush_logs(task_id, gen_task)
finally:
@@ -1141,7 +1338,59 @@ def generate_video(self, task_id: str) -> dict:
except Exception:
logger.warning("[task_id=%s] 记录失败日志异常", task_id, exc_info=True)
_update_task_status(task_id, "mark_failed", error_message=str(error))
_update_task_status(
task_id,
"mark_failed",
error_message=str(error),
error_info=error_info,
)
# ── 自动重试逻辑 ──────────────────────────────────────────────────
try:
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
_s = SessionLocal()
try:
_r = SQLAlchemyGenerationTaskRepository(_s)
_task = _r.get(task_id)
if _task and _task.auto_retry_enabled and _task.auto_retry_max > 0:
current_retry = _task.retry_count or 0
if current_retry < _task.auto_retry_max:
logger.info(
"[task_id=%s] 触发自动重试: 当前重试次数=%d, 最大重试次数=%d",
task_id,
current_retry,
_task.auto_retry_max,
)
# 计算退避延迟(指数退避,基础5s,最大60s)
backoff_seconds = min(5 * (2**current_retry), 60)
# 原地重试
_task.mark_pending_from_failed()
_r.update(_task)
# 延迟重新入队
celery_app.send_task(
"worker.generate_video",
args=[task_id],
countdown=backoff_seconds,
)
logger.info(
"[task_id=%s] 自动重试已入队: 延迟=%ds, 第%d次重试",
task_id,
backoff_seconds,
current_retry + 1,
)
finally:
_s.close()
except Exception as retry_err:
logger.warning(
"[task_id=%s] 自动重试逻辑执行失败: %s",
task_id,
retry_err,
exc_info=True,
)
return {
"status": "failed",
"task_id": task_id,
+48
View File
@@ -858,6 +858,22 @@
"type": "VARCHAR(20)",
"unique": false
},
{
"index": false,
"name": "transition_duration",
"nullable": false,
"primary_key": false,
"type": "FLOAT",
"unique": false
},
{
"index": false,
"name": "playback_speed",
"nullable": false,
"primary_key": false,
"type": "FLOAT",
"unique": false
},
{
"index": true,
"name": "status",
@@ -1493,6 +1509,38 @@
"type": "TEXT",
"unique": false
},
{
"index": false,
"name": "error_info",
"nullable": true,
"primary_key": false,
"type": "JSON",
"unique": false
},
{
"index": false,
"name": "retry_count",
"nullable": false,
"primary_key": false,
"type": "INTEGER",
"unique": false
},
{
"index": false,
"name": "auto_retry_enabled",
"nullable": false,
"primary_key": false,
"type": "BOOLEAN",
"unique": false
},
{
"index": false,
"name": "auto_retry_max",
"nullable": false,
"primary_key": false,
"type": "INTEGER",
"unique": false
},
{
"index": false,
"name": "started_at",
+113
View File
@@ -0,0 +1,113 @@
"""Mock ASR 服务 — 用于测试和开发环境。
生成模拟的字幕时间轴不依赖真实ASR服务
"""
from __future__ import annotations
import re
from pathlib import Path
from typing import Optional
from packages.domain.subtitle import (
SubtitleSegment,
SubtitleTimeline,
SubtitleWord,
)
from packages.ports.asr_service import ASRService, ASRServiceError
class MockASRService(ASRService):
"""Mock ASR 服务,生成模拟字幕数据。
如果 audio_path 对应的目录下有同名 .txt 文件
就读取该文件内容作为字幕文本按时间均匀分段
否则生成默认的测试字幕
"""
def __init__(self, mock_text: Optional[str] = None):
self._mock_text = mock_text
def transcribe(
self,
audio_path: Path,
language: Optional[str] = None,
with_word_timestamps: bool = True,
) -> SubtitleTimeline:
if not audio_path.exists():
raise ASRServiceError(f"音频文件不存在: {audio_path}", provider="mock")
# 尝试读取同名 txt 文件作为字幕文本
text = self._mock_text
if text is None:
txt_path = audio_path.with_suffix(".txt")
if txt_path.exists():
text = txt_path.read_text(encoding="utf-8").strip()
else:
text = "这是一段测试字幕。它用于验证ASR自动字幕功能是否正常工作。每一句话都会被正确地分段并显示在视频底部。字幕的样式可以根据用户的喜好进行自定义调整。"
# 估算音频时长(用ffmpeg probe或者直接假设)
# mock模式下按字数估算,每秒4个字
total_duration = max(5.0, len(text) / 4.0)
segments = self._text_to_segments(text, total_duration, with_word_timestamps)
return SubtitleTimeline(
segments=segments,
language=language or "zh",
total_duration=total_duration,
)
def _text_to_segments(
self,
text: str,
total_duration: float,
with_word_timestamps: bool,
) -> list[SubtitleSegment]:
"""将文本按句切分成带时间轴的字幕片段。"""
# 按句末标点拆分
sentences = re.split(r"(?<=[。!?!?])", text)
sentences = [s.strip() for s in sentences if s.strip()]
if not sentences:
sentences = [text]
total_chars = sum(len(s) for s in sentences)
if total_chars == 0:
return []
segments = []
current_time = 0.0
for sentence in sentences:
char_count = len(sentence)
duration = total_duration * (char_count / total_chars)
end_time = current_time + duration
words: list[SubtitleWord] = []
if with_word_timestamps:
# 每个字作为一个词级单元(中文按字,英文按词)
word_time = current_time
word_duration = duration / char_count
for char in sentence:
words.append(
SubtitleWord(
text=char,
start=word_time,
end=word_time + word_duration,
)
)
word_time += word_duration
segments.append(
SubtitleSegment(
text=sentence,
start=current_time,
end=end_time,
words=words,
)
)
current_time = end_time
return segments
@@ -44,11 +44,61 @@ class InMemoryAssetRepository:
return False
def batch_delete(self, asset_ids: list[str]) -> int:
"""批量删除素材,返回实际删除数量。"""
"""批量删除素材(软删除,标记 status=deleted,返回实际影响数量。"""
from datetime import datetime, timezone
from packages.domain import AssetStatus
count = 0
for aid in asset_ids:
if aid in self._assets:
del self._assets[aid]
asset = self._assets.get(aid)
if asset and asset.status != AssetStatus.DELETED:
asset.status = AssetStatus.DELETED
asset.updated_at = datetime.now(timezone.utc)
count += 1
return count
def batch_update_metadata(self, asset_ids: list[str], metadata_patch: dict[str, object]) -> int:
"""批量更新素材 metadata(合并 patch),返回实际影响数量。"""
from datetime import datetime, timezone
count = 0
for aid in asset_ids:
asset = self._assets.get(aid)
if asset:
asset.metadata = {**asset.metadata, **metadata_patch}
asset.updated_at = datetime.now(timezone.utc)
count += 1
return count
def batch_add_tags(self, asset_ids: list[str], tag_ids: list[str]) -> int:
"""批量给素材添加标签(合并去重),返回实际影响数量。"""
from datetime import datetime, timezone
count = 0
for aid in asset_ids:
asset = self._assets.get(aid)
if asset:
changed = False
for tid in tag_ids:
if tid not in asset.tag_ids:
asset.tag_ids.append(tid)
changed = True
if changed:
asset.updated_at = datetime.now(timezone.utc)
count += 1
return count
def batch_replace_tags(self, asset_ids: list[str], tag_ids: list[str]) -> int:
"""批量替换素材标签(全量覆盖),返回实际影响数量。"""
from datetime import datetime, timezone
count = 0
for aid in asset_ids:
asset = self._assets.get(aid)
if asset:
asset.tag_ids = list(tag_ids)
asset.updated_at = datetime.now(timezone.utc)
count += 1
return count
+82 -2
View File
@@ -127,10 +127,90 @@ class SQLAlchemyAssetRepository:
return False
def batch_delete(self, asset_ids: list[str]) -> int:
"""批量删除素材,返回实际删除数量。"""
"""批量删除素材(软删除,标记 status=deleted,返回实际影响数量。"""
if not asset_ids:
return 0
count = self.session.query(AssetModel).filter(AssetModel.id.in_(asset_ids)).delete(synchronize_session=False)
from datetime import datetime, timezone
now = datetime.now(timezone.utc)
count = (
self.session.query(AssetModel)
.filter(AssetModel.id.in_(asset_ids), AssetModel.status != "deleted")
.update({AssetModel.status: "deleted", AssetModel.updated_at: now}, synchronize_session=False)
)
self.session.commit()
return count
def batch_update_metadata(self, asset_ids: list[str], metadata_patch: dict[str, object]) -> int:
"""批量更新素材 metadata(合并 patch),返回实际影响数量。"""
if not asset_ids:
return 0
from datetime import datetime, timezone
now = datetime.now(timezone.utc)
# 逐条读取 + 合并 + 更新,保证 JSON 合并正确
models = self.session.query(AssetModel).filter(AssetModel.id.in_(asset_ids)).all()
count = 0
for model in models:
existing = {}
if model.classification_result:
try:
existing = json.loads(model.classification_result)
except Exception:
existing = {}
merged = {**existing, **metadata_patch}
model.classification_result = json.dumps(merged, ensure_ascii=False)
model.updated_at = now
count += 1
self.session.commit()
return count
def batch_add_tags(self, asset_ids: list[str], tag_ids: list[str]) -> int:
"""批量给素材添加标签(合并去重),返回实际影响数量。"""
if not asset_ids or not tag_ids:
return 0
from datetime import datetime, timezone
now = datetime.now(timezone.utc)
clean_tag_ids = list(set(tag_ids))
count = 0
for aid in asset_ids:
# 查询现有标签
existing = {
row.tag_id
for row in self.session.query(AssetTagModel.tag_id).filter(AssetTagModel.asset_id == aid).all()
}
new_tags = [t for t in clean_tag_ids if t not in existing]
if new_tags:
for tid in new_tags:
self.session.add(AssetTagModel(asset_id=aid, tag_id=tid))
# 更新 updated_at
self.session.query(AssetModel).filter(AssetModel.id == aid).update(
{AssetModel.updated_at: now}, synchronize_session=False
)
count += 1
self.session.commit()
return count
def batch_replace_tags(self, asset_ids: list[str], tag_ids: list[str]) -> int:
"""批量替换素材标签(全量覆盖),返回实际影响数量。"""
if not asset_ids:
return 0
from datetime import datetime, timezone
now = datetime.now(timezone.utc)
clean_tag_ids = list(set(tag_ids))
count = 0
for aid in asset_ids:
# 先删再加
self.session.query(AssetTagModel).filter(AssetTagModel.asset_id == aid).delete(synchronize_session=False)
for tid in clean_tag_ids:
self.session.add(AssetTagModel(asset_id=aid, tag_id=tid))
# 更新 updated_at
self.session.query(AssetModel).filter(AssetModel.id == aid).update(
{AssetModel.updated_at: now}, synchronize_session=False
)
count += 1
self.session.commit()
return count
+6
View File
@@ -54,6 +54,8 @@ class SQLAlchemyEditPlanClipRepository:
start_time=clip.start_time,
duration=clip.duration,
transition_effect=clip.transition_effect,
transition_duration=clip.transition_duration,
playback_speed=clip.playback_speed,
status=clip.status,
config=clip.config,
)
@@ -76,6 +78,8 @@ class SQLAlchemyEditPlanClipRepository:
model.start_time = clip.start_time
model.duration = clip.duration
model.transition_effect = clip.transition_effect
model.transition_duration = clip.transition_duration
model.playback_speed = clip.playback_speed
model.status = clip.status
model.config = clip.config
model.updated_at = clip.updated_at
@@ -120,6 +124,8 @@ class SQLAlchemyEditPlanClipRepository:
start_time=model.start_time or 0.0,
duration=model.duration or 0.0,
transition_effect=model.transition_effect or "cut",
transition_duration=getattr(model, "transition_duration", 0.0) or 0.0,
playback_speed=model.playback_speed or 1.0,
status=EditPlanClipStatus(model.status) if model.status else EditPlanClipStatus.PENDING,
config=model.config or {},
created_at=model.created_at,
+57
View File
@@ -100,6 +100,63 @@ class SQLAlchemyGeneratedVideoRepository:
)
return [self._to_domain(model) for model in models]
def list_paginated(
self,
*,
project_id: str | None = None,
status: str | None = None,
review_status: str | None = None,
page: int = 1,
page_size: int = 20,
) -> tuple[list[GeneratedVideo], int]:
"""分页查询成片列表,支持按项目、状态、复核状态筛选。"""
query = self.session.query(GeneratedVideoModel)
if project_id:
query = query.filter(GeneratedVideoModel.project_id == project_id)
if status:
query = query.filter(GeneratedVideoModel.status == status)
if review_status:
query = query.filter(GeneratedVideoModel.review_status == review_status)
total = query.count()
models = (
query.order_by(GeneratedVideoModel.generated_at.desc())
.offset((page - 1) * page_size)
.limit(page_size)
.all()
)
return [self._to_domain(model) for model in models], total
def update_review_status(self, video_id: str, review_status: str) -> GeneratedVideo | None:
"""更新成片复核状态。"""
model = self.session.query(GeneratedVideoModel).filter(GeneratedVideoModel.id == video_id).first()
if model is None:
return None
model.review_status = review_status
self.session.add(model)
self.session.commit()
return self._to_domain(model)
def update_thumbnail(self, video_id: str, thumbnail_url: str) -> bool:
"""更新成片封面图URL。"""
model = self.session.query(GeneratedVideoModel).filter(GeneratedVideoModel.id == video_id).first()
if model is None:
return False
model.thumbnail_url = thumbnail_url
self.session.add(model)
self.session.commit()
return True
def get_by_ids(self, video_ids: list[str]) -> list[GeneratedVideo]:
"""批量获取成片记录。"""
if not video_ids:
return []
models = self.session.query(GeneratedVideoModel).filter(GeneratedVideoModel.id.in_(video_ids)).all()
return [self._to_domain(model) for model in models]
@staticmethod
def _to_domain(model: GeneratedVideoModel) -> GeneratedVideo:
return GeneratedVideo(
@@ -21,6 +21,10 @@ def _to_domain(model: GenerationTaskModel) -> GenerationTask:
progress=model.progress,
result_count=int(model.result_count or 0),
error_message=model.error_message,
error_info=dict(model.error_info) if model.error_info else {},
retry_count=model.retry_count or 0,
auto_retry_enabled=bool(model.auto_retry_enabled),
auto_retry_max=model.auto_retry_max or 0,
started_at=model.started_at,
completed_at=model.completed_at,
created_by_user_id=model.created_by_user_id,
@@ -51,6 +55,10 @@ class SQLAlchemyGenerationTaskRepository:
progress=task.progress,
result_count=task.result_count,
error_message=task.error_message,
error_info=task.error_info or None,
retry_count=task.retry_count or 0,
auto_retry_enabled=task.auto_retry_enabled,
auto_retry_max=task.auto_retry_max or 0,
started_at=task.started_at,
completed_at=task.completed_at,
created_by_user_id=task.created_by_user_id,
@@ -127,6 +135,68 @@ class SQLAlchemyGenerationTaskRepository:
)
return [_to_domain(m) for m in models]
def list_by_user_filtered(
self,
user_id: str,
*,
status: str | None = None,
limit: int | None = None,
offset: int = 0,
) -> list[GenerationTask]:
"""按用户+状态筛选任务列表。"""
query = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.created_by_user_id == user_id)
if status:
query = query.filter(GenerationTaskModel.status == status)
query = query.order_by(GenerationTaskModel.created_at.desc())
if offset:
query = query.offset(offset)
if limit:
query = query.limit(limit)
return [_to_domain(m) for m in query.all()]
def count_by_user_filtered(
self,
user_id: str,
*,
status: str | None = None,
) -> int:
"""按用户+状态筛选计数。"""
query = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.created_by_user_id == user_id)
if status:
query = query.filter(GenerationTaskModel.status == status)
return query.count()
def list_by_project_filtered(
self,
project_id: str,
*,
status: str | None = None,
limit: int | None = None,
offset: int = 0,
) -> list[GenerationTask]:
"""按项目+状态筛选任务列表。"""
query = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.project_id == project_id)
if status:
query = query.filter(GenerationTaskModel.status == status)
query = query.order_by(GenerationTaskModel.created_at.desc())
if offset:
query = query.offset(offset)
if limit:
query = query.limit(limit)
return [_to_domain(m) for m in query.all()]
def count_by_project_filtered(
self,
project_id: str,
*,
status: str | None = None,
) -> int:
"""按项目+状态筛选计数。"""
query = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.project_id == project_id)
if status:
query = query.filter(GenerationTaskModel.status == status)
return query.count()
def update(self, task: GenerationTask) -> GenerationTask:
model = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.id == task.id).first()
if model is None:
@@ -143,6 +213,10 @@ class SQLAlchemyGenerationTaskRepository:
model.progress = task.progress
model.result_count = task.result_count
model.error_message = task.error_message
model.error_info = task.error_info or None
model.retry_count = task.retry_count or 0
model.auto_retry_enabled = task.auto_retry_enabled
model.auto_retry_max = task.auto_retry_max or 0
model.started_at = task.started_at
model.completed_at = task.completed_at
model.source_edit_plan_id = task.source_edit_plan_id or None
@@ -196,6 +196,8 @@ class EditPlanClipModel(Base):
start_time = Column(Float, nullable=False, default=0.0)
duration = Column(Float, nullable=False, default=0.0)
transition_effect = Column(String(20), nullable=False, default="cut")
transition_duration = Column(Float, nullable=False, default=0.0)
playback_speed = Column(Float, nullable=False, default=1.0)
status = Column(String(20), nullable=False, default="pending", index=True)
config = Column(JSON, nullable=False, default=dict)
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
@@ -250,6 +252,10 @@ class GenerationTaskModel(Base):
progress = Column(Float, nullable=False, default=0.0)
result_count = Column(Float, nullable=False, default=0)
error_message = Column(Text, nullable=False, default="")
error_info = Column(JSON, nullable=True)
retry_count = Column(Integer, nullable=False, default=0)
auto_retry_enabled = Column(Boolean, nullable=False, default=False)
auto_retry_max = Column(Integer, nullable=False, default=0)
started_at = Column(DateTime, nullable=True)
completed_at = Column(DateTime, nullable=True)
created_by_user_id = Column(String(36), nullable=False, default="", index=True)
+118 -18
View File
@@ -2,11 +2,14 @@
from __future__ import annotations
import uuid
from typing import List, Optional
from sqlalchemy import func
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import (
EditPlanModel,
TemplateCategoryModel,
TemplateModel,
TemplateSegmentModel,
@@ -28,18 +31,26 @@ class SQLAlchemyTemplateRepository:
*,
skip: int = 0,
limit: int = 50,
category: Optional[str] = None,
tag: Optional[str] = None,
keyword: Optional[str] = None,
mode: Optional[str] = None,
) -> List[Template]:
models = (
self.session.query(TemplateModel)
.filter(
TemplateModel.user_id == user_id,
TemplateModel.is_active.is_(True),
)
.order_by(TemplateModel.created_at.desc())
.offset(skip)
.limit(limit)
.all()
query = self.session.query(TemplateModel).filter(
TemplateModel.user_id == user_id,
TemplateModel.is_active.is_(True),
)
if category:
query = query.filter(TemplateModel.category == category)
if mode:
query = query.filter(TemplateModel.mode == mode)
if keyword:
like_pattern = f"%{keyword}%"
query = query.filter(TemplateModel.name.like(like_pattern))
if tag:
# JSON 数组包含指定标签(MySQL JSON_CONTAINS / SQLite json_each 兼容写法用 LIKE
query = query.filter(TemplateModel.tags.like(f'%"{tag}"%'))
models = query.order_by(TemplateModel.created_at.desc()).offset(skip).limit(limit).all()
templates = [self._model_to_entity(m) for m in models]
# 批量加载所有 segments,避免 N+1 查询
if templates:
@@ -142,15 +153,77 @@ class SQLAlchemyTemplateRepository:
self.session.commit()
return True
def count_by_user(self, user_id: str) -> int:
return (
self.session.query(TemplateModel)
.filter(
TemplateModel.user_id == user_id,
TemplateModel.is_active.is_(True),
)
.count()
def count_by_user(
self,
user_id: str,
*,
category: Optional[str] = None,
tag: Optional[str] = None,
keyword: Optional[str] = None,
mode: Optional[str] = None,
) -> int:
query = self.session.query(TemplateModel).filter(
TemplateModel.user_id == user_id,
TemplateModel.is_active.is_(True),
)
if category:
query = query.filter(TemplateModel.category == category)
if mode:
query = query.filter(TemplateModel.mode == mode)
if keyword:
query = query.filter(TemplateModel.name.like(f"%{keyword}%"))
if tag:
query = query.filter(TemplateModel.tags.like(f'%"{tag}"%'))
return query.count()
def copy_template(self, template_id: str, user_id: str, new_name: str) -> Template:
"""复制模板(含所有 segments)。"""
source = self.get(template_id, user_id)
if source is None:
raise ValueError(f"Template {template_id} not found")
new_id = str(uuid.uuid4())
new_template = Template(
id=new_id,
user_id=user_id,
name=new_name,
mode=source.mode,
category=source.category,
tags=list(source.tags),
title_config=dict(source.title_config),
subtitle_config=dict(source.subtitle_config),
bgm_config=dict(source.bgm_config),
estimated_duration=source.estimated_duration,
is_active=True,
)
created = self.create(new_template)
# 复制 segments
new_segments: List[TemplateSegment] = []
for seg in source.segments:
new_seg = TemplateSegment(
id=str(uuid.uuid4()),
template_id=new_id,
segment_order=seg.segment_order,
duration_min=seg.duration_min,
duration_max=seg.duration_max,
material_type=seg.material_type,
)
new_segments.append(new_seg)
model = TemplateSegmentModel(
id=new_seg.id,
template_id=new_seg.template_id,
segment_order=new_seg.segment_order,
duration_min=new_seg.duration_min,
duration_max=new_seg.duration_max,
material_type=new_seg.material_type,
)
self.session.add(model)
if new_segments:
self.session.commit()
created.segments = new_segments
return created
# ── Segments ──
@@ -234,6 +307,33 @@ class SQLAlchemyTemplateRepository:
self.session.commit()
return True
# ── Tags ──
def list_tags(self, user_id: str) -> List[str]:
"""获取用户所有模板的标签(去重)。"""
models = (
self.session.query(TemplateModel)
.filter(
TemplateModel.user_id == user_id,
TemplateModel.is_active.is_(True),
TemplateModel.tags.isnot(None),
)
.all()
)
tags_set: set[str] = set()
for m in models:
if m.tags:
for t in m.tags:
if t:
tags_set.add(t)
return sorted(tags_set)
# ── Usage Stats ──
def get_usage_count(self, template_id: str) -> int:
"""获取模板被使用的次数(关联的剪辑计划数量)。"""
return self.session.query(EditPlanModel).filter(EditPlanModel.template_id == template_id).count()
# ── Mapping helpers ──
@staticmethod
+19
View File
@@ -103,6 +103,25 @@ class SQLAlchemyTitleLibraryRepository:
self.session.commit()
return True
def increment_usage_count(self, title_id: str, user_id: str, increment: int = 1) -> bool:
"""递增标题使用次数。返回是否成功。"""
from sqlalchemy import func
model = (
self.session.query(TitleLibraryModel)
.filter(
TitleLibraryModel.id == title_id,
TitleLibraryModel.user_id == user_id,
)
.first()
)
if model is None:
return False
model.usage_count = (model.usage_count or 0) + increment
model.updated_at = func.now()
self.session.commit()
return True
def count_by_user(self, user_id: str, is_active: bool = True) -> int:
return (
self.session.query(TitleLibraryModel)
+231
View File
@@ -0,0 +1,231 @@
"""Mock TTS 服务实现.
使用 FFmpeg 合成简单音频模拟人声
- 不同音色用不同的基频sine 波频率
- 语速通过 atempo 调整
- 语调通过 asetrate 调整
- 加一点 tremolo 效果让声音更自然
用于开发测试不依赖外部 TTS 服务
"""
from __future__ import annotations
import logging
import subprocess
import tempfile
from pathlib import Path
from packages.domain.voice_presets import get_voice, list_voices
from packages.ports.tts_service import TtsError, TtsService
logger = logging.getLogger(__name__)
# Mock 时长估算:每字约 0.3 秒(中文)
_CHARS_PER_SECOND = 3.3
class MockTtsService(TtsService):
"""Mock TTS 服务 — 用 FFmpeg 合成测试音频."""
def __init__(self, ffmpeg_bin: str = "ffmpeg") -> None:
self._ffmpeg_bin = ffmpeg_bin
@property
def provider_name(self) -> str:
return "mock"
def available_voices(self) -> list[str]:
return [v.voice_id for v in list_voices(provider="mock")]
def synthesize(
self,
text: str,
*,
voice_id: str = "",
speed: float = 1.0,
pitch: float = 0.0,
output_path: Path | None = None,
sample_rate: int = 22050,
format: str = "wav",
) -> Path:
"""合成 Mock 音频.
FFmpeg sine 波合成带轻微调制的音频模拟人声
时长根据文本长度估算
"""
if not text.strip():
raise TtsError("文本不能为空")
# 语速边界
if speed <= 0:
speed = 1.0
speed = max(0.5, min(2.0, speed))
# 语调边界
pitch = max(-12, min(12, pitch))
# 解析音色
voice = get_voice(voice_id) if voice_id else get_voice("female_warm")
if voice is None:
voice = get_voice("female_warm")
# 计算基频(从 provider_voice_id 里提取,或者按音色默认)
base_freq = self._extract_freq(voice.provider_voice_id, voice.gender.value)
# 计算时长(按文本长度)
duration = self.estimate_duration(text, speed=speed)
duration = max(0.5, duration) # 最短 0.5 秒
# 输出路径
if output_path is None:
suffix = f".{format}"
tmp = tempfile.NamedTemporaryFile(suffix=suffix, delete=False)
tmp.close()
output_path = Path(tmp.name)
output_path.parent.mkdir(parents=True, exist_ok=True)
try:
self._synthesize_with_ffmpeg(
output_path=output_path,
base_freq=base_freq,
duration=duration,
speed=speed,
pitch=pitch,
sample_rate=sample_rate,
format=format,
)
except Exception as e:
logger.error("Mock TTS 合成失败: %s", e)
raise TtsError(f"Mock TTS 合成失败: {e}") from e
return output_path
def estimate_duration(self, text: str, *, speed: float = 1.0) -> float:
"""估算音频时长.
按中文字符数估算每字约 0.3
"""
if not text:
return 0.0
# 去除空白后的字符数
char_count = len([c for c in text if not c.isspace()])
if char_count == 0:
return 0.0
base_duration = char_count / _CHARS_PER_SECOND
return base_duration / max(0.1, speed)
def _extract_freq(self, provider_voice_id: str, gender: str) -> float:
"""从 provider_voice_id 提取基频,或按性别给默认值."""
if provider_voice_id.startswith("sine_"):
try:
return float(provider_voice_id.split("_")[1])
except (IndexError, ValueError):
pass
# 按性别给默认基频
if gender == "male":
return 120.0
elif gender == "child":
return 350.0
else: # female
return 220.0
def _synthesize_with_ffmpeg(
self,
*,
output_path: Path,
base_freq: float,
duration: float,
speed: float,
pitch: float,
sample_rate: int,
format: str,
) -> None:
"""使用 FFmpeg 合成音频.
效果链
1. sine 波生成基频
2. tremolo 增加轻微颤音
3. aeval 模拟简单的音色变化让声音不那么单调
4. atempo 调整语速
5. asetrate 调整语调
6. volume 调整音量
"""
# 语调频率偏移因子(每半音 = 2^(1/12) ≈ 1.05946
pitch_factor = 2 ** (pitch / 12)
# 颤音参数
tremolo_freq = 5.0 # 5Hz 颤音
tremolo_depth = 0.3 # 30% 深度
# 构建滤镜链
filters: list[str] = []
# 生成基频 + 泛音(让声音更丰富)
# 用多个 sine 波叠加模拟更自然的音色
filter_parts = []
# 主音 + 轻微频率调制
filter_parts.append(f"sine=frequency={base_freq}:duration={duration}:sample_rate={sample_rate}")
# 颤音效果
filter_parts.append(f"tremolo=f={tremolo_freq}:d={tremolo_depth}")
# 语速调整(同时调整时长)
if abs(speed - 1.0) > 0.01:
filter_parts.append(f"atempo={speed:.3f}")
# 语调调整(通过采样率变化实现,同时补偿时长)
if abs(pitch) > 0.01:
new_rate = int(sample_rate * pitch_factor)
filter_parts.append(f"asetrate={new_rate}")
filter_parts.append(f"aresample={sample_rate}")
# 音量包络:淡入淡出
fade_in = min(0.05, duration * 0.1)
fade_out = min(0.1, duration * 0.2)
filter_parts.append(f"afade=t=in:d={fade_in}")
filter_parts.append(f"afade=t=out:st={max(0, duration - fade_out)}:d={fade_out}")
# 音量调整到合适大小
filter_parts.append("volume=0.3")
filter_complex = ",".join(filter_parts)
# 编码参数
if format == "mp3":
codec_args = ["-acodec", "libmp3lame", "-b:a", "128k"]
else:
codec_args = ["-acodec", "pcm_s16le"]
command = [
self._ffmpeg_bin,
"-y",
"-f",
"lavfi",
"-i",
filter_complex,
*codec_args,
"-ar",
str(sample_rate),
"-ac",
"1",
str(output_path),
]
logger.debug("Mock TTS FFmpeg 命令: %s", " ".join(command))
result = subprocess.run(
command,
capture_output=True,
text=True,
timeout=max(30, duration * 2 + 10),
)
if result.returncode != 0:
raise TtsError(f"FFmpeg 合成失败: {result.stderr[-500:]}")
if not output_path.exists() or output_path.stat().st_size == 0:
raise TtsError("输出文件为空或不存在")
+12
View File
@@ -21,13 +21,19 @@ from .duplication import (
from .generated_videos import (
GetGeneratedVideoDownloadUrlUseCase,
GetGeneratedVideoUseCase,
GetVideosByIdsUseCase,
ListGeneratedVideosByTaskUseCase,
ListGeneratedVideosPaginatedUseCase,
ListGeneratedVideosUseCase,
UpdateVideoReviewStatusUseCase,
)
from .generation_tasks import (
CreateGenerationTaskCommand,
CreateGenerationTaskUseCase,
GetGenerationTaskUseCase,
ListGenerationTasksResult,
ListUserTasksFilteredUseCase,
RetryGenerationTaskUseCase,
)
from .ingest_jobs import SubmitIngestJobCommand, SubmitIngestJobUseCase
from .jobs import (
@@ -65,6 +71,9 @@ __all__ = [
"CreateGenerationTaskCommand",
"CreateGenerationTaskUseCase",
"GetGenerationTaskUseCase",
"ListGenerationTasksResult",
"ListUserTasksFilteredUseCase",
"RetryGenerationTaskUseCase",
"CreateJobCommand",
"CreateJobUseCase",
"CreateProjectCommand",
@@ -76,6 +85,7 @@ __all__ = [
"GetDuplicationDetailUseCase",
"GetGeneratedVideoDownloadUrlUseCase",
"GetGeneratedVideoUseCase",
"GetVideosByIdsUseCase",
"GetJobStatisticsUseCase",
"GetJobUseCase",
"GetProjectUseCase",
@@ -83,6 +93,7 @@ __all__ = [
"ListAssetsUseCase",
"ListDuplicationRecordsUseCase",
"ListGeneratedVideosByTaskUseCase",
"ListGeneratedVideosPaginatedUseCase",
"ListGeneratedVideosUseCase",
"ListJobsUseCase",
"ListProjectsUseCase",
@@ -95,6 +106,7 @@ __all__ = [
"SubmitJobUseCase",
"UpdateJobProgressCommand",
"UpdateJobProgressUseCase",
"UpdateVideoReviewStatusUseCase",
"UploadForDuplicationCommand",
"UploadForDuplicationUseCase",
]
+46
View File
@@ -14,6 +14,32 @@ class ListGeneratedVideosUseCase:
return self.generated_video_repository.list_by_project(project_id.strip())
class ListGeneratedVideosPaginatedUseCase:
def __init__(self, generated_video_repository: GeneratedVideoRepository):
self.generated_video_repository = generated_video_repository
def execute(
self,
*,
project_id: str | None = None,
status: str | None = None,
review_status: str | None = None,
page: int = 1,
page_size: int = 20,
) -> tuple[list[GeneratedVideo], int]:
if page < 1:
page = 1
if page_size < 1 or page_size > 100:
page_size = 20
return self.generated_video_repository.list_paginated(
project_id=project_id,
status=status,
review_status=review_status,
page=page,
page_size=page_size,
)
class GetGeneratedVideoUseCase:
def __init__(self, generated_video_repository: GeneratedVideoRepository):
self.generated_video_repository = generated_video_repository
@@ -41,3 +67,23 @@ class GetGeneratedVideoDownloadUrlUseCase:
if item is None:
return None
return item.file_url
class UpdateVideoReviewStatusUseCase:
def __init__(self, generated_video_repository: GeneratedVideoRepository):
self.generated_video_repository = generated_video_repository
def execute(self, video_id: str, review_status: str) -> GeneratedVideo | None:
if not video_id.strip():
raise ValueError("video_id 不能为空")
if review_status not in ("pending_review", "approved", "rejected"):
raise ValueError(f"无效的 review_status: {review_status}")
return self.generated_video_repository.update_review_status(video_id.strip(), review_status)
class GetVideosByIdsUseCase:
def __init__(self, generated_video_repository: GeneratedVideoRepository):
self.generated_video_repository = generated_video_repository
def execute(self, video_ids: list[str]) -> list[GeneratedVideo]:
return self.generated_video_repository.get_by_ids(video_ids)
+66 -2
View File
@@ -21,6 +21,8 @@ class CreateGenerationTaskCommand:
source_edit_plan_id: str = ""
asset_select_mode: str = ""
batch_id: str = ""
auto_retry_enabled: bool = False
auto_retry_max: int = 0
class CreateGenerationTaskUseCase:
@@ -42,12 +44,12 @@ class CreateGenerationTaskUseCase:
progress=0.0,
result_count=0,
error_message="",
started_at=None,
completed_at=None,
created_by_user_id=command.created_by_user_id,
source_edit_plan_id=command.source_edit_plan_id,
asset_select_mode=command.asset_select_mode,
batch_id=command.batch_id,
auto_retry_enabled=command.auto_retry_enabled,
auto_retry_max=command.auto_retry_max,
)
return self.generation_task_repository.create(task)
@@ -58,3 +60,65 @@ class GetGenerationTaskUseCase:
def execute(self, task_id: str) -> GenerationTask | None:
return self.generation_task_repository.get(task_id)
@dataclass(slots=True)
class ListTasksFilter:
"""任务列表筛选条件。"""
status: str | None = None # pending, running, completed, failed, cancelled
@dataclass(slots=True)
class ListGenerationTasksResult:
"""带筛选和分页的任务列表结果。"""
items: list[GenerationTask]
total: int
class ListUserTasksFilteredUseCase:
"""按用户+筛选条件查询任务列表。"""
def __init__(self, generation_task_repository: GenerationTaskRepository):
self.generation_task_repository = generation_task_repository
def execute(
self,
user_id: str,
*,
status: str | None = None,
limit: int | None = None,
offset: int = 0,
) -> ListGenerationTasksResult:
items = self.generation_task_repository.list_by_user_filtered(
user_id,
status=status,
limit=limit,
offset=offset,
)
total = self.generation_task_repository.count_by_user_filtered(
user_id,
status=status,
)
return ListGenerationTasksResult(items=items, total=total)
class RetryGenerationTaskUseCase:
"""原地重试失败的任务(重置状态+递增retry_count)。
与创建新任务不同复用同一个 task_id保留历史关联
"""
def __init__(self, generation_task_repository: GenerationTaskRepository):
self.generation_task_repository = generation_task_repository
def execute(self, task_id: str) -> GenerationTask:
task = self.generation_task_repository.get(task_id)
if task is None:
raise ValueError(f"任务不存在: {task_id}")
if not task.is_failed:
raise ValueError(f"只有失败状态的任务才能重试,当前状态: {task.status.value}")
task.mark_pending_from_failed()
self.generation_task_repository.update(task)
return task
+15
View File
@@ -49,6 +49,21 @@ class CreateCategoryCommand:
name: str
@dataclass
class CopyTemplateCommand:
template_id: str
user_id: str
new_name: str
@dataclass
class ListTemplatesFilter:
category: Optional[str] = None
tag: Optional[str] = None
keyword: Optional[str] = None
mode: Optional[str] = None
@dataclass
class ValidateTemplateCommand:
template_id: str
+74 -1
View File
@@ -7,8 +7,10 @@ from dataclasses import dataclass, field
from typing import List, Optional
from packages.application.template.commands import (
CopyTemplateCommand,
CreateCategoryCommand,
CreateTemplateCommand,
ListTemplatesFilter,
UpdateTemplateCommand,
ValidateTemplateCommand,
)
@@ -102,8 +104,40 @@ class ListTemplatesUseCase:
*,
skip: int = 0,
limit: int = 50,
filter: Optional[ListTemplatesFilter] = None,
) -> List[Template]:
return self.repository.list_by_user(user_id, skip=skip, limit=limit)
if filter is None:
return self.repository.list_by_user(user_id, skip=skip, limit=limit)
return self.repository.list_by_user(
user_id,
skip=skip,
limit=limit,
category=filter.category,
tag=filter.tag,
keyword=filter.keyword,
mode=filter.mode,
)
class CountTemplatesUseCase:
def __init__(self, repository: TemplateRepositoryPort) -> None:
self.repository = repository
def execute(
self,
user_id: str,
*,
filter: Optional[ListTemplatesFilter] = None,
) -> int:
if filter is None:
return self.repository.count_by_user(user_id)
return self.repository.count_by_user(
user_id,
category=filter.category,
tag=filter.tag,
keyword=filter.keyword,
mode=filter.mode,
)
class GetTemplateUseCase:
@@ -175,6 +209,23 @@ class DeleteTemplateUseCase:
return self.repository.delete(template_id, user_id)
class CopyTemplateUseCase:
def __init__(self, repository: TemplateRepositoryPort) -> None:
self.repository = repository
def execute(self, command: CopyTemplateCommand) -> Template:
existing = self.repository.get(command.template_id, command.user_id)
if existing is None:
raise NotFoundError(f"Template {command.template_id} not found")
if not command.new_name or not command.new_name.strip():
raise ValidationError("新模板名称不能为空")
return self.repository.copy_template(
command.template_id,
command.user_id,
command.new_name.strip(),
)
# ── Validate template ──
@@ -258,3 +309,25 @@ class DeleteCategoryUseCase:
def execute(self, category_id: str, user_id: str) -> bool:
return self.repository.delete_category(category_id, user_id)
# ── Tags ──
class ListTagsUseCase:
def __init__(self, repository: TemplateRepositoryPort) -> None:
self.repository = repository
def execute(self, user_id: str) -> List[str]:
return self.repository.list_tags(user_id)
# ── Usage Stats ──
class GetTemplateUsageUseCase:
def __init__(self, repository: TemplateRepositoryPort) -> None:
self.repository = repository
def execute(self, template_id: str) -> int:
return self.repository.get_usage_count(template_id)
+7
View File
@@ -1,11 +1,14 @@
"""Title library application module."""
from packages.application.title_library.commands import IncrementTitleUsageCommand, PickTitleCommand
from packages.application.title_library.use_cases import (
CreateTitleLibraryUseCase,
DeleteTitleLibraryUseCase,
GetTitleLibraryUseCase,
IncrementTitleUsageUseCase,
ListTitleLibraryUseCase,
NotFoundError,
PickTitleUseCase,
QuotaExceededError,
UpdateTitleLibraryUseCase,
)
@@ -14,7 +17,11 @@ __all__ = [
"CreateTitleLibraryUseCase",
"DeleteTitleLibraryUseCase",
"GetTitleLibraryUseCase",
"IncrementTitleUsageUseCase",
"IncrementTitleUsageCommand",
"ListTitleLibraryUseCase",
"PickTitleUseCase",
"PickTitleCommand",
"UpdateTitleLibraryUseCase",
"QuotaExceededError",
"NotFoundError",
+14
View File
@@ -28,3 +28,17 @@ class UpdateTitleLibraryCommand:
tags: Optional[List[str]] = None
is_active: Optional[bool] = None
metadata_: Optional[dict] = None
@dataclass
class IncrementTitleUsageCommand:
title_id: str
user_id: str
increment: int = 1
@dataclass
class PickTitleCommand:
user_id: str
category: Optional[str] = None
exclude_ids: List[str] = field(default_factory=list)
+64
View File
@@ -8,6 +8,8 @@ from typing import List, Optional
from packages.adapters.sqlalchemy_impl.title_library_repository import SQLAlchemyTitleLibraryRepository
from packages.application.title_library.commands import (
CreateTitleLibraryCommand,
IncrementTitleUsageCommand,
PickTitleCommand,
UpdateTitleLibraryCommand,
)
from packages.domain.quota import QuotaDimension, quota_checker
@@ -100,6 +102,68 @@ class DeleteTitleLibraryUseCase:
return self.repository.delete(title_id, user_id)
class IncrementTitleUsageUseCase:
"""递增标题使用次数。用于生成视频成功后,更新标题的使用统计。"""
def __init__(self, repository: SQLAlchemyTitleLibraryRepository) -> None:
self.repository = repository
def execute(self, command: IncrementTitleUsageCommand) -> bool:
if command.increment <= 0:
return False
return self.repository.increment_usage_count(
command.title_id,
command.user_id,
increment=command.increment,
)
class PickTitleUseCase:
"""智能选择一个标题。
策略
1. 可选按 category 过滤
2. 排除指定的 title_ids如本轮已用过的
3. 按使用次数升序取最少的前 5
4. 从中随机选一个增加多样性
5. 无可用标题时返回 None
"""
_CANDIDATE_POOL_SIZE = 5
def __init__(self, repository: SQLAlchemyTitleLibraryRepository) -> None:
self.repository = repository
def execute(self, command: PickTitleCommand) -> TitleLibraryItem | None:
import random
# 取该用户所有活跃标题(或指定分类)
all_titles = self.repository.list_by_user(
command.user_id,
category=command.category,
is_active=True,
skip=0,
limit=500, # 取足够多的候选
)
if not all_titles:
return None
# 排除已使用/指定排除的
exclude_set = set(command.exclude_ids or [])
candidates = [t for t in all_titles if t.id not in exclude_set]
if not candidates:
# 排除后没了,就从全部里选
candidates = all_titles
# 按使用次数升序,取最少的前 N 个
candidates.sort(key=lambda t: t.usage_count)
pool = candidates[: self._CANDIDATE_POOL_SIZE]
# 随机选一个
return random.choice(pool)
class QuotaExceededError(Exception):
def __init__(self, dimension: str, limit: float, used: float) -> None:
self.dimension = dimension
+31 -2
View File
@@ -112,14 +112,32 @@ class SubtitleConfig(BaseModel):
color: str = Field(default="#ffffff", description="文字颜色 (HEX)")
size: int = Field(default=24, ge=12, le=60, description="字号")
animation: TextAnimation = Field(default=TextAnimation.FADE_IN, description="入场动画")
# ASR 自动字幕
auto_generated: bool = Field(default=False, description="是否启用ASR自动生成字幕")
language: str = Field(default="", description="字幕语言,空字符串表示自动检测(如 zh/en/ja)")
max_chars_per_line: int = Field(default=20, ge=8, le=40, description="每行最多字符数")
min_chars_per_segment: int = Field(default=8, ge=2, le=20, description="每段最少字符数(低于则合并)")
class BGMConfig(BaseModel):
"""BGM 配置"""
enabled: bool = Field(default=False, description="是否启用 BGM")
source: BGMSource = Field(default=BGMSource.LIBRARY, description="BGM 来源")
asset_id: str = Field(default="", description="BGM 素材 ID")
volume: float = Field(default=0.3, ge=0.0, le=1.0, description="音量 (0.0 ~ 1.0)")
asset_id: str = Field(default="", description="BGM 素材 ID(来源为 library/upload 时使用)")
preset_id: str = Field(default="", description="预设 BGM ID(来源为 ai_recommend 或使用内置库时使用)")
audio_url: str = Field(default="", description="BGM 音频 URL(外部直链,优先级最高)")
volume: float = Field(default=0.3, ge=0.0, le=1.0, description="BGM 音量 (0.0 ~ 1.0)")
fade_in: float = Field(default=0.0, ge=0.0, le=30.0, description="淡入时长(秒)")
fade_out: float = Field(default=0.0, ge=0.0, le=30.0, description="淡出时长(秒)")
loop_enabled: bool = Field(default=True, description="BGM 是否循环播放以铺满整个视频时长")
sidechain_enabled: bool = Field(default=False, description="是否启用人声闪避(有人声时 BGM 自动降低音量)")
sidechain_ratio: float = Field(
default=0.3, ge=0.0, le=1.0, description="人声闪避时 BGM 音量降低比例(0.3 = 降低30%"
)
sidechain_attack: float = Field(default=0.02, ge=0.001, le=1.0, description="人声闪避攻击时间(秒)")
sidechain_release: float = Field(default=0.5, ge=0.01, le=5.0, description="人声闪避释放时间(秒)")
sidechain_threshold: float = Field(default=-25.0, ge=-60.0, le=0.0, description="人声闪避触发阈值(dB")
# ── 完整 config 模型 ─────────────────────────────────────────────────────────
@@ -185,9 +203,20 @@ DEFAULT_EDIT_PLAN_CONFIG: dict = {
"animation": "fade_in",
},
"bgm": {
"enabled": False,
"source": "library",
"asset_id": "",
"preset_id": "",
"audio_url": "",
"volume": 0.3,
"fade_in": 0.0,
"fade_out": 0.0,
"loop_enabled": True,
"sidechain_enabled": False,
"sidechain_ratio": 0.3,
"sidechain_attack": 0.02,
"sidechain_release": 0.5,
"sidechain_threshold": -25.0,
},
"editing_mode": "one_take",
}
+13
View File
@@ -50,6 +50,8 @@ class EditPlanClip:
start_time: float = 0.0
duration: float = 0.0
transition_effect: str = "cut"
transition_duration: float = 0.0 # 0 表示使用全局默认值
playback_speed: float = 1.0 # 0 或 1.0 表示原速,范围 0.25~4.0
status: EditPlanClipStatus = EditPlanClipStatus.PENDING
config: dict[str, Any] = field(default_factory=dict)
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
@@ -68,6 +70,8 @@ class EditPlanClip:
start_time: float = 0.0,
duration: float = 0.0,
transition_effect: str = "cut",
transition_duration: float = 0.0,
playback_speed: float = 1.0,
config: dict[str, Any] | None = None,
) -> EditPlanClip:
"""创建剪辑计划片段"""
@@ -79,6 +83,13 @@ class EditPlanClip:
raise ValueError("start_time 不能为负数")
if duration < 0:
raise ValueError("duration 不能为负数")
# 速度边界钳制
if playback_speed <= 0:
playback_speed = 1.0
elif playback_speed < 0.25:
playback_speed = 0.25
elif playback_speed > 4.0:
playback_speed = 4.0
return cls(
id=uuid4().hex,
@@ -91,6 +102,8 @@ class EditPlanClip:
start_time=start_time,
duration=duration,
transition_effect=transition_effect.strip() or "cut",
transition_duration=max(0.0, transition_duration),
playback_speed=playback_speed,
status=EditPlanClipStatus.PENDING,
config=config or {},
)
+1
View File
@@ -133,6 +133,7 @@ class AssetStatus(StrEnum):
READY = "ready"
PROCESSING = "processing"
ERROR = "error"
DELETED = "deleted"
@classmethod
def _missing_(cls, value: object) -> "AssetStatus":
+23 -3
View File
@@ -80,6 +80,10 @@ class GenerationTask:
progress: float = 0.0
result_count: int = 0
error_message: str = ""
error_info: dict = field(default_factory=dict)
retry_count: int = 0
auto_retry_enabled: bool = False
auto_retry_max: int = 0
started_at: datetime | None = None
completed_at: datetime | None = None
source_edit_plan_id: str = ""
@@ -105,6 +109,8 @@ class GenerationTask:
source_edit_plan_id: str = "",
asset_select_mode: str = "",
batch_id: str = "",
auto_retry_enabled: bool = False,
auto_retry_max: int = 0,
) -> "GenerationTask":
if not project_id.strip() and not template_id.strip():
raise ValueError("project_id 或 template_id 至少需要提供一个")
@@ -124,6 +130,8 @@ class GenerationTask:
source_edit_plan_id=source_edit_plan_id.strip(),
asset_select_mode=asset_select_mode,
batch_id=batch_id,
auto_retry_enabled=auto_retry_enabled,
auto_retry_max=auto_retry_max,
)
# ── 状态查询 ────────────────────────────────────────────────────────────
@@ -203,13 +211,14 @@ class GenerationTask:
self.result_count = result_count
self.error_message = ""
def mark_failed(self, error_message: str) -> None:
def mark_failed(self, error_message: str, error_info: dict | None = None) -> None:
"""标记为失败(pending / running → failed)。
设置 error_messagecompleted_at
设置 error_messageerror_infocompleted_at
Args:
error_message: 错误信息
error_info: 结构化错误信息error_type, stack_trace, stage, failed_at等
Raises:
ValueError: 当前状态不允许转换到 failed
@@ -217,6 +226,14 @@ class GenerationTask:
self.transition_to(GenerationTaskStatus.FAILED)
self.error_message = error_message
self.completed_at = datetime.now(timezone.utc)
if error_info is not None:
self.error_info = error_info
else:
self.error_info = {
"error_type": "UnknownError",
"message": error_message,
"failed_at": datetime.now(timezone.utc).isoformat(),
}
def mark_cancelled(self) -> None:
"""标记为已取消(pending / running → cancelled)。
@@ -269,7 +286,8 @@ class GenerationTask:
def mark_pending_from_failed(self) -> None:
"""从失败状态重置为待处理(用于重试)。
清除 error_messagestarted_atcompleted_atprogress
清除 error_messageerror_infostarted_atcompleted_atprogress
递增 retry_count
Raises:
ValueError: 当前状态不是 failed
@@ -278,7 +296,9 @@ class GenerationTask:
raise ValueError(f"只有 failed 状态的任务可以重置为 pending,当前状态: {self.status.value}")
self.transition_to(GenerationTaskStatus.PENDING)
self.error_message = ""
self.error_info = {}
self.started_at = None
self.completed_at = None
self.progress = 0.0
self.result_count = 0
self.retry_count += 1
+160
View File
@@ -0,0 +1,160 @@
"""预设 BGM 库 — 免费可商用背景音乐清单.
按风格分类存储在 OSS CDN
实际音频文件由运维统一上传这里只维护元数据清单
"""
from __future__ import annotations
from dataclasses import dataclass, field
@dataclass(frozen=True)
class PresetBGM:
"""预设 BGM 条目"""
id: str
name: str
style: str # 风格分类:upbeat/relax/tech/commerce/emotional/cinematic
duration: float # 时长(秒)
artist: str = ""
description: str = ""
tags: list[str] = field(default_factory=list)
audio_url: str = "" # CDN/OSS 地址,空字符串表示待部署
# ── 预设库清单 ────────────────────────────────────────────────────────────────
PRESET_BGM_LIBRARY: list[PresetBGM] = [
# 轻快 upbeat
PresetBGM(
id="bgm_upbeat_001",
name="阳光清晨",
style="upbeat",
duration=120.0,
artist="免费商用音乐库",
description="轻快明亮的吉他+钢琴,适合vlog、生活记录",
tags=["轻快", "阳光", "吉他", "vlog"],
),
PresetBGM(
id="bgm_upbeat_002",
name="活力节拍",
style="upbeat",
duration=95.0,
artist="免费商用音乐库",
description="电子鼓点+合成器,节奏明快,适合运动、产品展示",
tags=["轻快", "电子", "活力", "运动"],
),
PresetBGM(
id="bgm_upbeat_003",
name="夏日漫步",
style="upbeat",
duration=110.0,
artist="免费商用音乐库",
description="Ukulele+口哨,轻松愉悦,适合旅行、美食",
tags=["轻快", "夏日", "ukulele", "旅行"],
),
# 治愈 relax
PresetBGM(
id="bgm_relax_001",
name="静谧时光",
style="relax",
duration=180.0,
artist="免费商用音乐库",
description="温柔钢琴独奏,治愈系,适合读书、冥想",
tags=["治愈", "钢琴", "安静", "冥想"],
),
PresetBGM(
id="bgm_relax_002",
name="雨后森林",
style="relax",
duration=150.0,
artist="免费商用音乐库",
description="自然白噪音+轻柔吉他,放松减压",
tags=["治愈", "自然", "放松", "环境音"],
),
PresetBGM(
id="bgm_relax_003",
name="月光奏鸣曲",
style="relax",
duration=200.0,
artist="古典音乐(公版)",
description="贝多芬经典钢琴作品,公版免费",
tags=["治愈", "古典", "钢琴", "优雅"],
),
# 科技 tech
PresetBGM(
id="bgm_tech_001",
name="未来科技",
style="tech",
duration=85.0,
artist="免费商用音乐库",
description="电子合成器+科技鼓点,适合数码产品、科技解说",
tags=["科技", "电子", "未来感", "数码"],
),
PresetBGM(
id="bgm_tech_002",
name="数据脉冲",
style="tech",
duration=100.0,
artist="免费商用音乐库",
description="极简电子节奏,适合数据分析、AI类视频",
tags=["科技", "极简", "数据", "AI"],
),
# 电商 commerce
PresetBGM(
id="bgm_commerce_001",
name="心动时刻",
style="commerce",
duration=75.0,
artist="免费商用音乐库",
description="时尚动感节奏,适合商品展示、带货视频",
tags=["电商", "时尚", "动感", "带货"],
),
PresetBGM(
id="bgm_commerce_002",
name="品质生活",
style="commerce",
duration=90.0,
artist="免费商用音乐库",
description="高级感轻音乐,适合品牌宣传、高端产品",
tags=["电商", "高端", "品牌", "品质"],
),
]
# ── 风格分类字典 ──────────────────────────────────────────────────────────────
BGM_STYLES: dict[str, str] = {
"upbeat": "轻快",
"relax": "治愈",
"tech": "科技",
"commerce": "电商",
"emotional": "情感",
"cinematic": "电影",
}
# ── 工具函数 ──────────────────────────────────────────────────────────────────
def get_preset_bgm(bgm_id: str) -> PresetBGM | None:
"""按 ID 获取预设 BGM。"""
for bgm in PRESET_BGM_LIBRARY:
if bgm.id == bgm_id:
return bgm
return None
def list_preset_bgm_by_style(style: str) -> list[PresetBGM]:
"""按风格筛选预设 BGM。"""
return [bgm for bgm in PRESET_BGM_LIBRARY if bgm.style == style]
def search_preset_bgm(keyword: str) -> list[PresetBGM]:
"""按关键词搜索预设 BGM(名称+标签+描述)。"""
kw = keyword.lower()
results = []
for bgm in PRESET_BGM_LIBRARY:
if kw in bgm.name.lower() or kw in bgm.description.lower() or any(kw in tag.lower() for tag in bgm.tags):
results.append(bgm)
return results
+183
View File
@@ -0,0 +1,183 @@
"""字幕领域模型 — 带时间轴的字幕片段。"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import List
@dataclass
class SubtitleWord:
"""单个词级别的字幕单元,带精确时间戳。"""
text: str
start: float # 秒
end: float # 秒
@property
def duration(self) -> float:
return max(0.0, self.end - self.start)
@dataclass
class SubtitleSegment:
"""一段字幕(一句话),带时间轴和词级信息。"""
text: str
start: float # 秒
end: float # 秒
words: List[SubtitleWord] = field(default_factory=list)
@property
def duration(self) -> float:
return max(0.0, self.end - self.start)
@property
def char_count(self) -> int:
return len(self.text)
@dataclass
class SubtitleTimeline:
"""完整的字幕时间轴,由多个片段组成。"""
segments: List[SubtitleSegment] = field(default_factory=list)
language: str = "zh" # zh / en / ja 等
total_duration: float = 0.0 # 音频总时长(秒)
@property
def segment_count(self) -> int:
return len(self.segments)
@property
def total_chars(self) -> int:
return sum(s.char_count for s in self.segments)
def merge_short_segments(self, min_chars: int = 8) -> SubtitleTimeline:
"""合并过短的字幕片段,避免字幕跳动太快。"""
if len(self.segments) <= 1:
return self
merged: List[SubtitleSegment] = []
buffer: List[SubtitleSegment] = []
for seg in self.segments:
buffer.append(seg)
total_chars = sum(s.char_count for s in buffer)
if total_chars >= min_chars:
merged.append(self._merge_segments(buffer))
buffer = []
# 剩余的合并到最后一个或单独成段
if buffer:
if merged and sum(s.char_count for s in buffer) < min_chars:
# 太少了,合并到上一段
last = merged.pop()
merged.append(self._merge_segments([last] + buffer))
else:
merged.append(self._merge_segments(buffer))
return SubtitleTimeline(
segments=merged,
language=self.language,
total_duration=self.total_duration,
)
def split_long_segments(self, max_chars: int = 20) -> SubtitleTimeline:
"""拆分过长的字幕片段,按语义断句。"""
new_segments: List[SubtitleSegment] = []
for seg in self.segments:
if seg.char_count <= max_chars:
new_segments.append(seg)
continue
# 按标点符号拆分
parts = self._split_text_by_punctuation(seg.text, max_chars)
if len(parts) == 1:
new_segments.append(seg)
continue
# 按字数比例分配时间
total_chars = seg.char_count
current_time = seg.start
word_idx = 0
all_words = seg.words.copy()
for part in parts:
part_chars = len(part)
part_duration = seg.duration * (part_chars / total_chars)
part_end = min(current_time + part_duration, seg.end)
# 收集对应时间段的词
part_words = []
while word_idx < len(all_words) and all_words[word_idx].start < part_end:
part_words.append(all_words[word_idx])
word_idx += 1
new_segments.append(
SubtitleSegment(
text=part,
start=current_time,
end=part_end,
words=part_words,
)
)
current_time = part_end
return SubtitleTimeline(
segments=new_segments,
language=self.language,
total_duration=self.total_duration,
)
@staticmethod
def _merge_segments(segments: List[SubtitleSegment]) -> SubtitleSegment:
if not segments:
return SubtitleSegment(text="", start=0, end=0)
return SubtitleSegment(
text="".join(s.text for s in segments),
start=segments[0].start,
end=segments[-1].end,
words=[w for s in segments for w in s.words],
)
@staticmethod
def _split_text_by_punctuation(text: str, max_chars: int) -> List[str]:
"""按标点符号智能拆分长文本。"""
# 中文常见句末标点
sentence_end = "。!?!?"
clause_pause = ",;:,;:"
parts: List[str] = []
current = ""
for char in text:
current += char
if len(current) >= max_chars:
# 超过长度,找最近的标点断开
break_idx = -1
for i in range(len(current) - 1, -1, -1):
if current[i] in sentence_end or current[i] in clause_pause:
break_idx = i + 1
break
if break_idx > 0:
parts.append(current[:break_idx])
current = current[break_idx:]
else:
# 没有标点,硬切
parts.append(current[:max_chars])
current = current[max_chars:]
elif char in sentence_end:
# 句末标点,如果长度够就断开
if len(current) >= max_chars // 2:
parts.append(current)
current = ""
if current:
parts.append(current)
return parts
+102
View File
@@ -0,0 +1,102 @@
"""TTS 配音配置模型."""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any, Optional
@dataclass(slots=True)
class TtsConfig:
"""TTS 配音配置.
Attributes:
enabled: 是否启用配音
voice_id: 音色 ID
speed: 语速 (0.5 ~ 2.0)
pitch: 语调 (-12 ~ 12 半音)
volume: 音量 (0.0 ~ 1.0)
text: 配音文本整段配音时使用
align_mode: 对齐模式 - "subtitle"=按字幕对齐 / "full"=整段配音
overlap_mode: 与原音的叠加模式 - "replace"=替换 / "mix"=混音
"""
enabled: bool = False
voice_id: str = ""
speed: float = 1.0
pitch: float = 0.0
volume: float = 0.8
text: str = ""
align_mode: str = "full" # subtitle / full
overlap_mode: str = "replace" # replace / mix
@classmethod
def parse(cls, data: Optional[dict[str, Any]]) -> "TtsConfig":
"""从 dict 解析配置,无效值回退到默认."""
if not data or not isinstance(data, dict):
return cls()
enabled = data.get("enabled", False)
if not isinstance(enabled, bool):
enabled = False
if not enabled:
return cls(enabled=False)
voice_id = data.get("voice_id", "")
if not isinstance(voice_id, str):
voice_id = ""
speed = data.get("speed", 1.0)
if not isinstance(speed, (int, float)):
speed = 1.0
pitch = data.get("pitch", 0.0)
if not isinstance(pitch, (int, float)):
pitch = 0.0
volume = data.get("volume", 0.8)
if not isinstance(volume, (int, float)):
volume = 0.8
text = data.get("text", "")
if not isinstance(text, str):
text = ""
align_mode = data.get("align_mode", "full")
if align_mode not in ("subtitle", "full"):
align_mode = "full"
overlap_mode = data.get("overlap_mode", "replace")
if overlap_mode not in ("replace", "mix"):
overlap_mode = "replace"
config = cls(
enabled=enabled,
voice_id=voice_id,
speed=float(speed),
pitch=float(pitch),
volume=float(volume),
text=text,
align_mode=align_mode,
overlap_mode=overlap_mode,
)
config._clamp()
return config
def _clamp(self) -> None:
"""边界钳制."""
if self.speed < 0.5:
self.speed = 0.5
elif self.speed > 2.0:
self.speed = 2.0
if self.pitch < -12:
self.pitch = -12
elif self.pitch > 12:
self.pitch = 12
if self.volume < 0.0:
self.volume = 0.0
elif self.volume > 1.0:
self.volume = 1.0
+222
View File
@@ -0,0 +1,222 @@
"""配音引擎音色预设.
CosyVoice preset_voices 区分
- preset_voices.py: CosyVoice 真实音色阿里云
- voice_presets.py: 配音引擎通用音色预设 mock/后续接入的真实 TTS
"""
from __future__ import annotations
import sys
from dataclasses import dataclass
if sys.version_info >= (3, 11):
from enum import StrEnum
else:
from enum import Enum
class StrEnum(str, Enum):
pass
class VoiceGender(StrEnum):
"""音色性别."""
MALE = "male"
FEMALE = "female"
CHILD = "child"
class VoiceStyle(StrEnum):
"""音色风格."""
STABLE = "stable" # 沉稳
LIVELY = "lively" # 活泼
CUSTOMER_SERVICE = "customer_service" # 客服
NARRATION = "narration" # 旁白
NEWS = "news" # 新闻
STORY = "story" # 故事
@dataclass(slots=True)
class VoicePreset:
"""音色预设.
Attributes:
voice_id: 音色唯一标识
name: 音色名称
gender: 性别
style: 风格
description: 描述
provider: 供应商mock/aliyun/xunfei
provider_voice_id: 供应商侧音色 ID
default_speed: 默认语速
default_pitch: 默认语调
sample_rate: 采样率
language: 语言
"""
voice_id: str
name: str
gender: VoiceGender = VoiceGender.FEMALE
style: VoiceStyle = VoiceStyle.NARRATION
description: str = ""
provider: str = "mock"
provider_voice_id: str = ""
default_speed: float = 1.0
default_pitch: float = 0.0
sample_rate: int = 22050
language: str = "zh-CN"
# ─── Mock 音色预设列表 ──────────────────────────────────────
MOCK_VOICES: list[VoicePreset] = [
VoicePreset(
voice_id="female_warm",
name="温暖女声",
gender=VoiceGender.FEMALE,
style=VoiceStyle.NARRATION,
description="温柔温暖的女声,适合情感类、生活类视频",
provider="mock",
provider_voice_id="sine_220",
default_speed=1.0,
default_pitch=0.0,
sample_rate=22050,
language="zh-CN",
),
VoicePreset(
voice_id="male_stable",
name="沉稳男声",
gender=VoiceGender.MALE,
style=VoiceStyle.STABLE,
description="沉稳厚重的男声,适合商务、知识类视频",
provider="mock",
provider_voice_id="sine_110",
default_speed=0.9,
default_pitch=0.0,
sample_rate=22050,
language="zh-CN",
),
VoicePreset(
voice_id="female_lively",
name="活泼女声",
gender=VoiceGender.FEMALE,
style=VoiceStyle.LIVELY,
description="明亮活泼的女声,适合vlog、美食、旅行类视频",
provider="mock",
provider_voice_id="sine_280",
default_speed=1.2,
default_pitch=2.0,
sample_rate=22050,
language="zh-CN",
),
VoicePreset(
voice_id="child_cute",
name="可爱童声",
gender=VoiceGender.CHILD,
style=VoiceStyle.STORY,
description="清脆可爱的童声,适合儿童教育、动画类视频",
provider="mock",
provider_voice_id="sine_380",
default_speed=1.0,
default_pitch=4.0,
sample_rate=22050,
language="zh-CN",
),
VoicePreset(
voice_id="female_service",
name="客服女声",
gender=VoiceGender.FEMALE,
style=VoiceStyle.CUSTOMER_SERVICE,
description="专业清晰的客服女声,适合产品介绍、教程类视频",
provider="mock",
provider_voice_id="sine_250",
default_speed=1.0,
default_pitch=1.0,
sample_rate=22050,
language="zh-CN",
),
VoicePreset(
voice_id="male_news",
name="新闻男声",
gender=VoiceGender.MALE,
style=VoiceStyle.NEWS,
description="字正腔圆的新闻播报声,适合资讯、时政类视频",
provider="mock",
provider_voice_id="sine_140",
default_speed=1.0,
default_pitch=0.0,
sample_rate=22050,
language="zh-CN",
),
VoicePreset(
voice_id="female_soft",
name="轻柔女声",
gender=VoiceGender.FEMALE,
style=VoiceStyle.STORY,
description="轻柔舒缓的女声,适合睡前故事、冥想类视频",
provider="mock",
provider_voice_id="sine_180",
default_speed=0.8,
default_pitch=0.0,
sample_rate=22050,
language="zh-CN",
),
VoicePreset(
voice_id="male_magnetic",
name="磁性男声",
gender=VoiceGender.MALE,
style=VoiceStyle.STORY,
description="低沉磁性的男声,适合电影解说、读书类视频",
provider="mock",
provider_voice_id="sine_90",
default_speed=0.85,
default_pitch=-2.0,
sample_rate=22050,
language="zh-CN",
),
]
# voice_id → VoicePreset
_MOCK_VOICE_MAP: dict[str, VoicePreset] = {v.voice_id: v for v in MOCK_VOICES}
def get_voice(voice_id: str, *, provider: str = "mock") -> VoicePreset | None:
"""根据 voice_id 获取音色预设."""
if provider == "mock":
return _MOCK_VOICE_MAP.get(voice_id)
return None
def list_voices(
*,
gender: str | None = None,
style: str | None = None,
provider: str | None = None,
keyword: str | None = None,
) -> list[VoicePreset]:
"""按条件筛选音色列表."""
# 目前只有 mock 音色
result = list(MOCK_VOICES)
if provider and provider != "mock":
return []
if gender:
result = [v for v in result if v.gender.value == gender]
if style:
result = [v for v in result if v.style.value == style]
if keyword:
kw = keyword.lower()
result = [v for v in result if kw in v.name.lower() or kw in v.description.lower() or kw in v.voice_id.lower()]
return result
def get_default_voice() -> VoicePreset:
"""获取默认音色."""
return MOCK_VOICES[0]
+47
View File
@@ -0,0 +1,47 @@
"""ASR(语音识别)服务接口 — Port 层。"""
from __future__ import annotations
from abc import ABC, abstractmethod
from pathlib import Path
from typing import Optional
from packages.domain.subtitle import SubtitleTimeline
class ASRService(ABC):
"""ASR 服务抽象接口。
不同的 ASR 后端Whisper阿里云腾讯云等实现此接口
上层业务代码只依赖接口不依赖具体实现
"""
@abstractmethod
def transcribe(
self,
audio_path: Path,
language: Optional[str] = None,
with_word_timestamps: bool = True,
) -> SubtitleTimeline:
"""将音频文件转写为带时间轴的字幕。
Args:
audio_path: 音频文件路径支持 wav/mp3/m4a 等常见格式
language: 指定语言代码zh/en/ja None 表示自动检测
with_word_timestamps: 是否返回词级时间戳
Returns:
SubtitleTimeline 字幕时间轴对象
Raises:
ASRServiceError: 识别服务调用失败
"""
...
class ASRServiceError(Exception):
"""ASR 服务调用异常。"""
def __init__(self, message: str, provider: str = "unknown"):
self.provider = provider
super().__init__(f"[{provider}] {message}")
+16 -1
View File
@@ -52,7 +52,22 @@ class AssetRepository(ABC):
@abstractmethod
def batch_delete(self, asset_ids: list[str]) -> int:
"""批量删除素材,返回实际删除数量。"""
"""批量删除素材(软删除,标记 status=deleted,返回实际影响数量。"""
pass
@abstractmethod
def batch_update_metadata(self, asset_ids: list[str], metadata_patch: dict[str, object]) -> int:
"""批量更新素材 metadata(合并 patch),返回实际影响数量。"""
pass
@abstractmethod
def batch_add_tags(self, asset_ids: list[str], tag_ids: list[str]) -> int:
"""批量给素材添加标签(合并去重),返回实际影响数量。"""
pass
@abstractmethod
def batch_replace_tags(self, asset_ids: list[str], tag_ids: list[str]) -> int:
"""批量替换素材标签(全量覆盖),返回实际影响数量。"""
pass
@abstractmethod
+16
View File
@@ -15,3 +15,19 @@ class GeneratedVideoRepository(Protocol):
def list_by_generation_task(self, generation_task_id: str) -> list[GeneratedVideo]: ...
def list_by_batch(self, batch_id: str) -> list[GeneratedVideo]: ...
def list_paginated(
self,
*,
project_id: str | None = None,
status: str | None = None,
review_status: str | None = None,
page: int = 1,
page_size: int = 20,
) -> tuple[list[GeneratedVideo], int]: ...
def update_review_status(self, video_id: str, review_status: str) -> GeneratedVideo | None: ...
def update_thumbnail(self, video_id: str, thumbnail_url: str) -> bool: ...
def get_by_ids(self, video_ids: list[str]) -> list[GeneratedVideo]: ...
@@ -24,4 +24,36 @@ class GenerationTaskRepository(Protocol):
def list_by_source_edit_plan(self, plan_id: str) -> list[GenerationTask]: ...
def list_by_user_filtered(
self,
user_id: str,
*,
status: str | None = None,
limit: int | None = None,
offset: int = 0,
) -> list[GenerationTask]: ...
def count_by_user_filtered(
self,
user_id: str,
*,
status: str | None = None,
) -> int: ...
def list_by_project_filtered(
self,
project_id: str,
*,
status: str | None = None,
limit: int | None = None,
offset: int = 0,
) -> list[GenerationTask]: ...
def count_by_project_filtered(
self,
project_id: str,
*,
status: str | None = None,
) -> int: ...
def update(self, task: GenerationTask) -> GenerationTask: ...
+23 -2
View File
@@ -8,12 +8,31 @@ from packages.domain.template import Template, TemplateCategory, TemplateSegment
class TemplateRepositoryPort(Protocol):
def list_by_user(self, user_id: str, *, skip: int = 0, limit: int = 50) -> List[Template]: ...
def list_by_user(
self,
user_id: str,
*,
skip: int = 0,
limit: int = 50,
category: Optional[str] = None,
tag: Optional[str] = None,
keyword: Optional[str] = None,
mode: Optional[str] = None,
) -> List[Template]: ...
def get(self, template_id: str, user_id: str) -> Optional[Template]: ...
def create(self, template: Template) -> Template: ...
def update(self, template: Template) -> Template: ...
def delete(self, template_id: str, user_id: str) -> bool: ...
def count_by_user(self, user_id: str) -> int: ...
def count_by_user(
self,
user_id: str,
*,
category: Optional[str] = None,
tag: Optional[str] = None,
keyword: Optional[str] = None,
mode: Optional[str] = None,
) -> int: ...
def copy_template(self, template_id: str, user_id: str, new_name: str) -> Template: ...
def list_segments(self, template_id: str) -> List[TemplateSegment]: ...
def create_segments(self, segments: List[TemplateSegment]) -> List[TemplateSegment]: ...
def delete_segments_by_template(self, template_id: str) -> int: ...
@@ -21,3 +40,5 @@ class TemplateRepositoryPort(Protocol):
def create_category(self, category: TemplateCategory) -> TemplateCategory: ...
def get_category(self, category_id: str, user_id: str) -> Optional[TemplateCategory]: ...
def delete_category(self, category_id: str, user_id: str) -> bool: ...
def list_tags(self, user_id: str) -> List[str]: ...
def get_usage_count(self, template_id: str) -> int: ...
+78
View File
@@ -0,0 +1,78 @@
"""TTS 服务抽象接口 (Port).
新增 TTS 供应商时实现本接口即可
"""
from __future__ import annotations
from abc import ABC, abstractmethod
from pathlib import Path
class TtsService(ABC):
"""TTS 服务抽象基类.
所有 TTS 供应商Mock / 阿里云 / 讯飞 都需要实现本接口
"""
@abstractmethod
def synthesize(
self,
text: str,
*,
voice_id: str = "",
speed: float = 1.0,
pitch: float = 0.0,
output_path: Path | None = None,
sample_rate: int = 22050,
format: str = "wav",
) -> Path:
"""文本转语音合成.
Args:
text: 输入文本
voice_id: 音色 ID
speed: 语速 (0.5 ~ 2.0)
pitch: 语调半音-12 ~ 12
output_path: 输出文件路径None 则自动生成
sample_rate: 采样率
format: 输出格式 (wav/mp3)
Returns:
输出音频文件路径
Raises:
TtsError: 合成失败
"""
...
@abstractmethod
def estimate_duration(self, text: str, *, speed: float = 1.0) -> float:
"""预估音频时长(秒).
用于在实际合成前估算时长方便时间轴对齐
Args:
text: 输入文本
speed: 语速
Returns:
预估时长
"""
...
@property
@abstractmethod
def provider_name(self) -> str:
"""供应商名称."""
...
def available_voices(self) -> list[str]:
"""支持的音色 ID 列表."""
return []
class TtsError(Exception):
"""TTS 合成异常."""
pass
+4 -7
View File
@@ -63,8 +63,7 @@ BRANCH_NAME="${GITHUB_REF_NAME:-${CI_COMMIT_BRANCH:-unknown}}"
if [ "$USE_CACHE" -eq 1 ]; then
docker buildx build \
--build-arg APP_VERSION="$VERSION" \
--cache-from "type=registry,ref=${CACHE_REGISTRY}/api-cache:${CACHE_TAG},ignore-error=true" \
--cache-to "type=registry,ref=${CACHE_REGISTRY}/api-cache:${CACHE_TAG},mode=max" \
--cache-from "type=registry,ref=${CACHE_REGISTRY}/api-cache:${CACHE_TAG_PRIMARY},ignore-error=true" \
-f infra/docker/api.Dockerfile \
-t "$API_IMAGE" -t "$API_LATEST" \
--load \
@@ -123,8 +122,7 @@ echo "=== Building Worker image ==="
if [ "$USE_CACHE" -eq 1 ]; then
docker buildx build \
--build-arg APP_VERSION="$VERSION" \
--cache-from "type=registry,ref=${CACHE_REGISTRY}/worker-cache:${CACHE_TAG},ignore-error=true" \
--cache-to "type=registry,ref=${CACHE_REGISTRY}/worker-cache:${CACHE_TAG},mode=max" \
--cache-from "type=registry,ref=${CACHE_REGISTRY}/worker-cache:${CACHE_TAG_PRIMARY},ignore-error=true" \
-f infra/docker/worker.Dockerfile \
-t "$WORKER_IMAGE" -t "$WORKER_LATEST" \
--load \
@@ -152,8 +150,7 @@ test -f apps/web/dist/index.html
if [ "$USE_CACHE" -eq 1 ]; then
docker buildx build \
--cache-from "type=registry,ref=${CACHE_REGISTRY}/web-cache:${CACHE_TAG},ignore-error=true" \
--cache-to "type=registry,ref=${CACHE_REGISTRY}/web-cache:${CACHE_TAG},mode=max" \
--cache-from "type=registry,ref=${CACHE_REGISTRY}/web-cache:${CACHE_TAG_PRIMARY},ignore-error=true" \
-f infra/docker/web-artifact.Dockerfile \
--build-arg "NGINX_CONF=$NGINX_CONF_FILE" \
-t "$WEB_IMAGE" \
@@ -182,4 +179,4 @@ else
fi
echo "=== Build complete ==="
docker images | grep "xiaoxia-saas.*:$VERSION"
docker images | grep "xiaoxia-saas" | grep "$VERSION" || true
+117
View File
@@ -0,0 +1,117 @@
#!/bin/bash
# 灰度发布脚本:通过Nginx权重调整流量比例
# 用法: ./scripts/gray_deploy.sh <版本号> <灰度百分比>
#
# 需要在目标服务器上执行,或通过SSH执行
# 前提:服务器上运行两个版本的容器(stable + canary),Nginx做加权轮询
set -euo pipefail
VERSION="${1:-}"
GRAY_PCT="${2:-10}"
if [[ -z "$VERSION" ]]; then
echo "用法: $0 <版本号> [灰度百分比]"
echo "示例: $0 v0.1.129 5"
exit 1
fi
STABLE_VERSION="${STABLE_VERSION:-current}"
REGISTRY="${REGISTRY:-git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas}"
echo "=========================================="
echo " 灰度发布"
echo " 新版本: $VERSION"
echo " 灰度比例: ${GRAY_PCT}%"
echo " 稳定版本: $STABLE_VERSION"
echo "=========================================="
# 1. 拉取新版本镜像
echo ""
echo ">>> 拉取新版本镜像..."
for component in api worker web; do
echo " 拉取 $component:$VERSION ..."
docker pull "${REGISTRY}-${component}:${VERSION}" 2>&1 | tail -1
done
# 2. 启动灰度版本容器(canary)
echo ""
echo ">>> 启动灰度版本容器..."
# API canary
CANARY_API_NAME="saas-api-canary"
if docker ps -a --format '{{.Names}}' | grep -q "^${CANARY_API_NAME}$"; then
echo " 停止旧 canary 容器..."
docker stop "$CANARY_API_NAME" 2>/dev/null || true
docker rm "$CANARY_API_NAME" 2>/dev/null || true
fi
echo " 启动 api canary..."
docker run -d \
--name "$CANARY_API_NAME" \
--network saas-network \
-e DATABASE_URL="${DATABASE_URL}" \
-e REDIS_URL="${REDIS_URL}" \
-e FEATURE_FLAG_PROVIDER=redis \
--restart unless-stopped \
"${REGISTRY}-api:${VERSION}"
# Worker canary
CANARY_WORKER_NAME="saas-worker-canary"
if docker ps -a --format '{{.Names}}' | grep -q "^${CANARY_WORKER_NAME}$"; then
echo " 停止旧 worker canary..."
docker stop "$CANARY_WORKER_NAME" 2>/dev/null || true
docker rm "$CANARY_WORKER_NAME" 2>/dev/null || true
fi
echo " 启动 worker canary..."
docker run -d \
--name "$CANARY_WORKER_NAME" \
--network saas-network \
-e DATABASE_URL="${DATABASE_URL}" \
-e REDIS_URL="${REDIS_URL}" \
--restart unless-stopped \
"${REGISTRY}-worker:${VERSION}"
# 3. 等待容器健康
echo ""
echo ">>> 等待容器健康..."
sleep 10
if ! docker ps --format '{{.Names}} {{.Status}}' | grep -q "$CANARY_API_NAME"; then
echo "错误: API canary 容器未运行"
docker logs "$CANARY_API_NAME" --tail 20
exit 1
fi
echo " ✅ API canary 运行中"
# 4. 更新Nginx权重
echo ""
echo ">>> 更新Nginx权重 (稳定: $((100-GRAY_PCT))% / 灰度: ${GRAY_PCT}%)..."
NGINX_CONF="${NGINX_CONF:-/etc/nginx/conf.d/saas-api.conf}"
if [[ -f "$NGINX_CONF" ]]; then
# 备份
cp "$NGINX_CONF" "${NGINX_CONF}.bak.$(date +%Y%m%d%H%M%S)"
# 更新 upstream 权重(需要根据实际配置调整)
echo " 请手动更新 Nginx upstream 配置中的权重"
echo " 示例配置:"
cat <<EOF
upstream saas_api_backend {
server saas-api:8000 weight=$((100-GRAY_PCT));
server saas-api-canary:8000 weight=${GRAY_PCT};
}
EOF
nginx -t && nginx -s reload
echo " ✅ Nginx 已reload"
else
echo " 警告: Nginx 配置文件不存在 ($NGINX_CONF)"
echo " 请手动配置灰度流量权重"
fi
echo ""
echo "=========================================="
echo " ✅ 灰度发布完成"
echo " 新版本: $VERSION (${GRAY_PCT}%流量)"
echo " 监控: Grafana / 日志"
echo "=========================================="
+126
View File
@@ -0,0 +1,126 @@
#!/bin/bash
# 一键发布脚本:打tag → 触发生产镜像构建 → 部署到灰度
# 用法: ./scripts/release.sh v0.1.129 [--gray 5]
set -euo pipefail
SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
REPO_ROOT="$(cd "$SCRIPT_DIR/.." && pwd)"
usage() {
echo "用法: $0 <版本号> [--gray 百分比] [--no-deploy]"
echo ""
echo "示例:"
echo " $0 v0.1.129 # 打tag + 全量发布"
echo " $0 v0.1.129 --gray 5 # 打tag + 5%灰度发布"
echo " $0 v0.1.129 --no-deploy # 只打tag,不部署"
exit 1
}
# 参数解析
VERSION=""
GRAY_PCT=0
DEPLOY=true
while [[ $# -gt 0 ]]; do
case "$1" in
--gray)
GRAY_PCT="$2"
shift 2
;;
--no-deploy)
DEPLOY=false
shift
;;
-h|--help)
usage
;;
v*)
VERSION="$1"
shift
;;
*)
echo "未知参数: $1"
usage
;;
esac
done
if [[ -z "$VERSION" ]]; then
echo "错误: 请指定版本号(如 v0.1.129"
usage
fi
echo "=========================================="
echo " 发布版本: $VERSION"
echo " 灰度比例: ${GRAY_PCT}%"
echo " 自动部署: $DEPLOY"
echo "=========================================="
cd "$REPO_ROOT"
# 1. 确认在 develop 分支
CURRENT_BRANCH=$(git rev-parse --abbrev-ref HEAD)
if [[ "$CURRENT_BRANCH" != "develop" ]]; then
echo "错误: 请切换到 develop 分支后再发布"
exit 1
fi
# 2. 拉取最新代码
echo ""
echo ">>> 拉取最新代码..."
git pull origin develop
# 3. 生成 CHANGELOG
echo ""
echo ">>> 生成 CHANGELOG..."
if [[ -f "scripts/generate_changelog.py" ]]; then
PREV_TAG=$(git describe --tags --abbrev=0 HEAD^ 2>/dev/null || echo "")
if [[ -n "$PREV_TAG" ]]; then
python3 scripts/generate_changelog.py \
--from-tag "$PREV_TAG" \
--to-tag HEAD \
--gitea-token "${GITEA_TOKEN:-}" \
--output /tmp/changelog_$$.md
echo "CHANGELOG 已生成到 /tmp/changelog_$$.md"
fi
fi
# 4. 打tag
echo ""
echo ">>> 打 tag $VERSION ..."
if git rev-parse "$VERSION" >/dev/null 2>&1; then
echo "警告: tag $VERSION 已存在,跳过打tag"
else
git tag -a "$VERSION" -m "Release $VERSION"
git push origin "$VERSION"
echo "Tag $VERSION 已推送,触发生产镜像构建..."
fi
# 5. 等待镜像构建
if [[ "$DEPLOY" == "true" ]]; then
echo ""
echo ">>> 等待镜像构建完成(约10-15分钟)..."
echo " 镜像: git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas-{api,worker,web}:$VERSION"
# 这里可以加镜像存在性检查
echo " (镜像构建由CI自动完成,请在Gitea Actions中确认)"
fi
# 6. 灰度部署
if [[ "$DEPLOY" == "true" && "$GRAY_PCT" -gt 0 ]]; then
echo ""
echo ">>> 灰度部署: ${GRAY_PCT}% 流量到 $VERSION"
if [[ -f "scripts/gray_deploy.sh" ]]; then
./scripts/gray_deploy.sh "$VERSION" "$GRAY_PCT"
else
echo "警告: gray_deploy.sh 不存在,跳过灰度部署"
fi
fi
echo ""
echo "=========================================="
echo " ✅ 发布流程完成"
echo " 版本: $VERSION"
echo " 灰度: ${GRAY_PCT}%"
echo "=========================================="
+42
View File
@@ -0,0 +1,42 @@
#!/bin/bash
# 灰度回滚脚本:切回稳定版本流量
# 用法: ./scripts/rollback.sh [稳定版本号]
set -euo pipefail
STABLE_VERSION="${1:-current}"
echo "=========================================="
echo " 灰度回滚"
echo " 切回稳定版本: $STABLE_VERSION"
echo "=========================================="
# 1. 恢复Nginx全量到稳定版本
echo ""
echo ">>> 恢复Nginx全量流量到稳定版本..."
NGINX_CONF="${NGINX_CONF:-/etc/nginx/conf.d/saas-api.conf}"
if [[ -f "$NGINX_CONF" ]]; then
# 找最近的备份
LATEST_BAK=$(ls -t "${NGINX_CONF}".bak.* 2>/dev/null | head -1)
if [[ -n "$LATEST_BAK" ]]; then
cp "$LATEST_BAK" "$NGINX_CONF"
echo " 从备份恢复: $LATEST_BAK"
else
echo " 未找到备份,请手动移除 canary upstream"
fi
nginx -t && nginx -s reload
echo " ✅ Nginx 已回滚"
fi
# 2. 停止灰度版本容器(保留30分钟以便排查)
echo ""
echo ">>> 灰度版本容器将在30分钟后停止(便于排查问题)"
echo " 立即停止请执行: docker stop saas-api-canary saas-worker-canary"
echo ""
echo "=========================================="
echo " ✅ 回滚完成"
echo " 流量已全部切回稳定版本"
echo "=========================================="
+1
View File
@@ -15,3 +15,4 @@ exclude =
per-file-ignores =
*/__init__.py:F401,F403,F405
tests/*:E402,F401,F841
packages/ports/*:E301,E704
+63 -8
View File
@@ -129,10 +129,58 @@ class StubAssetRepository:
return False
def batch_delete(self, asset_ids: list[str]) -> int:
"""软删除:标记 status=deleted。"""
from datetime import datetime, timezone
from packages.domain import AssetStatus
count = 0
for aid in asset_ids:
if aid in self._assets:
del self._assets[aid]
asset = self._assets.get(aid)
if asset and asset.status != AssetStatus.DELETED:
asset.status = AssetStatus.DELETED
asset.updated_at = datetime.now(timezone.utc)
count += 1
return count
def batch_update_metadata(self, asset_ids: list[str], metadata_patch: dict) -> int:
from datetime import datetime, timezone
count = 0
for aid in asset_ids:
asset = self._assets.get(aid)
if asset:
asset.metadata = {**asset.metadata, **metadata_patch}
asset.updated_at = datetime.now(timezone.utc)
count += 1
return count
def batch_add_tags(self, asset_ids: list[str], tag_ids: list[str]) -> int:
from datetime import datetime, timezone
count = 0
for aid in asset_ids:
asset = self._assets.get(aid)
if asset:
changed = False
for tid in tag_ids:
if tid not in asset.tag_ids:
asset.tag_ids.append(tid)
changed = True
if changed:
asset.updated_at = datetime.now(timezone.utc)
count += 1
return count
def batch_replace_tags(self, asset_ids: list[str], tag_ids: list[str]) -> int:
from datetime import datetime, timezone
count = 0
for aid in asset_ids:
asset = self._assets.get(aid)
if asset:
asset.tag_ids = list(tag_ids)
asset.updated_at = datetime.now(timezone.utc)
count += 1
return count
@@ -613,17 +661,22 @@ class TestBatchDeleteAssets:
return ids
def test_batch_delete_success(self, client):
"""批量删除成功。"""
"""批量删除成功(软删除)"""
ids = self._create_assets(client, 3)
resp = client.post(
"/api/v1/assets/batch-delete",
json={"ids": ids[:2]},
json={"asset_ids": ids[:2]},
)
assert resp.status_code == 200
data = resp.json()
assert data["deleted_count"] == 2
assert data["success_count"] == 2
assert len(data["failed_ids"]) == 0
# 软删除:记录仍在,status 变为 deleted
for aid in ids[:2]:
r = client.get(f"/api/v1/assets/{aid}")
assert r.status_code == 200
assert r.json()["status"] == "deleted"
def test_batch_delete_with_nonexistent_ids(self, client):
"""批量删除包含不存在的 ID,失败的计入 failed_ids。"""
@@ -632,18 +685,20 @@ class TestBatchDeleteAssets:
resp = client.post(
"/api/v1/assets/batch-delete",
json={"ids": ids},
json={"asset_ids": ids},
)
assert resp.status_code == 200
data = resp.json()
assert data["deleted_count"] == 2
assert data["success_count"] == 2
assert "nonexistent-id" in data["failed_ids"]
assert "nonexistent-id" in data["failed_details"]
assert data["failed_details"]["nonexistent-id"] == "not_found"
def test_batch_delete_empty_list_returns_422(self, client):
"""空列表返回 422。"""
resp = client.post(
"/api/v1/assets/batch-delete",
json={"ids": []},
json={"asset_ids": []},
)
assert resp.status_code == 422
+64
View File
@@ -138,6 +138,70 @@ class StubGenerationTaskRepository:
def list_by_source_edit_plan(self, plan_id: str) -> list[GenerationTask]:
return [t for t in self._tasks.values() if getattr(t, "source_edit_plan_id", "") == plan_id]
def list_by_user_filtered(
self,
user_id: str,
*,
status: str | None = None,
limit: int | None = None,
offset: int = 0,
) -> list:
"""按用户+状态筛选任务列表(stub实现)。"""
items = [t for t in self._tasks.values() if t.created_by_user_id == user_id]
if status:
items = [t for t in items if str(t.status) == status]
# 按创建时间倒序
items.sort(key=lambda t: t.created_at or "", reverse=True)
if offset:
items = items[offset:]
if limit is not None:
items = items[:limit]
return items
def count_by_user_filtered(
self,
user_id: str,
*,
status: str | None = None,
) -> int:
"""按用户+状态筛选计数(stub实现)。"""
items = [t for t in self._tasks.values() if t.created_by_user_id == user_id]
if status:
items = [t for t in items if str(t.status) == status]
return len(items)
def list_by_project_filtered(
self,
project_id: str,
*,
status: str | None = None,
limit: int | None = None,
offset: int = 0,
) -> list:
"""按项目+状态筛选任务列表(stub实现)。"""
items = [t for t in self._tasks.values() if t.project_id == project_id]
if status:
items = [t for t in items if str(t.status) == status]
# 按创建时间倒序
items.sort(key=lambda t: t.created_at or "", reverse=True)
if offset:
items = items[offset:]
if limit is not None:
items = items[:limit]
return items
def count_by_project_filtered(
self,
project_id: str,
*,
status: str | None = None,
) -> int:
"""按项目+状态筛选计数(stub实现)。"""
items = [t for t in self._tasks.values() if t.project_id == project_id]
if status:
items = [t for t in items if str(t.status) == status]
return len(items)
class StubGeneratedVideoRepository:
def __init__(self, videos: dict[str, GeneratedVideo] | None = None):
+70 -4
View File
@@ -103,6 +103,70 @@ class StubGenerationTaskRepository:
def list_by_source_edit_plan(self, plan_id: str) -> list[GenerationTask]:
return [t for t in self._tasks.values() if getattr(t, "source_edit_plan_id", "") == plan_id]
def list_by_user_filtered(
self,
user_id: str,
*,
status: str | None = None,
limit: int | None = None,
offset: int = 0,
) -> list:
"""按用户+状态筛选任务列表(stub实现)。"""
items = [t for t in self._tasks.values() if t.created_by_user_id == user_id]
if status:
items = [t for t in items if str(t.status) == status]
# 按创建时间倒序
items.sort(key=lambda t: t.created_at or "", reverse=True)
if offset:
items = items[offset:]
if limit is not None:
items = items[:limit]
return items
def count_by_user_filtered(
self,
user_id: str,
*,
status: str | None = None,
) -> int:
"""按用户+状态筛选计数(stub实现)。"""
items = [t for t in self._tasks.values() if t.created_by_user_id == user_id]
if status:
items = [t for t in items if str(t.status) == status]
return len(items)
def list_by_project_filtered(
self,
project_id: str,
*,
status: str | None = None,
limit: int | None = None,
offset: int = 0,
) -> list:
"""按项目+状态筛选任务列表(stub实现)。"""
items = [t for t in self._tasks.values() if t.project_id == project_id]
if status:
items = [t for t in items if str(t.status) == status]
# 按创建时间倒序
items.sort(key=lambda t: t.created_at or "", reverse=True)
if offset:
items = items[offset:]
if limit is not None:
items = items[:limit]
return items
def count_by_project_filtered(
self,
project_id: str,
*,
status: str | None = None,
) -> int:
"""按项目+状态筛选计数(stub实现)。"""
items = [t for t in self._tasks.values() if t.project_id == project_id]
if status:
items = [t for t in items if str(t.status) == status]
return len(items)
class StubIngestJobRepository:
def __init__(self, jobs: dict[str, IngestJob] | None = None):
@@ -470,8 +534,8 @@ class TestRetryProjectTask:
assert data["task_type"] == "generation"
assert data["status"] == "pending"
assert "current_step" in data
# 验证新任务的 ID 不同于原任务
assert data["source_id"] != "gen-failed-1"
# 原地重试:source_id 保持不变(复用同一个任务
assert data["source_id"] == "gen-failed-1"
# 验证 Celery 任务被发送
assert mock_celery.send_task.called
@@ -630,11 +694,13 @@ class TestTaskCenterCrossEndpoint:
retry_resp = tc.post("/tasks/gen-fail-cross/retry")
assert retry_resp.status_code == 200
# 3. 再次列出,应有2个任务(旧的failed + 新的pending
# 3. 再次列出:原地重试,任务数不变(仍是1个),但状态变为 pending
list_resp2 = tc.get("/tasks")
assert list_resp2.status_code == 200
items2 = list_resp2.json()["items"]
assert len(items2) == 2
assert len(items2) == 1
assert items2[0]["status"] == "pending"
assert items2[0]["source_id"] == "gen-fail-cross"
test_app.dependency_overrides.clear()
+201
View File
@@ -0,0 +1,201 @@
# 测试素材创建工具
`create_test_asset.py` 是一个自动化测试辅助工具,用于快速创建 `ready` 状态的视频素材,跳过正常的上传和转码流程,直接指定已存在于 OSS 的文件来生成可用素材。
## 适用场景
- 渲染对比测试:快速创建测试素材用于生成任务
- 性能测试:批量创建素材模拟真实场景
- 开发调试:无需真实上传文件即可测试素材相关功能
## 前置条件
1. API 服务正在运行
2. 有有效的登录 token
3. 指定的 `storage_key` 对应的文件已存在于 OSS 中
4. 用户对指定项目有访问权限
## 使用方法
### 基本用法
```bash
# 设置环境变量(可选)
export API_BASE_URL=http://localhost:8000
export API_TOKEN=your_token_here
# 创建测试素材
python tests/render_compare/create_test_asset.py \
--project-id proj_xxx \
--name "测试素材-30s" \
--storage-key "assets/test/sample_30s.mp4" \
--duration 30 \
--file-size 10485760
```
### 完整参数示例
```bash
python tests/render_compare/create_test_asset.py \
--base-url http://localhost:8000 \
--token eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9... \
--project-id proj_a1b2c3d4e5f6 \
--name "1080p测试视频-60s" \
--storage-key "test-assets/1080p_60fps_60s.mp4" \
--duration 60 \
--width 1920 \
--height 1080 \
--fps 60 \
--mime-type "video/mp4" \
--file-size 52428800 \
--codec "h264" \
--kind video
```
### 在脚本中捕获 asset_id
```bash
# 最后一行输出为 asset_id,方便脚本捕获
ASSET_ID=$(python tests/render_compare/create_test_asset.py \
--project-id proj_xxx \
--name "测试素材" \
--storage-key "test/video.mp4" \
2>&1 | tail -1)
echo "创建的素材 ID: $ASSET_ID"
```
## 参数说明
| 参数 | 环境变量 | 必填 | 默认值 | 说明 |
|------|----------|------|--------|------|
| `--base-url` | `API_BASE_URL` | 否 | `http://localhost:8000` | API 服务地址 |
| `--token` | `API_TOKEN` | 是 | - | 登录认证 token |
| `--project-id` | - | 是 | - | 项目 ID |
| `--name` | - | 是 | - | 素材名称 |
| `--storage-key` | - | 是 | - | OSS storage_key(文件需已存在) |
| `--duration` | - | 否 | `30` | 视频时长(秒) |
| `--width` | - | 否 | `1280` | 视频宽度 |
| `--height` | - | 否 | `720` | 视频高度 |
| `--fps` | - | 否 | `25` | 帧率 |
| `--mime-type` | - | 否 | `video/mp4` | MIME 类型 |
| `--file-size` | - | 否 | `0` | 文件大小(字节) |
| `--codec` | - | 否 | - | 视频编码 |
| `--kind` | - | 否 | `video` | 素材库类型 (video/voice/image) |
## API 调用流程
脚本会依次调用以下 API
### 1. 确保默认素材库存在
```
POST /api/v1/asset-libraries/ensure-default
Content-Type: application/json
Authorization: Bearer {token}
{
"project_id": "proj_xxx",
"kind": "video"
}
```
- 如果项目下已有对应类型的素材库,直接返回第一个
- 如果不存在,自动创建默认名称的素材库
### 2. 创建素材
```
POST /api/v1/assets
Content-Type: application/json
Authorization: Bearer {token}
{
"project_id": "proj_xxx",
"library_id": "lib_xxx",
"name": "测试素材",
"storage_key": "assets/test/video.mp4",
"mime_type": "video/mp4",
"file_size": 10485760,
"duration": 30,
"width": 1280,
"height": 720,
"fps": 25,
"status": "ready",
"classification_status": "pending",
"metadata": {}
}
```
**关键点**
- `status: "ready"` 直接跳过转码流程,立即可用
- `storage_key` 必须对应 OSS 中真实存在的文件,否则播放会失败
- `uploaded_by_user_id` 由 API 自动设置为当前登录用户
## 输出说明
脚本输出分为三部分:
1. **参数回显**:确认输入的参数是否正确
2. **执行日志**:显示每一步的执行情况
3. **结果输出**
- 素材详细信息(asset_id、状态、文件 URL 等)
- 调用示例
- 最后一行为纯 asset_id,方便脚本捕获
## 错误排查
### 常见错误
#### 401 Unauthorized
- 检查 token 是否正确且未过期
- 确认 Authorization header 格式为 `Bearer {token}`
#### 403 Forbidden
- 确认用户对该项目有访问权限
- 检查项目 ID 是否正确
#### 404 Project not found
- 项目 ID 错误或项目不存在
#### 422 Validation Error
- 检查参数格式是否正确
- 查看响应中的 detail 字段了解具体错误
### 调试技巧
脚本在 API 失败时会打印完整的响应内容,包括:
- HTTP 状态码
- 错误原因
- 响应 body 详情
如果遇到问题,请检查:
1. API 服务是否正常运行
2. base-url 是否正确(注意端口号)
3. token 是否有效
4. storage_key 对应的文件是否存在于 OSS
## 与渲染对比测试配合使用
```bash
# 1. 创建测试素材
ASSET_ID=$(python tests/render_compare/create_test_asset.py \
--project-id proj_xxx \
--name "渲染对比测试素材" \
--storage-key "test-assets/base_1080p_30s.mp4" \
--duration 30 --width 1920 --height 1080 \
2>&1 | tail -1)
# 2. 使用该素材运行渲染对比测试
python tests/render_compare/runner.py \
--project-id proj_xxx \
--asset-id $ASSET_ID \
--scenario quality_test
```
## 注意事项
1. **storage_key 必须真实存在**:脚本不会上传文件,只是创建数据库记录。如果 OSS 中没有对应文件,素材虽然状态是 ready,但无法正常播放。
2. **素材参数要准确**duration、width、height、fps 等参数应与实际文件一致,否则可能导致后续渲染或分析出现偏差。
3. **权限检查**:确保 token 对应用户有项目的素材创建权限。
4. **清理测试数据**:测试完成后记得清理不需要的测试素材,避免占用资源。
+255
View File
@@ -0,0 +1,255 @@
#!/usr/bin/env python3
"""
自动化测试用一键创建 ready 状态的视频素材
跳过转码和上传流程直接指定 storage_key 创建可用素材
供渲染对比测试等场景快速生成测试素材
用法示例
python create_test_asset.py \
--base-url http://localhost:8000 \
--token xxx \
--project-id proj_xxx \
--name "测试素材-30s" \
--storage-key "assets/test/video.mp4" \
--duration 30 \
--file-size 10485760
"""
import argparse
import json
import os
import sys
import urllib.error
import urllib.request
def parse_args():
parser = argparse.ArgumentParser(description="创建 ready 状态的测试素材(跳过转码上传)")
parser.add_argument(
"--base-url",
default=os.environ.get("API_BASE_URL", "http://localhost:8000"),
help="API 地址,默认 http://localhost:8000(或环境变量 API_BASE_URL",
)
parser.add_argument(
"--token",
default=os.environ.get("API_TOKEN", ""),
help="登录 token(或环境变量 API_TOKEN",
)
parser.add_argument(
"--project-id",
required=True,
help="项目 ID",
)
parser.add_argument(
"--name",
required=True,
help="素材名称",
)
parser.add_argument(
"--storage-key",
required=True,
help="OSS storage_key(文件必须已存在于 OSS)",
)
parser.add_argument(
"--duration",
type=float,
default=30.0,
help="视频时长(秒),默认 30",
)
parser.add_argument(
"--width",
type=int,
default=1280,
help="视频宽度,默认 1280",
)
parser.add_argument(
"--height",
type=int,
default=720,
help="视频高度,默认 720",
)
parser.add_argument(
"--fps",
type=float,
default=25.0,
help="帧率,默认 25",
)
parser.add_argument(
"--mime-type",
default="video/mp4",
help="MIME 类型,默认 video/mp4",
)
parser.add_argument(
"--file-size",
type=int,
default=0,
help="文件大小(字节),默认 0",
)
parser.add_argument(
"--codec",
default=None,
help="视频编码,可选",
)
parser.add_argument(
"--kind",
default="video",
choices=["video", "voice", "image"],
help="素材库类型,默认 video",
)
return parser.parse_args()
def api_request(base_url: str, token: str, method: str, path: str, body: dict | None = None) -> dict:
"""发送 API 请求,返回 JSON 响应。"""
url = f"{base_url.rstrip('/')}{path}"
data = json.dumps(body).encode("utf-8") if body else None
headers = {
"Content-Type": "application/json",
}
if token:
headers["Authorization"] = f"Bearer {token}"
req = urllib.request.Request(url, data=data, method=method, headers=headers)
try:
with urllib.request.urlopen(req) as resp:
resp_body = resp.read().decode("utf-8")
return json.loads(resp_body) if resp_body else {}
except urllib.error.HTTPError as e:
error_body = e.read().decode("utf-8", errors="replace")
print(f"[ERROR] API 请求失败: {method} {url}", file=sys.stderr)
print(f" HTTP {e.code}: {e.reason}", file=sys.stderr)
print(f" 响应内容: {error_body}", file=sys.stderr)
sys.exit(1)
except urllib.error.URLError as e:
print(f"[ERROR] 网络错误: {method} {url}", file=sys.stderr)
print(f" 原因: {e.reason}", file=sys.stderr)
sys.exit(1)
def ensure_default_library(base_url: str, token: str, project_id: str, kind: str) -> str:
"""确保项目有默认素材库,返回 library_id。"""
print(f"[1/2] 确保默认 {kind} 素材库存在...")
result = api_request(
base_url,
token,
"POST",
"/api/v1/asset-libraries/ensure-default",
body={"project_id": project_id, "kind": kind},
)
library_id = result.get("id")
library_name = result.get("name")
print(f" 素材库: {library_name} (id: {library_id})")
return library_id
def create_asset(
base_url: str,
token: str,
project_id: str,
library_id: str,
name: str,
storage_key: str,
mime_type: str,
file_size: int,
duration: float,
width: int,
height: int,
fps: float,
codec: str | None,
) -> dict:
"""创建 ready 状态的素材。"""
print("[2/2] 创建 ready 状态素材...")
body = {
"project_id": project_id,
"library_id": library_id,
"name": name,
"storage_key": storage_key,
"mime_type": mime_type,
"file_size": file_size,
"duration": duration,
"width": width,
"height": height,
"fps": fps,
"status": "ready",
"classification_status": "pending",
"metadata": {},
}
if codec:
body["codec"] = codec
result = api_request(
base_url,
token,
"POST",
"/api/v1/assets",
body=body,
)
return result
def main():
args = parse_args()
if not args.token:
print("[ERROR] 请通过 --token 参数或 API_TOKEN 环境变量提供登录 token", file=sys.stderr)
sys.exit(1)
print("=" * 60)
print("创建测试素材工具")
print("=" * 60)
print(f" API 地址: {args.base_url}")
print(f" 项目 ID: {args.project_id}")
print(f" 素材名称: {args.name}")
print(f" storage_key: {args.storage_key}")
print(f" 分辨率: {args.width}x{args.height} @ {args.fps}fps")
print(f" 时长: {args.duration}s")
print(f" 文件大小: {args.file_size} bytes")
print("=" * 60)
print()
# Step 1: 确保默认素材库存在
library_id = ensure_default_library(args.base_url, args.token, args.project_id, args.kind)
# Step 2: 创建素材
asset = create_asset(
base_url=args.base_url,
token=args.token,
project_id=args.project_id,
library_id=library_id,
name=args.name,
storage_key=args.storage_key,
mime_type=args.mime_type,
file_size=args.file_size,
duration=args.duration,
width=args.width,
height=args.height,
fps=args.fps,
codec=args.codec,
)
asset_id = asset.get("id")
print()
print("=" * 60)
print("✓ 素材创建成功!")
print("=" * 60)
print(f" asset_id: {asset_id}")
print(f" 状态: {asset.get('status')}")
print(f" 素材库 ID: {asset.get('library_id')}")
print(f" 文件 URL: {asset.get('file_url', 'N/A')}")
print("=" * 60)
print()
print("调用示例:")
print(f" export ASSET_ID={asset_id}")
print(" # 在生成任务中使用:")
print(" # --asset-id $ASSET_ID")
print()
# 输出 asset_id 到 stdout(方便脚本捕获)
print(asset_id)
if __name__ == "__main__":
main()
+327
View File
@@ -0,0 +1,327 @@
"""ASR 自动字幕集成测试 — 验证渲染管道接入 ASR 的完整链路。"""
from __future__ import annotations
import tempfile
from dataclasses import dataclass
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
from video_processing.unified_render_service import UnifiedRenderService
from packages.adapters.asr.mock_asr_service import MockASRService
from packages.domain.subtitle import SubtitleSegment, SubtitleTimeline
# ── Fixtures ──────────────────────────────────────────────────────────────────
@dataclass
class FakeClip:
"""模拟 EditPlanClip。"""
id: str
plan_id: str = "plan_001"
asset_id: str = "asset_001"
clip_type: str = "video"
start_time: float = 0.0
duration: float = 10.0
layer: int = 0
role: str = "main"
config: dict = None
@dataclass
class FakePlan:
"""模拟 EditPlan。"""
id: str = "plan_001"
config: dict = None
@pytest.fixture
def work_dir():
with tempfile.TemporaryDirectory() as tmpdir:
yield Path(tmpdir)
@pytest.fixture
def test_video_path():
"""用 ffmpeg 生成一个5秒的测试视频(带音频)。"""
import subprocess
with tempfile.TemporaryDirectory() as tmpdir:
video_path = Path(tmpdir) / "test.mp4"
cmd = [
"ffmpeg",
"-y",
"-f",
"lavfi",
"-i",
"testsrc=duration=5:size=320x240:rate=30",
"-f",
"lavfi",
"-i",
"sine=frequency=440:duration=5",
"-c:v",
"libx264",
"-preset",
"ultrafast",
"-c:a",
"aac",
"-shortest",
str(video_path),
]
result = subprocess.run(cmd, capture_output=True, timeout=30)
if result.returncode != 0:
pytest.skip(f"ffmpeg 不可用或生成测试视频失败: {result.stderr[:200]}")
yield video_path
# ── 测试:ASR 服务接入 ────────────────────────────────────────────────────────
class TestASRServiceIntegration:
def test_asr_service_in_init(self):
"""验证 asr_service 参数正确传递。"""
plan = FakePlan(config={})
service = UnifiedRenderService(
plan=plan,
clips=[],
asset_path_map={},
work_dir=Path("/tmp"),
asr_service=MockASRService(),
)
assert service.asr_service is not None
assert isinstance(service.asr_service, MockASRService)
def test_no_asr_service_default(self):
"""验证不传 asr_service 时默认 None。"""
plan = FakePlan(config={})
service = UnifiedRenderService(
plan=plan,
clips=[],
asset_path_map={},
work_dir=Path("/tmp"),
)
assert service.asr_service is None
def test_maybe_generate_ass_auto_subtitle_with_asr(self, work_dir, test_video_path):
"""验证 ASR 自动字幕模式:有 asr_service + auto_generated=true 时生成 ASS。"""
plan = FakePlan(
config={
"subtitle": {
"enabled": True,
"auto_generated": True,
"position": "bottom",
}
}
)
clips = [
FakeClip(id="clip_1", asset_id="asset_1", duration=5.0),
]
asset_map = {"asset_1": test_video_path}
service = UnifiedRenderService(
plan=plan,
clips=clips,
asset_path_map=asset_map,
work_dir=work_dir,
output_width=320,
output_height=240,
asr_service=MockASRService(mock_text="这是ASR自动生成的测试字幕。用来验证渲染管道是否正常接入。"),
)
ass_path = service._maybe_generate_ass(5.0)
assert ass_path is not None
assert ass_path.exists()
content = ass_path.read_text(encoding="utf-8")
assert "[Events]" in content
assert "Dialogue:" in content
assert "ASR" in content
def test_maybe_generate_ass_no_asr_service_skip_auto(self, work_dir):
"""验证没有 asr_service 时,即使 auto_generated=true 也不生成 ASR 字幕。"""
plan = FakePlan(
config={
"subtitle": {
"enabled": True,
"auto_generated": True,
"text": "",
}
}
)
service = UnifiedRenderService(
plan=plan,
clips=[],
asset_path_map={},
work_dir=work_dir,
asr_service=None, # 没有 ASR 服务
)
ass_path = service._maybe_generate_ass(5.0)
# 没有 ASR 服务 + 没有静态字幕文本 → 返回 None
assert ass_path is None
def test_maybe_generate_ass_auto_mode_ignores_text(self, work_dir):
"""验证 ASR 模式下即使有 text 字段也走 ASR(ASR无结果则无字幕)。"""
plan = FakePlan(
config={
"subtitle": {
"enabled": True,
"auto_generated": True,
"text": "静态字幕文本", # ASR模式下忽略此字段
}
}
)
mock_asr = MockASRService()
mock_asr.transcribe = MagicMock(side_effect=mock_asr.transcribe)
service = UnifiedRenderService(
plan=plan,
clips=[],
asset_path_map={},
work_dir=work_dir,
asr_service=mock_asr,
)
ass_path = service._maybe_generate_ass(5.0)
# ASR模式下无素材 → 无结果 → 返回None(不fallback到静态text
assert ass_path is None
def test_maybe_generate_ass_title_still_works(self, work_dir):
"""验证 ASR 模式下不影响 title 的处理(两者独立)。"""
plan = FakePlan(
config={
"title": {
"enabled": True,
"text": "视频标题",
"position": "top",
},
"subtitle": {
"enabled": False, # 字幕关闭
"auto_generated": True,
},
}
)
mock_asr = MockASRService()
service = UnifiedRenderService(
plan=plan,
clips=[],
asset_path_map={},
work_dir=work_dir,
asr_service=mock_asr,
)
ass_path = service._maybe_generate_ass(5.0)
assert ass_path is not None
content = ass_path.read_text(encoding="utf-8")
assert "视频标题" in content
def test_asr_failure_does_not_block(self, work_dir, test_video_path):
"""验证 ASR 失败时不阻断主流程,降级为无字幕。"""
plan = FakePlan(
config={
"subtitle": {
"enabled": True,
"auto_generated": True,
}
}
)
clips = [FakeClip(id="clip_1", asset_id="asset_1", duration=5.0)]
asset_map = {"asset_1": test_video_path}
# ASR 服务总是抛异常
bad_asr = MockASRService()
bad_asr.transcribe = MagicMock(side_effect=RuntimeError("ASR service down"))
service = UnifiedRenderService(
plan=plan,
clips=clips,
asset_path_map=asset_map,
work_dir=work_dir,
output_width=320,
output_height=240,
asr_service=bad_asr,
)
# 应该不抛异常,返回 None(降级)
ass_path = service._maybe_generate_ass(5.0)
assert ass_path is None # ASR 失败 → 无字幕
def test_auto_subtitle_disabled(self, work_dir):
"""验证 subtitle.enabled=false 时即使 auto_generated=true 也不生成。"""
plan = FakePlan(
config={
"subtitle": {
"enabled": False,
"auto_generated": True,
}
}
)
mock_asr = MockASRService()
mock_asr.transcribe = MagicMock()
service = UnifiedRenderService(
plan=plan,
clips=[],
asset_path_map={},
work_dir=work_dir,
asr_service=mock_asr,
)
ass_path = service._maybe_generate_ass(5.0)
assert ass_path is None
mock_asr.transcribe.assert_not_called()
# ── 测试:SubtitleConfig 扩展 ────────────────────────────────────────────────
class TestSubtitleConfigExtension:
def test_config_has_auto_generated_field(self):
"""验证 SubtitleConfig 有 auto_generated 字段。"""
from packages.domain.config_schemas import SubtitleConfig
config = SubtitleConfig()
assert hasattr(config, "auto_generated")
assert config.auto_generated is False # 默认关闭
def test_config_default_values(self):
"""验证新增字段的默认值。"""
from packages.domain.config_schemas import SubtitleConfig
config = SubtitleConfig()
assert config.auto_generated is False
assert config.language == ""
assert config.max_chars_per_line == 20
assert config.min_chars_per_segment == 8
def test_config_custom_values(self):
"""验证可以自定义 ASR 相关字段。"""
from packages.domain.config_schemas import SubtitleConfig
config = SubtitleConfig(
auto_generated=True,
language="zh",
max_chars_per_line=15,
min_chars_per_segment=5,
)
assert config.auto_generated is True
assert config.language == "zh"
assert config.max_chars_per_line == 15
assert config.min_chars_per_segment == 5
def test_config_validation_max_chars(self):
"""验证 max_chars_per_line 的范围校验。"""
from pydantic import ValidationError
from packages.domain.config_schemas import SubtitleConfig
with pytest.raises(ValidationError):
SubtitleConfig(max_chars_per_line=5) # 小于8
with pytest.raises(ValidationError):
SubtitleConfig(max_chars_per_line=50) # 大于40
+34 -11
View File
@@ -1,4 +1,7 @@
"""批量删除素材 + 分页优化 单元测试。"""
"""批量删除素材 + 分页优化 单元测试。
注意batch_delete 现在是软删除标记 status=deleted不是硬删除
"""
import sys
from pathlib import Path
@@ -10,7 +13,7 @@ from packages.domain import Asset, AssetStatus
class TestBatchDelete:
"""batch_delete 仓储方法测试。"""
"""batch_delete 仓储方法测试(软删除)"""
def _make_repo_with_assets(self):
repo = InMemoryAssetRepository()
@@ -26,7 +29,8 @@ class TestBatchDelete:
repo.create(asset)
return repo
def test_batch_delete_removes_multiple(self):
def test_batch_delete_marks_deleted_status(self):
"""软删除:status 变为 deleted,记录仍然存在。"""
repo = InMemoryAssetRepository()
assets = []
for i in range(5):
@@ -36,6 +40,7 @@ class TestBatchDelete:
name=f"voice_{i}.mp3",
storage_key=f"uploads/voice_{i}.mp3",
mime_type="audio/mpeg",
status=AssetStatus.READY,
)
repo.create(asset)
assets.append(asset)
@@ -44,13 +49,13 @@ class TestBatchDelete:
deleted_count = repo.batch_delete(ids_to_delete)
assert deleted_count == 3
# 验证确实被删了
assert repo.get(assets[0].id) is None
assert repo.get(assets[2].id) is None
assert repo.get(assets[4].id) is None
# 验证其他还在
assert repo.get(assets[1].id) is not None
assert repo.get(assets[3].id) is not None
# 软删除:记录仍在,status 变为 deleted
assert repo.get(assets[0].id).status == AssetStatus.DELETED
assert repo.get(assets[2].id).status == AssetStatus.DELETED
assert repo.get(assets[4].id).status == AssetStatus.DELETED
# 未删除的保持 ready
assert repo.get(assets[1].id).status == AssetStatus.READY
assert repo.get(assets[3].id).status == AssetStatus.READY
def test_batch_delete_empty_list(self):
repo = self._make_repo_with_assets()
@@ -69,9 +74,27 @@ class TestBatchDelete:
name="voice.mp3",
storage_key="uploads/voice.mp3",
mime_type="audio/mpeg",
status=AssetStatus.READY,
)
repo.create(asset)
deleted = repo.batch_delete([asset.id, "nonexistent"])
assert deleted == 1
assert repo.get(asset.id) is None
assert repo.get(asset.id).status == AssetStatus.DELETED
def test_batch_delete_idempotent(self):
"""重复删除已删除的素材不重复计数。"""
repo = InMemoryAssetRepository()
asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
name="voice.mp3",
storage_key="uploads/voice.mp3",
mime_type="audio/mpeg",
status=AssetStatus.READY,
)
repo.create(asset)
assert repo.batch_delete([asset.id]) == 1
assert repo.batch_delete([asset.id]) == 0
assert repo.get(asset.id).status == AssetStatus.DELETED
+288
View File
@@ -0,0 +1,288 @@
"""素材批量操作单元测试:软删除、批量打标签、批量分类、批量智能视图标记。"""
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
from packages.adapters.in_memory.asset_repository import InMemoryAssetRepository
from packages.domain import Asset, AssetStatus
class TestBatchSoftDelete:
"""batch_delete 软删除测试。"""
def _make_assets(self, repo: InMemoryAssetRepository, count: int = 5) -> list[Asset]:
assets = []
for i in range(count):
asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
name=f"asset_{i}.mp4",
storage_key=f"uploads/asset_{i}.mp4",
mime_type="video/mp4",
status=AssetStatus.READY,
)
repo.create(asset)
assets.append(asset)
return assets
def test_batch_soft_delete_marks_status_deleted(self):
"""软删除:status 变为 deleted,记录仍然存在。"""
repo = InMemoryAssetRepository()
assets = self._make_assets(repo, 3)
ids_to_delete = [assets[0].id, assets[2].id]
count = repo.batch_delete(ids_to_delete)
assert count == 2
# 记录仍在,只是 status 变了
assert repo.get(assets[0].id) is not None
assert repo.get(assets[0].id).status == AssetStatus.DELETED
assert repo.get(assets[2].id).status == AssetStatus.DELETED
# 未删除的保持原样
assert repo.get(assets[1].id).status == AssetStatus.READY
def test_batch_soft_delete_idempotent(self):
"""重复删除已删除的素材,计数不增加。"""
repo = InMemoryAssetRepository()
assets = self._make_assets(repo, 2)
count1 = repo.batch_delete([assets[0].id])
count2 = repo.batch_delete([assets[0].id])
assert count1 == 1
assert count2 == 0
assert repo.get(assets[0].id).status == AssetStatus.DELETED
def test_batch_soft_delete_empty_list(self):
repo = InMemoryAssetRepository()
self._make_assets(repo, 3)
assert repo.batch_delete([]) == 0
def test_batch_soft_delete_nonexistent_ids(self):
repo = InMemoryAssetRepository()
self._make_assets(repo, 3)
assert repo.batch_delete(["nonexistent-1", "nonexistent-2"]) == 0
class TestBatchUpdateMetadata:
"""batch_update_metadata 批量更新 metadata 测试。"""
def test_batch_update_category(self):
"""批量修改分类(metadata.category)。"""
repo = InMemoryAssetRepository()
assets = []
for i in range(3):
asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
name=f"v{i}.mp4",
storage_key=f"v{i}.mp4",
mime_type="video/mp4",
metadata={"existing_key": "existing_value"},
status=AssetStatus.READY,
)
repo.create(asset)
assets.append(asset)
ids = [a.id for a in assets]
count = repo.batch_update_metadata(ids, {"category": "person"})
assert count == 3
for a in assets:
updated = repo.get(a.id)
assert updated.metadata["category"] == "person"
assert updated.metadata["existing_key"] == "existing_value" # 合并而非覆盖
def test_batch_update_smart_view(self):
"""批量设置智能视图标记。"""
repo = InMemoryAssetRepository()
assets = []
for i in range(4):
asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
name=f"v{i}.mp4",
storage_key=f"v{i}.mp4",
mime_type="video/mp4",
status=AssetStatus.READY,
)
repo.create(asset)
assets.append(asset)
# 标记前2个为 recommended
count = repo.batch_update_metadata([assets[0].id, assets[1].id], {"smart_view": "recommended"})
assert count == 2
assert repo.get(assets[0].id).metadata["smart_view"] == "recommended"
assert repo.get(assets[1].id).metadata["smart_view"] == "recommended"
# 其余不变
assert "smart_view" not in repo.get(assets[2].id).metadata
# 再标记后2个为 high_risk
count2 = repo.batch_update_metadata([assets[2].id, assets[3].id], {"smart_view": "high_risk"})
assert count2 == 2
assert repo.get(assets[2].id).metadata["smart_view"] == "high_risk"
assert repo.get(assets[3].id).metadata["smart_view"] == "high_risk"
def test_batch_update_metadata_partial_existing(self):
"""部分素材存在时,只更新存在的。"""
repo = InMemoryAssetRepository()
asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
name="v.mp4",
storage_key="v.mp4",
mime_type="video/mp4",
status=AssetStatus.READY,
)
repo.create(asset)
count = repo.batch_update_metadata([asset.id, "nonexistent"], {"category": "scenic"})
assert count == 1
assert repo.get(asset.id).metadata["category"] == "scenic"
def test_batch_update_metadata_empty_list(self):
repo = InMemoryAssetRepository()
assert repo.batch_update_metadata([], {"category": "x"}) == 0
class TestBatchAddTags:
"""batch_add_tags 批量添加标签测试。"""
def test_batch_add_tags_merges_and_dedups(self):
"""添加模式:合并去重,已有标签不重复添加。"""
repo = InMemoryAssetRepository()
assets = []
for i in range(3):
asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
name=f"v{i}.mp4",
storage_key=f"v{i}.mp4",
mime_type="video/mp4",
status=AssetStatus.READY,
)
asset.add_tag("tag-existing")
repo.create(asset)
assets.append(asset)
ids = [a.id for a in assets]
count = repo.batch_add_tags(ids, ["tag-1", "tag-2", "tag-existing"])
assert count == 3 # 都有新增标签,所以都算变更
for a in assets:
updated = repo.get(a.id)
assert set(updated.tag_ids) == {"tag-existing", "tag-1", "tag-2"}
def test_batch_add_tags_no_change_when_all_exist(self):
"""所有标签都已存在时,返回0。"""
repo = InMemoryAssetRepository()
asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
name="v.mp4",
storage_key="v.mp4",
mime_type="video/mp4",
status=AssetStatus.READY,
)
asset.add_tag("tag-a")
asset.add_tag("tag-b")
repo.create(asset)
count = repo.batch_add_tags([asset.id], ["tag-a", "tag-b"])
assert count == 0
def test_batch_add_tags_empty_input(self):
repo = InMemoryAssetRepository()
assert repo.batch_add_tags([], ["tag-1"]) == 0
assert repo.batch_add_tags(["aid"], []) == 0
class TestBatchReplaceTags:
"""batch_replace_tags 批量替换标签测试。"""
def test_batch_replace_tags_full_override(self):
"""替换模式:全量覆盖原有标签。"""
repo = InMemoryAssetRepository()
assets = []
for i in range(3):
asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
name=f"v{i}.mp4",
storage_key=f"v{i}.mp4",
mime_type="video/mp4",
status=AssetStatus.READY,
)
asset.add_tag(f"old-{i}")
asset.add_tag("old-common")
repo.create(asset)
assets.append(asset)
ids = [a.id for a in assets]
count = repo.batch_replace_tags(ids, ["new-1", "new-2"])
assert count == 3
for a in assets:
updated = repo.get(a.id)
assert set(updated.tag_ids) == {"new-1", "new-2"}
def test_batch_replace_tags_empty_tags_clears_all(self):
"""替换为空列表:清空所有标签。"""
repo = InMemoryAssetRepository()
asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
name="v.mp4",
storage_key="v.mp4",
mime_type="video/mp4",
status=AssetStatus.READY,
)
asset.add_tag("tag-a")
asset.add_tag("tag-b")
repo.create(asset)
count = repo.batch_replace_tags([asset.id], [])
assert count == 1
assert repo.get(asset.id).tag_ids == []
def test_batch_replace_tags_empty_assets(self):
repo = InMemoryAssetRepository()
assert repo.batch_replace_tags([], ["tag-1"]) == 0
class TestBatchOperationLimits:
"""批量操作上限与边界测试。"""
def test_large_batch_operations(self):
"""大量素材的批量操作(验证性能基本可用)。"""
repo = InMemoryAssetRepository()
assets = []
for i in range(50):
asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
name=f"v{i}.mp4",
storage_key=f"v{i}.mp4",
mime_type="video/mp4",
status=AssetStatus.READY,
)
repo.create(asset)
assets.append(asset)
ids = [a.id for a in assets]
# 批量打标签
count = repo.batch_add_tags(ids, ["bulk-tag"])
assert count == 50
# 批量分类
count = repo.batch_update_metadata(ids, {"category": "scenic"})
assert count == 50
# 批量软删除
count = repo.batch_delete(ids)
assert count == 50
for a in assets:
assert repo.get(a.id).status == AssetStatus.DELETED
+359
View File
@@ -0,0 +1,359 @@
"""BGM 混音单元测试.
测试
- BGMConfig 配置解析与边界值
- 预设 BGM 库查询
- BGM 音频生成端到端 ffmpeg
- BGM + 主音频混音端到端 ffmpeg
- 淡入淡出效果
- 音量边界0 1
- sidechain 人声闪避
"""
import sys
import tempfile
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker"))
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
import pytest
from video_processing.bgm_mixer import BGMConfig, build_bgm_only, mix_bgm_with_main, prepare_bgm_track
from video_processing.render_audio import RenderContext
# ── Fixtures ──────────────────────────────────────────────────────────────────
@pytest.fixture
def work_dir(tmp_path):
return tmp_path
@pytest.fixture
def ctx(work_dir):
return RenderContext(work_dir=work_dir, plan_id="test_plan")
@pytest.fixture
def main_audio_path(work_dir):
"""生成 10 秒测试主音频(正弦波模拟人声)。"""
import subprocess
path = work_dir / "main.aac"
# 生成 10 秒 440Hz 正弦波模拟主音频
subprocess.run(
[
"ffmpeg",
"-y",
"-f",
"lavfi",
"-i",
"sine=frequency=440:duration=10:sample_rate=44100",
"-c:a",
"aac",
"-b:a",
"128k",
str(path),
],
capture_output=True,
check=True,
timeout=30,
)
return str(path)
@pytest.fixture
def bgm_audio_path(work_dir):
"""生成 5 秒测试 BGM(更低频率模拟背景音乐)。"""
import subprocess
path = work_dir / "bgm.aac"
# 生成 5 秒 220Hz 正弦波模拟 BGM
subprocess.run(
[
"ffmpeg",
"-y",
"-f",
"lavfi",
"-i",
"sine=frequency=220:duration=5:sample_rate=44100",
"-c:a",
"aac",
"-b:a",
"128k",
str(path),
],
capture_output=True,
check=True,
timeout=30,
)
return str(path)
# ── BGMConfig 测试 ───────────────────────────────────────────────────────────
class TestBGMConfig:
"""BGMConfig 配置解析测试。"""
def test_default_values(self):
cfg = BGMConfig(bgm_path="/tmp/bgm.mp3")
assert cfg.volume == 0.3
assert cfg.fade_in == 0.0
assert cfg.fade_out == 0.0
assert cfg.loop_enabled is True
assert cfg.sidechain_enabled is False
assert cfg.sidechain_ratio == 0.3
def test_from_config_dict(self):
config_dict = {
"enabled": True,
"volume": 0.5,
"fade_in": 2.0,
"fade_out": 3.0,
"loop_enabled": False,
"sidechain_enabled": True,
"sidechain_ratio": 0.5,
}
cfg = BGMConfig.from_config_dict("/bgm.mp3", config_dict)
assert cfg.bgm_path == "/bgm.mp3"
assert cfg.volume == 0.5
assert cfg.fade_in == 2.0
assert cfg.fade_out == 3.0
assert cfg.loop_enabled is False
assert cfg.sidechain_enabled is True
assert cfg.sidechain_ratio == 0.5
def test_volume_clamped_by_config_schema(self):
"""音量边界由 Pydantic Schema 在入口层保证,内部直接使用。"""
from packages.domain.config_schemas import BGMConfig as BGMConfigSchema
# 边界值测试
cfg = BGMConfigSchema(enabled=True, volume=0.0)
assert cfg.volume == 0.0
cfg = BGMConfigSchema(enabled=True, volume=1.0)
assert cfg.volume == 1.0
def test_fade_boundaries(self):
from packages.domain.config_schemas import BGMConfig as BGMConfigSchema
# 0 是合法值
cfg = BGMConfigSchema(fade_in=0, fade_out=0)
assert cfg.fade_in == 0.0
assert cfg.fade_out == 0.0
# ── 预设 BGM 库测试 ─────────────────────────────────────────────────────────
class TestPresetBGM:
"""预设 BGM 库查询测试。"""
def test_total_count(self):
from packages.domain.preset_bgm import PRESET_BGM_LIBRARY
assert len(PRESET_BGM_LIBRARY) >= 10
def test_get_preset_by_id(self):
from packages.domain.preset_bgm import get_preset_bgm
bgm = get_preset_bgm("bgm_upbeat_001")
assert bgm is not None
assert bgm.name == "阳光清晨"
assert bgm.style == "upbeat"
def test_get_preset_not_found(self):
from packages.domain.preset_bgm import get_preset_bgm
assert get_preset_bgm("nonexistent") is None
def test_list_by_style(self):
from packages.domain.preset_bgm import list_preset_bgm_by_style
upbeat = list_preset_bgm_by_style("upbeat")
assert len(upbeat) >= 3
assert all(b.style == "upbeat" for b in upbeat)
def test_search_by_keyword(self):
from packages.domain.preset_bgm import search_preset_bgm
results = search_preset_bgm("钢琴")
assert len(results) >= 2
assert any("钢琴" in b.tags for b in results)
def test_all_presets_have_basic_fields(self):
from packages.domain.preset_bgm import PRESET_BGM_LIBRARY
for bgm in PRESET_BGM_LIBRARY:
assert bgm.id, f"{bgm.name} 缺少 id"
assert bgm.name, "缺少 name"
assert bgm.style, f"{bgm.name} 缺少 style"
assert bgm.duration > 0, f"{bgm.name} 时长无效"
# ── BGM 处理端到端测试 ──────────────────────────────────────────────────────
class TestPrepareBGMTrack:
"""prepare_bgm_track 端到端测试。"""
def test_bgm_without_loop_short_duration(self, ctx, bgm_audio_path):
"""BGM 比目标时长短且不循环 → 截断到目标时长(但前面没有足够内容)。"""
bgm = BGMConfig(bgm_path=bgm_audio_path, volume=0.5, loop_enabled=False)
result = prepare_bgm_track(ctx, bgm, target_duration=3.0)
assert result.exists()
assert result.stat().st_size > 0
def test_bgm_with_loop_longer_duration(self, ctx, bgm_audio_path):
"""BGM 比目标时长短,循环铺满。"""
bgm = BGMConfig(bgm_path=bgm_audio_path, volume=0.3, loop_enabled=True)
# BGM 5 秒,目标 12 秒,需要循环 3 次
result = prepare_bgm_track(ctx, bgm, target_duration=12.0)
assert result.exists()
assert result.stat().st_size > 0
def test_bgm_fade_in_and_fade_out(self, ctx, bgm_audio_path):
"""BGM 淡入淡出效果。"""
bgm = BGMConfig(
bgm_path=bgm_audio_path,
volume=0.5,
fade_in=1.0,
fade_out=1.0,
loop_enabled=False,
)
result = prepare_bgm_track(ctx, bgm, target_duration=4.0)
assert result.exists()
assert result.stat().st_size > 0
def test_volume_zero(self, ctx, bgm_audio_path):
"""音量为 0 时仍能正常处理。"""
bgm = BGMConfig(bgm_path=bgm_audio_path, volume=0.0, loop_enabled=False)
result = prepare_bgm_track(ctx, bgm, target_duration=3.0)
assert result.exists()
assert result.stat().st_size > 0
def test_volume_one(self, ctx, bgm_audio_path):
"""音量为 1(最大)时正常处理。"""
bgm = BGMConfig(bgm_path=bgm_audio_path, volume=1.0, loop_enabled=False)
result = prepare_bgm_track(ctx, bgm, target_duration=3.0)
assert result.exists()
assert result.stat().st_size > 0
class TestMixBGMMain:
"""BGM + 主音频混音端到端测试。"""
def test_simple_mix(self, ctx, main_audio_path, bgm_audio_path):
"""普通 amix 混音(无 sidechain)。"""
bgm = BGMConfig(
bgm_path=bgm_audio_path,
volume=0.3,
loop_enabled=True,
sidechain_enabled=False,
)
result = mix_bgm_with_main(ctx, Path(main_audio_path), bgm, target_duration=8.0)
assert result.exists()
assert result.stat().st_size > 0
def test_sidechain_mix(self, ctx, main_audio_path, bgm_audio_path):
"""sidechain 人声闪避混音。"""
bgm = BGMConfig(
bgm_path=bgm_audio_path,
volume=0.5,
loop_enabled=True,
sidechain_enabled=True,
sidechain_ratio=0.3,
sidechain_threshold=-25.0,
sidechain_attack=0.02,
sidechain_release=0.5,
)
result = mix_bgm_with_main(ctx, Path(main_audio_path), bgm, target_duration=8.0)
assert result.exists()
assert result.stat().st_size > 0
def test_sidechain_max_ratio(self, ctx, main_audio_path, bgm_audio_path):
"""sidechain 最大闪避比例。"""
bgm = BGMConfig(
bgm_path=bgm_audio_path,
volume=0.5,
loop_enabled=True,
sidechain_enabled=True,
sidechain_ratio=0.9, # 降低 90%
)
result = mix_bgm_with_main(ctx, Path(main_audio_path), bgm, target_duration=5.0)
assert result.exists()
assert result.stat().st_size > 0
class TestBuildBGMOnly:
"""纯 BGM 模式测试。"""
def test_build_bgm_only(self, ctx, bgm_audio_path):
"""只有 BGM、没有主音频时生成纯 BGM 音频。"""
bgm = BGMConfig(
bgm_path=bgm_audio_path,
volume=0.3,
fade_in=1.0,
fade_out=1.0,
loop_enabled=True,
)
result = build_bgm_only(ctx, bgm, target_duration=15.0)
assert result.exists()
assert result.stat().st_size > 0
# ── Config Schema 集成测试 ───────────────────────────────────────────────────
class TestConfigSchemaIntegration:
"""config schema 与渲染配置的集成测试。"""
def test_full_bgm_config(self):
"""完整 BGM 配置能正确解析。"""
from packages.domain.config_schemas import EditPlanConfigSchema, normalize_plan_config
config = normalize_plan_config(
{
"bgm": {
"enabled": True,
"source": "library",
"asset_id": "bgm-asset-001",
"volume": 0.4,
"fade_in": 2.5,
"fade_out": 3.0,
"loop_enabled": True,
"sidechain_enabled": True,
"sidechain_ratio": 0.4,
}
}
)
bgm = config["bgm"]
assert bgm["enabled"] is True
assert bgm["volume"] == 0.4
assert bgm["fade_in"] == 2.5
assert bgm["fade_out"] == 3.0
assert bgm["loop_enabled"] is True
assert bgm["sidechain_enabled"] is True
assert bgm["sidechain_ratio"] == 0.4
# 默认值保留
assert bgm["sidechain_attack"] == 0.02
assert bgm["sidechain_release"] == 0.5
assert bgm["sidechain_threshold"] == -25.0
def test_bgm_disabled_by_default(self):
"""默认 BGM 是关闭的。"""
from packages.domain.config_schemas import normalize_plan_config
config = normalize_plan_config({})
assert config["bgm"]["enabled"] is False
+540
View File
@@ -0,0 +1,540 @@
"""绿幕抠像 + 音频降噪引擎 单元测试."""
from __future__ import annotations
import pytest
from video_processing.chroma_key_engine import (
CHROMA_KEY_PRESETS,
ChromaKeyConfig,
ChromaKeyEngine,
apply_chroma_key_if_needed,
)
from video_processing.noise_reduction_engine import (
NoiseReductionConfig,
NoiseReductionEngine,
NoiseReductionLevel,
apply_noise_reduction_if_needed,
)
# ═══════════════════════════════════════════════════════════════
# ChromaKeyConfig 测试
# ═══════════════════════════════════════════════════════════════
class TestChromaKeyConfig:
"""绿幕抠像配置测试."""
def test_default_disabled(self):
"""默认配置是禁用的."""
config = ChromaKeyConfig()
assert config.enabled is False
assert config.has_effect() is False
def test_from_dict_none(self):
"""传入 None 返回禁用配置."""
config = ChromaKeyConfig.from_dict(None)
assert config.enabled is False
assert config.has_effect() is False
def test_from_dict_empty(self):
"""传入空 dict 返回禁用配置."""
config = ChromaKeyConfig.from_dict({})
assert config.enabled is False
def test_from_dict_disabled(self):
"""enabled=false 时禁用."""
config = ChromaKeyConfig.from_dict({"enabled": False})
assert config.enabled is False
assert config.has_effect() is False
def test_from_dict_enabled_defaults(self):
"""只开启,使用默认参数."""
config = ChromaKeyConfig.from_dict({"enabled": True})
assert config.enabled is True
assert config.key_color == "#00FF00"
assert config.similarity == 0.3
assert config.blend == 0.1
assert config.spill_suppress == 0.0
assert config.has_effect() is True
def test_from_dict_custom_params(self):
"""自定义所有参数."""
config = ChromaKeyConfig.from_dict(
{
"enabled": True,
"key_color": "#0000FF",
"similarity": 0.5,
"blend": 0.2,
"spill_suppress": 0.3,
}
)
assert config.key_color == "#0000FF"
assert config.similarity == 0.5
assert config.blend == 0.2
assert config.spill_suppress == 0.3
def test_similarity_clamp(self):
"""similarity 越界自动钳制."""
# 低于最小值
config = ChromaKeyConfig.from_dict({"enabled": True, "similarity": 0})
assert config.similarity == 0.01
# 高于最大值
config = ChromaKeyConfig.from_dict({"enabled": True, "similarity": 2.0})
assert config.similarity == 1.0
def test_blend_clamp(self):
"""blend 越界自动钳制."""
config = ChromaKeyConfig.from_dict({"enabled": True, "blend": -0.5})
assert config.blend == 0.0
config = ChromaKeyConfig.from_dict({"enabled": True, "blend": 2.0})
assert config.blend == 1.0
def test_spill_suppress_clamp(self):
"""spill_suppress 越界自动钳制."""
config = ChromaKeyConfig.from_dict({"enabled": True, "spill_suppress": -0.1})
assert config.spill_suppress == 0.0
config = ChromaKeyConfig.from_dict({"enabled": True, "spill_suppress": 2.0})
assert config.spill_suppress == 1.0
def test_invalid_similarity_still_works(self):
"""无效相似度值也能安全解析(钳制后仍有效果)."""
config = ChromaKeyConfig.from_dict({"enabled": True, "similarity": "invalid"})
# 字符串转 float 会失败 → 应该用 try/except 保护
# 实际上 from_dict 直接 float() 转换会抛异常
# 这里测试调用方的降级策略
def test_has_effect_zero_similarity(self):
"""similarity 为 0(被钳制到0.01)时仍然有效果."""
config = ChromaKeyConfig(enabled=True, similarity=0.0)
# 注意:直接构造不走 from_dict 的钳制逻辑
assert config.similarity == 0.0
assert config.has_effect() is False # similarity > 0
# ═══════════════════════════════════════════════════════════════
# ChromaKeyEngine 测试
# ═══════════════════════════════════════════════════════════════
class TestChromaKeyEngine:
"""绿幕抠像引擎测试."""
def test_build_filter_basic(self):
"""基础抠像滤镜构建."""
config = ChromaKeyConfig(enabled=True, key_color="#00FF00", similarity=0.3, blend=0.1)
engine = ChromaKeyEngine(config)
result = engine.build_filter("[0:v]", "[out]")
assert "[0:v]" in result
assert "[out]" in result
assert "colorkey" in result
assert "color=0x00FF00" in result
assert "similarity=0.3" in result
assert "blend=0.1" in result
def test_build_filter_no_effect(self):
"""无效果时返回 copy."""
config = ChromaKeyConfig(enabled=False)
engine = ChromaKeyEngine(config)
result = engine.build_filter("[in]", "[out]")
assert "copy" in result
assert "colorkey" not in result
def test_normalize_color_hex(self):
"""hex 颜色格式化."""
engine = ChromaKeyEngine(ChromaKeyConfig(enabled=True))
assert engine._normalize_color("#00FF00") == "0x00FF00"
assert engine._normalize_color("#00ff00") == "0x00FF00"
assert engine._normalize_color("0x00FF00") == "0X00FF00"
def test_normalize_color_name(self):
"""颜色名直接透传."""
engine = ChromaKeyEngine(ChromaKeyConfig(enabled=True))
assert engine._normalize_color("green") == "green"
assert engine._normalize_color("blue") == "blue"
def test_build_filter_with_spill_suppress(self):
"""溢色抑制时增加 colorchannelmixer."""
config = ChromaKeyConfig(enabled=True, key_color="#00FF00", similarity=0.3, blend=0.1, spill_suppress=0.5)
engine = ChromaKeyEngine(config)
result = engine.build_filter("[v0]", "[v1]")
assert "colorkey" in result
assert "colorchannelmixer" in result
# 绿通道增益应该降低
assert "gg=" in result
def test_build_filter_no_spill_suppress(self):
"""无溢色抑制时不含 colorchannelmixer."""
config = ChromaKeyConfig(enabled=True, key_color="#00FF00", similarity=0.3, blend=0.1, spill_suppress=0.0)
engine = ChromaKeyEngine(config)
result = engine.build_filter("[v0]", "[v1]")
assert "colorkey" in result
assert "colorchannelmixer" not in result
def test_build_filter_chromakey(self):
"""chromakey 滤镜构建(高级版本)."""
config = ChromaKeyConfig(enabled=True, key_color="#00FF00", similarity=0.3, blend=0.1)
engine = ChromaKeyEngine(config)
result = engine.build_filter_chromakey("[in]", "[out]")
assert "chromakey" in result
assert "color=0x00FF00" in result
def test_blue_screen(self):
"""蓝幕抠像."""
config = ChromaKeyConfig.from_dict({"enabled": True, "key_color": "#0000FF", "similarity": 0.3})
engine = ChromaKeyEngine(config)
result = engine.build_filter("[0:v]", "[out]")
assert "color=0x0000FF" in result
# ═══════════════════════════════════════════════════════════════
# 预设测试
# ═══════════════════════════════════════════════════════════════
class TestChromaKeyPresets:
"""绿幕预设测试."""
def test_presets_exist(self):
"""预设列表包含常见预设."""
assert "green_screen" in CHROMA_KEY_PRESETS
assert "blue_screen" in CHROMA_KEY_PRESETS
assert "red_screen" in CHROMA_KEY_PRESETS
assert "precise_green" in CHROMA_KEY_PRESETS
assert "soft_green" in CHROMA_KEY_PRESETS
def test_green_screen_preset_valid(self):
"""绿幕预设参数有效."""
preset = CHROMA_KEY_PRESETS["green_screen"]
config = ChromaKeyConfig.from_dict({"enabled": True, **preset})
assert config.has_effect() is True
assert config.key_color == "#00FF00"
assert 0.01 <= config.similarity <= 1.0
def test_blue_screen_preset_valid(self):
"""蓝幕预设参数有效."""
preset = CHROMA_KEY_PRESETS["blue_screen"]
config = ChromaKeyConfig.from_dict({"enabled": True, **preset})
assert config.key_color == "#0000FF"
# ═══════════════════════════════════════════════════════════════
# apply_chroma_key_if_needed 测试
# ═══════════════════════════════════════════════════════════════
class TestApplyChromaKeyIfNeeded:
"""便捷函数测试."""
def test_no_chroma_key_in_config(self):
"""没有 chroma_key 配置时返回 None."""
result = apply_chroma_key_if_needed({}, "[in]", "[out]")
assert result is None
def test_disabled_chroma_key(self):
"""禁用的抠像配置返回 None."""
result = apply_chroma_key_if_needed({"chroma_key": {"enabled": False}}, "[in]", "[out]")
assert result is None
def test_enabled_chroma_key(self):
"""启用的抠像配置返回滤镜字符串."""
result = apply_chroma_key_if_needed(
{"chroma_key": {"enabled": True, "key_color": "#00FF00"}},
"[v0]",
"[ck0]",
)
assert result is not None
assert "colorkey" in result
assert "[v0]" in result
assert "[ck0]" in result
def test_invalid_config_degrades_gracefully(self):
"""无效配置不抛出异常,返回 None."""
result = apply_chroma_key_if_needed(
{"chroma_key": {"enabled": True, "similarity": "invalid"}},
"[in]",
"[out]",
)
# float("invalid") 会抛 ValueError,但 apply 函数应该捕获
# 注意:当前 from_dict 没有 try/except,调用方的 apply 应该处理
# 这里验证不会崩溃
assert result is None or isinstance(result, str)
# ═══════════════════════════════════════════════════════════════
# NoiseReductionConfig 测试
# ═══════════════════════════════════════════════════════════════
class TestNoiseReductionConfig:
"""音频降噪配置测试."""
def test_default_disabled(self):
"""默认配置是禁用的."""
config = NoiseReductionConfig()
assert config.enabled is False
assert config.has_effect() is False
def test_from_dict_none(self):
"""传入 None 返回禁用配置."""
config = NoiseReductionConfig.from_dict(None)
assert config.enabled is False
assert config.has_effect() is False
def test_from_dict_empty(self):
"""传入空 dict 返回禁用配置."""
config = NoiseReductionConfig.from_dict({})
assert config.enabled is False
def test_from_dict_enabled_default(self):
"""只开启,使用默认参数."""
config = NoiseReductionConfig.from_dict({"enabled": True})
assert config.enabled is True
assert config.level == NoiseReductionLevel.MEDIUM
assert config.has_effect() is True
def test_from_dict_low_level(self):
"""低降噪等级."""
config = NoiseReductionConfig.from_dict({"enabled": True, "level": "low"})
assert config.level == NoiseReductionLevel.LOW
assert config.get_effective_noise_floor() == -35.0
def test_from_dict_medium_level(self):
"""中降噪等级."""
config = NoiseReductionConfig.from_dict({"enabled": True, "level": "medium"})
assert config.level == NoiseReductionLevel.MEDIUM
assert config.get_effective_noise_floor() == -25.0
def test_from_dict_high_level(self):
"""高降噪等级."""
config = NoiseReductionConfig.from_dict({"enabled": True, "level": "high"})
assert config.level == NoiseReductionLevel.HIGH
assert config.get_effective_noise_floor() == -15.0
def test_from_dict_custom_level(self):
"""自定义降噪等级."""
config = NoiseReductionConfig.from_dict({"enabled": True, "level": "custom", "noise_floor": -30.0})
assert config.level == NoiseReductionLevel.CUSTOM
assert config.get_effective_noise_floor() == -30.0
def test_invalid_level_falls_back_to_medium(self):
"""无效等级回退到 medium."""
config = NoiseReductionConfig.from_dict({"enabled": True, "level": "ultra"})
assert config.level == NoiseReductionLevel.MEDIUM
def test_noise_floor_clamp(self):
"""noise_floor 越界自动钳制."""
# 低于最小值
config = NoiseReductionConfig.from_dict({"enabled": True, "level": "custom", "noise_floor": -100})
assert config.noise_floor == -60.0
# 高于最大值
config = NoiseReductionConfig.from_dict({"enabled": True, "level": "custom", "noise_floor": 0})
assert config.noise_floor == -5.0
def test_voice_enhance(self):
"""人声增强开关."""
config = NoiseReductionConfig.from_dict({"enabled": True, "voice_enhance": True})
assert config.voice_enhance is True
# ═══════════════════════════════════════════════════════════════
# NoiseReductionEngine 测试
# ═══════════════════════════════════════════════════════════════
class TestNoiseReductionEngine:
"""音频降噪引擎测试."""
def test_build_filter_basic(self):
"""基础降噪滤镜构建."""
config = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.MEDIUM)
engine = NoiseReductionEngine(config)
result = engine.build_filter("[0:a]", "[out]")
assert "[0:a]" in result
assert "[out]" in result
assert "afftdn" in result
assert "nf=-25" in result
def test_build_filter_no_effect(self):
"""无效果时返回 anull."""
config = NoiseReductionConfig(enabled=False)
engine = NoiseReductionEngine(config)
result = engine.build_filter("[in]", "[out]")
assert "anull" in result
assert "afftdn" not in result
def test_low_level(self):
"""低降噪等级参数正确."""
config = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.LOW)
engine = NoiseReductionEngine(config)
result = engine.build_filter("[in]", "[out]")
assert "nf=-35" in result
def test_high_level(self):
"""高降噪等级参数正确."""
config = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.HIGH)
engine = NoiseReductionEngine(config)
result = engine.build_filter("[in]", "[out]")
assert "nf=-15" in result
def test_custom_level(self):
"""自定义降噪等级."""
config = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.CUSTOM, noise_floor=-40.0)
engine = NoiseReductionEngine(config)
result = engine.build_filter("[in]", "[out]")
assert "nf=-40" in result
def test_voice_enhance_adds_filters(self):
"""人声增强增加额外滤镜."""
config = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.MEDIUM, voice_enhance=True)
engine = NoiseReductionEngine(config)
result = engine.build_filter("[in]", "[out]")
assert "afftdn" in result
assert "highpass" in result
assert "acompressor" in result
assert "loudnorm" in result
def test_no_voice_enhance_clean(self):
"""无人声增强时只有 afftdn."""
config = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.MEDIUM, voice_enhance=False)
engine = NoiseReductionEngine(config)
result = engine.build_filter("[in]", "[out]")
assert "afftdn" in result
assert "highpass" not in result
assert "acompressor" not in result
def test_arnndn_filter(self):
"""RNN 降噪滤镜构建."""
config = NoiseReductionConfig(enabled=True, level=NoiseReductionLevel.MEDIUM)
engine = NoiseReductionEngine(config)
result = engine.build_filter_arnndn("[in]", "[out]", "/models/rnnoise.rnnn")
assert "arnndn" in result
assert "m=/models/rnnoise.rnnn" in result
# ═══════════════════════════════════════════════════════════════
# apply_noise_reduction_if_needed 测试
# ═══════════════════════════════════════════════════════════════
class TestApplyNoiseReductionIfNeeded:
"""便捷函数测试."""
def test_none_config(self):
"""None 配置返回 None."""
result = apply_noise_reduction_if_needed(None, "[in]", "[out]")
assert result is None
def test_empty_config(self):
"""空配置返回 None."""
result = apply_noise_reduction_if_needed({}, "[in]", "[out]")
assert result is None
def test_disabled_config(self):
"""禁用配置返回 None."""
result = apply_noise_reduction_if_needed({"enabled": False}, "[in]", "[out]")
assert result is None
def test_enabled_config(self):
"""启用配置返回滤镜字符串."""
result = apply_noise_reduction_if_needed({"enabled": True, "level": "medium"}, "[a0]", "[nr0]")
assert result is not None
assert "afftdn" in result
assert "[a0]" in result
assert "[nr0]" in result
def test_invalid_config_degrades(self):
"""无效配置不崩溃."""
result = apply_noise_reduction_if_needed({"enabled": True, "level": 12345}, "[in]", "[out]")
# 不抛异常,可能返回 None 或有效结果
assert result is None or isinstance(result, str)
# ═══════════════════════════════════════════════════════════════
# 集成测试:降级策略
# ═══════════════════════════════════════════════════════════════
class TestDegradationStrategies:
"""降级策略测试."""
def test_chroma_key_none_config_safe(self):
"""绿幕:None 配置安全."""
# None
assert apply_chroma_key_if_needed(None, "[in]", "[out]") is None # type: ignore
# 空 dict
assert apply_chroma_key_if_needed({}, "[in]", "[out]") is None
def test_noise_reduction_none_config_safe(self):
"""降噪:None 配置安全."""
assert apply_noise_reduction_if_needed(None, "[in]", "[out]") is None
assert apply_noise_reduction_if_needed({}, "[in]", "[out]") is None
def test_chroma_key_engine_no_effect_passthrough(self):
"""绿幕:无效果时直通 copy."""
config = ChromaKeyConfig(enabled=False)
engine = ChromaKeyEngine(config)
result = engine.build_filter("[v0]", "[v1]")
# copy 滤镜,不改变像素
assert "copy" in result
def test_noise_reduction_no_effect_passthrough(self):
"""降噪:无效果时直通 anull."""
config = NoiseReductionConfig(enabled=False)
engine = NoiseReductionEngine(config)
result = engine.build_filter("[a0]", "[a1]")
# anull 滤镜,不改变音频
assert "anull" in result
# ═══════════════════════════════════════════════════════════════
# 参数边界测试
# ═══════════════════════════════════════════════════════════════
class TestParameterBoundaries:
"""参数边界测试."""
@pytest.mark.parametrize(
"similarity,expected",
[
(0.0, 0.01), # 低于最小值 → 钳制到 min
(0.01, 0.01), # 最小值
(0.5, 0.5), # 中间值
(1.0, 1.0), # 最大值
(2.0, 1.0), # 超过最大值 → 钳制到 max
],
)
def test_similarity_boundaries(self, similarity, expected):
"""similarity 边界值测试."""
config = ChromaKeyConfig.from_dict({"enabled": True, "similarity": similarity})
assert abs(config.similarity - expected) < 0.001
@pytest.mark.parametrize(
"blend,expected",
[
(-1.0, 0.0),
(0.0, 0.0),
(0.5, 0.5),
(1.0, 1.0),
(2.0, 1.0),
],
)
def test_blend_boundaries(self, blend, expected):
"""blend 边界值测试."""
config = ChromaKeyConfig.from_dict({"enabled": True, "blend": blend})
assert abs(config.blend - expected) < 0.001
@pytest.mark.parametrize(
"noise_floor,expected",
[
(-100, -60.0),
(-60, -60.0),
(-30, -30.0),
(-5, -5.0),
(0, -5.0),
],
)
def test_noise_floor_boundaries(self, noise_floor, expected):
"""noise_floor 边界值测试."""
config = NoiseReductionConfig.from_dict({"enabled": True, "level": "custom", "noise_floor": noise_floor})
assert abs(config.noise_floor - expected) < 0.001
+572
View File
@@ -0,0 +1,572 @@
"""滤镜调色引擎单元测试."""
from __future__ import annotations
import pytest
from video_processing.color_grade_engine import (
DEFAULT_PARAMS,
PARAM_RANGES,
PRESET_BW,
PRESET_CINEMA,
PRESET_COOL,
PRESET_DISPLAY_NAMES,
PRESET_FILM,
PRESET_FRESH,
PRESET_JAPANESE,
PRESET_PARAMS,
PRESET_VINTAGE,
PRESET_WARM,
ColorGradeConfig,
ColorGradeEngine,
get_preset_names,
get_preset_params,
)
# ── 预设常量测试 ──────────────────────────────────────────────────────────────
class TestPresetConstants:
"""预设常量完整性测试."""
def test_eight_presets_defined(self):
"""应该有8种预设."""
assert len(PRESET_PARAMS) == 8
assert len(PRESET_DISPLAY_NAMES) == 8
def test_all_presets_have_display_names(self):
"""每个预设都应该有中文显示名."""
for key in PRESET_PARAMS:
assert key in PRESET_DISPLAY_NAMES
assert PRESET_DISPLAY_NAMES[key] # 非空
def test_preset_params_have_all_keys(self):
"""每个预设应该包含所有5个参数."""
required_keys = {"brightness", "contrast", "saturation", "temperature", "hue"}
for key, params in PRESET_PARAMS.items():
assert required_keys.issubset(params.keys()), f"预设 {key} 缺少参数"
def test_preset_params_in_valid_range(self):
"""所有预设参数应该在合法范围内."""
for preset_name, params in PRESET_PARAMS.items():
for param_name, value in params.items():
min_val, max_val = PARAM_RANGES[param_name]
assert (
min_val <= value <= max_val
), f"预设 {preset_name}{param_name}={value} 超出范围 [{min_val}, {max_val}]"
def test_black_white_has_zero_saturation(self):
"""黑白预设饱和度应该为0."""
assert PRESET_PARAMS[PRESET_BW]["saturation"] == 0
def test_warm_preset_has_positive_temperature(self):
"""暖色预设色温应该为正."""
assert PRESET_PARAMS[PRESET_WARM]["temperature"] > 0
def test_cool_preset_has_negative_temperature(self):
"""冷色预设色温应该为负."""
assert PRESET_PARAMS[PRESET_COOL]["temperature"] < 0
# ── ColorGradeConfig.from_dict 测试 ───────────────────────────────────────────
class TestColorGradeConfigFromDict:
"""配置字典解析测试."""
def test_none_config(self):
"""None返回disabled."""
config = ColorGradeConfig.from_dict(None)
assert not config.enabled
def test_empty_dict(self):
"""空字典返回disabled."""
config = ColorGradeConfig.from_dict({})
assert not config.enabled
def test_enabled_false(self):
"""enabled=False返回disabled."""
config = ColorGradeConfig.from_dict({"enabled": False})
assert not config.enabled
def test_enabled_only(self):
"""只开enabled,无预设无自定义参数."""
config = ColorGradeConfig.from_dict({"enabled": True})
assert config.enabled
assert config.preset == ""
assert config.brightness is None
assert config.contrast is None
assert config.saturation is None
assert config.temperature is None
assert config.hue is None
def test_with_preset(self):
"""指定预设."""
config = ColorGradeConfig.from_dict({"enabled": True, "preset": PRESET_FRESH})
assert config.enabled
assert config.preset == PRESET_FRESH
def test_invalid_preset_ignored(self):
"""无效预设名应该被忽略."""
config = ColorGradeConfig.from_dict({"enabled": True, "preset": "invalid_preset"})
assert config.preset == "" # 被清空
def test_with_custom_params(self):
"""自定义参数覆盖."""
config = ColorGradeConfig.from_dict(
{
"enabled": True,
"brightness": 20,
"contrast": -10,
"saturation": 150,
"temperature": 25,
"hue": 30,
}
)
assert config.enabled
assert config.brightness == 20
assert config.contrast == -10
assert config.saturation == 150
assert config.temperature == 25
assert config.hue == 30
def test_string_numeric_values(self):
"""字符串形式的数字应该能解析."""
config = ColorGradeConfig.from_dict(
{
"enabled": True,
"brightness": "20.5",
"saturation": "150",
}
)
assert config.brightness == 20.5
assert config.saturation == 150.0
def test_invalid_value_returns_none(self):
"""无效值应该返回None(不覆盖)."""
config = ColorGradeConfig.from_dict(
{
"enabled": True,
"brightness": "not_a_number",
}
)
assert config.brightness is None
# ── ColorGradeConfig.resolve_params 测试 ──────────────────────────────────────
class TestResolveParams:
"""参数解析与边界钳制测试."""
def test_default_params_when_empty(self):
"""无预设无自定义时返回默认值."""
config = ColorGradeConfig(enabled=True)
params = config.resolve_params()
for key, val in DEFAULT_PARAMS.items():
assert params[key] == val
def test_preset_params_applied(self):
"""预设参数应该被应用."""
config = ColorGradeConfig(enabled=True, preset=PRESET_FRESH)
params = config.resolve_params()
preset = PRESET_PARAMS[PRESET_FRESH]
for key, val in preset.items():
assert params[key] == val
def test_custom_overrides_preset(self):
"""自定义参数应该覆盖预设值."""
config = ColorGradeConfig(
enabled=True,
preset=PRESET_FRESH,
brightness=50, # 覆盖预设的8
)
params = config.resolve_params()
assert params["brightness"] == 50
# 其他参数还是预设值
assert params["contrast"] == PRESET_PARAMS[PRESET_FRESH]["contrast"]
def test_clamp_brightness_high(self):
"""亮度超过上限应该被钳制."""
config = ColorGradeConfig(enabled=True, brightness=200)
params = config.resolve_params()
assert params["brightness"] == 100
def test_clamp_brightness_low(self):
"""亮度低于下限应该被钳制."""
config = ColorGradeConfig(enabled=True, brightness=-200)
params = config.resolve_params()
assert params["brightness"] == -100
def test_clamp_saturation_low(self):
"""饱和度低于0应该被钳制到0."""
config = ColorGradeConfig(enabled=True, saturation=-50)
params = config.resolve_params()
assert params["saturation"] == 0
def test_clamp_saturation_high(self):
"""饱和度超过200应该被钳制."""
config = ColorGradeConfig(enabled=True, saturation=300)
params = config.resolve_params()
assert params["saturation"] == 200
def test_clamp_hue_high(self):
"""色调超过180应该被钳制."""
config = ColorGradeConfig(enabled=True, hue=270)
params = config.resolve_params()
assert params["hue"] == 180
def test_clamp_hue_low(self):
"""色调低于-180应该被钳制."""
config = ColorGradeConfig(enabled=True, hue=-270)
params = config.resolve_params()
assert params["hue"] == -180
def test_clamp_contrast(self):
"""对比度越界应该被钳制."""
config = ColorGradeConfig(enabled=True, contrast=150)
params = config.resolve_params()
assert params["contrast"] == 100
config2 = ColorGradeConfig(enabled=True, contrast=-150)
params2 = config2.resolve_params()
assert params2["contrast"] == -100
def test_clamp_temperature(self):
"""色温越界应该被钳制."""
config = ColorGradeConfig(enabled=True, temperature=150)
params = config.resolve_params()
assert params["temperature"] == 100
def test_preset_with_clamping(self):
"""预设+自定义覆盖,自定义值超范围仍需钳制."""
config = ColorGradeConfig(
enabled=True,
preset=PRESET_FRESH,
brightness=999, # 超范围
)
params = config.resolve_params()
assert params["brightness"] == 100 # 被钳制
# ── ColorGradeConfig.has_effect 测试 ──────────────────────────────────────────
class TestHasEffect:
"""是否有实际效果判断测试."""
def test_disabled_has_no_effect(self):
"""disabled的配置has_effect应该返回False."""
config = ColorGradeConfig(enabled=False)
assert not config.has_effect()
def test_default_params_no_effect(self):
"""所有参数都是默认值时应该返回False."""
config = ColorGradeConfig(enabled=True)
assert not config.has_effect()
def test_brightness_change_has_effect(self):
"""亮度变化应该有效果."""
config = ColorGradeConfig(enabled=True, brightness=10)
assert config.has_effect()
def test_saturation_100_no_effect(self):
"""饱和度100是默认值,无效果."""
config = ColorGradeConfig(enabled=True, saturation=100)
assert not config.has_effect()
def test_saturation_not_100_has_effect(self):
"""饱和度不等于100有效果."""
config = ColorGradeConfig(enabled=True, saturation=99)
assert config.has_effect()
def test_preset_has_effect(self):
"""预设通常有效果."""
for preset in PRESET_PARAMS:
config = ColorGradeConfig(enabled=True, preset=preset)
assert config.has_effect(), f"预设 {preset} 应该有效果"
def test_custom_zero_override_no_effect(self):
"""用预设但所有自定义值都设为默认值抵消 → 应该has_effect看实际值."""
# 黑白预设饱和度=0,如果手动覆盖饱和度=100、其他都=默认值,则可能无效果
config = ColorGradeConfig(
enabled=True,
preset=PRESET_BW,
brightness=0,
contrast=0,
saturation=100,
temperature=0,
hue=0,
)
assert not config.has_effect()
# ── ColorGradeEngine 参数映射测试 ─────────────────────────────────────────────
class TestParameterMapping:
"""FFmpeg参数映射测试."""
def test_brightness_mapping_zero(self):
"""亮度0 → 0.0."""
assert ColorGradeEngine._map_brightness(0) == 0.0
def test_brightness_mapping_max(self):
"""亮度100 → 1.0."""
assert ColorGradeEngine._map_brightness(100) == 1.0
def test_brightness_mapping_min(self):
"""亮度-100 → -1.0."""
assert ColorGradeEngine._map_brightness(-100) == -1.0
def test_contrast_mapping_zero(self):
"""对比度0 → 1.0(原始)."""
assert ColorGradeEngine._map_contrast(0) == 1.0
def test_contrast_mapping_positive(self):
"""正对比度应该 > 1.0."""
assert ColorGradeEngine._map_contrast(50) == 1.5
assert ColorGradeEngine._map_contrast(100) == 2.0
def test_contrast_mapping_negative(self):
"""负对比度应该 < 1.0."""
assert ColorGradeEngine._map_contrast(-50) == 0.5
assert ColorGradeEngine._map_contrast(-100) == 0.0
def test_saturation_mapping_default(self):
"""饱和度100 → 1.0."""
assert ColorGradeEngine._map_saturation(100) == 1.0
def test_saturation_mapping_zero(self):
"""饱和度0 → 0.0(黑白)."""
assert ColorGradeEngine._map_saturation(0) == 0.0
def test_saturation_mapping_double(self):
"""饱和度200 → 2.0."""
assert ColorGradeEngine._map_saturation(200) == 2.0
def test_temperature_warm(self):
"""暖色温应该红+蓝-."""
red, green, blue = ColorGradeEngine._map_temperature(100)
assert red > 0
assert blue < 0
def test_temperature_cool(self):
"""冷色温应该红-蓝+."""
red, green, blue = ColorGradeEngine._map_temperature(-100)
assert red < 0
assert blue > 0
def test_temperature_zero(self):
"""色温0应该全0."""
red, green, blue = ColorGradeEngine._map_temperature(0)
assert red == 0
assert green == 0
assert blue == 0
def test_hue_mapping_passthrough(self):
"""色调直接透传."""
assert ColorGradeEngine._map_hue(0) == 0
assert ColorGradeEngine._map_hue(90) == 90
assert ColorGradeEngine._map_hue(-45) == -45
# ── ColorGradeEngine.build_filter 测试 ────────────────────────────────────────
class TestBuildFilter:
"""滤镜字符串构建测试."""
def test_disabled_returns_empty(self):
"""disabled配置返回空."""
config = ColorGradeConfig(enabled=False)
result = ColorGradeEngine.build_filter(config)
assert result == ""
def test_no_effect_returns_empty(self):
"""无效果的配置返回空."""
config = ColorGradeConfig(enabled=True)
result = ColorGradeEngine.build_filter(config)
assert result == ""
def test_brightness_only(self):
"""只有亮度调整."""
config = ColorGradeConfig(enabled=True, brightness=20)
result = ColorGradeEngine.build_filter(config)
assert "eq=" in result
assert "brightness=" in result
assert "contrast=" not in result
assert "saturation=" not in result
def test_contrast_only(self):
"""只有对比度调整."""
config = ColorGradeConfig(enabled=True, contrast=30)
result = ColorGradeEngine.build_filter(config)
assert "eq=" in result
assert "contrast=" in result
def test_saturation_only(self):
"""只有饱和度调整."""
config = ColorGradeConfig(enabled=True, saturation=50)
result = ColorGradeEngine.build_filter(config)
assert "eq=" in result
assert "saturation=" in result
def test_temperature_only(self):
"""只有色温调整."""
config = ColorGradeConfig(enabled=True, temperature=20)
result = ColorGradeEngine.build_filter(config)
assert "colorbalance=" in result
# 暖色调应该有红通道调整
assert "rs=" in result
def test_hue_only(self):
"""只有色调调整."""
config = ColorGradeConfig(enabled=True, hue=30)
result = ColorGradeEngine.build_filter(config)
assert "hue=h=" in result
def test_with_input_output_labels(self):
"""带输入输出标签."""
config = ColorGradeConfig(enabled=True, brightness=10)
result = ColorGradeEngine.build_filter(config, input_label="[0:v]", output_label="[out]")
assert result.startswith("[0:v]")
assert result.endswith("[out]")
def test_preset_fresh_filter(self):
"""清新预设应该生成eq滤镜."""
config = ColorGradeConfig(enabled=True, preset=PRESET_FRESH)
result = ColorGradeEngine.build_filter(config)
assert "eq=" in result
# 清新预设饱和度>100,应该有saturation
assert "saturation=" in result
def test_preset_bw_filter(self):
"""黑白预设应该有saturation=0."""
config = ColorGradeConfig(enabled=True, preset=PRESET_BW)
result = ColorGradeEngine.build_filter(config)
assert "saturation=0.0" in result
def test_combined_params(self):
"""多个参数组合."""
config = ColorGradeConfig(
enabled=True,
brightness=15,
contrast=20,
saturation=130,
temperature=10,
hue=5,
)
result = ColorGradeEngine.build_filter(config)
# 应该有三个滤镜用逗号连接
assert "eq=" in result
assert "colorbalance=" in result
assert "hue=" in result
# 逗号分隔
assert "," in result
def test_filter_chain_order(self):
"""滤镜顺序应该是 eq → colorbalance → hue."""
config = ColorGradeConfig(
enabled=True,
brightness=10,
temperature=10,
hue=10,
)
result = ColorGradeEngine.build_filter(config)
eq_pos = result.find("eq=")
cb_pos = result.find("colorbalance=")
hue_pos = result.find("hue=")
assert eq_pos < cb_pos < hue_pos
def test_zero_temperature_no_colorbalance(self):
"""色温为0不应该有colorbalance滤镜."""
config = ColorGradeConfig(enabled=True, temperature=0, brightness=10)
result = ColorGradeEngine.build_filter(config)
assert "colorbalance" not in result
def test_zero_hue_no_hue_filter(self):
"""色调为0不应该有hue滤镜."""
config = ColorGradeConfig(enabled=True, hue=0, brightness=10)
result = ColorGradeEngine.build_filter(config)
assert "hue=" not in result
def test_all_presets_generate_valid_filter(self):
"""所有预设都应该能生成有效的非空滤镜."""
for preset_name in PRESET_PARAMS:
config = ColorGradeConfig(enabled=True, preset=preset_name)
result = ColorGradeEngine.build_filter(config)
assert result, f"预设 {preset_name} 应该生成非空滤镜"
# 不应该有语法错误(连续冒号、空参数等)
assert "::" not in result
assert result[0] != ":"
assert result[-1] != ":"
# ── 便捷函数测试 ──────────────────────────────────────────────────────────────
class TestHelperFunctions:
"""便捷函数测试."""
def test_get_preset_names_returns_eight(self):
"""应该返回8个预设."""
names = get_preset_names()
assert len(names) == 8
# 每个是 (key, display_name) 元组
for key, display in names:
assert key in PRESET_PARAMS
assert isinstance(display, str)
assert display
def test_get_preset_params_valid(self):
"""获取有效预设的参数."""
params = get_preset_params(PRESET_FRESH)
assert params is not None
assert params == PRESET_PARAMS[PRESET_FRESH]
def test_get_preset_params_invalid(self):
"""获取无效预设返回None."""
params = get_preset_params("nonexistent")
assert params is None
# ── 分段调色(不同clip不同滤镜)概念验证 ──────────────────────────────────────
class TestPerClipGrading:
"""分段调色概念验证 — 不同配置生成不同滤镜."""
def test_different_presets_different_filters(self):
"""不同预设应该生成不同的滤镜字符串."""
configs = [
ColorGradeConfig(enabled=True, preset=PRESET_FRESH),
ColorGradeConfig(enabled=True, preset=PRESET_VINTAGE),
ColorGradeConfig(enabled=True, preset=PRESET_BW),
]
filters = [ColorGradeEngine.build_filter(c) for c in configs]
# 三个滤镜应该各不相同
assert len(set(filters)) == 3
def test_same_preset_same_filter(self):
"""相同配置应该生成相同滤镜(确定性)."""
config1 = ColorGradeConfig(enabled=True, preset=PRESET_CINEMA)
config2 = ColorGradeConfig(enabled=True, preset=PRESET_CINEMA)
assert ColorGradeEngine.build_filter(config1) == ColorGradeEngine.build_filter(config2)
def test_custom_override_changes_filter(self):
"""自定义覆盖应该改变滤镜."""
base = ColorGradeConfig(enabled=True, preset=PRESET_FILM)
modified = ColorGradeConfig(enabled=True, preset=PRESET_FILM, brightness=50)
assert ColorGradeEngine.build_filter(base) != ColorGradeEngine.build_filter(modified)
def test_clips_with_and_without_grading(self):
"""有的clip有调色有的没有,生成结果不同."""
with_grade = ColorGradeConfig(enabled=True, preset=PRESET_WARM)
without_grade = ColorGradeConfig(enabled=False)
filter_with = ColorGradeEngine.build_filter(with_grade, "[0:v]", "[v0]")
filter_without = ColorGradeEngine.build_filter(without_grade, "[0:v]", "[v0]")
assert filter_with # 有调色应该非空
# 无调色但带标签时应该走 copy 直通(保证标签传递)
assert "[0:v]copy[v0]" in filter_without
+860
View File
@@ -0,0 +1,860 @@
"""封面生成 + 视频倒放 + 贴纸叠加 单元测试.
覆盖三个新渲染能力的核心场景和降级逻辑
"""
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.cover_generator import (
DEFAULT_COVER_HEIGHT,
DEFAULT_COVER_WIDTH,
CoverGenerator,
generate_cover_from_plan,
)
from video_processing.reverse_engine import ReverseConfig, ReverseEngine
from video_processing.sticker_engine import (
POSITION_PRESETS,
STICKER_CATEGORIES,
ImageStickerConfig,
StickerEngine,
TextStickerConfig,
get_sticker_categories,
parse_stickers_from_config,
)
from video_processing.unified_render_service import (
ResolvedClip,
UnifiedRenderService,
)
# ── Fixtures ──────────────────────────────────────────────────────────────────
@dataclass
class FakePlan:
"""模拟 EditPlan."""
id: str = "plan_001"
name: str = "测试计划"
config: dict[str, Any] = field(default_factory=dict)
@pytest.fixture
def sample_video(tmp_path):
"""创建一个测试视频文件(空文件,仅用于路径测试)."""
video_path = tmp_path / "test_video.mp4"
video_path.write_bytes(b"fake video data")
return video_path
@pytest.fixture
def sample_image(tmp_path):
"""创建一个测试图片文件."""
img_path = tmp_path / "sticker.png"
img_path.write_bytes(b"fake png data")
return img_path
# ═══════════════════════════════════════════════════════════════════════════════
# 一、视频倒放引擎测试
# ═══════════════════════════════════════════════════════════════════════════════
class TestReverseConfig:
"""ReverseConfig 配置解析测试."""
def test_default_disabled(self):
"""默认配置为关闭."""
config = ReverseConfig.from_dict(None)
assert config.enabled is False
assert config.reverse_video is True
assert config.reverse_audio is True
def test_empty_dict(self):
"""空字典视为关闭."""
config = ReverseConfig.from_dict({})
assert config.enabled is False
def test_enabled(self):
"""启用倒放."""
config = ReverseConfig.from_dict({"enabled": True})
assert config.enabled is True
assert config.reverse_video is True
assert config.reverse_audio is True
def test_video_only(self):
"""只倒放视频."""
config = ReverseConfig.from_dict(
{
"enabled": True,
"reverse_video": True,
"reverse_audio": False,
}
)
assert config.enabled is True
assert config.reverse_video is True
assert config.reverse_audio is False
def test_audio_only(self):
"""只倒放音频."""
config = ReverseConfig.from_dict(
{
"enabled": True,
"reverse_video": False,
"reverse_audio": True,
}
)
assert config.reverse_video is False
assert config.reverse_audio is True
def test_invalid_config_fallback(self):
"""无效配置降级为默认."""
config = ReverseConfig.from_dict("invalid") # type: ignore
assert config.enabled is False
def test_none_config(self):
"""None 配置."""
config = ReverseConfig.from_dict(None)
assert config.enabled is False
class TestReverseEngine:
"""ReverseEngine 滤镜生成测试."""
def test_video_reverse_filter(self):
"""视频倒放滤镜生成."""
config = ReverseConfig(enabled=True, reverse_video=True)
f = ReverseEngine.build_video_filter(config, duration=10.0)
assert f == "reverse"
def test_video_disabled(self):
"""视频倒放关闭时返回空."""
config = ReverseConfig(enabled=False)
f = ReverseEngine.build_video_filter(config, duration=10.0)
assert f == ""
def test_video_disabled_flag(self):
"""启用但 reverse_video=False."""
config = ReverseConfig(enabled=True, reverse_video=False)
f = ReverseEngine.build_video_filter(config, duration=10.0)
assert f == ""
def test_audio_reverse_filter(self):
"""音频倒放滤镜生成."""
config = ReverseConfig(enabled=True, reverse_audio=True)
f = ReverseEngine.build_audio_filter(config, duration=10.0)
assert f == "areverse"
def test_audio_disabled(self):
"""音频倒放关闭."""
config = ReverseConfig(enabled=False)
f = ReverseEngine.build_audio_filter(config, duration=10.0)
assert f == ""
def test_long_video_safety_limit(self):
"""超长视频安全限制:跳过倒放."""
config = ReverseConfig(enabled=True)
f = ReverseEngine.build_video_filter(config, duration=200.0)
assert f == "" # 超过 MAX_SAFE_DURATION
def test_long_audio_safety_limit(self):
"""超长音频安全限制."""
config = ReverseConfig(enabled=True)
f = ReverseEngine.build_audio_filter(config, duration=200.0)
assert f == ""
def test_duration_zero(self):
"""时长为0时正常返回."""
config = ReverseConfig(enabled=True)
f = ReverseEngine.build_video_filter(config, duration=0.0)
assert f == "reverse"
# ═══════════════════════════════════════════════════════════════════════════════
# 二、贴纸引擎测试
# ═══════════════════════════════════════════════════════════════════════════════
class TestStickerPosition:
"""贴纸位置计算测试."""
def test_presets_exist(self):
"""9宫格预设存在."""
assert "top_left" in POSITION_PRESETS
assert "center" in POSITION_PRESETS
assert "bottom_right" in POSITION_PRESETS
assert len(POSITION_PRESETS) == 9
def test_resolve_position_center(self):
"""居中位置计算."""
sticker = ImageStickerConfig(position="center")
x, y = StickerEngine._resolve_position(sticker, 1000, 1000, 200, 200)
assert abs(x - 400) < 1 # (1000-200)/2 = 400
assert abs(y - 400) < 1
def test_resolve_position_top_left(self):
"""左上角位置."""
sticker = ImageStickerConfig(position="top_left")
x, y = StickerEngine._resolve_position(sticker, 1000, 1000, 100, 100)
assert x == 0 # 0.05*1000 - 50 = 0 (clamped)
assert y == 0
def test_custom_position_percent(self):
"""自定义百分比位置."""
sticker = ImageStickerConfig(
position="center",
x=30.0,
y=70.0,
x_unit="percent",
y_unit="percent",
)
x, y = StickerEngine._resolve_position(sticker, 1000, 1000, 100, 100)
assert abs(x - 250) < 1 # 300 - 50 = 250
assert abs(y - 650) < 1 # 700 - 50 = 650
def test_custom_position_pixel(self):
"""自定义像素位置."""
sticker = ImageStickerConfig(
position="center",
x=100.0,
y=200.0,
x_unit="pixel",
y_unit="pixel",
)
x, y = StickerEngine._resolve_position(sticker, 1000, 1000, 50, 50)
assert abs(x - 75) < 1 # 100 - 25 = 75
assert abs(y - 175) < 1 # 200 - 25 = 175
def test_position_clamped(self):
"""位置钳制在画布内."""
sticker = ImageStickerConfig(
position="center",
x=-10.0,
y=-10.0,
x_unit="pixel",
y_unit="pixel",
)
x, y = StickerEngine._resolve_position(sticker, 1000, 1000, 50, 50)
assert x >= 0
assert y >= 0
class TestTextSticker:
"""文字贴纸测试."""
def test_drawtext_filter_basic(self):
"""基础文字贴纸滤镜生成."""
sticker = TextStickerConfig(
enabled=True,
text="Hello World",
font_size=36,
font_color="#FFFFFF",
position="center",
)
f = StickerEngine._build_drawtext_filter(sticker, "[in]", "[out]", 1080, 1920)
assert "drawtext" in f
assert "Hello World" in f
assert "fontsize=36" in f
assert "[in]" in f
assert "[out]" in f
def test_drawtext_with_stroke(self):
"""带描边的文字贴纸."""
sticker = TextStickerConfig(
enabled=True,
text="Test",
stroke_width=3,
stroke_color="#FF0000",
)
f = StickerEngine._build_drawtext_filter(sticker, "[in]", "[out]", 1080, 1920)
assert "borderw=3" in f
assert "bordercolor=#FF0000" in f
def test_drawtext_with_shadow(self):
"""带阴影的文字贴纸."""
sticker = TextStickerConfig(
enabled=True,
text="Shadow",
shadow_x=4,
shadow_y=4,
shadow_alpha=0.5,
)
f = StickerEngine._build_drawtext_filter(sticker, "[in]", "[out]", 1080, 1920)
assert "shadowx=4" in f
assert "shadowy=4" in f
def test_drawtext_time_range(self):
"""带时间范围的文字贴纸."""
sticker = TextStickerConfig(
enabled=True,
text="Timed",
start_time=2.0,
duration=3.0,
)
f = StickerEngine._build_drawtext_filter(sticker, "[in]", "[out]", 1080, 1920)
assert "enable='between(t,2.0,5.0)'" in f
def test_drawtext_empty_text(self):
"""空文字直通."""
sticker = TextStickerConfig(enabled=True, text="")
f = StickerEngine._build_drawtext_filter(sticker, "[in]", "[out]", 1080, 1920)
assert "[in]copy[out]" in f
def test_drawtext_with_fade(self):
"""带淡入淡出的文字贴纸."""
sticker = TextStickerConfig(
enabled=True,
text="Fade",
start_time=1.0,
duration=5.0,
fade_in=0.5,
fade_out=0.5,
)
f = StickerEngine._build_drawtext_filter(sticker, "[in]", "[out]", 1080, 1920)
assert "alpha=" in f
class TestImageSticker:
"""图片贴纸测试."""
def test_image_sticker_overlay(self, sample_image):
"""图片贴纸 overlay 滤镜生成."""
result = StickerEngine.build_sticker_chain(
stickers=[
{
"type": "image",
"image_path": str(sample_image),
"position": "top_right",
"scale": 0.5,
"opacity": 0.8,
"z_index": 10,
}
],
input_label="[base]",
output_label="[final]",
canvas_w=1080,
canvas_h=1920,
)
assert result.filter_str != ""
assert "overlay" in result.filter_str
assert len(result.extra_inputs) == 1
assert result.extra_inputs[0] == str(sample_image)
def test_image_sticker_missing_file(self):
"""图片贴纸素材不存在时跳过."""
result = StickerEngine.build_sticker_chain(
stickers=[
{
"type": "image",
"image_path": "/nonexistent/image.png",
"position": "center",
}
],
input_label="[in]",
output_label="[out]",
canvas_w=1080,
canvas_h=1920,
)
# 素材不存在,跳过,返回直通
assert "[in]copy[out]" in result.filter_str
assert len(result.extra_inputs) == 0
def test_mixed_stickers(self, sample_image):
"""混合贴纸:图片 + 文字."""
result = StickerEngine.build_sticker_chain(
stickers=[
{
"type": "image",
"image_path": str(sample_image),
"position": "top_left",
"z_index": 5,
},
{
"type": "text",
"text": "Hello",
"position": "bottom_center",
"z_index": 10,
},
],
input_label="[in]",
output_label="[out]",
canvas_w=1080,
canvas_h=1920,
)
assert "overlay" in result.filter_str
assert "drawtext" in result.filter_str
assert len(result.extra_inputs) == 1
def test_sticker_z_index_order(self, sample_image):
"""贴纸按 z_index 排序."""
result = StickerEngine.build_sticker_chain(
stickers=[
{"type": "text", "text": "Top", "z_index": 20, "position": "center"},
{"type": "text", "text": "Bottom", "z_index": 5, "position": "center"},
],
input_label="[in]",
output_label="[out]",
canvas_w=1080,
canvas_h=1920,
)
# z_index 小的先叠加,大的后叠加(在上面)
assert result.filter_str.count("drawtext") == 2
def test_empty_stickers(self):
"""空贴纸列表."""
result = StickerEngine.build_sticker_chain(
stickers=[],
input_label="[in]",
output_label="[out]",
canvas_w=1080,
canvas_h=1920,
)
assert "[in]copy[out]" in result.filter_str
assert result.extra_inputs == []
def test_invalid_sticker_skipped(self):
"""无效贴纸配置跳过."""
result = StickerEngine.build_sticker_chain(
stickers=[{"invalid": "data"}],
input_label="[in]",
output_label="[out]",
canvas_w=1080,
canvas_h=1920,
)
# 解析失败,跳过,直通
assert "[in]copy[out]" in result.filter_str
class TestStickerHelpers:
"""贴纸辅助函数测试."""
def test_parse_stickers_empty(self):
"""空配置解析."""
assert parse_stickers_from_config(None) == []
assert parse_stickers_from_config({}) == []
def test_parse_stickers_list(self):
"""正常贴纸列表解析."""
config = {"stickers": [{"type": "text", "text": "A"}, {"type": "text", "text": "B"}]}
result = parse_stickers_from_config(config)
assert len(result) == 2
def test_parse_stickers_not_list(self):
"""非列表类型返回空."""
config = {"stickers": "not a list"}
assert parse_stickers_from_config(config) == []
def test_get_categories(self):
"""贴纸分类列表."""
cats = get_sticker_categories()
assert len(cats) == len(STICKER_CATEGORIES)
assert cats[0][0] == "emoji"
# ═══════════════════════════════════════════════════════════════════════════════
# 三、封面生成器测试
# ═══════════════════════════════════════════════════════════════════════════════
class TestCoverGenerator:
"""CoverGenerator 测试."""
def test_default_dimensions(self):
"""默认封面尺寸."""
assert DEFAULT_COVER_WIDTH == 1080
assert DEFAULT_COVER_HEIGHT == 1920
@patch("video_processing.cover_generator.run_ffmpeg")
@patch("video_processing.cover_generator.probe_video_info")
def test_extract_frame_basic(self, mock_probe, mock_run, sample_video, tmp_path):
"""基础抽帧测试."""
mock_probe.return_value = {"duration": 30.0}
# mock run_ffmpeg 实际创建输出文件
def fake_run_ffmpeg(cmd):
# 找到输出路径并创建文件
output_path = Path(cmd[-1])
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_bytes(b"fake jpeg data")
mock_run.side_effect = fake_run_ffmpeg
output = tmp_path / "cover.jpg"
result = CoverGenerator.extract_frame(
sample_video,
output,
time_sec=2.0,
)
assert result == output
mock_run.assert_called_once()
# 检查命令参数
cmd = mock_run.call_args[0][0]
assert "-ss" in cmd
assert "2.000" in cmd
assert "-vframes" in cmd
assert "1" in cmd
@patch("video_processing.cover_generator.run_ffmpeg")
@patch("video_processing.cover_generator.probe_video_info")
def test_extract_frame_time_clamped(self, mock_probe, mock_run, sample_video, tmp_path):
"""抽帧时间超过视频长度时钳制."""
mock_probe.return_value = {"duration": 10.0}
def fake_run_ffmpeg(cmd):
output_path = Path(cmd[-1])
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_bytes(b"fake jpeg data")
mock_run.side_effect = fake_run_ffmpeg
output = tmp_path / "cover.jpg"
CoverGenerator.extract_frame(
sample_video,
output,
time_sec=100.0, # 超过视频时长
)
cmd = mock_run.call_args[0][0]
ss_idx = cmd.index("-ss")
time_val = float(cmd[ss_idx + 1])
# 应该被钳制到中间帧(5秒左右)
assert time_val <= 10.0
@patch("video_processing.cover_generator.run_ffmpeg")
@patch("video_processing.cover_generator.probe_video_info")
def test_extract_frame_negative_time(self, mock_probe, mock_run, sample_video, tmp_path):
"""负时间钳制到0."""
mock_probe.return_value = {"duration": 30.0}
def fake_run_ffmpeg(cmd):
output_path = Path(cmd[-1])
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_bytes(b"fake jpeg data")
mock_run.side_effect = fake_run_ffmpeg
output = tmp_path / "cover.jpg"
CoverGenerator.extract_frame(
sample_video,
output,
time_sec=-5.0,
)
cmd = mock_run.call_args[0][0]
ss_idx = cmd.index("-ss")
time_val = float(cmd[ss_idx + 1])
assert time_val >= 0
def test_extract_frame_file_not_found(self, tmp_path):
"""视频文件不存在抛异常."""
with pytest.raises(FileNotFoundError):
CoverGenerator.extract_frame(
"/nonexistent/video.mp4",
tmp_path / "cover.jpg",
)
@patch("video_processing.cover_generator.CoverGenerator.extract_frame")
@patch("video_processing.cover_generator.probe_video_info")
def test_smart_cover_3_frames(self, mock_probe, mock_extract, sample_video, tmp_path):
"""智能封面抽取3帧选最佳."""
mock_probe.return_value = {"duration": 30.0}
# 创建三个大小不同的临时文件(模拟清晰度不同)
def create_frame(video_path, output_path, **kwargs):
# 第二帧最大(最清晰)
p = Path(output_path)
p.parent.mkdir(parents=True, exist_ok=True)
if "candidate_1" in str(p):
p.write_bytes(b"x" * 10000) # 最大 = 最清晰
elif "candidate_0" in str(p):
p.write_bytes(b"x" * 1000)
else:
p.write_bytes(b"x" * 5000)
return p
mock_extract.side_effect = create_frame
output = tmp_path / "smart_cover.jpg"
result = CoverGenerator.extract_smart_cover(
sample_video,
output,
frame_count=3,
)
assert result == output
assert output.exists()
# 应该选最大的那个文件(candidate_1
assert output.stat().st_size == 10000
@patch("video_processing.cover_generator.CoverGenerator.extract_frame")
@patch("video_processing.cover_generator.probe_video_info")
def test_smart_cover_fallback(self, mock_probe, mock_extract, sample_video, tmp_path):
"""智能封面全部失败时降级."""
mock_probe.return_value = {"duration": 0.0} # 时长为0
output = tmp_path / "cover.jpg"
output.write_bytes(b"x" * 100)
mock_extract.return_value = output
result = CoverGenerator.extract_smart_cover(sample_video, output, frame_count=3)
assert result == output
@patch("video_processing.cover_generator.run_ffmpeg")
def test_custom_cover(self, mock_run, sample_image, tmp_path):
"""自定义封面处理."""
output = tmp_path / "custom_cover.jpg"
result = CoverGenerator.process_custom_cover(
sample_image,
output,
)
assert result == output
mock_run.assert_called_once()
cmd = mock_run.call_args[0][0]
assert str(sample_image) in cmd
def test_custom_cover_not_found(self, tmp_path):
"""自定义封面文件不存在."""
with pytest.raises(FileNotFoundError):
CoverGenerator.process_custom_cover(
"/nonexistent/img.png",
tmp_path / "cover.jpg",
)
@patch("video_processing.cover_generator.CoverGenerator.extract_frame")
def test_generate_cover_time_mode(self, mock_extract, sample_video, tmp_path):
"""统一入口 - time 模式."""
output = tmp_path / "cover.jpg"
mock_extract.return_value = output
result = CoverGenerator.generate_cover(
sample_video,
output,
mode="time",
time_sec=3.0,
)
assert result == output
mock_extract.assert_called_once()
@patch("video_processing.cover_generator.CoverGenerator.extract_smart_cover")
def test_generate_cover_smart_mode(self, mock_smart, sample_video, tmp_path):
"""统一入口 - smart 模式."""
output = tmp_path / "cover.jpg"
mock_smart.return_value = output
result = CoverGenerator.generate_cover(
sample_video,
output,
mode="smart",
)
assert result == output
mock_smart.assert_called_once()
@patch("video_processing.cover_generator.CoverGenerator.process_custom_cover")
def test_generate_cover_custom_mode(self, mock_custom, sample_video, sample_image, tmp_path):
"""统一入口 - custom 模式."""
output = tmp_path / "cover.jpg"
mock_custom.return_value = output
result = CoverGenerator.generate_cover(
sample_video,
output,
mode="custom",
custom_image=sample_image,
)
assert result == output
mock_custom.assert_called_once()
class TestGenerateCoverFromPlan:
"""从 plan 配置生成封面测试."""
@patch("video_processing.cover_generator.CoverGenerator.extract_smart_cover")
def test_smart_mode_from_plan(self, mock_smart, sample_video, tmp_path):
"""plan 配置 smart 模式."""
plan = FakePlan(id="plan_001", config={"cover_config": {"mode": "smart"}})
mock_smart.return_value = tmp_path / "cover.jpg"
(tmp_path / "cover.jpg").write_bytes(b"test")
result = generate_cover_from_plan(plan, sample_video, tmp_path)
assert result is not None
def test_no_cover_config(self, sample_video, tmp_path):
"""没有封面配置时返回 None."""
plan = FakePlan(id="plan_001", config={})
result = generate_cover_from_plan(plan, sample_video, tmp_path)
assert result is None
def test_none_config(self, sample_video, tmp_path):
"""config 为 None."""
plan = FakePlan(id="plan_001", config=None) # type: ignore
result = generate_cover_from_plan(plan, sample_video, tmp_path)
assert result is None
# ═══════════════════════════════════════════════════════════════════════════════
# 四、UnifiedRenderService 集成测试
# ═══════════════════════════════════════════════════════════════════════════════
def _make_clip(clip_id="c1", asset_id="a1", path=Path("/fake/video.mp4"), clip_type="main", config=None):
"""创建测试用 ResolvedClip."""
return ResolvedClip(
clip_id=clip_id,
asset_id=asset_id,
local_path=path,
clip_type=clip_type,
order=0,
start_time=0.0,
duration=0.0,
transition_effect="cut",
config=config or {},
actual_duration=10.0,
)
def _make_service(plan, clips, asset_path_map=None, work_dir=None, tmp_path=None):
"""创建测试用 UnifiedRenderService."""
from pathlib import Path as P
work_dir = work_dir or (tmp_path or P("/tmp")) / "render_test"
work_dir.mkdir(exist_ok=True, parents=True)
return UnifiedRenderService(
plan=plan,
clips=clips,
asset_path_map=asset_path_map or {},
work_dir=work_dir,
output_width=1080,
output_height=1920,
output_fps=30,
transition_duration=0.5,
)
class TestReverseIntegration:
"""倒放功能集成测试."""
@patch("video_processing.unified_render_service.probe_video_info")
@patch("video_processing.unified_render_service.run_ffmpeg")
def test_reverse_in_filter_complex(self, mock_run, mock_probe, tmp_path):
"""filter_complex 路径中包含倒放滤镜."""
mock_probe.return_value = {"duration": 10.0, "has_audio": True, "width": 1920, "height": 1080}
mock_run.return_value = None
plan = FakePlan(id="p1")
clip = _make_clip(config={"reverse": {"enabled": True}})
clip.actual_duration = 5.0
# 两个 clip 触发 filter_complex 路径
clip2 = _make_clip(clip_id="c2", config={})
clip2.actual_duration = 5.0
clip2.order = 1
service = _make_service(plan, [clip, clip2], tmp_path=tmp_path)
# 直接测 _build_filter_complex
from video_processing.unified_render_service import RenderLayer
layer = RenderLayer(role="main", clips=[clip, clip2])
filter_str, inputs = service._build_filter_complex([layer])
assert "reverse" in filter_str
def test_can_use_pass_through_with_reverse(self, tmp_path):
"""倒放不影响直通模式判断(只有贴纸才禁用)."""
plan = FakePlan(id="p1")
clip = _make_clip(config={"reverse": {"enabled": True}})
clip.actual_duration = 5.0
service = _make_service(plan, [clip], tmp_path=tmp_path)
from video_processing.unified_render_service import RenderLayer
layer = RenderLayer(role="main", clips=[clip])
layers = [layer]
assert service._can_use_pass_through(layers) is True
class TestStickerIntegration:
"""贴纸功能集成测试."""
def test_can_use_pass_through_with_stickers(self, tmp_path):
"""有贴纸时禁用直通模式."""
plan = FakePlan(id="p1", config={"stickers": [{"type": "text", "text": "Hello", "position": "center"}]})
clip = _make_clip()
clip.actual_duration = 5.0
service = _make_service(plan, [clip], tmp_path=tmp_path)
from video_processing.unified_render_service import RenderLayer
layer = RenderLayer(role="main", clips=[clip])
layers = [layer]
assert service._can_use_pass_through(layers) is False
def test_can_use_pass_through_no_stickers(self, tmp_path):
"""无贴纸时直通模式正常."""
plan = FakePlan(id="p1", config={})
clip = _make_clip()
clip.actual_duration = 5.0
service = _make_service(plan, [clip], tmp_path=tmp_path)
from video_processing.unified_render_service import RenderLayer
layer = RenderLayer(role="main", clips=[clip])
layers = [layer]
assert service._can_use_pass_through(layers) is True
def test_build_sticker_filters_text(self, tmp_path):
"""文字贴纸滤镜构建."""
plan = FakePlan(
id="p1", config={"stickers": [{"type": "text", "text": "Hello", "position": "top_center", "z_index": 10}]}
)
service = _make_service(plan, [], tmp_path=tmp_path)
filter_str, extra_inputs = service._build_sticker_filters("in", "out")
assert "drawtext" in filter_str
assert len(extra_inputs) == 0
def test_build_sticker_filters_empty(self, tmp_path):
"""无贴纸返回空."""
plan = FakePlan(id="p1", config={})
service = _make_service(plan, [], tmp_path=tmp_path)
filter_str, extra_inputs = service._build_sticker_filters("in", "out")
assert filter_str == ""
assert extra_inputs == []
def test_build_sticker_filters_image(self, sample_image, tmp_path):
"""图片贴纸滤镜构建 + 额外输入."""
plan = FakePlan(
id="p1",
config={
"stickers": [
{
"type": "image",
"image_path": str(sample_image),
"position": "bottom_right",
"z_index": 5,
}
]
},
)
service = _make_service(plan, [], tmp_path=tmp_path)
filter_str, extra_inputs = service._build_sticker_filters("in", "out")
assert "overlay" in filter_str
assert len(extra_inputs) == 1
@@ -193,6 +193,70 @@ class StubGenerationTaskRepository:
items.sort(key=lambda t: t.created_at, reverse=True)
return items[:limit]
def list_by_user_filtered(
self,
user_id: str,
*,
status: str | None = None,
limit: int | None = None,
offset: int = 0,
) -> list:
"""按用户+状态筛选任务列表(stub实现)。"""
items = [t for t in self._store.values() if t.created_by_user_id == user_id]
if status:
items = [t for t in items if str(t.status) == status]
# 按创建时间倒序
items.sort(key=lambda t: t.created_at or "", reverse=True)
if offset:
items = items[offset:]
if limit is not None:
items = items[:limit]
return items
def count_by_user_filtered(
self,
user_id: str,
*,
status: str | None = None,
) -> int:
"""按用户+状态筛选计数(stub实现)。"""
items = [t for t in self._store.values() if t.created_by_user_id == user_id]
if status:
items = [t for t in items if str(t.status) == status]
return len(items)
def list_by_project_filtered(
self,
project_id: str,
*,
status: str | None = None,
limit: int | None = None,
offset: int = 0,
) -> list:
"""按项目+状态筛选任务列表(stub实现)。"""
items = [t for t in self._store.values() if t.project_id == project_id]
if status:
items = [t for t in items if str(t.status) == status]
# 按创建时间倒序
items.sort(key=lambda t: t.created_at or "", reverse=True)
if offset:
items = items[offset:]
if limit is not None:
items = items[:limit]
return items
def count_by_project_filtered(
self,
project_id: str,
*,
status: str | None = None,
) -> int:
"""按项目+状态筛选计数(stub实现)。"""
items = [t for t in self._store.values() if t.project_id == project_id]
if status:
items = [t for t in items if str(t.status) == status]
return len(items)
# ── Fixtures ──────────────────────────────────────────────────────────────────
+64
View File
@@ -200,6 +200,70 @@ class StubGenerationTaskRepository:
def count_pending_total(self) -> int:
return 0
def list_by_user_filtered(
self,
user_id: str,
*,
status: str | None = None,
limit: int | None = None,
offset: int = 0,
) -> list:
"""按用户+状态筛选任务列表(stub实现)。"""
items = [t for t in self._tasks.values() if t.created_by_user_id == user_id]
if status:
items = [t for t in items if str(t.status) == status]
# 按创建时间倒序
items.sort(key=lambda t: t.created_at or "", reverse=True)
if offset:
items = items[offset:]
if limit is not None:
items = items[:limit]
return items
def count_by_user_filtered(
self,
user_id: str,
*,
status: str | None = None,
) -> int:
"""按用户+状态筛选计数(stub实现)。"""
items = [t for t in self._tasks.values() if t.created_by_user_id == user_id]
if status:
items = [t for t in items if str(t.status) == status]
return len(items)
def list_by_project_filtered(
self,
project_id: str,
*,
status: str | None = None,
limit: int | None = None,
offset: int = 0,
) -> list:
"""按项目+状态筛选任务列表(stub实现)。"""
items = [t for t in self._tasks.values() if t.project_id == project_id]
if status:
items = [t for t in items if str(t.status) == status]
# 按创建时间倒序
items.sort(key=lambda t: t.created_at or "", reverse=True)
if offset:
items = items[offset:]
if limit is not None:
items = items[:limit]
return items
def count_by_project_filtered(
self,
project_id: str,
*,
status: str | None = None,
) -> int:
"""按项目+状态筛选计数(stub实现)。"""
items = [t for t in self._tasks.values() if t.project_id == project_id]
if status:
items = [t for t in items if str(t.status) == status]
return len(items)
# ---------------------------------------------------------------------------
# Service factory
@@ -64,6 +64,38 @@ class StubGenerationTaskRepository:
def count_pending_total(self):
return 0
def list_by_user_filtered(self, user_id, *, status=None, limit=None, offset=0):
items = [t for t in self._tasks.values() if getattr(t, "created_by_user_id", None) == user_id]
if status:
items = [t for t in items if getattr(t, "status", None) == status]
if offset:
items = items[offset:]
if limit is not None:
items = items[:limit]
return items
def count_by_user_filtered(self, user_id, *, status=None):
items = [t for t in self._tasks.values() if getattr(t, "created_by_user_id", None) == user_id]
if status:
items = [t for t in items if getattr(t, "status", None) == status]
return len(items)
def list_by_project_filtered(self, project_id, *, status=None, limit=None, offset=0):
items = [t for t in self._tasks.values() if getattr(t, "project_id", None) == project_id]
if status:
items = [t for t in items if getattr(t, "status", None) == status]
if offset:
items = items[offset:]
if limit is not None:
items = items[:limit]
return items
def count_by_project_filtered(self, project_id, *, status=None):
items = [t for t in self._tasks.values() if getattr(t, "project_id", None) == project_id]
if status:
items = [t for t in items if getattr(t, "status", None) == status]
return len(items)
class StubGeneratedVideoRepository:
def __init__(self, videos=None):

Some files were not shown because too many files have changed in this diff Show More