Compare commits

..

6 Commits

Author SHA1 Message Date
xiaoxia b430c6dfe2 chore: remove temporary fix workflow
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 1m43s
CI/CD Pipeline / Frontend Lint (push) Successful in 2m32s
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 1m57s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 1m54s
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Successful in 1m50s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 2m58s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (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 / Build Production Runtime Images (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 / Build & Push Staging (Watchtower auto-deploy) (push) Successful in 5m11s
CI/CD Pipeline / Staging API Integration Tests (push) Has been skipped
CI/CD Pipeline / Staging E2E Tests (push) Has been skipped
2026-07-12 20:45:34 +08:00
CI Bot 423342085c style: fix isort imports + black formatting
CI/CD Pipeline / Validate Code Quality And Tests (push) Successful in 1m57s
CI/CD Pipeline / Frontend Lint (push) Successful in 2m31s
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Failing after 3m22s
CI/CD Pipeline / Staging E2E Tests (push) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (push) Has been skipped
CI/CD Pipeline / Integration Tests (push) Successful in 2m6s
2026-07-12 20:32:08 +08:00
xiaoxia c1784f4232 ci: update fix workflow to also fix isort
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Failing after 27s
Fix Black + iSort Formatting / fix-format (push) Successful in 46s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 1m50s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 2m37s
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
2026-07-12 20:31:20 +08:00
xiaoxia 6a89870ec7 ci: add black fix workflow
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Failing after 35s
Fix Black Formatting / fix-black (push) Successful in 56s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 2m6s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 2m43s
CI/CD Pipeline / Build Production Runtime Images (pull_request) Has been skipped
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (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
2026-07-12 20:26:42 +08:00
xiaoxia f1ba9f8b64 style: fix black formatting for test_render_adapter.py (line-length=120)
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Failing after 25s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 1m52s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 3m5s
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
2026-07-12 20:17:33 +08:00
xiaoxia 70ce5c57e1 style: fix black formatting for test_render_adapter.py
CI/CD Pipeline / Validate Code Quality And Tests (pull_request) Failing after 42s
CI/CD Pipeline / Frontend Lint (pull_request) Successful in 2m33s
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Build Production Runtime Images (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / Integration Tests (pull_request) Successful in 2m9s
2026-07-12 20:04:45 +08:00
92 changed files with 1223 additions and 6716 deletions
+1 -4
View File
@@ -3,7 +3,6 @@
# ==================== 应用配置 ====================
APP_NAME=小虾 SaaS
APP_BASE_URL=http://localhost:3000
APP_ENV=development
# ==================== 数据库配置 ====================
DATABASE_URL=postgresql://xiaoxia_user:your_password@localhost:5432/xiaoxia_saas
@@ -36,8 +35,7 @@ ENVIRONMENT=development
DEBUG=true
# ==================== CORS 配置 ====================
# 逗号分隔的域名列表(Settings 读取 CORS_ORIGINS_RAW
CORS_ORIGINS_RAW=http://localhost:3000,http://localhost:5173
CORS_ORIGINS=["http://localhost:3000","http://localhost:5173"]
# ==================== 阿里云 OSS 配置 ====================
OSS_ENDPOINT=oss-cn-hangzhou.aliyuncs.com
@@ -51,7 +49,6 @@ OSS_BUCKET_NAME=xiaoxia-autocut
# cosyvoice-v3-plus (高质量,系统音色少)
# cosyvoice-v3.5-flash / cosyvoice-v3.5-plus (仅支持克隆/设计音色,无系统音色)
# 音色: v3系列系统音色带 _v3 后缀,如 longxiaochun_v3, longxiaoxia_v3, longanyang (无后缀)
# 注意:COSYVOICE_* 变量由 packages/shared/config.py 的 SharedSettings 读取
COSYVOICE_API_KEY=your-cosyvoice-api-key
COSYVOICE_BASE_URL=https://dashscope.aliyuncs.com/api/v1
COSYVOICE_MODEL=cosyvoice-v3-flash
File diff suppressed because one or more lines are too long
+1
View File
@@ -0,0 +1 @@
"""API application package."""
+1
View File
@@ -0,0 +1 @@
"""API package."""
-10
View File
@@ -8,12 +8,10 @@ from app.api.routes.dashboard import router as dashboard_router
from app.api.routes.duplication import router as duplication_router
from app.api.routes.edit_plans import router as edit_plans_router
from app.api.routes.edit_templates import router as edit_templates_router
from app.api.routes.feature_flags import router as feature_flags_router
from app.api.routes.generated_videos import router as generated_videos_router
from app.api.routes.generation_tasks import router as generation_tasks_router
from app.api.routes.health import router as health_check_router
from app.api.routes.ingest_jobs import router as ingest_jobs_router
from app.api.routes.internal_render import router as internal_render_router
from app.api.routes.jobs import router as jobs_router
from app.api.routes.projects import router as projects_router
from app.api.routes.recipes import router as recipes_router
@@ -153,11 +151,3 @@ api_router.include_router(
prefix="/tts",
tags=["TTS"],
)
api_router.include_router(
feature_flags_router,
tags=["Internal"],
)
api_router.include_router(
internal_render_router,
tags=["Internal"],
)
-48
View File
@@ -1,48 +0,0 @@
"""路由层共享辅助函数 — 消除跨文件重复定义。"""
from typing import Any
from fastapi import HTTPException, status
from packages.application import GetProjectUseCase
from packages.ports.user_repository import UserRepository
def check_project_access(project_id: str, user_id: str, project_repository) -> None:
"""检查用户是否有项目访问权限。
合并自 asset_libraries.py / edit_plans.py 的同名函数。
- 空 project_id 直接放行(兼容 edit_plans 中 project_id 可选的场景)
- 错误信息使用中文,与项目其他路由保持一致
"""
if not project_id or not project_id.strip():
return
project = project_repository.find_by_id(project_id)
if project is None:
raise HTTPException(status_code=404, detail="项目不存在")
if not project.can_access(user_id):
raise HTTPException(status_code=403, detail="无权访问该项目")
def get_user_plan(user_id: str, user_repository: UserRepository) -> str:
"""获取用户的订阅计划名称。"""
user = user_repository.find_by_id(user_id)
if user is None:
return "free"
return getattr(user, "subscription_plan", "free") or "free"
def require_project_and_library(
project_id: str,
library_id: str,
project_repository: Any,
asset_library_repository: Any,
) -> None:
"""Verify project and asset library exist."""
project = GetProjectUseCase(project_repository).execute(project_id)
if project is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
libraries = asset_library_repository.find_by_project(project_id)
if not any(item.id == library_id for item in libraries):
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Asset library not found")
+10 -3
View File
@@ -22,11 +22,18 @@ from packages.application import (
)
from packages.domain import AssetLibrary, AssetLibraryKind
from ._helpers import check_project_access
router = APIRouter()
def _check_project_access(project_id: str, user_id: str, project_repository) -> None:
"""检查用户是否有项目访问权限"""
project = project_repository.find_by_id(project_id)
if project is None:
raise HTTPException(status_code=404, detail=f"Project {project_id} not found")
if not project.can_access(user_id):
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Access denied to project")
def _to_asset_library_response(item) -> AssetLibraryResponse:
return AssetLibraryResponse(
id=item.id,
@@ -161,7 +168,7 @@ def delete_asset_library(
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="素材库不存在")
# 权限校验:检查用户是否有项目访问权限
check_project_access(library.project_id, authenticated_user.user.id, project_repository)
_check_project_access(library.project_id, authenticated_user.user.id, project_repository)
# 删除库内所有素材(无 FK 级联,需手动清理)
assets_in_library = asset_repository.find_by_library(library_id)
+19 -13
View File
@@ -27,8 +27,6 @@ from packages.application import (
)
from packages.domain import AssetStatus, ClassificationStatus
from app.api.routes._helpers import check_project_access
logger = logging.getLogger(__name__)
router = APIRouter()
@@ -74,6 +72,14 @@ def _to_asset_response(item, storage_service=None) -> AssetResponse:
)
def _check_project_access(project_id: str, user_id: str, project_repository) -> None:
"""检查用户是否有项目访问权限"""
project = project_repository.find_by_id(project_id)
if project is None:
raise HTTPException(status_code=404, detail=f"Project {project_id} not found")
if not project.can_access(user_id):
raise HTTPException(status_code=403, detail="Access denied to project")
@router.get("", response_model=ListAssetsResponse)
def list_assets(
@@ -130,7 +136,7 @@ def list_assets(
library = asset_library_repository.get(library_id)
if library is None:
raise HTTPException(status_code=404, detail=f"AssetLibrary {library_id} not found")
check_project_access(library.project_id, user_id, project_repository)
_check_project_access(library.project_id, user_id, project_repository)
if ft:
items = asset_repository.find_by_library_and_file_type(library_id, ft, skip=skip, limit=limit)
total = asset_repository.count_by_project(library.project_id) if not kind else len(items)
@@ -146,7 +152,7 @@ def list_assets(
# 模式2:指定 project_id
if project_id:
check_project_access(project_id, user_id, project_repository)
_check_project_access(project_id, user_id, project_repository)
if ft:
# 无直接方法,加载后按 file_type 过滤(仍比全量加载好)
all_items = asset_repository.find_by_project(project_id)
@@ -204,13 +210,13 @@ def list_assets(
library = asset_library_repository.get(library_id)
if library is None:
raise HTTPException(status_code=404, detail=f"AssetLibrary {library_id} not found")
check_project_access(library.project_id, user_id, project_repository)
_check_project_access(library.project_id, user_id, project_repository)
if kind:
all_items = asset_repository.find_by_library_and_file_type(library_id, kind_to_file_type[kind])
else:
all_items = asset_repository.find_by_library(library_id)
elif project_id:
check_project_access(project_id, user_id, project_repository)
_check_project_access(project_id, user_id, project_repository)
all_items = asset_repository.find_by_project(project_id)
else:
try:
@@ -256,7 +262,7 @@ def update_asset_review_status(
item = asset_repository.get(asset_id)
if item is None:
raise HTTPException(status_code=404, detail=f"Asset {asset_id} not found")
check_project_access(item.project_id, authenticated_user.user.id, project_repository)
_check_project_access(item.project_id, authenticated_user.user.id, project_repository)
_apply_asset_review_status(item, request.review_status)
updated = asset_repository.update(item)
return _to_asset_response(updated)
@@ -280,7 +286,7 @@ def batch_delete_assets(
failed_ids.append(asset_id)
continue
try:
check_project_access(item.project_id, user_id, project_repository)
_check_project_access(item.project_id, user_id, project_repository)
deleted_ids.append(asset_id)
except HTTPException:
failed_ids.append(asset_id)
@@ -301,7 +307,7 @@ def get_asset(
item = asset_repository.find_by_id(asset_id)
if item is None:
raise HTTPException(status_code=404, detail=f"Asset {asset_id} not found")
check_project_access(item.project_id, authenticated_user.user.id, project_repository)
_check_project_access(item.project_id, authenticated_user.user.id, project_repository)
return _to_asset_response(item)
@@ -316,7 +322,7 @@ def update_asset(
item = asset_repository.find_by_id(asset_id)
if item is None:
raise HTTPException(status_code=404, detail=f"Asset {asset_id} not found")
check_project_access(item.project_id, authenticated_user.user.id, project_repository)
_check_project_access(item.project_id, authenticated_user.user.id, project_repository)
# 合并可修改字段
if request.name is not None:
@@ -340,7 +346,7 @@ def delete_asset(
item = asset_repository.find_by_id(asset_id)
if item is None:
raise HTTPException(status_code=404, detail=f"Asset {asset_id} not found")
check_project_access(item.project_id, authenticated_user.user.id, project_repository)
_check_project_access(item.project_id, authenticated_user.user.id, project_repository)
asset_repository.delete(asset_id)
@@ -357,7 +363,7 @@ def tag_asset(
item = asset_repository.find_by_id(asset_id)
if item is None:
raise HTTPException(status_code=404, detail=f"Asset {asset_id} not found")
check_project_access(item.project_id, authenticated_user.user.id, project_repository)
_check_project_access(item.project_id, authenticated_user.user.id, project_repository)
for tag_id in request.tag_ids:
tag = tag_repository.get(tag_id)
if tag is None:
@@ -381,7 +387,7 @@ def untag_asset(
item = asset_repository.find_by_id(asset_id)
if item is None:
raise HTTPException(status_code=404, detail=f"Asset {asset_id} not found")
check_project_access(item.project_id, authenticated_user.user.id, project_repository)
_check_project_access(item.project_id, authenticated_user.user.id, project_repository)
item.remove_tag(tag_id)
asset_repository.update(item)
+3 -2
View File
@@ -206,7 +206,7 @@ async def verify_email_post(
return _verify_email_token(request.token, user_repository)
@router.post("/forgot-password", response_model=MessageResponse, status_code=status.HTTP_202_ACCEPTED)
@router.post("/password/forgot", response_model=MessageResponse, status_code=status.HTTP_202_ACCEPTED)
async def forgot_password(
request: PasswordResetRequestModel,
user_repository: UserRepository = Depends(get_user_repository),
@@ -223,7 +223,7 @@ async def forgot_password(
return MessageResponse(message="如果账户存在,密码重置邮件已发送")
@router.post("/reset-password", response_model=MessageResponse)
@router.post("/password/reset", response_model=MessageResponse)
async def reset_password(
request: ResetPasswordModel,
user_repository: UserRepository = Depends(get_user_repository),
@@ -243,6 +243,7 @@ async def logout(
current_user: AuthenticatedUser = Depends(get_current_user),
):
"""登出 - 将当前 token 加入黑名单"""
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
if credentials:
try:
+19 -3
View File
@@ -14,6 +14,7 @@ from typing import Any
from uuid import uuid4
from app.auth import AuthenticatedUser, get_current_user
from app.config import get_settings
from app.core.celery_app import celery_app
from app.core.storage import OSSStorageService, get_storage_service
from app.dependencies import (
@@ -34,8 +35,6 @@ from fastapi.params import File
from packages.application import GetProjectUseCase, SubmitIngestJobCommand, SubmitIngestJobUseCase
from app.api.routes._helpers import require_project_and_library
router = APIRouter()
logger = logging.getLogger(__name__)
@@ -114,6 +113,22 @@ def _atomic_check_and_record(upload_id: str, chunk_index: int) -> bool:
fcntl.flock(f.fileno(), fcntl.LOCK_UN)
def _require_project_and_library(
project_id: str,
library_id: str,
project_repository: Any,
asset_library_repository: Any,
) -> None:
"""Verify project and asset library exist"""
project = GetProjectUseCase(project_repository).execute(project_id)
if project is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
libraries = asset_library_repository.find_by_project(project_id)
if not any(item.id == library_id for item in libraries):
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Asset library not found")
def _load_upload_meta(upload_id: str) -> dict[str, Any]:
"""Load upload metadata"""
meta_path = _get_upload_meta_path(upload_id)
@@ -191,6 +206,7 @@ async def init_chunked_upload(
asset_library_repository: Any = Depends(get_asset_library_repository),
) -> ChunkedUploadInitResponse:
"""Initialize chunked upload"""
settings = get_settings()
# Validate file size
if request.file_size > MAX_FILE_SIZE:
@@ -205,7 +221,7 @@ async def init_chunked_upload(
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
# Verify asset library
require_project_and_library(
_require_project_and_library(
request.project_id,
request.library_id,
project_repository,
+31 -16
View File
@@ -32,6 +32,12 @@ from fastapi import APIRouter, Depends, HTTPException, Query, status
from pydantic import BaseModel, Field
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.asset_library_repository import (
SQLAlchemyAssetLibraryRepository,
)
from packages.adapters.sqlalchemy_impl.asset_repository import (
SQLAlchemyAssetRepository,
)
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
SQLAlchemyGenerationTaskRepository,
)
@@ -45,8 +51,6 @@ from packages.application.generation_tasks import (
CreateGenerationTaskCommand,
CreateGenerationTaskUseCase,
)
from ._helpers import check_project_access
from packages.domain.config_schemas import normalize_plan_config
from packages.domain.edit_plan import EditPlan, EditPlanStatus
@@ -243,6 +247,17 @@ class GenerateFromTemplateResponse(BaseModel):
# ── Helpers ───────────────────────────────────────────────────────────────────
def _check_project_access(project_id: str, user_id: str, project_repository: Any) -> None:
"""校验用户对项目的访问权限(参照 assets.py 的 can_access 模式)"""
if not project_id or not project_id.strip():
return
project = project_repository.find_by_id(project_id)
if project is None:
raise HTTPException(status_code=404, detail="项目不存在")
if not project.can_access(user_id):
raise HTTPException(status_code=403, detail="无权访问该项目")
def _to_response(p: EditPlan) -> EditPlanResponse:
return EditPlanResponse(
id=p.id,
@@ -296,7 +311,7 @@ def list_plans(
# 项目鉴权:如果指定了 project_id,校验用户是否有权访问
if project_id:
check_project_access(project_id, current_user.user.id, project_repository)
_check_project_access(project_id, current_user.user.id, project_repository)
skip = (page - 1) * page_size
plans = svc.list_plans(
@@ -338,7 +353,7 @@ def get_plan(
)
# 项目鉴权
if plan.project_id:
check_project_access(plan.project_id, current_user.user.id, project_repository)
_check_project_access(plan.project_id, current_user.user.id, project_repository)
return _to_response(plan)
@@ -354,7 +369,7 @@ def create_plan(
project_id = (body.project_id or "").strip()
# 项目鉴权
if project_id:
check_project_access(project_id, current_user.user.id, project_repository)
_check_project_access(project_id, current_user.user.id, project_repository)
svc = EditPlanService(db)
# 标准化 config,填充 cover/title/subtitle/bgm 默认值
normalized_config = normalize_plan_config(body.config)
@@ -396,7 +411,7 @@ def update_plan(
if existing is None:
raise HTTPException(status_code=404, detail=f"剪辑计划不存在: {plan_id}")
if existing.project_id:
check_project_access(existing.project_id, current_user.user.id, project_repository)
_check_project_access(existing.project_id, current_user.user.id, project_repository)
# 基础字段更新
try:
@@ -450,7 +465,7 @@ def delete_plan(
# 项目鉴权
existing = svc.get_plan(plan_id)
if existing and existing.project_id:
check_project_access(existing.project_id, current_user.user.id, project_repository)
_check_project_access(existing.project_id, current_user.user.id, project_repository)
deleted = svc.delete_plan(plan_id)
if not deleted:
raise HTTPException(
@@ -492,7 +507,7 @@ def generate_plan(
if plan_check is None:
raise HTTPException(status_code=404, detail=f"剪辑计划不存在: {plan_id}")
if plan_check.project_id:
check_project_access(plan_check.project_id, current_user.user.id, project_repository)
_check_project_access(plan_check.project_id, current_user.user.id, project_repository)
# ── 自动兜底 1: draft → editing ──────────────────────────────────────
if plan_check.status == EditPlanStatus.DRAFT:
@@ -695,7 +710,7 @@ def generate_plan(
except HTTPException:
# 已处理的 HTTP 异常直接透传
raise
except Exception:
except Exception as exc:
logger.exception("触发剪辑计划生成失败: plan_id=%s", plan_id)
# 尝试将计划标记为失败(RENDERING → FAILED 是合法的状态流转)
try:
@@ -734,7 +749,7 @@ def get_generation_status(
plan = gen_status["plan"]
# 项目鉴权
if plan.project_id:
check_project_access(plan.project_id, current_user.user.id, project_repository)
_check_project_access(plan.project_id, current_user.user.id, project_repository)
clips = gen_status["clips"]
clip_items = [
@@ -776,7 +791,7 @@ def list_plan_generations(
# 验证计划存在 + 项目鉴权
plan = svc.get_plan_or_raise(plan_id)
if plan.project_id:
check_project_access(plan.project_id, current_user.user.id, project_repository)
_check_project_access(plan.project_id, current_user.user.id, project_repository)
gen_task_repo = SQLAlchemyGenerationTaskRepository(db)
tasks = gen_task_repo.list_by_source_edit_plan(plan_id)
@@ -845,7 +860,7 @@ def ai_recommend_clips(
# 项目鉴权
if plan.project_id:
check_project_access(plan.project_id, current_user.user.id, project_repository)
_check_project_access(plan.project_id, current_user.user.id, project_repository)
# 验证状态:只允许 draft 或 editing
plan_status = plan.status.value if hasattr(plan.status, "value") else plan.status
@@ -894,7 +909,7 @@ def ai_recommend_clips(
config=normalized_config,
total_duration=result["total_duration"],
)
except Exception:
except Exception as exc:
logger.exception("AI 推荐写入失败,plan_id=%s 数据可能不一致", plan_id)
# 尝试回滚未提交的变更
try:
@@ -975,7 +990,7 @@ def generate_cover(
# 项目鉴权
if plan.project_id:
check_project_access(plan.project_id, current_user.user.id, project_repository)
_check_project_access(plan.project_id, current_user.user.id, project_repository)
# 调用 AI 封面生成服务
from apps.worker.worker_app.tasks.ai_tasks import run_generate_cover
@@ -1095,7 +1110,7 @@ def get_plan_timeline(
plan = svc.get_plan_or_raise(plan_id)
# 项目鉴权
if plan.project_id:
check_project_access(plan.project_id, current_user.user.id, project_repository)
_check_project_access(plan.project_id, current_user.user.id, project_repository)
clips = svc.list_clips(plan_id=plan_id, skip=0, limit=200)
# 按 order 排序
@@ -1156,7 +1171,7 @@ def generate_from_template(
# 项目鉴权
if body.project_id:
check_project_access(body.project_id, current_user.user.id, project_repository)
_check_project_access(body.project_id, current_user.user.id, project_repository)
template_svc = EditTemplateService(db)
-195
View File
@@ -1,195 +0,0 @@
"""Feature Flag 内部管理接口。
通过内部 API Key 鉴权,支持查看和修改 Feature Flag 配置。
主要用于灰度发布期间的动态开关控制。
API:
GET /api/v1/internal/feature-flags - 列出所有 flag
GET /api/v1/internal/feature-flags/{name} - 查看单个 flag
PUT /api/v1/internal/feature-flags/{name} - 设置 flag 配置
DELETE /api/v1/internal/feature-flags/{name} - 删除 flag
鉴权:X-API-Key header,走内部 API Key 验证
"""
from __future__ import annotations
import logging
from typing import Optional
from app.api.routes.auth import _verify_internal_api_key
from app.config import settings
from fastapi import APIRouter, Depends, HTTPException, Query, status
from pydantic import BaseModel, Field
from packages.adapters.redis.feature_flag_store import (
FEATURE_FLAG_REDIS_PREFIX,
FeatureFlagConfig,
RedisFeatureFlagStore,
)
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/internal/feature-flags", tags=["Internal"])
# 允许管理的 flag 白名单(防止误操作其他系统 flag)
ALLOWED_FLAGS = {
"render_engine",
}
def _get_feature_flag_store() -> RedisFeatureFlagStore:
"""获取 Feature Flag 存储实例。"""
return RedisFeatureFlagStore(redis_url=settings.REDIS_URL)
class FeatureFlagUpdateRequest(BaseModel):
"""Feature Flag 更新请求体。"""
enabled: bool = Field(..., description="是否启用")
percentage: int = Field(0, ge=0, le=100, description="灰度百分比 (0-100)")
whitelist: list[str] = Field(default_factory=list, description="白名单列表(如 user_id")
class FeatureFlagResponse(BaseModel):
"""Feature Flag 响应。"""
name: str
enabled: bool
percentage: int
whitelist: list[str]
@classmethod
def from_config(cls, config: FeatureFlagConfig) -> "FeatureFlagResponse":
return cls(
name=config.name,
enabled=config.enabled,
percentage=config.percentage,
whitelist=sorted(config.whitelist),
)
class FeatureFlagCheckResponse(BaseModel):
"""Flag 激活检查响应。"""
name: str
active: bool
identifier: Optional[str] = None
def _validate_flag_name(name: str) -> None:
"""校验 flag 名称是否在允许列表中。"""
if name not in ALLOWED_FLAGS:
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail=f"Unsupported flag: {name}. Allowed: {sorted(ALLOWED_FLAGS)}",
)
@router.get("", response_model=list[FeatureFlagResponse])
async def list_feature_flags(
_: bool = Depends(_verify_internal_api_key),
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
):
"""列出所有 Feature Flag。"""
try:
flags = store.list_all()
# 同时返回预定义的 flag(即使未设置也显示默认值)
result = []
for name in sorted(ALLOWED_FLAGS):
config = flags.get(name) or FeatureFlagConfig(name=name, enabled=False)
result.append(FeatureFlagResponse.from_config(config))
# 加上已存在但不在白名单中的 flag(只读展示)
for name, config in flags.items():
if name not in ALLOWED_FLAGS:
result.append(FeatureFlagResponse.from_config(config))
return sorted(result, key=lambda x: x.name)
except Exception as exc:
logger.error("Failed to list feature flags: %s", exc)
raise HTTPException(status_code=500, detail=f"Failed to list flags: {exc}")
@router.get("/{name}", response_model=FeatureFlagResponse)
async def get_feature_flag(
name: str,
_: bool = Depends(_verify_internal_api_key),
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
):
"""获取单个 Feature Flag 配置。"""
try:
config = store.get(name)
return FeatureFlagResponse.from_config(config)
except Exception as exc:
logger.error("Failed to get feature flag %s: %s", name, exc)
raise HTTPException(status_code=500, detail=f"Failed to get flag: {exc}")
@router.get("/{name}/check", response_model=FeatureFlagCheckResponse)
async def check_feature_flag(
name: str,
identifier: Optional[str] = Query(None, description="标识符,如 user_id"),
_: bool = Depends(_verify_internal_api_key),
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
):
"""检查某个标识符是否命中 Feature Flag。"""
try:
active = store.is_active(name, identifier=identifier)
return FeatureFlagCheckResponse(name=name, active=active, identifier=identifier)
except Exception as exc:
logger.error("Failed to check feature flag %s: %s", name, exc)
raise HTTPException(status_code=500, detail=f"Failed to check flag: {exc}")
@router.put("/{name}", response_model=FeatureFlagResponse)
async def update_feature_flag(
name: str,
request: FeatureFlagUpdateRequest,
_: bool = Depends(_verify_internal_api_key),
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
):
"""更新 Feature Flag 配置。
只允许修改 ALLOWED_FLAGS 列表中的 flag。
"""
_validate_flag_name(name)
try:
config = FeatureFlagConfig(
name=name,
enabled=request.enabled,
percentage=request.percentage,
whitelist=set(request.whitelist),
)
store.set(config)
logger.info(
"Feature flag updated: name=%s enabled=%s percentage=%d whitelist=%d",
name,
config.enabled,
config.percentage,
len(config.whitelist),
)
return FeatureFlagResponse.from_config(config)
except Exception as exc:
logger.error("Failed to update feature flag %s: %s", name, exc)
raise HTTPException(status_code=500, detail=f"Failed to update flag: {exc}")
@router.delete("/{name}", status_code=status.HTTP_204_NO_CONTENT)
async def delete_feature_flag(
name: str,
_: bool = Depends(_verify_internal_api_key),
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
):
"""删除 Feature Flag。
只允许删除 ALLOWED_FLAGS 列表中的 flag。
"""
_validate_flag_name(name)
try:
deleted = store.delete(name)
logger.info("Feature flag deleted: name=%s deleted=%s", name, deleted)
return None
except Exception as exc:
logger.error("Failed to delete feature flag %s: %s", name, exc)
raise HTTPException(status_code=500, detail=f"Failed to delete flag: {exc}")
+11 -4
View File
@@ -32,8 +32,6 @@ from app.schemas.generation_task import (
)
from fastapi import APIRouter, Depends, HTTPException
from app.api.routes._helpers import check_project_access
from packages.application import (
CreateGenerationTaskCommand,
CreateGenerationTaskUseCase,
@@ -45,6 +43,15 @@ logger = logging.getLogger(__name__)
router = APIRouter()
def _check_project_access(project_id: str, user_id: str, project_repository) -> None:
"""检查用户是否有项目访问权限"""
project = project_repository.find_by_id(project_id)
if project is None:
raise HTTPException(status_code=404, detail=f"Project {project_id} not found")
if not project.can_access(user_id):
raise HTTPException(status_code=403, detail="Access denied to project")
def _to_generation_task_response(task) -> GenerationTaskResponse:
return GenerationTaskResponse(
id=task.id,
@@ -332,7 +339,7 @@ def get_generation_task(
if task is None:
raise HTTPException(status_code=404, detail=f"GenerationTask {task_id} not found")
if task.project_id:
check_project_access(task.project_id, authenticated_user.user.id, project_repository)
_check_project_access(task.project_id, authenticated_user.user.id, project_repository)
return _to_generation_task_response(task)
@@ -349,7 +356,7 @@ def list_generation_results(
if task is None:
raise HTTPException(status_code=404, detail=f"GenerationTask {task_id} not found")
if task.project_id:
check_project_access(task.project_id, authenticated_user.user.id, project_repository)
_check_project_access(task.project_id, authenticated_user.user.id, project_repository)
use_case = ListGeneratedVideosByTaskUseCase(generated_video_repository)
items = use_case.execute(task_id)
responses = []
-120
View File
@@ -1,120 +0,0 @@
"""渲染结果内部下载接口。
通过内部 API Key 鉴权,为灰度对比工具等内部系统提供渲染结果下载能力。
API:
GET /api/v1/internal/render/videos/{video_id}/download-url - 获取单个视频下载URL
GET /api/v1/internal/render/tasks/{task_id}/videos - 获取任务下所有视频及下载URL
鉴权:X-API-Key header,走内部 API Key 验证
"""
from __future__ import annotations
import logging
from typing import Any
from app.api.routes.auth import _verify_internal_api_key
from app.core.storage import OSSStorageService, get_storage_service
from app.dependencies import get_generated_video_repository
from fastapi import APIRouter, Depends, HTTPException, Query
from pydantic import BaseModel
logger = logging.getLogger(__name__)
router = APIRouter(prefix="/internal/render", tags=["Internal"])
class InternalRenderVideoItem(BaseModel):
"""内部渲染视频项。"""
video_id: str
generation_task_id: str
project_id: str
name: str
file_url: str
file_size: int | None = None
duration: float | None = None
width: int | None = None
height: int | None = None
fps: float | None = None
status: str
download_url: str
class InternalRenderTaskVideosResponse(BaseModel):
"""任务下所有渲染视频响应。"""
task_id: str
count: int
videos: list[InternalRenderVideoItem]
class InternalRenderDownloadUrlResponse(BaseModel):
"""单个视频下载URL响应。"""
video_id: str
download_url: str
def _video_to_item(video: Any, download_url: str) -> InternalRenderVideoItem:
"""将 GeneratedVideo 领域对象转为响应项。"""
return InternalRenderVideoItem(
video_id=video.id,
generation_task_id=video.generation_task_id,
project_id=video.project_id,
name=video.name,
file_url=video.file_url,
file_size=getattr(video, "file_size", None),
duration=getattr(video, "duration", None),
width=getattr(video, "width", None),
height=getattr(video, "height", None),
fps=getattr(video, "fps", None),
status=video.status,
download_url=download_url,
)
@router.get("/videos/{video_id}/download-url", response_model=InternalRenderDownloadUrlResponse)
def get_render_video_download_url(
video_id: str,
_: bool = Depends(_verify_internal_api_key),
generated_video_repository: Any = Depends(get_generated_video_repository),
storage_service: OSSStorageService = Depends(get_storage_service),
) -> InternalRenderDownloadUrlResponse:
"""获取单个渲染视频的下载URL(预签名)。"""
video = generated_video_repository.get(video_id)
if video is None:
raise HTTPException(status_code=404, detail=f"GeneratedVideo {video_id} not found")
download_url = storage_service.get_download_url(video.file_url, expires_seconds=86400)
logger.info("内部渲染下载URL生成: video_id=%s", video_id)
return InternalRenderDownloadUrlResponse(video_id=video_id, download_url=download_url)
@router.get("/tasks/{task_id}/videos", response_model=InternalRenderTaskVideosResponse)
def get_render_task_videos(
task_id: str,
status: str | None = Query(None, description="按状态筛选,如 completed/failed"),
_: bool = Depends(_verify_internal_api_key),
generated_video_repository: Any = Depends(get_generated_video_repository),
storage_service: OSSStorageService = Depends(get_storage_service),
) -> InternalRenderTaskVideosResponse:
"""获取生成任务下所有渲染视频及下载URL。"""
videos = generated_video_repository.list_by_generation_task(task_id)
# 状态筛选
if status:
videos = [v for v in videos if v.status == status]
items = []
for video in videos:
download_url = storage_service.get_download_url(video.file_url, expires_seconds=86400)
items.append(_video_to_item(video, download_url))
logger.info("内部渲染任务视频查询: task_id=%s count=%d", task_id, len(items))
return InternalRenderTaskVideosResponse(
task_id=task_id,
count=len(items),
videos=items,
)
+13 -6
View File
@@ -20,7 +20,7 @@ from typing import Any
from app.auth import AuthenticatedUser, get_current_user
from app.core.celery_app import celery_app
from app.dependencies import get_job_repository, get_project_repository
from app.dependencies import get_db_session, get_job_repository, get_project_repository
from app.schemas.job import (
CompleteJobRequest,
CreateJobRequest,
@@ -51,8 +51,6 @@ from packages.application.jobs import (
)
from packages.domain.job import JobType
from app.api.routes._helpers import check_project_access
logger = logging.getLogger(__name__)
router = APIRouter()
@@ -68,6 +66,15 @@ _JOB_TYPE_TO_CELERY_TASK: dict[str, str] = {
}
def _check_project_access(project_id: str, user_id: str, project_repository) -> None:
"""检查用户是否有项目访问权限。"""
project = project_repository.find_by_id(project_id)
if project is None:
raise HTTPException(status_code=404, detail=f"Project {project_id} not found")
if not project.can_access(user_id):
raise HTTPException(status_code=403, detail="Access denied to project")
# ── 创建任务 ──────────────────────────────────────────────────────────────────
@@ -82,7 +89,7 @@ def create_job(
创建后任务处于 pending 状态,需要调用 /submit 提交执行。
"""
check_project_access(request.project_id, authenticated_user.user.id, project_repository)
_check_project_access(request.project_id, authenticated_user.user.id, project_repository)
# 校验 job_type
try:
@@ -175,7 +182,7 @@ def list_project_jobs(
offset: int = Query(default=0, ge=0),
) -> ListJobsResponse:
"""获取项目下的任务列表。"""
check_project_access(project_id, authenticated_user.user.id, project_repository)
_check_project_access(project_id, authenticated_user.user.id, project_repository)
use_case = ListJobsUseCase(job_repo)
jobs = use_case.execute(
@@ -197,7 +204,7 @@ def get_job_statistics(
project_repository: Any = Depends(get_project_repository),
) -> JobStatisticsResponse:
"""获取项目任务统计摘要。"""
check_project_access(project_id, authenticated_user.user.id, project_repository)
_check_project_access(project_id, authenticated_user.user.id, project_repository)
use_case = GetJobStatisticsUseCase(job_repo)
stats = use_case.execute(project_id)
+8 -3
View File
@@ -33,8 +33,6 @@ from packages.application.recipe.use_cases import (
)
from packages.ports.user_repository import UserRepository
from app.api.routes._helpers import get_user_plan
router = APIRouter()
@@ -42,6 +40,13 @@ def _get_recipe_repository(session: Session = Depends(get_db_session)) -> SQLAlc
return SQLAlchemyRecipeRepository(session)
def _get_user_plan(user_id: str, user_repository: UserRepository) -> str:
user = user_repository.find_by_id(user_id)
if user is None:
return "free"
return getattr(user, "subscription_plan", "free") or "free"
def _item_to_response(item) -> RecipeItemResponse:
return RecipeItemResponse(
id=item.id,
@@ -189,7 +194,7 @@ def use_recipe(
user_repository: UserRepository = Depends(get_user_repository),
) -> UseRecipeResponse:
user_id = authenticated_user.user.id
plan_name = get_user_plan(user_id, user_repository)
plan_name = _get_user_plan(user_id, user_repository)
use_case = UseRecipeUseCase(recipe_repository)
try:
result = use_case.execute(recipe_id, user_id, user_plan=plan_name)
+1 -1
View File
@@ -232,7 +232,7 @@ async def payment_callback(
# 创建账单记录
record_id = uuid.uuid4().hex
repo.create(
record = repo.create(
{
"id": record_id,
"user_id": user_id,
+8 -3
View File
@@ -28,8 +28,6 @@ from packages.application.title_library.use_cases import (
)
from packages.ports.user_repository import UserRepository
from app.api.routes._helpers import get_user_plan
router = APIRouter()
@@ -53,6 +51,13 @@ def _to_response(item) -> TitleLibraryItemResponse:
)
def _get_user_plan(user_id: str, user_repository: UserRepository) -> str:
user = user_repository.find_by_id(user_id)
if user is None:
return "free"
return getattr(user, "subscription_plan", "free") or "free"
@router.get("", response_model=ListTitleLibraryResponse)
def list_titles(
category: Optional[str] = Query(None),
@@ -93,7 +98,7 @@ def create_title(
user_repository: UserRepository = Depends(get_user_repository),
) -> TitleLibraryItemResponse:
user_id = authenticated_user.user.id
plan_name = get_user_plan(user_id, user_repository)
plan_name = _get_user_plan(user_id, user_repository)
command = CreateTitleLibraryCommand(
user_id=user_id,
name=request.name,
+21 -7
View File
@@ -1,5 +1,5 @@
import logging
from typing import Any
from typing import Annotated, Any
from uuid import uuid4
from app.auth import AuthenticatedUser, get_current_user
@@ -17,13 +17,12 @@ from app.schemas.upload import (
DirectUploadCompleteResponse,
DirectUploadPrepareRequest,
DirectUploadPrepareResponse,
UploadAssetRequest,
UploadAssetResponse,
)
from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile, status
from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase
from app.api.routes._helpers import require_project_and_library
from packages.application import GetProjectUseCase, SubmitIngestJobCommand, SubmitIngestJobUseCase
logger = logging.getLogger(__name__)
@@ -81,6 +80,21 @@ def _validate_mime_type(content_type: str | None) -> str:
return base_type
def _require_project_and_library(
project_id: str,
library_id: str,
project_repository: Any,
asset_library_repository: Any,
) -> None:
project = GetProjectUseCase(project_repository).execute(project_id)
if project is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
libraries = asset_library_repository.find_by_project(project_id)
if not any(item.id == library_id for item in libraries):
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Asset library not found")
def _submit_ingest_job(
project_id: str,
library_id: str,
@@ -121,7 +135,7 @@ async def prepare_direct_upload(
# P2-5: 服务端验证 MIME 类型
validated_content_type = _validate_mime_type(request.content_type)
require_project_and_library(
_require_project_and_library(
request.project_id,
request.library_id,
project_repository,
@@ -169,7 +183,7 @@ async def complete_direct_upload(
storage_service: OSSStorageService = Depends(get_storage_service),
) -> DirectUploadCompleteResponse:
"""确认浏览器直传完成并创建导入任务。"""
require_project_and_library(
_require_project_and_library(
request.project_id,
request.library_id,
project_repository,
@@ -238,7 +252,7 @@ async def upload_asset(
storage_service: OSSStorageService = Depends(get_storage_service),
) -> UploadAssetResponse:
"""上传素材文件并触发导入流水线。"""
require_project_and_library(project_id, library_id, project_repository, asset_library_repository)
_require_project_and_library(project_id, library_id, project_repository, asset_library_repository)
# ── 素材去重检测:上传前检查同素材库 + 同 file_hash ──
if file_hash:
+1
View File
@@ -28,6 +28,7 @@ from packages.application.voice_clone.use_cases import (
VoiceCloneNotRetryableError,
)
from packages.application.voice_clone.workflow import (
VoiceCloneWorkflowError,
VoiceCloneWorkflowService,
)
+8 -3
View File
@@ -39,8 +39,6 @@ from packages.application.voice_library.use_cases import (
from packages.domain.preset_voices import PRESET_VOICES
from packages.ports.user_repository import UserRepository
from app.api.routes._helpers import get_user_plan
router = APIRouter()
@@ -127,6 +125,13 @@ def _preset_to_unified_response(preset) -> UnifiedVoiceItemResponse:
)
def _get_user_plan(user_id: str, user_repository: UserRepository) -> str:
user = user_repository.find_by_id(user_id)
if user is None:
return "free"
return getattr(user, "subscription_plan", "free") or "free"
# ==================== 统一配音列表(预置 + 克隆)====================
@@ -266,7 +271,7 @@ def create_voice(
sign_url=Depends(get_audio_url_signer),
) -> VoiceLibraryItemResponse:
user_id = authenticated_user.user.id
plan_name = get_user_plan(user_id, user_repository)
plan_name = _get_user_plan(user_id, user_repository)
command = CreateVoiceLibraryCommand(
user_id=user_id,
name=request.name,
+1 -6
View File
@@ -25,7 +25,7 @@ class Settings(BaseSettings):
DATABASE_POOL_SIZE: int = 20
DATABASE_MAX_OVERFLOW: int = 10 # 调整为合理值:pool_size(20) + max_overflow(10) = 最大30连接
DATABASE_POOL_TIMEOUT: int = 30
DATABASE_POOL_RECYCLE: int = 3600
DATABASE_POOL_RECYLE: int = 3600
USE_IN_MEMORY_DB: bool = False
AUTO_CREATE_SCHEMA: bool = False
@@ -41,11 +41,6 @@ class Settings(BaseSettings):
# 密钥轮换天数(到达此天数后建议更换密钥)
SECRET_ROTATION_DAYS: int = 90
# JWT 算法与过期时间(与 .env.example 对齐)
JWT_ALGORITHM: str = "HS256"
JWT_ACCESS_TOKEN_EXPIRE_MINUTES: int = 30
JWT_REFRESH_TOKEN_EXPIRE_DAYS: int = 30
@field_validator("JWT_SECRET_KEY", mode="before")
@classmethod
def validate_jwt_secret_key(cls, v):
+1
View File
@@ -0,0 +1 @@
"""Core configuration package."""
+12
View File
@@ -50,8 +50,20 @@ from packages.adapters.sqlalchemy_impl.voice_clone_profile_repository import (
from packages.adapters.sqlalchemy_impl.voice_library_repository import (
SQLAlchemyVoiceLibraryRepository,
)
from packages.ports.asset_library_repository import AssetLibraryRepository
from packages.ports.asset_repository import AssetRepository
from packages.ports.classification_job_repository import ClassificationJobRepository
from packages.ports.duplication_repository import DuplicationRecordRepository
from packages.ports.generated_video_repository import GeneratedVideoRepository
from packages.ports.generation_task_repository import GenerationTaskRepository
from packages.ports.ingest_job_repository import IngestJobRepository
from packages.ports.job_repository import JobRepository
from packages.ports.project_repository import ProjectRepository
from packages.ports.tag_repository import TagRepository
from packages.ports.title_library_repository import TitleLibraryRepository
from packages.ports.user_repository import UserRepository
from packages.ports.voice_clone_profile_repository import VoiceCloneProfileRepository
from packages.ports.voice_library_repository import VoiceLibraryRepository
_engine, _SessionLocal = build_session_factory(settings.DATABASE_URL)
+1 -1
View File
@@ -10,7 +10,7 @@ from __future__ import annotations
from app.auth import AuthenticatedUser
from app.auth import get_current_user as get_authenticated_user
from app.dependencies import get_user_repository
from fastapi import Depends, HTTPException
from fastapi import Depends
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
from packages.domain.entities import User
+1 -1
View File
@@ -6,7 +6,7 @@ import logging
import time
from typing import Callable
from fastapi import Request
from fastapi import Request, Response
from starlette.middleware.base import BaseHTTPMiddleware
logger = logging.getLogger(__name__)
+1 -1
View File
@@ -2,7 +2,7 @@
from __future__ import annotations
from typing import Optional
from typing import List, Optional
from pydantic import BaseModel, Field
+1 -1
View File
@@ -24,7 +24,7 @@ from packages.adapters.sqlalchemy_impl import (
)
from packages.domain.asset import AssetType
from packages.domain.classification import AssetClassification
from packages.domain.edit_plan_clip import EditPlanClip
from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus
logger = logging.getLogger(__name__)
@@ -18,6 +18,7 @@ from packages.adapters.sqlalchemy_impl import (
)
from packages.domain.edit_plan import EditPlan, EditPlanStatus
from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus
from packages.domain.generation_task import GenerationTaskStatus
logger = logging.getLogger(__name__)
+1
View File
@@ -13,6 +13,7 @@ from __future__ import annotations
import logging
from typing import Any
from sqlalchemy.orm import Session
from packages.application.jobs import (
CancelJobUseCase,
@@ -13,7 +13,7 @@
from __future__ import annotations
import logging
from typing import Any, List
from typing import Any, List, Optional
from sqlalchemy.orm import Session
@@ -16,6 +16,7 @@ FFmpeg 视频合成编排服务:
from __future__ import annotations
import logging
import shutil
from dataclasses import dataclass, field
from typing import Any
@@ -27,7 +28,7 @@ from packages.adapters.sqlalchemy_impl.edit_plan_clip_repository import (
from packages.adapters.sqlalchemy_impl.edit_plan_repository import (
SQLAlchemyEditPlanRepository,
)
from packages.domain.edit_plan import EditPlanStatus
from packages.domain.edit_plan import EditPlan, EditPlanStatus
from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus
from packages.domain.template_clip_config import TransitionEffect
+42
View File
@@ -0,0 +1,42 @@
/**
* API
* Phase 1
*/
import apiClient from "./client";
/** 仪表盘概览数据 */
export interface DashboardOverview {
/** 素材总数 */
total_assets: number;
/** 已用存储(字节) */
used_storage_bytes: number;
/** 总标题数 */
total_titles: number;
/** 总配音数 */
total_voices: number;
/** 生成任务总数 */
total_tasks: number;
/** 成品总数 */
total_products: number;
/** 最近生成任务 */
recent_tasks: Array<{
id: string;
task_type: string;
status: string;
progress: number;
user_message: string;
created_at: string;
}>;
/** 订阅信息 */
subscription: {
plan: "free" | "pro" | "enterprise";
status: "active" | "inactive" | "expired";
expires_at?: string;
};
}
/** 获取仪表盘概览数据 */
export const getDashboardOverview = async (): Promise<DashboardOverview> => {
const response = await apiClient.get("/dashboard/overview");
return response.data;
};
@@ -0,0 +1,365 @@
/* V21 业务组件统一样式 */
/* ==================== 按钮 ==================== */
.xx-primary-btn {
background: var(--gradient-primary) !important;
color: var(--text-inverse) !important;
border: none !important;
border-radius: var(--radius-md) !important;
padding: 10px 20px !important;
font-weight: var(--font-weight-bold) !important;
box-shadow: var(--shadow-primary) !important;
transition: var(--transition-all) !important;
cursor: pointer;
height: auto !important;
}
.xx-primary-btn:hover {
box-shadow: var(--shadow-hover) !important;
transform: translateY(-1px);
}
.xx-ghost-btn {
background: transparent !important;
color: var(--primary-color) !important;
border: 2px solid var(--primary-color) !important;
border-radius: var(--radius-md) !important;
padding: var(--space-sm) 18px !important;
font-weight: var(--font-weight-bold) !important;
transition: var(--transition-all) !important;
cursor: pointer;
height: auto !important;
}
.xx-ghost-btn:hover {
background: var(--primary-soft) !important;
}
/* ==================== 卡片 ==================== */
.xx-card {
background: var(--bg-elevated);
border: 1px solid var(--border-color);
border-radius: var(--radius-xl);
box-shadow: var(--shadow-card);
padding: var(--space-lg);
margin-bottom: 20px;
transition: all var(--transition-slow);
}
.xx-card:hover {
box-shadow: var(--shadow-md);
transform: translateY(-2px);
}
/* ==================== 页面结构 ==================== */
.xx-page {
max-width: 1200px;
margin: 0 auto;
padding: var(--space-lg);
}
.xx-page-head {
display: flex;
justify-content: space-between;
align-items: flex-start;
gap: 18px;
margin-bottom: 28px;
}
.xx-page-head h2 {
font-size: 26px;
font-weight: var(--font-weight-extrabold);
color: var(--text-primary);
margin: 0 0 var(--space-sm);
}
.xx-page-head p {
font-size: var(--font-size-base);
color: var(--text-secondary);
margin: 0;
}
/* ==================== 表格样式 ==================== */
.xx-table-card {
background: var(--bg-elevated);
border: 1px solid var(--border-color);
border-radius: var(--radius-xl);
box-shadow: var(--shadow-card);
padding: 20px;
overflow: hidden;
}
/* 表格包装器 */
.xx-table-wrapper {
border-radius: var(--radius-lg);
overflow: hidden;
}
/* ==================== 标签/Tag ==================== */
.xx-tag {
padding: var(--space-xs) 12px;
border-radius: var(--radius-xs);
font-size: 13px;
font-weight: var(--font-weight-medium);
}
.xx-tag-indigo {
background: var(--primary-soft);
color: var(--primary-color);
border: 1px solid var(--color-primary-200);
}
.xx-tag-success {
background: var(--success-soft);
color: var(--color-secondary-500);
border: 1px solid var(--success-border);
}
.xx-tag-warning {
background: var(--warning-soft);
color: var(--accent-dark);
border: 1px solid var(--color-accent-200);
}
.xx-tag-error {
background: var(--error-soft);
color: var(--error-color);
border: 1px solid var(--error-border);
}
/* ==================== 搜索栏 ==================== */
.xx-search-bar {
margin-bottom: 20px;
}
.xx-search-input {
width: 100%;
padding: 12px 18px;
border: 2px solid var(--border-color);
border-radius: var(--radius-md);
font-size: var(--font-size-base);
background: var(--bg-primary);
transition: var(--transition-all);
outline: none;
}
.xx-search-input:focus {
border-color: var(--primary-color);
box-shadow: 0 0 0 4px
color-mix(in srgb, var(--primary-color) 10%, transparent);
}
/* ==================== Modal ==================== */
.xx-modal .ant-modal-content {
border-radius: var(--radius-xl);
padding: var(--space-lg);
}
.xx-modal .ant-modal-header {
border-radius: var(--radius-xl) var(--radius-xl) 0 0;
padding: 20px var(--space-lg);
border-bottom: 1px solid var(--border-color);
}
.xx-modal .ant-modal-title {
font-size: var(--font-size-lg);
font-weight: var(--font-weight-bold);
color: var(--text-primary);
}
.xx-modal .ant-modal-footer {
border-top: 1px solid var(--border-color);
padding: var(--space-md) var(--space-lg);
}
/* ==================== 空状态 ==================== */
.xx-empty-state {
text-align: center;
padding: var(--space-3xl) var(--space-lg);
color: var(--text-secondary);
}
.xx-empty-state-icon {
font-size: 48px;
margin-bottom: var(--space-md);
}
/* ==================== 网格布局 ==================== */
.xx-grid-2 {
display: grid;
grid-template-columns: repeat(2, 1fr);
gap: 20px;
}
.xx-grid-3 {
display: grid;
grid-template-columns: repeat(3, 1fr);
gap: 20px;
}
.xx-grid-4 {
display: grid;
grid-template-columns: repeat(4, 1fr);
gap: 20px;
}
@media (max-width: 768px) {
.xx-grid-2,
.xx-grid-3,
.xx-grid-4 {
grid-template-columns: 1fr;
}
}
/* ==================== 配额展示 ==================== */
.xx-quota-item {
padding: 20px;
background: var(--bg-primary);
border: 1px solid var(--border-color);
border-radius: var(--radius-lg);
transition: var(--transition-all);
}
.xx-quota-item:hover {
border-color: var(--primary-color);
box-shadow: 0 8px 24px
color-mix(in srgb, var(--primary-color) 10%, transparent);
}
/* ==================== 进度条 ==================== */
.xx-progress {
margin-top: 12px;
}
/* ==================== Ant Design 覆盖样式 ==================== */
/* Table overrides */
.ant-table-wrapper .ant-table-thead > tr > th {
background: var(--bg-secondary) !important;
font-weight: var(--font-weight-bold) !important;
color: var(--text-primary) !important;
border-bottom: 2px solid var(--border-color) !important;
padding: 14px var(--space-md) !important;
}
.ant-table-wrapper .ant-table-tbody > tr > td {
padding: 14px var(--space-md) !important;
border-bottom: 1px solid var(--color-gray-100) !important;
}
.ant-table-wrapper .ant-table-tbody > tr:hover > td {
background: var(--color-gray-50) !important;
}
/* Card overrides */
.ant-card {
border-radius: var(--radius-xl) !important;
border: 1px solid var(--border-color) !important;
}
.ant-card-head {
border-bottom: 1px solid var(--border-color) !important;
min-height: 52px !important;
padding: 0 var(--space-lg) !important;
}
.ant-card-head-title {
font-weight: var(--font-weight-bold) !important;
font-size: var(--font-size-md) !important;
color: var(--text-primary) !important;
}
.ant-card-body {
padding: 20px var(--space-lg) !important;
}
/* Modal overrides */
.ant-modal-content {
border-radius: var(--radius-xl) !important;
overflow: hidden;
}
.ant-modal-header {
padding: 20px var(--space-lg) !important;
background: var(--bg-primary) !important;
}
.ant-modal-title {
font-weight: var(--font-weight-bold) !important;
font-size: var(--font-size-lg) !important;
color: var(--text-primary) !important;
}
.ant-modal-body {
padding: var(--space-lg) !important;
}
.ant-modal-footer {
padding: var(--space-md) var(--space-lg) !important;
}
/* Button overrides */
.ant-btn-primary {
background: var(--gradient-primary) !important;
border: none !important;
border-radius: var(--radius-md) !important;
box-shadow: var(--shadow-primary) !important;
height: auto !important;
padding: 10px 20px !important;
font-weight: var(--font-weight-bold) !important;
}
.ant-btn-primary:hover {
background: var(--gradient-primary) !important;
box-shadow: var(--shadow-hover) !important;
transform: translateY(-1px);
}
/* Tag overrides */
.ant-tag {
border-radius: var(--radius-xs) !important;
padding: var(--space-xs) 12px !important;
font-weight: var(--font-weight-medium) !important;
}
/* Select overrides */
.ant-select-selector {
border-radius: var(--radius-md) !important;
border-color: var(--border-color) !important;
}
.ant-select:not(.ant-select-disabled):hover .ant-select-selector {
border-color: var(--primary-color) !important;
}
.ant-select-focused .ant-select-selector {
border-color: var(--primary-color) !important;
box-shadow: 0 0 0 3px
color-mix(in srgb, var(--primary-color) 10%, transparent) !important;
}
/* Input overrides */
.ant-input {
border-radius: var(--radius-md) !important;
border-color: var(--border-color) !important;
padding: 10px 14px !important;
}
.ant-input:hover {
border-color: var(--primary-color) !important;
}
.ant-input:focus {
border-color: var(--primary-color) !important;
box-shadow: 0 0 0 3px
color-mix(in srgb, var(--primary-color) 10%, transparent) !important;
}
/* Progress overrides */
.ant-progress-inner {
background: var(--color-gray-100) !important;
border-radius: var(--radius-xs) !important;
}
.ant-progress-bg {
border-radius: var(--radius-xs) !important;
}
+211
View File
@@ -0,0 +1,211 @@
/**
*
* Header Sidebar
*/
import React from "react";
import {
DashboardOutlined,
VideoCameraOutlined,
FileOutlined,
AudioOutlined,
FileTextOutlined,
TrophyOutlined,
AppstoreOutlined,
HistoryOutlined,
ControlOutlined,
CrownOutlined,
ScanOutlined,
EditOutlined,
FolderOutlined,
} from "@ant-design/icons";
/** 导航项定义 */
export interface NavItem {
key: string;
label: string;
path: string;
icon: React.ReactNode;
}
/** 导航分组定义 */
export interface NavGroup {
title: string;
items: NavItem[];
}
/**
* Header 使
*/
export const NAV_ITEMS: NavItem[] = [
{
key: "dashboard",
label: "概览",
path: "/app/dashboard",
icon: <DashboardOutlined />,
},
{
key: "assets",
label: "素材库",
path: "/app/assets",
icon: <FileOutlined />,
},
{
key: "titles",
label: "标题库",
path: "/app/titles",
icon: <FileTextOutlined />,
},
{
key: "voices",
label: "配音库",
path: "/app/voices",
icon: <AudioOutlined />,
},
{
key: "voice-clone",
label: "我的音色",
path: "/app/voice-clone",
icon: <AudioOutlined />,
},
{
key: "voice-materials",
label: "配音素材库",
path: "/app/voice-materials",
icon: <AudioOutlined />,
},
{
key: "templates",
label: "模板库",
path: "/app/templates",
icon: <AppstoreOutlined />,
},
{
key: "editing-planner",
label: "剪辑编辑器",
path: "/app/editing-planner",
icon: <EditOutlined />,
},
{
key: "my-templates",
label: "我的模板",
path: "/app/my-templates",
icon: <FolderOutlined />,
},
{
key: "generate",
label: "一键生成",
path: "/app/generate",
icon: <VideoCameraOutlined />,
},
{
key: "history",
label: "任务历史",
path: "/app/history",
icon: <HistoryOutlined />,
},
{
key: "products",
label: "成品库",
path: "/app/products",
icon: <TrophyOutlined />,
},
{
key: "duplication",
label: "查重",
path: "/app/duplication",
icon: <ScanOutlined />,
},
];
/**
* Sidebar 使
*/
export const NAV_GROUPS: NavGroup[] = [
{
title: "创作工具",
items: [
{
key: "dashboard",
label: "首页",
path: "/app/dashboard",
icon: <DashboardOutlined />,
},
{
key: "generate",
label: "一键生成",
path: "/app/generate",
icon: <VideoCameraOutlined />,
},
],
},
{
title: "资源管理",
items: [
{
key: "assets",
label: "素材库",
path: "/app/assets",
icon: <FileOutlined />,
},
{
key: "voices",
label: "配音库",
path: "/app/voices",
icon: <AudioOutlined />,
},
{
key: "voice-clone",
label: "我的音色",
path: "/app/voice-clone",
icon: <AudioOutlined />,
},
{
key: "voice-materials",
label: "配音素材库",
path: "/app/voice-materials",
icon: <AudioOutlined />,
},
{
key: "titles",
label: "标题库",
path: "/app/titles",
icon: <FileTextOutlined />,
},
{
key: "products",
label: "成片库",
path: "/app/products",
icon: <TrophyOutlined />,
},
{
key: "templates",
label: "模板库",
path: "/app/templates",
icon: <AppstoreOutlined />,
},
],
},
{
title: "系统",
items: [
{
key: "history",
label: "任务历史",
path: "/app/history",
icon: <HistoryOutlined />,
},
{
key: "admin",
label: "控制台",
path: "/app/admin",
icon: <ControlOutlined />,
},
{
key: "subscription",
label: "订阅管理",
path: "/app/subscription",
icon: <CrownOutlined />,
},
],
},
];
+40
View File
@@ -75,6 +75,35 @@
margin: 0;
}
/* V21 卡片 */
.xx-card {
background: rgba(255, 255, 255, 0.94);
border: 1px solid rgba(226, 232, 240, 0.95);
border-radius: var(--radius-lg);
box-shadow: 0 10px 30px rgba(15, 23, 42, 0.06);
padding: 24px;
transition: all 0.3s;
}
.xx-card:hover {
box-shadow: 0 16px 40px rgba(15, 23, 42, 0.08);
}
.xx-card .ant-card-head {
border-bottom: 1px solid rgba(226, 232, 240, 0.8);
padding: 20px 24px;
}
.xx-card .ant-card-head-title {
font-weight: 800;
font-size: 17px;
color: var(--slate, #0f172a);
}
.xx-card .ant-card-body {
padding: 24px;
}
/* 统计卡片网格 - 4列 */
.xx-grid-4 {
display: grid;
@@ -273,6 +302,17 @@
box-shadow: 0 0 0 3px rgba(79, 70, 229, 0.1) !important;
}
/* V21 Select */
.xx-select {
border-radius: var(--radius-md) !important;
}
.xx-select:hover,
.xx-select:focus {
border-color: var(--indigo, #4f46e5) !important;
box-shadow: 0 0 0 3px rgba(79, 70, 229, 0.1) !important;
}
/* V21 Tag */
.xx-tag {
border-radius: var(--radius-xs) !important;
+48
View File
@@ -612,6 +612,54 @@
flex: 1;
}
/* ============================================================
按钮匹配原型 .btn .ghost / .btn .primary
============================================================ */
.xx-btn {
display: inline-flex;
align-items: center;
justify-content: center;
gap: 6px;
height: 42px;
padding: 0 20px;
border-radius: var(--radius-sm);
font-size: 14px;
font-weight: 600;
cursor: pointer;
transition: all 0.15s ease;
border: none;
outline: none;
white-space: nowrap;
}
.xx-btn:disabled {
opacity: 0.5;
cursor: not-allowed;
}
.xx-btn-primary {
background: var(--gradient-primary);
color: var(--text-inverse);
box-shadow: 0 14px 26px rgba(79, 70, 229, 0.22);
}
.xx-btn-primary:hover:not(:disabled) {
transform: translateY(-2px);
box-shadow: 0 18px 34px rgba(79, 70, 229, 0.28);
}
.xx-btn-ghost {
background: var(--bg-primary);
border: 1px solid var(--border-color);
color: var(--text-secondary);
}
.xx-btn-ghost:hover:not(:disabled) {
border-color: var(--info-border);
color: var(--primary-dark);
background: var(--primary-soft);
}
/* ============================================================
右侧预览区 generate-preview
============================================================ */
+7
View File
@@ -170,6 +170,13 @@ export const router = createBrowserRouter([
Component: m.default,
})),
},
{
path: "my-voices",
lazy: () =>
import("@/pages/my-voices/MyVoices").then((m) => ({
Component: m.default,
})),
},
{
path: "accounts",
lazy: () =>
+10 -6
View File
@@ -1,17 +1,21 @@
"""
视频处理模块
轻量工具ffmpeg_utils / oss_helpers / dedup_helpers顶层直接导出
无额外依赖渲染相关组件UnifiedRenderService / RenderAdapter /
VideoProcessor 按需从子模块导入避免 __init__ 阶段引入
packages / DB 等重依赖
"""
# 共享工具模块(零外部依赖,供 editing_modes / generation / edit_plan_generation 等复用)
# 共享工具模块(供 editing_modes / generation / edit_plan_generation 等复用)
from . import dedup_helpers, ffmpeg_utils, oss_helpers
from .processor import VideoProcessor, VideoResult
from .render_adapter import RenderAdapter, RenderAdapterResult
from .unified_render_service import RenderResult, UnifiedRenderService
__all__ = [
"VideoProcessor",
"VideoResult",
"ffmpeg_utils",
"oss_helpers",
"dedup_helpers",
"UnifiedRenderService",
"RenderResult",
"RenderAdapter",
"RenderAdapterResult",
]
+3 -7
View File
@@ -93,16 +93,12 @@ class VideoFingerprint:
resolution: tuple[int, int]
def to_dict(self) -> dict:
# 注意:color_histograms 里的值可能是 np.float32(来自 cv2.normalize),
# 直接存进 dict 后 SQLAlchemy JSON 序列化会报 "float32 is not JSON serializable"。
# 这里统一转成 Python 原生 float。
native_histograms = [[float(v) for v in hist] for hist in self.color_histograms]
return {
"md5": self.md5,
"keyframe_phashes": self.keyframe_phashes,
"color_histograms": native_histograms,
"duration": float(self.duration),
"resolution": [int(self.resolution[0]), int(self.resolution[1])],
"color_histograms": self.color_histograms,
"duration": self.duration,
"resolution": list(self.resolution),
}
+10 -79
View File
@@ -42,10 +42,6 @@ XFADE_TRANSITION_MAP: dict[str, str] = {
DEFAULT_TRANSITION_DURATION = 0.5
# FFmpeg 执行默认超时(秒),防止 FFmpeg hang 住导致 worker 永久阻塞
# 默认 30 分钟,足够处理大部分短视频渲染;超长视频可单独传参覆盖
DEFAULT_FFMPEG_TIMEOUT = 1800
# ── FFmpeg 执行 ───────────────────────────────────────────────────────────────
@@ -54,14 +50,12 @@ def run_ffmpeg(
command: list[str],
*,
capture_output: bool = True,
timeout: int | None = DEFAULT_FFMPEG_TIMEOUT,
) -> tuple[str, str]:
"""执行 FFmpeg 命令。
Args:
command: 完整的 ffmpeg 命令列表 "ffmpeg" 本身
capture_output: 是否捕获 stdout/stderr
timeout: 超时时间默认 1800s30分钟None 表示不设超时不推荐
Returns:
(stdout, stderr) 元组
@@ -69,7 +63,6 @@ def run_ffmpeg(
Raises:
subprocess.CalledProcessError: 命令执行失败时抛出
异常信息包含完整 stderr 以便排查
subprocess.TimeoutExpired: 超时未完成时抛出FFmpeg 进程会被 kill
"""
try:
result = subprocess.run( # nosec B603
@@ -78,16 +71,8 @@ def run_ffmpeg(
stdout=subprocess.PIPE if capture_output else None,
stderr=subprocess.PIPE if capture_output else None,
text=True,
timeout=timeout,
)
return (result.stdout or "", result.stderr or "")
except subprocess.TimeoutExpired as e:
logger.error(
"FFmpeg 命令超时 (%ds): command=%s",
timeout or -1,
" ".join(str(c) for c in command[:20]),
)
raise
except subprocess.CalledProcessError as e:
# 把完整 stderr 打到日志,方便排查 exit code 183 等问题
stderr_text = (e.stderr or "").strip()
@@ -100,41 +85,6 @@ def run_ffmpeg(
raise
def probe_has_audio(local_path: str | Path) -> bool:
"""探测文件是否包含音频流。
Args:
local_path: 本地文件路径
Returns:
True 表示有音频流或探测失败保守返回False 表示确认无音频流
"""
try:
result = subprocess.run( # nosec B603
[
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"a:0",
"-show_entries",
"stream=codec_type",
"-of",
"default=noprint_wrappers=1:nokey=1",
str(local_path),
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=10,
)
return result.stdout.strip() == "audio"
except Exception:
# 探测失败保守返回 True,让 FFmpeg 自己处理(避免误删音频)
return True
def probe_duration(local_path: str | Path) -> float:
"""用 ffprobe 获取视频时长(秒)。
@@ -163,14 +113,10 @@ def probe_duration(local_path: str | Path) -> float:
def probe_video_info(video_path: str) -> dict[str, Any]:
"""获取视频信息(宽、高、时长、fps、编码、像素格式)。
"""获取视频信息(宽、高、时长、fps)。
Returns:
{
"width": int, "height": int, "duration": float, "fps": float,
"video_codec": str, "audio_codec": str, "pix_fmt": str,
"has_audio": bool,
}
{"width": int, "height": int, "duration": float, "fps": float}
失败时返回默认值
"""
try:
@@ -179,8 +125,10 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"v:0",
"-show_entries",
"stream=width,height,r_frame_rate,duration,codec_name,codec_type,pix_fmt",
"stream=width,height,r_frame_rate,duration",
"-show_entries",
"format=duration",
"-of",
@@ -191,25 +139,19 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=15,
)
import json
info = json.loads(result.stdout)
streams = info.get("streams", [])
stream = info.get("streams", [{}])[0]
fmt = info.get("format", {})
video_stream = next((s for s in streams if s.get("codec_type") == "video"), {})
audio_stream = next((s for s in streams if s.get("codec_type") == "audio"), {})
width = int(video_stream.get("width", DEFAULT_OUTPUT_WIDTH))
height = int(video_stream.get("height", DEFAULT_OUTPUT_HEIGHT))
video_codec = video_stream.get("codec_name", "") or ""
pix_fmt = video_stream.get("pix_fmt", "") or ""
width = int(stream.get("width", DEFAULT_OUTPUT_WIDTH))
height = int(stream.get("height", DEFAULT_OUTPUT_HEIGHT))
# 解析帧率
fps_str = video_stream.get("r_frame_rate", "25/1")
fps_str = stream.get("r_frame_rate", "25/1")
if "/" in fps_str:
num, den = fps_str.split("/")
fps = float(num) / float(den) if float(den) > 0 else DEFAULT_FPS
@@ -217,20 +159,13 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
fps = float(fps_str) if fps_str else DEFAULT_FPS
# 时长
duration = float(fmt.get("duration", 0)) or float(video_stream.get("duration", 0))
has_audio = bool(audio_stream)
audio_codec = audio_stream.get("codec_name", "") or ""
duration = float(fmt.get("duration", 0)) or float(stream.get("duration", 0))
return {
"width": width,
"height": height,
"duration": duration,
"fps": round(fps, 2),
"video_codec": video_codec,
"audio_codec": audio_codec,
"pix_fmt": pix_fmt,
"has_audio": has_audio,
}
except Exception as e:
logger.warning("获取视频信息失败: %s, error: %s", video_path, e)
@@ -239,10 +174,6 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
"height": DEFAULT_OUTPUT_HEIGHT,
"duration": 0.0,
"fps": DEFAULT_FPS,
"video_codec": "",
"audio_codec": "",
"pix_fmt": "",
"has_audio": True,
}
+10 -82
View File
@@ -9,7 +9,6 @@ from __future__ import annotations
import hashlib
import logging
import os
import threading
from pathlib import Path
from typing import Optional
from urllib.parse import urlparse
@@ -18,13 +17,6 @@ import oss2
logger = logging.getLogger(__name__)
# OSS 上传配置
OSS_CONNECT_TIMEOUT = 10 # 连接超时(秒),防止 TCP 握手挂死
OSS_UPLOAD_TOTAL_TIMEOUT = 300 # 单文件上传总超时(秒),防止网络慢时无限卡住
OSS_MULTIPART_THRESHOLD = 100 * 1024 * 1024 # 分片上传阈值:100MB 以上走分片
OSS_PART_SIZE = 8 * 1024 * 1024 # 分片大小:8MB
OSS_MULTIPART_NUM_THREADS = 3 # 分片上传并发数
# ── OSS 配置 ──────────────────────────────────────────────────────────────────
@@ -51,9 +43,6 @@ def oss_bucket() -> oss2.Bucket | None:
P0-2 修复endpoint 不带 scheme 时自动补 https:// 前缀
确保 sign_url 等依赖 scheme 的方法返回 HTTPS URL
P0-staging 修复增加 connect_timeout=10s防止网络抖动时
TCP 握手阶段无限挂死导致 worker 进程卡死
Returns:
oss2.Bucket 实例配置缺失时返回 None
"""
@@ -64,12 +53,7 @@ def oss_bucket() -> oss2.Bucket | None:
# endpoint 无 scheme 时补 https://,与 API 端 storage.py 保持一致
if not endpoint.startswith(("http://", "https://")):
endpoint = f"https://{endpoint}"
return oss2.Bucket(
oss2.Auth(access_key_id, access_key_secret),
endpoint,
bucket_name,
connect_timeout=OSS_CONNECT_TIMEOUT,
)
return oss2.Bucket(oss2.Auth(access_key_id, access_key_secret), endpoint, bucket_name)
def normalize_storage_key(storage_key_or_url: str) -> str:
@@ -112,9 +96,6 @@ def download_asset(asset_storage_key: str, local_path: Path) -> bool:
def upload_to_oss(local_path: Path, storage_key: str) -> str | None:
"""上传文件到 OSS,返回公开 URL。
大文件>100MB自动走分片上传降低内存峰值减少 OOM 风险
上传加总超时保护默认 300s防止网络异常时无限挂死
Args:
local_path: 本地文件路径
storage_key: 目标存储键
@@ -125,71 +106,18 @@ def upload_to_oss(local_path: Path, storage_key: str) -> str | None:
bucket = oss_bucket()
if bucket is None:
return None
result: dict = {"url": None, "error": None, "file_size": 0}
done = threading.Event()
def _do_upload():
try:
# 尝试获取文件大小,用于分片判断和日志;stat 失败时 fallback 走普通上传
try:
file_size = local_path.stat().st_size
result["file_size"] = file_size
use_multipart = file_size >= OSS_MULTIPART_THRESHOLD
except OSError:
use_multipart = False
file_size = 0
if use_multipart:
# 分片上传:降低内存峰值,每片 8MB,3 线程并发
logger.info(
"大文件分片上传: storage_key=%s, size=%.1fMB, part_size=%dMB, threads=%d",
storage_key[:80],
file_size / 1024 / 1024,
OSS_PART_SIZE // 1024 // 1024,
OSS_MULTIPART_NUM_THREADS,
)
oss2.resumable_upload(
bucket,
storage_key,
str(local_path),
multipart_threshold=OSS_MULTIPART_THRESHOLD,
part_size=OSS_PART_SIZE,
num_threads=OSS_MULTIPART_NUM_THREADS,
)
else:
bucket.put_object_from_file(storage_key, str(local_path))
# 构造返回 URL
settings = oss_settings()
if settings:
_, _, endpoint, bucket_name = settings
endpoint_clean = endpoint.replace("https://", "").replace("http://", "")
result["url"] = f"https://{bucket_name}.{endpoint_clean}/{storage_key}"
except Exception as e:
result["error"] = e
logger.exception("上传 OSS 失败: %s", storage_key)
finally:
done.set()
upload_thread = threading.Thread(target=_do_upload, daemon=True)
upload_thread.start()
finished = done.wait(timeout=OSS_UPLOAD_TOTAL_TIMEOUT)
if not finished:
logger.error(
"OSS 上传超时(%.0fs),强制中止: storage_key=%s, size=%.1fMB",
OSS_UPLOAD_TOTAL_TIMEOUT,
storage_key[:80],
result["file_size"] / 1024 / 1024 if result["file_size"] else 0,
)
try:
bucket.put_object_from_file(storage_key, str(local_path))
settings = oss_settings()
if settings:
_, _, endpoint, bucket_name = settings
endpoint_clean = endpoint.replace("https://", "").replace("http://", "")
return f"https://{bucket_name}.{endpoint_clean}/{storage_key}"
return None
if result["error"]:
except Exception:
logger.exception("上传 OSS 失败: %s", storage_key)
return None
return result["url"]
def get_signed_download_url(storage_key_or_url: str, expires_seconds: int = 3600) -> str | None:
"""生成预签名下载 URL(用于私有 bucket 的 URL 校验或临时下载)。
+12 -22
View File
@@ -21,10 +21,17 @@ from typing import Any, Callable
from sqlalchemy.orm import Session
from video_processing.oss_helpers import download_asset, upload_to_oss
from video_processing.unified_render_service import RenderResult, UnifiedRenderService
from video_processing.unified_render_service import (
RenderResult,
UnifiedRenderService,
)
from packages.adapters.sqlalchemy_impl.edit_plan_clip_repository import SQLAlchemyEditPlanClipRepository
from packages.adapters.sqlalchemy_impl.edit_plan_repository import SQLAlchemyEditPlanRepository
from packages.adapters.sqlalchemy_impl.edit_plan_clip_repository import (
SQLAlchemyEditPlanClipRepository,
)
from packages.adapters.sqlalchemy_impl.edit_plan_repository import (
SQLAlchemyEditPlanRepository,
)
from packages.domain.edit_plan import EditPlan, EditPlanStatus
from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus
@@ -133,7 +140,7 @@ class RenderAdapter:
)
logger.info(
"开始渲染: plan_id=%s job_id=%s ready_clips=%d engine=unified",
"开始渲染: plan_id=%s job_id=%s ready_clips=%d",
plan_id,
job_id,
len(ready_clips),
@@ -169,18 +176,6 @@ class RenderAdapter:
self._report_progress(progress_cb, 100.0, "渲染完成")
logger.info(
"[render-adapter] render success: plan_id=%s job_id=%s engine=unified "
"duration=%.2fs file_size=%d resolution=%dx%d clip_count=%d",
plan_id,
job_id,
result.duration,
result.file_size,
result.width,
result.height,
len(ready_clips),
)
return RenderAdapterResult(
success=True,
output_url=output_url or "",
@@ -193,12 +188,7 @@ class RenderAdapter:
)
except Exception as exc:
logger.exception(
"[render-adapter] render failed: plan_id=%s job_id=%s engine=unified error=%s",
plan_id,
job_id,
str(exc)[:200],
)
logger.exception("渲染失败: plan_id=%s", plan_id)
return RenderAdapterResult(
success=False,
error_message=str(exc)[:500],
@@ -1,204 +0,0 @@
"""渲染引擎 Feature Flag 解析器。
封装渲染引擎选择逻辑支持
- 环境变量作为默认值RENDER_ENGINE=legacy/unified
- Redis Feature Flag 运行时覆盖白名单 + 百分比 + 全局开关
- 定时刷新支持热更新不重启 worker
使用方式
resolver = RenderEngineResolver(redis_url="redis://...", default_engine="legacy")
engine = resolver.get_engine(user_id="user123")
# engine: "legacy" 或 "unified"
"""
from __future__ import annotations
import logging
import threading
from typing import Optional
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
FeatureFlagStore,
InMemoryFeatureFlagStore,
RedisFeatureFlagStore,
)
logger = logging.getLogger(__name__)
# Feature Flag 名称常量
FLAG_RENDER_ENGINE = "render_engine"
# 引擎常量
ENGINE_LEGACY = "legacy"
ENGINE_UNIFIED = "unified"
VALID_ENGINES = {ENGINE_LEGACY, ENGINE_UNIFIED}
class RenderEngineResolver:
"""渲染引擎选择器。
判定逻辑从高到低
1. Redis flag 白名单匹配 unified
2. Redis flag 百分比命中 unified
3. Redis flag 全局开启100% unified
4. 环境变量默认值 legacy / unified
Redis 不可用时自动降级到环境变量默认值不影响业务
"""
def __init__(
self,
default_engine: str = ENGINE_LEGACY,
redis_url: Optional[str] = None,
refresh_interval: float = 30.0,
store: Optional[FeatureFlagStore] = None,
) -> None:
"""
Args:
default_engine: 环境变量默认的引擎名legacy / unified
redis_url: Redis 连接 URL None 时使用内存实现测试用
refresh_interval: Redis flag 配置刷新间隔
store: 直接传入 store 实例测试用优先级高于 redis_url
"""
self._default_engine = default_engine.lower() if default_engine else ENGINE_LEGACY
if self._default_engine not in VALID_ENGINES:
logger.warning(
"Invalid default engine '%s', fallback to '%s'",
self._default_engine,
ENGINE_LEGACY,
)
self._default_engine = ENGINE_LEGACY
if store is not None:
self._store = store
elif redis_url:
self._store = RedisFeatureFlagStore(redis_url=redis_url)
else:
self._store = InMemoryFeatureFlagStore()
logger.info("No Redis configured, using in-memory feature flag store")
self._refresh_interval = refresh_interval
self._lock = threading.Lock()
self._cached_config: Optional[FeatureFlagConfig] = None
self._last_refresh: float = 0.0
def _maybe_refresh(self) -> None:
"""惰性刷新配置,超过刷新间隔时从存储重新读取。"""
import time
now = time.time()
if now - self._last_refresh < self._refresh_interval:
return
try:
config = self._store.get(FLAG_RENDER_ENGINE)
with self._lock:
self._cached_config = config
self._last_refresh = now
except Exception as exc:
logger.warning("Failed to refresh render engine flag: %s", exc)
# 刷新失败时保留旧缓存,不中断业务
if self._cached_config is None:
# 首次就读失败,设一个默认值
with self._lock:
self._cached_config = FeatureFlagConfig(name=FLAG_RENDER_ENGINE)
self._last_refresh = now
def _get_config(self) -> FeatureFlagConfig:
"""获取当前 flag 配置(带缓存)。"""
if self._cached_config is None:
self._maybe_refresh()
else:
self._maybe_refresh()
return self._cached_config or FeatureFlagConfig(name=FLAG_RENDER_ENGINE)
def get_engine(self, user_id: Optional[str] = None) -> str:
"""获取当前应该使用的渲染引擎。
Args:
user_id: 用户ID用于白名单匹配和百分比哈希
None 时只看全局开关
Returns:
"legacy" "unified"
"""
config = self._get_config()
# 全局关闭 → 用默认值
if not config.enabled:
return self._default_engine
# 白名单匹配 / 百分比命中 → unified
if config.is_active(user_id):
return ENGINE_UNIFIED
# 未命中灰度 → 用默认值
return self._default_engine
def should_use_unified(self, user_id: Optional[str] = None) -> bool:
"""便捷方法:是否应该使用统一渲染引擎。"""
return self.get_engine(user_id) == ENGINE_UNIFIED
def force_refresh(self) -> None:
"""强制立即刷新配置(用于管理接口修改后立即生效)。"""
self._last_refresh = 0.0
if isinstance(self._store, RedisFeatureFlagStore):
self._store.invalidate_cache(FLAG_RENDER_ENGINE)
self._maybe_refresh()
def get_config_snapshot(self) -> dict:
"""获取当前配置快照(用于管理接口展示)。"""
config = self._get_config()
return {
"flag_name": FLAG_RENDER_ENGINE,
"default_engine": self._default_engine,
"enabled": config.enabled,
"percentage": config.percentage,
"whitelist": sorted(config.whitelist),
"refresh_interval": self._refresh_interval,
"last_refresh": self._last_refresh,
}
def set_flag(self, config: FeatureFlagConfig) -> None:
"""设置 flag 配置(管理接口用)。"""
config.name = FLAG_RENDER_ENGINE
self._store.set(config)
self.force_refresh()
# 全局单例
_resolver: Optional[RenderEngineResolver] = None
_resolver_lock = threading.Lock()
def get_render_engine_resolver() -> RenderEngineResolver:
"""获取全局单例(基于 worker 配置)。"""
global _resolver
if _resolver is not None:
return _resolver
with _resolver_lock:
if _resolver is not None:
return _resolver
try:
from worker_app.core.config import get_settings
settings = get_settings()
redis_url = getattr(settings, "redis_url", None) or getattr(settings, "broker_url", None)
default = getattr(settings, "render_engine", ENGINE_LEGACY)
_resolver = RenderEngineResolver(
default_engine=default,
redis_url=redis_url,
)
logger.info(
"RenderEngineResolver initialized: default=%s, redis=%s",
default,
bool(redis_url),
)
except Exception as exc:
logger.warning("Failed to init RenderEngineResolver from settings: %s", exc)
_resolver = RenderEngineResolver(default_engine=ENGINE_LEGACY)
return _resolver
@@ -22,8 +22,8 @@
from __future__ import annotations
import logging
import os
import subprocess
import time
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any
@@ -295,7 +295,7 @@ 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 # noqa: E501
Format: Name, Fontname, Fontsize, PrimaryColour, SecondaryColour, OutlineColour, BackColour, Bold, Italic, Underline, StrikeOut, ScaleX, ScaleY, Spacing, Angle, BorderStyle, Outline, Shadow, Alignment, MarginL, MarginR, MarginV, Encoding
{chr(10).join(styles)}
[Events]
@@ -356,8 +356,6 @@ def _resolve_layer_role(clip_type: str, config: dict[str, Any]) -> str:
# main type
if role == "b_roll":
return "broll"
if role == "audio":
return "audio"
return "main"
@@ -417,16 +415,9 @@ class UnifiedRenderService:
1. 视频主渲染直通或完整链路
2. 如有 title/subtitle叠加 ASS 字幕
音频后处理
1. 主图层音频 concat 拼接
2. 独立音频轨 amix 混入
3. 合并到输出视频
Raises:
ValueError: 没有可渲染的片段时抛出
"""
t_start = time.time()
# 1. 解析 clips → ResolvedClips(跳过无素材的 clip
resolved = self._resolve_clips()
if not resolved:
@@ -441,102 +432,18 @@ class UnifiedRenderService:
# 4. 生成 ASS 字幕文件(如果有 title/subtitle 配置)
ass_path = self._maybe_generate_ass(video_duration)
# 灰度埋点:开始渲染
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",
self.plan.id,
len(resolved),
layer_roles,
clip_counts,
)
# 5. 视频主渲染
t_video_start = time.time()
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)
pass_through_has_audio = False
used_stream_copy = False
if is_pass_through:
# 先尝试 stream copy 优化(无重编码,性能提升 10 倍+)
# 条件不满足或失败时回退到带滤镜的直通渲染
stream_copy_ok = self._try_render_stream_copy(
layers, output_path, ass_path=ass_path, video_duration=video_duration
)
if stream_copy_ok:
used_stream_copy = True
# stream copy 模式下,直接探测输出是否有音频
clip = layers[0].clips[0]
info = probe_video_info(str(clip.local_path))
pass_through_has_audio = info.get("has_audio", True)
else:
# 回退到带滤镜的直通渲染
pass_through_has_audio = self._render_pass_through(
layers, output_path, ass_path=ass_path, video_duration=video_duration
)
# 5. 视频主渲染
if self._can_use_pass_through(layers):
self._render_pass_through(layers, output_path, ass_path=ass_path)
else:
filter_complex, input_args = self._build_filter_complex(layers, ass_path=ass_path)
self._execute_ffmpeg(filter_complex, input_args, video_only_path)
self._execute_ffmpeg(filter_complex, input_args, output_path)
t_video_end = time.time()
video_render_ms = int((t_video_end - t_video_start) * 1000)
logger.info(
"[unified-render] video render done: plan_id=%s duration_ms=%d pass_through=%s stream_copy=%s",
self.plan.id,
video_render_ms,
is_pass_through,
used_stream_copy,
)
# 6. 音频后处理混音(直通场景已合并处理,跳过)
t_audio_start = time.time()
audio_mix_ms = 0
has_audio = False
if is_pass_through:
# 直通场景已在一次调用中完成视频+音频
has_audio = pass_through_has_audio
else:
audio_path = self._mix_audio(layers, video_duration)
t_audio_end = time.time()
audio_mix_ms = int((t_audio_end - t_audio_start) * 1000)
has_audio = audio_path is not None
if has_audio:
logger.info(
"[unified-render] audio mix done: plan_id=%s duration_ms=%d",
self.plan.id,
audio_mix_ms,
)
# 7. 合并音视频
self._merge_audio_video(video_only_path, audio_path, output_path)
else:
# 无音频,直接用无声视频
import shutil
shutil.copy2(video_only_path, output_path)
# 8. 探测输出
# 6. 探测输出
duration, file_size, width, height = self._probe_output(output_path)
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 "
"output_duration=%.2fs output_size=%d resolution=%dx%d has_audio=%s",
self.plan.id,
t_total,
video_render_ms,
audio_mix_ms if has_audio else 0,
duration,
file_size,
width,
height,
has_audio,
)
return RenderResult(
output_path=output_path,
duration=duration,
@@ -563,7 +470,14 @@ class UnifiedRenderService:
if not main_layer or not main_layer.clips:
return 0.0
total = sum(UnifiedRenderService._clip_effective_duration(c) for c in main_layer.clips)
total = sum(
(
min(c.duration, c.actual_duration)
if c.duration > 0 and c.actual_duration > 0
else (c.duration if c.duration > 0 else c.actual_duration)
)
for c in main_layer.clips
)
# 减去转场重叠时间(粗略估算)
n_clips = len(main_layer.clips)
@@ -632,207 +546,31 @@ class UnifiedRenderService:
return False
return True
def _can_use_stream_copy(
self,
clip: ResolvedClip,
*,
ass_path: Path | None = None,
video_duration: float = 0.0,
) -> tuple[bool, str]:
"""判断是否可以走 stream copy(流拷贝,不重编码)。
性能提升10 倍以上典型场景从 20s 1-2s
条件
1. 视频编码为 h264输出目标也是 h264
2. 像素格式为 yuv420p
3. 分辨率与输出一致不需要 scale/crop
4. 帧率与输出一致误差 < 0.1fps
5. 无字幕叠加字幕需要滤镜
6. trim 需求 trim 后恰好等于原时长
7. 无转场无特效 clip 直通已保证
Returns:
(是否可以 copy, 原因说明)
"""
# 有字幕 → 需要滤镜 → 不能 copy
if ass_path is not None:
return False, "有字幕叠加"
# 探测输入视频参数
info = probe_video_info(str(clip.local_path))
# 编码必须是 h264
if info.get("video_codec", "") != "h264":
return False, f"视频编码不是h264: {info.get('video_codec', 'unknown')}"
# 像素格式必须是 yuv420p
if info.get("pix_fmt", "") != "yuv420p":
return False, f"像素格式不是yuv420p: {info.get('pix_fmt', 'unknown')}"
# 分辨率必须一致
if info.get("width", 0) != self.output_width or info.get("height", 0) != self.output_height:
return False, (
f"分辨率不匹配: "
f"{info.get('width', 0)}x{info.get('height', 0)} "
f"vs {self.output_width}x{self.output_height}"
)
# 帧率必须一致(误差 < 0.1fps
fps_diff = abs(info.get("fps", 0) - self.output_fps)
if fps_diff > 0.1:
return False, f"帧率不匹配: {info.get('fps', 0)} vs {self.output_fps}"
# 检查是否需要 trim
effective_duration = UnifiedRenderService._clip_effective_duration(clip)
if effective_duration > 0:
# 有 trim 需求但视频时长足够,可用 -ss/-t 实现 copy trim
input_duration = info.get("duration", 0)
if input_duration <= 0:
return False, "无法探测输入时长"
# trim 起始点 + 目标时长 <= 输入时长
start_time = getattr(clip, "start_time", 0) or 0
if start_time + effective_duration > input_duration + 0.1:
return False, "trim 超出输入时长"
# video_duration 截断
if video_duration > 0 and effective_duration > 0:
final_duration = min(effective_duration, video_duration)
if final_duration != effective_duration:
# 也需要截断,但 -t 可以 copy 模式下用
pass
return True, "所有条件满足"
def _try_render_stream_copy(
self,
layers: list[RenderLayer],
output_path: Path,
*,
ass_path: Path | None = None,
video_duration: float = 0.0,
) -> bool:
"""尝试 stream copy 渲染,成功返回 True,失败返回 False(调用方回退到重编码)。
stream copy 模式不重编码直接拷贝视频/音频流性能提升 10 +
仅用于单 clip 直通场景且满足 copy 条件
"""
clip = layers[0].clips[0]
role = layers[0].role
# 判断是否满足 copy 条件
can_copy, reason = self._can_use_stream_copy(clip, ass_path=ass_path, video_duration=video_duration)
if not can_copy:
logger.info(
"[unified-render] stream_copy 跳过: plan_id=%s reason=%s",
self.plan.id,
reason,
)
return False
# 构建 copy 命令
command = [
FFMPEG_BIN,
"-y",
]
# trim 支持(-ss 放在 -i 前 = input seeking,速度更快但精度稍差;
# 放在 -i 后 = output seeking,精度高但慢)
# 这里用 output seeking 保证精度,反正 copy 模式已经很快了
start_time = getattr(clip, "start_time", 0) or 0
effective_duration = UnifiedRenderService._clip_effective_duration(clip)
command.extend(["-i", str(clip.local_path)])
if start_time > 0:
command.extend(["-ss", f"{start_time:.3f}"])
# 计算最终时长
final_duration = effective_duration
if video_duration > 0 and (final_duration <= 0 or final_duration > video_duration):
final_duration = video_duration
if final_duration > 0:
command.extend(["-t", f"{final_duration:.3f}"])
# 流拷贝
command.extend(
[
"-c:v",
"copy",
"-c:a",
"copy",
"-movflags",
"+faststart",
str(output_path),
]
)
logger.info(
"[unified-render] stream_copy 渲染: plan_id=%s clip=%s role=%s duration=%.2fs",
self.plan.id,
clip.clip_id,
role,
final_duration,
)
try:
run_ffmpeg(command)
# 验证输出文件存在且有大小
if output_path.exists() and output_path.stat().st_size > 0:
logger.info(
"[unified-render] stream_copy 成功: plan_id=%s size=%d",
self.plan.id,
output_path.stat().st_size,
)
return True
else:
logger.warning("[unified-render] stream_copy 输出为空: plan_id=%s", self.plan.id)
return False
except (subprocess.CalledProcessError, subprocess.TimeoutExpired) as e:
logger.warning(
"[unified-render] stream_copy 失败,回退到重编码: plan_id=%s error=%s",
self.plan.id,
str(e)[:200],
)
# 清理可能的损坏输出文件
if output_path.exists():
try:
output_path.unlink()
except OSError:
pass
return False
def _render_pass_through(
self,
layers: list[RenderLayer],
output_path: Path,
*,
ass_path: Path | None = None,
video_duration: float = 0.0,
) -> bool:
"""单图层单 clip 直通渲染(使用 -vf 而非 -filter_complex),一次性输出带音频的最终视频。
self, layers: list[RenderLayer], output_path: Path, *, ass_path: Path | None = None
) -> None:
"""单图层单 clip 直通渲染(使用 -vf 而非 -filter_complex)。
性能优化
- 避免 filter_complex 的解析和调度开销单clip场景性能提升 ~30%
- 视频+音频一次FFmpeg调用完成省去后续音频提取+音视频合并两次调用
性能优化避免 filter_complex 的解析和调度开销
对于一镜到底场景性能提升 ~30%接近链路A水平
Args:
layers: 图层列表只有1个图层1个clip
output_path: 输出文件路径
ass_path: ASS 字幕文件路径有则叠加字幕
video_duration: 视频总时长用于截断音频0表示不额外截断
Returns:
True 表示输出包含音频近似判断实际以输出文件为准
"""
clip = layers[0].clips[0]
role = layers[0].role
# 构建视频滤镜链(与 _build_filter_complex 中预处理逻辑一致)
# 构建滤镜链(与 _build_filter_complex 中预处理逻辑一致)
filters: list[str] = []
# trim
effective_duration = UnifiedRenderService._clip_effective_duration(clip)
effective_duration = 0.0
if clip.duration > 0:
effective_duration = min(clip.duration, clip.actual_duration) if clip.actual_duration > 0 else clip.duration
elif clip.actual_duration > 0:
effective_duration = clip.actual_duration
if effective_duration > 0:
filters.append(f"trim=duration={effective_duration}")
@@ -854,16 +592,12 @@ class UnifiedRenderService:
# 字幕叠加
if ass_path is not None:
# ASS 文件路径需要转义:Windows 反斜杠转正斜杠,冒号转义
ass_filter_path = str(ass_path).replace("\\", "/").replace(":", "\\:")
filters.append(f"subtitles='{ass_filter_path}'")
vf_str = ",".join(filters)
# 最终输出时长:取 clip 有效时长和 video_duration 的较小值
final_duration = effective_duration
if video_duration > 0 and (final_duration <= 0 or final_duration > video_duration):
final_duration = video_duration
command = [
FFMPEG_BIN,
"-y",
@@ -881,27 +615,16 @@ class UnifiedRenderService:
"yuv420p",
"-movflags",
"+faststart",
"-an", # 直通模式暂不处理音频,音频统一在后续混音阶段处理
str(output_path),
]
# 音频处理:background 通常是图片无音频,跳过;其他编码为 aac
# background 以外的视频素材,默认带音频
has_audio = role != "background"
if has_audio:
command.extend(["-c:a", "aac", "-b:a", "128k"])
# 统一截断时长(同时作用于视频和音频)
if final_duration > 0:
command.extend(["-t", f"{final_duration:.3f}"])
command.append(str(output_path))
logger.info(
"直通渲染: plan_id=%s clip=%s role=%s duration=%.2fs has_audio=%s",
"直通渲染: plan_id=%s clip=%s role=%s duration=%.2fs",
self.plan.id,
clip.clip_id,
role,
effective_duration,
has_audio,
)
try:
run_ffmpeg(command)
@@ -915,8 +638,6 @@ class UnifiedRenderService:
)
raise
return has_audio
# ── 内部方法 ──────────────────────────────────────────────────────────────
def _resolve_clips(self) -> list[ResolvedClip]:
@@ -980,6 +701,7 @@ class UnifiedRenderService:
# 计算 PiP 位置
pip_width = int(self.output_width * _PIP_SCALE)
pip_height = int(self.output_height * _PIP_SCALE)
margin = 20 # 边距
if "overlay" in layer_map:
@@ -1037,7 +759,14 @@ class UnifiedRenderService:
filters: list[str] = []
# trim — 始终将输出截断到有效时长,防止 xfade offset 与实际时长不匹配
effective_duration = UnifiedRenderService._clip_effective_duration(clip)
# 有效时长 = min(指定时长, 实际时长);若均未设置则跳过
effective_duration = 0.0
if clip.duration > 0:
effective_duration = (
min(clip.duration, clip.actual_duration) if clip.actual_duration > 0 else clip.duration
)
elif clip.actual_duration > 0:
effective_duration = clip.actual_duration
if effective_duration > 0:
filters.append(f"trim=duration={effective_duration}")
@@ -1074,7 +803,14 @@ class UnifiedRenderService:
layer_clip_indices = [all_clips.index(c) for c in layer.clips]
layer_labels = [preprocessed_labels[i] for i in layer_clip_indices]
# 使用 trim 后的有效时长,与 Step 1 的 trim=duration 保持一致
layer_durations = [UnifiedRenderService._clip_effective_duration(all_clips[i]) for i in layer_clip_indices]
layer_durations = []
for i in layer_clip_indices:
c = all_clips[i]
if c.duration > 0:
eff = min(c.duration, c.actual_duration) if c.actual_duration > 0 else c.duration
else:
eff = c.actual_duration if c.actual_duration > 0 else 0.0
layer_durations.append(eff)
layer_transitions = [all_clips[i].transition_effect for i in layer_clip_indices]
if len(layer_labels) == 1:
@@ -1194,7 +930,7 @@ class UnifiedRenderService:
raise
def _probe_output(self, output_path: Path) -> tuple[float, int, int, int]:
"""探测输出文件的时长、大小、宽高.
"""探测输出文件的时长、大小、宽高
Returns:
(duration, file_size, width, height)
@@ -1207,304 +943,3 @@ class UnifiedRenderService:
info["width"],
info["height"],
)
# ── 音频后处理 ────────────────────────────────────────────────────────
def _mix_audio(self, layers: list[RenderLayer], video_duration: float) -> Path | None:
"""音频后处理混音.
处理逻辑
1. 主音频源按优先级查找main > brollbackground 不参与主音频通常是图片无音轨
2. 主图层音频按顺序 concat 拼接
3. 独立音频轨audio role amix 混入
4. 输出时长截断到 video_duration
5. 无音频流的 clip 会被自动跳过避免 FFmpeg 引用 [i:a] 失败
Args:
layers: 图层列表
video_duration: 视频总时长用于截断音频
Returns:
混音后的音频文件路径无音频时返回 None
"""
# 按优先级精确查找主音频图层:main > broll
# background 不参与主音频(通常是静态图片,无音轨)
layer_map = {layer.role: layer for layer in layers}
main_layer = None
for role in ("main", "broll"):
if role in layer_map and layer_map[role].clips:
main_layer = layer_map[role]
break
main_clips: list[ResolvedClip] = main_layer.clips if main_layer else []
# 没有主视频图层时兜底:检查 overlay/corner_voice 层是否有带音频的素材
if not main_clips:
for role in ("overlay", "corner_voice"):
if role in layer_map and layer_map[role].clips:
main_clips = layer_map[role].clips
break
# 收集独立音频轨
audio_clips: list[ResolvedClip] = []
if "audio" in layer_map:
audio_clips = layer_map["audio"].clips
# ── 防御:过滤掉无音频流的 clip ──
# 源视频可能没有音频流(如静音视频、纯图片转的视频),直接引用 [i:a] 会导致 FFmpeg 失败
main_clips = [c for c in main_clips if self._clip_has_audio(c)]
audio_clips = [c for c in audio_clips if self._clip_has_audio(c)]
if not main_clips and not audio_clips:
return None
# 构建音频处理命令
output_path = self.work_dir / f"audio_{self.plan.id}.aac"
# 简单场景:只有主图层 + 无独立音频 → 直接从视频提取音频并拼接
if main_clips and not audio_clips:
self._concat_main_audio(main_clips, output_path, video_duration)
return output_path
# 有独立音频轨 → amix 混音
self._mix_with_independent_audio(main_clips, audio_clips, output_path, video_duration)
return output_path
def _concat_main_audio(self, clips: list[ResolvedClip], output_path: Path, video_duration: float) -> None:
"""主图层音频 concat 拼接(对齐链路A行为).
每个 clip 提取音频 trim 按顺序 concat
"""
if len(clips) == 1:
# 单 clip,直接提取音频,截断到 min(clip有效时长, 视频总时长)
clip = clips[0]
effective_duration = self._clip_effective_duration(clip)
# 最终时长:取 clip 有效时长和视频总时长的较小值
# (视频总时长由主图层决定,但单 clip 场景下两者应该一致,仍做保护)
final_duration = effective_duration
if video_duration > 0 and (final_duration <= 0 or final_duration > video_duration):
final_duration = video_duration
command = [
FFMPEG_BIN,
"-y",
"-i",
str(clip.local_path),
"-vn",
"-acodec",
"aac",
"-b:a",
"128k",
]
if final_duration > 0:
command.extend(["-t", f"{final_duration:.3f}"])
command.append(str(output_path))
run_ffmpeg(command)
return
# 多 clip,用 filter_complex concat
input_args: list[str] = []
filter_parts: list[str] = []
for i, clip in enumerate(clips):
input_args.extend(["-i", str(clip.local_path)])
effective_duration = self._clip_effective_duration(clip)
if effective_duration > 0:
filter_parts.append(f"[{i}:a]atrim=0:{effective_duration:.3f},asetpts=PTS-STARTPTS[a{i}]")
else:
filter_parts.append(f"[{i}:a]asetpts=PTS-STARTPTS[a{i}]")
audio_labels = "".join(f"[a{i}]" for i in range(len(clips)))
filter_parts.append(f"{audio_labels}concat=n={len(clips)}:v=0:a=1[outa]")
# 截断到视频总时长
if video_duration > 0:
filter_parts.append(f"[outa]atrim=0:{video_duration:.3f}[final_audio]")
final_label = "final_audio"
else:
final_label = "outa"
filter_complex = ";".join(filter_parts)
command = [
FFMPEG_BIN,
"-y",
*input_args,
"-filter_complex",
filter_complex,
"-map",
f"[{final_label}]",
"-acodec",
"aac",
"-b:a",
"128k",
str(output_path),
]
run_ffmpeg(command)
def _mix_with_independent_audio(
self,
main_clips: list[ResolvedClip],
audio_clips: list[ResolvedClip],
output_path: Path,
video_duration: float,
) -> None:
"""主音频 + 独立音频轨 amix 混音.
Args:
main_clips: 主视频 clips提取音频后 concat
audio_clips: 独立音频轨 clips
output_path: 输出路径
video_duration: 视频总时长
"""
input_args: list[str] = []
filter_parts: list[str] = []
mix_labels: list[str] = []
input_idx = 0
# 1. 主图层音频 concat
if main_clips:
for clip in main_clips:
input_args.extend(["-i", str(clip.local_path)])
effective_duration = self._clip_effective_duration(clip)
if effective_duration > 0:
filter_parts.append(
f"[{input_idx}:a]atrim=0:{effective_duration:.3f},asetpts=PTS-STARTPTS[ma{input_idx}]"
)
else:
filter_parts.append(f"[{input_idx}:a]asetpts=PTS-STARTPTS[ma{input_idx}]")
input_idx += 1
if len(main_clips) == 1:
mix_labels.append("ma0")
else:
main_labels = "".join(f"[ma{i}]" for i in range(len(main_clips)))
filter_parts.append(f"{main_labels}concat=n={len(main_clips)}:v=0:a=1[main_audio]")
mix_labels.append("main_audio")
# 2. 独立音频轨
for j, clip in enumerate(audio_clips):
input_args.extend(["-i", str(clip.local_path)])
effective_duration = self._clip_effective_duration(clip)
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("asetpts=PTS-STARTPTS")
if volume != 1.0:
filters.append(f"volume={volume}")
filter_parts.append(f"[{input_idx}:a]{','.join(filters)}[{label}]")
mix_labels.append(label)
input_idx += 1
# 3. amix 混音
mix_inputs = "".join(f"[{label}]" for label in mix_labels)
n_inputs = len(mix_labels)
# normalized=0 保持音量,duration=shortest 取最短
filter_parts.append(f"{mix_inputs}amix=inputs={n_inputs}:duration=longest:normalize=0[mixed_audio]")
# 4. 截断到视频时长
if video_duration > 0:
filter_parts.append(f"[mixed_audio]atrim=0:{video_duration:.3f}[final_audio]")
final_label = "final_audio"
else:
final_label = "mixed_audio"
filter_complex = ";".join(filter_parts)
command = [
FFMPEG_BIN,
"-y",
*input_args,
"-filter_complex",
filter_complex,
"-map",
f"[{final_label}]",
"-acodec",
"aac",
"-b:a",
"128k",
str(output_path),
]
logger.info(
"音频混音: plan_id=%s main_clips=%d audio_clips=%d",
self.plan.id,
len(main_clips),
len(audio_clips),
)
try:
run_ffmpeg(command)
except subprocess.CalledProcessError as e:
logger.error(
"音频混音失败: plan_id=%s exit_code=%d\nfilter_complex:\n%s",
self.plan.id,
e.returncode,
filter_complex[:3000],
)
raise
def _merge_audio_video(self, video_path: Path, audio_path: Path, output_path: Path) -> None:
"""将音频合并到视频中(视频流拷贝,音频直接复用).
Args:
video_path: 无声视频路径
audio_path: 音频文件路径
output_path: 输出文件路径
"""
command = [
FFMPEG_BIN,
"-y",
"-i",
str(video_path),
"-i",
str(audio_path),
"-c:v",
"copy",
"-c:a",
"aac",
"-b:a",
"128k",
"-map",
"0:v:0",
"-map",
"1:a:0",
"-shortest",
"-movflags",
"+faststart",
str(output_path),
]
logger.info("合并音视频: plan_id=%s", self.plan.id)
try:
run_ffmpeg(command)
except subprocess.CalledProcessError as e:
logger.error(
"合并音视频失败: plan_id=%s exit_code=%d",
self.plan.id,
e.returncode,
)
raise
@staticmethod
def _clip_effective_duration(clip: ResolvedClip) -> float:
"""计算 clip 的有效时长."""
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
def _clip_has_audio(self, clip: ResolvedClip) -> bool:
"""探测 clip 是否有音频流(带缓存).
避免同一个 clip 被多次 ffprobe 探测
"""
if not hasattr(self, "_audio_cache"):
self._audio_cache: dict[str, bool] = {}
key = str(clip.local_path)
if key not in self._audio_cache:
from .ffmpeg_utils import probe_has_audio
self._audio_cache[key] = probe_has_audio(clip.local_path)
return self._audio_cache[key]
-1
View File
@@ -17,7 +17,6 @@ class WorkerSettings(BaseSettings):
database_pool_recycle: int = 3600
environment: str = "development"
auto_create_schema: bool = False
redis_url: str = "redis://redis:6379/0"
# 渲染引擎选择:legacy=旧VideoComposeServiceunified=新UnifiedRenderService
render_engine: str = "legacy"
+4 -11
View File
@@ -61,12 +61,10 @@ def compose_video(self, job_id: str, **kwargs):
return {"status": "error", "message": "Missing plan_id"}
# 判断使用哪个渲染引擎
# 优先级:Redis Feature Flag(白名单 > 百分比) > 环境变量默认
from video_processing.render_engine_resolver import get_render_engine_resolver
from worker_app.core.config import get_settings as get_worker_settings
resolver = get_render_engine_resolver()
user_id = job.created_by_user_id or None
engine = resolver.get_engine(user_id=user_id)
worker_settings = get_worker_settings()
engine = (worker_settings.render_engine or "legacy").lower()
if engine == "unified":
return _compose_with_unified_engine(self, job_service, job, plan_id, db)
@@ -209,12 +207,7 @@ def _compose_with_unified_engine(task, job_service, job, plan_id: str, db) -> di
}
job_service.complete_job(job_id, result=result_data)
logger.info(
"视频合成完成(unified): job_id=%s plan_id=%s duration=%.2fs",
job_id,
plan_id,
result.duration,
)
logger.info("视频合成完成(unified): job_id=%s plan_id=%s duration=%.2fs", job_id, plan_id, result.duration)
return {"status": "completed", "job_id": job_id, "result": result_data}
+101 -325
View File
@@ -1,18 +1,13 @@
"""剪辑计划渲染任务 — 支持 Feature Flag 灰度.
"""剪辑计划渲染任务 — Phase 8 任务 2.05.
Celery 任务 worker.render_edit_plan:
1. 加载 EditPlan + EditPlanClips
2. 根据 Feature Flag 选择渲染引擎legacy / unified
3. 下载各片段素材 + 渲染
2. 下载各片段素材
3. 使用 UnifiedRenderService 按时间线+图层渲染
4. 上传渲染结果到 OSS
5. 创建 GeneratedVideo 记录 + 查重
6. 更新 EditPlan / EditPlanClip 状态
7. 更新 GenerationTask 进度
渲染引擎灰度
- Feature Flag (render_engine) 控制
- legacy: VideoComposeService + FFmpeg filter_complex
- unified: UnifiedRenderService 图层架构
"""
from __future__ import annotations
@@ -68,268 +63,14 @@ def _get_repos():
# ── Celery Task ───────────────────────────────────────────────────────────────
def _resolve_render_engine(user_id: str) -> str:
"""根据 Feature Flag 决定使用哪个渲染引擎。
Returns:
"legacy" "unified"
"""
try:
from video_processing.render_engine_resolver import get_render_engine_resolver
resolver = get_render_engine_resolver()
return resolver.get_engine(user_id=user_id)
except Exception as exc:
logger.warning("获取渲染引擎配置失败,fallback 到 legacy: %s", exc)
return "legacy"
def _mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, error_msg: str):
"""统一的计划失败标记工具。"""
plan = plan_repo.get(plan_id)
if plan and plan.status.value == "rendering":
plan.mark_failed()
plan_repo.update(plan)
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task and gen_task.status.value != "failed":
gen_task.status = "failed"
gen_task.error_message = error_msg
gen_task.completed_at = datetime.now(timezone.utc)
gen_task_repo.update(gen_task)
def _finalize_render_success(
plan,
plan_repo,
clip_repo,
gen_task_repo,
db,
plan_id: str,
output_url: str,
storage_key: str,
duration: float,
file_size: int,
width: int,
height: int,
rendered_clip_ids: list[str],
failed_clip_ids: list[str],
generation_task_id: str,
output_path: Path,
engine: str,
) -> dict:
"""渲染成功后的统一收尾:查重 + 更新状态 + 返回结果。"""
# 创建 GeneratedVideo 记录 + 查重
project_id = plan.project_id or ""
batch_id = plan.config.get("batch_id", "")
mode = plan.config.get("mode", "edit_plan")
if generation_task_id and project_id:
try:
create_video_record_and_dedup(
generation_task_id=generation_task_id,
project_id=project_id,
batch_id=batch_id,
file_url=output_url or "",
file_size=file_size,
duration=duration,
video_path=str(output_path),
mode=mode,
session=db,
width=width,
height=height,
fps=OUTPUT_FPS,
)
except Exception as dedup_err:
logger.warning("查重失败(不影响渲染结果): %s", dedup_err)
# 更新片段状态为 rendered
for clip_id in rendered_clip_ids:
clip = clip_repo.get(clip_id)
if clip and clip.status.value == "ready":
clip.mark_rendered()
clip_repo.update(clip)
# 更新 EditPlan 状态为 completed
plan.config["rendered_url"] = output_url or ""
plan.config["rendered_storage_key"] = storage_key
plan.mark_completed()
plan_repo.update(plan)
# 更新 GenerationTask 状态为 completed
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task:
gen_task.status = "completed"
gen_task.progress = 100.0
gen_task.result_count = len(rendered_clip_ids)
gen_task.completed_at = datetime.now(timezone.utc)
gen_task_repo.update(gen_task)
logger.info(
"剪辑计划渲染完成: plan_id=%s engine=%s rendered=%d failed=%d duration=%.1fs",
plan_id,
engine,
len(rendered_clip_ids),
len(failed_clip_ids),
duration,
)
return {
"status": "completed",
"plan_id": plan_id,
"rendered_count": len(rendered_clip_ids),
"failed_count": len(failed_clip_ids),
"output_url": output_url,
"duration": duration,
}
def _render_with_unified(
plan,
clips,
asset_path_map: dict[str, Path],
tmpdir_path: Path,
rendered_clip_ids: list[str],
plan_id: str,
generation_task_id: str,
plan_repo,
clip_repo,
gen_task_repo,
db,
) -> dict:
"""统一渲染引擎路径(UnifiedRenderService 图层架构)。"""
render_service = UnifiedRenderService(
plan=plan,
clips=clips,
asset_path_map=asset_path_map,
work_dir=tmpdir_path,
output_width=OUTPUT_WIDTH,
output_height=OUTPUT_HEIGHT,
output_fps=int(OUTPUT_FPS),
)
try:
render_result = render_service.render()
except Exception as render_err:
logger.error("渲染失败(unified): %s%s", plan_id, render_err)
_mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, f"渲染失败: {render_err}")
return {"status": "error", "message": f"渲染失败: {render_err}"}
output_path = render_result.output_path
# 上传到 OSS
storage_key = f"rendered/{plan_id}/output.mp4"
output_url = upload_to_oss(output_path, storage_key)
failed_clip_ids: list[str] = []
return _finalize_render_success(
plan=plan,
plan_repo=plan_repo,
clip_repo=clip_repo,
gen_task_repo=gen_task_repo,
db=db,
plan_id=plan_id,
output_url=output_url or "",
storage_key=storage_key,
duration=render_result.duration,
file_size=render_result.file_size,
width=render_result.width,
height=render_result.height,
rendered_clip_ids=rendered_clip_ids,
failed_clip_ids=failed_clip_ids,
generation_task_id=generation_task_id,
output_path=output_path,
engine="unified",
)
def _render_with_legacy(
plan,
clips,
rendered_clip_ids: list[str],
failed_clip_ids: list[str],
tmpdir_path: Path,
plan_id: str,
generation_task_id: str,
plan_repo,
clip_repo,
gen_task_repo,
db,
) -> dict:
"""旧引擎路径(VideoComposeService + FFmpeg filter_complex)。"""
import os
import subprocess
from apps.api.app.services.video_compose_service import VideoComposeService
compose_svc = VideoComposeService(db)
# 校验合成条件
validation = compose_svc.validate_compose(plan_id)
if not validation.valid:
error_msg = "; ".join(validation.errors)
logger.error("合成校验失败(legacy): %s%s", plan_id, error_msg)
_mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, f"合成校验失败: {error_msg}")
return {"status": "error", "message": error_msg}
# 构建 FFmpeg 命令
output_dir = os.environ.get("VIDEO_OUTPUT_DIR", str(tmpdir_path))
output_path = Path(output_dir) / f"{plan_id}.mp4"
compose_cmd = compose_svc.build_compose_command(plan_id, str(output_path))
logger.info("执行 FFmpeg (legacy): plan_id=%s", plan_id)
try:
subprocess.run(
compose_cmd.command,
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=3600,
)
except subprocess.CalledProcessError as e:
error_msg = f"FFmpeg 执行失败: {e.stderr[:500]}"
logger.error("FFmpeg 执行失败(legacy): %s%s", plan_id, error_msg)
_mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, error_msg)
return {"status": "error", "message": error_msg}
# 获取文件大小
file_size = output_path.stat().st_size if output_path.exists() else 0
duration = compose_cmd.estimated_duration or 0.0
# 上传到 OSS
storage_key = f"rendered/{plan_id}/output.mp4"
output_url = upload_to_oss(output_path, storage_key)
return _finalize_render_success(
plan=plan,
plan_repo=plan_repo,
clip_repo=clip_repo,
gen_task_repo=gen_task_repo,
db=db,
plan_id=plan_id,
output_url=output_url or "",
storage_key=storage_key,
duration=duration,
file_size=file_size,
width=OUTPUT_WIDTH,
height=OUTPUT_HEIGHT,
rendered_clip_ids=rendered_clip_ids,
failed_clip_ids=failed_clip_ids,
generation_task_id=generation_task_id,
output_path=output_path,
engine="legacy",
)
@celery_app.task(name="worker.render_edit_plan", bind=True, max_retries=2)
def render_edit_plan(self, plan_id: str) -> dict:
"""渲染剪辑计划
流程
1. 加载 EditPlan + EditPlanClips
2. 根据 Feature Flag 选择渲染引擎legacy / unified
3. 下载素材 + 渲染
2. 下载各片段素材到临时目录构建 asset_path_map
3. 使用 UnifiedRenderService 按时间线+图层渲染
4. 上传渲染结果到 OSS
5. 创建 GeneratedVideo 记录 + 查重
6. 更新 EditPlan completed, EditPlanClips rendered
@@ -338,7 +79,6 @@ def render_edit_plan(self, plan_id: str) -> dict:
logger.info("开始渲染剪辑计划: plan_id=%s", plan_id)
generation_task_id = ""
engine = "legacy"
for repos in _get_repos():
plan_repo, clip_repo, gen_task_repo, db = repos
@@ -353,12 +93,7 @@ def render_edit_plan(self, plan_id: str) -> dict:
# 获取 generation_task_id(提前读取,确保 except 块可用)
generation_task_id = plan.config.get("generation_task_id", "")
# 2. 选择渲染引擎(Feature Flag 灰度控制
user_id = plan.created_by_user_id or ""
engine = _resolve_render_engine(user_id)
logger.info("剪辑计划渲染引擎: plan_id=%s engine=%s user_id=%s", plan_id, engine, user_id)
# 3. 加载片段列表(按 order 排序)
# 2. 加载片段列表(按 order 排序
clips = clip_repo.list_by_plan(plan_id, skip=0, limit=10000)
if not clips:
logger.warning("剪辑计划没有片段: %s", plan_id)
@@ -381,15 +116,6 @@ def render_edit_plan(self, plan_id: str) -> dict:
rendered_clip_ids: list[str] = []
failed_clip_ids: list[str] = []
# 预先批量查询所有素材的 storage_keyfile_url
from packages.adapters.sqlalchemy_impl.models import AssetModel
clip_asset_ids = [c.asset_id for c in clips if c.asset_id]
asset_storage_map: dict[str, str] = {}
if clip_asset_ids:
assets = db.query(AssetModel).filter(AssetModel.id.in_(clip_asset_ids)).all()
asset_storage_map = {a.id: a.file_url for a in assets if a.file_url}
for clip in clips:
if not clip.asset_id:
# 没有素材的片段跳过,标记为失败
@@ -403,22 +129,10 @@ def render_edit_plan(self, plan_id: str) -> dict:
rendered_clip_ids.append(clip.id)
continue
storage_key = asset_storage_map.get(clip.asset_id)
if not storage_key:
logger.warning(
"片段素材无 storage_key,跳过: clip_id=%s asset_id=%s",
clip.id,
clip.asset_id,
)
clip.mark_failed()
clip_repo.update(clip)
failed_clip_ids.append(clip.id)
continue
# 下载素材
ext = Path(storage_key).suffix or ".mp4"
ext = Path(clip.asset_id).suffix or ".mp4"
local_path = tmpdir_path / f"clip_{clip.order:04d}{ext}"
if download_asset(storage_key, local_path):
if download_asset(clip.asset_id, local_path):
asset_path_map[clip.asset_id] = local_path
rendered_clip_ids.append(clip.id)
else:
@@ -439,38 +153,100 @@ def render_edit_plan(self, plan_id: str) -> dict:
gen_task_repo.update(gen_task)
return {"status": "error", "message": "所有片段素材下载失败"}
# 4. 根据引擎选择渲染方式
if engine == "unified":
result = _render_with_unified(
plan=plan,
clips=clips,
asset_path_map=asset_path_map,
tmpdir_path=tmpdir_path,
rendered_clip_ids=rendered_clip_ids,
plan_id=plan_id,
generation_task_id=generation_task_id,
plan_repo=plan_repo,
clip_repo=clip_repo,
gen_task_repo=gen_task_repo,
db=db,
)
else:
result = _render_with_legacy(
plan=plan,
clips=clips,
rendered_clip_ids=rendered_clip_ids,
failed_clip_ids=failed_clip_ids,
tmpdir_path=tmpdir_path,
plan_id=plan_id,
generation_task_id=generation_task_id,
plan_repo=plan_repo,
clip_repo=clip_repo,
gen_task_repo=gen_task_repo,
db=db,
)
# 4. 使用 UnifiedRenderService 渲染
render_service = UnifiedRenderService(
plan=plan,
clips=clips,
asset_path_map=asset_path_map,
work_dir=tmpdir_path,
output_width=OUTPUT_WIDTH,
output_height=OUTPUT_HEIGHT,
output_fps=int(OUTPUT_FPS),
)
result["engine"] = engine
return result
try:
render_result = render_service.render()
except Exception as render_err:
logger.error("渲染失败: %s%s", plan_id, render_err)
plan.mark_failed()
plan_repo.update(plan)
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task:
gen_task.status = "failed"
gen_task.error_message = f"渲染失败: {render_err}"
gen_task.completed_at = datetime.now(timezone.utc)
gen_task_repo.update(gen_task)
return {"status": "error", "message": f"渲染失败: {render_err}"}
output_path = render_result.output_path
# 5. 上传到 OSS
storage_key = f"rendered/{plan_id}/output.mp4"
output_url = upload_to_oss(output_path, storage_key)
# 6. 创建 GeneratedVideo 记录 + 查重
project_id = plan.project_id or ""
batch_id = plan.config.get("batch_id", "")
mode = plan.config.get("mode", "edit_plan")
if generation_task_id and project_id:
try:
create_video_record_and_dedup(
generation_task_id=generation_task_id,
project_id=project_id,
batch_id=batch_id,
file_url=output_url or "",
file_size=render_result.file_size,
duration=render_result.duration,
video_path=str(output_path),
mode=mode,
session=db,
width=render_result.width,
height=render_result.height,
fps=OUTPUT_FPS,
)
except Exception as dedup_err:
logger.warning("查重失败(不影响渲染结果): %s", dedup_err)
# 7. 更新片段状态为 rendered
for clip_id in rendered_clip_ids:
clip = clip_repo.get(clip_id)
if clip and clip.status.value == "ready":
clip.mark_rendered()
clip_repo.update(clip)
# 8. 更新 EditPlan 状态为 completed
plan.config["rendered_url"] = output_url or ""
plan.config["rendered_storage_key"] = storage_key
plan.mark_completed()
plan_repo.update(plan)
# 9. 更新 GenerationTask 状态为 completed
if generation_task_id:
gen_task = gen_task_repo.get(generation_task_id)
if gen_task:
gen_task.status = "completed"
gen_task.progress = 100.0
gen_task.result_count = len(rendered_clip_ids)
gen_task.completed_at = datetime.now(timezone.utc)
gen_task_repo.update(gen_task)
logger.info(
"剪辑计划渲染完成: plan_id=%s rendered=%d failed=%d duration=%.1fs",
plan_id,
len(rendered_clip_ids),
len(failed_clip_ids),
render_result.duration,
)
return {
"status": "completed",
"plan_id": plan_id,
"rendered_count": len(rendered_clip_ids),
"failed_count": len(failed_clip_ids),
"output_url": output_url,
"duration": render_result.duration,
}
except Exception as exc:
logger.exception("渲染剪辑计划异常: %s", plan_id)
+23 -200
View File
@@ -113,7 +113,6 @@ from video_processing.oss_helpers import (
get_signed_download_url,
upload_to_oss,
)
from video_processing.render_engine_resolver import ENGINE_LEGACY, ENGINE_UNIFIED
from video_processing.unified_render_service import UnifiedRenderService
# ── 虚拟 Plan / Clip(内存中构建,不写数据库) ────────────────────────────────
@@ -163,15 +162,13 @@ def _build_plan_and_clips_from_task(
"""
plan = _VirtualPlan(id=task_id, name=f"Generated-{task_id[:8]}")
# 为每个下载路径生成合成 asset_id,并预探测素材时长
# 为每个下载路径生成合成 asset_id
asset_path_map: dict[str, Path] = {}
path_to_asset_id: dict[Path, str] = {}
path_duration: dict[Path, float] = {}
for i, p in enumerate(downloaded_paths):
asset_id = f"gen_{task_id[:8]}_{i:03d}{p.suffix or '.mp4'}"
asset_path_map[asset_id] = p
path_to_asset_id[p] = asset_id
path_duration[p] = probe_duration(p)
clips: list[_VirtualClip] = []
n = len(downloaded_paths)
@@ -187,7 +184,6 @@ def _build_plan_and_clips_from_task(
clip_type=clip_type,
order=i,
asset_id=path_to_asset_id[p],
duration=path_duration[p],
)
)
elif mode == "voice_over":
@@ -200,7 +196,6 @@ def _build_plan_and_clips_from_task(
clip_type="main",
order=i,
asset_id=path_to_asset_id[p],
duration=path_duration[p],
config={"role": "b_roll"},
)
)
@@ -220,7 +215,6 @@ def _build_plan_and_clips_from_task(
clip_type=clip_type,
order=i,
asset_id=path_to_asset_id[p],
duration=path_duration[p],
)
)
else:
@@ -233,7 +227,6 @@ def _build_plan_and_clips_from_task(
clip_type="main",
order=i,
asset_id=path_to_asset_id[p],
duration=path_duration[p],
)
)
@@ -574,148 +567,6 @@ def _validate_template_exists(template_id: str) -> None:
session.close()
# ── 渲染引擎选择 ─────────────────────────────────────────────────────────────
def _resolve_render_engine(user_id: str) -> str:
"""根据 Feature Flag 决定使用哪个渲染引擎。
Returns:
"legacy" "unified"
"""
try:
from video_processing.render_engine_resolver import get_render_engine_resolver
resolver = get_render_engine_resolver()
return resolver.get_engine(user_id=user_id)
except Exception as exc:
logger.warning("获取渲染引擎配置失败,fallback 到 unified: %s", exc)
return ENGINE_UNIFIED
# ── 旧引擎渲染(FFmpeg filter_complex) ────────────────────────────────────────
def _render_with_legacy_engine(
task_id: str,
virtual_clips: list[_VirtualClip],
asset_path_map: dict[str, Path],
work_dir: Path,
output_path: Path,
) -> tuple[float, int]:
"""旧引擎渲染路径:手动构建 FFmpeg filter_complex 命令。
说明generate_video 任务使用虚拟 clips EditPlan 数据库记录
因此无法直接复用 VideoComposeService这里手动构建等价的 filter_complex
命令与旧引擎行为一致scale crop setpts trim setpts
fps 归一化保持原帧率
支持模式one_take / pip / voice_over / voice_pip
- 所有模式统一走 concat 滤镜与旧引擎多片段逻辑一致
Returns:
(duration_seconds, file_size_bytes)
"""
import subprocess
main_clips = [
c
for c in virtual_clips
if c.clip_type in ("main", "b_roll", "background")
or (c.clip_type == "main" and c.config.get("role") == "b_roll")
]
if not main_clips:
main_clips = virtual_clips[:1]
input_args: list[str] = []
video_filters: list[str] = []
audio_filters: list[str] = []
for i, clip in enumerate(main_clips):
local_path = asset_path_map.get(clip.asset_id)
if not local_path:
continue
input_args.extend(["-i", str(local_path)])
duration = clip.duration or 0.0
# 视频滤镜:scale → crop → setpts → trim → setpts(与旧引擎一致)
vf = (
f"[{i}:v]"
f"scale={OUTPUT_WIDTH}:{OUTPUT_HEIGHT}:force_original_aspect_ratio=increase,"
f"crop={OUTPUT_WIDTH}:{OUTPUT_HEIGHT},"
f"setpts=PTS-STARTPTS,"
f"trim=0:{duration:.3f},"
f"setpts=PTS-STARTPTS"
f"[v{i}]"
)
video_filters.append(vf)
# 音频滤镜:atrim → asetpts
af = f"[{i}:a]atrim=0:{duration:.3f},asetpts=PTS-STARTPTS[a{i}]"
audio_filters.append(af)
n = len(main_clips)
if n == 1:
video_label = "[v0]"
audio_label = "[a0]"
else:
# concat 视频
v_inputs = "".join(f"[v{i}]" for i in range(n))
video_filters.append(f"{v_inputs}concat=n={n}:v=1:a=0[outv]")
# concat 音频
a_inputs = "".join(f"[a{i}]" for i in range(n))
audio_filters.append(f"{a_inputs}concat=n={n}:v=0:a=1[outa]")
video_label = "[outv]"
audio_label = "[outa]"
# 组装 filter_complex
fc_parts = video_filters + audio_filters
filter_complex = ";".join(fc_parts)
command = [
FFMPEG_BIN,
"-y",
*input_args,
"-filter_complex",
filter_complex,
"-map",
video_label,
"-map",
audio_label,
"-c:v",
"libx264",
"-crf",
"23",
"-preset",
"medium",
"-c:a",
"aac",
"-b:a",
"192k",
"-movflags",
"+faststart",
str(output_path),
]
logger.info("[task_id=%s] [渲染] legacy 引擎 FFmpeg 开始: clips=%d", task_id, n)
try:
run_ffmpeg(command)
except subprocess.CalledProcessError as e:
logger.error(
"[task_id=%s] [渲染] legacy 引擎 FFmpeg 失败: %s\nfilter_complex: %s",
task_id,
e,
filter_complex[:500],
)
raise
file_size = output_path.stat().st_size if output_path.exists() else 0
duration = probe_duration(output_path)
return duration, file_size
# ── Celery Task ──────────────────────────────────────────────────────────────
@@ -873,59 +724,31 @@ def generate_video(self, task_id: str) -> dict:
)
_flush_logs(task_id, gen_task)
# 3. 根据 Feature Flag 选择渲染引擎
user_id = getattr(gen_task, "created_by_user_id", "") if gen_task else ""
engine = _resolve_render_engine(user_id) if user_id else ENGINE_UNIFIED
logger.info("[task_id=%s] [渲染] 引擎选择: %s (user_id=%s)", task_id, engine, user_id)
# 使用 UnifiedRenderService 渲染
logger.info("[task_id=%s] [渲染] FFmpeg 渲染开始", task_id)
render_start = time.monotonic()
render_output_path = temp_path / f"rendered-{task_id}.mp4"
if engine == ENGINE_LEGACY:
# 旧引擎:filter_complex + concat(保持原帧率,无 fps 归一化)
render_duration, render_file_size = _render_with_legacy_engine(
task_id=task_id,
virtual_clips=virtual_clips,
asset_path_map=asset_path_map,
work_dir=temp_path,
output_path=render_output_path,
)
render_elapsed = time.monotonic() - render_start
logger.info(
"[task_id=%s] [渲染] legacy 引擎完成: 耗时=%.1fs, 时长=%.2fs",
task_id,
render_elapsed,
render_duration,
)
else:
# 新引擎:UnifiedRenderService 图层架构
logger.info("[task_id=%s] [渲染] unified 引擎 FFmpeg 渲染开始", task_id)
render_service = UnifiedRenderService(
plan=virtual_plan,
clips=virtual_clips,
asset_path_map=asset_path_map,
work_dir=temp_path,
output_width=OUTPUT_WIDTH,
output_height=OUTPUT_HEIGHT,
output_fps=int(OUTPUT_FPS),
)
render_result = render_service.render()
render_output_path = render_result.output_path
render_duration = render_result.duration
render_file_size = render_result.file_size
render_elapsed = time.monotonic() - render_start
logger.info(
"[task_id=%s] [渲染] unified 引擎完成: 耗时=%.1fs",
task_id,
render_elapsed,
)
render_service = UnifiedRenderService(
plan=virtual_plan,
clips=virtual_clips,
asset_path_map=asset_path_map,
work_dir=temp_path,
output_width=OUTPUT_WIDTH,
output_height=OUTPUT_HEIGHT,
output_fps=int(OUTPUT_FPS),
)
render_result = render_service.render()
render_elapsed = time.monotonic() - render_start
logger.info(
"[task_id=%s] [渲染] FFmpeg 渲染完成: 耗时=%.1fs",
task_id,
render_elapsed,
)
if gen_task:
gen_task.append_log(
"渲染",
f"引擎={engine}, 耗时={render_elapsed:.1f}s",
f"FFmpeg 渲染完成, 耗时={render_elapsed:.1f}s",
duration=round(render_elapsed, 2),
engine=engine,
)
_flush_logs(task_id, gen_task)
@@ -933,14 +756,14 @@ def generate_video(self, task_id: str) -> dict:
if audio_path:
final_path = temp_path / f"final-{task_id}.mp4"
try:
_mux_audio_track(render_output_path, audio_path, final_path)
_mux_audio_track(render_result.output_path, audio_path, final_path)
# 混音成功,使用混音后的文件
output_path = final_path
except Exception as mux_err:
logger.warning("[task_id=%s] [混音] 音频混合失败,使用无音频版本: %s", task_id, mux_err)
output_path = render_output_path
output_path = render_result.output_path
else:
output_path = render_output_path
output_path = render_result.output_path
file_size = output_path.stat().st_size
duration = probe_duration(output_path)
+1
View File
@@ -0,0 +1 @@
"""Packages root."""
+1
View File
@@ -0,0 +1 @@
"""Adapters package for external implementations."""
+1 -16
View File
@@ -1,9 +1,3 @@
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
FeatureFlagStore,
InMemoryFeatureFlagStore,
RedisFeatureFlagStore,
)
from packages.adapters.redis.session_store import (
NoopSessionStore,
RedisConfig,
@@ -11,13 +5,4 @@ from packages.adapters.redis.session_store import (
get_session_store,
)
__all__ = [
"FeatureFlagConfig",
"FeatureFlagStore",
"InMemoryFeatureFlagStore",
"NoopSessionStore",
"RedisConfig",
"RedisFeatureFlagStore",
"SessionStore",
"get_session_store",
]
__all__ = ["NoopSessionStore", "RedisConfig", "SessionStore", "get_session_store"]
@@ -1,259 +0,0 @@
"""Feature Flag 存储实现。
支持两种后端
- RedisFeatureFlagStore生产环境使用支持多实例共享热更新
- InMemoryFeatureFlagStore测试/开发环境使用纯内存
支持的 Flag 类型
- 全局开关enabled: bool
- 白名单whitelist: Set[str] user_id 列表
- 百分比切流percentage: 0-100基于标识符哈希取模
判定优先级白名单 > 百分比 > 全局开关
"""
from __future__ import annotations
import hashlib
import json
import logging
import threading
import time
from abc import ABC, abstractmethod
from dataclasses import dataclass, field
from typing import Optional, Set
logger = logging.getLogger(__name__)
# Redis key 前缀
FEATURE_FLAG_REDIS_PREFIX = "feature_flag:"
@dataclass
class FeatureFlagConfig:
"""单个 Feature Flag 的配置。"""
name: str
enabled: bool = False
percentage: int = 0 # 0-100
whitelist: Set[str] = field(default_factory=set)
def to_dict(self) -> dict:
return {
"name": self.name,
"enabled": self.enabled,
"percentage": self.percentage,
"whitelist": sorted(self.whitelist),
}
@classmethod
def from_dict(cls, data: dict) -> "FeatureFlagConfig":
return cls(
name=data["name"],
enabled=bool(data.get("enabled", False)),
percentage=int(data.get("percentage", 0)),
whitelist=set(data.get("whitelist", [])),
)
def is_active(self, identifier: Optional[str] = None) -> bool:
"""判断当前 flag 是否激活。
判定优先级
1. 全局关闭 False
2. 白名单匹配 True
3. 百分比命中 True
4. 其他 False
Args:
identifier: 用于白名单匹配和百分比哈希的标识符 user_id
None 时只看全局开关 + 百分比百分比用随机值
"""
if not self.enabled:
return False
# 白名单:精确匹配
if identifier and identifier in self.whitelist:
return True
# 百分比:0 直接 False100 直接 True
if self.percentage <= 0:
# 没有白名单且百分比为0 → 未启用
return False
if self.percentage >= 100:
return True
# 基于 identifier 做哈希取模,确保同一用户始终落在同一侧
if identifier:
hash_val = int(
hashlib.md5(f"{self.name}:{identifier}".encode("utf-8")).hexdigest(), 16 # nosec B324
) # nosec B324 - 用于哈希取模做百分比切流,非安全用途
return (hash_val % 100) < self.percentage
# 无 identifier 且百分比在 0-100 之间 → 按比例随机(不保证一致性)
import random
return random.randint(0, 99) < self.percentage
class FeatureFlagStore(ABC):
"""Feature Flag 存储抽象接口。"""
@abstractmethod
def get(self, name: str) -> FeatureFlagConfig:
"""获取指定 flag 的配置,不存在则返回默认配置(关闭状态)。"""
...
@abstractmethod
def set(self, config: FeatureFlagConfig) -> None:
"""设置 flag 配置。"""
...
@abstractmethod
def delete(self, name: str) -> bool:
"""删除 flag,返回是否成功删除。"""
...
@abstractmethod
def list_all(self) -> dict[str, FeatureFlagConfig]:
"""列出所有 flag。"""
...
def is_active(self, name: str, identifier: Optional[str] = None) -> bool:
"""便捷方法:判断 flag 是否激活。"""
return self.get(name).is_active(identifier)
class InMemoryFeatureFlagStore(FeatureFlagStore):
"""内存实现,用于测试和本地开发。"""
def __init__(self) -> None:
self._flags: dict[str, FeatureFlagConfig] = {}
self._lock = threading.Lock()
def get(self, name: str) -> FeatureFlagConfig:
with self._lock:
return self._flags.get(name, FeatureFlagConfig(name=name, enabled=False))
def set(self, config: FeatureFlagConfig) -> None:
with self._lock:
self._flags[config.name] = config
def delete(self, name: str) -> bool:
with self._lock:
if name in self._flags:
del self._flags[name]
return True
return False
def list_all(self) -> dict[str, FeatureFlagConfig]:
with self._lock:
return dict(self._flags)
class RedisFeatureFlagStore(FeatureFlagStore):
"""Redis 实现,支持多实例共享配置。
每个 flag 存在一个独立的 Redis hash key
Key: feature_flag:{name}
Fields: enabled, percentage, whitelist(JSON array)
"""
def __init__(self, redis_url: str, key_prefix: str = FEATURE_FLAG_REDIS_PREFIX) -> None:
import redis as redis_lib
self._redis = redis_lib.from_url(redis_url, decode_responses=True)
self._key_prefix = key_prefix
# 本地缓存 + TTL,减少 Redis 调用
self._cache: dict[str, tuple[FeatureFlagConfig, float]] = {}
self._cache_ttl = 5.0 # 秒,默认5秒本地缓存
self._lock = threading.Lock()
def _redis_key(self, name: str) -> str:
return f"{self._key_prefix}{name}"
def _parse_whitelist(self, raw: Optional[str]) -> Set[str]:
if not raw:
return set()
try:
data = json.loads(raw)
return set(data) if isinstance(data, list) else set()
except (json.JSONDecodeError, TypeError):
return set()
def get(self, name: str) -> FeatureFlagConfig:
now = time.time()
# 先查本地缓存
with self._lock:
cached = self._cache.get(name)
if cached and now - cached[1] < self._cache_ttl:
return cached[0]
# 从 Redis 读取
try:
key = self._redis_key(name)
data = self._redis.hgetall(key)
if not data:
config = FeatureFlagConfig(name=name, enabled=False)
else:
config = FeatureFlagConfig(
name=name,
enabled=(data.get("enabled", "0") in ("1", "true", "True")),
percentage=int(data.get("percentage", 0)),
whitelist=self._parse_whitelist(data.get("whitelist")),
)
# 写入本地缓存
with self._lock:
self._cache[name] = (config, now)
return config
except Exception as exc:
logger.warning("Failed to get feature flag %s from Redis: %s", name, exc)
# Redis 不可用时返回默认值(关闭),不影响业务
return FeatureFlagConfig(name=name, enabled=False)
def set(self, config: FeatureFlagConfig) -> None:
key = self._redis_key(config.name)
self._redis.hset(
key,
mapping={
"enabled": "1" if config.enabled else "0",
"percentage": str(config.percentage),
"whitelist": json.dumps(sorted(config.whitelist), ensure_ascii=False),
},
)
# 失效本地缓存
with self._lock:
self._cache.pop(config.name, None)
def delete(self, name: str) -> bool:
key = self._redis_key(name)
result = self._redis.delete(key)
with self._lock:
self._cache.pop(name, None)
return bool(result)
def list_all(self) -> dict[str, FeatureFlagConfig]:
pattern = f"{self._key_prefix}*"
result: dict[str, FeatureFlagConfig] = {}
try:
cursor = 0
while True:
cursor, keys = self._redis.scan(cursor=cursor, match=pattern, count=100)
for key in keys:
name = key[len(self._key_prefix) :]
result[name] = self.get(name)
if cursor == 0:
break
except Exception as exc:
logger.warning("Failed to list feature flags from Redis: %s", exc)
return result
def invalidate_cache(self, name: Optional[str] = None) -> None:
"""手动失效本地缓存。"""
with self._lock:
if name:
self._cache.pop(name, None)
else:
self._cache.clear()
+1 -1
View File
@@ -1,6 +1,6 @@
from datetime import datetime, timezone
from sqlalchemy import JSON, Boolean, Column, DateTime, Float, Integer, String, Text, UniqueConstraint
from sqlalchemy import JSON, Boolean, Column, DateTime, Float, Integer, String, Text, UniqueConstraint, create_engine
from sqlalchemy.orm import declarative_base
Base = declarative_base()
+1
View File
@@ -27,6 +27,7 @@ from .generated_videos import (
from .generation_tasks import (
CreateGenerationTaskCommand,
CreateGenerationTaskUseCase,
GetGenerationTaskUseCase,
)
from .ingest_jobs import SubmitIngestJobCommand, SubmitIngestJobUseCase
from .jobs import (
+2 -1
View File
@@ -12,9 +12,10 @@ JWT 处理器委托层
payload = jwt_handler.verify_access_token(token)
"""
from datetime import datetime, timedelta
from typing import Any, Dict, Optional
from packages.application.auth.jwt_service import JWTConfig, JWTService
from packages.application.auth.jwt_service import JWTConfig, JWTService, TokenType
class JWTHandler:
+4 -4
View File
@@ -208,10 +208,10 @@ def _get_jwt_service():
kw = dict(secret_key=settings.JWT_SECRET_KEY)
if hasattr(settings, "JWT_ALGORITHM"):
kw["algorithm"] = settings.JWT_ALGORITHM
if hasattr(settings, "JWT_ACCESS_TOKEN_EXPIRE_MINUTES"):
kw["access_token_expire_minutes"] = settings.JWT_ACCESS_TOKEN_EXPIRE_MINUTES
if hasattr(settings, "JWT_REFRESH_TOKEN_EXPIRE_DAYS"):
kw["refresh_token_expire_days"] = settings.JWT_REFRESH_TOKEN_EXPIRE_DAYS
if hasattr(settings, "ACCESS_TOKEN_EXPIRE_MINUTES"):
kw["access_token_expire_minutes"] = settings.ACCESS_TOKEN_EXPIRE_MINUTES
if hasattr(settings, "REFRESH_TOKEN_EXPIRE_DAYS"):
kw["refresh_token_expire_days"] = settings.REFRESH_TOKEN_EXPIRE_DAYS
_jwt_service_instance = JWTService(JWTConfig(**kw))
return _jwt_service_instance
+1 -1
View File
@@ -293,7 +293,7 @@ class LogoutUseCase:
try:
if request.logout_all_devices:
# 删除所有设备的 session
self.session_store.delete_all_user_sessions(request.user_id)
count = self.session_store.delete_all_user_sessions(request.user_id)
return True, None
else:
# 删除当前 session
@@ -85,6 +85,8 @@ class PasswordHasher:
True 如果需要重新哈希
"""
try:
hashed_bytes = hashed_password.encode("utf-8")
current_rounds = bcrypt.getsalt(hashed_bytes)
# 提取当前的 cost factor
# bcrypt hash 格式: $2b$rounds$salt+hash
@@ -3,7 +3,7 @@
"""
import secrets
from datetime import datetime, timezone
from datetime import datetime, timedelta, timezone
from typing import Optional
from uuid import uuid4
+1 -1
View File
@@ -3,7 +3,7 @@
"""
from math import ceil
from typing import Generic, List, TypeVar
from typing import Generic, List, Optional, TypeVar
from pydantic import BaseModel, Field
+2
View File
@@ -7,7 +7,9 @@ from __future__ import annotations
import logging
from dataclasses import dataclass, field
from datetime import datetime, timezone
from typing import Any
from uuid import uuid4
from packages.domain.job import Job, JobStatus, JobType
from packages.ports.job_repository import JobRepository
+1
View File
@@ -9,6 +9,7 @@ from typing import List, Optional
from packages.adapters.sqlalchemy_impl.recipe_repository import SQLAlchemyRecipeRepository
from packages.application.recipe.commands import (
CreateRecipeCommand,
RecipeItemCommand,
UpdateRecipeCommand,
)
from packages.domain.recipe import Recipe, RecipeItem
+1
View File
@@ -0,0 +1 @@
"""TTS Job application layer."""
@@ -148,6 +148,7 @@ class TTSStreamingService:
# 并发合成所有分段,按顺序流式推送
queue: asyncio.Queue[tuple[int, Optional[bytes], Optional[str]]] = asyncio.Queue()
completed_count = 0
async def _synthesize_one(idx: int, seg_text: str) -> None:
"""合成单个分段并放入队列。"""
+1 -1
View File
@@ -24,7 +24,7 @@ from packages.application.cosyvoice_service import (
CosyVoiceError,
CosyVoiceService,
)
from packages.application.tts_job.audio_merger import AudioMerger
from packages.application.tts_job.audio_merger import AudioMergeError, AudioMerger
from packages.application.tts_job.text_splitter import split_text
from packages.domain.tts_job import TTSJob, TTSJobStatus
from packages.ports.tts_job_repository import TTSJobRepository
@@ -2,6 +2,7 @@
from __future__ import annotations
import uuid
from typing import List, Optional
from packages.domain.voice_clone_profile import VoiceCloneProfile
+3 -2
View File
@@ -10,7 +10,7 @@
from __future__ import annotations
import logging
from typing import Optional
from typing import Any, Optional
from packages.application.cosyvoice_service import (
CosyVoiceAuthError,
@@ -21,8 +21,9 @@ from packages.application.voice_clone.use_cases import (
CreateVoiceCloneUseCase,
RetryVoiceCloneUseCase,
VoiceCloneNotFoundError,
VoiceCloneNotRetryableError,
)
from packages.domain.voice_clone_profile import VoiceCloneProfile
from packages.domain.voice_clone_profile import VoiceCloneProfile, VoiceCloneStatus
from packages.ports.voice_clone_profile_repository import VoiceCloneProfileRepository
logger = logging.getLogger(__name__)
+1
View File
@@ -13,6 +13,7 @@ else:
pass
from typing import Any
from uuid import uuid4
+1 -1
View File
@@ -10,7 +10,7 @@ from __future__ import annotations
import copy
from enum import Enum
from typing import Optional
from typing import List, Optional
from pydantic import BaseModel, Field
+1 -1
View File
@@ -24,7 +24,7 @@ from __future__ import annotations
import logging
from dataclasses import dataclass, field
from typing import Dict, Optional
from typing import Any, Dict, Optional, Set
logger = logging.getLogger(__name__)
+1 -1
View File
@@ -14,7 +14,7 @@ from __future__ import annotations
import logging
from dataclasses import dataclass, field
from enum import Enum
from typing import Any, Callable, Dict, List, Optional
from typing import Any, Callable, Dict, List, Optional, Set
logger = logging.getLogger(__name__)
+1 -1
View File
@@ -2,7 +2,7 @@
from abc import ABC, abstractmethod
from packages.domain import AssetLibrary
from packages.domain import AssetLibrary, AssetLibraryKind
class AssetLibraryRepository(ABC):
+3 -61
View File
@@ -30,8 +30,7 @@ fi
# ---- Registry 配置 ----
REGISTRY="${REGISTRY:-git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas}"
CACHE_REGISTRY="${CACHE_REGISTRY:-git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas}"
# 主缓存 tag:develop 分支构建时写入,所有分支读取
CACHE_TAG_PRIMARY="${CACHE_TAG:-develop}"
CACHE_TAG="${CACHE_TAG:-release}"
API_IMAGE="xiaoxia-saas-api:$VERSION"
WORKER_IMAGE="xiaoxia-saas-worker:$VERSION"
@@ -46,7 +45,6 @@ REGISTRY_WEB="${REGISTRY}/xiaoxia-saas-web:$VERSION"
USE_CACHE=0
USE_PUSH=0
CACHE_WRITE=0
# 检查 buildx 和 Registry 认证
if docker buildx version >/dev/null 2>&1; then
@@ -56,10 +54,7 @@ if docker buildx version >/dev/null 2>&1; then
docker buildx use default 2>/dev/null || true
fi
# ---- 缓存读写策略(按分支隔离)----
# 默认只读不写,防止 feature 分支污染主缓存
# 只有 develop/main 分支才写回缓存
BRANCH_NAME="${GITHUB_REF_NAME:-${CI_COMMIT_BRANCH:-unknown}}"
echo "=== Building API image ==="
if [ "$USE_CACHE" -eq 1 ]; then
docker buildx build \
--build-arg APP_VERSION="$VERSION" \
@@ -73,52 +68,6 @@ else
docker build --pull=false --build-arg APP_VERSION="$VERSION" -f infra/docker/api.Dockerfile -t "$API_IMAGE" -t "$API_LATEST" .
fi
build_with_cache() {
# usage: build_with_cache <image_name> <dockerfile> <extra_args...>
IMG_NAME="$1"
DOCKERFILE="$2"
shift 2
EXTRA_ARGS="$*"
CACHE_FROM="type=registry,ref=${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY},ignore-error=true"
if [ "$CACHE_WRITE" -eq 1 ]; then
CACHE_TO="type=registry,ref=${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY},mode=max"
echo " cache: read+write from ${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY}"
else
CACHE_TO=""
echo " cache: read-only from ${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY}"
fi
if [ "$USE_CACHE" -eq 1 ]; then
if [ -n "$CACHE_TO" ]; then
docker buildx build \
$EXTRA_ARGS \
--cache-from "$CACHE_FROM" \
--cache-to "$CACHE_TO" \
-f "$DOCKERFILE" \
-t "$IMG_NAME:$VERSION" \
--load \
.
else
docker buildx build \
$EXTRA_ARGS \
--cache-from "$CACHE_FROM" \
-f "$DOCKERFILE" \
-t "$IMG_NAME:$VERSION" \
--load \
.
fi
else
docker build --pull=false $EXTRA_ARGS -f "$DOCKERFILE" -t "$IMG_NAME:$VERSION" .
fi
}
echo "=== Building API image ==="
build_with_cache "api" "infra/docker/api.Dockerfile" \
"--build-arg APP_VERSION=$VERSION"
docker tag "$API_IMAGE" "$API_LATEST"
echo "=== Building Worker image ==="
if [ "$USE_CACHE" -eq 1 ]; then
docker buildx build \
@@ -134,16 +83,9 @@ else
fi
echo "=== Building Web image (with buildx cache) ==="
# 先构建前端产物(使用持久化 npm 缓存卷)
NPM_CACHE_VOLUME="xiaoxia-npm-cache"
if ! docker volume inspect "$NPM_CACHE_VOLUME" >/dev/null 2>&1; then
docker volume create "$NPM_CACHE_VOLUME" >/dev/null
echo " Created npm cache volume: $NPM_CACHE_VOLUME"
fi
# 先构建前端产物
docker run --rm \
-v "$PWD:/workspace" \
-v "$NPM_CACHE_VOLUME:/workspace/apps/web/node_modules" \
-w /workspace/apps/web \
docker.m.daocloud.io/library/node:20 \
sh -lc "npm ci && npm run build"
-1
View File
@@ -247,4 +247,3 @@ def main() -> int:
if __name__ == "__main__":
sys.exit(main())
-256
View File
@@ -1,256 +0,0 @@
#!/bin/sh
# ============================================================
# Production 部署脚本 - Registry 拉取方式
# ============================================================
# 原始位置:原内嵌在 .gitea/workflows/ci-cd.yml 的 deploy-production Job 中
# 以 base64 编码存储在 DEPLOY_B64 变量中,通过 SSH 管道传送到服务器执行
#
# 功能:
# 1. 登录 Gitea Registry
# 2. Pull api / worker / web 三个镜像
# 3. 备份旧前端静态资源(兼容 CDN 缓存,防止 404)
# 4. 检查基础设施容器(PostgreSQL / Redis
# 5. 执行数据库 Migration
# 6. 停止并重新启动三个业务容器
# 7. 健康检查等待就绪
# 8. 清理旧镜像
#
# 依赖的环境变量(由 CI 通过 SSH 传入):
# IMAGE_TAG - 镜像版本标签(如 v0.1.127,对应 git tag
# REGISTRY_TOKEN - Gitea Registry 访问令牌
# ============================================================
set -eu
IMAGE_TAG="${IMAGE_TAG:-}"
REGISTRY="${REGISTRY:-git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas}"
REGISTRY_USER="${REGISTRY_USER:-xiaoxia}"
REGISTRY_TOKEN="${REGISTRY_TOKEN:-}"
ENV_FILE="${ENV_FILE:-/var/lib/xiaoxia-saas-production/.env}"
GENERATED_DIR="${GENERATED_DIR:-/var/lib/xiaoxia-saas-production/generated}"
LEGACY_ASSETS_DIR="${LEGACY_ASSETS_DIR:-/var/lib/xiaoxia-saas-production/legacy-assets}"
if [ -z "$IMAGE_TAG" ]; then
echo "ERROR: IMAGE_TAG is required"
exit 1
fi
test -f "$ENV_FILE"
mkdir -p "$GENERATED_DIR"
mkdir -p "$LEGACY_ASSETS_DIR"
# ---- 登录 Registry ----
if [ -n "$REGISTRY_TOKEN" ]; then
echo "Logging in to registry: $REGISTRY"
REGISTRY_HOST=$(echo "$REGISTRY" | cut -d/ -f1)
printf %s "$REGISTRY_TOKEN" | docker login "$REGISTRY_HOST" -u "$REGISTRY_USER" --password-stdin 2>/dev/null || {
echo "WARN: docker login failed, will try to pull anyway"
}
fi
# ---- Pull 三个镜像 ----
REGISTRY_API="${REGISTRY}/xiaoxia-saas-api:${IMAGE_TAG}"
REGISTRY_WORKER="${REGISTRY}/xiaoxia-saas-worker:${IMAGE_TAG}"
REGISTRY_WEB="${REGISTRY}/xiaoxia-saas-web:${IMAGE_TAG}"
LOCAL_API="xiaoxia-saas-api:${IMAGE_TAG}"
LOCAL_WORKER="xiaoxia-saas-worker:${IMAGE_TAG}"
LOCAL_WEB="xiaoxia-saas-web:${IMAGE_TAG}"
echo "Pulling API image..."
docker pull "$REGISTRY_API"
echo "Pulling Worker image..."
docker pull "$REGISTRY_WORKER"
echo "Pulling Web image..."
docker pull "$REGISTRY_WEB"
# ---- Re-tag 成本地镜像名 ----
docker tag "$REGISTRY_API" "$LOCAL_API"
docker tag "$REGISTRY_WORKER" "$LOCAL_WORKER"
docker tag "$REGISTRY_WEB" "$LOCAL_WEB"
echo "All images pulled and tagged."
# ---- 备份旧前端 assets(生产环境访问量大,防止 CDN 缓存命中 404) ----
echo "Backing up legacy assets from current web container..."
if docker inspect xiaoxia-web-production >/dev/null 2>&1; then
_tmpdir="/tmp/legacy-assets-$$"
rm -rf "$_tmpdir"
mkdir -p "$_tmpdir"
docker cp xiaoxia-web-production:/usr/share/nginx/html/assets/. "$_tmpdir/" 2>/dev/null || true
# 复制到 LEGACY_ASSETS_DIR(下一次部署时作为 fallback 挂载)
if [ -d "$_tmpdir" ] && [ "$(ls -A "$_tmpdir" 2>/dev/null)" ]; then
cp -an "$_tmpdir"/. "$LEGACY_ASSETS_DIR"/ 2>/dev/null || true
echo "Legacy assets backed up: $(ls "$_tmpdir" | wc -l) files"
fi
rm -rf "$_tmpdir"
else
echo "No existing web container, skipping legacy assets backup"
fi
# 清理超过 7 天的旧 assets 文件
if [ -d "$LEGACY_ASSETS_DIR" ]; then
find "$LEGACY_ASSETS_DIR" -type f -mtime +7 -delete 2>/dev/null || true
echo "Legacy assets cleanup done (retain 7 days)"
fi
# ---- 检查基础设施容器状态 ----
echo "Checking infrastructure containers..."
for c in xiaoxia-postgres-production xiaoxia-redis-production; do
if ! docker inspect "$c" >/dev/null 2>&1; then
echo "ERROR: Required container not found: $c"
exit 1
fi
state=$(docker inspect -f '{{.State.Status}}' "$c")
if [ "$state" != "running" ]; then
echo "ERROR: Container not running: $c ($state)"
exit 1
fi
done
# ---- 创建生产网络 ----
docker network create xiaoxia-net-production 2>/dev/null || true
# ---- 执行数据库 Migration ----
echo "Running database migrations..."
docker run --rm \
--env-file "$ENV_FILE" \
--network xiaoxia-net-production \
-e APP_ENV=production \
"$LOCAL_API" sh -c "cd /app && alembic upgrade head"
echo "Migrations completed."
# ---- 停止旧容器 ----
echo "Stopping old containers..."
docker rm -f xiaoxia-api-production 2>/dev/null || true
docker rm -f xiaoxia-worker-production 2>/dev/null || true
docker rm -f xiaoxia-web-production 2>/dev/null || true
# ---- 启动参数(日志配置统一) ----
LOG_OPTS="--log-driver json-file --log-opt max-size=50m --log-opt max-file=3"
# ---- 启动 API 容器 ----
echo "Starting API container..."
docker run -d \
--name xiaoxia-api-production \
--env-file "$ENV_FILE" \
--network xiaoxia-net-production \
-p 127.0.0.1:8001:8000 \
-e APP_ENV=production \
-e APP_VERSION="$IMAGE_TAG" \
-e GENERATED_FILES_DIR=/app/generated \
-e GENERATED_FILES_URL_PREFIX=/generated-files \
-e PUBLIC_API_BASE_URL=https://api.xiaoxiajianji.com \
-v "$GENERATED_DIR:/app/generated" \
--restart unless-stopped \
--cpus 2 \
--memory 2g \
--health-cmd "python3 -c \"import urllib.request; urllib.request.urlopen('http://localhost:8000/health', timeout=5)\"" \
--health-interval 30s \
--health-timeout 10s \
--health-retries 3 \
--health-start-period 40s \
$LOG_OPTS \
"$LOCAL_API"
# ---- 启动 Worker 容器 ----
echo "Starting Worker container..."
docker run -d \
--name xiaoxia-worker-production \
--env-file "$ENV_FILE" \
--network xiaoxia-net-production \
-e APP_ENV=production \
-e APP_VERSION="$IMAGE_TAG" \
-e WORKER_CONCURRENCY=1 \
-e WORKER_MAX_TASKS_PER_CHILD=100 \
-e GENERATED_FILES_DIR=/app/generated \
-e GENERATED_FILES_URL_PREFIX=/generated-files \
-e PUBLIC_API_BASE_URL=https://api.xiaoxiajianji.com \
-v "$GENERATED_DIR:/app/generated" \
--restart unless-stopped \
--cpus 2 \
--memory 2g \
--health-cmd "sh -c \"grep -q celery /proc/1/cmdline || exit 1\"" \
--health-interval 30s \
--health-timeout 10s \
--health-retries 3 \
--health-start-period 30s \
$LOG_OPTS \
"$LOCAL_WORKER"
# ---- 启动 Web 容器 ----
# Legacy assets 挂载到 /usr/share/nginx/html/assets-legacy/assets/
# nginx 配置中 assets location 有 fallback 规则
LEGACY_VOLUME=""
if [ -d "$LEGACY_ASSETS_DIR" ] && [ "$(ls -A "$LEGACY_ASSETS_DIR" 2>/dev/null)" ]; then
LEGACY_VOLUME="-v ${LEGACY_ASSETS_DIR}:/usr/share/nginx/html/assets-legacy/assets:ro"
echo "Web container: legacy assets mounted (fallback)"
else
echo "Web container: no legacy assets to mount"
fi
echo "Starting Web container..."
docker run -d \
--name xiaoxia-web-production \
--network xiaoxia-net-production \
-p 127.0.0.1:3002:80 \
--restart unless-stopped \
--cpus 0.5 \
--memory 512m \
$LEGACY_VOLUME \
--health-cmd "wget --spider -q http://127.0.0.1:80" \
--health-interval 30s \
--health-timeout 5s \
--health-retries 3 \
$LOG_OPTS \
"$LOCAL_WEB"
# ---- 等待 API 就绪 ----
echo "Waiting for API to become healthy..."
i=0
while [ "$i" -lt 40 ]; do
if curl -sf --max-time 5 http://127.0.0.1:8001/health >/dev/null 2>&1; then
echo "API is healthy!"
break
fi
i=$((i + 1))
echo " Waiting... ($i/40)"
sleep 3
done
if [ "$i" -ge 40 ]; then
echo "ERROR: API did not become healthy within 120s"
docker logs --tail 50 xiaoxia-api-production
exit 1
fi
# ---- 等待 Web 就绪 ----
echo "Waiting for Web to become healthy..."
i=0
while [ "$i" -lt 15 ]; do
if curl -sf --max-time 5 http://127.0.0.1:3002/ >/dev/null 2>&1; then
echo "Web is healthy!"
break
fi
i=$((i + 1))
echo " Waiting... ($i/15)"
sleep 2
done
if [ "$i" -ge 15 ]; then
echo "ERROR: Web did not become healthy within 30s"
docker logs --tail 30 xiaoxia-web-production
exit 1
fi
# ---- 清理旧镜像 ----
echo "Cleaning up old images..."
docker image prune -af --filter "until=168h" 2>/dev/null || true
docker builder prune -af --filter "until=168h" 2>/dev/null || true
echo ""
echo "=== Production deployment complete ==="
echo "API: http://127.0.0.1:8001"
echo "Web: http://127.0.0.1:3002"
echo "Version: $IMAGE_TAG"
docker ps --format "table {{.Names}}\t{{.Status}}\t{{.Image}}" | grep production
-219
View File
@@ -1,219 +0,0 @@
#!/bin/sh
# ============================================================
# Staging 部署脚本 - Registry 拉取方式(替代 Watchtower
# ============================================================
# 说明:原 Staging 使用 Watchtower 轮询 :staging 标签自动更新
# 本脚本替代 Watchtower,由 CI 主动触发部署,优势:
# - 部署时机精确可控,无需固定 sleep 等待
# - 部署完成立即健康检查,失败可快速回滚
# - 部署日志完整记录在 CI 中
#
# 功能:
# 1. 登录 Gitea Registry
# 2. Pull api / worker / web 三个镜像(用 commit SHA 作为 tag
# 3. 执行数据库 Migration
# 4. 停止并重新启动三个业务容器
# 5. 健康检查等待就绪
#
# 依赖的环境变量(由 CI 通过 SSH 传入):
# IMAGE_TAG - 镜像标签(一般为 commit SHA)
# REGISTRY_TOKEN - Gitea Registry 访问令牌
# ============================================================
set -eu
IMAGE_TAG="${IMAGE_TAG:-}"
REGISTRY="${REGISTRY:-git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas}"
REGISTRY_USER="${REGISTRY_USER:-xiaoxia}"
REGISTRY_TOKEN="${REGISTRY_TOKEN:-}"
ENV_FILE="${ENV_FILE:-/var/lib/xiaoxia-saas-staging/.env}"
GENERATED_DIR="${GENERATED_DIR:-/var/lib/xiaoxia-saas-staging/generated}"
if [ -z "$IMAGE_TAG" ]; then
echo "ERROR: IMAGE_TAG is required"
exit 1
fi
test -f "$ENV_FILE"
mkdir -p "$GENERATED_DIR"
# ---- 登录 Registry ----
if [ -n "$REGISTRY_TOKEN" ]; then
echo "Logging in to registry: $REGISTRY"
REGISTRY_HOST=$(echo "$REGISTRY" | cut -d/ -f1)
printf %s "$REGISTRY_TOKEN" | docker login "$REGISTRY_HOST" -u "$REGISTRY_USER" --password-stdin 2>/dev/null || {
echo "WARN: docker login failed, will try to pull anyway"
}
fi
# ---- Pull 三个镜像 ----
REGISTRY_API="${REGISTRY}/xiaoxia-saas-api:${IMAGE_TAG}"
REGISTRY_WORKER="${REGISTRY}/xiaoxia-saas-worker:${IMAGE_TAG}"
REGISTRY_WEB="${REGISTRY}/xiaoxia-saas-web:${IMAGE_TAG}"
LOCAL_API="xiaoxia-saas-api:staging-${IMAGE_TAG}"
LOCAL_WORKER="xiaoxia-saas-worker:staging-${IMAGE_TAG}"
LOCAL_WEB="xiaoxia-saas-web:staging-${IMAGE_TAG}"
echo "Pulling API image..."
docker pull "$REGISTRY_API"
echo "Pulling Worker image..."
docker pull "$REGISTRY_WORKER"
echo "Pulling Web image..."
docker pull "$REGISTRY_WEB"
# ---- Re-tag 成本地镜像名 ----
docker tag "$REGISTRY_API" "$LOCAL_API"
docker tag "$REGISTRY_WORKER" "$LOCAL_WORKER"
docker tag "$REGISTRY_WEB" "$LOCAL_WEB"
echo "All images pulled and tagged."
# ---- 检查基础设施容器 ----
echo "Checking infrastructure containers..."
for c in xiaoxia-postgres-staging xiaoxia-redis-staging; do
if ! docker inspect "$c" >/dev/null 2>&1; then
echo "ERROR: Required container not found: $c"
exit 1
fi
state=$(docker inspect -f '{{.State.Status}}' "$c")
if [ "$state" != "running" ]; then
echo "ERROR: Container not running: $c ($state)"
exit 1
fi
done
# ---- 创建 Staging 网络 ----
docker network create xiaoxia-net-staging 2>/dev/null || true
# ---- 执行数据库 Migration ----
echo "Running database migrations..."
docker run --rm \
--env-file "$ENV_FILE" \
--network xiaoxia-net-staging \
-e APP_ENV=staging \
"$LOCAL_API" sh -c "cd /app && alembic upgrade head"
echo "Migrations completed."
# ---- 停止旧容器 ----
echo "Stopping old containers..."
docker rm -f xiaoxia-api-staging 2>/dev/null || true
docker rm -f xiaoxia-worker-staging 2>/dev/null || true
docker rm -f xiaoxia-web-staging 2>/dev/null || true
# ---- 启动参数(日志配置统一) ----
LOG_OPTS="--log-driver json-file --log-opt max-size=50m --log-opt max-file=3"
# ---- 启动 API 容器 ----
echo "Starting API container..."
docker run -d \
--name xiaoxia-api-staging \
--env-file "$ENV_FILE" \
--network xiaoxia-net-staging \
-p 127.0.0.1:8000:8000 \
-e APP_ENV=staging \
-e APP_VERSION="$IMAGE_TAG" \
-e GENERATED_FILES_DIR=/app/generated \
-e GENERATED_FILES_URL_PREFIX=/generated-files \
-e PUBLIC_API_BASE_URL=https://staging-api.xiaoxiajianji.com \
-v "$GENERATED_DIR:/app/generated" \
--restart unless-stopped \
--cpus 1 \
--memory 1g \
--health-cmd "python3 -c \"import urllib.request; urllib.request.urlopen('http://localhost:8000/health', timeout=5)\"" \
--health-interval 30s \
--health-timeout 10s \
--health-retries 3 \
--health-start-period 30s \
$LOG_OPTS \
"$LOCAL_API"
# ---- 启动 Worker 容器 ----
echo "Starting Worker container..."
docker run -d \
--name xiaoxia-worker-staging \
--env-file "$ENV_FILE" \
--network xiaoxia-net-staging \
-e APP_ENV=staging \
-e APP_VERSION="$IMAGE_TAG" \
-e WORKER_CONCURRENCY=1 \
-e WORKER_MAX_TASKS_PER_CHILD=100 \
-e GENERATED_FILES_DIR=/app/generated \
-e GENERATED_FILES_URL_PREFIX=/generated-files \
-e PUBLIC_API_BASE_URL=https://staging-api.xiaoxiajianji.com \
-v "$GENERATED_DIR:/app/generated" \
--restart unless-stopped \
--cpus 1 \
--memory 1g \
--health-cmd "sh -c \"grep -q celery /proc/1/cmdline || exit 1\"" \
--health-interval 30s \
--health-timeout 10s \
--health-retries 3 \
--health-start-period 30s \
$LOG_OPTS \
"$LOCAL_WORKER"
# ---- 启动 Web 容器 ----
echo "Starting Web container..."
docker run -d \
--name xiaoxia-web-staging \
--network xiaoxia-net-staging \
-p 127.0.0.1:3001:80 \
--restart unless-stopped \
--cpus 0.5 \
--memory 256m \
--health-cmd "wget --spider -q http://127.0.0.1:80" \
--health-interval 30s \
--health-timeout 5s \
--health-retries 3 \
$LOG_OPTS \
"$LOCAL_WEB"
# ---- 等待 API 就绪 ----
echo "Waiting for API to become healthy..."
i=0
while [ "$i" -lt 30 ]; do
if curl -sf --max-time 5 http://127.0.0.1:8000/health >/dev/null 2>&1; then
echo "API is healthy!"
break
fi
i=$((i + 1))
echo " Waiting... ($i/30)"
sleep 2
done
if [ "$i" -ge 30 ]; then
echo "ERROR: API did not become healthy within 60s"
docker logs --tail 50 xiaoxia-api-staging
exit 1
fi
# ---- 等待 Web 就绪 ----
echo "Waiting for Web to become healthy..."
i=0
while [ "$i" -lt 15 ]; do
if curl -sf --max-time 5 http://127.0.0.1:3001/ >/dev/null 2>&1; then
echo "Web is healthy!"
break
fi
i=$((i + 1))
echo " Waiting... ($i/15)"
sleep 2
done
if [ "$i" -ge 15 ]; then
echo "ERROR: Web did not become healthy within 30s"
docker logs --tail 30 xiaoxia-web-staging
exit 1
fi
# ---- 清理旧镜像 ----
echo "Cleaning up old images..."
docker image prune -af --filter "until=72h" 2>/dev/null || true
echo ""
echo "=== Staging deployment complete ==="
echo "API: http://127.0.0.1:8000"
echo "Web: http://127.0.0.1:3001"
echo "Version: $IMAGE_TAG"
docker ps --format "table {{.Names}}\t{{.Status}}\t{{.Image}}" | grep staging
-136
View File
@@ -1,136 +0,0 @@
# 灰度对比测试工具
用于统一渲染引擎灰度发布期间的新旧引擎对比验证。
## 能力
- **像素对比**:基于 FFmpeg SSIM + PSNR 双指标,评估视频画质差异
- **音频对比**:基于差值音频 RMS,评估音频波形差异
- **批量对比**10个预设场景覆盖 P0/P1/P2 优先级
- **HTML 报告**:可视化对比结果,包含画质、音频、性能三维度
- **两种切换方式**:支持 engine 参数直传 或 Feature Flag 白名单切换
## 目录结构
```
tests/render_compare/
├── __init__.py # 包导出
├── README.md # 本文档
├── video_diff.py # 视频像素对比(SSIM + PSNR
├── audio_diff.py # 音频对比(差值 RMS)
├── scenarios.py # 预定义对比场景(10个)
└── runner.py # 批量对比执行器 + HTML 报告生成
```
## 快速开始
### 环境要求
- FFmpeg 4.4+(需带 ssim 和 psnr 滤镜)
- Python 3.10+
- httpxAPI 调用)
### 配置环境变量
```bash
export STAGING_API_URL=https://api.staging.example.com
export STAGING_API_KEY=your_api_key
export STAGING_INTERNAL_API_KEY=your_internal_key # 可选,Feature Flag 模式需要
```
### 运行对比
```bash
# 运行所有 P0 场景(最核心的5个)
python -m tests.render_compare.runner --priority P0 --output ./report/
# 运行 P0 + P1 场景
python -m tests.render_compare.runner --priority P1 --output ./report/
# 只跑指定场景
python -m tests.render_compare.runner --scenarios simple_pass_through,subtitle_rendering
# 使用 Feature Flag 方式切换引擎(需要 internal key
python -m tests.render_compare.runner --priority P0 --flag-mode
# 自定义阈值
python -m tests.render_compare.runner --priority P0 --ssim-threshold 0.95 --psnr-threshold 30
```
## 对比场景
| ID | 名称 | 优先级 | 验证点 |
|----|------|--------|--------|
| simple_pass_through | 简单直通 | P0 | 直通优化路径正确性 |
| multi_clip_transition | 多clip转场 | P0 | 转场效果 + concat |
| subtitle_rendering | 字幕渲染 | P0 | ASS字幕渲染 |
| independent_audio_track | 独立音频轨 | P0 | 音频混音(amix) |
| no_audio_video | 无音轨视频 | P0 | 无音轨防御逻辑 |
| picture_in_picture | 画中画 | P1 | overlay 图层 |
| multi_layer_mix | 多图层混合 | P1 | 多图层复杂场景 |
| image_background | 图片背景 | P1 | background 层 + 无音频 |
| long_video_stress | 长视频压力 | P2 | 多clip性能 |
| vertical_portrait | 竖屏9:16 | P2 | scale 策略(铺满裁剪) |
## 验收标准(建议)
### 视频质量
- **平均 SSIM >= 0.90**:通过(有微小差异但视觉可接受)
- **平均 SSIM >= 0.95**:优秀(视觉几乎无差异)
- **平均 PSNR >= 25 dB**:通过
- **分辨率一致 + 时长差 < 0.1s**:通过
### 音频质量
- **相似度 >= 0.85**:通过
- **采样率/声道数一致**:通过
### 性能
- **平均性能差异在 ±10% 以内**:可接受
- **直通场景新引擎更快**(预期 +30%
## API 约定
Runner 默认假设渲染 API 支持以下接口:
### 提交任务
```
POST /api/v1/render/compose
Authorization: Bearer {api_key}
Body: { ...plan_payload, "engine": "legacy" | "unified" }
Response: { "task_id": "xxx" }
```
### 查询状态
```
GET /api/v1/tasks/{task_id}
Response: { "status": "completed", "output_url": "...", "duration_sec": 5.2 }
```
### Feature Flagflag-mode
```
PUT /api/v1/internal/feature-flags/render_engine
X-API-Key: {internal_key}
Body: { "enabled": true, "percentage": 100 }
```
如果你的 API 接口不同,请修改 `StagingAPI` 类中的对应方法。
## 故障排查
### 对比失败定位指南
1. **像素差异大(SSIM < 0.90**
- 检查分辨率是否一致
- 检查帧率是否一致
- 用 `save_diff_frame` 生成差异帧可视化
- 检查转场效果(slideup/slidedown 是新引擎独有)
2. **音频不一致**
- 检查音频编码参数(码率、采样率)
- 检查主音频源优先级(main > broll
- 用 ffprobe 对比两视频音频流参数
3. **渲染失败**
- 检查日志:`[unified-render] render failed`
- 检查素材是否完整下载
- 检查 FFmpeg 命令是否正确
-26
View File
@@ -1,26 +0,0 @@
"""灰度对比测试工具包.
用于新旧渲染引擎的批量对比测试包含
- video_diff: 视频像素对比SSIM + PSNR
- audio_diff: 音频对比差值 RMS
- scenarios: 预定义对比场景
- runner: 批量对比执行器 + HTML 报告
"""
from .audio_diff import AudioDiffResult, compute_audio_diff, extract_audio, probe_duration, probe_has_audio
from .scenarios import SCENARIOS, CompareScenario, get_scenarios_by_priority
from .video_diff import VideoDiffResult, compute_video_diff, save_diff_frame
__all__ = [
"VideoDiffResult",
"compute_video_diff",
"save_diff_frame",
"AudioDiffResult",
"compute_audio_diff",
"extract_audio",
"probe_has_audio",
"probe_duration",
"SCENARIOS",
"CompareScenario",
"get_scenarios_by_priority",
]
-322
View File
@@ -1,322 +0,0 @@
"""音频对比工具 — 基于 FFmpeg 的音频质量对比.
使用以下指标评估两段音频的相似度
1. 波形差异RMS 差值
2. 频谱相似度FFT 分帧比较
3. 时长差异
对比方式
- 直接对两个音频做 `ametadata=select='gt(scene\\,0.3)'` 过于复杂
- 简化方案 `amerge` + `astats` 计算差值音频的 RMS
更精确的方案已实现
- 将两轨音频做差amix=0:weights='1 -1' 实际上用 pan 更简单
- 对差值音频做 astats获取差值的 RMS峰值等指标
"""
from __future__ import annotations
import json
import re
import shutil
import subprocess # nosec B404
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any
FFMPEG_BIN: str = shutil.which("ffmpeg") or "ffmpeg"
FFPROBE_BIN: str = shutil.which("ffprobe") or "ffprobe"
@dataclass
class AudioDiffResult:
"""音频对比结果."""
audio_a: str
audio_b: str
duration_a: float
duration_b: float
duration_diff: float
sample_rate_match: bool
channels_match: bool
diff_rms_db: float # 差值音频的 RMS(dB,越低越相似)
diff_peak_db: float # 差值音频的峰值(dB,越低越相似)
similarity_score: float # 综合相似度评分 [0, 1],1 = 完全一致
passed: bool
def to_dict(self) -> dict[str, Any]:
return asdict(self)
def probe_duration(file_path: str) -> float:
"""探测文件时长(秒),失败返回 0."""
try:
result = subprocess.run( # nosec B603
[
FFPROBE_BIN,
"-v",
"error",
"-show_entries",
"format=duration",
"-of",
"default=noprint_wrappers=1:nokey=1",
str(file_path),
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=10,
)
return round(float(result.stdout.strip()), 3)
except Exception:
return 0.0
def probe_has_audio(file_path: str | Path) -> bool:
"""探测文件是否包含音频流."""
try:
result = subprocess.run( # nosec B603
[
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"a:0",
"-show_entries",
"stream=codec_type",
"-of",
"default=noprint_wrappers=1:nokey=1",
str(file_path),
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=10,
)
return result.stdout.strip() == "audio"
except Exception:
return False # 探测失败保守返回 False,避免误判有音频
def compute_audio_diff(
audio_a: str | Path,
audio_b: str | Path,
*,
similarity_threshold: float = 0.90,
duration_tolerance: float = 0.1,
) -> AudioDiffResult:
"""计算两段音频的差异.
方案 pan 滤镜将两轨相减对差值音频做 astats 分析
Args:
audio_a: 音频A基线
audio_b: 音频B对比
similarity_threshold: 相似度合格阈值
duration_tolerance: 时长容忍度
Returns:
AudioDiffResult 对比结果
"""
dur_a = probe_duration(str(audio_a))
dur_b = probe_duration(str(audio_b))
duration_diff = abs(dur_a - dur_b)
# 获取音频元信息
info_a = _probe_audio_info(str(audio_a))
info_b = _probe_audio_info(str(audio_b))
sample_rate_match = info_a["sample_rate"] == info_b["sample_rate"]
channels_match = info_a["channels"] == info_b["channels"]
# 相减后分析差值
# 取较短时长做对比
min_dur = min(dur_a, dur_b)
if min_dur <= 0:
return AudioDiffResult(
audio_a=str(audio_a),
audio_b=str(audio_b),
duration_a=dur_a,
duration_b=dur_b,
duration_diff=duration_diff,
sample_rate_match=sample_rate_match,
channels_match=channels_match,
diff_rms_db=-999.0,
diff_peak_db=-999.0,
similarity_score=0.0,
passed=False,
)
# 做差值音频:a - b
# 注意:amix 会自动按输入数归一化音量(除以N),
# 所以 a + (-1)*b 经过 amix=inputs=2 后整体音量会减半(-6dB)。
# 加 volume=2 补偿回来,确保差值 RMS 反映真实差异幅度。
command = [
FFMPEG_BIN,
"-i",
str(audio_a),
"-i",
str(audio_b),
"-filter_complex",
# 第2轨反相 → amix混合 → volume=2补偿amix的自动缩放
"[1:a]volume=-1[inv];[0:a][inv]amix=inputs=2:duration=shortest:dropout_transition=0,volume=2[diff]",
"-map",
"[diff]",
"-f",
"null",
"-af",
"astats=metadata=1:reset=0",
"-",
]
try:
result = subprocess.run( # nosec B603
command,
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=120,
)
stderr = result.stderr or ""
except subprocess.CalledProcessError as e:
# 如果音频格式不兼容,返回失败
return AudioDiffResult(
audio_a=str(audio_a),
audio_b=str(audio_b),
duration_a=dur_a,
duration_b=dur_b,
duration_diff=duration_diff,
sample_rate_match=sample_rate_match,
channels_match=channels_match,
diff_rms_db=999.0,
diff_peak_db=999.0,
similarity_score=0.0,
passed=False,
)
diff_rms_db, diff_peak_db = _parse_astats(stderr)
# 相似度评分:基于差值 RMS
# 差值 RMS -60dB → 相似度 ~1.0(几乎无声差)
# 差值 RMS -20dB → 相似度 ~0.5(有明显差异)
# 差值 RMS 0dB → 相似度 ~0.0(完全相反)
if diff_rms_db <= -60:
similarity_score = 1.0
elif diff_rms_db >= 0:
similarity_score = 0.0
else:
# 线性映射:-60dB → 1.0, 0dB → 0.0
similarity_score = max(0.0, min(1.0, 1.0 + diff_rms_db / 60.0))
passed = (
duration_diff <= duration_tolerance
and sample_rate_match
and channels_match
and similarity_score >= similarity_threshold
)
return AudioDiffResult(
audio_a=str(audio_a),
audio_b=str(audio_b),
duration_a=round(dur_a, 3),
duration_b=round(dur_b, 3),
duration_diff=round(duration_diff, 3),
sample_rate_match=sample_rate_match,
channels_match=channels_match,
diff_rms_db=round(diff_rms_db, 2),
diff_peak_db=round(diff_peak_db, 2),
similarity_score=round(similarity_score, 4),
passed=passed,
)
def _probe_audio_info(file_path: str) -> dict[str, int]:
"""探测音频元信息."""
try:
result = subprocess.run( # nosec B603
[
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"a:0",
"-show_entries",
"stream=sample_rate,channels",
"-of",
"json",
file_path,
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=10,
)
info = json.loads(result.stdout)
stream = info.get("streams", [{}])[0]
return {
"sample_rate": int(stream.get("sample_rate", 44100)),
"channels": int(stream.get("channels", 2)),
}
except Exception:
return {"sample_rate": 0, "channels": 0}
def _parse_astats(stderr: str) -> tuple[float, float]:
"""从 astats 输出中解析 RMS 和峰值.
astats 输出格式 stderr :
[Parsed_astats_1 @ 0x...] Channel: 1
[Parsed_astats_1 @ 0x...] ...
[Parsed_astats_1 @ 0x...] Overall
[Parsed_astats_1 @ 0x...] DC offset: 0.000000
[Parsed_astats_1 @ 0x...] Min level: -0.123456
[Parsed_astats_1 @ 0x...] Max level: 0.789012
[Parsed_astats_1 @ 0x...] Peak level dB: -2.01
[Parsed_astats_1 @ 0x...] RMS level dB: -10.56
...
"""
lines = stderr.split("\n")
rms_db = -999.0
peak_db = -999.0
for line in lines:
# 找 Overall 部分的统计(双声道时取整体值)
rms_match = re.search(r"RMS level dB:\s*(-?\d+\.?\d*)", line)
peak_match = re.search(r"Peak level dB:\s*(-?\d+\.?\d*)", line)
if rms_match:
rms_db = float(rms_match.group(1))
if peak_match:
peak_db = float(peak_match.group(1))
return rms_db, peak_db
def extract_audio(video_path: str | Path, output_path: str | Path) -> Path:
"""从视频中提取音频(AAC 格式).
Args:
video_path: 视频文件路径
output_path: 输出音频路径
Returns:
输出音频文件路径
"""
command = [
FFMPEG_BIN,
"-y",
"-i",
str(video_path),
"-vn",
"-acodec",
"aac",
"-b:a",
"128k",
str(output_path),
]
subprocess.run(command, check=True, capture_output=True, timeout=120) # nosec B603
return Path(output_path)
-628
View File
@@ -1,628 +0,0 @@
"""灰度对比测试 Runner — 新旧引擎批量对比 + 报告生成.
使用方法
# 配置环境变量
export STAGING_API_URL=https://api.staging.example.com
export STAGING_API_KEY=your_key
# 运行全部 P0 场景
python -m tests.render_compare.runner --priority P0 --output ./report/
# 只跑指定场景
python -m tests.render_compare.runner --scenario simple_pass_through,subtitle_rendering
对比流程
1. 对每个场景分别提交到 legacy unified 引擎通过 Feature Flag 白名单/百分比控制
- 方式A通过内部 API 临时切换 flag需要 admin key
- 方式B提交任务时指定 engine 参数如果 API 支持
2. 等待任务完成下载输出视频
3. 像素对比SSIM + PSNR+ 音频对比差值RMS
4. 生成 HTML 对比报告
注意默认假设 API 支持 `engine` 参数来指定渲染引擎
如果不支持需要先通过内部 API 切换 Feature Flag然后提交任务
"""
from __future__ import annotations
import argparse
import json
import os
import sys
import time
from dataclasses import dataclass, field
from datetime import datetime
from pathlib import Path
from typing import Any
import httpx
# 确保项目根目录在 path 中
sys.path.insert(0, str(Path(__file__).parent.parent.parent))
from .audio_diff import AudioDiffResult, compute_audio_diff
from .scenarios import SCENARIOS, CompareScenario, get_scenarios_by_priority
from .video_diff import VideoDiffResult, compute_video_diff
@dataclass
class ScenarioResult:
"""单个场景的对比结果."""
scenario: CompareScenario
legacy_task_id: str = ""
unified_task_id: str = ""
legacy_video_path: str = ""
unified_video_path: str = ""
legacy_duration_sec: float = 0.0
unified_duration_sec: float = 0.0
video_diff: VideoDiffResult | None = None
audio_diff: AudioDiffResult | None = None
legacy_success: bool = False
unified_success: bool = False
error: str = ""
@property
def passed(self) -> bool:
if not (self.legacy_success and self.unified_success):
return False
if self.video_diff and not self.video_diff.passed:
return False
if self.audio_diff and not self.audio_diff.passed:
return False
return True
class StagingAPI:
"""Staging 环境 API 客户端."""
def __init__(self, base_url: str, api_key: str, internal_api_key: str = ""):
self.base_url = base_url.rstrip("/")
self.api_key = api_key
self.internal_api_key = internal_api_key
self.client = httpx.Client(timeout=30.0)
def _headers(self, internal: bool = False) -> dict[str, str]:
headers = {"Authorization": f"Bearer {self.api_key}"}
if internal and self.internal_api_key:
headers["X-API-Key"] = self.internal_api_key
return headers
def submit_render_task(self, plan_payload: dict[str, Any], engine: str = "") -> str:
"""提交渲染任务,返回 task_id.
Args:
plan_payload: EditPlan payload
engine: 可选指定引擎"legacy" / "unified"
Returns:
task_id
"""
url = f"{self.base_url}/api/v1/render/compose"
payload = dict(plan_payload)
if engine:
payload["engine"] = engine
resp = self.client.post(url, json=payload, headers=self._headers())
resp.raise_for_status()
data = resp.json()
return data.get("task_id") or data.get("id", "")
def get_task_status(self, task_id: str) -> dict[str, Any]:
"""获取任务状态."""
url = f"{self.base_url}/api/v1/tasks/{task_id}"
resp = self.client.get(url, headers=self._headers())
resp.raise_for_status()
return resp.json()
def wait_for_task(self, task_id: str, timeout: float = 300.0, poll_interval: float = 3.0) -> dict[str, Any]:
"""等待任务完成.
Returns:
最终任务状态
Raises:
TimeoutError: 超时
"""
start = time.time()
while time.time() - start < timeout:
status = self.get_task_status(task_id)
state = status.get("status", "")
if state in ("completed", "success", "done", "failed", "error"):
return status
time.sleep(poll_interval)
raise TimeoutError(f"Task {task_id} timed out after {timeout}s")
def set_feature_flag(self, flag_name: str, enabled: bool, percentage: int = 0, whitelist: list[str] | None = None):
"""通过内部 API 设置 Feature Flag.
用于不支持 engine 参数的场景切换全局灰度比例
"""
if not self.internal_api_key:
raise ValueError("internal_api_key is required for feature flag operations")
url = f"{self.base_url}/api/v1/internal/feature-flags/{flag_name}"
body: dict[str, Any] = {"enabled": enabled, "percentage": percentage}
if whitelist is not None:
body["whitelist"] = whitelist
resp = self.client.put(url, json=body, headers=self._headers(internal=True))
resp.raise_for_status()
return resp.json()
def get_feature_flag(self, flag_name: str) -> dict[str, Any]:
"""获取 Feature Flag 配置."""
if not self.internal_api_key:
raise ValueError("internal_api_key is required")
url = f"{self.base_url}/api/v1/internal/feature-flags/{flag_name}"
resp = self.client.get(url, headers=self._headers(internal=True))
resp.raise_for_status()
return resp.json()
def download_video(self, video_url: str, output_path: str | Path) -> Path:
"""下载视频文件."""
output_path = Path(output_path)
output_path.parent.mkdir(parents=True, exist_ok=True)
with self.client.stream("GET", video_url, timeout=60.0) as resp:
resp.raise_for_status()
with open(output_path, "wb") as f:
for chunk in resp.iter_bytes():
f.write(chunk)
return output_path
class CompareRunner:
"""新旧引擎对比 Runner."""
# 全局默认阈值(唯一真实来源,所有入口统一引用)
DEFAULT_SSIM_THRESHOLD: float = 0.95
DEFAULT_PSNR_THRESHOLD: float = 28.0
DEFAULT_AUDIO_SIMILARITY_THRESHOLD: float = 0.90
DEFAULT_DURATION_TOLERANCE: float = 0.1
DEFAULT_TASK_TIMEOUT: float = 300.0
def __init__(
self,
api: StagingAPI,
output_dir: Path,
*,
ssim_threshold: float | None = None,
psnr_threshold: float | None = None,
audio_similarity_threshold: float | None = None,
task_timeout: float | None = None,
flag_mode: bool = False, # 是否使用 Feature Flag 方式切换引擎
duration_tolerance: float | None = None,
):
self.api = api
self.output_dir = output_dir
self.ssim_threshold = ssim_threshold if ssim_threshold is not None else self.DEFAULT_SSIM_THRESHOLD
self.psnr_threshold = psnr_threshold if psnr_threshold is not None else self.DEFAULT_PSNR_THRESHOLD
self.audio_similarity_threshold = (
audio_similarity_threshold
if audio_similarity_threshold is not None
else self.DEFAULT_AUDIO_SIMILARITY_THRESHOLD
)
self.duration_tolerance = (
duration_tolerance if duration_tolerance is not None else self.DEFAULT_DURATION_TOLERANCE
)
self.task_timeout = task_timeout if task_timeout is not None else self.DEFAULT_TASK_TIMEOUT
self.flag_mode = flag_mode
self.results: list[ScenarioResult] = []
# flag_mode 下保存原始配置,测试结束后恢复(防污染线上)
self._original_flag_config: dict[str, Any] | None = None
def run_scenario(self, scenario: CompareScenario) -> ScenarioResult:
"""运行单个场景对比."""
print(f"\n{'='*60}")
print(f"[{scenario.priority}] {scenario.id}: {scenario.name}")
print(f" {scenario.description}")
result = ScenarioResult(scenario=scenario)
scenario_dir = self.output_dir / scenario.id
scenario_dir.mkdir(parents=True, exist_ok=True)
try:
# 1. 提交两个引擎的任务
legacy_task_id = self._submit_with_engine(scenario, "legacy")
unified_task_id = self._submit_with_engine(scenario, "unified")
result.legacy_task_id = legacy_task_id
result.unified_task_id = unified_task_id
print(f" legacy task: {legacy_task_id}")
print(f" unified task: {unified_task_id}")
# 2. 等待完成
print(" waiting for legacy...", end="", flush=True)
legacy_status = self.api.wait_for_task(legacy_task_id, timeout=self.task_timeout)
result.legacy_success = legacy_status.get("status") in ("completed", "success", "done")
legacy_video_url = legacy_status.get("output_url", "") or legacy_status.get("video_url", "")
print(f" {'' if result.legacy_success else ''} ({legacy_status.get('duration_sec', '?')}s)")
print(" waiting for unified...", end="", flush=True)
unified_status = self.api.wait_for_task(unified_task_id, timeout=self.task_timeout)
result.unified_success = unified_status.get("status") in ("completed", "success", "done")
unified_video_url = unified_status.get("output_url", "") or unified_status.get("video_url", "")
print(f" {'' if result.unified_success else ''} ({unified_status.get('duration_sec', '?')}s)")
result.legacy_duration_sec = float(legacy_status.get("duration_sec", 0))
result.unified_duration_sec = float(unified_status.get("duration_sec", 0))
if not (result.legacy_success and result.unified_success):
result.error = f"Legacy success={result.legacy_success}, Unified success={result.unified_success}"
print(" ⚠️ 任务未全部成功,跳过对比")
return result
# 3. 下载视频
print(" downloading...", end="", flush=True)
legacy_path = self.api.download_video(legacy_video_url, scenario_dir / "legacy.mp4")
unified_path = self.api.download_video(unified_video_url, scenario_dir / "unified.mp4")
result.legacy_video_path = str(legacy_path)
result.unified_video_path = str(unified_path)
print("")
# 4. 像素对比
print(" computing video diff...", end="", flush=True)
result.video_diff = compute_video_diff(
legacy_path,
unified_path,
ssim_threshold=self.ssim_threshold,
psnr_threshold=self.psnr_threshold,
duration_tolerance=self.duration_tolerance,
)
print(
f" SSIM={result.video_diff.avg_ssim:.4f} PSNR={result.video_diff.avg_psnr:.2f}dB {'' if result.video_diff.passed else ''}"
)
# 5. 音频对比(仅当都有音频时)
from .audio_diff import probe_has_audio
legacy_has_audio = probe_has_audio(legacy_path)
unified_has_audio = probe_has_audio(unified_path)
if legacy_has_audio and unified_has_audio:
print(" computing audio diff...", end="", flush=True)
result.audio_diff = compute_audio_diff(
legacy_path,
unified_path,
similarity_threshold=self.audio_similarity_threshold,
)
print(
f" similarity={result.audio_diff.similarity_score:.4f} {'' if result.audio_diff.passed else ''}"
)
elif legacy_has_audio != unified_has_audio:
result.error = f"音频不一致: legacy_has_audio={legacy_has_audio}, unified_has_audio={unified_has_audio}"
print(f" ⚠️ 音频不一致: legacy={legacy_has_audio}, unified={unified_has_audio}")
else:
print(" audio: both silent (skip)")
except Exception as e:
result.error = str(e)
print(f" ❌ 错误: {e}")
self.results.append(result)
return result
def _submit_with_engine(self, scenario: CompareScenario, engine: str) -> str:
"""提交指定引擎的任务.
如果 flag_mode=True通过 Feature Flag 切换否则通过 engine 参数
"""
if self.flag_mode:
# 先设置 flag(用白名单方式,确保只有当前测试用户命中)
percentage = 0 if engine == "legacy" else 100
self.api.set_feature_flag("render_engine", enabled=True, percentage=percentage)
time.sleep(1) # 给 worker 一点时间刷新配置
return self.api.submit_render_task(scenario.plan_payload)
else:
return self.api.submit_render_task(scenario.plan_payload, engine=engine)
def run_all(self, scenarios: list[CompareScenario]) -> list[ScenarioResult]:
"""运行所有场景.
flag_mode=True 测试开始前保存原始 Feature Flag 配置
结束后无论成功失败自动恢复避免污染线上环境
"""
print(f"\n灰度对比测试开始 - {len(scenarios)} 个场景")
print(f"输出目录: {self.output_dir}")
print(f"视频阈值: SSIM>={self.ssim_threshold}, PSNR>={self.psnr_threshold}dB")
print(f"音频阈值: similarity>={self.audio_similarity_threshold}")
# flag_mode:保存原始配置,测试结束后恢复(防污染)
if self.flag_mode:
try:
self._original_flag_config = self.api.get_feature_flag("render_engine")
print(f" [flag_mode] 已保存原始配置: {self._original_flag_config}")
except Exception as e:
print(f" ⚠️ [flag_mode] 保存原始配置失败: {e}")
print(" 为避免污染线上,将中止测试。请检查 internal_api_key 配置。")
return self.results
try:
for i, scenario in enumerate(scenarios):
print(f"\n进度: {i+1}/{len(scenarios)}")
self.run_scenario(scenario)
finally:
# 始终恢复原始 flag 配置
if self.flag_mode and self._original_flag_config:
try:
orig = self._original_flag_config
self.api.set_feature_flag(
"render_engine",
enabled=orig.get("enabled", False),
percentage=orig.get("percentage", 0),
whitelist=orig.get("whitelist"),
)
print("\n[flag_mode] ✅ 已恢复原始 Feature Flag 配置")
except Exception as e:
print(f"\n[flag_mode] ❌ 恢复 Feature Flag 失败: {e}")
print(" 请手动检查并恢复 render_engine flag 配置!")
return self.results
def summary(self) -> dict[str, Any]:
"""生成汇总统计."""
total = len(self.results)
passed = sum(1 for r in self.results if r.passed)
failed = total - passed
# 性能对比
perf_diffs = []
for r in self.results:
if r.legacy_success and r.unified_success and r.legacy_duration_sec > 0:
diff_pct = (r.unified_duration_sec - r.legacy_duration_sec) / r.legacy_duration_sec * 100
perf_diffs.append(diff_pct)
avg_perf_diff = sum(perf_diffs) / len(perf_diffs) if perf_diffs else 0.0
return {
"total": total,
"passed": passed,
"failed": failed,
"pass_rate": f"{passed/total*100:.1f}%" if total > 0 else "0%",
"avg_perf_diff_pct": round(avg_perf_diff, 2),
"scenarios": [self._result_to_dict(r) for r in self.results],
"timestamp": datetime.now().isoformat(),
"ssim_threshold": self.ssim_threshold,
"psnr_threshold": self.psnr_threshold,
"audio_threshold": self.audio_similarity_threshold,
}
def _result_to_dict(self, r: ScenarioResult) -> dict[str, Any]:
return {
"id": r.scenario.id,
"name": r.scenario.name,
"priority": r.scenario.priority,
"passed": r.passed,
"legacy_success": r.legacy_success,
"unified_success": r.unified_success,
"legacy_duration_sec": r.legacy_duration_sec,
"unified_duration_sec": r.unified_duration_sec,
"video_diff": r.video_diff.to_dict() if r.video_diff else None,
"audio_diff": r.audio_diff.to_dict() if r.audio_diff else None,
"error": r.error,
}
def generate_html_report(summary: dict[str, Any], output_path: Path):
"""生成 HTML 对比报告."""
scenarios = summary["scenarios"]
# 按通过/失败分组
passed_list = [s for s in scenarios if s["passed"]]
failed_list = [s for s in scenarios if not s["passed"]]
# 构建场景卡片
scenario_cards = ""
for s in scenarios:
status_class = "pass" if s["passed"] else "fail"
status_text = "✅ 通过" if s["passed"] else "❌ 失败"
vdiff = s.get("video_diff") or {}
adiff = s.get("audio_diff") or {}
video_info = ""
if vdiff:
video_info = f"""
<div class="metric-row">
<span>SSIM:</span>
<span class="{'good' if vdiff.get('avg_ssim', 0) >= 0.95 else 'warn'}">{vdiff.get('avg_ssim', 0):.4f}</span>
</div>
<div class="metric-row">
<span>PSNR:</span>
<span>{vdiff.get('avg_psnr', 0):.2f} dB</span>
</div>
<div class="metric-row">
<span>时长差:</span>
<span>{vdiff.get('duration_diff', 0):.3f}s</span>
</div>
"""
audio_info = ""
if adiff:
audio_info = f"""
<div class="metric-row">
<span>音频相似度:</span>
<span class="{'good' if adiff.get('similarity_score', 0) >= 0.9 else 'warn'}">{adiff.get('similarity_score', 0):.4f}</span>
</div>
<div class="metric-row">
<span>差值 RMS:</span>
<span>{adiff.get('diff_rms_db', 0):.2f} dB</span>
</div>
"""
perf_info = ""
if s["legacy_duration_sec"] and s["unified_duration_sec"]:
diff = s["unified_duration_sec"] - s["legacy_duration_sec"]
pct = diff / s["legacy_duration_sec"] * 100 if s["legacy_duration_sec"] else 0
trend = "🔴" if pct > 10 else ("🟡" if pct > 0 else "🟢")
perf_info = f"""
<div class="perf-row">
<span>Legacy: {s['legacy_duration_sec']:.2f}s</span>
<span>Unified: {s['unified_duration_sec']:.2f}s</span>
<span>{trend} {pct:+.1f}%</span>
</div>
"""
error_info = f'<div class="error-box">{s["error"]}</div>' if s["error"] else ""
scenario_cards += f"""
<div class="card {status_class}">
<div class="card-header">
<span class="badge">{s['priority']}</span>
<span class="scenario-name">{s['name']}</span>
<span class="status {status_class}">{status_text}</span>
</div>
<div class="card-body">
<div class="grid-2">
<div>
<h4>视频质量</h4>
{video_info or '<p class="muted">无数据</p>'}
</div>
<div>
<h4>音频质量</h4>
{audio_info or '<p class="muted">无音频或跳过</p>'}
</div>
</div>
<div>
<h4>性能对比</h4>
{perf_info or '<p class="muted">无数据</p>'}
</div>
{error_info}
</div>
</div>
"""
html = f"""<!DOCTYPE html>
<html lang="zh-CN">
<head>
<meta charset="UTF-8">
<meta name="viewport" content="width=device-width, initial-scale=1.0">
<title>统一渲染引擎灰度对比报告</title>
<style>
* {{ box-sizing: border-box; margin: 0; padding: 0; }}
body {{ font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, sans-serif; background: #f5f5f5; color: #333; padding: 20px; }}
.container {{ max-width: 1200px; margin: 0 auto; }}
h1 {{ margin-bottom: 20px; font-size: 24px; }}
.summary {{ background: white; border-radius: 12px; padding: 24px; margin-bottom: 24px; display: flex; gap: 32px; flex-wrap: wrap; }}
.summary-item {{ text-align: center; }}
.summary-item .value {{ font-size: 32px; font-weight: bold; margin-bottom: 4px; }}
.summary-item .label {{ color: #666; font-size: 14px; }}
.pass .value {{ color: #10b981; }}
.fail .value {{ color: #ef4444; }}
.card {{ background: white; border-radius: 12px; margin-bottom: 16px; overflow: hidden; border-left: 4px solid #10b981; }}
.card.fail {{ border-left-color: #ef4444; }}
.card-header {{ padding: 16px 20px; background: #fafafa; display: flex; align-items: center; gap: 12px; border-bottom: 1px solid #eee; }}
.badge {{ background: #e5e7eb; color: #374151; padding: 2px 8px; border-radius: 4px; font-size: 12px; font-weight: 600; }}
.scenario-name {{ flex: 1; font-weight: 600; }}
.status {{ font-weight: 600; }}
.status.pass {{ color: #10b981; }}
.status.fail {{ color: #ef4444; }}
.card-body {{ padding: 20px; }}
.grid-2 {{ display: grid; grid-template-columns: 1fr 1fr; gap: 24px; margin-bottom: 16px; }}
h4 {{ margin-bottom: 12px; color: #374151; font-size: 14px; }}
.metric-row {{ display: flex; justify-content: space-between; padding: 6px 0; font-size: 14px; }}
.metric-row .good {{ color: #10b981; font-weight: 600; }}
.metric-row .warn {{ color: #f59e0b; font-weight: 600; }}
.perf-row {{ display: flex; gap: 24px; padding: 8px 0; font-size: 14px; background: #f9fafb; padding: 12px; border-radius: 8px; }}
.error-box {{ background: #fef2f2; color: #dc2626; padding: 12px; border-radius: 8px; margin-top: 12px; font-size: 13px; }}
.muted {{ color: #9ca3af; font-size: 14px; }}
.timestamp {{ text-align: center; color: #9ca3af; font-size: 12px; margin-top: 24px; }}
</style>
</head>
<body>
<div class="container">
<h1>🎬 统一渲染引擎灰度对比报告</h1>
<div class="summary">
<div class="summary-item">
<div class="value">{summary['total']}</div>
<div class="label">总场景数</div>
</div>
<div class="summary-item pass">
<div class="value">{summary['passed']}</div>
<div class="label">通过</div>
</div>
<div class="summary-item fail">
<div class="value">{summary['failed']}</div>
<div class="label">失败</div>
</div>
<div class="summary-item">
<div class="value">{summary['pass_rate']}</div>
<div class="label">通过率</div>
</div>
<div class="summary-item">
<div class="value {'good' if summary['avg_perf_diff_pct'] <= 0 else 'warn'}" style="font-size: 24px; color: {'#10b981' if summary['avg_perf_diff_pct'] <= 0 else '#f59e0b'}">{summary['avg_perf_diff_pct']:+.1f}%</div>
<div class="label">平均性能差异</div>
</div>
</div>
{scenario_cards}
<div class="timestamp">生成时间: {summary['timestamp']}</div>
</div>
</body>
</html>"""
output_path.parent.mkdir(parents=True, exist_ok=True)
output_path.write_text(html, encoding="utf-8")
return output_path
def main():
parser = argparse.ArgumentParser(description="统一渲染引擎灰度对比测试")
parser.add_argument("--priority", default="P0", choices=["P0", "P1", "P2"], help="最低优先级")
parser.add_argument("--scenarios", default="", help="指定场景ID,逗号分隔")
parser.add_argument("--output", default="./gray_compare_report", help="输出目录")
parser.add_argument("--ssim-threshold", type=float, default=None, help="SSIM阈值(默认0.95")
parser.add_argument("--psnr-threshold", type=float, default=None, help="PSNR阈值(dB)(默认28.0")
parser.add_argument("--audio-threshold", type=float, default=None, help="音频相似度阈值(默认0.90")
parser.add_argument("--flag-mode", action="store_true", help="使用Feature Flag方式切换引擎")
parser.add_argument("--task-timeout", type=float, default=300.0, help="单任务超时时间(秒)")
args = parser.parse_args()
base_url = os.environ.get("STAGING_API_URL", "")
api_key = os.environ.get("STAGING_API_KEY", "")
internal_key = os.environ.get("STAGING_INTERNAL_API_KEY", "")
if not base_url or not api_key:
print("❌ 请设置环境变量 STAGING_API_URL 和 STAGING_API_KEY")
sys.exit(1)
# 选择场景
if args.scenarios:
scenario_ids = [s.strip() for s in args.scenarios.split(",")]
selected = [s for s in SCENARIOS if s.id in scenario_ids]
if not selected:
print(f"❌ 未找到匹配的场景: {scenario_ids}")
print(f"可用场景: {[s.id for s in SCENARIOS]}")
sys.exit(1)
else:
selected = get_scenarios_by_priority(args.priority)
output_dir = Path(args.output).resolve()
output_dir.mkdir(parents=True, exist_ok=True)
api = StagingAPI(base_url, api_key, internal_key)
runner = CompareRunner(
api,
output_dir,
ssim_threshold=args.ssim_threshold,
psnr_threshold=args.psnr_threshold,
audio_similarity_threshold=args.audio_threshold,
flag_mode=args.flag_mode,
task_timeout=args.task_timeout,
)
runner.run_all(selected)
# 生成报告
summary = runner.summary()
# JSON 报告
json_path = output_dir / "report.json"
json_path.write_text(json.dumps(summary, indent=2, ensure_ascii=False), encoding="utf-8")
# HTML 报告
html_path = output_dir / "report.html"
generate_html_report(summary, html_path)
print(f"\n{'='*60}")
print(f"对比完成: {summary['passed']}/{summary['total']} 通过 ({summary['pass_rate']})")
print(f"报告: {html_path}")
print(f"JSON: {json_path}")
if __name__ == "__main__":
main()
-257
View File
@@ -1,257 +0,0 @@
"""灰度对比测试场景定义 — 覆盖典型渲染场景.
每个场景对应一个 EditPlan用于新旧引擎对比
覆盖场景
1. 简单直通单clip无特效
2. 多clip转场fade + slide
3. 画中画main + overlay
4. 字幕渲染ASS字幕
5. 独立音频轨主视频 + BGM
6. 多图层混合main + broll + overlay + audio
7. 背景图片 + 主视频图片背景无音频
8. 无音频视频纯画面验证无音轨防御
9. 长视频10+ clip压力测试
10. 分辨率非标竖屏9:16验证scale策略
"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Any
@dataclass
class CompareScenario:
"""对比测试场景."""
id: str
name: str
description: str
priority: str # P0 / P1 / P2
plan_payload: dict[str, Any] # EditPlan JSON payload(提交给 API 的数据)
expected: dict[str, Any] = field(default_factory=dict) # 预期结果
SCENARIOS: list[CompareScenario] = [
CompareScenario(
id="simple_pass_through",
name="简单直通",
description="单主clip,无转场无特效,验证直通优化路径",
priority="P0",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 5.0,
"order": 0,
}
],
},
),
CompareScenario(
id="multi_clip_transition",
name="多clip转场",
description="3个clipfade + slideleft 转场",
priority="P0",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 3.0,
"order": 0,
"transition_effect": "cut",
},
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 3.0,
"order": 1,
"transition_effect": "fade",
},
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 3.0,
"order": 2,
"transition_effect": "slideleft",
},
],
},
),
CompareScenario(
id="picture_in_picture",
name="画中画",
description="主视频 + 角落小窗(corner_voice",
priority="P1",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{"clip_type": "main", "asset_id": "sample_5s.mp4", "duration": 5.0, "order": 0},
{"clip_type": "corner_voice", "asset_id": "sample_5s.mp4", "duration": 5.0, "order": 0},
],
},
),
CompareScenario(
id="subtitle_rendering",
name="字幕渲染",
description="主视频 + ASS字幕",
priority="P0",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 5.0,
"order": 0,
"config": {"subtitles": [{"text": "测试字幕 Test Subtitle", "start_time": 0, "end_time": 5.0}]},
}
],
},
),
CompareScenario(
id="independent_audio_track",
name="独立音频轨",
description="主视频(带音频)+ 独立BGM轨,验证音频混音",
priority="P0",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{"clip_type": "main", "asset_id": "sample_5s.mp4", "duration": 5.0, "order": 0},
{
"clip_type": "main",
"asset_id": "sample_bgm.mp3",
"duration": 5.0,
"order": 0,
"config": {"role": "audio", "volume": 0.5},
},
],
},
),
CompareScenario(
id="multi_layer_mix",
name="多图层混合",
description="main + broll + overlay + audio 四图层",
priority="P1",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 4.0,
"order": 0,
"transition_effect": "fade",
},
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 4.0,
"order": 1,
"transition_effect": "slideup",
},
{"clip_type": "broll", "asset_id": "sample_broll.mp4", "duration": 8.0, "order": 0},
{"clip_type": "overlay", "asset_id": "sample_overlay.png", "duration": 8.0, "order": 0},
{
"clip_type": "main",
"asset_id": "sample_bgm.mp3",
"duration": 8.0,
"order": 0,
"config": {"role": "audio", "volume": 0.3},
},
],
},
),
CompareScenario(
id="image_background",
name="图片背景",
description="background图片层 + 主视频,验证背景层无音频",
priority="P1",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{"clip_type": "background", "asset_id": "sample_bg.jpg", "duration": 5.0, "order": 0},
{"clip_type": "main", "asset_id": "sample_5s.mp4", "duration": 5.0, "order": 0},
],
},
),
CompareScenario(
id="no_audio_video",
name="无音轨视频",
description="源视频无音频流,验证无音轨防御逻辑",
priority="P0",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{"clip_type": "main", "asset_id": "sample_silent_5s.mp4", "duration": 5.0, "order": 0},
],
},
),
CompareScenario(
id="long_video_stress",
name="长视频压力",
description="10个clip + 多种转场,性能压力测试",
priority="P2",
plan_payload={
"width": 1280,
"height": 720,
"fps": 25,
"clips": [
{
"clip_type": "main",
"asset_id": "sample_5s.mp4",
"duration": 3.0,
"order": i,
"transition_effect": ["cut", "fade", "slideleft", "slidedown", "dissolve"][i % 5],
}
for i in range(10)
],
},
),
CompareScenario(
id="vertical_portrait",
name="竖屏9:16",
description="竖屏分辨率,验证scale策略(铺满裁剪)",
priority="P2",
plan_payload={
"width": 720,
"height": 1280,
"fps": 25,
"clips": [
{"clip_type": "main", "asset_id": "sample_5s.mp4", "duration": 5.0, "order": 0},
],
},
),
]
def get_scenarios_by_priority(min_priority: str = "P2") -> list[CompareScenario]:
"""按优先级过滤场景.
P0 包含 P0
P1 包含 P0 + P1
P2 包含全部
"""
priority_order = {"P0": 0, "P1": 1, "P2": 2}
threshold = priority_order.get(min_priority, 2)
return [s for s in SCENARIOS if priority_order.get(s.priority, 2) <= threshold]
-283
View File
@@ -1,283 +0,0 @@
"""视频对比工具 — 基于 FFmpeg 的像素级质量对比.
使用 SSIM + PSNR 双指标评估两个视频的相似度
- SSIM (Structural Similarity): 结构相似性范围 [0, 1]越接近 1 越相似
- PSNR (Peak Signal-to-Noise Ratio): 峰值信噪比单位 dB越高越好
灰度验收标准
- 平均 SSIM >= 0.95 视觉上几乎无差异P0 场景必达
- 最低 SSIM >= 0.90 最严重帧差异可接受
- 平均 PSNR >= 28dB 质量达标
"""
from __future__ import annotations
import json
import re
import shutil
import subprocess # nosec B404
from dataclasses import asdict, dataclass
from pathlib import Path
from typing import Any
FFMPEG_BIN: str = shutil.which("ffmpeg") or "ffmpeg"
FFPROBE_BIN: str = shutil.which("ffprobe") or "ffprobe"
@dataclass
class VideoDiffResult:
"""视频对比结果."""
video_a: str
video_b: str
width: int
height: int
duration_a: float
duration_b: float
avg_ssim: float
min_ssim: float
avg_psnr: float # dB
min_psnr: float
frame_count: int
duration_diff: float # 时长差(秒)
resolution_match: bool
passed: bool # 是否通过阈值
def to_dict(self) -> dict[str, Any]:
return asdict(self)
def probe_video_info(video_path: str) -> dict[str, Any]:
"""获取视频信息(宽、高、时长、fps."""
try:
result = subprocess.run( # nosec B603
[
FFPROBE_BIN,
"-v",
"error",
"-select_streams",
"v:0",
"-show_entries",
"stream=width,height,r_frame_rate,duration",
"-show_entries",
"format=duration",
"-of",
"json",
video_path,
],
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=10,
)
info = json.loads(result.stdout)
stream = info.get("streams", [{}])[0]
fmt = info.get("format", {})
width = int(stream.get("width", 1280))
height = int(stream.get("height", 720))
fps_str = stream.get("r_frame_rate", "25/1")
if "/" in fps_str:
num, den = fps_str.split("/")
fps = float(num) / float(den) if float(den) > 0 else 25.0
else:
fps = float(fps_str) if fps_str else 25.0
duration = float(fmt.get("duration", 0)) or float(stream.get("duration", 0))
return {"width": width, "height": height, "duration": duration, "fps": round(fps, 2)}
except Exception:
return {"width": 1280, "height": 720, "duration": 0.0, "fps": 25.0}
def compute_video_diff(
video_a: str | Path,
video_b: str | Path,
*,
ssim_threshold: float = 0.95,
psnr_threshold: float = 28.0,
duration_tolerance: float = 0.1,
) -> VideoDiffResult:
"""计算两个视频的像素差异.
使用 FFmpeg ssim + psnr 滤镜一次性计算两个指标
Args:
video_a: 视频A路径基线
video_b: 视频B路径对比
ssim_threshold: SSIM 合格阈值默认 0.90
psnr_threshold: PSNR 合格阈值默认 25dB
duration_tolerance: 时长容忍度默认 0.1s
Returns:
VideoDiffResult 对比结果
Raises:
subprocess.CalledProcessError: FFmpeg 执行失败
"""
info_a = probe_video_info(str(video_a))
info_b = probe_video_info(str(video_b))
duration_diff = abs(info_a["duration"] - info_b["duration"])
resolution_match = info_a["width"] == info_b["width"] and info_a["height"] == info_b["height"]
# ssim 和 psnr 的 stats_file 都输出到 stdout
# 用行格式区分:SSIM 行含 "All:"PSNR 行含 "psnr_avg:"
command = [
FFMPEG_BIN,
"-i",
str(video_a),
"-i",
str(video_b),
"-lavfi",
"[0:v][1:v]ssim=stats_file=-[out1];[0:v][1:v]psnr=stats_file=-[out2]",
"-f",
"null",
"-",
]
result = subprocess.run( # nosec B603
command,
check=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
timeout=300,
)
# 逐帧统计在 stdoutstats_file=-),汇总日志在 stderr
stats_stdout = result.stdout or ""
avg_ssim, min_ssim = _parse_ssim_stats(stats_stdout)
avg_psnr, min_psnr = _parse_psnr_stats(stats_stdout)
frame_count = _count_frames(result.stderr or "")
passed = (
resolution_match
and duration_diff <= duration_tolerance
and avg_ssim >= ssim_threshold
and avg_psnr >= psnr_threshold
)
return VideoDiffResult(
video_a=str(video_a),
video_b=str(video_b),
width=info_a["width"],
height=info_a["height"],
duration_a=round(info_a["duration"], 3),
duration_b=round(info_b["duration"], 3),
avg_ssim=round(avg_ssim, 6),
min_ssim=round(min_ssim, 6),
avg_psnr=round(avg_psnr, 3),
min_psnr=round(min_psnr, 3),
frame_count=frame_count,
duration_diff=round(duration_diff, 3),
resolution_match=resolution_match,
passed=passed,
)
def _parse_ssim_stats(stats_output: str) -> tuple[float, float]:
"""从 SSIM stats_file 输出中解析逐帧 SSIM.
FFmpeg ssim 滤镜 stats_file 输出格式每行一帧:
n:1 Y:0.987654 U:0.991234 V:0.990000 All:0.989000 (19.585642)
n:2 Y:0.986543 U:0.990123 V:0.988888 All:0.987654 (19.123456)
...
Returns:
(avg_ssim, min_ssim)
"""
ssim_values: list[float] = []
for line in stats_output.split("\n"):
# 匹配 stats_file 格式:n:数字 ... All:数字
if not line.startswith("n:"):
continue
match = re.search(r"All:(\d+\.\d+)", line)
if match:
ssim_values.append(float(match.group(1)))
if not ssim_values:
return 0.0, 0.0
avg_ssim = sum(ssim_values) / len(ssim_values)
min_ssim = min(ssim_values)
return avg_ssim, min_ssim
def _parse_psnr_stats(stats_output: str) -> tuple[float, float]:
"""从 PSNR stats_file 输出中解析逐帧 PSNR.
FFmpeg psnr 滤镜 stats_file 输出格式每行一帧:
n:1 mse_avg:100.23 mse_y:150.12 mse_u:50.34 mse_v:80.56 psnr_avg:28.12 psnr_y:26.34 psnr_u:31.12 psnr_v:29.08
n:2 ...
Returns:
(avg_psnr, min_psnr) avg_psnr 是逐帧 psnr_avg 的均值min_psnr 是逐帧最小值
"""
psnr_values: list[float] = []
for line in stats_output.split("\n"):
if not line.startswith("n:"):
continue
match = re.search(r"psnr_avg:(\d+\.\d+)", line)
if match:
psnr_values.append(float(match.group(1)))
if not psnr_values:
return 0.0, 0.0
avg_psnr = sum(psnr_values) / len(psnr_values)
min_psnr = min(psnr_values)
return avg_psnr, min_psnr
def _count_frames(stderr: str) -> int:
"""从 FFmpeg 输出中统计帧数."""
match = re.search(r"frame=\s*(\d+)", stderr)
return int(match.group(1)) if match else 0
def save_diff_frame(
video_a: str | Path,
video_b: str | Path,
output_path: str | Path,
*,
timestamp: float = 1.0,
) -> Path:
"""生成差异帧可视化图(红绿色差).
使用 blend 滤镜生成差异可视化图差异越大越亮
Args:
video_a: 视频A
video_b: 视频B
output_path: 输出图片路径
timestamp: 截取的时间点
Returns:
输出图片路径
"""
command = [
FFMPEG_BIN,
"-y",
"-ss",
str(timestamp),
"-i",
str(video_a),
"-ss",
str(timestamp),
"-i",
str(video_b),
"-lavfi",
"[0:v][1:v]blend=all_mode=difference,eq=contrast=5:brightness=0.5[diff]",
"-map",
"[diff]",
"-vframes",
"1",
str(output_path),
]
subprocess.run(command, check=True, capture_output=True, timeout=60) # nosec B603
return Path(output_path)
-2
View File
@@ -63,8 +63,6 @@ class StubEditPlan:
template_id: str = "tmpl-001"
status: Any = None
config: dict = field(default_factory=dict)
project_id: str = ""
created_by_user_id: str = "user-001"
def mark_failed(self):
self.status = _StubStatus("failed")
-424
View File
@@ -1,424 +0,0 @@
"""Feature Flag 单元测试。
测试 FeatureFlagConfigInMemoryFeatureFlagStoreRenderEngineResolver 的核心逻辑
"""
from __future__ import annotations
import time
from unittest.mock import MagicMock, patch
import pytest
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
InMemoryFeatureFlagStore,
)
# ── FeatureFlagConfig 测试 ──────────────────────────────────────────────────
class TestFeatureFlagConfig:
"""FeatureFlagConfig 核心逻辑测试。"""
def test_default_disabled(self):
"""默认配置为关闭状态。"""
config = FeatureFlagConfig(name="test_flag")
assert config.enabled is False
assert config.percentage == 0
assert config.whitelist == set()
assert config.is_active() is False
assert config.is_active("user1") is False
def test_global_enabled_100_percent(self):
"""100% + 启用 = 全部命中。"""
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=100)
assert config.is_active() is True
assert config.is_active("user1") is True
assert config.is_active("any_user") is True
def test_global_enabled_0_percent_no_whitelist(self):
"""启用但 0% 且无白名单 = 不命中。"""
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=0)
assert config.is_active() is False
assert config.is_active("user1") is False
def test_whitelist_takes_priority(self):
"""白名单优先级高于百分比。"""
config = FeatureFlagConfig(
name="test_flag",
enabled=True,
percentage=0,
whitelist={"user1", "user2"},
)
assert config.is_active("user1") is True
assert config.is_active("user2") is True
assert config.is_active("user3") is False
def test_whitelist_with_percentage(self):
"""白名单用户即使百分比为0也命中,非白名单按百分比。"""
config = FeatureFlagConfig(
name="test_flag",
enabled=True,
percentage=100, # 100% 所有人命中
whitelist={"user1"},
)
assert config.is_active("user1") is True
assert config.is_active("user999") is True # 100% 命中
def test_percentage_consistency_same_user(self):
"""同一用户多次调用结果一致(哈希确定性)。"""
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=50)
results = [config.is_active("user_fixed") for _ in range(100)]
assert all(r == results[0] for r in results)
def test_percentage_different_users_distributed(self):
"""不同用户分布大致符合百分比(统计检验,宽松阈值)。"""
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=50)
active_count = sum(1 for i in range(1000) if config.is_active(f"user_{i}"))
# 50% 上下浮动 10% 都算合理
assert 400 <= active_count <= 600, f"Expected ~500, got {active_count}"
def test_percentage_boundary_0_and_100(self):
"""0% 和 100% 的边界情况。"""
config_0 = FeatureFlagConfig(name="test", enabled=True, percentage=0)
config_100 = FeatureFlagConfig(name="test", enabled=True, percentage=100)
for i in range(100):
assert config_0.is_active(f"user_{i}") is False
assert config_100.is_active(f"user_{i}") is True
def test_disabled_ignores_all_other_settings(self):
"""关闭时忽略白名单和百分比。"""
config = FeatureFlagConfig(
name="test_flag",
enabled=False,
percentage=100,
whitelist={"user1"},
)
assert config.is_active("user1") is False
assert config.is_active() is False
def test_none_identifier_with_percentage(self):
"""无 identifier 时按随机比例(0% 和 100% 是确定的)。"""
config_0 = FeatureFlagConfig(name="test", enabled=True, percentage=0)
config_100 = FeatureFlagConfig(name="test", enabled=True, percentage=100)
assert config_0.is_active(None) is False
assert config_100.is_active(None) is True
def test_to_dict_and_from_dict(self):
"""序列化和反序列化对称。"""
original = FeatureFlagConfig(
name="test_flag",
enabled=True,
percentage=30,
whitelist={"user_a", "user_b", "user_c"},
)
data = original.to_dict()
restored = FeatureFlagConfig.from_dict(data)
assert restored.name == original.name
assert restored.enabled == original.enabled
assert restored.percentage == original.percentage
assert restored.whitelist == original.whitelist
def test_from_dict_with_missing_fields(self):
"""from_dict 缺失字段时使用默认值。"""
config = FeatureFlagConfig.from_dict({"name": "minimal"})
assert config.name == "minimal"
assert config.enabled is False
assert config.percentage == 0
assert config.whitelist == set()
# ── InMemoryFeatureFlagStore 测试 ───────────────────────────────────────────
class TestInMemoryFeatureFlagStore:
"""内存存储实现测试。"""
def test_get_nonexistent_returns_default(self):
"""获取不存在的 flag 返回默认配置(关闭)。"""
store = InMemoryFeatureFlagStore()
config = store.get("nonexistent")
assert config.name == "nonexistent"
assert config.enabled is False
def test_set_and_get(self):
"""设置后可以读取。"""
store = InMemoryFeatureFlagStore()
config = FeatureFlagConfig(name="test", enabled=True, percentage=50, whitelist={"u1"})
store.set(config)
got = store.get("test")
assert got.enabled is True
assert got.percentage == 50
assert got.whitelist == {"u1"}
def test_delete_existing(self):
"""删除存在的 flag 返回 True。"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="test", enabled=True))
assert store.delete("test") is True
assert store.get("test").enabled is False
def test_delete_nonexistent(self):
"""删除不存在的 flag 返回 False。"""
store = InMemoryFeatureFlagStore()
assert store.delete("nonexistent") is False
def test_list_all(self):
"""列出所有 flag。"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="flag_a", enabled=True))
store.set(FeatureFlagConfig(name="flag_b", percentage=10))
all_flags = store.list_all()
assert len(all_flags) == 2
assert "flag_a" in all_flags
assert "flag_b" in all_flags
assert all_flags["flag_a"].enabled is True
def test_is_active_convenience(self):
"""is_active 便捷方法。"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render", enabled=True, percentage=0, whitelist={"vip_user"}))
assert store.is_active("render", "vip_user") is True
assert store.is_active("render", "normal_user") is False
assert store.is_active("nonexistent") is False
# ── RenderEngineResolver 测试 ───────────────────────────────────────────────
class TestRenderEngineResolver:
"""渲染引擎选择器测试。"""
def test_default_legacy_when_flag_disabled(self):
"""flag 关闭时使用默认引擎(legacy)。"""
store = InMemoryFeatureFlagStore()
resolver = self._make_resolver(store=store, default="legacy")
assert resolver.get_engine() == "legacy"
assert resolver.get_engine("user1") == "legacy"
def test_default_unified_when_flag_disabled(self):
"""flag 关闭但默认值是 unified 时返回 unified。"""
store = InMemoryFeatureFlagStore()
resolver = self._make_resolver(store=store, default="unified")
assert resolver.get_engine() == "unified"
def test_whitelist_user_uses_unified(self):
"""白名单用户走新引擎。"""
store = InMemoryFeatureFlagStore()
store.set(
FeatureFlagConfig(
name="render_engine",
enabled=True,
percentage=0,
whitelist={"beta_tester"},
)
)
resolver = self._make_resolver(store=store, default="legacy")
assert resolver.get_engine("beta_tester") == "unified"
assert resolver.get_engine("normal_user") == "legacy"
def test_100_percent_all_unified(self):
"""100% 时所有用户走新引擎。"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render_engine", enabled=True, percentage=100))
resolver = self._make_resolver(store=store, default="legacy")
for i in range(50):
assert resolver.get_engine(f"user_{i}") == "unified"
def test_invalid_default_engine_fallback(self):
"""无效默认值回退到 legacy。"""
store = InMemoryFeatureFlagStore()
resolver = self._make_resolver(store=store, default="invalid_value")
assert resolver.get_engine() == "legacy"
def test_should_use_unified_helper(self):
"""should_use_unified 便捷方法。"""
store = InMemoryFeatureFlagStore()
store.set(
FeatureFlagConfig(
name="render_engine",
enabled=True,
percentage=0,
whitelist={"user_a"},
)
)
resolver = self._make_resolver(store=store)
assert resolver.should_use_unified("user_a") is True
assert resolver.should_use_unified("user_b") is False
def test_config_snapshot(self):
"""配置快照。"""
store = InMemoryFeatureFlagStore()
store.set(
FeatureFlagConfig(
name="render_engine",
enabled=True,
percentage=30,
whitelist={"u1", "u2"},
)
)
resolver = self._make_resolver(store=store)
snapshot = resolver.get_config_snapshot()
assert snapshot["flag_name"] == "render_engine"
assert snapshot["enabled"] is True
assert snapshot["percentage"] == 30
assert snapshot["whitelist"] == ["u1", "u2"]
def test_set_flag_updates_config(self):
"""通过 set_flag 修改后立即生效。"""
store = InMemoryFeatureFlagStore()
resolver = self._make_resolver(store=store, default="legacy")
# 初始:关闭
assert resolver.get_engine("user1") == "legacy"
# 开启 100%
resolver.set_flag(FeatureFlagConfig(name="render_engine", enabled=True, percentage=100))
assert resolver.get_engine("user1") == "unified"
# 关闭
resolver.set_flag(FeatureFlagConfig(name="render_engine", enabled=False))
assert resolver.get_engine("user1") == "legacy"
def test_force_refresh(self):
"""强制刷新不报错。"""
store = InMemoryFeatureFlagStore()
resolver = self._make_resolver(store=store)
resolver.force_refresh() # 不抛异常即可
def test_does_not_affect_in_flight_tasks(self):
"""
热更新不影响在途任务验证
任务开始时确定引擎中途配置变更不改变当前任务的引擎选择
这是通过"每次调用 get_engine 时读取当前配置"来保证的
任务开始时调用一次拿到结果之后不再变化
"""
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render_engine", enabled=True, percentage=100))
resolver = self._make_resolver(store=store, default="legacy")
# 模拟任务开始时获取引擎
engine_at_start = resolver.get_engine("user1")
assert engine_at_start == "unified"
# 任务进行中关闭 flag
store.set(FeatureFlagConfig(name="render_engine", enabled=False))
resolver.force_refresh()
# 在途任务持有的 engine_at_start 仍然是 unified(不随配置变化)
assert engine_at_start == "unified"
# 新任务会拿到 legacy
assert resolver.get_engine("user1") == "legacy"
# ── 辅助方法 ──
@staticmethod
def _make_resolver(store=None, default="legacy"):
from apps.worker.video_processing.render_engine_resolver import (
RenderEngineResolver,
)
return RenderEngineResolver(
default_engine=default,
store=store or InMemoryFeatureFlagStore(),
refresh_interval=9999, # 测试时禁用自动刷新
)
# ── RedisFeatureFlagStore 降级测试(无 Redis 环境) ───────────────────────
class TestRedisStoreDegradation:
"""Redis 不可用时的降级行为测试。"""
def test_get_returns_default_when_redis_unavailable(self):
"""Redis 连接失败时返回默认关闭配置,不抛异常。"""
import importlib
from packages.adapters.redis import feature_flag_store as ff_module
# 模拟 redis 模块不存在的场景不好做,这里直接测试异常捕获逻辑
store = ff_module.RedisFeatureFlagStore.__new__(ff_module.RedisFeatureFlagStore)
store._redis = MagicMock()
store._redis.hgetall.side_effect = ConnectionError("Redis down")
store._key_prefix = ff_module.FEATURE_FLAG_REDIS_PREFIX
store._cache = {}
store._cache_ttl = 5.0
import threading
store._lock = threading.Lock()
config = store.get("render_engine")
assert config.enabled is False
assert config.name == "render_engine"
def test_list_all_returns_empty_on_redis_error(self):
"""Redis 错误时 list_all 返回空字典。"""
import importlib
from packages.adapters.redis import feature_flag_store as ff_module
store = ff_module.RedisFeatureFlagStore.__new__(ff_module.RedisFeatureFlagStore)
store._redis = MagicMock()
store._redis.scan.side_effect = ConnectionError("Redis down")
store._key_prefix = ff_module.FEATURE_FLAG_REDIS_PREFIX
store._cache = {}
store._cache_ttl = 5.0
import threading
store._lock = threading.Lock()
result = store.list_all()
assert result == {}
class TestRedisStoreListAll:
"""RedisFeatureFlagStore list_all 正常路径测试。"""
def _make_store(self):
from packages.adapters.redis import feature_flag_store as ff_module
store = ff_module.RedisFeatureFlagStore.__new__(ff_module.RedisFeatureFlagStore)
store._redis = MagicMock()
store._key_prefix = ff_module.FEATURE_FLAG_REDIS_PREFIX
store._cache = {}
store._cache_ttl = 5.0
import threading
store._lock = threading.Lock()
return store
def test_list_all_scan_with_match_param(self):
"""list_all 调用 redis.scan 时使用正确的 match 参数名。"""
store = self._make_store()
prefix = store._key_prefix
# 模拟 scan 返回 2 个 key,分 2 次游标
store._redis.scan.side_effect = [
(10, [f"{prefix}render_engine", f"{prefix}other_flag"]),
(0, []),
]
# 模拟 hgetall 返回配置
store._redis.hgetall.return_value = {
b"enabled": b"true",
b"percentage": b"50",
b"whitelist": b'["user1","user2"]',
}
result = store.list_all()
# 验证 scan 被调用了 2 次(游标遍历)
assert store._redis.scan.call_count == 2
# 验证参数名是 match(不是 match_pattern
first_call_kwargs = store._redis.scan.call_args_list[0][1]
assert "match" in first_call_kwargs
assert "match_pattern" not in first_call_kwargs
assert first_call_kwargs["match"] == f"{prefix}*"
# 验证返回了 2 个 flag
assert len(result) == 2
assert "render_engine" in result
assert "other_flag" in result
@@ -1,93 +0,0 @@
"""FFmpeg 超时保护测试。
验证 run_ffmpeg / probe_video_info 的超时保护机制
防止 FFmpeg hang 住导致 worker 永久阻塞
"""
from __future__ import annotations
import subprocess
from unittest.mock import MagicMock, patch
import pytest
from video_processing.ffmpeg_utils import (
DEFAULT_FFMPEG_TIMEOUT,
probe_video_info,
run_ffmpeg,
)
# ── run_ffmpeg 超时保护 ──────────────────────────────────────────────────────
class TestRunFFmpegTimeout:
"""run_ffmpeg 超时保护测试。"""
def test_default_timeout_is_set(self):
"""默认超时应为 1800 秒(30分钟)。"""
assert DEFAULT_FFMPEG_TIMEOUT == 1800
def test_timeout_expired_is_raised(self):
"""超时未完成时 TimeoutExpired 异常被传播。"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_run.side_effect = subprocess.TimeoutExpired(cmd=["ffmpeg", "test"], timeout=1)
with pytest.raises(subprocess.TimeoutExpired):
run_ffmpeg(["ffmpeg", "test"])
def test_custom_timeout(self):
"""支持自定义超时时间。"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_run.side_effect = subprocess.TimeoutExpired(cmd=["ffmpeg"], timeout=5)
with pytest.raises(subprocess.TimeoutExpired):
run_ffmpeg(["ffmpeg", "test"], timeout=5)
def test_none_timeout_disables_protection(self):
"""timeout=None 可以禁用超时保护(不推荐)。"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_result = MagicMock()
mock_result.stdout = ""
mock_result.stderr = ""
mock_run.return_value = mock_result
run_ffmpeg(["ffmpeg", "test"], timeout=None)
# 验证 timeout=None 被传递
call_kwargs = mock_run.call_args.kwargs
assert call_kwargs["timeout"] is None
def test_called_process_error_still_raised(self):
"""超时异常不影响原有 CalledProcessError 的抛出。"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_run.side_effect = subprocess.CalledProcessError(returncode=1, cmd=["ffmpeg"], stderr="error msg")
with pytest.raises(subprocess.CalledProcessError):
run_ffmpeg(["ffmpeg", "test"])
# ── probe_video_info 超时保护 ────────────────────────────────────────────────
class TestProbeVideoInfoTimeout:
"""probe_video_info 超时保护测试。"""
def test_probe_uses_timeout(self):
"""probe_video_info 调用 ffprobe 时应设置 timeout=15。"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_run.side_effect = subprocess.TimeoutExpired(cmd=["ffprobe"], timeout=15)
# 超时异常被捕获,返回默认值
result = probe_video_info("/tmp/test.mp4")
assert result["width"] == 1280 # DEFAULT_OUTPUT_WIDTH
assert result["height"] == 720 # DEFAULT_OUTPUT_HEIGHT
def test_probe_success(self):
"""正常情况应解析 ffprobe JSON 输出。"""
fake_output = """
{
"streams": [{"width": 1920, "height": 1080, "codec_type": "video", "r_frame_rate": "30/1", "duration": "10.5"}],
"format": {"duration": "10.5"}
}
"""
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
mock_result = MagicMock()
mock_result.stdout = fake_output
mock_run.return_value = mock_result
result = probe_video_info("/tmp/test.mp4")
assert result["width"] == 1920
assert result["height"] == 1080
assert abs(result["duration"] - 10.5) < 0.01
-338
View File
@@ -1,338 +0,0 @@
"""generate_video 任务 Feature Flag 灰度引擎选择单元测试.
覆盖
- _resolve_render_engine 正常返回 unified / legacy
- Feature Flag 不可用时 fallback unified
- 白名单 / 百分比 / 全局开关各场景
- _render_with_legacy_engine 命令构建与输出验证
"""
from __future__ import annotations
import os
import sys
from datetime import datetime, timezone
from types import ModuleType
from typing import Any
from unittest.mock import MagicMock, patch
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
from pathlib import Path
import pytest
# ── Mock worker 模块以避免数据库连接 ──────────────────────────────────────────
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "worker"))
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", ".."))
_mock_db_mod = ModuleType("worker_app.db")
_mock_db_mod.SessionLocal = MagicMock()
sys.modules.setdefault("worker_app.db", _mock_db_mod)
_mock_celery_mod = ModuleType("worker_app.celery_app")
_mock_celery_app = MagicMock()
_mock_celery_app.task = lambda **kwargs: lambda fn: fn
_mock_celery_mod.celery_app = _mock_celery_app
sys.modules.setdefault("worker_app.celery_app", _mock_celery_mod)
# Mock worker_app.core.config 避免 settings 加载
_mock_config_mod = ModuleType("worker_app.core.config")
_mock_settings = MagicMock()
_mock_settings.redis_url = None
_mock_settings.render_engine = "unified"
_mock_config_mod.get_settings = lambda: _mock_settings
sys.modules.setdefault("worker_app.core", ModuleType("worker_app.core"))
sys.modules.setdefault("worker_app.core.config", _mock_config_mod)
# ── 测试用数据类 ──────────────────────────────────────────────────────────────
class _TestClip:
def __init__(self, asset_id, duration=30.0, clip_type="main", config=None, order=0):
self.id = f"clip_{asset_id}"
self.plan_id = "test-plan"
self.clip_type = clip_type
self.order = order
self.asset_id = asset_id
self.duration = duration
self.config = config or {}
self.start_time = 0.0
self.transition_effect = "cut"
# ── RenderEngineResolver 基础行为测试 ───────────────────────────────────────
def test_resolver_unified_when_enabled_100_percent():
"""flag 全局开启(percentage=100)时,返回 unified。"""
from video_processing.render_engine_resolver import RenderEngineResolver
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
InMemoryFeatureFlagStore,
)
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render_engine", enabled=True, percentage=100))
resolver = RenderEngineResolver(default_engine="legacy", store=store)
assert resolver.get_engine(user_id="user-123") == "unified"
def test_resolver_legacy_when_flag_disabled():
"""flag 全局关闭时,返回默认引擎 legacy。"""
from video_processing.render_engine_resolver import RenderEngineResolver
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
InMemoryFeatureFlagStore,
)
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render_engine", enabled=False, percentage=100))
resolver = RenderEngineResolver(default_engine="legacy", store=store)
assert resolver.get_engine(user_id="user-123") == "legacy"
def test_resolver_whitelist_overrides_percentage_0():
"""白名单用户即使 percentage=0 也走 unified。"""
from video_processing.render_engine_resolver import RenderEngineResolver
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
InMemoryFeatureFlagStore,
)
store = InMemoryFeatureFlagStore()
store.set(
FeatureFlagConfig(
name="render_engine",
enabled=True,
percentage=0,
whitelist={"user-vip"},
)
)
resolver = RenderEngineResolver(default_engine="legacy", store=store)
assert resolver.get_engine(user_id="user-vip") == "unified"
assert resolver.get_engine(user_id="user-other") == "legacy"
def test_resolver_percentage_0_all_legacy():
"""percentage=0 且无白名单时,全部走 legacy。"""
from video_processing.render_engine_resolver import RenderEngineResolver
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
InMemoryFeatureFlagStore,
)
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render_engine", enabled=True, percentage=0))
resolver = RenderEngineResolver(default_engine="legacy", store=store)
for i in range(50):
assert resolver.get_engine(user_id=f"user-{i}") == "legacy"
def test_resolver_default_unified_when_flag_off():
"""默认引擎设为 unified 且 flag 关闭时,返回 unified。"""
from video_processing.render_engine_resolver import RenderEngineResolver
from packages.adapters.redis.feature_flag_store import (
FeatureFlagConfig,
InMemoryFeatureFlagStore,
)
store = InMemoryFeatureFlagStore()
store.set(FeatureFlagConfig(name="render_engine", enabled=False, percentage=0))
resolver = RenderEngineResolver(default_engine="unified", store=store)
assert resolver.get_engine(user_id="user-123") == "unified"
# ── _render_with_legacy_engine 集成测试 ──────────────────────────────────────
def test_legacy_engine_single_clip_keeps_original_fps():
"""单 clip 场景:输出保持原帧率(不做 fps 归一化),分辨率缩放正确。"""
import subprocess
import tempfile
from video_processing.ffmpeg_utils import probe_video_info
from apps.worker.worker_app.tasks.generation import _render_with_legacy_engine
with tempfile.TemporaryDirectory() as tmpdir:
tmp_path = Path(tmpdir)
input_path = tmp_path / "input.mp4"
output_path = tmp_path / "output.mp4"
# 生成 1 秒 30fps 测试视频(带音频)
subprocess.run(
[
"ffmpeg",
"-y",
"-f",
"lavfi",
"-i",
"color=c=red:s=640x360:d=1:r=30",
"-f",
"lavfi",
"-i",
"anullsrc=r=44100:cl=stereo:d=1",
"-c:v",
"libx264",
"-pix_fmt",
"yuv420p",
"-c:a",
"aac",
"-shortest",
str(input_path),
],
check=True,
capture_output=True,
)
clip = _TestClip(asset_id="asset-1", duration=1.0)
asset_path_map = {"asset-1": input_path}
duration, file_size = _render_with_legacy_engine(
task_id="test-task",
virtual_clips=[clip],
asset_path_map=asset_path_map,
work_dir=tmp_path,
output_path=output_path,
)
assert output_path.exists()
assert file_size > 0
assert duration > 0
# 旧引擎保持原帧率(30fps),不做 fps 归一化
info = probe_video_info(str(output_path))
assert abs(info.get("fps", 0) - 30.0) < 0.5
assert info.get("width") == 1280
assert info.get("height") == 720
def test_legacy_engine_two_clips_concat_duration():
"""多 clip 场景:concat 后时长为两片段之和。"""
import subprocess
import tempfile
from video_processing.ffmpeg_utils import probe_duration
from apps.worker.worker_app.tasks.generation import _render_with_legacy_engine
with tempfile.TemporaryDirectory() as tmpdir:
tmp_path = Path(tmpdir)
input1 = tmp_path / "input1.mp4"
input2 = tmp_path / "input2.mp4"
output_path = tmp_path / "output.mp4"
for idx, inp in enumerate([input1, input2]):
color = "red" if idx == 0 else "blue"
subprocess.run(
[
"ffmpeg",
"-y",
"-f",
"lavfi",
"-i",
f"color=c={color}:s=640x360:d=1:r=30",
"-f",
"lavfi",
"-i",
"anullsrc=r=44100:cl=stereo:d=1",
"-c:v",
"libx264",
"-pix_fmt",
"yuv420p",
"-c:a",
"aac",
"-shortest",
str(inp),
],
check=True,
capture_output=True,
)
clip1 = _TestClip(asset_id="asset-1", duration=1.0, clip_type="main", order=0)
clip2 = _TestClip(asset_id="asset-2", duration=1.0, clip_type="main", order=1)
asset_path_map = {"asset-1": input1, "asset-2": input2}
duration, file_size = _render_with_legacy_engine(
task_id="test-task",
virtual_clips=[clip1, clip2],
asset_path_map=asset_path_map,
work_dir=tmp_path,
output_path=output_path,
)
assert output_path.exists()
assert file_size > 0
assert abs(duration - 2.0) < 0.2
def test_legacy_engine_broll_mode_supported():
"""b_roll 类型的 clip 也被正确识别为主图层并渲染。"""
import subprocess
import tempfile
from apps.worker.worker_app.tasks.generation import _render_with_legacy_engine
with tempfile.TemporaryDirectory() as tmpdir:
tmp_path = Path(tmpdir)
input_path = tmp_path / "input.mp4"
output_path = tmp_path / "output.mp4"
subprocess.run(
[
"ffmpeg",
"-y",
"-f",
"lavfi",
"-i",
"color=c=green:s=640x360:d=1:r=30",
"-f",
"lavfi",
"-i",
"anullsrc=r=44100:cl=stereo:d=1",
"-c:v",
"libx264",
"-pix_fmt",
"yuv420p",
"-c:a",
"aac",
"-shortest",
str(input_path),
],
check=True,
capture_output=True,
)
clip = _TestClip(
asset_id="asset-1",
duration=1.0,
clip_type="main",
config={"role": "b_roll"},
)
asset_path_map = {"asset-1": input_path}
duration, file_size = _render_with_legacy_engine(
task_id="test-task",
virtual_clips=[clip],
asset_path_map=asset_path_map,
work_dir=tmp_path,
output_path=output_path,
)
assert output_path.exists()
assert file_size > 0
assert duration > 0
-190
View File
@@ -1,190 +0,0 @@
"""渲染结果内部下载接口单元测试。
测试 internal_render 路由的核心逻辑mock repository storage 依赖
"""
from __future__ import annotations
from unittest.mock import MagicMock
import pytest
from app.api.routes.internal_render import (
InternalRenderDownloadUrlResponse,
InternalRenderTaskVideosResponse,
_video_to_item,
get_render_task_videos,
get_render_video_download_url,
)
# ── Helpers ────────────────────────────────────────────────────────────────
class MockVideo:
"""模拟 GeneratedVideo 领域对象。"""
def __init__(self, **kwargs):
self.id = kwargs.get("id", "video-1")
self.generation_task_id = kwargs.get("generation_task_id", "task-1")
self.project_id = kwargs.get("project_id", "proj-1")
self.name = kwargs.get("name", "test_video.mp4")
self.file_url = kwargs.get("file_url", "videos/test/output.mp4")
self.file_size = kwargs.get("file_size", 1024000)
self.duration = kwargs.get("duration", 30.5)
self.width = kwargs.get("width", 1080)
self.height = kwargs.get("height", 1920)
self.fps = kwargs.get("fps", 30.0)
self.status = kwargs.get("status", "completed")
# ── _video_to_item 测试 ────────────────────────────────────────────────────
class TestVideoToItem:
"""测试视频对象转响应项。"""
def test_basic_conversion(self):
video = MockVideo(id="v1", generation_task_id="t1", status="completed")
item = _video_to_item(video, "https://oss.example.com/download?v1")
assert item.video_id == "v1"
assert item.generation_task_id == "t1"
assert item.status == "completed"
assert item.download_url == "https://oss.example.com/download?v1"
def test_missing_optional_fields(self):
"""缺可选字段时返回 None。"""
video = MockVideo()
# 去掉可选字段
del video.file_size
del video.duration
item = _video_to_item(video, "https://example.com/dl")
assert item.file_size is None
assert item.duration is None
assert item.width == 1080 # 还在
# ── 路由函数测试 ────────────────────────────────────────────────────────────
class TestGetRenderVideoDownloadUrl:
"""测试单个视频下载URL接口。"""
def test_video_exists(self):
video = MockVideo(id="v-abc", file_url="videos/abc/out.mp4")
mock_repo = MagicMock()
mock_repo.get.return_value = video
mock_storage = MagicMock()
mock_storage.get_download_url.return_value = "https://oss.test/signed?v=abc"
result = get_render_video_download_url(
video_id="v-abc",
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
assert isinstance(result, InternalRenderDownloadUrlResponse)
assert result.video_id == "v-abc"
assert result.download_url == "https://oss.test/signed?v=abc"
mock_repo.get.assert_called_once_with("v-abc")
mock_storage.get_download_url.assert_called_once()
def test_video_not_found_raises_404(self):
from fastapi import HTTPException
mock_repo = MagicMock()
mock_repo.get.return_value = None
mock_storage = MagicMock()
with pytest.raises(HTTPException) as exc_info:
get_render_video_download_url(
video_id="nonexistent",
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
assert exc_info.value.status_code == 404
def test_download_url_long_expiry(self):
"""过期时间应为 24 小时(86400s)。"""
video = MockVideo(id="v1")
mock_repo = MagicMock()
mock_repo.get.return_value = video
mock_storage = MagicMock()
mock_storage.get_download_url.return_value = "https://oss.test/signed"
get_render_video_download_url(
video_id="v1",
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
# 验证 expires_seconds=86400
call_kwargs = mock_storage.get_download_url.call_args
assert call_kwargs.kwargs.get("expires_seconds") == 86400 or call_kwargs[1].get("expires_seconds") == 86400
class TestGetRenderTaskVideos:
"""测试任务视频列表接口。"""
def test_list_multiple_videos(self):
videos = [
MockVideo(id="v1", status="completed"),
MockVideo(id="v2", status="completed"),
MockVideo(id="v3", status="failed"),
]
mock_repo = MagicMock()
mock_repo.list_by_generation_task.return_value = videos
mock_storage = MagicMock()
mock_storage.get_download_url.return_value = "https://oss.test/signed"
result = get_render_task_videos(
task_id="task-1",
status=None,
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
assert isinstance(result, InternalRenderTaskVideosResponse)
assert result.task_id == "task-1"
assert result.count == 3
assert len(result.videos) == 3
def test_filter_by_status(self):
videos = [
MockVideo(id="v1", status="completed"),
MockVideo(id="v2", status="completed"),
MockVideo(id="v3", status="failed"),
]
mock_repo = MagicMock()
mock_repo.list_by_generation_task.return_value = videos
mock_storage = MagicMock()
mock_storage.get_download_url.return_value = "https://oss.test/signed"
result = get_render_task_videos(
task_id="task-1",
status="completed",
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
assert result.count == 2
assert all(v.status == "completed" for v in result.videos)
def test_empty_task(self):
mock_repo = MagicMock()
mock_repo.list_by_generation_task.return_value = []
mock_storage = MagicMock()
result = get_render_task_videos(
task_id="empty-task",
status=None,
_=True,
generated_video_repository=mock_repo,
storage_service=mock_storage,
)
assert result.count == 0
assert result.videos == []
-236
View File
@@ -1,236 +0,0 @@
"""P0-stagingOSS 上传崩溃修复测试.
测试
1. oss_bucket() 传递 connect_timeout 参数
2. upload_to_oss() 小文件走 put_object_from_file大文件走分片上传
3. upload_to_oss() 超时保护超过总超时返回 None
4. upload_to_oss() 异常时返回 None
"""
from __future__ import annotations
import os
import tempfile
import time
from pathlib import Path
from unittest.mock import MagicMock, patch
import pytest
# ── oss_bucket connect_timeout 测试 ───────────────────────────────────────────
class TestOSSBucketConnectTimeout:
"""测试 oss_bucket() 传递 connect_timeout 参数."""
def test_oss_bucket_has_connect_timeout(self):
"""oss_bucket 应传递 connect_timeout=10s 参数."""
from video_processing.oss_helpers import oss_bucket
mock_bucket_instance = MagicMock()
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket_instance) as mock_bucket_cls,
):
bucket = oss_bucket()
assert bucket is mock_bucket_instance
# 验证 connect_timeout 关键字参数
call_kwargs = mock_bucket_cls.call_args[1]
assert "connect_timeout" in call_kwargs, "oss_bucket 应传递 connect_timeout 参数"
assert (
call_kwargs["connect_timeout"] == 10
), f"connect_timeout 应为 10,实际为 {call_kwargs['connect_timeout']}"
def test_oss_bucket_no_config_returns_none(self):
"""OSS 配置缺失时返回 None."""
from video_processing.oss_helpers import oss_bucket
with patch.dict(os.environ, {}, clear=True):
bucket = oss_bucket()
assert bucket is None
# ── upload_to_oss 分片上传测试 ────────────────────────────────────────────────
class TestUploadToOSSMultipart:
"""测试 upload_to_oss() 根据文件大小选择上传方式."""
def _create_temp_file(self, size_bytes: int) -> Path:
"""创建指定大小的临时文件."""
tmp = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
tmp.write(b"x" * size_bytes)
tmp.close()
return Path(tmp.name)
def test_small_file_uses_put_object(self):
"""小文件(<100MB)走 put_object_from_file."""
from video_processing.oss_helpers import upload_to_oss
small_file = self._create_temp_file(10 * 1024 * 1024) # 10MB
try:
mock_bucket = MagicMock()
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
patch("video_processing.oss_helpers.oss2.resumable_upload") as mock_resumable,
):
url = upload_to_oss(small_file, "test/small.mp4")
# 验证调用了 put_object_from_file
mock_bucket.put_object_from_file.assert_called_once()
# 验证没调用分片上传
mock_resumable.assert_not_called()
# 验证返回 URL
assert url == "https://test-bucket.oss-cn-hangzhou.aliyuncs.com/test/small.mp4"
finally:
small_file.unlink()
def test_large_file_uses_resumable_upload(self):
"""大文件(>=100MB)走 resumable_upload 分片上传."""
from video_processing.oss_helpers import upload_to_oss
large_file = self._create_temp_file(100 * 1024 * 1024) # 100MB
try:
mock_bucket = MagicMock()
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
patch("video_processing.oss_helpers.oss2.resumable_upload") as mock_resumable,
):
url = upload_to_oss(large_file, "test/large.mp4")
# 验证调用了分片上传
mock_resumable.assert_called_once()
# 验证没调用 put_object_from_file
mock_bucket.put_object_from_file.assert_not_called()
# 验证分片参数
call_kwargs = mock_resumable.call_args[1]
assert call_kwargs["multipart_threshold"] == 100 * 1024 * 1024
assert call_kwargs["part_size"] == 8 * 1024 * 1024
assert call_kwargs["num_threads"] == 3
# 验证返回 URL
assert url == "https://test-bucket.oss-cn-hangzhou.aliyuncs.com/test/large.mp4"
finally:
large_file.unlink()
# ── upload_to_oss 超时测试 ────────────────────────────────────────────────────
class TestUploadToOSSTimeout:
"""测试 upload_to_oss() 超时保护."""
def test_upload_timeout_returns_none(self):
"""上传超过总超时时返回 None."""
from video_processing.oss_helpers import upload_to_oss
small_file = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
small_file.write(b"x" * 1024) # 1KB
small_file.close()
file_path = Path(small_file.name)
def slow_upload(*args, **kwargs):
"""模拟慢速上传,超过超时时间."""
time.sleep(2)
mock_bucket = MagicMock()
mock_bucket.put_object_from_file.side_effect = slow_upload
try:
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
patch("video_processing.oss_helpers.OSS_UPLOAD_TOTAL_TIMEOUT", 1), # 1秒超时
):
url = upload_to_oss(file_path, "test/slow.mp4")
# 超时应返回 None
assert url is None, "上传超时应返回 None"
finally:
file_path.unlink()
def test_upload_exception_returns_none(self):
"""上传异常时返回 None."""
from video_processing.oss_helpers import upload_to_oss
small_file = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
small_file.write(b"x" * 1024)
small_file.close()
file_path = Path(small_file.name)
mock_bucket = MagicMock()
mock_bucket.put_object_from_file.side_effect = RuntimeError("Network error")
try:
with (
patch.dict(
os.environ,
{
"OSS_ACCESS_KEY_ID": "test-key",
"OSS_ACCESS_KEY_SECRET": "test-secret",
"OSS_ENDPOINT": "oss-cn-hangzhou.aliyuncs.com",
"OSS_BUCKET_NAME": "test-bucket",
},
),
patch("video_processing.oss_helpers.oss2.Auth"),
patch("video_processing.oss_helpers.oss2.Bucket", return_value=mock_bucket),
):
url = upload_to_oss(file_path, "test/error.mp4")
assert url is None, "上传异常应返回 None"
finally:
file_path.unlink()
def test_upload_no_bucket_returns_none(self):
"""OSS 未配置时返回 None."""
from video_processing.oss_helpers import upload_to_oss
small_file = tempfile.NamedTemporaryFile(delete=False, suffix=".mp4")
small_file.write(b"x" * 1024)
small_file.close()
file_path = Path(small_file.name)
try:
with patch.dict(os.environ, {}, clear=True):
url = upload_to_oss(file_path, "test/noconfig.mp4")
assert url is None
finally:
file_path.unlink()
-782
View File
@@ -573,8 +573,6 @@ class TestPassThrough:
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch.object(svc, "_render_pass_through") as mock_pass,
patch.object(svc, "_execute_ffmpeg") as mock_exec,
patch.object(svc, "_mix_audio", return_value=None),
patch("shutil.copy2"),
patch.object(svc, "_probe_output", return_value=(5.0, 1024, 1280, 720)),
):
result = svc.render()
@@ -600,8 +598,6 @@ class TestPassThrough:
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch.object(svc, "_render_pass_through") as mock_pass,
patch.object(svc, "_execute_ffmpeg") as mock_exec,
patch.object(svc, "_mix_audio", return_value=None),
patch("shutil.copy2"),
patch.object(svc, "_probe_output", return_value=(5.5, 2048, 1280, 720)),
):
result = svc.render()
@@ -820,8 +816,6 @@ class TestRender:
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch.object(svc, "_render_pass_through") as mock_pass,
patch.object(svc, "_mix_audio", return_value=None),
patch("shutil.copy2"),
patch.object(svc, "_probe_output", return_value=(5.0, 1024, 1280, 720)),
):
result = svc.render()
@@ -849,8 +843,6 @@ class TestRender:
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch.object(svc, "_execute_ffmpeg") as mock_exec,
patch.object(svc, "_mix_audio", return_value=None),
patch("shutil.copy2"),
patch.object(svc, "_probe_output", return_value=(5.5, 2048, 1280, 720)),
):
result = svc.render()
@@ -861,777 +853,3 @@ class TestRender:
assert result.width == 1280
assert result.height == 720
mock_exec.assert_called_once()
# ── 测试音频后处理 ──────────────────────────────────────────────────────────
class TestAudioMixing:
"""测试音频后处理混音功能。"""
def test_clip_effective_duration_with_both(self):
"""指定时长和实际时长都有时取较小值。"""
clip = ResolvedClip(
clip_id="c1",
asset_id="a1",
local_path=Path("/tmp/c1.mp4"),
clip_type="main",
order=0,
duration=3.0,
actual_duration=5.0,
)
assert UnifiedRenderService._clip_effective_duration(clip) == 3.0
def test_clip_effective_duration_only_actual(self):
"""只有实际时长时用实际时长。"""
clip = ResolvedClip(
clip_id="c1",
asset_id="a1",
local_path=Path("/tmp/c1.mp4"),
clip_type="main",
order=0,
duration=0.0,
actual_duration=5.0,
)
assert UnifiedRenderService._clip_effective_duration(clip) == 5.0
def test_clip_effective_duration_only_specified(self):
"""只有指定时长时用指定时长。"""
clip = ResolvedClip(
clip_id="c1",
asset_id="a1",
local_path=Path("/tmp/c1.mp4"),
clip_type="main",
order=0,
duration=3.0,
actual_duration=0.0,
)
assert UnifiedRenderService._clip_effective_duration(clip) == 3.0
def test_clip_effective_duration_zero(self):
"""都没有时返回0。"""
clip = ResolvedClip(
clip_id="c1",
asset_id="a1",
local_path=Path("/tmp/c1.mp4"),
clip_type="main",
order=0,
duration=0.0,
actual_duration=0.0,
)
assert UnifiedRenderService._clip_effective_duration(clip) == 0.0
def test_mix_audio_single_main_clip(self):
"""单主clip时直接提取音频。"""
clips = [_make_clip("c1", "main", order=0, duration=5.0)]
asset_paths = {"asset_c1.mp4": Path("/tmp/asset_c1.mp4")}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch("video_processing.unified_render_service.run_ffmpeg") as mock_run,
):
layers = svc._group_clips_into_layers(svc._resolve_clips())
result = svc._mix_audio(layers, 5.0)
assert result is not None
assert result.name == "audio_plan_001.aac"
mock_run.assert_called_once()
# 验证命令包含 -vn(无视频)和 aac 编码
cmd = mock_run.call_args[0][0]
assert "-vn" in cmd
assert "aac" in cmd
def test_mix_audio_multi_main_clips(self):
"""多主clip时用concat拼接音频。"""
clips = [
_make_clip("c1", "main", order=0, duration=3.0),
_make_clip("c2", "main", order=1, duration=2.0),
]
asset_paths = {
"asset_c1.mp4": Path("/tmp/asset_c1.mp4"),
"asset_c2.mp4": Path("/tmp/asset_c2.mp4"),
}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch("video_processing.unified_render_service.run_ffmpeg") as mock_run,
):
layers = svc._group_clips_into_layers(svc._resolve_clips())
result = svc._mix_audio(layers, 4.5)
assert result is not None
mock_run.assert_called_once()
cmd = mock_run.call_args[0][0]
# 验证有 filter_complex 和 concat
assert "-filter_complex" in cmd
cmd_str = " ".join(cmd)
assert "concat=n=2:v=0:a=1" in cmd_str
def test_mix_audio_with_independent_audio_track(self):
"""有独立音频轨时用amix混音。"""
clips = [
_make_clip("c1", "main", order=0, duration=5.0),
_make_clip(
"bgm1",
"main",
order=0,
duration=5.0,
config={"role": "audio", "volume": 0.5},
),
]
asset_paths = {
"asset_c1.mp4": Path("/tmp/asset_c1.mp4"),
"asset_bgm1.mp4": Path("/tmp/asset_bgm1.mp4"),
}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch("video_processing.unified_render_service.run_ffmpeg") as mock_run,
):
layers = svc._group_clips_into_layers(svc._resolve_clips())
result = svc._mix_audio(layers, 5.0)
assert result is not None
mock_run.assert_called_once()
cmd = mock_run.call_args[0][0]
cmd_str = " ".join(cmd)
assert "amix" in cmd_str
assert "volume=0.5" in cmd_str
def test_mix_audio_no_audio_returns_none(self):
"""没有音频素材时返回None。"""
# 构造一个没有音频的场景(比如纯文字)
clips = [_make_clip("t1", "title", order=0, duration=3.0)]
clips[0].asset_id = "" # 无素材
asset_paths: dict[str, Path] = {}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=3.0),
):
# 没有素材的clip会被跳过,layers为空
resolved = svc._resolve_clips()
layers = svc._group_clips_into_layers(resolved)
result = svc._mix_audio(layers, 3.0)
assert result is None
def test_mix_audio_background_not_used_as_main(self):
"""background 图层不参与主音频,main 优先级更高。"""
clips = [
_make_clip("bg1", "background", order=0, duration=5.0),
_make_clip("c1", "main", order=0, duration=5.0),
]
asset_paths = {
"asset_bg1.mp4": Path("/tmp/asset_bg1.mp4"),
"asset_c1.mp4": Path("/tmp/asset_c1.mp4"),
}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch("video_processing.unified_render_service.run_ffmpeg") as mock_run,
):
layers = svc._group_clips_into_layers(svc._resolve_clips())
result = svc._mix_audio(layers, 5.0)
assert result is not None
mock_run.assert_called_once()
cmd = mock_run.call_args[0][0]
# 验证主音频源是 main 的 c1,不是 background 的 bg1
# 单 main clip 走直接提取路径,输入文件应该只有 c1
cmd_str = " ".join(cmd)
assert "asset_c1.mp4" in cmd_str
assert "asset_bg1.mp4" not in cmd_str
def test_mix_audio_main_priority_over_broll(self):
"""main 图层优先级高于 broll。"""
clips = [
_make_clip("b1", "b_roll", order=0, duration=5.0),
_make_clip("c1", "main", order=0, duration=5.0),
]
asset_paths = {
"asset_b1.mp4": Path("/tmp/asset_b1.mp4"),
"asset_c1.mp4": Path("/tmp/asset_c1.mp4"),
}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch("video_processing.unified_render_service.run_ffmpeg") as mock_run,
):
layers = svc._group_clips_into_layers(svc._resolve_clips())
result = svc._mix_audio(layers, 5.0)
assert result is not None
mock_run.assert_called_once()
cmd = mock_run.call_args[0][0]
cmd_str = " ".join(cmd)
# 主音频源应该是 main 的 c1,不是 broll 的 b1
assert "asset_c1.mp4" in cmd_str
assert "asset_b1.mp4" not in cmd_str
def test_mix_audio_broll_used_when_no_main(self):
"""没有 main 时,broll 作为主音频源。"""
clips = [_make_clip("b1", "b_roll", order=0, duration=5.0)]
asset_paths = {"asset_b1.mp4": Path("/tmp/asset_b1.mp4")}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch("video_processing.unified_render_service.run_ffmpeg") as mock_run,
):
layers = svc._group_clips_into_layers(svc._resolve_clips())
result = svc._mix_audio(layers, 5.0)
assert result is not None
mock_run.assert_called_once()
cmd = mock_run.call_args[0][0]
assert "-vn" in cmd
assert "asset_b1.mp4" in " ".join(cmd)
def test_mix_audio_single_clip_truncated_to_video_duration(self):
"""单clip音频截断到 video_durationvideo_duration < clip有效时长)。"""
clips = [_make_clip("c1", "main", order=0, duration=10.0)]
asset_paths = {"asset_c1.mp4": Path("/tmp/asset_c1.mp4")}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=10.0),
patch("video_processing.unified_render_service.run_ffmpeg") as mock_run,
):
layers = svc._group_clips_into_layers(svc._resolve_clips())
# video_duration 只有 3.0,小于 clip 的 10.0
result = svc._mix_audio(layers, 3.0)
assert result is not None
mock_run.assert_called_once()
cmd = mock_run.call_args[0][0]
# 验证 -t 参数是 3.0 不是 10.0
t_index = cmd.index("-t")
assert t_index >= 0
t_value = float(cmd[t_index + 1])
assert t_value == 3.0
def test_merge_audio_video(self):
"""合并音视频命令正确。"""
svc = _make_service([], {})
video_path = Path("/tmp/video.mp4")
audio_path = Path("/tmp/audio.aac")
output_path = Path("/tmp/output.mp4")
with patch("video_processing.unified_render_service.run_ffmpeg") as mock_run:
svc._merge_audio_video(video_path, audio_path, output_path)
mock_run.assert_called_once()
cmd = mock_run.call_args[0][0]
assert "-c:v" in cmd
assert "copy" in cmd
assert "-map" in cmd
assert "-shortest" in cmd
def test_render_calls_audio_mixing(self):
"""多clip完整render流程会调用音频混音(非直通路径)。"""
clips = [
_make_clip("c1", "main", order=0, duration=3.0),
_make_clip("c2", "main", order=1, duration=2.0),
]
asset_paths = {
"asset_c1.mp4": Path("/tmp/asset_c1.mp4"),
"asset_c2.mp4": Path("/tmp/asset_c2.mp4"),
}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch.object(svc, "_execute_ffmpeg"),
patch.object(svc, "_mix_audio", return_value=Path("/tmp/audio.aac")) as mock_mix,
patch.object(svc, "_merge_audio_video") as mock_merge,
patch.object(svc, "_probe_output", return_value=(5.0, 1024, 1280, 720)),
):
result = svc.render()
mock_mix.assert_called_once()
mock_merge.assert_called_once()
assert result.duration == 5.0
def test_render_without_audio_copies_video(self):
"""多clip无音频时走copy路径(非直通路径)。"""
clips = [
_make_clip("c1", "main", order=0, duration=3.0),
_make_clip("c2", "main", order=1, duration=2.0),
]
asset_paths = {
"asset_c1.mp4": Path("/tmp/asset_c1.mp4"),
"asset_c2.mp4": Path("/tmp/asset_c2.mp4"),
}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch.object(svc, "_execute_ffmpeg"),
patch.object(svc, "_mix_audio", return_value=None),
patch("shutil.copy2") as mock_copy,
patch.object(svc, "_probe_output", return_value=(5.0, 1024, 1280, 720)),
):
result = svc.render()
mock_copy.assert_called_once()
assert result.duration == 5.0
def test_render_pass_through_skips_audio_mix(self):
"""直通场景下视频+音频一次完成,跳过 _mix_audio 和 _merge_audio_video。"""
clips = [_make_clip("c1", "main", order=0, duration=5.0)]
asset_paths = {"asset_c1.mp4": Path("/tmp/asset_c1.mp4")}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch.object(svc, "_render_pass_through", return_value=True) as mock_pt,
patch.object(svc, "_mix_audio") as mock_mix,
patch.object(svc, "_merge_audio_video") as mock_merge,
patch("shutil.copy2") as mock_copy,
patch.object(svc, "_probe_output", return_value=(5.0, 1024, 1280, 720)),
):
result = svc.render()
# 直通场景调用了 _render_pass_through,跳过了 _mix_audio / _merge / copy
mock_pt.assert_called_once()
mock_mix.assert_not_called()
mock_merge.assert_not_called()
mock_copy.assert_not_called()
assert result.duration == 5.0
def test_pass_through_main_has_aac_audio(self):
"""直通main/broll场景输出带aac音频。"""
clips = [_make_clip("c1", "main", order=0, duration=5.0)]
asset_paths = {"asset_c1.mp4": Path("/tmp/asset_c1.mp4")}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch("video_processing.unified_render_service.run_ffmpeg") as mock_run,
):
layers = svc._group_clips_into_layers(svc._resolve_clips())
result = svc._render_pass_through(layers, Path("/tmp/out.mp4"), video_duration=5.0)
assert result is True # main 类型返回有音频
mock_run.assert_called_once()
cmd = mock_run.call_args[0][0]
assert "-an" not in cmd # 不再是无声
assert "aac" in cmd # 有aac音频编码
assert "-b:a" in cmd
def test_pass_through_background_no_audio(self):
"""直通background场景不带音频(图片素材)。"""
clips = [_make_clip("bg1", "background", order=0, duration=5.0)]
asset_paths = {"asset_bg1.mp4": Path("/tmp/asset_bg1.mp4")}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch("video_processing.unified_render_service.run_ffmpeg") as mock_run,
):
layers = svc._group_clips_into_layers(svc._resolve_clips())
result = svc._render_pass_through(layers, Path("/tmp/out.mp4"))
assert result is False # background 返回无音频
mock_run.assert_called_once()
cmd = mock_run.call_args[0][0]
assert "aac" not in cmd # 没有音频编码参数
# ── 无音轨视频防御测试 ──
def test_mix_audio_main_no_audio_stream_returns_none(self):
"""主图层clip无音频流且无独立音频轨时,返回None(不报错)。"""
clips = [_make_clip("c1", "main", order=0, duration=5.0)]
asset_paths = {"asset_c1.mp4": Path("/tmp/asset_c1.mp4")}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch("video_processing.ffmpeg_utils.probe_has_audio", return_value=False),
patch("video_processing.unified_render_service.run_ffmpeg") as mock_run,
):
layers = svc._group_clips_into_layers(svc._resolve_clips())
result = svc._mix_audio(layers, 5.0)
assert result is None
# 没有音频流时不应调用 FFmpeg
mock_run.assert_not_called()
def test_mix_audio_partial_clips_no_audio_filtered(self):
"""部分主图层clip无音频流时,过滤掉无音轨的,剩余有音频的正常concat。"""
clips = [
_make_clip("c1", "main", order=0, duration=3.0), # 无音频
_make_clip("c2", "main", order=1, duration=2.0), # 有音频
]
asset_paths = {
"asset_c1.mp4": Path("/tmp/asset_c1.mp4"),
"asset_c2.mp4": Path("/tmp/asset_c2.mp4"),
}
svc = _make_service(clips, asset_paths)
# 模拟:c1 无音频,c2 有音频
def fake_has_audio(path):
return "c2" in str(path)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch("video_processing.ffmpeg_utils.probe_has_audio", side_effect=fake_has_audio),
patch("video_processing.unified_render_service.run_ffmpeg") as mock_run,
):
layers = svc._group_clips_into_layers(svc._resolve_clips())
result = svc._mix_audio(layers, 5.0)
assert result is not None
mock_run.assert_called_once()
cmd = mock_run.call_args[0][0]
cmd_str = " ".join(cmd)
# 只剩 1 个有效音频 clip,走单clip路径(-vn),不走 filter_complex concat
assert "-vn" in cmd
assert "concat=n=2" not in cmd_str
def test_mix_audio_all_main_no_audio_but_independent_track(self):
"""主图层全部无音频,但有独立音频轨时,正常走amix混音。"""
clips = [
_make_clip("c1", "main", order=0, duration=5.0), # 无音频
_make_clip(
"bgm1",
"main",
order=0,
duration=5.0,
config={"role": "audio", "volume": 0.5},
), # 独立音频轨(有音频)
]
asset_paths = {
"asset_c1.mp4": Path("/tmp/asset_c1.mp4"),
"asset_bgm1.mp4": Path("/tmp/asset_bgm1.mp4"),
}
svc = _make_service(clips, asset_paths)
def fake_has_audio(path):
return "bgm" in str(path)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch("video_processing.ffmpeg_utils.probe_has_audio", side_effect=fake_has_audio),
patch("video_processing.unified_render_service.run_ffmpeg") as mock_run,
):
layers = svc._group_clips_into_layers(svc._resolve_clips())
result = svc._mix_audio(layers, 5.0)
assert result is not None
mock_run.assert_called_once()
cmd = mock_run.call_args[0][0]
cmd_str = " ".join(cmd)
# 只有独立音频轨参与混音,amix 输入数=1
assert "amix=inputs=1" in cmd_str
def test_mix_audio_both_no_audio_returns_none(self):
"""主图层和独立音频轨都无音频时,返回None。"""
clips = [
_make_clip("c1", "main", order=0, duration=5.0),
_make_clip(
"bgm1",
"main",
order=0,
duration=5.0,
config={"role": "audio"},
),
]
asset_paths = {
"asset_c1.mp4": Path("/tmp/asset_c1.mp4"),
"asset_bgm1.mp4": Path("/tmp/asset_bgm1.mp4"),
}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch("video_processing.ffmpeg_utils.probe_has_audio", return_value=False),
patch("video_processing.unified_render_service.run_ffmpeg") as mock_run,
):
layers = svc._group_clips_into_layers(svc._resolve_clips())
result = svc._mix_audio(layers, 5.0)
assert result is None
mock_run.assert_not_called()
def test_clip_has_audio_cache(self):
"""_clip_has_audio 带缓存,同一clip只探测一次。"""
clips = [_make_clip("c1", "main", order=0, duration=5.0)]
asset_paths = {"asset_c1.mp4": Path("/tmp/asset_c1.mp4")}
svc = _make_service(clips, asset_paths)
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
):
resolved = svc._resolve_clips()
clip = resolved[0]
with patch("video_processing.ffmpeg_utils.probe_has_audio", return_value=True) as mock_probe:
# 调用 3 次
r1 = svc._clip_has_audio(clip)
r2 = svc._clip_has_audio(clip)
r3 = svc._clip_has_audio(clip)
assert r1 is True and r2 is True and r3 is True
# 实际只探测了 1 次
assert mock_probe.call_count == 1
# ── 测试 stream copy 流拷贝优化 ───────────────────────────────────────────────
class TestStreamCopy:
"""stream copy 流拷贝优化测试。"""
def _make_single_clip_service(self):
clips = [_make_clip("c1", "main", order=0, duration=5.0)]
svc = _make_service(clips)
with _patch_path_exists(), patch("video_processing.unified_render_service.probe_duration", return_value=5.0):
resolved = svc._resolve_clips()
layers = svc._group_clips_into_layers(resolved)
return svc, resolved[0], layers
def test_can_use_stream_copy_all_conditions_met(self):
"""所有条件满足 → 可以 stream copy。"""
svc, clip, layers = self._make_single_clip_service()
probe_result = {
"width": 1280,
"height": 720,
"fps": 25.0,
"video_codec": "h264",
"pix_fmt": "yuv420p",
"duration": 5.0,
"has_audio": True,
"audio_codec": "aac",
}
with patch(
"video_processing.unified_render_service.probe_video_info",
return_value=probe_result,
):
can_copy, reason = svc._can_use_stream_copy(clip, ass_path=None, video_duration=0)
assert can_copy is True
assert "所有条件满足" in reason
def test_cannot_copy_with_subtitles(self):
"""有字幕 → 不能 stream copy。"""
svc, clip, layers = self._make_single_clip_service()
can_copy, reason = svc._can_use_stream_copy(clip, ass_path=Path("/tmp/sub.ass"), video_duration=0)
assert can_copy is False
assert "字幕" in reason
def test_cannot_copy_wrong_codec(self):
"""编码不是 h264 → 不能 stream copy。"""
svc, clip, layers = self._make_single_clip_service()
probe_result = {
"width": 1280,
"height": 720,
"fps": 25.0,
"video_codec": "hevc",
"pix_fmt": "yuv420p",
"duration": 5.0,
"has_audio": True,
"audio_codec": "aac",
}
with patch(
"video_processing.unified_render_service.probe_video_info",
return_value=probe_result,
):
can_copy, reason = svc._can_use_stream_copy(clip, ass_path=None, video_duration=0)
assert can_copy is False
assert "编码" in reason
def test_cannot_copy_wrong_resolution(self):
"""分辨率不匹配 → 不能 stream copy。"""
svc, clip, layers = self._make_single_clip_service()
probe_result = {
"width": 1920,
"height": 1080,
"fps": 25.0,
"video_codec": "h264",
"pix_fmt": "yuv420p",
"duration": 5.0,
"has_audio": True,
"audio_codec": "aac",
}
with patch(
"video_processing.unified_render_service.probe_video_info",
return_value=probe_result,
):
can_copy, reason = svc._can_use_stream_copy(clip, ass_path=None, video_duration=0)
assert can_copy is False
assert "分辨率" in reason
def test_cannot_copy_wrong_fps(self):
"""帧率不匹配 → 不能 stream copy。"""
svc, clip, layers = self._make_single_clip_service()
probe_result = {
"width": 1280,
"height": 720,
"fps": 30.0,
"video_codec": "h264",
"pix_fmt": "yuv420p",
"duration": 5.0,
"has_audio": True,
"audio_codec": "aac",
}
with patch(
"video_processing.unified_render_service.probe_video_info",
return_value=probe_result,
):
can_copy, reason = svc._can_use_stream_copy(clip, ass_path=None, video_duration=0)
assert can_copy is False
assert "帧率" in reason
def test_cannot_copy_wrong_pix_fmt(self):
"""像素格式不匹配 → 不能 stream copy。"""
svc, clip, layers = self._make_single_clip_service()
probe_result = {
"width": 1280,
"height": 720,
"fps": 25.0,
"video_codec": "h264",
"pix_fmt": "yuv422p",
"duration": 5.0,
"has_audio": True,
"audio_codec": "aac",
}
with patch(
"video_processing.unified_render_service.probe_video_info",
return_value=probe_result,
):
can_copy, reason = svc._can_use_stream_copy(clip, ass_path=None, video_duration=0)
assert can_copy is False
assert "像素格式" in reason
def test_try_render_stream_copy_success(self):
"""stream copy 渲染成功 → 返回 True。"""
svc, clip, layers = self._make_single_clip_service()
probe_result = {
"width": 1280,
"height": 720,
"fps": 25.0,
"video_codec": "h264",
"pix_fmt": "yuv420p",
"duration": 5.0,
"has_audio": True,
"audio_codec": "aac",
}
output_path = Path("/tmp/test_output.mp4")
def fake_stat():
m = MagicMock()
m.st_size = 1024000
return m
with (
patch(
"video_processing.unified_render_service.probe_video_info",
return_value=probe_result,
),
patch("video_processing.unified_render_service.run_ffmpeg") as mock_run,
patch("pathlib.Path.exists", return_value=True),
patch("pathlib.Path.stat", side_effect=fake_stat),
):
result = svc._try_render_stream_copy(layers, output_path, ass_path=None, video_duration=0)
assert result is True
mock_run.assert_called_once()
cmd = mock_run.call_args[0][0]
assert "-c:v" in cmd
assert "copy" in cmd
assert "-c:a" in cmd
def test_try_render_stream_copy_fallback_on_ffmpeg_error(self):
"""stream copy FFmpeg 失败 → 返回 False(调用方回退到重编码)。"""
svc, clip, layers = self._make_single_clip_service()
probe_result = {
"width": 1280,
"height": 720,
"fps": 25.0,
"video_codec": "h264",
"pix_fmt": "yuv420p",
"duration": 5.0,
"has_audio": True,
"audio_codec": "aac",
}
output_path = Path("/tmp/test_output.mp4")
import subprocess as sp
with (
patch(
"video_processing.unified_render_service.probe_video_info",
return_value=probe_result,
),
patch(
"video_processing.unified_render_service.run_ffmpeg",
side_effect=sp.CalledProcessError(1, ["ffmpeg"], stderr="copy failed"),
),
patch("pathlib.Path.exists", return_value=False),
):
result = svc._try_render_stream_copy(layers, output_path, ass_path=None, video_duration=0)
assert result is False
def test_render_uses_stream_copy_when_eligible(self):
"""完整渲染流程:满足条件时走 stream copy。"""
clips = [_make_clip("c1", "main", order=0, duration=5.0)]
svc = _make_service(clips)
probe_result = {
"width": 1280,
"height": 720,
"fps": 25.0,
"video_codec": "h264",
"pix_fmt": "yuv420p",
"duration": 5.0,
"has_audio": True,
"audio_codec": "aac",
}
def fake_stat():
m = MagicMock()
m.st_size = 1024000
return m
with (
_patch_path_exists(),
patch("video_processing.unified_render_service.probe_duration", return_value=5.0),
patch(
"video_processing.unified_render_service.probe_video_info",
return_value=probe_result,
),
patch("video_processing.unified_render_service.run_ffmpeg") as mock_run,
patch("pathlib.Path.stat", side_effect=fake_stat),
patch("shutil.copy2"),
):
result = svc.render()
assert mock_run.call_count == 1
cmd = mock_run.call_args[0][0]
assert "copy" in cmd
assert isinstance(result.output_path, Path)