Compare commits

...

15 Commits

Author SHA1 Message Date
xiaoxia-bot 7dbff690cd feat: 转场特效引擎 — TransitionEngine + 14种转场预设 + 时长边界校验 + 降级策略
CI/CD Pipeline / Unit Tests (pull_request) Successful in 1m14s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 1m23s
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Successful in 3m11s
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) Failing after 1m32s
2026-07-14 10:59:17 +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
95 changed files with 13563 additions and 193 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")
+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)
+2
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,7 @@ class _PlanClipItem(BaseModel):
start_time: float
duration: float
transition_effect: str
transition_duration: float
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
+6
View File
@@ -281,6 +281,7 @@ class EditPlanService:
start_time: float = 0.0,
duration: float = 0.0,
transition_effect: str = "cut",
transition_duration: float = 0.0,
config: Optional[dict[str, Any]] = None,
) -> EditPlanClip:
"""创建片段
@@ -301,6 +302,7 @@ class EditPlanService:
start_time=start_time,
duration=duration,
transition_effect=transition_effect,
transition_duration=transition_duration,
config=config,
)
created = self._clip_repo.create(clip)
@@ -324,6 +326,7 @@ class EditPlanService:
start_time: Optional[float] = None,
duration: Optional[float] = None,
transition_effect: Optional[str] = None,
transition_duration: Optional[float] = None,
config: Optional[dict[str, Any]] = None,
) -> EditPlanClip:
"""更新片段
@@ -346,6 +349,9 @@ 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
),
status=existing.status,
config=config if config is not None else existing.config,
created_at=existing.created_at,
+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)
+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
+94 -7
View File
@@ -35,6 +35,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 +72,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 +84,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 +124,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 +141,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,6 +223,7 @@ def concat_main_audio(
# 单 clip,直接提取音频,截断到 min(clip有效时长, 视频总时长)
clip = clips[0]
effective_duration = clip_effective_duration(clip)
trim_start = getattr(clip, "start_time", 0) or 0
# 最终时长:取 clip 有效时长和视频总时长的较小值
# (视频总时长由主图层决定,但单 clip 场景下两者应该一致,仍做保护)
final_duration = effective_duration
@@ -162,6 +241,8 @@ def concat_main_audio(
"-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))
@@ -175,8 +256,11 @@ def concat_main_audio(
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
if effective_duration > 0:
filter_parts.append(f"[{i}:a]atrim=0:{effective_duration:.3f},asetpts=PTS-STARTPTS[a{i}]")
filter_parts.append(
f"[{i}:a]atrim=start={trim_start:.3f}:duration={effective_duration:.3f}," f"asetpts=PTS-STARTPTS[a{i}]"
)
else:
filter_parts.append(f"[{i}:a]asetpts=PTS-STARTPTS[a{i}]")
@@ -236,9 +320,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 +341,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}")
+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
@@ -28,19 +28,29 @@ from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
from video_processing.chroma_key_engine import apply_chroma_key_if_needed
from video_processing.color_grade_engine import ColorGradeConfig, ColorGradeEngine
from video_processing.ffmpeg_utils import (
DEFAULT_FPS,
DEFAULT_OUTPUT_HEIGHT,
DEFAULT_OUTPUT_WIDTH,
DEFAULT_TRANSITION_DURATION,
FFMPEG_BIN,
build_xfade_filter_chain,
probe_duration,
probe_video_info,
run_ffmpeg,
)
from video_processing.intro_outro_engine import IntroOutroConfig, IntroOutroEngine
from video_processing.pip_engine import PiPConfig, PiPEngine, PiPLayerConfig
from video_processing.render_audio import RenderContext, merge_audio_video, mix_audio
from video_processing.render_subtitles import generate_ass_subtitles
from video_processing.subtitle_generator import generate_ass_from_timeline
from video_processing.transition_engine import TransitionEngine
from video_processing.trim_engine import TrimConfig, TrimEngine, extract_trim_from_clip_config
from video_processing.tts_engine import TtsEngine
from video_processing.watermark_engine import WatermarkConfig, WatermarkEngine
from packages.domain.tts_config import TtsConfig
logger = logging.getLogger(__name__)
@@ -60,10 +70,12 @@ class ResolvedClip:
start_time: float = 0.0
duration: float = 0.0 # 0 表示使用素材完整时长
transition_effect: str = "cut"
transition_duration: float = 0.0 # 0 表示使用全局默认值
config: dict[str, Any] = field(default_factory=dict)
# 运行时填充
actual_duration: float = 0.0 # 素材实际时长(probe 后填充)
trim_config: TrimConfig | None = None # 解析后的裁剪配置(运行时填充)
@dataclass
@@ -158,6 +170,8 @@ class UnifiedRenderService:
output_height: int = DEFAULT_OUTPUT_HEIGHT,
output_fps: int = DEFAULT_FPS,
transition_duration: float = DEFAULT_TRANSITION_DURATION,
asr_service: Any = None, # ASRService 实例,用于自动生成字幕
bgm_path: str | None = None, # BGM 本地文件路径
):
self.plan = plan
self.clips = clips
@@ -167,6 +181,9 @@ class UnifiedRenderService:
self.output_height = output_height
self.output_fps = output_fps
self.transition_duration = transition_duration
self.asr_service = asr_service
self.bgm_path = bgm_path
self._transition_engine = TransitionEngine(default_duration=transition_duration)
def render(self) -> RenderResult:
"""执行渲染,返回 RenderResult.
@@ -200,18 +217,27 @@ class UnifiedRenderService:
# 3. 计算视频总时长(用于字幕显示时长)
video_duration = self._estimate_total_duration(layers)
# 3.5 TTS 配音生成(如果配置了)
self._maybe_add_voiceover_layer(layers, video_duration=video_duration)
# 4. 生成 ASS 字幕文件(如果有 title/subtitle 配置)
ass_path = self._maybe_generate_ass(video_duration)
# 4.5 解析画中画配置
pip_config = PiPConfig.from_dict((self.plan.config or {}).get("pip_config"))
pip_sources = self._resolve_pip_sources(pip_config) if pip_config.enabled else []
has_pip = len(pip_sources) > 0
# 灰度埋点:开始渲染
layer_roles = [layer.role for layer in layers]
clip_counts = {layer.role: len(layer.clips) for layer in layers}
logger.info(
"[unified-render] start render: plan_id=%s clip_count=%d layers=%s clip_counts=%s",
"[unified-render] start render: plan_id=%s clip_count=%d layers=%s clip_counts=%s pip_layers=%d",
self.plan.id,
len(resolved),
layer_roles,
clip_counts,
len(pip_sources),
)
# 5. 视频主渲染
@@ -219,7 +245,8 @@ class UnifiedRenderService:
video_only_path = self.work_dir / f"rendered_{self.plan.id}_video.mp4"
output_path = self.work_dir / f"rendered_{self.plan.id}.mp4"
is_pass_through = self._can_use_pass_through(layers)
# 有画中画时不走直通(需要额外图层叠加)
is_pass_through = self._can_use_pass_through(layers) and not has_pip
pass_through_has_audio = False
used_stream_copy = False
@@ -245,6 +272,11 @@ class UnifiedRenderService:
)
else:
filter_complex, input_args = self._build_filter_complex(layers, ass_path=ass_path)
# 追加画中画滤镜
if has_pip:
filter_complex, input_args = self._append_pip_filters(filter_complex, input_args, pip_sources)
self._execute_ffmpeg(filter_complex, input_args, video_only_path)
t_video_end = time.time()
@@ -265,9 +297,60 @@ class UnifiedRenderService:
if is_pass_through:
# 直通场景已在一次调用中完成视频+音频
has_audio = pass_through_has_audio
# 直通模式下也支持 BGM 混音:提取音频 → 混 BGM → 合并回视频
if self.bgm_path and pass_through_has_audio:
config = self.plan.config or {}
bgm_config = config.get("bgm", {}) or {}
if bgm_config.get("enabled", False):
ctx = RenderContext(work_dir=self.work_dir, plan_id=self.plan.id)
from video_processing.bgm_mixer import BGMConfig, mix_bgm_with_main
bgm_cfg = BGMConfig.from_config_dict(self.bgm_path, bgm_config)
# 从直通输出中提取音频
main_audio_path = self.work_dir / f"pass_through_audio_{self.plan.id}.aac"
extract_cmd = [
FFMPEG_BIN,
"-y",
"-i",
str(output_path),
"-vn",
"-acodec",
"aac",
"-b:a",
"128k",
str(main_audio_path),
]
try:
from video_processing.ffmpeg_utils import run_ffmpeg
run_ffmpeg(extract_cmd)
final_audio = mix_bgm_with_main(ctx, main_audio_path, bgm_cfg, video_duration)
# 合并回视频
bgm_output = self.work_dir / f"rendered_{self.plan.id}_bgm.mp4"
merge_audio_video(ctx, output_path, final_audio, bgm_output)
output_path = bgm_output
logger.info("[unified-render] pass-through BGM mix done: plan_id=%s", self.plan.id)
except Exception:
logger.exception(
"[unified-render] pass-through BGM mix failed, skipping: plan_id=%s", self.plan.id
)
else:
ctx = RenderContext(work_dir=self.work_dir, plan_id=self.plan.id)
audio_path = mix_audio(ctx, layers, video_duration)
config = self.plan.config or {}
bgm_config = config.get("bgm", {}) or {}
noise_reduction_config = config.get("audio_noise_reduction")
ctx = RenderContext(
work_dir=self.work_dir,
plan_id=self.plan.id,
noise_reduction_config=noise_reduction_config,
)
audio_path = mix_audio(
ctx,
layers,
video_duration,
bgm_path=self.bgm_path,
bgm_config=bgm_config,
)
t_audio_end = time.time()
audio_mix_ms = int((t_audio_end - t_audio_start) * 1000)
has_audio = audio_path is not None
@@ -288,6 +371,85 @@ class UnifiedRenderService:
# 8. 探测输出
duration, file_size, width, height = self._probe_output(output_path)
# 9. 片头片尾拼接(后处理)
intro_outro_config = IntroOutroConfig.from_dict((self.plan.config or {}).get("intro_outro"))
if intro_outro_config.has_intro or intro_outro_config.has_outro:
io_valid, io_err = intro_outro_config.validate()
if io_valid:
final_with_io = self.work_dir / f"rendered_{self.plan.id}_with_io.mp4"
intro_path = None
outro_path = None
# 生成片头
if intro_outro_config.has_intro:
intro_path = self.work_dir / f"intro_{self.plan.id}.mp4"
intro_ok = False
if intro_outro_config.intro_type == "video":
import shutil
src = Path(intro_outro_config.intro_video_path)
if src.exists():
shutil.copy2(src, intro_path)
intro_ok = True
else:
logger.warning("片头视频不存在,跳过片头: %s", src)
elif intro_outro_config.intro_type == "text":
intro_ok = IntroOutroEngine.generate_text_intro(
intro_path,
intro_outro_config,
self.output_width,
self.output_height,
self.output_fps,
)
if not intro_ok:
intro_path = None
# 生成片尾
if intro_outro_config.has_outro:
outro_path = self.work_dir / f"outro_{self.plan.id}.mp4"
outro_ok = False
if intro_outro_config.outro_type == "video":
import shutil
src = Path(intro_outro_config.outro_video_path)
if src.exists():
shutil.copy2(src, outro_path)
outro_ok = True
else:
logger.warning("片尾视频不存在,跳过片尾: %s", src)
elif intro_outro_config.outro_type in ("text", "follow"):
outro_ok = IntroOutroEngine.generate_text_outro(
outro_path,
intro_outro_config,
self.output_width,
self.output_height,
self.output_fps,
)
if not outro_ok:
outro_path = None
# 拼接
if intro_path or outro_path:
concat_ok = IntroOutroEngine.concat_with_intro_outro(
output_path,
intro_path,
outro_path,
final_with_io,
transition_duration=intro_outro_config.transition_duration,
transition_effect=intro_outro_config.transition_effect,
)
if concat_ok and final_with_io.exists():
output_path = final_with_io
# 重新探测
duration, file_size, width, height = self._probe_output(output_path)
logger.info("[unified-render] 片头片尾拼接完成: plan_id=%s", self.plan.id)
else:
logger.warning("[unified-render] 片头片尾拼接失败,使用原视频: plan_id=%s", self.plan.id)
else:
logger.warning("[unified-render] 片头片尾配置无效,跳过: %s", io_err)
t_total = int((time.time() - t_start) * 1000)
logger.info(
"[unified-render] render done: plan_id=%s total_ms=%d video_ms=%d audio_ms=%d "
@@ -341,6 +503,10 @@ class UnifiedRenderService:
def _maybe_generate_ass(self, video_duration: float) -> Path | None:
"""根据 plan.config 生成 ASS 字幕文件。
支持两种字幕模式:
1. 静态字幕 — title/subtitle 配置了 text 时,生成整段静态字幕
2. ASR 自动字幕 — subtitle.auto_generated=true 时,从音频自动识别生成时间轴字幕
Returns:
ASS 文件路径,没有字幕时返回 None
"""
@@ -352,15 +518,46 @@ class UnifiedRenderService:
subtitle_enabled = subtitle_cfg.get("enabled", True)
title_text = title_cfg.get("text", "") or ""
subtitle_text = subtitle_cfg.get("text", "") or ""
auto_generated = subtitle_cfg.get("auto_generated", False)
has_title = title_enabled and bool(title_text.strip())
has_subtitle = subtitle_enabled and bool(subtitle_text.strip())
has_static_subtitle = subtitle_enabled and bool(subtitle_text.strip())
has_auto_subtitle = subtitle_enabled and auto_generated and self.asr_service is not None
if not has_title and not has_subtitle:
if not has_title and not has_static_subtitle and not has_auto_subtitle:
return None
ass_path = self.work_dir / f"subtitles_{self.plan.id}.ass"
# ASR 自动字幕模式
if has_auto_subtitle:
try:
timeline = self._generate_asr_subtitles(video_duration, subtitle_cfg)
if timeline and timeline.segments:
generate_ass_from_timeline(
ass_path,
timeline,
video_width=self.output_width,
video_height=self.output_height,
subtitle_config=subtitle_cfg,
)
logger.info(
"ASR自动字幕生成完成: plan_id=%s segments=%d duration=%.1fs",
self.plan.id,
timeline.segment_count,
video_duration,
)
return ass_path
else:
# ASR 无结果,不生成字幕
logger.info("ASR自动字幕无识别结果,跳过字幕: plan_id=%s", self.plan.id)
return None
except Exception:
# ASR 失败降级:不生成字幕,不阻断主流程
logger.warning("ASR自动字幕生成失败,跳过字幕", exc_info=True)
return None
# 静态字幕模式(原有逻辑)
generate_ass_subtitles(
ass_path,
video_width=self.output_width,
@@ -376,11 +573,173 @@ class UnifiedRenderService:
"生成字幕: plan_id=%s title=%s subtitle=%s ass=%s",
self.plan.id,
has_title,
has_subtitle,
has_static_subtitle,
ass_path,
)
return ass_path
def _generate_asr_subtitles(self, video_duration: float, subtitle_cfg: dict) -> Any: # SubtitleTimeline
"""从视频素材音频中自动识别生成字幕时间轴。
MVP 版本:使用第一个有音频的素材做ASR,然后按比例映射到整个视频时长。
后续优化:支持多片段拼接后的完整音频ASR。
"""
from packages.domain.subtitle import SubtitleTimeline
# 找第一个有本地路径的素材
first_asset_path = None
for clip in self.clips:
asset_id = getattr(clip, "asset_id", None)
if asset_id and asset_id in self.asset_path_map:
first_asset_path = self.asset_path_map[asset_id]
break
if first_asset_path is None:
logger.warning("ASR字幕生成失败:找不到可用素材音频")
return SubtitleTimeline(segments=[], total_duration=video_duration)
# 提取素材音频为 wav(16kHz单声道,ASR友好格式)
audio_path = self.work_dir / f"asr_audio_{self.plan.id}.wav"
try:
self._extract_audio(first_asset_path, audio_path)
except Exception:
logger.warning("ASR音频提取失败", exc_info=True)
return SubtitleTimeline(segments=[], total_duration=video_duration)
if not audio_path.exists():
return SubtitleTimeline(segments=[], total_duration=video_duration)
# 调用 ASR 服务
language = subtitle_cfg.get("language", "") or None
timeline = self.asr_service.transcribe(
audio_path,
language=language,
with_word_timestamps=True,
)
# 字幕后处理:合并短片段 + 拆分长片段
min_chars = int(subtitle_cfg.get("min_chars_per_segment", 8))
max_chars = int(subtitle_cfg.get("max_chars_per_line", 20))
if timeline.segments:
timeline = timeline.merge_short_segments(min_chars=min_chars)
timeline = timeline.split_long_segments(max_chars=max_chars)
# 清理临时音频文件
try:
audio_path.unlink(missing_ok=True)
except Exception:
pass
return timeline
def _extract_audio(self, video_path: Path, output_path: Path) -> None:
"""从视频中提取音频为16kHz单声道wav(ASR友好格式)。"""
import subprocess
cmd = [
"ffmpeg",
"-y",
"-i",
str(video_path),
"-vn",
"-acodec",
"pcm_s16le",
"-ar",
"16000",
"-ac",
"1",
str(output_path),
]
result = subprocess.run(
cmd,
capture_output=True,
text=True,
timeout=120,
)
if result.returncode != 0:
raise RuntimeError(f"音频提取失败: {result.stderr[:200]}")
def _maybe_add_voiceover_layer(
self,
layers: list[RenderLayer],
*,
video_duration: float,
) -> bool:
"""根据 plan.config 生成 TTS 配音,加到 audio 图层.
Returns:
是否成功添加了配音音轨
"""
config = self.plan.config or {}
tts_cfg = config.get("tts", {}) or {}
tts_config = TtsConfig.parse(tts_cfg)
if not tts_config.enabled:
return False
try:
from apps.worker.services.tts_service_factory import get_tts_service
tts_service = get_tts_service()
tts_engine = TtsEngine(tts_service, self.work_dir / "tts")
# 整段配音模式
result = tts_engine.generate_full_voiceover(tts_config, total_duration=video_duration)
if not result.success or not result.segments:
logger.warning("TTS 配音生成失败,跳过: %s", result.error_message)
return False
# 获取主音轨图层(用于判断 replace 模式下是否静音原音)
# 这里只处理混音添加,replace 模式在外部处理
# 找到或创建 audio 图层
audio_layer = None
for layer in layers:
if layer.role == "audio":
audio_layer = layer
break
if audio_layer is None:
from video_processing.unified_render_service import _LAYER_Z_INDEX # type: ignore
z_index = _LAYER_Z_INDEX.get("audio", 2)
audio_layer = RenderLayer(role="audio", z_index=z_index)
layers.append(audio_layer)
# 把配音片段作为 audio clip 加入
for seg in result.segments:
if seg.audio_path is None:
continue
vo_clip = ResolvedClip(
clip_id=f"tts_{seg.start_time:.3f}",
asset_id="tts_voiceover",
local_path=seg.audio_path,
clip_type="audio",
order=len(audio_layer.clips),
start_time=seg.start_time,
duration=seg.duration,
config={"volume": tts_config.volume, "tts": True},
actual_duration=seg.duration,
)
audio_layer.clips.append(vo_clip)
logger.info(
"TTS 配音已添加: plan_id=%s voice_id=%s segments=%d total_%.2fs",
self.plan.id,
tts_config.voice_id,
len(result.segments),
result.total_duration,
)
return True
except Exception as e:
logger.warning("TTS 配音异常,跳过: %s", e)
return False
def _can_use_pass_through(self, layers: list[RenderLayer]) -> bool:
"""判断是否可以走直通优化路径。
@@ -618,6 +977,25 @@ class UnifiedRenderService:
filters.append(f"scale={self.output_width}:{self.output_height}" ":force_original_aspect_ratio=increase")
filters.append(f"crop={self.output_width}:{self.output_height}")
# 调色滤镜
color_grade = ColorGradeConfig.from_dict(clip.config.get("color_grade"))
if color_grade.enabled and color_grade.has_effect():
grade_filter = ColorGradeEngine.build_filter(color_grade)
if grade_filter:
filters.append(grade_filter)
# chroma key 绿幕抠像
try:
from video_processing.chroma_key_engine import ChromaKeyConfig, ChromaKeyEngine
ck_config = ChromaKeyConfig.from_dict(clip.config.get("chroma_key"))
if ck_config.has_effect():
ck_engine = ChromaKeyEngine(ck_config)
ck_full = ck_engine.build_filter("[in]", "[out]")
ck_filter_part = ck_full[len("[in]") : -len("[out]")]
filters.append(ck_filter_part)
except Exception as e:
logger.warning("[unified-render] chroma key 直通模式应用失败,跳过: %s", e)
filters.append("setpts=PTS-STARTPTS")
filters.append(f"fps={self.output_fps}")
filters.append("format=yuv420p")
@@ -657,6 +1035,24 @@ class UnifiedRenderService:
# background 以外的视频素材,默认带音频
has_audio = role != "background"
if has_audio:
# 检查是否需要音频降噪
af_parts: list[str] = []
try:
from video_processing.noise_reduction_engine import NoiseReductionConfig, NoiseReductionEngine
plan_config = getattr(self.plan, "config", {}) or {}
nr_config = NoiseReductionConfig.from_dict(plan_config.get("audio_noise_reduction"))
if nr_config.has_effect():
nr_engine = NoiseReductionEngine(nr_config)
nr_full = nr_engine.build_filter("[in]", "[out]")
nr_filter_part = nr_full[len("[in]") : -len("[out]")]
af_parts.append(nr_filter_part)
except Exception as e:
logger.warning("[unified-render] 直通模式音频降噪应用失败,跳过: %s", e)
if af_parts:
command.extend(["-af", ",".join(af_parts)])
command.extend(["-c:a", "aac", "-b:a", "128k"])
# 统一截断时长(同时作用于视频和音频)
@@ -693,6 +1089,7 @@ class UnifiedRenderService:
"""将 EditPlanClip 列表解析为 ResolvedClip 列表。
跳过 asset_id 为空或在 asset_path_map 中找不到的片段。
支持多段裁剪:一个 clip 配置了 trim_segments 时会展开为多个 ResolvedClip。
"""
resolved: list[ResolvedClip] = []
for clip in self.clips:
@@ -712,17 +1109,76 @@ class UnifiedRenderService:
except Exception:
actual_duration = clip.duration or 5.0
# 检查是否有多段裁剪配置
clip_config = clip.config or {}
trim_segments = TrimEngine.parse_segments_from_config(clip_config)
if trim_segments and len(trim_segments) > 1:
# 多段裁剪:展开为多个 clip
resolved_segments = TrimEngine.resolve_segments(trim_segments, actual_duration)
for i, seg in enumerate(resolved_segments):
# 每个段生成一个独立的 ResolvedClip
seg_clip_id = f"{clip.id}_seg_{seg.segment_id}"
seg_order = clip.order + seg.order * 0.001 + i * 0.0001 # 保持排序
seg_start = seg.trim.start_time
seg_duration = seg.trim.duration
rc = ResolvedClip(
clip_id=seg_clip_id,
asset_id=asset_id,
local_path=local_path,
clip_type=clip.clip_type,
order=seg_order,
start_time=seg_start,
duration=seg_duration,
transition_effect=clip.transition_effect or "cut",
config={**clip_config, "_segment_id": seg.segment_id},
actual_duration=actual_duration,
trim_config=seg.trim,
)
resolved.append(rc)
continue
# 单段裁剪(或无裁剪)
# 解析裁剪配置:config 优先,否则用 clip.start_time + clip.duration
trim_config = extract_trim_from_clip_config(clip_config)
if trim_config is None and (clip.start_time > 0 or clip.duration > 0):
# 用旧字段构造
trim_config = TrimConfig(
start_time=clip.start_time,
duration=clip.duration,
)
# 钳制到实际素材时长
effective_trim: TrimConfig | None = None
final_start = clip.start_time
final_duration = clip.duration
if trim_config is not None and actual_duration > 0:
effective_trim = trim_config.validate_and_resolve(actual_duration)
if effective_trim.is_valid:
final_start = effective_trim.start_time
final_duration = effective_trim.duration
else:
# 裁剪无效 → 使用完整素材
logger.warning("裁剪配置无效,使用完整素材: clip_id=%s", clip.id)
effective_trim = None
final_start = 0.0
final_duration = actual_duration
rc = ResolvedClip(
clip_id=clip.id,
asset_id=asset_id,
local_path=local_path,
clip_type=clip.clip_type,
order=clip.order,
start_time=clip.start_time,
duration=clip.duration,
start_time=final_start,
duration=final_duration,
transition_effect=clip.transition_effect or "cut",
config=clip.config or {},
transition_duration=getattr(clip, "transition_duration", 0.0) or 0.0,
config=clip_config,
actual_duration=actual_duration,
trim_config=effective_trim,
)
resolved.append(rc)
@@ -797,7 +1253,7 @@ class UnifiedRenderService:
filter_parts: list[str] = []
# Step 1: 预处理每个 clip — scale + setpts
# Step 1: 预处理每个 clip — trim + scale + setpts
# 为每个 clip 生成预处理后的标签 [v0], [v1], ...
preprocessed_labels: list[str] = []
for i, clip in enumerate(all_clips):
@@ -806,11 +1262,15 @@ class UnifiedRenderService:
filters: list[str] = []
# trim — 始终将输出截断到有效时长,防止 xfade offset 与实际时长不匹配
# trim — 裁剪到指定区间,精确到帧
effective_duration = UnifiedRenderService._clip_effective_duration(clip)
trim_start = getattr(clip, "start_time", 0) or 0
if effective_duration > 0:
filters.append(f"trim=duration={effective_duration}")
if trim_start > 0:
filters.append(f"trim=start={trim_start:.3f}:duration={effective_duration:.3f}")
else:
filters.append(f"trim=duration={effective_duration:.3f}")
filters.append("setpts=PTS-STARTPTS")
# scale
@@ -831,6 +1291,26 @@ class UnifiedRenderService:
)
filters.append(f"crop={self.output_width}:{self.output_height}")
# 调色滤镜(每个 clip 独立的 color grade 配置)
color_grade = ColorGradeConfig.from_dict(clip.config.get("color_grade"))
if color_grade.enabled and color_grade.has_effect():
grade_filter = ColorGradeEngine.build_filter(color_grade)
if grade_filter:
filters.append(grade_filter)
# chroma key 绿幕抠像(在 scale 之后,fps 之前)
try:
from video_processing.chroma_key_engine import ChromaKeyConfig, ChromaKeyEngine
ck_config = ChromaKeyConfig.from_dict(clip.config.get("chroma_key"))
if ck_config.has_effect():
ck_engine = ChromaKeyEngine(ck_config)
# 提取滤镜部分(不带输入输出标签)
ck_full = ck_engine.build_filter("[in]", "[out]")
ck_filter_part = ck_full[len("[in]") : -len("[out]")]
filters.append(ck_filter_part)
except Exception as e:
logger.warning("[unified-render] chroma key 应用失败,跳过 clip=%s: %s", clip.clip_id, e)
filters.append("setpts=PTS-STARTPTS")
filters.append(f"fps={self.output_fps}")
@@ -846,18 +1326,25 @@ class UnifiedRenderService:
# 使用 trim 后的有效时长,与 Step 1 的 trim=duration 保持一致
layer_durations = [UnifiedRenderService._clip_effective_duration(all_clips[i]) for i in layer_clip_indices]
layer_transitions = [all_clips[i].transition_effect for i in layer_clip_indices]
layer_transition_durations = [all_clips[i].transition_duration for i in layer_clip_indices]
if len(layer_labels) == 1:
# 单 clip 层,直接使用预处理标签
layer_output_labels[layer.role] = layer_labels[0]
else:
# 多 clip 层,用 xfade 串联
# 多 clip 层,用 TransitionEngine 构建转场链
out_label = f"{layer.role}_merged"
xfade_filter, _ = build_xfade_filter_chain(
# 计算该层使用的转场时长(取首个非零值,否则用默认)
layer_dur = 0.0
for d in layer_transition_durations:
if d > 0:
layer_dur = d
break
xfade_filter, _ = self._transition_engine.build_xfade_chain(
clip_durations=layer_durations,
clip_video_labels=layer_labels,
transitions=layer_transitions,
transition_duration=self.transition_duration,
transition_duration=layer_dur if layer_dur > 0 else None,
output_label=out_label,
)
if xfade_filter:
@@ -904,6 +1391,66 @@ class UnifiedRenderService:
filter_parts.append(f"[{final_video_label}][{overlay_label}]" f"overlay={x}:{y}[{combined_label}]")
final_video_label = combined_label
# 叠加水印(在字幕之前)
watermark_config = WatermarkConfig.from_dict((self.plan.config or {}).get("watermark"))
if watermark_config is not None:
wm_valid, wm_err = watermark_config.validate()
if wm_valid:
wm_label = "watermarked"
if watermark_config.mode == "image":
# 图片水印:检查图片是否存在
wm_path = Path(watermark_config.image_path)
if wm_path.exists():
# 图片水印需要额外输入,放在 filter 开头
wm_idx = len(all_clips) # 水印图是最后一个输入
wm_scale = int(self.output_width * watermark_config.scale)
# 透明度
wm_filters = f"scale={wm_scale}:-1"
if watermark_config.opacity < 1.0:
wm_filters += f",format=rgba,colorchannelmixer=aa={watermark_config.opacity}"
filter_parts.insert(0, f"[{wm_idx}:v]{wm_filters}[wm_scaled]")
input_args.extend(["-i", str(wm_path)])
# 位置计算(水印高度用 scale 后的宽度近似)
wm_h = wm_scale # 近似(正方形假设)
x, y = WatermarkEngine.calc_position(
watermark_config.position,
self.output_width,
self.output_height,
wm_scale,
wm_h,
watermark_config.margin_x,
watermark_config.margin_y,
)
# 滚动水印
if watermark_config.scroll:
x_expr = f"W-mod({watermark_config.scroll_speed}*t\\,W+w)"
overlay = f"[{final_video_label}][wm_scaled]overlay=x={x_expr}:y={y}[{wm_label}]"
else:
overlay = f"[{final_video_label}][wm_scaled]overlay=x={x}:y={y}[{wm_label}]"
filter_parts.append(overlay)
final_video_label = wm_label
else:
logger.warning("水印图片不存在,跳过水印: %s", wm_path)
elif watermark_config.mode == "text":
# 文字水印
try:
text_wm = WatermarkEngine.build_text_watermark_filter(
f"[{final_video_label}]",
f"[{wm_label}]",
watermark_config,
self.output_width,
self.output_height,
)
filter_parts.append(text_wm)
final_video_label = wm_label
except Exception as e:
logger.warning("文字水印构建失败,跳过: %s", e)
# 叠加字幕(如有)+ 最终像素格式
if ass_path is not None:
ass_filter_path = str(ass_path).replace("\\", "/").replace(":", "\\:")
@@ -984,3 +1531,96 @@ class UnifiedRenderService:
if clip.duration > 0:
return min(clip.duration, clip.actual_duration) if clip.actual_duration > 0 else clip.duration
return clip.actual_duration if clip.actual_duration > 0 else 0.0
# ── 画中画(PiP)相关方法 ──────────────────────────────────────────────────
def _resolve_pip_sources(self, pip_config: PiPConfig) -> list[tuple[str, PiPLayerConfig, Path]]:
"""解析画中画图层的素材源,返回可用的图层列表.
降级策略:素材不存在或无效的图层自动跳过,不阻断渲染。
Returns:
[(input_label_placeholder, layer_config, local_path), ...]
input_label 在 build_pip_filters 中会用实际的输入索引替换
"""
if not pip_config.enabled:
return []
engine = PiPEngine(
output_width=self.output_width,
output_height=self.output_height,
output_fps=self.output_fps,
)
result = []
for i, layer in enumerate(pip_config.layers):
path = engine.validate_layer_source(layer, self.asset_path_map)
if path is None:
logger.warning("PiP图层素材不可用,跳过: layer_index=%d source=%s", i, layer.source)
continue
# 标签占位,实际输入索引由 build_pip_filters 内部管理
result.append((f"pip_src_{i}", layer, path))
return result
def _append_pip_filters(
self,
filter_complex: str,
input_args: list[str],
pip_sources: list[tuple[str, Any, Path]],
) -> tuple[str, list[str]]:
"""将画中画滤镜追加到 filter_complex 末尾.
处理逻辑:
1. 将原 final_video 标签重命名为 pip_base(作为PiP的底层视频)
2. 追加 PiP 预处理和 overlay 滤镜
3. PiP 最终输出命名为 final_video
Args:
filter_complex: 原 filter_complex 字符串
input_args: 原输入参数列表
pip_sources: PiP 素材列表 [(label, layer_config, path), ...]
Returns:
(new_filter_complex, new_input_args)
"""
if not pip_sources:
return filter_complex, input_args
pip_engine = PiPEngine(
output_width=self.output_width,
output_height=self.output_height,
output_fps=self.output_fps,
)
# 1. 将原 final_video 改为 pip_base
new_filter = filter_complex.replace("[final_video]", "[pip_base]")
# 2. 构建 PiP 滤镜链
# 主输入数量 = len(input_args) // 2(每个输入占 "-i path" 两个参数)
base_input_idx = len(input_args) // 2
pip_filter_parts, pip_input_args, final_label = pip_engine.build_pip_filters(
base_label="pip_base",
pip_sources=pip_sources,
base_input_idx=base_input_idx,
)
if not pip_filter_parts:
# 没有有效PiP滤镜,恢复原标签
return filter_complex, input_args
# 3. 追加 PiP 滤镜 + 最终格式转换(输出为 final_video
pip_filter_str = ";".join(pip_filter_parts)
final_format = f"[{final_label}]format=yuv420p[final_video]"
new_filter = f"{new_filter};{pip_filter_str};{final_format}"
# 4. 追加输入参数
new_input_args = list(input_args) + pip_input_args
logger.info(
"[unified-render] appended PiP filters: layers=%d new_inputs=%d",
len(pip_sources),
len(pip_input_args) // 2,
)
return new_filter, new_input_args
+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,
+40
View File
@@ -858,6 +858,14 @@
"type": "VARCHAR(20)",
"unique": false
},
{
"index": false,
"name": "transition_duration",
"nullable": false,
"primary_key": false,
"type": "FLOAT",
"unique": false
},
{
"index": true,
"name": "status",
@@ -1493,6 +1501,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
+3
View File
@@ -54,6 +54,7 @@ class SQLAlchemyEditPlanClipRepository:
start_time=clip.start_time,
duration=clip.duration,
transition_effect=clip.transition_effect,
transition_duration=clip.transition_duration,
status=clip.status,
config=clip.config,
)
@@ -76,6 +77,7 @@ 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.status = clip.status
model.config = clip.config
model.updated_at = clip.updated_at
@@ -120,6 +122,7 @@ 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,
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,7 @@ 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)
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 +251,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",
}
+3
View File
@@ -50,6 +50,7 @@ class EditPlanClip:
start_time: float = 0.0
duration: float = 0.0
transition_effect: str = "cut"
transition_duration: float = 0.0 # 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 +69,7 @@ class EditPlanClip:
start_time: float = 0.0,
duration: float = 0.0,
transition_effect: str = "cut",
transition_duration: float = 0.0,
config: dict[str, Any] | None = None,
) -> EditPlanClip:
"""创建剪辑计划片段"""
@@ -91,6 +93,7 @@ class EditPlanClip:
start_time=start_time,
duration=duration,
transition_effect=transition_effect.strip() or "cut",
transition_duration=max(0.0, transition_duration),
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_message、completed_at。
设置 error_message、error_info、completed_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_message、started_at、completed_at、progress
清除 error_message、error_info、started_at、completed_at、progress
递增 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
+1 -1
View File
@@ -179,4 +179,4 @@ else
fi
echo "=== Build complete ==="
docker images | grep "xiaoxia-saas.*:$VERSION"
docker images | grep "xiaoxia-saas" | grep "$VERSION" || true
+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
+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
+102
View File
@@ -453,3 +453,105 @@ class TestFullFlow:
task.mark_cancelled()
assert task.status == GenerationTaskStatus.CANCELLED
assert task.is_terminal
# ── 错误信息与重试(任务中心升级) ──────────────────────────────────────────
class TestErrorInfo:
"""测试 error_info 结构化错误信息。"""
def test_mark_failed_default_error_info(self) -> None:
"""mark_failed 不传 error_info 时自动生成默认结构。"""
task = _make_task()
task.mark_processing()
task.mark_failed("something went wrong")
assert task.is_failed
assert task.error_message == "something went wrong"
assert task.error_info["error_type"] == "UnknownError"
assert task.error_info["message"] == "something went wrong"
assert "failed_at" in task.error_info
def test_mark_failed_with_custom_error_info(self) -> None:
"""mark_failed 传自定义 error_info。"""
task = _make_task()
task.mark_processing()
info = {
"error_type": "FFmpegError",
"message": "Invalid data found",
"stack_trace": "Traceback...",
"stage": "render",
"failed_at": "2026-01-01T00:00:00+00:00",
}
task.mark_failed("Invalid data found", error_info=info)
assert task.error_info == info
def test_error_info_cleared_on_retry(self) -> None:
"""重试时 error_info 被清空。"""
task = _make_task()
task.mark_processing()
task.mark_failed("oops")
assert task.error_info # 失败时有值
task.mark_pending_from_failed()
assert task.error_info == {}
assert task.status == GenerationTaskStatus.PENDING
class TestRetryCount:
"""测试 retry_count 重试次数。"""
def test_default_retry_count_is_zero(self) -> None:
"""新任务 retry_count 默认 0。"""
task = _make_task()
assert task.retry_count == 0
def test_retry_increments_count(self) -> None:
"""每次失败后重试,retry_count +1。"""
task = _make_task()
task.mark_processing()
task.mark_failed("fail 1")
task.mark_pending_from_failed()
assert task.retry_count == 1
task.mark_processing()
task.mark_failed("fail 2")
task.mark_pending_from_failed()
assert task.retry_count == 2
def test_completed_does_not_affect_retry_count(self) -> None:
"""正常完成不改变 retry_count。"""
task = _make_task()
task.mark_processing()
task.mark_completed()
assert task.retry_count == 0
class TestAutoRetryConfig:
"""测试自动重试配置。"""
def test_default_auto_retry_disabled(self) -> None:
"""默认关闭自动重试。"""
task = _make_task()
assert task.auto_retry_enabled is False
assert task.auto_retry_max == 0
def test_create_with_auto_retry(self) -> None:
"""create 工厂方法支持 auto_retry 参数。"""
task = GenerationTask.create(
project_id="proj-1",
asset_library_id="lib-1",
auto_retry_enabled=True,
auto_retry_max=3,
)
assert task.auto_retry_enabled is True
assert task.auto_retry_max == 3
def test_auto_retry_max_default_zero(self) -> None:
"""auto_retry_max 默认 0 表示不自动重试。"""
task = GenerationTask.create(
project_id="proj-1",
asset_library_id="lib-1",
auto_retry_enabled=True,
)
assert task.auto_retry_enabled is True
assert task.auto_retry_max == 0
+597
View File
@@ -0,0 +1,597 @@
"""画中画(PiP)引擎单元测试."""
from __future__ import annotations
from pathlib import Path
from unittest.mock import patch
import pytest
from video_processing.pip_engine import (
ANIMATION_FADE,
ANIMATION_SLIDE_BOTTOM,
ANIMATION_SLIDE_LEFT,
ANIMATION_SLIDE_RIGHT,
ANIMATION_SLIDE_TOP,
POSITION_BOTTOM_LEFT,
POSITION_BOTTOM_RIGHT,
POSITION_CENTER,
POSITION_TOP_LEFT,
POSITION_TOP_RIGHT,
PiPConfig,
PiPEngine,
PiPLayerConfig,
)
# ── PiPLayerConfig.validate 测试 ──────────────────────────────────────────────
class TestPiPLayerConfigValidate:
"""PiP图层配置校验测试."""
def test_valid_config(self):
"""正常配置应该通过校验."""
layer = PiPLayerConfig(source="asset_001")
ok, err = layer.validate()
assert ok
assert err == ""
def test_empty_source(self):
"""空source应该失败."""
layer = PiPLayerConfig(source="")
ok, err = layer.validate()
assert not ok
assert "source" in err
def test_invalid_position(self):
"""无效位置应该失败."""
layer = PiPLayerConfig(source="asset_001", position="invalid_pos")
ok, err = layer.validate()
assert not ok
assert "position" in err
def test_custom_position_valid(self):
"""custom位置应该通过."""
layer = PiPLayerConfig(source="asset_001", position="custom", x=100, y=50)
ok, err = layer.validate()
assert ok
def test_opacity_out_of_range_high(self):
"""opacity超过1应该失败."""
layer = PiPLayerConfig(source="asset_001", opacity=1.5)
ok, err = layer.validate()
assert not ok
assert "opacity" in err
def test_opacity_out_of_range_low(self):
"""opacity小于0应该失败."""
layer = PiPLayerConfig(source="asset_001", opacity=-0.5)
ok, err = layer.validate()
assert not ok
assert "opacity" in err
def test_opacity_boundary_values(self):
"""opacity边界值应该通过."""
for val in [0.0, 0.5, 1.0]:
layer = PiPLayerConfig(source="asset_001", opacity=val)
ok, _ = layer.validate()
assert ok
def test_negative_corner_radius(self):
"""负圆角应该失败."""
layer = PiPLayerConfig(source="asset_001", corner_radius=-5)
ok, err = layer.validate()
assert not ok
assert "corner_radius" in err
def test_negative_start_time(self):
"""负开始时间应该失败."""
layer = PiPLayerConfig(source="asset_001", start_time=-1.0)
ok, err = layer.validate()
assert not ok
assert "start_time" in err
def test_negative_duration(self):
"""负持续时间应该失败."""
layer = PiPLayerConfig(source="asset_001", duration=-5.0)
ok, err = layer.validate()
assert not ok
assert "duration" in err
def test_invalid_animation_in(self):
"""无效入场动画应该失败."""
layer = PiPLayerConfig(source="asset_001", animation_in="spin")
ok, err = layer.validate()
assert not ok
assert "入场动画" in err
def test_all_valid_animations(self):
"""所有有效动画类型应该通过."""
for anim in [
ANIMATION_FADE,
ANIMATION_SLIDE_LEFT,
ANIMATION_SLIDE_RIGHT,
ANIMATION_SLIDE_TOP,
ANIMATION_SLIDE_BOTTOM,
]:
layer = PiPLayerConfig(source="asset_001", animation_in=anim, animation_out=anim)
ok, _ = layer.validate()
assert ok
def test_zero_duration_valid(self):
"""duration=0(全程显示)应该通过."""
layer = PiPLayerConfig(source="asset_001", duration=0.0)
ok, _ = layer.validate()
assert ok
# ── PiPConfig.from_dict 测试 ──────────────────────────────────────────────────
class TestPiPConfigFromDict:
"""PiP配置字典解析测试."""
def test_none_config(self):
"""None配置应该返回disabled."""
config = PiPConfig.from_dict(None)
assert not config.enabled
assert len(config.layers) == 0
def test_empty_config(self):
"""空字典应该返回disabled."""
config = PiPConfig.from_dict({})
assert not config.enabled
def test_enabled_false(self):
"""enabled=False应该返回disabled."""
config = PiPConfig.from_dict({"enabled": False, "layers": [{"source": "a"}]})
assert not config.enabled
def test_single_layer(self):
"""单图层解析."""
data = {
"enabled": True,
"layers": [
{
"source": "asset_001",
"position": POSITION_TOP_RIGHT,
"width": "30%",
"opacity": 0.9,
"corner_radius": 10,
"start_time": 2.0,
"duration": 5.0,
"z_index": 2,
}
],
}
config = PiPConfig.from_dict(data)
assert config.enabled
assert len(config.layers) == 1
layer = config.layers[0]
assert layer.source == "asset_001"
assert layer.position == POSITION_TOP_RIGHT
assert layer.width == "30%"
assert layer.opacity == 0.9
assert layer.corner_radius == 10
assert layer.start_time == 2.0
assert layer.duration == 5.0
assert layer.z_index == 2
def test_multiple_layers_sorted_by_z_index(self):
"""多图层应该按z_index排序."""
data = {
"enabled": True,
"layers": [
{"source": "asset_high", "z_index": 5},
{"source": "asset_low", "z_index": 1},
{"source": "asset_mid", "z_index": 3},
],
}
config = PiPConfig.from_dict(data)
assert len(config.layers) == 3
assert config.layers[0].source == "asset_low"
assert config.layers[1].source == "asset_mid"
assert config.layers[2].source == "asset_high"
def test_invalid_layer_skipped(self):
"""无效图层应该被跳过."""
data = {
"enabled": True,
"layers": [
{"source": "asset_good"},
{"source": "", "position": "invalid"}, # 空source
{"source": "asset_good2", "opacity": 2.0}, # opacity超范围
],
}
config = PiPConfig.from_dict(data)
# 第1个有效,第2、3个无效
assert len(config.layers) == 1
assert config.layers[0].source == "asset_good"
def test_all_invalid_layers_disabled(self):
"""所有图层都无效时enabled为False."""
data = {
"enabled": True,
"layers": [
{"source": ""},
{"source": ""},
],
}
config = PiPConfig.from_dict(data)
assert not config.enabled
assert len(config.layers) == 0
def test_default_values(self):
"""默认值应该正确."""
data = {
"enabled": True,
"layers": [{"source": "asset_001"}],
}
config = PiPConfig.from_dict(data)
layer = config.layers[0]
assert layer.position == POSITION_BOTTOM_RIGHT
assert layer.width == "25%"
assert layer.opacity == 1.0
assert layer.corner_radius == 0
assert layer.start_time == 0.0
assert layer.duration == 0.0
assert layer.z_index == 1
# ── PiPEngine 位置计算测试 ────────────────────────────────────────────────────
class TestPiPEnginePosition:
"""PiP引擎位置计算测试."""
@pytest.fixture
def engine(self):
return PiPEngine(output_width=1920, output_height=1080, output_fps=30)
def test_top_left_position(self, engine):
"""左上角位置."""
layer = PiPLayerConfig(source="a", position=POSITION_TOP_LEFT, margin=20)
x, y = engine._parse_position(layer, 480, 270)
assert x == 20
assert y == 20
def test_top_right_position(self, engine):
"""右上角位置."""
layer = PiPLayerConfig(source="a", position=POSITION_TOP_RIGHT, margin=20)
x, y = engine._parse_position(layer, 480, 270)
assert x == 1920 - 480 - 20
assert y == 20
def test_bottom_right_position(self, engine):
"""右下角位置(默认)."""
layer = PiPLayerConfig(source="a", position=POSITION_BOTTOM_RIGHT, margin=30)
x, y = engine._parse_position(layer, 480, 270)
assert x == 1920 - 480 - 30
assert y == 1080 - 270 - 30
def test_bottom_left_position(self, engine):
"""左下角位置."""
layer = PiPLayerConfig(source="a", position=POSITION_BOTTOM_LEFT, margin=15)
x, y = engine._parse_position(layer, 480, 270)
assert x == 15
assert y == 1080 - 270 - 15
def test_center_position(self, engine):
"""中心位置."""
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, margin=0)
x, y = engine._parse_position(layer, 480, 270)
assert x == (1920 - 480) // 2
assert y == (1080 - 270) // 2
def test_custom_position_pixel(self, engine):
"""自定义像素位置."""
layer = PiPLayerConfig(source="a", position="custom", x=100, y=200)
x, y = engine._parse_position(layer, 480, 270)
assert x == 100
assert y == 200
def test_custom_position_percentage(self, engine):
"""自定义百分比位置."""
layer = PiPLayerConfig(source="a", position="custom", x="50%", y="25%")
x, y = engine._parse_position(layer, 480, 270)
assert x == 1920 // 2
assert y == 1080 // 4
def test_top_center_position(self, engine):
"""顶部居中位置."""
layer = PiPLayerConfig(source="a", position="top_center", margin=10)
x, y = engine._parse_position(layer, 480, 270)
assert x == (1920 - 480) // 2
assert y == 10
def test_invalid_position_fallback(self, engine):
"""无效位置应该fallback到右下角."""
layer = PiPLayerConfig(source="a", position="unknown_position", margin=20)
# 直接测试_parse_position(注意:validate会拦截,但_parse_position自己也有fallback
x, y = engine._parse_position(layer, 480, 270)
assert x == 1920 - 480 - 20
assert y == 1080 - 270 - 20
# ── PiPEngine 尺寸解析测试 ────────────────────────────────────────────────────
class TestPiPEngineSize:
"""PiP引擎尺寸解析测试."""
@pytest.fixture
def engine(self):
return PiPEngine(output_width=1920, output_height=1080, output_fps=30)
def test_pixel_size_int(self, engine):
"""像素尺寸(整数)."""
assert engine._parse_size(500, 1920) == 500
def test_pixel_size_str(self, engine):
"""像素尺寸(字符串数字)."""
assert engine._parse_size("500", 1920) == 500
def test_percentage_size(self, engine):
"""百分比尺寸."""
assert engine._parse_size("50%", 1920) == 960
assert engine._parse_size("25%", 1920) == 480
def test_zero_size_default(self, engine):
"""0或无效值应该有最小值保护."""
assert engine._parse_size(0, 1920) == 1
assert engine._parse_size("", 1920) == 480 # 默认25%
def test_negative_size_default(self, engine):
"""负值应该取绝对值后至少为1."""
# _parse_size 用 max(1, value),负值会走 except 分支
result = engine._parse_size("-100", 1920)
# 会走ValueError分支,返回默认值
assert result > 0
# ── PiPEngine 滤镜构建测试 ────────────────────────────────────────────────────
class TestPiPEngineBuildFilters:
"""PiP引擎滤镜构建测试."""
@pytest.fixture
def engine(self):
return PiPEngine(output_width=1920, output_height=1080, output_fps=30)
@pytest.fixture
def fake_video(self, tmp_path):
"""创建一个假的视频文件路径."""
path = tmp_path / "test_video.mp4"
path.write_bytes(b"fake video data")
return path
def test_empty_sources(self, engine):
"""空素材列表应该返回空."""
filters, inputs, label = engine.build_pip_filters("base_label", [])
assert filters == []
assert inputs == []
assert label == "base_label"
def test_single_layer_basic(self, engine, fake_video):
"""单图层基础滤镜构建."""
layer = PiPLayerConfig(
source="asset_001",
position=POSITION_TOP_RIGHT,
width="25%",
)
sources = [("pip_src_0", layer, fake_video)]
filters, inputs, final_label = engine.build_pip_filters("base_video", sources, base_input_idx=3)
# 应该有2个滤镜: 预处理 + overlay
assert len(filters) == 2
# 输入参数应该有2个(-i + path)
assert len(inputs) == 2
assert inputs[0] == "-i"
assert inputs[1] == str(fake_video)
# 预处理滤镜应该使用正确的输入索引
assert "3:v" in filters[0]
# 应该包含scale
assert "scale=" in filters[0]
# 应该有pip_pre_0标签
assert "[pip_pre_0]" in filters[0]
# overlay滤镜
assert "overlay=" in filters[1]
assert "[base_video][pip_pre_0]" in filters[1]
def test_single_layer_final_label(self, engine, fake_video):
"""最终输出标签应该正确."""
layer = PiPLayerConfig(source="a", position=POSITION_CENTER)
sources = [("s0", layer, fake_video)]
_, _, final_label = engine.build_pip_filters("main_v", sources)
assert final_label == "pip_combined_0"
def test_multiple_layers(self, engine, fake_video):
"""多图层叠加."""
layer1 = PiPLayerConfig(source="a", position=POSITION_TOP_LEFT, z_index=1)
layer2 = PiPLayerConfig(source="b", position=POSITION_BOTTOM_RIGHT, z_index=2)
sources = [
("s0", layer1, fake_video),
("s1", layer2, fake_video),
]
filters, inputs, final_label = engine.build_pip_filters("base", sources, base_input_idx=0)
# 2层 × 2个滤镜(预处理+overlay)= 4个滤镜
assert len(filters) == 4
# 2个输入文件
assert len(inputs) == 4 # 2 × (-i + path)
# 输入索引应该连续
assert "0:v" in filters[0]
assert "1:v" in filters[2]
# 最终标签应该是第二个overlay的输出
assert final_label == "pip_combined_1"
def test_with_opacity(self, engine, fake_video):
"""透明度应该在滤镜中体现."""
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, opacity=0.5)
sources = [("s0", layer, fake_video)]
filters, _, _ = engine.build_pip_filters("base", sources)
pre_filter = filters[0]
assert "colorchannelmixer=aa=0.5" in pre_filter
assert "yuva420p" in pre_filter
def test_with_corner_radius(self, engine, fake_video):
"""圆角裁剪应该在滤镜中体现."""
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, corner_radius=20)
sources = [("s0", layer, fake_video)]
filters, _, _ = engine.build_pip_filters("base", sources)
pre_filter = filters[0]
assert "geq=" in pre_filter
def test_with_border(self, engine, fake_video):
"""边框应该在滤镜中体现."""
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, border_width=3, border_color="red")
sources = [("s0", layer, fake_video)]
filters, _, _ = engine.build_pip_filters("base", sources)
pre_filter = filters[0]
assert "pad=" in pre_filter
assert "red" in pre_filter
def test_timing_start_time_and_duration(self, engine, fake_video):
"""时间控制应该生成enable表达式."""
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, start_time=5.0, duration=10.0)
sources = [("s0", layer, fake_video)]
filters, _, _ = engine.build_pip_filters("base", sources)
overlay_filter = filters[1]
assert "enable=" in overlay_filter
assert "between(t,5.0,15.0)" in overlay_filter
def test_timing_start_time_only(self, engine, fake_video):
"""只有开始时间(全程显示到结束)."""
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, start_time=3.0, duration=0.0)
sources = [("s0", layer, fake_video)]
filters, _, _ = engine.build_pip_filters("base", sources)
overlay_filter = filters[1]
assert "enable=" in overlay_filter
assert "gte(t,3.0)" in overlay_filter
def test_no_timing_no_enable(self, engine, fake_video):
"""无时间限制时不应该有enable表达式."""
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, start_time=0.0, duration=0.0)
sources = [("s0", layer, fake_video)]
filters, _, _ = engine.build_pip_filters("base", sources)
overlay_filter = filters[1]
assert "enable=" not in overlay_filter
def test_fade_animation(self, engine, fake_video):
"""淡入淡出动画."""
layer = PiPLayerConfig(
source="a",
position=POSITION_CENTER,
animation_in=ANIMATION_FADE,
animation_out=ANIMATION_FADE,
duration=10.0,
animation_duration=0.8,
)
sources = [("s0", layer, fake_video)]
filters, _, _ = engine.build_pip_filters("base", sources)
pre_filter = filters[0]
assert "fade=t=in" in pre_filter
assert "fade=t=out" in pre_filter
assert "alpha=1" in pre_filter
def test_slide_animation_in(self, engine, fake_video):
"""滑入动画应该在overlay表达式中."""
layer = PiPLayerConfig(
source="a",
position=POSITION_CENTER,
animation_in=ANIMATION_SLIDE_LEFT,
animation_duration=0.5,
)
sources = [("s0", layer, fake_video)]
filters, _, _ = engine.build_pip_filters("base", sources)
overlay_filter = filters[1]
# x表达式应该包含动态变化
assert "overlay=" in overlay_filter
def test_full_opacity_no_alpha(self, engine, fake_video):
"""opacity=1时不应该有colorchannelmixer."""
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, opacity=1.0)
sources = [("s0", layer, fake_video)]
filters, _, _ = engine.build_pip_filters("base", sources)
pre_filter = filters[0]
assert "colorchannelmixer" not in pre_filter
def test_zero_corner_radius_no_geq(self, engine, fake_video):
"""corner_radius=0时不应该有geq滤镜."""
layer = PiPLayerConfig(source="a", position=POSITION_CENTER, corner_radius=0)
sources = [("s0", layer, fake_video)]
filters, _, _ = engine.build_pip_filters("base", sources)
pre_filter = filters[0]
assert "geq=" not in pre_filter
# ── PiPEngine 素材验证(降级策略)测试 ────────────────────────────────────────
class TestPiPEngineValidateSource:
"""PiP引擎素材验证与降级测试."""
@pytest.fixture
def engine(self):
return PiPEngine(output_width=1920, output_height=1080, output_fps=30)
def test_asset_id_in_map(self, engine, tmp_path):
"""asset_id在map中应该返回路径."""
asset_path = tmp_path / "test.mp4"
asset_path.write_bytes(b"data")
asset_map = {"asset_001": asset_path}
layer = PiPLayerConfig(source="asset_001", source_type="asset_id")
result = engine.validate_layer_source(layer, asset_map)
assert result == asset_path
def test_asset_id_not_in_map(self, engine):
"""asset_id不在map中应该返回None(降级)."""
layer = PiPLayerConfig(source="nonexistent", source_type="asset_id")
result = engine.validate_layer_source(layer, {})
assert result is None
def test_local_path_exists(self, engine, tmp_path):
"""本地路径存在应该返回."""
path = tmp_path / "video.mp4"
path.write_bytes(b"data")
layer = PiPLayerConfig(source=str(path), source_type="local_path")
result = engine.validate_layer_source(layer, {})
assert result == path
def test_local_path_not_exists(self, engine):
"""本地路径不存在应该返回None(降级)."""
layer = PiPLayerConfig(source="/nonexistent/path.mp4", source_type="local_path")
result = engine.validate_layer_source(layer, {})
assert result is None
def test_url_type_not_supported(self, engine):
"""URL类型暂时不支持,返回None."""
layer = PiPLayerConfig(source="http://example.com/video.mp4", source_type="url")
result = engine.validate_layer_source(layer, {})
assert result is None
def test_exception_handling(self, engine):
"""异常情况应该返回None(不阻断)."""
layer = PiPLayerConfig(source=None, source_type="local_path") # type: ignore
# 模拟异常情况
result = engine.validate_layer_source(layer, {})
assert result is None
+178
View File
@@ -0,0 +1,178 @@
"""字幕生成器 + Mock ASR 单元测试。"""
import tempfile
from pathlib import Path
import pytest
from apps.worker.video_processing.subtitle_generator import (
_wrap_text,
generate_ass_from_timeline,
)
from packages.adapters.asr.mock_asr_service import MockASRService
from packages.domain.subtitle import SubtitleSegment, SubtitleTimeline
from packages.ports.asr_service import ASRServiceError
class TestMockASRService:
def test_transcribe_with_mock_text(self):
service = MockASRService(mock_text="你好世界!这是一段测试语音识别的文字。用来验证Mock ASR是否正常工作。")
# 创建一个假的音频文件(mock不真的读内容)
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f:
f.write(b"fake audio data")
audio_path = Path(f.name)
try:
timeline = service.transcribe(audio_path, language="zh")
assert timeline is not None
assert timeline.language == "zh"
assert timeline.segment_count > 0
assert timeline.total_duration > 0
# 总字数应该对得上
assert timeline.total_chars == len("你好世界!这是一段测试语音识别的文字。用来验证Mock ASR是否正常工作。")
finally:
audio_path.unlink()
def test_transcribe_file_not_found(self):
service = MockASRService()
with pytest.raises(ASRServiceError):
service.transcribe(Path("/nonexistent/audio.wav"))
def test_transcribe_with_word_timestamps(self):
service = MockASRService(mock_text="你好世界!")
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f:
f.write(b"fake")
audio_path = Path(f.name)
try:
timeline = service.transcribe(audio_path, with_word_timestamps=True)
# 每段应该有词级时间戳
for seg in timeline.segments:
if seg.words:
assert len(seg.words) > 0
assert seg.words[0].start >= seg.start
assert seg.words[-1].end <= seg.end
finally:
audio_path.unlink()
def test_auto_detect_language(self):
service = MockASRService(mock_text="Hello world. This is a test.")
with tempfile.NamedTemporaryFile(suffix=".wav", delete=False) as f:
f.write(b"fake")
audio_path = Path(f.name)
try:
timeline = service.transcribe(audio_path, language=None)
# None 时默认 zh
assert timeline.language == "zh"
finally:
audio_path.unlink()
class TestWrapText:
def test_short_text_no_wrap(self):
result = _wrap_text("你好世界", 20)
assert result == ["你好世界"]
def test_wrap_at_punctuation(self):
result = _wrap_text("你好世界!这是一段很长的测试文字。", 10)
assert len(result) == 2
assert "" in result[0]
def test_hard_wrap_no_punctuation(self):
result = _wrap_text("一二三四五六七八九十一二三四五六七八九十", 10)
assert len(result) == 2
assert len(result[0]) == 10
assert len(result[1]) == 10
def test_exact_length(self):
result = _wrap_text("一二三四五六七八九十", 10)
assert len(result) == 1
class TestGenerateAssFromTimeline:
def test_generate_basic(self):
timeline = SubtitleTimeline(
segments=[
SubtitleSegment(text="你好世界", start=0.0, end=2.0),
SubtitleSegment(text="这是测试", start=2.0, end=4.0),
],
language="zh",
total_duration=4.0,
)
with tempfile.TemporaryDirectory() as tmpdir:
output_path = Path(tmpdir) / "test.ass"
result = generate_ass_from_timeline(
output_path,
timeline,
video_width=1920,
video_height=1080,
)
assert result.exists()
content = result.read_text(encoding="utf-8")
assert "[Script Info]" in content
assert "[V4+ Styles]" in content
assert "[Events]" in content
assert "你好世界" in content
assert "这是测试" in content
assert "PlayResX: 1920" in content
assert "PlayResY: 1080" in content
def test_empty_timeline(self):
timeline = SubtitleTimeline(segments=[])
with tempfile.TemporaryDirectory() as tmpdir:
output_path = Path(tmpdir) / "empty.ass"
result = generate_ass_from_timeline(output_path, timeline, video_width=1920, video_height=1080)
assert result.exists()
assert result.read_text(encoding="utf-8") == ""
def test_with_custom_style(self):
timeline = SubtitleTimeline(
segments=[SubtitleSegment(text="测试", start=0, end=1)],
total_duration=1.0,
)
with tempfile.TemporaryDirectory() as tmpdir:
output_path = Path(tmpdir) / "style.ass"
generate_ass_from_timeline(
output_path,
timeline,
video_width=1280,
video_height=720,
subtitle_config={
"font": "微软雅黑",
"size": 32,
"color": "#ff0000",
"position": "bottom",
},
)
content = output_path.read_text(encoding="utf-8")
assert "微软雅黑" in content
assert "32" in content
def test_time_format(self):
timeline = SubtitleTimeline(
segments=[
SubtitleSegment(text="测试", start=0.5, end=1.25),
SubtitleSegment(text="长字幕", start=3661.0, end=3662.5), # 超过1小时
],
total_duration=3662.5,
)
with tempfile.TemporaryDirectory() as tmpdir:
output_path = Path(tmpdir) / "time.ass"
generate_ass_from_timeline(output_path, timeline, video_width=1920, video_height=1080)
content = output_path.read_text(encoding="utf-8")
# 0:00:00.50 格式
assert "0:00:00.50" in content
assert "0:00:01.25" in content
# 1:01:01.00 格式(3661秒 = 1小时1分1秒)
assert "1:01:01.00" in content
+183
View File
@@ -0,0 +1,183 @@
"""字幕时间轴单元测试。"""
import pytest
from packages.domain.subtitle import (
SubtitleSegment,
SubtitleTimeline,
SubtitleWord,
)
class TestSubtitleWord:
def test_duration(self):
word = SubtitleWord(text="", start=1.0, end=1.5)
assert word.duration == pytest.approx(0.5)
def test_zero_duration(self):
word = SubtitleWord(text="", start=1.0, end=1.0)
assert word.duration == 0.0
class TestSubtitleSegment:
def test_duration(self):
seg = SubtitleSegment(text="你好世界", start=0.0, end=2.0)
assert seg.duration == pytest.approx(2.0)
def test_char_count(self):
seg = SubtitleSegment(text="你好世界", start=0.0, end=2.0)
assert seg.char_count == 4
class TestSubtitleTimeline:
def test_segment_count(self):
timeline = SubtitleTimeline(
segments=[
SubtitleSegment(text="第一段", start=0, end=1),
SubtitleSegment(text="第二段", start=1, end=2),
]
)
assert timeline.segment_count == 2
def test_total_chars(self):
timeline = SubtitleTimeline(
segments=[
SubtitleSegment(text="你好", start=0, end=1),
SubtitleSegment(text="世界", start=1, end=2),
]
)
assert timeline.total_chars == 4
class TestMergeShortSegments:
def test_no_merge_when_long_enough(self):
timeline = SubtitleTimeline(
segments=[
SubtitleSegment(text="这是第一段测试文字", start=0, end=2),
SubtitleSegment(text="这是第二段测试文字", start=2, end=4),
],
total_duration=4.0,
)
result = timeline.merge_short_segments(min_chars=8)
assert result.segment_count == 2
def test_merge_short_segments(self):
timeline = SubtitleTimeline(
segments=[
SubtitleSegment(text="你好", start=0, end=0.5),
SubtitleSegment(text="世界", start=0.5, end=1.0),
SubtitleSegment(text="这是一段长文字", start=1.0, end=3.0),
],
total_duration=3.0,
)
result = timeline.merge_short_segments(min_chars=4)
# 前两段合并(共4字),第三段保留
assert result.segment_count == 2
assert result.segments[0].text == "你好世界"
assert result.segments[0].start == 0
assert result.segments[0].end == 1.0
def test_merge_remaining_to_last(self):
timeline = SubtitleTimeline(
segments=[
SubtitleSegment(text="这是第一段测试文字", start=0, end=2),
SubtitleSegment(text="", start=2, end=2.2),
SubtitleSegment(text="", start=2.2, end=2.4),
],
total_duration=2.4,
)
result = timeline.merge_short_segments(min_chars=8)
# 最后两段字数不够,合并到上一段
assert result.segment_count == 1
assert result.segments[0].text == "这是第一段测试文字你好"
def test_single_segment_no_change(self):
timeline = SubtitleTimeline(
segments=[SubtitleSegment(text="你好", start=0, end=1)],
total_duration=1.0,
)
result = timeline.merge_short_segments(min_chars=8)
assert result.segment_count == 1
assert result.segments[0].text == "你好"
def test_empty_timeline(self):
timeline = SubtitleTimeline(segments=[])
result = timeline.merge_short_segments()
assert result.segment_count == 0
class TestSplitLongSegments:
def test_no_split_when_short_enough(self):
timeline = SubtitleTimeline(
segments=[SubtitleSegment(text="你好世界", start=0, end=1)],
total_duration=1.0,
)
result = timeline.split_long_segments(max_chars=20)
assert result.segment_count == 1
def test_split_by_punctuation(self):
timeline = SubtitleTimeline(
segments=[
SubtitleSegment(
text="这是第一段很长的测试文字。这是第二段很长的测试文字!这是第三段很长的测试文字?",
start=0,
end=6.0,
)
],
total_duration=6.0,
)
result = timeline.split_long_segments(max_chars=15)
# 按标点拆成3段
assert result.segment_count == 3
assert "" in result.segments[0].text
assert "" in result.segments[1].text
assert "" in result.segments[2].text
def test_hard_split_when_no_punctuation(self):
timeline = SubtitleTimeline(
segments=[
SubtitleSegment(
text="一二三四五六七八九十一二三四五六七八九十一二三四五六七八九十",
start=0,
end=6.0,
)
],
total_duration=6.0,
)
result = timeline.split_long_segments(max_chars=10)
assert result.segment_count == 3
assert len(result.segments[0].text) == 10
def test_time_proportional_split(self):
timeline = SubtitleTimeline(
segments=[
SubtitleSegment(
text="你好世界,这是一段测试文字。用来验证时间比例是否正确。",
start=0,
end=10.0,
)
],
total_duration=10.0,
)
result = timeline.split_long_segments(max_chars=10)
# 所有片段时间加起来应该等于总时长
total_time = sum(s.duration for s in result.segments)
assert total_time == pytest.approx(10.0, abs=0.1)
class TestTextSplitByPunctuation:
def test_basic_split(self):
parts = SubtitleTimeline._split_text_by_punctuation("你好世界!这是测试。", max_chars=10)
assert len(parts) == 2
assert parts[0] == "你好世界!"
assert parts[1] == "这是测试。"
def test_no_punctuation_hard_split(self):
parts = SubtitleTimeline._split_text_by_punctuation("一二三四五六七八九十一二三四五六七八九十", max_chars=10)
assert len(parts) == 2
assert len(parts[0]) == 10
def test_short_text_no_split(self):
parts = SubtitleTimeline._split_text_by_punctuation("你好世界", max_chars=10)
assert len(parts) == 1
assert parts[0] == "你好世界"
+231
View File
@@ -7,18 +7,24 @@ from unittest.mock import Mock
import pytest
from packages.application.template.commands import (
CopyTemplateCommand,
CreateCategoryCommand,
CreateTemplateCommand,
ListTemplatesFilter,
SegmentCommand,
UpdateTemplateCommand,
ValidateTemplateCommand,
)
from packages.application.template.use_cases import (
CopyTemplateUseCase,
CountTemplatesUseCase,
CreateCategoryUseCase,
CreateTemplateUseCase,
DeleteTemplateUseCase,
GetTemplateUsageUseCase,
GetTemplateUseCase,
ListCategoriesUseCase,
ListTagsUseCase,
ListTemplatesUseCase,
NotFoundError,
UpdateTemplateUseCase,
@@ -44,6 +50,9 @@ def _make_repo():
repo.create_category = Mock()
repo.get_category = Mock(return_value=None)
repo.delete_category = Mock(return_value=False)
repo.copy_template = Mock()
repo.list_tags = Mock(return_value=[])
repo.get_usage_count = Mock(return_value=0)
return repo
@@ -442,3 +451,225 @@ class TestGetTemplateUseCase:
result = use_case.execute("nonexistent", "user-001")
assert result is None
# ── CopyTemplateUseCase ──
class TestCopyTemplateUseCase:
@pytest.fixture
def repo(self):
repo = _make_repo()
source = _make_template(
id="tmpl-src",
name="源模板",
segments=[
TemplateSegment(
id="seg-1",
template_id="tmpl-src",
segment_order=0,
duration_min=5.0,
duration_max=10.0,
material_type=None,
),
],
)
repo.get = Mock(return_value=source)
def _copy_side_effect(template_id, user_id, new_name):
return _make_template(
id="tmpl-copied",
user_id=user_id,
name=new_name,
segments=[
TemplateSegment(
id="seg-copied",
template_id="tmpl-copied",
segment_order=0,
duration_min=5.0,
duration_max=10.0,
material_type=None,
)
],
)
repo.copy_template = Mock(side_effect=_copy_side_effect)
return repo
@pytest.fixture
def use_case(self, repo):
return CopyTemplateUseCase(repo)
def test_copy_success(self, use_case, repo):
command = CopyTemplateCommand(
template_id="tmpl-src",
user_id="user-001",
new_name="复制的模板",
)
result = use_case.execute(command)
assert result.id == "tmpl-copied"
assert result.name == "复制的模板"
assert len(result.segments) == 1
repo.copy_template.assert_called_once_with(
"tmpl-src",
"user-001",
"复制的模板",
)
def test_copy_not_found_raises(self, use_case, repo):
repo.get = Mock(return_value=None)
command = CopyTemplateCommand(
template_id="tmpl-nonexist",
user_id="user-001",
new_name="新名字",
)
with pytest.raises(NotFoundError):
use_case.execute(command)
def test_copy_empty_name_raises(self, use_case, repo):
command = CopyTemplateCommand(
template_id="tmpl-src",
user_id="user-001",
new_name=" ",
)
with pytest.raises(ValidationError):
use_case.execute(command)
# ── ListTemplatesUseCase (filter) ──
class TestListTemplatesUseCaseWithFilter:
def test_list_with_category_filter(self):
repo = _make_repo()
repo.list_by_user = Mock(return_value=[])
use_case = ListTemplatesUseCase(repo)
f = ListTemplatesFilter(category="vlog")
use_case.execute("user-001", skip=0, limit=10, filter=f)
repo.list_by_user.assert_called_once()
call_kwargs = repo.list_by_user.call_args
assert call_kwargs[1]["category"] == "vlog"
def test_list_with_tag_filter(self):
repo = _make_repo()
repo.list_by_user = Mock(return_value=[])
use_case = ListTemplatesUseCase(repo)
f = ListTemplatesFilter(tag="热门")
use_case.execute("user-001", filter=f)
call_kwargs = repo.list_by_user.call_args
assert call_kwargs[1]["tag"] == "热门"
def test_list_with_keyword_filter(self):
repo = _make_repo()
repo.list_by_user = Mock(return_value=[])
use_case = ListTemplatesUseCase(repo)
f = ListTemplatesFilter(keyword="vlog")
use_case.execute("user-001", filter=f)
call_kwargs = repo.list_by_user.call_args
assert call_kwargs[1]["keyword"] == "vlog"
def test_list_with_mode_filter(self):
repo = _make_repo()
repo.list_by_user = Mock(return_value=[])
use_case = ListTemplatesUseCase(repo)
f = ListTemplatesFilter(mode="one_take")
use_case.execute("user-001", filter=f)
call_kwargs = repo.list_by_user.call_args
assert call_kwargs[1]["mode"] == "one_take"
def test_list_without_filter_uses_defaults(self):
repo = _make_repo()
repo.list_by_user = Mock(return_value=[])
use_case = ListTemplatesUseCase(repo)
use_case.execute("user-001", skip=0, limit=50)
call_args = repo.list_by_user.call_args
assert call_args[0][0] == "user-001"
assert call_args[1]["skip"] == 0
assert call_args[1]["limit"] == 50
# ── CountTemplatesUseCase ──
class TestCountTemplatesUseCase:
def test_count_without_filter(self):
repo = _make_repo()
repo.count_by_user = Mock(return_value=5)
use_case = CountTemplatesUseCase(repo)
result = use_case.execute("user-001")
assert result == 5
repo.count_by_user.assert_called_once_with("user-001")
def test_count_with_filter(self):
repo = _make_repo()
repo.count_by_user = Mock(return_value=2)
use_case = CountTemplatesUseCase(repo)
f = ListTemplatesFilter(category="vlog", tag="热门")
result = use_case.execute("user-001", filter=f)
assert result == 2
call_kwargs = repo.count_by_user.call_args
assert call_kwargs[1]["category"] == "vlog"
assert call_kwargs[1]["tag"] == "热门"
# ── ListTagsUseCase ──
class TestListTagsUseCase:
def test_list_tags_returns_sorted(self):
repo = _make_repo()
repo.list_tags = Mock(return_value=["vlog", "热门", "教程"])
use_case = ListTagsUseCase(repo)
result = use_case.execute("user-001")
assert result == ["vlog", "热门", "教程"]
repo.list_tags.assert_called_once_with("user-001")
def test_list_tags_empty(self):
repo = _make_repo()
repo.list_tags = Mock(return_value=[])
use_case = ListTagsUseCase(repo)
result = use_case.execute("user-001")
assert result == []
# ── GetTemplateUsageUseCase ──
class TestGetTemplateUsageUseCase:
def test_get_usage_count(self):
repo = _make_repo()
repo.get_usage_count = Mock(return_value=3)
use_case = GetTemplateUsageUseCase(repo)
result = use_case.execute("tmpl-001")
assert result == 3
repo.get_usage_count.assert_called_once_with("tmpl-001")
def test_get_usage_zero(self):
repo = _make_repo()
repo.get_usage_count = Mock(return_value=0)
use_case = GetTemplateUsageUseCase(repo)
result = use_case.execute("tmpl-001")
assert result == 0
+484
View File
@@ -0,0 +1,484 @@
"""转场特效引擎单测 — Phase 8 智能增强."""
from __future__ import annotations
import pytest
from video_processing.transition_engine import (
CUT_TRANSITION,
DEFAULT_TRANSITION_DURATION,
MAX_TRANSITION_DURATION,
MIN_TRANSITION_DURATION,
TransitionConfig,
TransitionEngine,
TransitionType,
_normalize_transition_name,
)
# ── TransitionType 枚举测试 ──────────────────────────────────────────────────
class TestTransitionType:
"""TransitionType 枚举测试."""
def test_all_supported_count(self):
"""支持的转场类型数量(不含cut."""
supported = TransitionType.all_supported()
# 至少 8 种:fade, dissolve, slide*4, zoom, wipe*4, circlecrop, rectcrop
assert len(supported) >= 8
assert "fade" in supported
assert "dissolve" in supported
assert "zoom" in supported
assert "circlecrop" in supported
assert "rectcrop" in supported
def test_slide_directions(self):
"""四个方向的滑入转场都支持."""
assert TransitionType.is_supported("slideleft")
assert TransitionType.is_supported("slideright")
assert TransitionType.is_supported("slideup")
assert TransitionType.is_supported("slidedown")
def test_wipe_directions(self):
"""四个方向的擦除转场都支持."""
assert TransitionType.is_supported("wipeleft")
assert TransitionType.is_supported("wiperight")
assert TransitionType.is_supported("wipeup")
assert TransitionType.is_supported("wipedown")
def test_is_supported_case_insensitive(self):
"""大小写不敏感."""
assert TransitionType.is_supported("FADE")
assert TransitionType.is_supported("Fade")
assert TransitionType.is_supported("fade")
def test_is_supported_with_underscores(self):
"""下划线不影响判断."""
assert TransitionType.is_supported("slide_left")
assert TransitionType.is_supported("slide-left")
def test_is_supported_aliases(self):
"""别名支持."""
assert TransitionType.is_supported("crossfade")
assert TransitionType.is_supported("dissolve")
assert TransitionType.is_supported("zoomin")
assert TransitionType.is_supported("wipe")
def test_unsupported_transition(self):
"""不支持的转场返回 False."""
assert not TransitionType.is_supported("nonexistent_effect")
assert not TransitionType.is_supported("random_stuff")
assert not TransitionType.is_supported("")
def test_cut_not_in_supported(self):
"""硬切不在"支持的转场效果"列表中(它不是特效)."""
supported = TransitionType.all_supported()
assert "cut" not in supported
# ── 名称标准化测试 ────────────────────────────────────────────────────────────
class TestNormalizeTransitionName:
"""名称标准化函数测试."""
def test_lowercase(self):
"""大写转小写."""
assert _normalize_transition_name("FADE") == "fade"
assert _normalize_transition_name("Fade") == "fade"
def test_remove_underscores(self):
"""移除下划线."""
assert _normalize_transition_name("slide_left") == "slideleft"
assert _normalize_transition_name("slide_up") == "slideup"
def test_remove_hyphens(self):
"""移除连字符."""
assert _normalize_transition_name("slide-left") == "slideleft"
def test_mixed(self):
"""混合情况."""
assert _normalize_transition_name("Slide_Left") == "slideleft"
assert _normalize_transition_name("FADE-IN") == "fadein"
# ── TransitionConfig 测试 ────────────────────────────────────────────────────
class TestTransitionConfig:
"""TransitionConfig 配置解析测试."""
# ── 默认值 ──
def test_default_config(self):
"""默认配置是硬切."""
cfg = TransitionConfig.parse()
assert cfg.effect == CUT_TRANSITION
assert cfg.duration == DEFAULT_TRANSITION_DURATION
assert cfg.is_cut is True
def test_none_effect(self):
"""None effect 降级为 cut."""
cfg = TransitionConfig.parse(effect=None)
assert cfg.effect == CUT_TRANSITION
assert cfg.is_cut is True
def test_empty_effect(self):
"""空字符串 effect 降级为 cut."""
cfg = TransitionConfig.parse(effect="")
assert cfg.effect == CUT_TRANSITION
assert cfg.is_cut is True
# ── 有效转场类型 ──
def test_fade_effect(self):
"""fade 转场."""
cfg = TransitionConfig.parse(effect="fade")
assert cfg.effect == "fade"
assert cfg.is_cut is False
assert cfg.ffmpeg_transition == "fade"
def test_dissolve_effect(self):
"""dissolve 转场."""
cfg = TransitionConfig.parse(effect="dissolve")
assert cfg.effect == "dissolve"
assert cfg.ffmpeg_transition == "dissolve"
def test_zoom_effect(self):
"""zoom 转场 → FFmpeg zoomin."""
cfg = TransitionConfig.parse(effect="zoom")
assert cfg.effect == "zoom"
assert cfg.ffmpeg_transition == "zoomin"
def test_slide_left_alias(self):
"""slide_left 别名."""
cfg = TransitionConfig.parse(effect="slide_left")
assert cfg.effect == "slideleft"
assert cfg.ffmpeg_transition == "slideleft"
def test_wipe_alias(self):
"""wipe 别名 → 默认向左擦."""
cfg = TransitionConfig.parse(effect="wipe")
assert cfg.effect == "wipeleft"
assert cfg.ffmpeg_transition == "wipeleft"
def test_circlecrop_effect(self):
"""圆形扩散转场."""
cfg = TransitionConfig.parse(effect="circlecrop")
assert cfg.effect == "circlecrop"
assert cfg.ffmpeg_transition == "circlecrop"
def test_rectcrop_effect(self):
"""矩形扩散转场."""
cfg = TransitionConfig.parse(effect="rectcrop")
assert cfg.effect == "rectcrop"
assert cfg.ffmpeg_transition == "rectcrop"
# ── 降级策略 ──
def test_unsupported_fallback_to_cut(self):
"""不支持的转场自动降级为硬切,不阻断渲染."""
cfg = TransitionConfig.parse(effect="nonexistent_effect")
assert cfg.effect == CUT_TRANSITION
assert cfg.is_cut is True
def test_unsupported_whitespace_fallback(self):
"""带空格的不支持转场也降级."""
cfg = TransitionConfig.parse(effect=" bad effect ")
assert cfg.effect == CUT_TRANSITION
# ── 时长边界校验 ──
def test_default_duration(self):
"""默认时长 0.5s."""
cfg = TransitionConfig.parse(effect="fade")
assert cfg.duration == 0.5
def test_duration_within_range(self):
"""正常范围内的时长."""
cfg = TransitionConfig.parse(effect="fade", duration=1.0)
assert cfg.duration == 1.0
def test_duration_min_boundary(self):
"""最小值边界."""
cfg = TransitionConfig.parse(effect="fade", duration=MIN_TRANSITION_DURATION)
assert cfg.duration == MIN_TRANSITION_DURATION
def test_duration_max_boundary(self):
"""最大值边界."""
cfg = TransitionConfig.parse(effect="fade", duration=MAX_TRANSITION_DURATION)
assert cfg.duration == MAX_TRANSITION_DURATION
def test_duration_below_min_clamped(self):
"""低于最小值的时长被钳制."""
cfg = TransitionConfig.parse(effect="fade", duration=0.1)
assert cfg.duration == MIN_TRANSITION_DURATION
assert cfg.duration >= MIN_TRANSITION_DURATION
def test_duration_above_max_clamped(self):
"""高于最大值的时长被钳制."""
cfg = TransitionConfig.parse(effect="fade", duration=5.0)
assert cfg.duration == MAX_TRANSITION_DURATION
assert cfg.duration <= MAX_TRANSITION_DURATION
def test_duration_zero_default_for_effect(self):
"""有转场效果但 duration=0 时使用默认值."""
# 0.0 会被当作小于最小值钳制到 0.3
cfg = TransitionConfig.parse(effect="fade", duration=0.0)
assert cfg.duration == MIN_TRANSITION_DURATION
def test_duration_negative_clamped(self):
"""负时长被钳制到最小值."""
cfg = TransitionConfig.parse(effect="fade", duration=-1.0)
assert cfg.duration == MIN_TRANSITION_DURATION
def test_duration_none_uses_default(self):
"""None duration 使用默认值."""
cfg = TransitionConfig.parse(effect="fade", duration=None)
assert cfg.duration == DEFAULT_TRANSITION_DURATION
def test_duration_invalid_type(self):
"""无效类型的时长使用默认值."""
cfg = TransitionConfig.parse(effect="fade", duration="abc") # type: ignore
assert cfg.duration == DEFAULT_TRANSITION_DURATION
# ── cut 的 ffmpeg_transition ──
def test_cut_ffmpeg_transition_empty(self):
"""硬切没有对应的 FFmpeg xfade transition."""
cfg = TransitionConfig.parse(effect="cut")
assert cfg.ffmpeg_transition == ""
# ── TransitionEngine 测试 ────────────────────────────────────────────────────
class TestTransitionEngine:
"""TransitionEngine 转场引擎测试."""
def test_default_engine(self):
"""默认引擎初始化."""
engine = TransitionEngine()
assert engine is not None
def test_custom_default_duration(self):
"""自定义默认时长."""
engine = TransitionEngine(default_duration=1.0)
cfg = engine.resolve_config(effect="fade")
assert cfg.duration == 1.0
def test_resolve_config_fade(self):
"""解析 fade 配置."""
engine = TransitionEngine()
cfg = engine.resolve_config(effect="fade", duration=0.8)
assert cfg.effect == "fade"
assert cfg.duration == 0.8
def test_resolve_config_fallback(self):
"""不支持的转场降级."""
engine = TransitionEngine()
cfg = engine.resolve_config(effect="unknown_effect")
assert cfg.effect == CUT_TRANSITION
assert cfg.is_cut is True
def test_resolve_config_duration_clamp(self):
"""时长边界钳制."""
engine = TransitionEngine()
cfg = engine.resolve_config(effect="fade", duration=3.0)
assert cfg.duration == MAX_TRANSITION_DURATION
# ── 批量解析 ──
def test_resolve_clip_transitions_all_valid(self):
"""批量解析全部有效转场."""
engine = TransitionEngine()
configs = engine.resolve_clip_transitions(["cut", "fade", "dissolve", "slideleft"])
assert len(configs) == 4
assert configs[0].effect == "cut"
assert configs[0].is_cut is True
assert configs[1].effect == "fade"
assert configs[2].effect == "dissolve"
assert configs[3].effect == "slideleft"
def test_resolve_clip_transitions_with_fallback(self):
"""批量解析包含不支持的转场,自动降级."""
engine = TransitionEngine()
configs = engine.resolve_clip_transitions(["fade", "bad_effect", "dissolve", "worse_effect"])
assert len(configs) == 4
assert configs[0].effect == "fade"
assert configs[1].effect == "cut" # 降级
assert configs[2].effect == "dissolve"
assert configs[3].effect == "cut" # 降级
def test_resolve_clip_transitions_with_durations(self):
"""带时长校验的批量解析(转场时长不超过片段时长的一半)."""
engine = TransitionEngine(default_duration=1.0)
# 片段只有 1.0s,转场时长被限制在 0.5s
configs = engine.resolve_clip_transitions(
["fade", "dissolve"],
clip_durations=[1.0, 1.0],
)
assert len(configs) == 2
# 1.0s 默认值超过了片段时长的一半 (0.5s),所以被钳制
assert configs[0].duration <= 0.5
assert configs[1].duration <= 0.5
def test_resolve_clip_transitions_short_clip_min_bound(self):
"""超短片段的转场时长至少为最小值."""
engine = TransitionEngine()
configs = engine.resolve_clip_transitions(
["fade"],
clip_durations=[0.1], # 极短片段
)
assert len(configs) == 1
# 0.1 * 0.5 = 0.05 < MIN_TRANSITION_DURATION,所以用最小值
assert configs[0].duration == MIN_TRANSITION_DURATION
# ── xfade 滤镜链构建 ──
def test_build_xfade_single_clip(self):
"""单 clip 直接 copy."""
engine = TransitionEngine()
filter_str, total_dur = engine.build_xfade_chain(
clip_durations=[5.0],
clip_video_labels=["v0"],
transitions=["cut"],
output_label="outv",
)
assert "copy" in filter_str
assert "[outv]" in filter_str
assert total_dur == pytest.approx(5.0, abs=0.01)
def test_build_xfade_two_clips_fade(self):
"""两个 clip 之间 fade 转场."""
engine = TransitionEngine()
filter_str, total_dur = engine.build_xfade_chain(
clip_durations=[3.0, 4.0],
clip_video_labels=["v0", "v1"],
transitions=["cut", "fade"],
output_label="outv",
)
assert "xfade" in filter_str
assert "transition=fade" in filter_str
# 总时长 = 3 + 4 - transition_duration (0.5) = 6.5
assert total_dur == pytest.approx(6.5, abs=0.1)
def test_build_xfade_three_clips_mixed(self):
"""三个 clip 混合转场."""
engine = TransitionEngine()
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"],
output_label="outv",
)
assert "xfade" in filter_str
assert "transition=fade" in filter_str
assert "transition=dissolve" in filter_str
# 总时长 ≈ 3 + 4 + 5 - 2 * 0.5 = 11.0
assert total_dur == pytest.approx(11.0, abs=0.2)
def test_build_xfade_with_custom_duration(self):
"""自定义转场时长."""
engine = TransitionEngine(default_duration=0.5)
filter_str, total_dur = engine.build_xfade_chain(
clip_durations=[3.0, 4.0],
clip_video_labels=["v0", "v1"],
transitions=["cut", "fade"],
transition_duration=1.0,
output_label="outv",
)
assert "xfade" in filter_str
# 总时长 = 3 + 4 - 1.0 = 6.0
assert total_dur == pytest.approx(6.0, abs=0.1)
def test_build_xfade_zoom_transition(self):
"""zoom 转场滤镜构建."""
engine = TransitionEngine()
filter_str, _ = engine.build_xfade_chain(
clip_durations=[3.0, 4.0],
clip_video_labels=["v0", "v1"],
transitions=["cut", "zoom"],
)
assert "xfade" in filter_str
assert "transition=zoomin" in filter_str # zoom → zoomin
def test_build_xfade_slide_directions(self):
"""四个方向的滑入转场."""
engine = TransitionEngine()
for direction in ["slideleft", "slideright", "slideup", "slidedown"]:
filter_str, _ = engine.build_xfade_chain(
clip_durations=[3.0, 4.0],
clip_video_labels=["v0", "v1"],
transitions=["cut", direction],
)
assert f"transition={direction}" in filter_str
def test_build_xfade_fallback_transition(self):
"""不支持的转场降级后构建(降级为cut,等效于极短fade)."""
engine = TransitionEngine()
# bad_effect 降级为 cutcut 使用极短转场
filter_str, _ = engine.build_xfade_chain(
clip_durations=[3.0, 4.0],
clip_video_labels=["v0", "v1"],
transitions=["cut", "bad_effect"],
)
# 降级后是 cutcut 会被 xfade 层映射为 fade(因为 cut 不在 map 里)
# 但时长会很短,所以仍然有 xfade
assert "xfade" in filter_str
# ── 支持的转场列表 ──
def test_supported_transitions_list(self):
"""获取支持的转场列表(给 API 用)."""
transitions = TransitionEngine.supported_transitions()
assert len(transitions) >= 10 # cut + 至少 9 种特效
# 检查结构
for t in transitions:
assert "name" in t
assert "display_name" in t
assert "category" in t
# 检查分类
names = [t["name"] for t in transitions]
assert "cut" in names
assert "fade" in names
assert "zoom" in names
assert "circlecrop" in names
# ── 集成测试:与 UnifiedRenderService 协作 ────────────────────────────────────
class TestTransitionIntegration:
"""转场引擎与统一渲染服务的集成测试."""
def test_unified_render_service_has_transition_engine(self):
"""UnifiedRenderService 内部有 TransitionEngine 实例."""
from pathlib import Path
from video_processing.unified_render_service import UnifiedRenderService
# 构造最小化的服务实例
service = UnifiedRenderService(
plan=None,
clips=[],
asset_path_map={},
work_dir=Path("/tmp"),
)
assert hasattr(service, "_transition_engine")
assert isinstance(service._transition_engine, TransitionEngine)
def test_resolved_clip_has_transition_duration(self):
"""ResolvedClip 有 transition_duration 字段."""
from video_processing.unified_render_service import ResolvedClip
rc = ResolvedClip(
clip_id="test",
asset_id="asset1",
local_path=__file__, # 随便一个存在的路径
clip_type="main",
order=0,
transition_effect="fade",
transition_duration=0.8,
)
assert rc.transition_duration == 0.8
assert rc.transition_effect == "fade"
+268
View File
@@ -0,0 +1,268 @@
"""裁剪引擎单元测试."""
import sys
import unittest
from pathlib import Path
# 确保 apps/worker 在路径中
sys.path.insert(0, str(Path(__file__).parent.parent.parent / "apps" / "worker"))
from video_processing.trim_engine import (
MIN_TRIM_DURATION,
TrimConfig,
TrimEngine,
TrimSegment,
extract_trim_from_clip_config,
)
class TestTrimConfig(unittest.TestCase):
"""TrimConfig 单元测试."""
def test_from_dict_none(self):
"""空字典返回 None(不裁剪)."""
self.assertIsNone(TrimConfig.from_dict(None))
self.assertIsNone(TrimConfig.from_dict({}))
def test_from_dict_with_start(self):
"""只有 start_time."""
cfg = TrimConfig.from_dict({"start_time": 5.0})
self.assertIsNotNone(cfg)
self.assertEqual(cfg.start_time, 5.0)
self.assertEqual(cfg.end_time, 0.0)
self.assertEqual(cfg.duration, 0.0)
def test_from_dict_with_duration(self):
"""只有 duration."""
cfg = TrimConfig.from_dict({"duration": 10.0})
self.assertIsNotNone(cfg)
self.assertEqual(cfg.start_time, 0.0)
self.assertEqual(cfg.duration, 10.0)
def test_resolve_start_and_end(self):
"""start + end 推导 duration."""
cfg = TrimConfig(start_time=5.0, end_time=15.0)
resolved = cfg.validate_and_resolve(asset_duration=30.0)
self.assertEqual(resolved.start_time, 5.0)
self.assertEqual(resolved.end_time, 15.0)
self.assertAlmostEqual(resolved.duration, 10.0, places=3)
self.assertTrue(resolved.is_valid)
def test_resolve_start_and_duration(self):
"""start + duration 推导 end."""
cfg = TrimConfig(start_time=5.0, duration=10.0)
resolved = cfg.validate_and_resolve(asset_duration=30.0)
self.assertEqual(resolved.start_time, 5.0)
self.assertAlmostEqual(resolved.end_time, 15.0, places=3)
self.assertEqual(resolved.duration, 10.0)
def test_resolve_end_and_duration(self):
"""end + duration 推导 start."""
cfg = TrimConfig(end_time=20.0, duration=8.0)
resolved = cfg.validate_and_resolve(asset_duration=30.0)
self.assertAlmostEqual(resolved.start_time, 12.0, places=3)
self.assertEqual(resolved.end_time, 20.0)
self.assertEqual(resolved.duration, 8.0)
def test_resolve_only_start(self):
"""只有 start → 取到末尾."""
cfg = TrimConfig(start_time=10.0)
resolved = cfg.validate_and_resolve(asset_duration=30.0)
self.assertEqual(resolved.start_time, 10.0)
self.assertEqual(resolved.end_time, 30.0)
self.assertAlmostEqual(resolved.duration, 20.0, places=3)
def test_resolve_only_duration(self):
"""只有 duration → 从开头取."""
cfg = TrimConfig(duration=15.0)
resolved = cfg.validate_and_resolve(asset_duration=30.0)
self.assertEqual(resolved.start_time, 0.0)
self.assertAlmostEqual(resolved.end_time, 15.0, places=3)
self.assertEqual(resolved.duration, 15.0)
def test_boundary_clamp_end(self):
"""end 超出素材时长 → 钳制."""
cfg = TrimConfig(start_time=5.0, duration=30.0)
resolved = cfg.validate_and_resolve(asset_duration=20.0)
self.assertEqual(resolved.start_time, 5.0)
self.assertEqual(resolved.end_time, 20.0)
self.assertAlmostEqual(resolved.duration, 15.0, places=3)
def test_boundary_clamp_start_negative(self):
"""start 为负 → 钳制到 0."""
cfg = TrimConfig(start_time=-5.0, duration=10.0)
resolved = cfg.validate_and_resolve(asset_duration=30.0)
self.assertEqual(resolved.start_time, 0.0)
self.assertAlmostEqual(resolved.end_time, 10.0, places=3)
self.assertEqual(resolved.duration, 10.0)
def test_boundary_start_past_end(self):
"""start 超过素材总时长 → 钳制到末尾最小片段."""
cfg = TrimConfig(start_time=50.0, duration=5.0)
resolved = cfg.validate_and_resolve(asset_duration=30.0)
self.assertTrue(resolved.start_time < 30.0)
self.assertEqual(resolved.end_time, 30.0)
self.assertTrue(resolved.duration >= MIN_TRIM_DURATION)
def test_invalid_end_before_start(self):
"""end <= start → 无效."""
cfg = TrimConfig(start_time=15.0, end_time=10.0)
resolved = cfg.validate_and_resolve(asset_duration=30.0)
self.assertFalse(resolved.is_valid)
def test_zero_duration_invalid(self):
"""duration 为 0 → 无效."""
cfg = TrimConfig(start_time=5.0, duration=0.0)
resolved = cfg.validate_and_resolve(asset_duration=30.0)
# 只有 start 没有 duration → 会被推导为取到末尾
self.assertTrue(resolved.is_valid)
self.assertEqual(resolved.end_time, 30.0)
def test_is_noop(self):
"""is_noop 判断."""
noop = TrimConfig(start_time=0.0, end_time=0.0, duration=0.0)
self.assertTrue(noop.is_noop)
not_noop = TrimConfig(start_time=5.0, duration=10.0)
self.assertFalse(not_noop.is_noop)
def test_zero_asset_duration(self):
"""素材时长为 0 → 不裁剪."""
cfg = TrimConfig(start_time=5.0, duration=10.0)
resolved = cfg.validate_and_resolve(asset_duration=0.0)
self.assertTrue(resolved.is_noop)
def test_all_three_params_use_start_duration(self):
"""三个参数都给了 → 以 start + duration 为准."""
cfg = TrimConfig(start_time=5.0, end_time=20.0, duration=8.0)
resolved = cfg.validate_and_resolve(asset_duration=30.0)
# validate_and_resolve 中 start+end 优先于 start+duration
# 因为先检查的是 start>0 and end>0
self.assertAlmostEqual(resolved.duration, 15.0, places=3)
class TestTrimEngine(unittest.TestCase):
"""TrimEngine 单元测试."""
def test_build_video_trim_with_start_and_duration(self):
"""视频裁剪:start + duration."""
trim = TrimConfig(start_time=10.0, duration=5.0)
result = TrimEngine.build_video_trim_filter("[0:v]", trim, "[v0]")
self.assertIn("trim=start=10.000:duration=5.000", result)
self.assertIn("setpts=PTS-STARTPTS", result)
self.assertTrue(result.startswith("[0:v]"))
self.assertTrue(result.endswith("[v0]"))
def test_build_video_trim_duration_only(self):
"""视频裁剪:只有 duration."""
trim = TrimConfig(start_time=0.0, duration=8.0)
result = TrimEngine.build_video_trim_filter("[0:v]", trim, "[v0]")
self.assertIn("trim=duration=8.000", result)
self.assertNotIn("start=", result.split("setpts")[0])
def test_build_audio_trim_with_start(self):
"""音频裁剪:start + duration."""
trim = TrimConfig(start_time=3.0, duration=7.0)
result = TrimEngine.build_audio_trim_filter("[0:a]", trim, "[a0]")
self.assertIn("atrim=start=3.000:duration=7.000", result)
self.assertIn("asetpts=PTS-STARTPTS", result)
def test_build_audio_trim_noop(self):
"""音频裁剪:noop."""
trim = TrimConfig(start_time=0.0, end_time=0.0, duration=0.0)
result = TrimEngine.build_audio_trim_filter("[0:a]", trim, "[a0]")
self.assertIn("asetpts=PTS-STARTPTS", result)
self.assertNotIn("atrim=", result)
def test_resolve_segments(self):
"""多段裁剪解析."""
segments = [
TrimSegment(segment_id="s1", trim=TrimConfig(start_time=0.0, duration=5.0), order=0),
TrimSegment(segment_id="s2", trim=TrimConfig(start_time=10.0, duration=5.0), order=1),
TrimSegment(segment_id="s3", trim=TrimConfig(start_time=20.0, duration=5.0), order=2),
]
resolved = TrimEngine.resolve_segments(segments, asset_duration=30.0)
self.assertEqual(len(resolved), 3)
self.assertEqual(resolved[0].segment_id, "s1")
self.assertEqual(resolved[0].trim.duration, 5.0)
self.assertEqual(resolved[1].segment_id, "s2")
self.assertEqual(resolved[1].trim.start_time, 10.0)
self.assertEqual(resolved[2].trim.start_time, 20.0)
def test_resolve_segments_filter_invalid(self):
"""多段裁剪:过滤无效段."""
segments = [
TrimSegment(segment_id="good", trim=TrimConfig(start_time=0.0, duration=5.0), order=0),
TrimSegment(segment_id="bad", trim=TrimConfig(start_time=10.0, end_time=5.0), order=1), # end < start
]
resolved = TrimEngine.resolve_segments(segments, asset_duration=30.0)
self.assertEqual(len(resolved), 1)
self.assertEqual(resolved[0].segment_id, "good")
def test_resolve_segments_boundary_clamp(self):
"""多段裁剪:边界钳制."""
segments = [
TrimSegment(segment_id="s1", trim=TrimConfig(start_time=25.0, duration=10.0), order=0),
]
resolved = TrimEngine.resolve_segments(segments, asset_duration=30.0)
self.assertEqual(len(resolved), 1)
self.assertEqual(resolved[0].trim.end_time, 30.0)
self.assertAlmostEqual(resolved[0].trim.duration, 5.0, places=3)
def test_parse_segments_from_list(self):
"""从 config 解析多段配置."""
config = {
"trim_segments": [
{"segment_id": "intro", "start_time": 0, "duration": 3, "order": 0},
{"segment_id": "highlight", "start_time": 10, "duration": 5, "order": 1},
{"segment_id": "outro", "start_time": 50, "duration": 3, "order": 2},
]
}
segments = TrimEngine.parse_segments_from_config(config)
self.assertEqual(len(segments), 3)
self.assertEqual(segments[0].segment_id, "intro")
self.assertEqual(segments[1].trim.start_time, 10.0)
self.assertEqual(segments[2].trim.duration, 3.0)
def test_parse_segments_empty(self):
"""无裁剪配置 → 空列表."""
self.assertEqual(TrimEngine.parse_segments_from_config(None), [])
self.assertEqual(TrimEngine.parse_segments_from_config({}), [])
def test_parse_single_trim_legacy(self):
"""旧格式单段裁剪(trim_start/trim_duration."""
config = {"trim_start": 5.0, "trim_duration": 10.0}
segments = TrimEngine.parse_segments_from_config(config)
self.assertEqual(len(segments), 1)
self.assertEqual(segments[0].trim.start_time, 5.0)
self.assertEqual(segments[0].trim.duration, 10.0)
class TestExtractTrimFromClipConfig(unittest.TestCase):
"""extract_trim_from_clip_config 单元测试."""
def test_trim_subdict(self):
"""trim 子字典."""
config = {"trim": {"start_time": 5.0, "duration": 10.0}}
result = extract_trim_from_clip_config(config)
self.assertIsNotNone(result)
self.assertEqual(result.start_time, 5.0)
self.assertEqual(result.duration, 10.0)
def test_flat_fields(self):
"""扁平字段(trim_start/trim_end/trim_duration."""
config = {"trim_start": 2.0, "trim_end": 8.0}
result = extract_trim_from_clip_config(config)
self.assertIsNotNone(result)
self.assertEqual(result.start_time, 2.0)
self.assertEqual(result.end_time, 8.0)
def test_no_trim(self):
"""无裁剪配置."""
self.assertIsNone(extract_trim_from_clip_config(None))
self.assertIsNone(extract_trim_from_clip_config({}))
self.assertIsNone(extract_trim_from_clip_config({"other": "value"}))
if __name__ == "__main__":
unittest.main()
+413
View File
@@ -0,0 +1,413 @@
"""TTS 配音引擎单元测试."""
from pathlib import Path
import pytest
from apps.worker.video_processing.tts_engine import TtsEngine, VoiceoverResult, VoiceoverSegment
from packages.adapters.tts.mock_tts_service import MockTtsService
from packages.domain.tts_config import TtsConfig
from packages.domain.voice_presets import (
VoiceGender,
VoicePreset,
VoiceStyle,
get_default_voice,
get_voice,
list_voices,
)
from packages.ports.tts_service import TtsError, TtsService
# ─── TtsConfig 配置解析 ────────────────────────────────────
class TestTtsConfig:
def test_default_values(self):
config = TtsConfig()
assert config.enabled is False
assert config.voice_id == ""
assert config.speed == 1.0
assert config.pitch == 0.0
assert config.volume == 0.8
assert config.text == ""
assert config.align_mode == "full"
assert config.overlap_mode == "replace"
def test_parse_none(self):
config = TtsConfig.parse(None)
assert config.enabled is False
def test_parse_empty_dict(self):
config = TtsConfig.parse({})
assert config.enabled is False
def test_parse_enabled(self):
config = TtsConfig.parse({"enabled": True, "voice_id": "female_warm", "text": "你好"})
assert config.enabled is True
assert config.voice_id == "female_warm"
assert config.text == "你好"
def test_parse_not_enabled_ignores_other_fields(self):
config = TtsConfig.parse({"enabled": False, "voice_id": "test", "speed": 2.0})
assert config.enabled is False
assert config.voice_id == ""
assert config.speed == 1.0
def test_parse_speed_boundary(self):
config = TtsConfig.parse({"enabled": True, "speed": 0.1})
assert config.speed == 0.5
config = TtsConfig.parse({"enabled": True, "speed": 5.0})
assert config.speed == 2.0
def test_parse_pitch_boundary(self):
config = TtsConfig.parse({"enabled": True, "pitch": -20})
assert config.pitch == -12
config = TtsConfig.parse({"enabled": True, "pitch": 20})
assert config.pitch == 12
def test_parse_volume_boundary(self):
config = TtsConfig.parse({"enabled": True, "volume": -1.0})
assert config.volume == 0.0
config = TtsConfig.parse({"enabled": True, "volume": 2.0})
assert config.volume == 1.0
def test_parse_invalid_types(self):
config = TtsConfig.parse(
{
"enabled": True,
"speed": "fast",
"pitch": "high",
"volume": "loud",
"voice_id": 123,
"text": 456,
"align_mode": "invalid",
"overlap_mode": "invalid",
}
)
assert config.speed == 1.0
assert config.pitch == 0.0
assert config.volume == 0.8
assert config.voice_id == ""
assert config.text == ""
assert config.align_mode == "full"
assert config.overlap_mode == "replace"
# ─── 预设音色库 ───────────────────────────────────────────
class TestPresetVoices:
def test_list_voices_all(self):
voices = list_voices()
assert len(voices) >= 6
def test_get_voice_existing(self):
voice = get_voice("female_warm")
assert voice is not None
assert voice.voice_id == "female_warm"
assert voice.name == "温暖女声"
assert voice.gender == VoiceGender.FEMALE
def test_get_voice_not_found(self):
assert get_voice("nonexistent") is None
def test_get_default_voice(self):
voice = get_default_voice()
assert voice is not None
assert voice.provider == "mock"
def test_filter_by_gender(self):
female = list_voices(gender="female")
assert len(female) >= 2
for v in female:
assert v.gender == VoiceGender.FEMALE
male = list_voices(gender="male")
assert len(male) >= 2
for v in male:
assert v.gender == VoiceGender.MALE
def test_filter_by_style(self):
stable = list_voices(style="stable")
for v in stable:
assert v.style == VoiceStyle.STABLE
def test_filter_by_keyword(self):
result = list_voices(keyword="女声")
assert len(result) >= 1
for v in result:
assert "" in v.name
def test_voice_fields(self):
voice = get_voice("male_stable")
assert voice is not None
assert voice.name
assert voice.voice_id
assert voice.description
assert voice.sample_rate > 0
# ─── Mock TTS 服务 ────────────────────────────────────────
class TestMockTtsService:
def setup_method(self):
self.service = MockTtsService()
def test_provider_name(self):
assert self.service.provider_name == "mock"
def test_available_voices(self):
voices = self.service.available_voices()
assert len(voices) >= 6
def test_synthesize_success(self, tmp_path):
output = tmp_path / "test.wav"
result = self.service.synthesize(
"测试文本一二三四五",
voice_id="female_warm",
speed=1.0,
output_path=output,
)
assert result == output
assert result.exists()
assert result.stat().st_size > 0
def test_synthesize_different_voices(self, tmp_path):
voices = ["female_warm", "male_stable", "child_cute"]
for vid in voices:
output = tmp_path / f"{vid}.wav"
result = self.service.synthesize("测试", voice_id=vid, output_path=output)
assert result.exists()
def test_synthesize_speed_faster(self, tmp_path):
"""语速快应该时长短."""
out_slow = tmp_path / "slow.wav"
out_fast = tmp_path / "fast.wav"
text = "一二三四五六七八九十"
self.service.synthesize(text, speed=0.5, output_path=out_slow)
self.service.synthesize(text, speed=2.0, output_path=out_fast)
# 快速应该文件更小(时长短)
size_slow = out_slow.stat().st_size
size_fast = out_fast.stat().st_size
assert size_fast < size_slow
def test_synthesize_pitch_changes(self, tmp_path):
output = tmp_path / "high_pitch.wav"
result = self.service.synthesize("测试", voice_id="female_warm", pitch=6, output_path=output)
assert result.exists()
def test_synthesize_empty_text_raises(self):
with pytest.raises(TtsError):
self.service.synthesize("")
def test_synthesize_whitespace_text_raises(self):
with pytest.raises(TtsError):
self.service.synthesize(" ")
def test_estimate_duration(self):
dur = self.service.estimate_duration("一二三四五")
assert dur > 0
assert dur < 10 # 5个字应该少于10秒
def test_estimate_duration_speed(self):
text = "一二三四五六七八九十"
dur_normal = self.service.estimate_duration(text, speed=1.0)
dur_fast = self.service.estimate_duration(text, speed=2.0)
dur_slow = self.service.estimate_duration(text, speed=0.5)
assert dur_fast < dur_normal
assert dur_slow > dur_normal
def test_synthesize_unknown_voice_fallback(self, tmp_path):
output = tmp_path / "fallback.wav"
# 未知音色应该 fallback 到默认音色,不报错
result = self.service.synthesize("测试", voice_id="unknown_voice", output_path=output)
assert result.exists()
# ─── TtsEngine 配音引擎 ──────────────────────────────────
class TestTtsEngine:
def _make_engine(self, tmp_path):
service = MockTtsService()
work_dir = tmp_path / "tts_engine"
return TtsEngine(service, work_dir)
def test_generate_full_voiceover_disabled(self, tmp_path):
engine = self._make_engine(tmp_path)
config = TtsConfig(enabled=False)
result = engine.generate_full_voiceover(config)
assert result.success is False
def test_generate_full_voiceover_empty_text(self, tmp_path):
engine = self._make_engine(tmp_path)
config = TtsConfig(enabled=True, text="")
result = engine.generate_full_voiceover(config)
assert result.success is False
def test_generate_full_voiceover_success(self, tmp_path):
engine = self._make_engine(tmp_path)
config = TtsConfig(
enabled=True,
voice_id="female_warm",
text="这是一段测试配音文本",
)
result = engine.generate_full_voiceover(config)
assert result.success is True
assert len(result.segments) == 1
assert result.total_duration > 0
assert result.segments[0].audio_path is not None
assert result.segments[0].audio_path.exists()
assert result.segments[0].duration > 0
def test_generate_full_voiceover_with_speed(self, tmp_path):
engine = self._make_engine(tmp_path)
config_slow = TtsConfig(
enabled=True,
voice_id="female_warm",
text="测试文本一二三四五六七八九十",
speed=0.5,
)
config_fast = TtsConfig(
enabled=True,
voice_id="female_warm",
text="测试文本一二三四五六七八九十",
speed=2.0,
)
result_slow = engine.generate_full_voiceover(config_slow)
result_fast = engine.generate_full_voiceover(config_fast)
assert result_slow.success
assert result_fast.success
# 慢速时长 > 快速时长
assert result_slow.total_duration > result_fast.total_duration
def test_generate_subtitle_voiceover_empty_subtitles(self, tmp_path):
engine = self._make_engine(tmp_path)
config = TtsConfig(enabled=True, voice_id="female_warm")
result = engine.generate_subtitle_voiceover(config, [])
assert result.success is False
def test_generate_subtitle_voiceover_success(self, tmp_path):
engine = self._make_engine(tmp_path)
config = TtsConfig(enabled=True, voice_id="female_warm", align_mode="subtitle")
subtitles = [
{"text": "大家好", "start_time": 0, "end_time": 2},
{"text": "欢迎观看", "start_time": 2, "end_time": 4},
{"text": "今天的视频", "start_time": 4, "end_time": 6},
]
result = engine.generate_subtitle_voiceover(config, subtitles)
assert result.success is True
assert len(result.segments) == 3
# 每个片段的 start_time 应该对应字幕的开始时间
assert result.segments[0].start_time == 0
assert result.segments[1].start_time == 2
assert result.segments[2].start_time == 4
for seg in result.segments:
assert seg.audio_path is not None
assert seg.audio_path.exists()
assert seg.duration > 0
def test_generate_subtitle_voiceover_skips_empty(self, tmp_path):
engine = self._make_engine(tmp_path)
config = TtsConfig(enabled=True, voice_id="female_warm")
subtitles = [
{"text": "有文本", "start_time": 0, "end_time": 1},
{"text": "", "start_time": 1, "end_time": 2},
{"text": "也有文本", "start_time": 2, "end_time": 3},
]
result = engine.generate_subtitle_voiceover(config, subtitles)
assert result.success is True
assert len(result.segments) == 2 # 跳过了空文本
def test_generate_full_voiceover_failure_graceful(self, tmp_path, monkeypatch):
"""失败时优雅降级,不抛出异常."""
engine = self._make_engine(tmp_path)
def failing_synth(*args, **kwargs):
raise TtsError("模拟失败")
monkeypatch.setattr(engine._tts, "synthesize", failing_synth)
config = TtsConfig(enabled=True, voice_id="test", text="测试")
result = engine.generate_full_voiceover(config)
assert result.success is False
assert result.error_message
assert "模拟失败" in result.error_message
def test_generate_subtitle_voiceover_partial_failure(self, tmp_path, monkeypatch):
"""部分片段失败时跳过,其他正常生成."""
engine = self._make_engine(tmp_path)
original_synth = engine._tts.synthesize
call_count = [0]
def sometimes_fail(*args, **kwargs):
call_count[0] += 1
if call_count[0] == 2: # 第2个片段失败
raise TtsError("模拟失败")
return original_synth(*args, **kwargs)
monkeypatch.setattr(engine._tts, "synthesize", sometimes_fail)
config = TtsConfig(enabled=True, voice_id="female_warm")
subtitles = [
{"text": "第一段", "start_time": 0, "end_time": 2},
{"text": "第二段失败", "start_time": 2, "end_time": 4},
{"text": "第三段", "start_time": 4, "end_time": 6},
]
result = engine.generate_subtitle_voiceover(config, subtitles)
# 有部分成功就算成功
assert result.success is True
assert len(result.segments) == 2 # 跳过了失败的第2段
def test_build_audio_mix_filter_empty(self, tmp_path):
engine = self._make_engine(tmp_path)
result = VoiceoverResult(success=False)
filter_str, files = engine.build_audio_mix_filter(result, video_duration=10)
assert filter_str == ""
assert files == []
def test_build_audio_mix_filter_single(self, tmp_path):
engine = self._make_engine(tmp_path)
audio_file = tmp_path / "seg.wav"
audio_file.write_bytes(b"fake")
segment = VoiceoverSegment(
text="test",
start_time=1.0,
end_time=3.0,
audio_path=audio_file,
duration=2.0,
)
result = VoiceoverResult(success=True, segments=[segment], total_duration=3.0)
filter_str, files = engine.build_audio_mix_filter(result, video_duration=10)
assert len(files) == 1
assert "adelay" in filter_str
assert "1000" in filter_str # 1秒 = 1000ms
# ─── VoicePreset 数据类 ──────────────────────────────────
class TestVoicePreset:
def test_create_preset(self):
preset = VoicePreset(
voice_id="test_voice",
name="测试音色",
gender=VoiceGender.MALE,
style=VoiceStyle.NARRATION,
)
assert preset.voice_id == "test_voice"
assert preset.name == "测试音色"
assert preset.gender == VoiceGender.MALE
assert preset.style == VoiceStyle.NARRATION
assert preset.sample_rate == 22050
+216
View File
@@ -0,0 +1,216 @@
"""成片中心新功能单元测试。
覆盖:分页列表、复核状态更新、缩略图更新、批量获取、use case。
"""
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
import pytest
from sqlalchemy import create_engine
from sqlalchemy.orm import sessionmaker
from packages.adapters.sqlalchemy_impl.generated_video_repository import SQLAlchemyGeneratedVideoRepository
from packages.adapters.sqlalchemy_impl.models import Base
from packages.application.generated_videos import (
GetVideosByIdsUseCase,
ListGeneratedVideosPaginatedUseCase,
UpdateVideoReviewStatusUseCase,
)
from packages.domain import GeneratedVideo
def _repository():
engine = create_engine("sqlite:///:memory:")
Base.metadata.create_all(engine)
session = sessionmaker(bind=engine)()
return SQLAlchemyGeneratedVideoRepository(session)
def _create_video(repo, project_id="proj-1", status="completed", review_status="pending_review", idx=1):
video = GeneratedVideo.create(
project_id=project_id,
generation_task_id=f"task-{idx}",
name=f"video-{idx}.mp4",
file_url=f"generated/video-{idx}.mp4",
file_size=1024 * idx,
duration=10.0 * idx,
width=1280,
height=720,
fps=25.0,
)
video.status = status
video.review_status = review_status
repo.create(video)
return video
class TestGeneratedVideoRepository:
"""GeneratedVideoRepository 新方法测试。"""
def test_list_paginated_default(self):
repo = _repository()
for i in range(5):
_create_video(repo, idx=i)
items, total = repo.list_paginated(page=1, page_size=3)
assert total == 5
assert len(items) == 3
# 按 generated_at 倒序,最新的在前
assert items[0].name == "video-4.mp4"
def test_list_paginated_by_project(self):
repo = _repository()
_create_video(repo, project_id="proj-a", idx=1)
_create_video(repo, project_id="proj-a", idx=2)
_create_video(repo, project_id="proj-b", idx=3)
items, total = repo.list_paginated(project_id="proj-a")
assert total == 2
assert all(i.project_id == "proj-a" for i in items)
def test_list_paginated_by_status(self):
repo = _repository()
_create_video(repo, status="completed", idx=1)
_create_video(repo, status="completed", idx=2)
_create_video(repo, status="failed", idx=3)
items, total = repo.list_paginated(status="completed")
assert total == 2
assert all(i.status == "completed" for i in items)
def test_list_paginated_by_review_status(self):
repo = _repository()
_create_video(repo, review_status="pending_review", idx=1)
_create_video(repo, review_status="approved", idx=2)
_create_video(repo, review_status="rejected", idx=3)
items, total = repo.list_paginated(review_status="approved")
assert total == 1
assert items[0].review_status == "approved"
def test_list_paginated_multi_filter(self):
repo = _repository()
_create_video(repo, project_id="p1", status="completed", review_status="approved", idx=1)
_create_video(repo, project_id="p1", status="completed", review_status="pending_review", idx=2)
_create_video(repo, project_id="p2", status="completed", review_status="approved", idx=3)
items, total = repo.list_paginated(project_id="p1", review_status="approved")
assert total == 1
assert items[0].project_id == "p1"
assert items[0].review_status == "approved"
def test_update_review_status(self):
repo = _repository()
video = _create_video(repo, idx=1)
result = repo.update_review_status(video.id, "approved")
assert result is not None
assert result.review_status == "approved"
# 验证持久化
saved = repo.get(video.id)
assert saved.review_status == "approved"
def test_update_review_status_not_found(self):
repo = _repository()
result = repo.update_review_status("nonexistent", "approved")
assert result is None
def test_update_thumbnail(self):
repo = _repository()
video = _create_video(repo, idx=1)
assert video.thumbnail_url is None
ok = repo.update_thumbnail(video.id, "https://oss/thumb.jpg")
assert ok is True
saved = repo.get(video.id)
assert saved.thumbnail_url == "https://oss/thumb.jpg"
def test_update_thumbnail_not_found(self):
repo = _repository()
ok = repo.update_thumbnail("nonexistent", "https://oss/thumb.jpg")
assert ok is False
def test_get_by_ids(self):
repo = _repository()
v1 = _create_video(repo, idx=1)
v2 = _create_video(repo, idx=2)
v3 = _create_video(repo, idx=3)
result = repo.get_by_ids([v1.id, v3.id])
assert len(result) == 2
ids = {v.id for v in result}
assert v1.id in ids
assert v3.id in ids
def test_get_by_ids_empty(self):
repo = _repository()
result = repo.get_by_ids([])
assert result == []
class TestGeneratedVideoUseCases:
"""Use case 层测试。"""
def test_list_paginated_use_case(self):
repo = _repository()
for i in range(10):
_create_video(repo, idx=i)
use_case = ListGeneratedVideosPaginatedUseCase(repo)
items, total = use_case.execute(page=2, page_size=3)
assert total == 10
assert len(items) == 3
def test_list_paginated_use_case_page_clamp(self):
repo = _repository()
use_case = ListGeneratedVideosPaginatedUseCase(repo)
# page < 1 应该被修正为 1
items, total = use_case.execute(page=0, page_size=20)
assert total == 0
def test_list_paginated_use_case_page_size_clamp(self):
repo = _repository()
use_case = ListGeneratedVideosPaginatedUseCase(repo)
# page_size > 100 应该被修正为 20
for i in range(30):
_create_video(repo, idx=i)
items, total = use_case.execute(page=1, page_size=200)
assert total == 30
assert len(items) == 20 # clamp 到默认 20
def test_update_review_status_use_case(self):
repo = _repository()
video = _create_video(repo, idx=1)
use_case = UpdateVideoReviewStatusUseCase(repo)
result = use_case.execute(video.id, "approved")
assert result is not None
assert result.review_status == "approved"
def test_update_review_status_use_case_invalid_status(self):
repo = _repository()
video = _create_video(repo, idx=1)
use_case = UpdateVideoReviewStatusUseCase(repo)
with pytest.raises(ValueError, match="无效的 review_status"):
use_case.execute(video.id, "invalid_status")
def test_update_review_status_use_case_empty_id(self):
repo = _repository()
use_case = UpdateVideoReviewStatusUseCase(repo)
with pytest.raises(ValueError, match="video_id 不能为空"):
use_case.execute("", "approved")
def test_get_by_ids_use_case(self):
repo = _repository()
v1 = _create_video(repo, idx=1)
v2 = _create_video(repo, idx=2)
use_case = GetVideosByIdsUseCase(repo)
result = use_case.execute([v1.id, v2.id])
assert len(result) == 2
+344
View File
@@ -0,0 +1,344 @@
"""水印 + 片头片尾引擎单元测试."""
import sys
import unittest
from pathlib import Path
sys.path.insert(0, str(Path(__file__).parent.parent.parent / "apps" / "worker"))
from video_processing.intro_outro_engine import (
IntroOutroConfig,
IntroOutroEngine,
)
from video_processing.watermark_engine import (
WATERMARK_POSITIONS,
WatermarkConfig,
WatermarkEngine,
)
class TestWatermarkConfig(unittest.TestCase):
"""WatermarkConfig 单元测试."""
def test_from_dict_none_disabled(self):
"""空配置或未启用 → None."""
self.assertIsNone(WatermarkConfig.from_dict(None))
self.assertIsNone(WatermarkConfig.from_dict({}))
self.assertIsNone(WatermarkConfig.from_dict({"enabled": False}))
def test_from_dict_text_mode(self):
"""文字水印模式."""
cfg = WatermarkConfig.from_dict({
"enabled": True,
"mode": "text",
"text": "hello world",
"position": "top_left",
})
self.assertIsNotNone(cfg)
self.assertEqual(cfg.mode, "text")
self.assertEqual(cfg.text, "hello world")
self.assertEqual(cfg.position, "top_left")
def test_from_dict_image_missing_path(self):
"""图片水印缺路径 → None."""
cfg = WatermarkConfig.from_dict({
"enabled": True,
"mode": "image",
})
self.assertIsNone(cfg)
def test_from_dict_text_missing_text(self):
"""文字水印缺文字 → None."""
cfg = WatermarkConfig.from_dict({
"enabled": True,
"mode": "text",
})
self.assertIsNone(cfg)
def test_validate_text_valid(self):
"""文字水印合法配置."""
cfg = WatermarkConfig(
mode="text",
text="test",
position="bottom_right",
)
ok, err = cfg.validate()
self.assertTrue(ok)
self.assertEqual(err, "")
def test_validate_invalid_position(self):
"""非法位置."""
cfg = WatermarkConfig(mode="text", text="test", position="invalid")
ok, err = cfg.validate()
self.assertFalse(ok)
self.assertIn("不支持的位置", err)
def test_validate_opacity_out_of_range(self):
"""透明度超范围."""
cfg = WatermarkConfig(mode="text", text="test", opacity=1.5)
ok, err = cfg.validate()
self.assertFalse(ok)
def test_validate_image_missing_path(self):
"""图片水印缺路径."""
cfg = WatermarkConfig(mode="image")
ok, err = cfg.validate()
self.assertFalse(ok)
class TestWatermarkEnginePosition(unittest.TestCase):
"""水印位置计算单元测试."""
def setUp(self):
self.out_w = 1920
self.out_h = 1080
self.wm_w = 200
self.wm_h = 100
self.mx = 20
self.my = 20
def test_top_left(self):
x, y = WatermarkEngine.calc_position(
"top_left", self.out_w, self.out_h, self.wm_w, self.wm_h, self.mx, self.my
)
self.assertEqual(x, 20)
self.assertEqual(y, 20)
def test_top_center(self):
x, y = WatermarkEngine.calc_position(
"top_center", self.out_w, self.out_h, self.wm_w, self.wm_h, self.mx, self.my
)
self.assertEqual(x, (1920 - 200) // 2)
self.assertEqual(y, 20)
def test_top_right(self):
x, y = WatermarkEngine.calc_position(
"top_right", self.out_w, self.out_h, self.wm_w, self.wm_h, self.mx, self.my
)
self.assertEqual(x, 1920 - 200 - 20)
self.assertEqual(y, 20)
def test_center_left(self):
x, y = WatermarkEngine.calc_position(
"center_left", self.out_w, self.out_h, self.wm_w, self.wm_h, self.mx, self.my
)
self.assertEqual(x, 20)
self.assertEqual(y, (1080 - 100) // 2)
def test_center(self):
x, y = WatermarkEngine.calc_position(
"center", self.out_w, self.out_h, self.wm_w, self.wm_h, self.mx, self.my
)
self.assertEqual(x, (1920 - 200) // 2)
self.assertEqual(y, (1080 - 100) // 2)
def test_center_right(self):
x, y = WatermarkEngine.calc_position(
"center_right", self.out_w, self.out_h, self.wm_w, self.wm_h, self.mx, self.my
)
self.assertEqual(x, 1920 - 200 - 20)
self.assertEqual(y, (1080 - 100) // 2)
def test_bottom_left(self):
x, y = WatermarkEngine.calc_position(
"bottom_left", self.out_w, self.out_h, self.wm_w, self.wm_h, self.mx, self.my
)
self.assertEqual(x, 20)
self.assertEqual(y, 1080 - 100 - 20)
def test_bottom_center(self):
x, y = WatermarkEngine.calc_position(
"bottom_center", self.out_w, self.out_h, self.wm_w, self.wm_h, self.mx, self.my
)
self.assertEqual(x, (1920 - 200) // 2)
self.assertEqual(y, 1080 - 100 - 20)
def test_bottom_right(self):
x, y = WatermarkEngine.calc_position(
"bottom_right", self.out_w, self.out_h, self.wm_w, self.wm_h, self.mx, self.my
)
self.assertEqual(x, 1920 - 200 - 20)
self.assertEqual(y, 1080 - 100 - 20)
def test_default_fallback(self):
"""非法位置默认右下角."""
x, y = WatermarkEngine.calc_position(
"unknown", self.out_w, self.out_h, self.wm_w, self.wm_h, self.mx, self.my
)
self.assertEqual(x, 1920 - 200 - 20)
self.assertEqual(y, 1080 - 100 - 20)
def test_nine_positions_all_present(self):
"""9宫格位置都有定义."""
self.assertEqual(len(WATERMARK_POSITIONS), 9)
class TestWatermarkEngineFilters(unittest.TestCase):
"""水印滤镜构建单元测试."""
def test_text_watermark_filter(self):
"""文字水印滤镜构建."""
cfg = WatermarkConfig(
mode="text",
text="hello",
position="top_left",
font_size=24,
font_color="white",
opacity=0.8,
margin_x=10,
margin_y=10,
)
result = WatermarkEngine.build_text_watermark_filter(
"[in]", "[out]", cfg, 1920, 1080
)
self.assertTrue(result.startswith("[in]drawtext="))
self.assertIn("text='hello'", result)
self.assertIn("fontsize=24", result)
self.assertIn("fontcolor=white@0.8", result)
self.assertTrue(result.endswith("[out]"))
def test_text_watermark_scroll(self):
"""滚动文字水印."""
cfg = WatermarkConfig(
mode="text",
text="scroll",
position="bottom_left",
scroll=True,
scroll_speed=60,
)
result = WatermarkEngine.build_text_watermark_filter(
"[in]", "[out]", cfg, 1920, 1080
)
self.assertIn("mod(60*t", result)
class TestIntroOutroConfig(unittest.TestCase):
"""IntroOutroConfig 单元测试."""
def test_from_dict_disabled(self):
"""未启用 → 空配置."""
cfg = IntroOutroConfig.from_dict(None)
self.assertFalse(cfg.enabled)
self.assertFalse(cfg.has_intro)
self.assertFalse(cfg.has_outro)
def test_from_dict_intro_text(self):
"""文字片头配置."""
cfg = IntroOutroConfig.from_dict({
"enabled": True,
"intro": {
"type": "text",
"title": "欢迎观看",
"subtitle": "精彩内容马上开始",
"duration": 3.0,
"background": "#1a1a2e",
},
})
self.assertTrue(cfg.enabled)
self.assertTrue(cfg.has_intro)
self.assertFalse(cfg.has_outro)
self.assertEqual(cfg.intro_type, "text")
self.assertEqual(cfg.intro_title, "欢迎观看")
self.assertEqual(cfg.intro_duration, 3.0)
def test_from_dict_outro_video(self):
"""视频片尾配置."""
cfg = IntroOutroConfig.from_dict({
"enabled": True,
"outro": {
"type": "video",
"video_path": "/tmp/outro.mp4",
"duration": 5.0,
},
})
self.assertTrue(cfg.has_outro)
self.assertEqual(cfg.outro_type, "video")
self.assertEqual(cfg.outro_video_path, "/tmp/outro.mp4")
def test_validate_valid(self):
"""合法配置."""
cfg = IntroOutroConfig(
enabled=True,
intro_type="text",
intro_title="标题",
intro_duration=3.0,
outro_type="text",
outro_title="片尾",
outro_duration=3.0,
)
ok, err = cfg.validate()
self.assertTrue(ok)
def test_validate_video_intro_missing_path(self):
"""视频片头缺路径."""
cfg = IntroOutroConfig(
enabled=True,
intro_type="video",
intro_duration=3.0,
)
ok, err = cfg.validate()
self.assertFalse(ok)
self.assertIn("video_path", err)
def test_validate_text_intro_missing_title(self):
"""文字片头缺标题."""
cfg = IntroOutroConfig(
enabled=True,
intro_type="text",
intro_duration=3.0,
)
ok, err = cfg.validate()
self.assertFalse(ok)
def test_has_intro_false_when_none(self):
"""type=none 时 has_intro 为 False."""
cfg = IntroOutroConfig(enabled=True, intro_type="none")
self.assertFalse(cfg.has_intro)
def test_has_outro_follow_type(self):
"""follow 类型也算有片尾."""
cfg = IntroOutroConfig(enabled=True, outro_type="follow", outro_title="关注")
self.assertTrue(cfg.has_outro)
class TestIntroOutroEngineConcat(unittest.TestCase):
"""片头片尾拼接单元测试."""
def test_concat_no_intro_outro(self):
"""没有片头片尾 → 直接复制."""
import tempfile
with tempfile.TemporaryDirectory() as tmpdir:
main_video = Path(tmpdir) / "main.mp4"
output = Path(tmpdir) / "output.mp4"
# 创建空文件模拟
main_video.write_bytes(b"fake video data")
result = IntroOutroEngine.concat_with_intro_outro(
main_video, None, None, output
)
self.assertTrue(result)
self.assertTrue(output.exists())
self.assertEqual(main_video.read_bytes(), output.read_bytes())
def test_concat_intro_only_no_file(self):
"""只有片头但文件不存在 → 直接复制主视频."""
import tempfile
with tempfile.TemporaryDirectory() as tmpdir:
main_video = Path(tmpdir) / "main.mp4"
output = Path(tmpdir) / "output.mp4"
main_video.write_bytes(b"fake data")
# intro 路径不存在
intro = Path(tmpdir) / "nonexistent.mp4"
result = IntroOutroEngine.concat_with_intro_outro(
main_video, intro, None, output
)
self.assertTrue(result)
self.assertTrue(output.exists())
if __name__ == "__main__":
unittest.main()