Compare commits
34 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| c845ceb6ca | |||
| f92721e689 | |||
| 47699a2e93 | |||
| 73375c6639 | |||
| cc47c9f90f | |||
| c1e466f9c1 | |||
| 8598638e8f | |||
| c48ddeef7d | |||
| b87d7b763e | |||
| 9c6c477f55 | |||
| c2ebe9d254 | |||
| 1d06d2ddd2 | |||
| 5e704094f6 | |||
| ffd99ffeb0 | |||
| 1b2bccee6f | |||
| a0cac1b75d | |||
| bbe831f9e0 | |||
| 4c5ab7f80e | |||
| 9b2e782abd | |||
| 788559ff29 | |||
| bdf99bba39 | |||
| bfe8bfe2da | |||
| fdcf48103e | |||
| 9e87ac05c6 | |||
| 8427bb6852 | |||
| ad86f5bc79 | |||
| e39f8bacdd | |||
| 8883581e34 | |||
| c80a6935fd | |||
| a6ae041944 | |||
| 082c9f6f09 | |||
| da4b95a40a | |||
| 7c541910b3 | |||
| e5d627fc3e |
+4
-1
@@ -3,6 +3,7 @@
|
||||
# ==================== 应用配置 ====================
|
||||
APP_NAME=小虾 SaaS
|
||||
APP_BASE_URL=http://localhost:3000
|
||||
APP_ENV=development
|
||||
|
||||
# ==================== 数据库配置 ====================
|
||||
DATABASE_URL=postgresql://xiaoxia_user:your_password@localhost:5432/xiaoxia_saas
|
||||
@@ -35,7 +36,8 @@ ENVIRONMENT=development
|
||||
DEBUG=true
|
||||
|
||||
# ==================== CORS 配置 ====================
|
||||
CORS_ORIGINS=["http://localhost:3000","http://localhost:5173"]
|
||||
# 逗号分隔的域名列表(Settings 读取 CORS_ORIGINS_RAW)
|
||||
CORS_ORIGINS_RAW=http://localhost:3000,http://localhost:5173
|
||||
|
||||
# ==================== 阿里云 OSS 配置 ====================
|
||||
OSS_ENDPOINT=oss-cn-hangzhou.aliyuncs.com
|
||||
@@ -49,6 +51,7 @@ 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
|
||||
|
||||
+249
-78
File diff suppressed because one or more lines are too long
@@ -1 +0,0 @@
|
||||
"""API application package."""
|
||||
@@ -1 +0,0 @@
|
||||
"""API package."""
|
||||
@@ -8,10 +8,12 @@ from app.api.routes.dashboard import router as dashboard_router
|
||||
from app.api.routes.duplication import router as duplication_router
|
||||
from app.api.routes.edit_plans import router as edit_plans_router
|
||||
from app.api.routes.edit_templates import router as edit_templates_router
|
||||
from app.api.routes.feature_flags import router as feature_flags_router
|
||||
from app.api.routes.generated_videos import router as generated_videos_router
|
||||
from app.api.routes.generation_tasks import router as generation_tasks_router
|
||||
from app.api.routes.health import router as health_check_router
|
||||
from app.api.routes.ingest_jobs import router as ingest_jobs_router
|
||||
from app.api.routes.internal_render import router as internal_render_router
|
||||
from app.api.routes.jobs import router as jobs_router
|
||||
from app.api.routes.projects import router as projects_router
|
||||
from app.api.routes.recipes import router as recipes_router
|
||||
@@ -151,3 +153,11 @@ api_router.include_router(
|
||||
prefix="/tts",
|
||||
tags=["TTS"],
|
||||
)
|
||||
api_router.include_router(
|
||||
feature_flags_router,
|
||||
tags=["Internal"],
|
||||
)
|
||||
api_router.include_router(
|
||||
internal_render_router,
|
||||
tags=["Internal"],
|
||||
)
|
||||
|
||||
@@ -0,0 +1,48 @@
|
||||
"""路由层共享辅助函数 — 消除跨文件重复定义。"""
|
||||
|
||||
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")
|
||||
@@ -22,18 +22,11 @@ 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,
|
||||
@@ -168,7 +161,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)
|
||||
|
||||
@@ -27,6 +27,8 @@ 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()
|
||||
@@ -72,14 +74,6 @@ 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(
|
||||
@@ -136,7 +130,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)
|
||||
@@ -152,7 +146,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)
|
||||
@@ -210,13 +204,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:
|
||||
@@ -262,7 +256,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)
|
||||
@@ -286,7 +280,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)
|
||||
@@ -307,7 +301,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)
|
||||
|
||||
|
||||
@@ -322,7 +316,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:
|
||||
@@ -346,7 +340,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)
|
||||
|
||||
|
||||
@@ -363,7 +357,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:
|
||||
@@ -387,7 +381,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)
|
||||
|
||||
|
||||
@@ -206,7 +206,7 @@ async def verify_email_post(
|
||||
return _verify_email_token(request.token, user_repository)
|
||||
|
||||
|
||||
@router.post("/password/forgot", response_model=MessageResponse, status_code=status.HTTP_202_ACCEPTED)
|
||||
@router.post("/forgot-password", 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("/password/reset", response_model=MessageResponse)
|
||||
@router.post("/reset-password", response_model=MessageResponse)
|
||||
async def reset_password(
|
||||
request: ResetPasswordModel,
|
||||
user_repository: UserRepository = Depends(get_user_repository),
|
||||
@@ -243,7 +243,6 @@ async def logout(
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""登出 - 将当前 token 加入黑名单"""
|
||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||
|
||||
if credentials:
|
||||
try:
|
||||
|
||||
@@ -14,7 +14,6 @@ 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 (
|
||||
@@ -35,6 +34,8 @@ 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__)
|
||||
|
||||
@@ -113,22 +114,6 @@ 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)
|
||||
@@ -206,7 +191,6 @@ 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:
|
||||
@@ -221,7 +205,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,
|
||||
|
||||
@@ -32,12 +32,6 @@ 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,
|
||||
)
|
||||
@@ -51,6 +45,8 @@ 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
|
||||
|
||||
@@ -247,17 +243,6 @@ 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,
|
||||
@@ -311,7 +296,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(
|
||||
@@ -353,7 +338,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)
|
||||
|
||||
|
||||
@@ -369,7 +354,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)
|
||||
@@ -411,7 +396,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:
|
||||
@@ -465,7 +450,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(
|
||||
@@ -507,7 +492,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:
|
||||
@@ -710,7 +695,7 @@ def generate_plan(
|
||||
except HTTPException:
|
||||
# 已处理的 HTTP 异常直接透传
|
||||
raise
|
||||
except Exception as exc:
|
||||
except Exception:
|
||||
logger.exception("触发剪辑计划生成失败: plan_id=%s", plan_id)
|
||||
# 尝试将计划标记为失败(RENDERING → FAILED 是合法的状态流转)
|
||||
try:
|
||||
@@ -749,7 +734,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 = [
|
||||
@@ -791,7 +776,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)
|
||||
@@ -860,7 +845,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
|
||||
@@ -909,7 +894,7 @@ def ai_recommend_clips(
|
||||
config=normalized_config,
|
||||
total_duration=result["total_duration"],
|
||||
)
|
||||
except Exception as exc:
|
||||
except Exception:
|
||||
logger.exception("AI 推荐写入失败,plan_id=%s 数据可能不一致", plan_id)
|
||||
# 尝试回滚未提交的变更
|
||||
try:
|
||||
@@ -990,7 +975,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
|
||||
@@ -1110,7 +1095,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 排序
|
||||
@@ -1171,7 +1156,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)
|
||||
|
||||
|
||||
Executable
+195
@@ -0,0 +1,195 @@
|
||||
"""Feature Flag 内部管理接口。
|
||||
|
||||
通过内部 API Key 鉴权,支持查看和修改 Feature Flag 配置。
|
||||
主要用于灰度发布期间的动态开关控制。
|
||||
|
||||
API:
|
||||
GET /api/v1/internal/feature-flags - 列出所有 flag
|
||||
GET /api/v1/internal/feature-flags/{name} - 查看单个 flag
|
||||
PUT /api/v1/internal/feature-flags/{name} - 设置 flag 配置
|
||||
DELETE /api/v1/internal/feature-flags/{name} - 删除 flag
|
||||
|
||||
鉴权:X-API-Key header,走内部 API Key 验证
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Optional
|
||||
|
||||
from app.api.routes.auth import _verify_internal_api_key
|
||||
from app.config import settings
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from packages.adapters.redis.feature_flag_store import (
|
||||
FEATURE_FLAG_REDIS_PREFIX,
|
||||
FeatureFlagConfig,
|
||||
RedisFeatureFlagStore,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/internal/feature-flags", tags=["Internal"])
|
||||
|
||||
# 允许管理的 flag 白名单(防止误操作其他系统 flag)
|
||||
ALLOWED_FLAGS = {
|
||||
"render_engine",
|
||||
}
|
||||
|
||||
|
||||
def _get_feature_flag_store() -> RedisFeatureFlagStore:
|
||||
"""获取 Feature Flag 存储实例。"""
|
||||
return RedisFeatureFlagStore(redis_url=settings.REDIS_URL)
|
||||
|
||||
|
||||
class FeatureFlagUpdateRequest(BaseModel):
|
||||
"""Feature Flag 更新请求体。"""
|
||||
|
||||
enabled: bool = Field(..., description="是否启用")
|
||||
percentage: int = Field(0, ge=0, le=100, description="灰度百分比 (0-100)")
|
||||
whitelist: list[str] = Field(default_factory=list, description="白名单列表(如 user_id)")
|
||||
|
||||
|
||||
class FeatureFlagResponse(BaseModel):
|
||||
"""Feature Flag 响应。"""
|
||||
|
||||
name: str
|
||||
enabled: bool
|
||||
percentage: int
|
||||
whitelist: list[str]
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config: FeatureFlagConfig) -> "FeatureFlagResponse":
|
||||
return cls(
|
||||
name=config.name,
|
||||
enabled=config.enabled,
|
||||
percentage=config.percentage,
|
||||
whitelist=sorted(config.whitelist),
|
||||
)
|
||||
|
||||
|
||||
class FeatureFlagCheckResponse(BaseModel):
|
||||
"""Flag 激活检查响应。"""
|
||||
|
||||
name: str
|
||||
active: bool
|
||||
identifier: Optional[str] = None
|
||||
|
||||
|
||||
def _validate_flag_name(name: str) -> None:
|
||||
"""校验 flag 名称是否在允许列表中。"""
|
||||
if name not in ALLOWED_FLAGS:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail=f"Unsupported flag: {name}. Allowed: {sorted(ALLOWED_FLAGS)}",
|
||||
)
|
||||
|
||||
|
||||
@router.get("", response_model=list[FeatureFlagResponse])
|
||||
async def list_feature_flags(
|
||||
_: bool = Depends(_verify_internal_api_key),
|
||||
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
|
||||
):
|
||||
"""列出所有 Feature Flag。"""
|
||||
try:
|
||||
flags = store.list_all()
|
||||
# 同时返回预定义的 flag(即使未设置也显示默认值)
|
||||
result = []
|
||||
for name in sorted(ALLOWED_FLAGS):
|
||||
config = flags.get(name) or FeatureFlagConfig(name=name, enabled=False)
|
||||
result.append(FeatureFlagResponse.from_config(config))
|
||||
# 加上已存在但不在白名单中的 flag(只读展示)
|
||||
for name, config in flags.items():
|
||||
if name not in ALLOWED_FLAGS:
|
||||
result.append(FeatureFlagResponse.from_config(config))
|
||||
return sorted(result, key=lambda x: x.name)
|
||||
except Exception as exc:
|
||||
logger.error("Failed to list feature flags: %s", exc)
|
||||
raise HTTPException(status_code=500, detail=f"Failed to list flags: {exc}")
|
||||
|
||||
|
||||
@router.get("/{name}", response_model=FeatureFlagResponse)
|
||||
async def get_feature_flag(
|
||||
name: str,
|
||||
_: bool = Depends(_verify_internal_api_key),
|
||||
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
|
||||
):
|
||||
"""获取单个 Feature Flag 配置。"""
|
||||
try:
|
||||
config = store.get(name)
|
||||
return FeatureFlagResponse.from_config(config)
|
||||
except Exception as exc:
|
||||
logger.error("Failed to get feature flag %s: %s", name, exc)
|
||||
raise HTTPException(status_code=500, detail=f"Failed to get flag: {exc}")
|
||||
|
||||
|
||||
@router.get("/{name}/check", response_model=FeatureFlagCheckResponse)
|
||||
async def check_feature_flag(
|
||||
name: str,
|
||||
identifier: Optional[str] = Query(None, description="标识符,如 user_id"),
|
||||
_: bool = Depends(_verify_internal_api_key),
|
||||
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
|
||||
):
|
||||
"""检查某个标识符是否命中 Feature Flag。"""
|
||||
try:
|
||||
active = store.is_active(name, identifier=identifier)
|
||||
return FeatureFlagCheckResponse(name=name, active=active, identifier=identifier)
|
||||
except Exception as exc:
|
||||
logger.error("Failed to check feature flag %s: %s", name, exc)
|
||||
raise HTTPException(status_code=500, detail=f"Failed to check flag: {exc}")
|
||||
|
||||
|
||||
@router.put("/{name}", response_model=FeatureFlagResponse)
|
||||
async def update_feature_flag(
|
||||
name: str,
|
||||
request: FeatureFlagUpdateRequest,
|
||||
_: bool = Depends(_verify_internal_api_key),
|
||||
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
|
||||
):
|
||||
"""更新 Feature Flag 配置。
|
||||
|
||||
只允许修改 ALLOWED_FLAGS 列表中的 flag。
|
||||
"""
|
||||
_validate_flag_name(name)
|
||||
|
||||
try:
|
||||
config = FeatureFlagConfig(
|
||||
name=name,
|
||||
enabled=request.enabled,
|
||||
percentage=request.percentage,
|
||||
whitelist=set(request.whitelist),
|
||||
)
|
||||
store.set(config)
|
||||
logger.info(
|
||||
"Feature flag updated: name=%s enabled=%s percentage=%d whitelist=%d",
|
||||
name,
|
||||
config.enabled,
|
||||
config.percentage,
|
||||
len(config.whitelist),
|
||||
)
|
||||
return FeatureFlagResponse.from_config(config)
|
||||
except Exception as exc:
|
||||
logger.error("Failed to update feature flag %s: %s", name, exc)
|
||||
raise HTTPException(status_code=500, detail=f"Failed to update flag: {exc}")
|
||||
|
||||
|
||||
@router.delete("/{name}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
async def delete_feature_flag(
|
||||
name: str,
|
||||
_: bool = Depends(_verify_internal_api_key),
|
||||
store: RedisFeatureFlagStore = Depends(_get_feature_flag_store),
|
||||
):
|
||||
"""删除 Feature Flag。
|
||||
|
||||
只允许删除 ALLOWED_FLAGS 列表中的 flag。
|
||||
"""
|
||||
_validate_flag_name(name)
|
||||
|
||||
try:
|
||||
deleted = store.delete(name)
|
||||
logger.info("Feature flag deleted: name=%s deleted=%s", name, deleted)
|
||||
return None
|
||||
except Exception as exc:
|
||||
logger.error("Failed to delete feature flag %s: %s", name, exc)
|
||||
raise HTTPException(status_code=500, detail=f"Failed to delete flag: {exc}")
|
||||
@@ -32,6 +32,8 @@ 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,
|
||||
@@ -43,15 +45,6 @@ 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,
|
||||
@@ -339,7 +332,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)
|
||||
|
||||
|
||||
@@ -356,7 +349,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 = []
|
||||
|
||||
Executable
+120
@@ -0,0 +1,120 @@
|
||||
"""渲染结果内部下载接口。
|
||||
|
||||
通过内部 API Key 鉴权,为灰度对比工具等内部系统提供渲染结果下载能力。
|
||||
|
||||
API:
|
||||
GET /api/v1/internal/render/videos/{video_id}/download-url - 获取单个视频下载URL
|
||||
GET /api/v1/internal/render/tasks/{task_id}/videos - 获取任务下所有视频及下载URL
|
||||
|
||||
鉴权:X-API-Key header,走内部 API Key 验证
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from app.api.routes.auth import _verify_internal_api_key
|
||||
from app.core.storage import OSSStorageService, get_storage_service
|
||||
from app.dependencies import get_generated_video_repository
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
from pydantic import BaseModel
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
router = APIRouter(prefix="/internal/render", tags=["Internal"])
|
||||
|
||||
|
||||
class InternalRenderVideoItem(BaseModel):
|
||||
"""内部渲染视频项。"""
|
||||
|
||||
video_id: str
|
||||
generation_task_id: str
|
||||
project_id: str
|
||||
name: str
|
||||
file_url: str
|
||||
file_size: int | None = None
|
||||
duration: float | None = None
|
||||
width: int | None = None
|
||||
height: int | None = None
|
||||
fps: float | None = None
|
||||
status: str
|
||||
download_url: str
|
||||
|
||||
|
||||
class InternalRenderTaskVideosResponse(BaseModel):
|
||||
"""任务下所有渲染视频响应。"""
|
||||
|
||||
task_id: str
|
||||
count: int
|
||||
videos: list[InternalRenderVideoItem]
|
||||
|
||||
|
||||
class InternalRenderDownloadUrlResponse(BaseModel):
|
||||
"""单个视频下载URL响应。"""
|
||||
|
||||
video_id: str
|
||||
download_url: str
|
||||
|
||||
|
||||
def _video_to_item(video: Any, download_url: str) -> InternalRenderVideoItem:
|
||||
"""将 GeneratedVideo 领域对象转为响应项。"""
|
||||
return InternalRenderVideoItem(
|
||||
video_id=video.id,
|
||||
generation_task_id=video.generation_task_id,
|
||||
project_id=video.project_id,
|
||||
name=video.name,
|
||||
file_url=video.file_url,
|
||||
file_size=getattr(video, "file_size", None),
|
||||
duration=getattr(video, "duration", None),
|
||||
width=getattr(video, "width", None),
|
||||
height=getattr(video, "height", None),
|
||||
fps=getattr(video, "fps", None),
|
||||
status=video.status,
|
||||
download_url=download_url,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/videos/{video_id}/download-url", response_model=InternalRenderDownloadUrlResponse)
|
||||
def get_render_video_download_url(
|
||||
video_id: str,
|
||||
_: bool = Depends(_verify_internal_api_key),
|
||||
generated_video_repository: Any = Depends(get_generated_video_repository),
|
||||
storage_service: OSSStorageService = Depends(get_storage_service),
|
||||
) -> InternalRenderDownloadUrlResponse:
|
||||
"""获取单个渲染视频的下载URL(预签名)。"""
|
||||
video = generated_video_repository.get(video_id)
|
||||
if video is None:
|
||||
raise HTTPException(status_code=404, detail=f"GeneratedVideo {video_id} not found")
|
||||
|
||||
download_url = storage_service.get_download_url(video.file_url, expires_seconds=86400)
|
||||
logger.info("内部渲染下载URL生成: video_id=%s", video_id)
|
||||
return InternalRenderDownloadUrlResponse(video_id=video_id, download_url=download_url)
|
||||
|
||||
|
||||
@router.get("/tasks/{task_id}/videos", response_model=InternalRenderTaskVideosResponse)
|
||||
def get_render_task_videos(
|
||||
task_id: str,
|
||||
status: str | None = Query(None, description="按状态筛选,如 completed/failed"),
|
||||
_: bool = Depends(_verify_internal_api_key),
|
||||
generated_video_repository: Any = Depends(get_generated_video_repository),
|
||||
storage_service: OSSStorageService = Depends(get_storage_service),
|
||||
) -> InternalRenderTaskVideosResponse:
|
||||
"""获取生成任务下所有渲染视频及下载URL。"""
|
||||
videos = generated_video_repository.list_by_generation_task(task_id)
|
||||
|
||||
# 状态筛选
|
||||
if status:
|
||||
videos = [v for v in videos if v.status == status]
|
||||
|
||||
items = []
|
||||
for video in videos:
|
||||
download_url = storage_service.get_download_url(video.file_url, expires_seconds=86400)
|
||||
items.append(_video_to_item(video, download_url))
|
||||
|
||||
logger.info("内部渲染任务视频查询: task_id=%s count=%d", task_id, len(items))
|
||||
return InternalRenderTaskVideosResponse(
|
||||
task_id=task_id,
|
||||
count=len(items),
|
||||
videos=items,
|
||||
)
|
||||
@@ -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_db_session, get_job_repository, get_project_repository
|
||||
from app.dependencies import get_job_repository, get_project_repository
|
||||
from app.schemas.job import (
|
||||
CompleteJobRequest,
|
||||
CreateJobRequest,
|
||||
@@ -51,6 +51,8 @@ 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()
|
||||
@@ -66,15 +68,6 @@ _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")
|
||||
|
||||
|
||||
# ── 创建任务 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -89,7 +82,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:
|
||||
@@ -182,7 +175,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(
|
||||
@@ -204,7 +197,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)
|
||||
|
||||
@@ -33,6 +33,8 @@ 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()
|
||||
|
||||
|
||||
@@ -40,13 +42,6 @@ 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,
|
||||
@@ -194,7 +189,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)
|
||||
|
||||
@@ -232,7 +232,7 @@ async def payment_callback(
|
||||
|
||||
# 创建账单记录
|
||||
record_id = uuid.uuid4().hex
|
||||
record = repo.create(
|
||||
repo.create(
|
||||
{
|
||||
"id": record_id,
|
||||
"user_id": user_id,
|
||||
|
||||
@@ -28,6 +28,8 @@ 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()
|
||||
|
||||
|
||||
@@ -51,13 +53,6 @@ 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),
|
||||
@@ -98,7 +93,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,
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
import logging
|
||||
from typing import Annotated, Any
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
@@ -17,12 +17,13 @@ from app.schemas.upload import (
|
||||
DirectUploadCompleteResponse,
|
||||
DirectUploadPrepareRequest,
|
||||
DirectUploadPrepareResponse,
|
||||
UploadAssetRequest,
|
||||
UploadAssetResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, File, Form, HTTPException, UploadFile, status
|
||||
|
||||
from packages.application import GetProjectUseCase, SubmitIngestJobCommand, SubmitIngestJobUseCase
|
||||
from packages.application import SubmitIngestJobCommand, SubmitIngestJobUseCase
|
||||
|
||||
from app.api.routes._helpers import require_project_and_library
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -80,21 +81,6 @@ 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,
|
||||
@@ -135,7 +121,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,
|
||||
@@ -183,7 +169,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,
|
||||
@@ -252,7 +238,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:
|
||||
|
||||
@@ -28,7 +28,6 @@ from packages.application.voice_clone.use_cases import (
|
||||
VoiceCloneNotRetryableError,
|
||||
)
|
||||
from packages.application.voice_clone.workflow import (
|
||||
VoiceCloneWorkflowError,
|
||||
VoiceCloneWorkflowService,
|
||||
)
|
||||
|
||||
|
||||
@@ -39,6 +39,8 @@ 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()
|
||||
|
||||
|
||||
@@ -125,13 +127,6 @@ 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"
|
||||
|
||||
|
||||
# ==================== 统一配音列表(预置 + 克隆)====================
|
||||
|
||||
|
||||
@@ -271,7 +266,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,
|
||||
|
||||
@@ -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_RECYLE: int = 3600
|
||||
DATABASE_POOL_RECYCLE: int = 3600
|
||||
USE_IN_MEMORY_DB: bool = False
|
||||
AUTO_CREATE_SCHEMA: bool = False
|
||||
|
||||
@@ -41,6 +41,11 @@ 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):
|
||||
@@ -109,6 +114,9 @@ class Settings(BaseSettings):
|
||||
LOG_LEVEL: str = "INFO"
|
||||
CORS_ORIGINS_RAW: str = "http://localhost:3000,http://localhost:5173,http://localhost:8000"
|
||||
|
||||
# 渲染引擎选择:legacy=旧VideoComposeService,unified=新UnifiedRenderService
|
||||
RENDER_ENGINE: str = "legacy"
|
||||
|
||||
model_config = SettingsConfigDict(
|
||||
env_file=".env",
|
||||
env_file_encoding="utf-8",
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
"""Core configuration package."""
|
||||
@@ -50,20 +50,8 @@ 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)
|
||||
|
||||
|
||||
@@ -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
|
||||
from fastapi import Depends, HTTPException
|
||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||
|
||||
from packages.domain.entities import User
|
||||
|
||||
@@ -6,7 +6,7 @@ import logging
|
||||
import time
|
||||
from typing import Callable
|
||||
|
||||
from fastapi import Request, Response
|
||||
from fastapi import Request
|
||||
from starlette.middleware.base import BaseHTTPMiddleware
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import List, Optional
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
@@ -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, EditPlanClipStatus
|
||||
from packages.domain.edit_plan_clip import EditPlanClip
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -18,7 +18,6 @@ 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__)
|
||||
|
||||
|
||||
@@ -13,7 +13,6 @@ 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, Optional
|
||||
from typing import Any, List
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
|
||||
@@ -16,7 +16,6 @@ FFmpeg 视频合成编排服务:
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import shutil
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
@@ -28,7 +27,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 EditPlan, EditPlanStatus
|
||||
from packages.domain.edit_plan import EditPlanStatus
|
||||
from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus
|
||||
from packages.domain.template_clip_config import TransitionEffect
|
||||
|
||||
|
||||
@@ -1,42 +0,0 @@
|
||||
/**
|
||||
* 仪表盘 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;
|
||||
};
|
||||
@@ -1,365 +0,0 @@
|
||||
/* 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;
|
||||
}
|
||||
@@ -1,211 +0,0 @@
|
||||
/**
|
||||
* 统一导航配置
|
||||
* 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 />,
|
||||
},
|
||||
],
|
||||
},
|
||||
];
|
||||
@@ -75,35 +75,6 @@
|
||||
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;
|
||||
@@ -302,17 +273,6 @@
|
||||
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;
|
||||
|
||||
@@ -612,54 +612,6 @@
|
||||
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
|
||||
============================================================ */
|
||||
|
||||
@@ -170,13 +170,6 @@ export const router = createBrowserRouter([
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "my-voices",
|
||||
lazy: () =>
|
||||
import("@/pages/my-voices/MyVoices").then((m) => ({
|
||||
Component: m.default,
|
||||
})),
|
||||
},
|
||||
{
|
||||
path: "accounts",
|
||||
lazy: () =>
|
||||
|
||||
@@ -1,18 +1,17 @@
|
||||
"""
|
||||
视频处理模块
|
||||
|
||||
轻量工具(ffmpeg_utils / oss_helpers / dedup_helpers)顶层直接导出,
|
||||
无额外依赖。渲染相关组件(UnifiedRenderService / RenderAdapter /
|
||||
VideoProcessor 等)按需从子模块导入,避免 __init__ 阶段引入
|
||||
packages / DB 等重依赖。
|
||||
"""
|
||||
|
||||
# 共享工具模块(供 editing_modes / generation / edit_plan_generation 等复用)
|
||||
# 共享工具模块(零外部依赖,供 editing_modes / generation / edit_plan_generation 等复用)
|
||||
from . import dedup_helpers, ffmpeg_utils, oss_helpers
|
||||
from .processor import VideoProcessor, VideoResult
|
||||
from .unified_render_service import RenderResult, UnifiedRenderService
|
||||
|
||||
__all__ = [
|
||||
"VideoProcessor",
|
||||
"VideoResult",
|
||||
"ffmpeg_utils",
|
||||
"oss_helpers",
|
||||
"dedup_helpers",
|
||||
"UnifiedRenderService",
|
||||
"RenderResult",
|
||||
]
|
||||
|
||||
Regular → Executable
+7
-3
@@ -93,12 +93,16 @@ class VideoFingerprint:
|
||||
resolution: tuple[int, int]
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
# 注意:color_histograms 里的值可能是 np.float32(来自 cv2.normalize),
|
||||
# 直接存进 dict 后 SQLAlchemy JSON 序列化会报 "float32 is not JSON serializable"。
|
||||
# 这里统一转成 Python 原生 float。
|
||||
native_histograms = [[float(v) for v in hist] for hist in self.color_histograms]
|
||||
return {
|
||||
"md5": self.md5,
|
||||
"keyframe_phashes": self.keyframe_phashes,
|
||||
"color_histograms": self.color_histograms,
|
||||
"duration": self.duration,
|
||||
"resolution": list(self.resolution),
|
||||
"color_histograms": native_histograms,
|
||||
"duration": float(self.duration),
|
||||
"resolution": [int(self.resolution[0]), int(self.resolution[1])],
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -1,657 +0,0 @@
|
||||
"""
|
||||
视频剪辑模式处理器
|
||||
支持四种剪辑模式:一镜到底、画中画、口播、口播+画中画
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
|
||||
if sys.version_info >= (3, 11):
|
||||
from enum import StrEnum
|
||||
else:
|
||||
from enum import Enum
|
||||
|
||||
class StrEnum(str, Enum):
|
||||
pass
|
||||
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
from video_processing.ffmpeg_utils import FFMPEG_BIN, probe_video_info, run_ffmpeg
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# 从 domain 层导入 EditingMode,避免重复定义
|
||||
from packages.domain.editing_mode import EditingMode
|
||||
|
||||
|
||||
class PIPPosition(StrEnum):
|
||||
"""画中画位置枚举"""
|
||||
|
||||
TOP_LEFT = "top_left"
|
||||
TOP_RIGHT = "top_right"
|
||||
BOTTOM_LEFT = "bottom_left"
|
||||
BOTTOM_RIGHT = "bottom_right"
|
||||
|
||||
|
||||
@dataclass
|
||||
class EditingModeConfig:
|
||||
"""剪辑模式配置"""
|
||||
|
||||
mode: EditingMode
|
||||
output_width: int = 1280
|
||||
output_height: int = 720
|
||||
output_fps: int = 25
|
||||
pip_position: PIPPosition = PIPPosition.TOP_RIGHT
|
||||
pip_scale: float = 0.25 # 画中画占主画面的比例
|
||||
transition_duration: float = 0.5 # 转场时长(秒)
|
||||
output_codec: str = "libx264"
|
||||
output_preset: str = "medium"
|
||||
output_crf: int = 23
|
||||
|
||||
|
||||
class EditingModeProcessor:
|
||||
"""剪辑模式处理器"""
|
||||
|
||||
def __init__(self, config: EditingModeConfig, work_dir: Optional[str] = None):
|
||||
"""
|
||||
初始化剪辑模式处理器
|
||||
|
||||
Args:
|
||||
config: 剪辑模式配置
|
||||
work_dir: 工作目录,默认使用系统临时目录
|
||||
"""
|
||||
self.config = config
|
||||
self.work_dir = work_dir or tempfile.gettempdir()
|
||||
|
||||
def process(
|
||||
self,
|
||||
video_paths: list[str],
|
||||
audio_path: Optional[str] = None,
|
||||
output_path: Optional[str] = None,
|
||||
) -> str:
|
||||
"""
|
||||
根据模式处理视频,返回输出文件路径
|
||||
|
||||
Args:
|
||||
video_paths: 视频素材路径列表
|
||||
audio_path: 音频路径(用于口播模式)
|
||||
output_path: 输出文件路径,默认自动生成
|
||||
|
||||
Returns:
|
||||
输出文件路径
|
||||
"""
|
||||
if not video_paths:
|
||||
raise ValueError("video_paths cannot be empty")
|
||||
|
||||
self._validate_inputs(video_paths, audio_path)
|
||||
|
||||
if output_path is None:
|
||||
output_path = self._generate_output_path()
|
||||
|
||||
logger.info(f"Processing videos with mode: {self.config.mode}, count: {len(video_paths)}")
|
||||
|
||||
try:
|
||||
if self.config.mode == EditingMode.ONE_TAKE:
|
||||
return self._one_take(video_paths, output_path)
|
||||
elif self.config.mode == EditingMode.PIP:
|
||||
return self._pip(video_paths, output_path)
|
||||
elif self.config.mode == EditingMode.VOICE_OVER:
|
||||
return self._voice_over(video_paths, audio_path, output_path)
|
||||
elif self.config.mode == EditingMode.VOICE_PIP:
|
||||
return self._voice_pip(video_paths, audio_path, output_path)
|
||||
else:
|
||||
raise ValueError(f"Unsupported editing mode: {self.config.mode}")
|
||||
except Exception as e:
|
||||
logger.error(f"Error processing videos: {e}")
|
||||
raise
|
||||
|
||||
def _validate_inputs(self, video_paths: list[str], audio_path: Optional[str]) -> None:
|
||||
"""验证输入文件"""
|
||||
for path in video_paths:
|
||||
if not os.path.exists(path):
|
||||
raise FileNotFoundError(f"Video file not found: {path}")
|
||||
if not os.path.getsize(path) > 0:
|
||||
raise ValueError(f"Video file is empty: {path}")
|
||||
|
||||
if audio_path and not os.path.exists(audio_path):
|
||||
raise FileNotFoundError(f"Audio file not found: {audio_path}")
|
||||
|
||||
def _generate_output_path(self) -> str:
|
||||
"""生成输出文件路径"""
|
||||
os.makedirs(self.work_dir, exist_ok=True)
|
||||
return os.path.join(self.work_dir, f"output_{self.config.mode}_{os.getpid()}.mp4")
|
||||
|
||||
def _run_ffmpeg(self, command: list[str], capture_output: bool = True) -> tuple:
|
||||
"""执行 FFmpeg 命令 — 委托给共享 ffmpeg_utils.run_ffmpeg"""
|
||||
try:
|
||||
return run_ffmpeg(command, capture_output=capture_output)
|
||||
except RuntimeError as e:
|
||||
logger.error(f"FFmpeg error: {e}")
|
||||
raise
|
||||
|
||||
def _get_video_info(self, video_path: str) -> dict:
|
||||
"""获取视频信息 — 委托给共享 ffmpeg_utils.probe_video_info,补充 codec/size 字段"""
|
||||
try:
|
||||
info = probe_video_info(video_path)
|
||||
info["codec"] = "unknown"
|
||||
info["size"] = os.path.getsize(video_path) if os.path.exists(video_path) else 0
|
||||
return info
|
||||
except Exception as e:
|
||||
logger.warning(f"Failed to get video info for {video_path}: {e}")
|
||||
return {"width": 0, "height": 0, "fps": 25, "duration": 0, "codec": "unknown", "size": 0}
|
||||
|
||||
def _get_pip_position_offset(
|
||||
self, main_width: int, main_height: int, pip_width: int, pip_height: int
|
||||
) -> tuple[int, int]:
|
||||
"""获取画中画位置偏移量"""
|
||||
margin = 10
|
||||
position_offsets = {
|
||||
PIPPosition.TOP_LEFT: (margin, margin),
|
||||
PIPPosition.TOP_RIGHT: (main_width - pip_width - margin, margin),
|
||||
PIPPosition.BOTTOM_LEFT: (margin, main_height - pip_height - margin),
|
||||
PIPPosition.BOTTOM_RIGHT: (main_width - pip_width - margin, main_height - pip_height - margin),
|
||||
}
|
||||
return position_offsets.get(self.config.pip_position, position_offsets[PIPPosition.TOP_RIGHT])
|
||||
|
||||
def _normalize_video(self, input_path: str, output_path: str) -> dict:
|
||||
"""标准化视频格式:先统一帧率,再缩放/填充"""
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-i",
|
||||
input_path,
|
||||
"-r",
|
||||
str(self.config.output_fps), # 先统一帧率
|
||||
"-vf",
|
||||
f"scale={self.config.output_width}:{self.config.output_height}:force_original_aspect_ratio=decrease,pad={self.config.output_width}:{self.config.output_height}:(ow-iw)/2:(oh-ih)/2,setsar=1",
|
||||
"-r",
|
||||
str(self.config.output_fps),
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
"-movflags",
|
||||
"+faststart",
|
||||
"-an",
|
||||
output_path,
|
||||
]
|
||||
run_ffmpeg(command)
|
||||
return self._get_video_info(output_path)
|
||||
|
||||
def _one_take(self, video_paths: list[str], output_path: str) -> str:
|
||||
"""一镜到底模式:顺序拼接视频,添加淡入淡出转场"""
|
||||
if len(video_paths) == 1:
|
||||
return self._normalize_video(video_paths[0], output_path)
|
||||
|
||||
normalized_paths = []
|
||||
for i, path in enumerate(video_paths):
|
||||
normalized = os.path.join(self.work_dir, f"normalized_{i}_{os.getpid()}.mp4")
|
||||
self._normalize_video(path, normalized)
|
||||
normalized_paths.append(normalized)
|
||||
|
||||
durations = [self._get_video_info(p)["duration"] for p in normalized_paths]
|
||||
|
||||
if len(normalized_paths) <= 5:
|
||||
output_path = self._one_take_with_xfade(normalized_paths, durations, output_path)
|
||||
else:
|
||||
output_path = self._one_take_simple_concat(normalized_paths, output_path)
|
||||
|
||||
for p in normalized_paths:
|
||||
try:
|
||||
if p != output_path:
|
||||
os.remove(p)
|
||||
except Exception as e:
|
||||
logger.warning(f"Operation failed in apps/worker/video_processing/editing_modes.py: {e}", exc_info=True)
|
||||
|
||||
return output_path
|
||||
|
||||
def _one_take_with_xfade(self, normalized_paths: list[str], durations: list[float], output_path: str) -> str:
|
||||
"""使用 xfade 滤镜实现转场"""
|
||||
if len(normalized_paths) == 2:
|
||||
transition = self.config.transition_duration
|
||||
offset1 = durations[0] - transition / 2
|
||||
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-i",
|
||||
normalized_paths[0],
|
||||
"-i",
|
||||
normalized_paths[1],
|
||||
"-filter_complex",
|
||||
f"[0:v][1:v]xfade=transition=fade:duration={transition}:offset={offset1}[v]",
|
||||
"-map",
|
||||
"[v]",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
output_path,
|
||||
]
|
||||
run_ffmpeg(command)
|
||||
return output_path
|
||||
else:
|
||||
return self._one_take_simple_concat(normalized_paths, output_path)
|
||||
|
||||
def _one_take_simple_concat(self, normalized_paths: list[str], output_path: str) -> str:
|
||||
"""使用 concat demuxer 简单拼接"""
|
||||
concat_file = os.path.join(self.work_dir, f"concat_list_{os.getpid()}.txt")
|
||||
with open(concat_file, "w") as f:
|
||||
for path in normalized_paths:
|
||||
f.write(f"file '{os.path.abspath(path)}'\n")
|
||||
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-f",
|
||||
"concat",
|
||||
"-safe",
|
||||
"0",
|
||||
"-i",
|
||||
concat_file,
|
||||
"-c",
|
||||
"copy",
|
||||
output_path,
|
||||
]
|
||||
run_ffmpeg(command)
|
||||
|
||||
try:
|
||||
os.remove(concat_file)
|
||||
except Exception as e:
|
||||
logger.warning(f"Operation failed in apps/worker/video_processing/editing_modes.py: {e}", exc_info=True)
|
||||
|
||||
return output_path
|
||||
|
||||
def _pip(self, video_paths: list[str], output_path: str) -> str:
|
||||
"""画中画模式:主视频全屏,后续视频叠加在角落"""
|
||||
if not video_paths:
|
||||
raise ValueError("No video paths provided")
|
||||
|
||||
main_video = video_paths[0]
|
||||
main_normalized = os.path.join(self.work_dir, f"main_{os.getpid()}.mp4")
|
||||
main_info = self._normalize_video(main_video, main_normalized)
|
||||
|
||||
if len(video_paths) == 1:
|
||||
os.rename(main_normalized, output_path)
|
||||
return output_path
|
||||
|
||||
pip_width = int(self.config.output_width * self.config.pip_scale)
|
||||
pip_height = int(self.config.output_height * self.config.pip_scale)
|
||||
x_offset, y_offset = self._get_pip_position_offset(
|
||||
self.config.output_width, self.config.output_height, pip_width, pip_height
|
||||
)
|
||||
|
||||
pip_normalized = os.path.join(self.work_dir, f"pip_{os.getpid()}.mp4")
|
||||
pip_info = self._get_video_info(video_paths[1])
|
||||
|
||||
if pip_info["duration"] > main_info["duration"]:
|
||||
temp_pip = os.path.join(self.work_dir, f"pip_temp_{os.getpid()}.mp4")
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-i",
|
||||
video_paths[1],
|
||||
"-t",
|
||||
str(main_info["duration"]),
|
||||
"-vf",
|
||||
f"scale={pip_width}:{pip_height}",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
temp_pip,
|
||||
]
|
||||
run_ffmpeg(command)
|
||||
pip_normalized_input = temp_pip
|
||||
else:
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-i",
|
||||
video_paths[1],
|
||||
"-vf",
|
||||
f"scale={pip_width}:{pip_height}",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
pip_normalized,
|
||||
]
|
||||
run_ffmpeg(command)
|
||||
pip_normalized_input = pip_normalized
|
||||
|
||||
if main_info["duration"] > pip_info["duration"]:
|
||||
looped_pip = os.path.join(self.work_dir, f"pip_looped_{os.getpid()}.mp4")
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-stream_loop",
|
||||
"-1",
|
||||
"-i",
|
||||
pip_normalized_input,
|
||||
"-t",
|
||||
str(main_info["duration"]),
|
||||
"-vf",
|
||||
f"scale={pip_width}:{pip_height}",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
looped_pip,
|
||||
]
|
||||
run_ffmpeg(command)
|
||||
pip_normalized_input = looped_pip
|
||||
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-i",
|
||||
main_normalized,
|
||||
"-i",
|
||||
pip_normalized_input,
|
||||
"-filter_complex",
|
||||
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
|
||||
"-map",
|
||||
"[v]",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
output_path,
|
||||
]
|
||||
run_ffmpeg(command)
|
||||
|
||||
for temp_file in [main_normalized, pip_normalized]:
|
||||
if temp_file and temp_file != output_path:
|
||||
try:
|
||||
os.remove(temp_file)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Operation failed in apps/worker/video_processing/editing_modes.py: {e}", exc_info=True
|
||||
)
|
||||
|
||||
return output_path
|
||||
|
||||
def _voice_over(self, video_paths: list[str], audio_path: Optional[str], output_path: str) -> str:
|
||||
"""口播模式:背景画面 + 配音"""
|
||||
if not audio_path:
|
||||
raise ValueError("audio_path is required for VOICE_OVER mode")
|
||||
|
||||
if not video_paths:
|
||||
raise ValueError("No background video provided")
|
||||
|
||||
audio_info = self._get_video_info(audio_path)
|
||||
audio_duration = audio_info["duration"]
|
||||
|
||||
bg_normalized = os.path.join(self.work_dir, f"bg_{os.getpid()}.mp4")
|
||||
bg_info = self._normalize_video(video_paths[0], bg_normalized)
|
||||
|
||||
if bg_info["duration"] < audio_duration:
|
||||
looped_bg = os.path.join(self.work_dir, f"bg_looped_{os.getpid()}.mp4")
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-stream_loop",
|
||||
"-1",
|
||||
"-i",
|
||||
bg_normalized,
|
||||
"-t",
|
||||
str(audio_duration),
|
||||
"-vf",
|
||||
f"scale={self.config.output_width}:{self.config.output_height}",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
looped_bg,
|
||||
]
|
||||
run_ffmpeg(command)
|
||||
bg_normalized = looped_bg
|
||||
elif bg_info["duration"] > audio_duration:
|
||||
temp_bg = os.path.join(self.work_dir, f"bg_trimmed_{os.getpid()}.mp4")
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-i",
|
||||
bg_normalized,
|
||||
"-t",
|
||||
str(audio_duration),
|
||||
"-c:v",
|
||||
"copy",
|
||||
temp_bg,
|
||||
]
|
||||
run_ffmpeg(command)
|
||||
bg_normalized = temp_bg
|
||||
|
||||
blurred_bg = os.path.join(self.work_dir, f"bg_blurred_{os.getpid()}.mp4")
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-i",
|
||||
bg_normalized,
|
||||
"-vf",
|
||||
f"boxblur=5:5,scale={self.config.output_width}:{self.config.output_height}",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
blurred_bg,
|
||||
]
|
||||
run_ffmpeg(command)
|
||||
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-i",
|
||||
blurred_bg,
|
||||
"-i",
|
||||
audio_path,
|
||||
"-filter_complex",
|
||||
"[0:v]drawbox=x=0:y=0:w=iw:h=ih:color=black@0.3:t=fill[v]",
|
||||
"-map",
|
||||
"[v]",
|
||||
"-map",
|
||||
"1:a",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
"-shortest",
|
||||
output_path,
|
||||
]
|
||||
run_ffmpeg(command)
|
||||
|
||||
for temp_file in [bg_normalized, blurred_bg]:
|
||||
try:
|
||||
if temp_file != output_path:
|
||||
os.remove(temp_file)
|
||||
except Exception as e:
|
||||
logger.warning(f"Operation failed in apps/worker/video_processing/editing_modes.py: {e}", exc_info=True)
|
||||
|
||||
return output_path
|
||||
|
||||
def _voice_pip(self, video_paths: list[str], audio_path: Optional[str], output_path: str) -> str:
|
||||
"""口播+画中画模式:口播视频在角落,其他视频作为背景"""
|
||||
if not video_paths:
|
||||
raise ValueError("No video paths provided")
|
||||
|
||||
if len(video_paths) == 1:
|
||||
return self._normalize_video(video_paths[0], output_path)
|
||||
|
||||
voice_video = video_paths[0]
|
||||
bg_video = video_paths[1] if len(video_paths) > 1 else video_paths[0]
|
||||
|
||||
voice_normalized = os.path.join(self.work_dir, f"voice_{os.getpid()}.mp4")
|
||||
voice_info = self._normalize_video(voice_video, voice_normalized)
|
||||
|
||||
bg_normalized = os.path.join(self.work_dir, f"bg_{os.getpid()}.mp4")
|
||||
bg_info = self._normalize_video(bg_video, bg_normalized)
|
||||
|
||||
final_duration = min(voice_info["duration"], bg_info["duration"])
|
||||
|
||||
pip_width = int(self.config.output_width * self.config.pip_scale)
|
||||
pip_height = int(self.config.output_height * self.config.pip_scale)
|
||||
x_offset, y_offset = self._get_pip_position_offset(
|
||||
self.config.output_width, self.config.output_height, pip_width, pip_height
|
||||
)
|
||||
|
||||
voice_adjusted = os.path.join(self.work_dir, f"voice_adj_{os.getpid()}.mp4")
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-i",
|
||||
voice_normalized,
|
||||
"-t",
|
||||
str(final_duration),
|
||||
"-vf",
|
||||
f"scale={pip_width}:{pip_height}",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
voice_adjusted,
|
||||
]
|
||||
run_ffmpeg(command)
|
||||
|
||||
bg_adjusted = os.path.join(self.work_dir, f"bg_adj_{os.getpid()}.mp4")
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-i",
|
||||
bg_normalized,
|
||||
"-t",
|
||||
str(final_duration),
|
||||
"-c:v",
|
||||
"copy",
|
||||
bg_adjusted,
|
||||
]
|
||||
run_ffmpeg(command)
|
||||
|
||||
if audio_path:
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-i",
|
||||
bg_adjusted,
|
||||
"-i",
|
||||
voice_adjusted,
|
||||
"-i",
|
||||
audio_path,
|
||||
"-filter_complex",
|
||||
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
|
||||
"-map",
|
||||
"[v]",
|
||||
"-map",
|
||||
"2:a",
|
||||
"-shortest",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
output_path,
|
||||
]
|
||||
else:
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-i",
|
||||
bg_adjusted,
|
||||
"-i",
|
||||
voice_adjusted,
|
||||
"-filter_complex",
|
||||
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
|
||||
"-map",
|
||||
"[v]",
|
||||
"-map",
|
||||
"1:a",
|
||||
"-shortest",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
output_path,
|
||||
]
|
||||
run_ffmpeg(command)
|
||||
|
||||
for temp_file in [voice_normalized, voice_adjusted, bg_normalized, bg_adjusted]:
|
||||
try:
|
||||
if temp_file != output_path:
|
||||
os.remove(temp_file)
|
||||
except Exception as e:
|
||||
logger.warning(f"Operation failed in apps/worker/video_processing/editing_modes.py: {e}", exc_info=True)
|
||||
|
||||
return output_path
|
||||
|
||||
|
||||
def create_processor(mode: str, work_dir: Optional[str] = None, **kwargs) -> EditingModeProcessor:
|
||||
"""便捷工厂函数:创建剪辑模式处理器"""
|
||||
try:
|
||||
editing_mode = EditingMode(mode)
|
||||
except ValueError:
|
||||
raise ValueError(f"Invalid editing mode: {mode}. Valid modes: {[m.value for m in EditingMode]}")
|
||||
|
||||
config = EditingModeConfig(
|
||||
mode=editing_mode,
|
||||
output_width=kwargs.get("output_width", 1280),
|
||||
output_height=kwargs.get("output_height", 720),
|
||||
output_fps=kwargs.get("output_fps", 25),
|
||||
pip_position=PIPPosition(kwargs.get("pip_position", "top_right")),
|
||||
pip_scale=kwargs.get("pip_scale", 0.25),
|
||||
transition_duration=kwargs.get("transition_duration", 0.5),
|
||||
)
|
||||
|
||||
return EditingModeProcessor(config=config, work_dir=work_dir)
|
||||
Regular → Executable
+85
-13
@@ -1,8 +1,7 @@
|
||||
"""FFmpeg 工具函数 — 从 editing_modes.py / video_compose_service.py 提取的共享原语.
|
||||
"""FFmpeg 工具函数 — 共享原语.
|
||||
|
||||
提供 FFmpeg / FFprobe 调用、视频信息探测、视频标准化、xfade 转场滤镜构建
|
||||
等底层能力,供 EditingModeProcessor、VideoComposeService、UnifiedRenderService
|
||||
共同复用。
|
||||
等底层能力,供 UnifiedRenderService、VideoComposeService 等复用。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -32,6 +31,10 @@ XFADE_TRANSITION_MAP: dict[str, str] = {
|
||||
"slide_left": "slideleft",
|
||||
"slideright": "slideright",
|
||||
"slide_right": "slideright",
|
||||
"slideup": "slideup",
|
||||
"slide_up": "slideup",
|
||||
"slidedown": "slidedown",
|
||||
"slide_down": "slidedown",
|
||||
"dissolve": "dissolve",
|
||||
"wipe": "wipeleft",
|
||||
"wipeleft": "wipeleft",
|
||||
@@ -39,6 +42,10 @@ XFADE_TRANSITION_MAP: dict[str, str] = {
|
||||
|
||||
DEFAULT_TRANSITION_DURATION = 0.5
|
||||
|
||||
# FFmpeg 执行默认超时(秒),防止 FFmpeg hang 住导致 worker 永久阻塞
|
||||
# 默认 30 分钟,足够处理大部分短视频渲染;超长视频可单独传参覆盖
|
||||
DEFAULT_FFMPEG_TIMEOUT = 1800
|
||||
|
||||
|
||||
# ── FFmpeg 执行 ───────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -47,12 +54,14 @@ def run_ffmpeg(
|
||||
command: list[str],
|
||||
*,
|
||||
capture_output: bool = True,
|
||||
timeout: int | None = DEFAULT_FFMPEG_TIMEOUT,
|
||||
) -> tuple[str, str]:
|
||||
"""执行 FFmpeg 命令。
|
||||
|
||||
Args:
|
||||
command: 完整的 ffmpeg 命令列表(含 "ffmpeg" 本身)
|
||||
capture_output: 是否捕获 stdout/stderr
|
||||
timeout: 超时时间(秒),默认 1800s(30分钟);None 表示不设超时(不推荐)
|
||||
|
||||
Returns:
|
||||
(stdout, stderr) 元组
|
||||
@@ -60,6 +69,7 @@ def run_ffmpeg(
|
||||
Raises:
|
||||
subprocess.CalledProcessError: 命令执行失败时抛出,
|
||||
异常信息包含完整 stderr 以便排查。
|
||||
subprocess.TimeoutExpired: 超时未完成时抛出,FFmpeg 进程会被 kill。
|
||||
"""
|
||||
try:
|
||||
result = subprocess.run( # nosec B603
|
||||
@@ -68,8 +78,16 @@ def run_ffmpeg(
|
||||
stdout=subprocess.PIPE if capture_output else None,
|
||||
stderr=subprocess.PIPE if capture_output else None,
|
||||
text=True,
|
||||
timeout=timeout,
|
||||
)
|
||||
return (result.stdout or "", result.stderr or "")
|
||||
except subprocess.TimeoutExpired as e:
|
||||
logger.error(
|
||||
"FFmpeg 命令超时 (%ds): command=%s",
|
||||
timeout or -1,
|
||||
" ".join(str(c) for c in command[:20]),
|
||||
)
|
||||
raise
|
||||
except subprocess.CalledProcessError as e:
|
||||
# 把完整 stderr 打到日志,方便排查 exit code 183 等问题
|
||||
stderr_text = (e.stderr or "").strip()
|
||||
@@ -82,6 +100,41 @@ def run_ffmpeg(
|
||||
raise
|
||||
|
||||
|
||||
def probe_has_audio(local_path: str | Path) -> bool:
|
||||
"""探测文件是否包含音频流。
|
||||
|
||||
Args:
|
||||
local_path: 本地文件路径
|
||||
|
||||
Returns:
|
||||
True 表示有音频流(或探测失败保守返回),False 表示确认无音频流
|
||||
"""
|
||||
try:
|
||||
result = subprocess.run( # nosec B603
|
||||
[
|
||||
FFPROBE_BIN,
|
||||
"-v",
|
||||
"error",
|
||||
"-select_streams",
|
||||
"a:0",
|
||||
"-show_entries",
|
||||
"stream=codec_type",
|
||||
"-of",
|
||||
"default=noprint_wrappers=1:nokey=1",
|
||||
str(local_path),
|
||||
],
|
||||
check=True,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
timeout=10,
|
||||
)
|
||||
return result.stdout.strip() == "audio"
|
||||
except Exception:
|
||||
# 探测失败保守返回 True,让 FFmpeg 自己处理(避免误删音频)
|
||||
return True
|
||||
|
||||
|
||||
def probe_duration(local_path: str | Path) -> float:
|
||||
"""用 ffprobe 获取视频时长(秒)。
|
||||
|
||||
@@ -110,10 +163,14 @@ def probe_duration(local_path: str | Path) -> float:
|
||||
|
||||
|
||||
def probe_video_info(video_path: str) -> dict[str, Any]:
|
||||
"""获取视频信息(宽、高、时长、fps)。
|
||||
"""获取视频信息(宽、高、时长、fps、编码、像素格式)。
|
||||
|
||||
Returns:
|
||||
{"width": int, "height": int, "duration": float, "fps": float}
|
||||
{
|
||||
"width": int, "height": int, "duration": float, "fps": float,
|
||||
"video_codec": str, "audio_codec": str, "pix_fmt": str,
|
||||
"has_audio": bool,
|
||||
}
|
||||
失败时返回默认值。
|
||||
"""
|
||||
try:
|
||||
@@ -122,10 +179,8 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
|
||||
FFPROBE_BIN,
|
||||
"-v",
|
||||
"error",
|
||||
"-select_streams",
|
||||
"v:0",
|
||||
"-show_entries",
|
||||
"stream=width,height,r_frame_rate,duration",
|
||||
"stream=width,height,r_frame_rate,duration,codec_name,codec_type,pix_fmt",
|
||||
"-show_entries",
|
||||
"format=duration",
|
||||
"-of",
|
||||
@@ -136,19 +191,25 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
timeout=15,
|
||||
)
|
||||
|
||||
import json
|
||||
|
||||
info = json.loads(result.stdout)
|
||||
stream = info.get("streams", [{}])[0]
|
||||
streams = info.get("streams", [])
|
||||
fmt = info.get("format", {})
|
||||
|
||||
width = int(stream.get("width", DEFAULT_OUTPUT_WIDTH))
|
||||
height = int(stream.get("height", DEFAULT_OUTPUT_HEIGHT))
|
||||
video_stream = next((s for s in streams if s.get("codec_type") == "video"), {})
|
||||
audio_stream = next((s for s in streams if s.get("codec_type") == "audio"), {})
|
||||
|
||||
width = int(video_stream.get("width", DEFAULT_OUTPUT_WIDTH))
|
||||
height = int(video_stream.get("height", DEFAULT_OUTPUT_HEIGHT))
|
||||
video_codec = video_stream.get("codec_name", "") or ""
|
||||
pix_fmt = video_stream.get("pix_fmt", "") or ""
|
||||
|
||||
# 解析帧率
|
||||
fps_str = stream.get("r_frame_rate", "25/1")
|
||||
fps_str = video_stream.get("r_frame_rate", "25/1")
|
||||
if "/" in fps_str:
|
||||
num, den = fps_str.split("/")
|
||||
fps = float(num) / float(den) if float(den) > 0 else DEFAULT_FPS
|
||||
@@ -156,13 +217,20 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
|
||||
fps = float(fps_str) if fps_str else DEFAULT_FPS
|
||||
|
||||
# 时长
|
||||
duration = float(fmt.get("duration", 0)) or float(stream.get("duration", 0))
|
||||
duration = float(fmt.get("duration", 0)) or float(video_stream.get("duration", 0))
|
||||
|
||||
has_audio = bool(audio_stream)
|
||||
audio_codec = audio_stream.get("codec_name", "") or ""
|
||||
|
||||
return {
|
||||
"width": width,
|
||||
"height": height,
|
||||
"duration": duration,
|
||||
"fps": round(fps, 2),
|
||||
"video_codec": video_codec,
|
||||
"audio_codec": audio_codec,
|
||||
"pix_fmt": pix_fmt,
|
||||
"has_audio": has_audio,
|
||||
}
|
||||
except Exception as e:
|
||||
logger.warning("获取视频信息失败: %s, error: %s", video_path, e)
|
||||
@@ -171,6 +239,10 @@ def probe_video_info(video_path: str) -> dict[str, Any]:
|
||||
"height": DEFAULT_OUTPUT_HEIGHT,
|
||||
"duration": 0.0,
|
||||
"fps": DEFAULT_FPS,
|
||||
"video_codec": "",
|
||||
"audio_codec": "",
|
||||
"pix_fmt": "",
|
||||
"has_audio": True,
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -9,6 +9,7 @@ from __future__ import annotations
|
||||
import hashlib
|
||||
import logging
|
||||
import os
|
||||
import threading
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
from urllib.parse import urlparse
|
||||
@@ -17,6 +18,13 @@ import oss2
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# OSS 上传配置
|
||||
OSS_CONNECT_TIMEOUT = 10 # 连接超时(秒),防止 TCP 握手挂死
|
||||
OSS_UPLOAD_TOTAL_TIMEOUT = 300 # 单文件上传总超时(秒),防止网络慢时无限卡住
|
||||
OSS_MULTIPART_THRESHOLD = 100 * 1024 * 1024 # 分片上传阈值:100MB 以上走分片
|
||||
OSS_PART_SIZE = 8 * 1024 * 1024 # 分片大小:8MB
|
||||
OSS_MULTIPART_NUM_THREADS = 3 # 分片上传并发数
|
||||
|
||||
|
||||
# ── OSS 配置 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
@@ -43,6 +51,9 @@ def oss_bucket() -> oss2.Bucket | None:
|
||||
P0-2 修复:endpoint 不带 scheme 时自动补 https:// 前缀,
|
||||
确保 sign_url 等依赖 scheme 的方法返回 HTTPS URL。
|
||||
|
||||
P0-staging 修复:增加 connect_timeout=10s,防止网络抖动时
|
||||
TCP 握手阶段无限挂死,导致 worker 进程卡死。
|
||||
|
||||
Returns:
|
||||
oss2.Bucket 实例,配置缺失时返回 None。
|
||||
"""
|
||||
@@ -53,7 +64,12 @@ def oss_bucket() -> oss2.Bucket | None:
|
||||
# endpoint 无 scheme 时补 https://,与 API 端 storage.py 保持一致
|
||||
if not endpoint.startswith(("http://", "https://")):
|
||||
endpoint = f"https://{endpoint}"
|
||||
return oss2.Bucket(oss2.Auth(access_key_id, access_key_secret), endpoint, bucket_name)
|
||||
return oss2.Bucket(
|
||||
oss2.Auth(access_key_id, access_key_secret),
|
||||
endpoint,
|
||||
bucket_name,
|
||||
connect_timeout=OSS_CONNECT_TIMEOUT,
|
||||
)
|
||||
|
||||
|
||||
def normalize_storage_key(storage_key_or_url: str) -> str:
|
||||
@@ -96,6 +112,9 @@ def download_asset(asset_storage_key: str, local_path: Path) -> bool:
|
||||
def upload_to_oss(local_path: Path, storage_key: str) -> str | None:
|
||||
"""上传文件到 OSS,返回公开 URL。
|
||||
|
||||
大文件(>100MB)自动走分片上传,降低内存峰值,减少 OOM 风险。
|
||||
上传加总超时保护(默认 300s),防止网络异常时无限挂死。
|
||||
|
||||
Args:
|
||||
local_path: 本地文件路径
|
||||
storage_key: 目标存储键
|
||||
@@ -106,18 +125,71 @@ def upload_to_oss(local_path: Path, storage_key: str) -> str | None:
|
||||
bucket = oss_bucket()
|
||||
if bucket is None:
|
||||
return None
|
||||
try:
|
||||
bucket.put_object_from_file(storage_key, str(local_path))
|
||||
settings = oss_settings()
|
||||
if settings:
|
||||
_, _, endpoint, bucket_name = settings
|
||||
endpoint_clean = endpoint.replace("https://", "").replace("http://", "")
|
||||
return f"https://{bucket_name}.{endpoint_clean}/{storage_key}"
|
||||
|
||||
result: dict = {"url": None, "error": None, "file_size": 0}
|
||||
done = threading.Event()
|
||||
|
||||
def _do_upload():
|
||||
try:
|
||||
# 尝试获取文件大小,用于分片判断和日志;stat 失败时 fallback 走普通上传
|
||||
try:
|
||||
file_size = local_path.stat().st_size
|
||||
result["file_size"] = file_size
|
||||
use_multipart = file_size >= OSS_MULTIPART_THRESHOLD
|
||||
except OSError:
|
||||
use_multipart = False
|
||||
file_size = 0
|
||||
|
||||
if use_multipart:
|
||||
# 分片上传:降低内存峰值,每片 8MB,3 线程并发
|
||||
logger.info(
|
||||
"大文件分片上传: storage_key=%s, size=%.1fMB, part_size=%dMB, threads=%d",
|
||||
storage_key[:80],
|
||||
file_size / 1024 / 1024,
|
||||
OSS_PART_SIZE // 1024 // 1024,
|
||||
OSS_MULTIPART_NUM_THREADS,
|
||||
)
|
||||
oss2.resumable_upload(
|
||||
bucket,
|
||||
storage_key,
|
||||
str(local_path),
|
||||
multipart_threshold=OSS_MULTIPART_THRESHOLD,
|
||||
part_size=OSS_PART_SIZE,
|
||||
num_threads=OSS_MULTIPART_NUM_THREADS,
|
||||
)
|
||||
else:
|
||||
bucket.put_object_from_file(storage_key, str(local_path))
|
||||
|
||||
# 构造返回 URL
|
||||
settings = oss_settings()
|
||||
if settings:
|
||||
_, _, endpoint, bucket_name = settings
|
||||
endpoint_clean = endpoint.replace("https://", "").replace("http://", "")
|
||||
result["url"] = f"https://{bucket_name}.{endpoint_clean}/{storage_key}"
|
||||
except Exception as e:
|
||||
result["error"] = e
|
||||
logger.exception("上传 OSS 失败: %s", storage_key)
|
||||
finally:
|
||||
done.set()
|
||||
|
||||
upload_thread = threading.Thread(target=_do_upload, daemon=True)
|
||||
upload_thread.start()
|
||||
finished = done.wait(timeout=OSS_UPLOAD_TOTAL_TIMEOUT)
|
||||
|
||||
if not finished:
|
||||
logger.error(
|
||||
"OSS 上传超时(%.0fs),强制中止: storage_key=%s, size=%.1fMB",
|
||||
OSS_UPLOAD_TOTAL_TIMEOUT,
|
||||
storage_key[:80],
|
||||
result["file_size"] / 1024 / 1024 if result["file_size"] else 0,
|
||||
)
|
||||
return None
|
||||
except Exception:
|
||||
logger.exception("上传 OSS 失败: %s", storage_key)
|
||||
|
||||
if result["error"]:
|
||||
return None
|
||||
|
||||
return result["url"]
|
||||
|
||||
|
||||
def get_signed_download_url(storage_key_or_url: str, expires_seconds: int = 3600) -> str | None:
|
||||
"""生成预签名下载 URL(用于私有 bucket 的 URL 校验或临时下载)。
|
||||
|
||||
+299
@@ -0,0 +1,299 @@
|
||||
"""统一渲染引擎适配层 — Phase 2.
|
||||
|
||||
将 EditPlan + EditPlanClips(来自 DB)适配为 UnifiedRenderService 的输入格式,
|
||||
封装素材下载、渲染执行、结果上传的完整流程。
|
||||
|
||||
职责:
|
||||
1. 从 DB 读取 EditPlan + EditPlanClips
|
||||
2. 下载素材到本地,构建 asset_path_map
|
||||
3. 调用 UnifiedRenderService 执行渲染
|
||||
4. 上传渲染结果到 OSS
|
||||
5. 支持进度回调(对接 JobService)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any, Callable
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
from video_processing.oss_helpers import download_asset, upload_to_oss
|
||||
from video_processing.unified_render_service import RenderResult, UnifiedRenderService
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.edit_plan_clip_repository import SQLAlchemyEditPlanClipRepository
|
||||
from packages.adapters.sqlalchemy_impl.edit_plan_repository import SQLAlchemyEditPlanRepository
|
||||
from packages.domain.edit_plan import EditPlan, EditPlanStatus
|
||||
from packages.domain.edit_plan_clip import EditPlanClip, EditPlanClipStatus
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# ── 数据结构 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class RenderAdapterResult:
|
||||
"""渲染适配结果。"""
|
||||
|
||||
success: bool
|
||||
output_url: str = ""
|
||||
output_path: Path | None = None
|
||||
duration: float = 0.0
|
||||
file_size: int = 0
|
||||
width: int = 0
|
||||
height: int = 0
|
||||
clip_count: int = 0
|
||||
error_message: str = ""
|
||||
|
||||
|
||||
ProgressCallback = Callable[[float, str], None]
|
||||
"""进度回调:(progress_0_100, stage_description) → None"""
|
||||
|
||||
|
||||
# ── 适配层主体 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class RenderAdapter:
|
||||
"""统一渲染引擎适配层。
|
||||
|
||||
桥接 EditPlan 领域模型与 UnifiedRenderService 图层模型。
|
||||
|
||||
用法::
|
||||
|
||||
adapter = RenderAdapter(db)
|
||||
result = adapter.render_plan(
|
||||
plan_id=plan_id,
|
||||
job_id=job_id,
|
||||
progress_cb=lambda p, s: job_service.update_progress(job_id, p, s),
|
||||
)
|
||||
"""
|
||||
|
||||
def __init__(self, db: Session) -> None:
|
||||
self._db = db
|
||||
self._plan_repo = SQLAlchemyEditPlanRepository(db)
|
||||
self._clip_repo = SQLAlchemyEditPlanClipRepository(db)
|
||||
|
||||
# ── 公开方法 ──────────────────────────────────────────────────────────
|
||||
|
||||
def render_plan(
|
||||
self,
|
||||
plan_id: str,
|
||||
*,
|
||||
job_id: str = "",
|
||||
work_dir: Path | None = None,
|
||||
progress_cb: ProgressCallback | None = None,
|
||||
) -> RenderAdapterResult:
|
||||
"""渲染一个 EditPlan。
|
||||
|
||||
完整流程:
|
||||
1. 加载计划与片段
|
||||
2. 下载素材
|
||||
3. 执行统一渲染
|
||||
4. 上传结果
|
||||
|
||||
Args:
|
||||
plan_id: EditPlan ID
|
||||
job_id: 关联的 Job ID(用于结果存储路径)
|
||||
work_dir: 工作目录,不传则使用临时目录
|
||||
progress_cb: 进度回调函数
|
||||
|
||||
Returns:
|
||||
RenderAdapterResult
|
||||
"""
|
||||
temp_dir = None
|
||||
try:
|
||||
# 0. 准备工作目录
|
||||
if work_dir is None:
|
||||
temp_dir = tempfile.mkdtemp(prefix="render_")
|
||||
work_dir = Path(temp_dir)
|
||||
work_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
self._report_progress(progress_cb, 5.0, "加载剪辑计划")
|
||||
|
||||
# 1. 加载计划与片段
|
||||
plan = self._plan_repo.get(plan_id)
|
||||
if plan is None:
|
||||
return RenderAdapterResult(
|
||||
success=False,
|
||||
error_message=f"剪辑计划不存在: {plan_id}",
|
||||
)
|
||||
|
||||
clips = self._clip_repo.list_by_plan(plan_id, skip=0, limit=10000)
|
||||
ready_clips = [c for c in clips if c.status == EditPlanClipStatus.READY and c.asset_id]
|
||||
ready_clips.sort(key=lambda c: c.order)
|
||||
|
||||
if not ready_clips:
|
||||
return RenderAdapterResult(
|
||||
success=False,
|
||||
error_message="没有可渲染的就绪片段",
|
||||
clip_count=0,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"开始渲染: plan_id=%s job_id=%s ready_clips=%d engine=unified",
|
||||
plan_id,
|
||||
job_id,
|
||||
len(ready_clips),
|
||||
)
|
||||
|
||||
self._report_progress(progress_cb, 15.0, f"下载素材({len(ready_clips)} 个)")
|
||||
|
||||
# 2. 下载素材
|
||||
asset_path_map = self._download_assets(ready_clips, work_dir)
|
||||
if not asset_path_map:
|
||||
return RenderAdapterResult(
|
||||
success=False,
|
||||
error_message="所有素材下载失败",
|
||||
clip_count=len(ready_clips),
|
||||
)
|
||||
|
||||
self._report_progress(progress_cb, 40.0, "执行视频渲染")
|
||||
|
||||
# 3. 执行统一渲染
|
||||
render_svc = UnifiedRenderService(
|
||||
plan=plan,
|
||||
clips=ready_clips,
|
||||
asset_path_map=asset_path_map,
|
||||
work_dir=work_dir,
|
||||
)
|
||||
result = render_svc.render()
|
||||
|
||||
self._report_progress(progress_cb, 80.0, "上传渲染结果")
|
||||
|
||||
# 4. 上传结果
|
||||
storage_key = f"rendered/{plan_id}/{job_id or plan_id}.mp4"
|
||||
output_url = upload_to_oss(result.output_path, storage_key)
|
||||
|
||||
self._report_progress(progress_cb, 100.0, "渲染完成")
|
||||
|
||||
logger.info(
|
||||
"[render-adapter] render success: plan_id=%s job_id=%s engine=unified "
|
||||
"duration=%.2fs file_size=%d resolution=%dx%d clip_count=%d",
|
||||
plan_id,
|
||||
job_id,
|
||||
result.duration,
|
||||
result.file_size,
|
||||
result.width,
|
||||
result.height,
|
||||
len(ready_clips),
|
||||
)
|
||||
|
||||
return RenderAdapterResult(
|
||||
success=True,
|
||||
output_url=output_url or "",
|
||||
output_path=result.output_path,
|
||||
duration=result.duration,
|
||||
file_size=result.file_size,
|
||||
width=result.width,
|
||||
height=result.height,
|
||||
clip_count=len(ready_clips),
|
||||
)
|
||||
|
||||
except Exception as exc:
|
||||
logger.exception(
|
||||
"[render-adapter] render failed: plan_id=%s job_id=%s engine=unified error=%s",
|
||||
plan_id,
|
||||
job_id,
|
||||
str(exc)[:200],
|
||||
)
|
||||
return RenderAdapterResult(
|
||||
success=False,
|
||||
error_message=str(exc)[:500],
|
||||
)
|
||||
finally:
|
||||
# 清理临时目录
|
||||
if temp_dir:
|
||||
import shutil
|
||||
|
||||
try:
|
||||
shutil.rmtree(temp_dir, ignore_errors=True)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def validate_plan(self, plan_id: str) -> tuple[bool, list[str], list[str], int, int]:
|
||||
"""校验计划是否可渲染(兼容 VideoComposeService.validate_compose 接口)。
|
||||
|
||||
Returns:
|
||||
(valid, errors, warnings, ready_clip_count, total_clip_count)
|
||||
"""
|
||||
errors: list[str] = []
|
||||
warnings: list[str] = []
|
||||
|
||||
plan = self._plan_repo.get(plan_id)
|
||||
if plan is None:
|
||||
return False, [f"剪辑计划不存在: {plan_id}"], [], 0, 0
|
||||
|
||||
if plan.status not in (EditPlanStatus.EDITING, EditPlanStatus.RENDERING):
|
||||
errors.append(f"计划状态不正确,需要 editing 或 rendering,当前: {plan.status}")
|
||||
|
||||
clips = self._clip_repo.list_by_plan(plan_id, skip=0, limit=10000)
|
||||
if not clips:
|
||||
errors.append("计划没有任何片段")
|
||||
return False, errors, warnings, 0, 0
|
||||
|
||||
clips.sort(key=lambda c: c.order)
|
||||
|
||||
ready_count = 0
|
||||
pending_count = 0
|
||||
no_asset_count = 0
|
||||
|
||||
for clip in clips:
|
||||
if clip.status == EditPlanClipStatus.READY:
|
||||
ready_count += 1
|
||||
if not clip.asset_id:
|
||||
errors.append(f"片段 {clip.id} (order={clip.order}) 没有分配素材")
|
||||
no_asset_count += 1
|
||||
elif clip.status == EditPlanClipStatus.PENDING:
|
||||
pending_count += 1
|
||||
elif clip.status == EditPlanClipStatus.FAILED:
|
||||
warnings.append(f"片段 {clip.id} (order={clip.order}) 状态为 failed,已跳过")
|
||||
|
||||
if ready_count == 0:
|
||||
errors.append("没有就绪(ready)的片段可以合成")
|
||||
|
||||
if pending_count > 0:
|
||||
warnings.append(f"有 {pending_count} 个片段仍处于 pending 状态")
|
||||
|
||||
return len(errors) == 0, errors, warnings, ready_count, len(clips)
|
||||
|
||||
# ── 内部方法 ──────────────────────────────────────────────────────────
|
||||
|
||||
@staticmethod
|
||||
def _report_progress(progress_cb: ProgressCallback | None, progress: float, stage: str) -> None:
|
||||
"""上报进度。"""
|
||||
if progress_cb is not None:
|
||||
try:
|
||||
progress_cb(progress, stage)
|
||||
except Exception:
|
||||
logger.exception("进度回调失败")
|
||||
|
||||
@staticmethod
|
||||
def _download_assets(clips: list[EditPlanClip], work_dir: Path) -> dict[str, Path]:
|
||||
"""下载片段素材到本地,返回 asset_id → local_path 映射。
|
||||
|
||||
只保留下载成功的素材。
|
||||
"""
|
||||
asset_dir = work_dir / "assets"
|
||||
asset_dir.mkdir(exist_ok=True)
|
||||
|
||||
asset_path_map: dict[str, Path] = {}
|
||||
|
||||
for clip in clips:
|
||||
asset_id = clip.asset_id
|
||||
if not asset_id:
|
||||
continue
|
||||
|
||||
# 生成安全的本地文件名
|
||||
safe_name = f"clip_{clip.order:04d}_{abs(hash(asset_id)) % 100000:05d}.mp4"
|
||||
local_path = asset_dir / safe_name
|
||||
|
||||
if download_asset(asset_id, local_path):
|
||||
asset_path_map[asset_id] = local_path
|
||||
logger.debug("素材下载成功: clip_id=%s asset_id=%s", clip.id, asset_id[:60])
|
||||
else:
|
||||
logger.warning("素材下载失败: clip_id=%s asset_id=%s", clip.id, asset_id[:60])
|
||||
|
||||
return asset_path_map
|
||||
+204
@@ -0,0 +1,204 @@
|
||||
"""渲染引擎 Feature Flag 解析器。
|
||||
|
||||
封装渲染引擎选择逻辑,支持:
|
||||
- 环境变量作为默认值(RENDER_ENGINE=legacy/unified)
|
||||
- Redis Feature Flag 运行时覆盖(白名单 + 百分比 + 全局开关)
|
||||
- 定时刷新,支持热更新不重启 worker
|
||||
|
||||
使用方式:
|
||||
resolver = RenderEngineResolver(redis_url="redis://...", default_engine="legacy")
|
||||
engine = resolver.get_engine(user_id="user123")
|
||||
# engine: "legacy" 或 "unified"
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
from typing import Optional
|
||||
|
||||
from packages.adapters.redis.feature_flag_store import (
|
||||
FeatureFlagConfig,
|
||||
FeatureFlagStore,
|
||||
InMemoryFeatureFlagStore,
|
||||
RedisFeatureFlagStore,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Feature Flag 名称常量
|
||||
FLAG_RENDER_ENGINE = "render_engine"
|
||||
|
||||
# 引擎常量
|
||||
ENGINE_LEGACY = "legacy"
|
||||
ENGINE_UNIFIED = "unified"
|
||||
VALID_ENGINES = {ENGINE_LEGACY, ENGINE_UNIFIED}
|
||||
|
||||
|
||||
class RenderEngineResolver:
|
||||
"""渲染引擎选择器。
|
||||
|
||||
判定逻辑(从高到低):
|
||||
1. Redis flag 白名单匹配 → unified
|
||||
2. Redis flag 百分比命中 → unified
|
||||
3. Redis flag 全局开启(100%)→ unified
|
||||
4. 环境变量默认值 → legacy / unified
|
||||
|
||||
当 Redis 不可用时,自动降级到环境变量默认值,不影响业务。
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
default_engine: str = ENGINE_LEGACY,
|
||||
redis_url: Optional[str] = None,
|
||||
refresh_interval: float = 30.0,
|
||||
store: Optional[FeatureFlagStore] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Args:
|
||||
default_engine: 环境变量默认的引擎名(legacy / unified)
|
||||
redis_url: Redis 连接 URL,传 None 时使用内存实现(测试用)
|
||||
refresh_interval: Redis flag 配置刷新间隔(秒)
|
||||
store: 直接传入 store 实例(测试用,优先级高于 redis_url)
|
||||
"""
|
||||
self._default_engine = default_engine.lower() if default_engine else ENGINE_LEGACY
|
||||
if self._default_engine not in VALID_ENGINES:
|
||||
logger.warning(
|
||||
"Invalid default engine '%s', fallback to '%s'",
|
||||
self._default_engine,
|
||||
ENGINE_LEGACY,
|
||||
)
|
||||
self._default_engine = ENGINE_LEGACY
|
||||
|
||||
if store is not None:
|
||||
self._store = store
|
||||
elif redis_url:
|
||||
self._store = RedisFeatureFlagStore(redis_url=redis_url)
|
||||
else:
|
||||
self._store = InMemoryFeatureFlagStore()
|
||||
logger.info("No Redis configured, using in-memory feature flag store")
|
||||
|
||||
self._refresh_interval = refresh_interval
|
||||
self._lock = threading.Lock()
|
||||
self._cached_config: Optional[FeatureFlagConfig] = None
|
||||
self._last_refresh: float = 0.0
|
||||
|
||||
def _maybe_refresh(self) -> None:
|
||||
"""惰性刷新配置,超过刷新间隔时从存储重新读取。"""
|
||||
import time
|
||||
|
||||
now = time.time()
|
||||
if now - self._last_refresh < self._refresh_interval:
|
||||
return
|
||||
|
||||
try:
|
||||
config = self._store.get(FLAG_RENDER_ENGINE)
|
||||
with self._lock:
|
||||
self._cached_config = config
|
||||
self._last_refresh = now
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to refresh render engine flag: %s", exc)
|
||||
# 刷新失败时保留旧缓存,不中断业务
|
||||
if self._cached_config is None:
|
||||
# 首次就读失败,设一个默认值
|
||||
with self._lock:
|
||||
self._cached_config = FeatureFlagConfig(name=FLAG_RENDER_ENGINE)
|
||||
self._last_refresh = now
|
||||
|
||||
def _get_config(self) -> FeatureFlagConfig:
|
||||
"""获取当前 flag 配置(带缓存)。"""
|
||||
if self._cached_config is None:
|
||||
self._maybe_refresh()
|
||||
else:
|
||||
self._maybe_refresh()
|
||||
return self._cached_config or FeatureFlagConfig(name=FLAG_RENDER_ENGINE)
|
||||
|
||||
def get_engine(self, user_id: Optional[str] = None) -> str:
|
||||
"""获取当前应该使用的渲染引擎。
|
||||
|
||||
Args:
|
||||
user_id: 用户ID,用于白名单匹配和百分比哈希。
|
||||
传 None 时只看全局开关。
|
||||
|
||||
Returns:
|
||||
"legacy" 或 "unified"
|
||||
"""
|
||||
config = self._get_config()
|
||||
|
||||
# 全局关闭 → 用默认值
|
||||
if not config.enabled:
|
||||
return self._default_engine
|
||||
|
||||
# 白名单匹配 / 百分比命中 → unified
|
||||
if config.is_active(user_id):
|
||||
return ENGINE_UNIFIED
|
||||
|
||||
# 未命中灰度 → 用默认值
|
||||
return self._default_engine
|
||||
|
||||
def should_use_unified(self, user_id: Optional[str] = None) -> bool:
|
||||
"""便捷方法:是否应该使用统一渲染引擎。"""
|
||||
return self.get_engine(user_id) == ENGINE_UNIFIED
|
||||
|
||||
def force_refresh(self) -> None:
|
||||
"""强制立即刷新配置(用于管理接口修改后立即生效)。"""
|
||||
self._last_refresh = 0.0
|
||||
if isinstance(self._store, RedisFeatureFlagStore):
|
||||
self._store.invalidate_cache(FLAG_RENDER_ENGINE)
|
||||
self._maybe_refresh()
|
||||
|
||||
def get_config_snapshot(self) -> dict:
|
||||
"""获取当前配置快照(用于管理接口展示)。"""
|
||||
config = self._get_config()
|
||||
return {
|
||||
"flag_name": FLAG_RENDER_ENGINE,
|
||||
"default_engine": self._default_engine,
|
||||
"enabled": config.enabled,
|
||||
"percentage": config.percentage,
|
||||
"whitelist": sorted(config.whitelist),
|
||||
"refresh_interval": self._refresh_interval,
|
||||
"last_refresh": self._last_refresh,
|
||||
}
|
||||
|
||||
def set_flag(self, config: FeatureFlagConfig) -> None:
|
||||
"""设置 flag 配置(管理接口用)。"""
|
||||
config.name = FLAG_RENDER_ENGINE
|
||||
self._store.set(config)
|
||||
self.force_refresh()
|
||||
|
||||
|
||||
# 全局单例
|
||||
_resolver: Optional[RenderEngineResolver] = None
|
||||
_resolver_lock = threading.Lock()
|
||||
|
||||
|
||||
def get_render_engine_resolver() -> RenderEngineResolver:
|
||||
"""获取全局单例(基于 worker 配置)。"""
|
||||
global _resolver
|
||||
if _resolver is not None:
|
||||
return _resolver
|
||||
|
||||
with _resolver_lock:
|
||||
if _resolver is not None:
|
||||
return _resolver
|
||||
|
||||
try:
|
||||
from worker_app.core.config import get_settings
|
||||
|
||||
settings = get_settings()
|
||||
redis_url = getattr(settings, "redis_url", None) or getattr(settings, "broker_url", None)
|
||||
default = getattr(settings, "render_engine", ENGINE_LEGACY)
|
||||
_resolver = RenderEngineResolver(
|
||||
default_engine=default,
|
||||
redis_url=redis_url,
|
||||
)
|
||||
logger.info(
|
||||
"RenderEngineResolver initialized: default=%s, redis=%s",
|
||||
default,
|
||||
bool(redis_url),
|
||||
)
|
||||
except Exception as exc:
|
||||
logger.warning("Failed to init RenderEngineResolver from settings: %s", exc)
|
||||
_resolver = RenderEngineResolver(default_engine=ENGINE_LEGACY)
|
||||
|
||||
return _resolver
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,821 +0,0 @@
|
||||
"""
|
||||
视频合成服务
|
||||
支持多种剪辑模式和转场效果,包含完整的安全校验
|
||||
"""
|
||||
|
||||
import logging
|
||||
import os
|
||||
import subprocess
|
||||
import tempfile
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
|
||||
try:
|
||||
from enum import StrEnum
|
||||
except ImportError:
|
||||
|
||||
class StrEnum(str, Enum): # type: ignore[no-redef]
|
||||
"""Python 3.10 兼容的 StrEnum 回退实现。"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Optional
|
||||
|
||||
from packages.domain.editing_mode import EditingMode
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ========== 安全常量 ==========
|
||||
# 允许的输出目录白名单(使用环境变量或系统临时目录,避免硬编码 /tmp)
|
||||
_VIDEO_OUTPUT_DIR = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
|
||||
ALLOWED_OUTPUT_DIRS = [_VIDEO_OUTPUT_DIR, "/var/app/rendered"]
|
||||
|
||||
# 允许的输入路径前缀白名单
|
||||
ALLOWED_INPUT_PREFIXES = ("s3://", "oss://", "local://", "/var/storage/")
|
||||
|
||||
# 允许的转场效果白名单
|
||||
ALLOWED_TRANSITIONS = {
|
||||
"fade",
|
||||
"slideleft",
|
||||
"slideright",
|
||||
"dissolve",
|
||||
"wipeleft",
|
||||
"wiperight",
|
||||
"cut",
|
||||
"slideup",
|
||||
"slidedown",
|
||||
}
|
||||
|
||||
# 转场效果映射
|
||||
_XFADE_TRANSITION_MAP = {
|
||||
"fade": "fade",
|
||||
"slideleft": "slideleft",
|
||||
"slideright": "slideright",
|
||||
"dissolve": "dissolve",
|
||||
"wipeleft": "wipeleft",
|
||||
"wiperight": "wiperight",
|
||||
"cut": "cut",
|
||||
"slideup": "slideup",
|
||||
"slidedown": "slidedown",
|
||||
}
|
||||
|
||||
|
||||
class VideoComposeError(Exception):
|
||||
"""视频合成服务异常"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class PIPPosition(StrEnum):
|
||||
"""画中画位置枚举"""
|
||||
|
||||
TOP_LEFT = "top_left"
|
||||
TOP_RIGHT = "top_right"
|
||||
BOTTOM_LEFT = "bottom_left"
|
||||
BOTTOM_RIGHT = "bottom_right"
|
||||
|
||||
|
||||
@dataclass
|
||||
class Clip:
|
||||
"""视频片段"""
|
||||
|
||||
asset_id: str # 资源ID,对应输入路径
|
||||
start_time: float = 0.0
|
||||
duration: float = 0.0
|
||||
transition: str = "fade" # 转场效果
|
||||
|
||||
|
||||
@dataclass
|
||||
class EditingModeConfig:
|
||||
"""剪辑模式配置"""
|
||||
|
||||
mode: EditingMode
|
||||
output_width: int = 1280
|
||||
output_height: int = 720
|
||||
output_fps: int = 25
|
||||
pip_position: PIPPosition = PIPPosition.TOP_RIGHT
|
||||
pip_scale: float = 0.25 # 画中画占主画面的比例
|
||||
transition_duration: float = 0.5 # 转场时长(秒)
|
||||
output_codec: str = "libx264"
|
||||
output_preset: str = "medium"
|
||||
output_crf: int = 23
|
||||
|
||||
|
||||
class VideoComposeService:
|
||||
"""视频合成服务"""
|
||||
|
||||
def __init__(self, config: EditingModeConfig, work_dir: Optional[str] = None):
|
||||
"""
|
||||
初始化视频合成服务
|
||||
|
||||
Args:
|
||||
config: 剪辑模式配置
|
||||
work_dir: 工作目录,默认使用系统临时目录
|
||||
"""
|
||||
self.config = config
|
||||
self.work_dir = work_dir or tempfile.gettempdir()
|
||||
self._ffmpeg_bin = "ffmpeg"
|
||||
self._ffprobe_bin = "ffprobe"
|
||||
|
||||
def _validate_output_path(self, path: str) -> str:
|
||||
"""
|
||||
校验输出路径是否在允许范围内 (P0 修复)
|
||||
|
||||
防止路径穿越攻击,如 /app/config/../../../etc/passwd
|
||||
|
||||
Args:
|
||||
path: 用户提供的输出路径
|
||||
|
||||
Returns:
|
||||
标准化后的绝对路径
|
||||
|
||||
Raises:
|
||||
ValueError: 路径不在允许范围内
|
||||
"""
|
||||
abs_path = os.path.abspath(path)
|
||||
for allowed_dir in ALLOWED_OUTPUT_DIRS:
|
||||
allowed_abs = os.path.abspath(allowed_dir)
|
||||
if abs_path.startswith(allowed_abs):
|
||||
return abs_path
|
||||
raise ValueError(f"输出路径不在允许范围内: {path}")
|
||||
|
||||
def _validate_input_path(self, path: str) -> bool:
|
||||
"""
|
||||
校验输入路径格式是否合法 (P1-1 修复)
|
||||
|
||||
Args:
|
||||
path: 输入文件路径
|
||||
|
||||
Returns:
|
||||
是否合法
|
||||
"""
|
||||
return any(path.startswith(prefix) for prefix in ALLOWED_INPUT_PREFIXES)
|
||||
|
||||
def _validate_transition(self, transition: str) -> str:
|
||||
"""
|
||||
校验转场效果是否在白名单内 (P1-2 修复)
|
||||
|
||||
Args:
|
||||
transition: 转场效果名称
|
||||
|
||||
Returns:
|
||||
安全的转场效果名称
|
||||
"""
|
||||
if transition not in ALLOWED_TRANSITIONS:
|
||||
logger.warning(f"未知的转场效果 '{transition}',使用默认 'fade'")
|
||||
return "fade"
|
||||
return transition
|
||||
|
||||
def _get_validated_transition(self, transition: str) -> str:
|
||||
"""获取白名单校验后的转场效果名称"""
|
||||
return _XFADE_TRANSITION_MAP.get(self._validate_transition(transition), "fade")
|
||||
|
||||
def compose(self, clips: list[Clip], output_path: Optional[str] = None) -> str:
|
||||
"""
|
||||
合成视频
|
||||
|
||||
Args:
|
||||
clips: 视频片段列表,每个片段包含 asset_id 和转场配置
|
||||
output_path: 输出文件路径
|
||||
|
||||
Returns:
|
||||
输出文件路径
|
||||
"""
|
||||
if not clips:
|
||||
raise ValueError("clips 不能为空")
|
||||
|
||||
# P1-1: 校验所有输入路径
|
||||
for clip in clips:
|
||||
if not self._validate_input_path(clip.asset_id):
|
||||
raise ValueError(f"不合法的输入路径: {clip.asset_id}")
|
||||
|
||||
# 生成默认输出路径并校验
|
||||
if output_path is None:
|
||||
output_path = self._generate_output_path()
|
||||
|
||||
# P0: 校验输出路径
|
||||
validated_output = self._validate_output_path(output_path)
|
||||
|
||||
logger.info(f"合成视频,片段数: {len(clips)}, 输出: {validated_output}")
|
||||
|
||||
# 获取输入路径列表
|
||||
input_paths = [clip.asset_id for clip in clips]
|
||||
|
||||
try:
|
||||
if self.config.mode == EditingMode.ONE_TAKE:
|
||||
return self._one_take(input_paths, validated_output, clips)
|
||||
elif self.config.mode == EditingMode.PIP:
|
||||
return self._pip(input_paths, validated_output)
|
||||
elif self.config.mode == EditingMode.VOICE_OVER:
|
||||
return self._voice_over(input_paths, validated_output)
|
||||
elif self.config.mode == EditingMode.VOICE_PIP:
|
||||
return self._voice_pip(input_paths, validated_output)
|
||||
else:
|
||||
raise ValueError(f"不支持的剪辑模式: {self.config.mode}")
|
||||
except Exception as e:
|
||||
logger.error(f"视频合成失败: {e}")
|
||||
raise VideoComposeError(f"视频合成失败: {e}") from e
|
||||
|
||||
def _generate_output_path(self) -> str:
|
||||
"""生成输出文件路径"""
|
||||
os.makedirs(self.work_dir, exist_ok=True)
|
||||
return os.path.join(self.work_dir, f"output_{self.config.mode}_{os.getpid()}.mp4")
|
||||
|
||||
def _validate_inputs(self, video_paths: list[str], audio_path: Optional[str] = None) -> None:
|
||||
"""验证输入文件存在"""
|
||||
for path in video_paths:
|
||||
if not os.path.exists(path):
|
||||
raise FileNotFoundError(f"视频文件不存在: {path}")
|
||||
if not os.path.getsize(path) > 0:
|
||||
raise ValueError(f"视频文件为空: {path}")
|
||||
|
||||
if audio_path and not os.path.exists(audio_path):
|
||||
raise FileNotFoundError(f"音频文件不存在: {audio_path}")
|
||||
|
||||
def _run_ffmpeg(self, command: list[str], capture_output: bool = True) -> tuple:
|
||||
"""执行 FFmpeg 命令"""
|
||||
logger.debug(f"Running FFmpeg: {' '.join(command)}")
|
||||
try:
|
||||
result = subprocess.run(
|
||||
command,
|
||||
check=True,
|
||||
stdout=subprocess.PIPE if capture_output else None,
|
||||
stderr=subprocess.PIPE if capture_output else None,
|
||||
text=capture_output,
|
||||
)
|
||||
return result.stdout or "", result.stderr or ""
|
||||
except subprocess.CalledProcessError as e:
|
||||
stderr = e.stderr.decode() if e.stderr else str(e)
|
||||
logger.error(f"FFmpeg error: {stderr}")
|
||||
raise RuntimeError(f"FFmpeg 执行失败: {stderr}") from e
|
||||
|
||||
def _get_video_info(self, video_path: str) -> dict:
|
||||
"""获取视频信息"""
|
||||
try:
|
||||
result = subprocess.run(
|
||||
[
|
||||
self._ffprobe_bin,
|
||||
"-v",
|
||||
"error",
|
||||
"-show_entries",
|
||||
"stream=width,height,r_frame_rate,duration,codec_name",
|
||||
"-show_entries",
|
||||
"format=duration,size",
|
||||
"-of",
|
||||
"json",
|
||||
video_path,
|
||||
],
|
||||
check=True,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
)
|
||||
import json
|
||||
|
||||
data = json.loads(result.stdout)
|
||||
streams = data.get("streams", [{}])
|
||||
video_stream = next((s for s in streams if s.get("codec_type") == "video"), streams[0] if streams else {})
|
||||
fmt = data.get("format", {})
|
||||
|
||||
fps_str = video_stream.get("r_frame_rate", "25/1")
|
||||
fps_parts = fps_str.split("/")
|
||||
fps = float(fps_parts[0]) / float(fps_parts[1]) if len(fps_parts) == 2 else float(fps_parts[0])
|
||||
|
||||
return {
|
||||
"width": int(video_stream.get("width", 0)),
|
||||
"height": int(video_stream.get("height", 0)),
|
||||
"fps": fps,
|
||||
"duration": float(fmt.get("duration", 0)),
|
||||
"codec": video_stream.get("codec_name", "unknown"),
|
||||
"size": int(fmt.get("size", 0)),
|
||||
}
|
||||
except Exception as e:
|
||||
logger.warning(f"获取视频信息失败 {video_path}: {e}")
|
||||
return {"width": 0, "height": 0, "fps": 25, "duration": 0, "codec": "unknown", "size": 0}
|
||||
|
||||
def _get_pip_position_offset(
|
||||
self, main_width: int, main_height: int, pip_width: int, pip_height: int
|
||||
) -> tuple[int, int]:
|
||||
"""获取画中画位置偏移量"""
|
||||
margin = 10
|
||||
position_offsets = {
|
||||
PIPPosition.TOP_LEFT: (margin, margin),
|
||||
PIPPosition.TOP_RIGHT: (main_width - pip_width - margin, margin),
|
||||
PIPPosition.BOTTOM_LEFT: (margin, main_height - pip_height - margin),
|
||||
PIPPosition.BOTTOM_RIGHT: (main_width - pip_width - margin, main_height - pip_height - margin),
|
||||
}
|
||||
return position_offsets.get(self.config.pip_position, position_offsets[PIPPosition.TOP_RIGHT])
|
||||
|
||||
def _normalize_video(self, input_path: str, output_path: str) -> dict:
|
||||
"""标准化视频格式"""
|
||||
command = [
|
||||
self._ffmpeg_bin,
|
||||
"-y",
|
||||
"-i",
|
||||
input_path,
|
||||
"-r",
|
||||
str(self.config.output_fps),
|
||||
"-vf",
|
||||
f"scale={self.config.output_width}:{self.config.output_height}:force_original_aspect_ratio=decrease,pad={self.config.output_width}:{self.config.output_height}:(ow-iw)/2:(oh-ih)/2,setsar=1",
|
||||
"-r",
|
||||
str(self.config.output_fps),
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
"-movflags",
|
||||
"+faststart",
|
||||
"-an",
|
||||
output_path,
|
||||
]
|
||||
self._run_ffmpeg(command)
|
||||
return self._get_video_info(output_path)
|
||||
|
||||
def _one_take(self, video_paths: list[str], output_path: str, clips: list[Clip]) -> str:
|
||||
"""一镜到底模式"""
|
||||
if len(video_paths) == 1:
|
||||
return self._normalize_video(video_paths[0], output_path)
|
||||
|
||||
normalized_paths = []
|
||||
for i, path in enumerate(video_paths):
|
||||
normalized = os.path.join(self.work_dir, f"normalized_{i}_{os.getpid()}.mp4")
|
||||
self._normalize_video(path, normalized)
|
||||
normalized_paths.append(normalized)
|
||||
|
||||
durations = [self._get_video_info(p)["duration"] for p in normalized_paths]
|
||||
|
||||
if len(normalized_paths) <= 5:
|
||||
output_path = self._one_take_with_xfade(normalized_paths, durations, output_path, clips)
|
||||
else:
|
||||
output_path = self._one_take_simple_concat(normalized_paths, output_path)
|
||||
|
||||
for p in normalized_paths:
|
||||
try:
|
||||
if p != output_path:
|
||||
os.remove(p)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Operation failed in apps/worker/video_processing/video_compose_service.py: {e}", exc_info=True
|
||||
)
|
||||
|
||||
return output_path
|
||||
|
||||
def _one_take_with_xfade(
|
||||
self, normalized_paths: list[str], durations: list[float], output_path: str, clips: list[Clip]
|
||||
) -> str:
|
||||
"""使用 xfade 滤镜实现转场 (P1-2: 转场参数白名单校验)"""
|
||||
if len(normalized_paths) == 2:
|
||||
# 获取当前片段的转场效果并校验白名单
|
||||
transition = "fade"
|
||||
if len(clips) > 1:
|
||||
transition = self._get_validated_transition(clips[1].transition)
|
||||
|
||||
trans_duration = self.config.transition_duration
|
||||
offset1 = durations[0] - trans_duration / 2
|
||||
|
||||
command = [
|
||||
self._ffmpeg_bin,
|
||||
"-y",
|
||||
"-i",
|
||||
normalized_paths[0],
|
||||
"-i",
|
||||
normalized_paths[1],
|
||||
"-filter_complex",
|
||||
f"[0:v][1:v]xfade=transition={transition}:duration={trans_duration}:offset={offset1}[v]",
|
||||
"-map",
|
||||
"[v]",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
output_path,
|
||||
]
|
||||
self._run_ffmpeg(command)
|
||||
return output_path
|
||||
else:
|
||||
return self._one_take_simple_concat(normalized_paths, output_path)
|
||||
|
||||
def _one_take_simple_concat(self, normalized_paths: list[str], output_path: str) -> str:
|
||||
"""使用 concat demuxer 简单拼接"""
|
||||
concat_file = os.path.join(self.work_dir, f"concat_list_{os.getpid()}.txt")
|
||||
with open(concat_file, "w") as f:
|
||||
for path in normalized_paths:
|
||||
f.write(f"file '{os.path.abspath(path)}'\n")
|
||||
|
||||
command = [
|
||||
self._ffmpeg_bin,
|
||||
"-y",
|
||||
"-f",
|
||||
"concat",
|
||||
"-safe",
|
||||
"0",
|
||||
"-i",
|
||||
concat_file,
|
||||
"-c",
|
||||
"copy",
|
||||
output_path,
|
||||
]
|
||||
self._run_ffmpeg(command)
|
||||
|
||||
try:
|
||||
os.remove(concat_file)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Operation failed in apps/worker/video_processing/video_compose_service.py: {e}", exc_info=True
|
||||
)
|
||||
|
||||
return output_path
|
||||
|
||||
def _pip(self, video_paths: list[str], output_path: str) -> str:
|
||||
"""画中画模式"""
|
||||
if not video_paths:
|
||||
raise ValueError("No video paths provided")
|
||||
|
||||
main_video = video_paths[0]
|
||||
main_normalized = os.path.join(self.work_dir, f"main_{os.getpid()}.mp4")
|
||||
main_info = self._normalize_video(main_video, main_normalized)
|
||||
|
||||
if len(video_paths) == 1:
|
||||
os.rename(main_normalized, output_path)
|
||||
return output_path
|
||||
|
||||
pip_width = int(self.config.output_width * self.config.pip_scale)
|
||||
pip_height = int(self.config.output_height * self.config.pip_scale)
|
||||
x_offset, y_offset = self._get_pip_position_offset(
|
||||
self.config.output_width, self.config.output_height, pip_width, pip_height
|
||||
)
|
||||
|
||||
pip_normalized = os.path.join(self.work_dir, f"pip_{os.getpid()}.mp4")
|
||||
pip_info = self._get_video_info(video_paths[1])
|
||||
|
||||
if pip_info["duration"] > main_info["duration"]:
|
||||
temp_pip = os.path.join(self.work_dir, f"pip_temp_{os.getpid()}.mp4")
|
||||
command = [
|
||||
self._ffmpeg_bin,
|
||||
"-y",
|
||||
"-i",
|
||||
video_paths[1],
|
||||
"-t",
|
||||
str(main_info["duration"]),
|
||||
"-vf",
|
||||
f"scale={pip_width}:{pip_height}",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
temp_pip,
|
||||
]
|
||||
self._run_ffmpeg(command)
|
||||
pip_normalized_input = temp_pip
|
||||
else:
|
||||
command = [
|
||||
self._ffmpeg_bin,
|
||||
"-y",
|
||||
"-i",
|
||||
video_paths[1],
|
||||
"-vf",
|
||||
f"scale={pip_width}:{pip_height}",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
pip_normalized,
|
||||
]
|
||||
self._run_ffmpeg(command)
|
||||
pip_normalized_input = pip_normalized
|
||||
|
||||
if main_info["duration"] > pip_info["duration"]:
|
||||
looped_pip = os.path.join(self.work_dir, f"pip_looped_{os.getpid()}.mp4")
|
||||
command = [
|
||||
self._ffmpeg_bin,
|
||||
"-y",
|
||||
"-stream_loop",
|
||||
"-1",
|
||||
"-i",
|
||||
pip_normalized_input,
|
||||
"-t",
|
||||
str(main_info["duration"]),
|
||||
"-vf",
|
||||
f"scale={pip_width}:{pip_height}",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
looped_pip,
|
||||
]
|
||||
self._run_ffmpeg(command)
|
||||
pip_normalized_input = looped_pip
|
||||
|
||||
command = [
|
||||
self._ffmpeg_bin,
|
||||
"-y",
|
||||
"-i",
|
||||
main_normalized,
|
||||
"-i",
|
||||
pip_normalized_input,
|
||||
"-filter_complex",
|
||||
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
|
||||
"-map",
|
||||
"[v]",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
output_path,
|
||||
]
|
||||
self._run_ffmpeg(command)
|
||||
|
||||
for temp_file in [main_normalized, pip_normalized]:
|
||||
if temp_file and temp_file != output_path:
|
||||
try:
|
||||
os.remove(temp_file)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Operation failed in apps/worker/video_processing/video_compose_service.py: {e}", exc_info=True
|
||||
)
|
||||
|
||||
return output_path
|
||||
|
||||
def _voice_over(self, video_paths: list[str], audio_path: str, output_path: str) -> str:
|
||||
"""口播模式"""
|
||||
if not audio_path:
|
||||
raise ValueError("audio_path is required for VOICE_OVER mode")
|
||||
|
||||
if not video_paths:
|
||||
raise ValueError("No background video provided")
|
||||
|
||||
audio_info = self._get_video_info(audio_path)
|
||||
audio_duration = audio_info["duration"]
|
||||
|
||||
bg_normalized = os.path.join(self.work_dir, f"bg_{os.getpid()}.mp4")
|
||||
bg_info = self._normalize_video(video_paths[0], bg_normalized)
|
||||
|
||||
if bg_info["duration"] < audio_duration:
|
||||
looped_bg = os.path.join(self.work_dir, f"bg_looped_{os.getpid()}.mp4")
|
||||
command = [
|
||||
self._ffmpeg_bin,
|
||||
"-y",
|
||||
"-stream_loop",
|
||||
"-1",
|
||||
"-i",
|
||||
bg_normalized,
|
||||
"-t",
|
||||
str(audio_duration),
|
||||
"-vf",
|
||||
f"scale={self.config.output_width}:{self.config.output_height}",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
looped_bg,
|
||||
]
|
||||
self._run_ffmpeg(command)
|
||||
bg_normalized = looped_bg
|
||||
elif bg_info["duration"] > audio_duration:
|
||||
temp_bg = os.path.join(self.work_dir, f"bg_trimmed_{os.getpid()}.mp4")
|
||||
command = [
|
||||
self._ffmpeg_bin,
|
||||
"-y",
|
||||
"-i",
|
||||
bg_normalized,
|
||||
"-t",
|
||||
str(audio_duration),
|
||||
"-c:v",
|
||||
"copy",
|
||||
temp_bg,
|
||||
]
|
||||
self._run_ffmpeg(command)
|
||||
bg_normalized = temp_bg
|
||||
|
||||
blurred_bg = os.path.join(self.work_dir, f"bg_blurred_{os.getpid()}.mp4")
|
||||
command = [
|
||||
self._ffmpeg_bin,
|
||||
"-y",
|
||||
"-i",
|
||||
bg_normalized,
|
||||
"-vf",
|
||||
f"boxblur=5:5,scale={self.config.output_width}:{self.config.output_height}",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
blurred_bg,
|
||||
]
|
||||
self._run_ffmpeg(command)
|
||||
|
||||
command = [
|
||||
self._ffmpeg_bin,
|
||||
"-y",
|
||||
"-i",
|
||||
blurred_bg,
|
||||
"-i",
|
||||
audio_path,
|
||||
"-filter_complex",
|
||||
"[0:v]drawbox=x=0:y=0:w=iw:h=ih:color=black@0.3:t=fill[v]",
|
||||
"-map",
|
||||
"[v]",
|
||||
"-map",
|
||||
"1:a",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
"-shortest",
|
||||
output_path,
|
||||
]
|
||||
self._run_ffmpeg(command)
|
||||
|
||||
for temp_file in [bg_normalized, blurred_bg]:
|
||||
try:
|
||||
if temp_file != output_path:
|
||||
os.remove(temp_file)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Operation failed in apps/worker/video_processing/video_compose_service.py: {e}", exc_info=True
|
||||
)
|
||||
|
||||
return output_path
|
||||
|
||||
def _voice_pip(self, video_paths: list[str], audio_path: Optional[str], output_path: str) -> str:
|
||||
"""口播+画中画模式"""
|
||||
if not video_paths:
|
||||
raise ValueError("No video paths provided")
|
||||
|
||||
if len(video_paths) == 1:
|
||||
return self._normalize_video(video_paths[0], output_path)
|
||||
|
||||
voice_video = video_paths[0]
|
||||
bg_video = video_paths[1] if len(video_paths) > 1 else video_paths[0]
|
||||
|
||||
voice_normalized = os.path.join(self.work_dir, f"voice_{os.getpid()}.mp4")
|
||||
voice_info = self._normalize_video(voice_video, voice_normalized)
|
||||
|
||||
bg_normalized = os.path.join(self.work_dir, f"bg_{os.getpid()}.mp4")
|
||||
bg_info = self._normalize_video(bg_video, bg_normalized)
|
||||
|
||||
final_duration = min(voice_info["duration"], bg_info["duration"])
|
||||
|
||||
pip_width = int(self.config.output_width * self.config.pip_scale)
|
||||
pip_height = int(self.config.output_height * self.config.pip_scale)
|
||||
x_offset, y_offset = self._get_pip_position_offset(
|
||||
self.config.output_width, self.config.output_height, pip_width, pip_height
|
||||
)
|
||||
|
||||
voice_adjusted = os.path.join(self.work_dir, f"voice_adj_{os.getpid()}.mp4")
|
||||
command = [
|
||||
self._ffmpeg_bin,
|
||||
"-y",
|
||||
"-i",
|
||||
voice_normalized,
|
||||
"-t",
|
||||
str(final_duration),
|
||||
"-vf",
|
||||
f"scale={pip_width}:{pip_height}",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
voice_adjusted,
|
||||
]
|
||||
self._run_ffmpeg(command)
|
||||
|
||||
bg_adjusted = os.path.join(self.work_dir, f"bg_adj_{os.getpid()}.mp4")
|
||||
command = [
|
||||
self._ffmpeg_bin,
|
||||
"-y",
|
||||
"-i",
|
||||
bg_normalized,
|
||||
"-t",
|
||||
str(final_duration),
|
||||
"-c:v",
|
||||
"copy",
|
||||
bg_adjusted,
|
||||
]
|
||||
self._run_ffmpeg(command)
|
||||
|
||||
if audio_path:
|
||||
command = [
|
||||
self._ffmpeg_bin,
|
||||
"-y",
|
||||
"-i",
|
||||
bg_adjusted,
|
||||
"-i",
|
||||
voice_adjusted,
|
||||
"-i",
|
||||
audio_path,
|
||||
"-filter_complex",
|
||||
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
|
||||
"-map",
|
||||
"[v]",
|
||||
"-map",
|
||||
"2:a",
|
||||
"-shortest",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
output_path,
|
||||
]
|
||||
else:
|
||||
command = [
|
||||
self._ffmpeg_bin,
|
||||
"-y",
|
||||
"-i",
|
||||
bg_adjusted,
|
||||
"-i",
|
||||
voice_adjusted,
|
||||
"-filter_complex",
|
||||
f"[0:v][1:v]overlay={x_offset}:{y_offset}[v]",
|
||||
"-map",
|
||||
"[v]",
|
||||
"-map",
|
||||
"1:a",
|
||||
"-shortest",
|
||||
"-c:v",
|
||||
self.config.output_codec,
|
||||
"-preset",
|
||||
self.config.output_preset,
|
||||
"-crf",
|
||||
str(self.config.output_crf),
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
output_path,
|
||||
]
|
||||
self._run_ffmpeg(command)
|
||||
|
||||
for temp_file in [voice_normalized, voice_adjusted, bg_normalized, bg_adjusted]:
|
||||
try:
|
||||
if temp_file != output_path:
|
||||
os.remove(temp_file)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"Operation failed in apps/worker/video_processing/video_compose_service.py: {e}", exc_info=True
|
||||
)
|
||||
|
||||
return output_path
|
||||
|
||||
|
||||
def create_compose_service(mode: str, work_dir: Optional[str] = None, **kwargs) -> VideoComposeService:
|
||||
"""便捷工厂函数:创建视频合成服务"""
|
||||
try:
|
||||
editing_mode = EditingMode(mode)
|
||||
except ValueError:
|
||||
raise ValueError(f"无效的剪辑模式: {mode}. 有效模式: {[m.value for m in EditingMode]}")
|
||||
|
||||
config = EditingModeConfig(
|
||||
mode=editing_mode,
|
||||
output_width=kwargs.get("output_width", 1280),
|
||||
output_height=kwargs.get("output_height", 720),
|
||||
output_fps=kwargs.get("output_fps", 25),
|
||||
pip_position=PIPPosition(kwargs.get("pip_position", "top_right")),
|
||||
pip_scale=kwargs.get("pip_scale", 0.25),
|
||||
transition_duration=kwargs.get("transition_duration", 0.5),
|
||||
)
|
||||
|
||||
return VideoComposeService(config=config, work_dir=work_dir)
|
||||
Regular → Executable
+4
@@ -17,6 +17,10 @@ class WorkerSettings(BaseSettings):
|
||||
database_pool_recycle: int = 3600
|
||||
environment: str = "development"
|
||||
auto_create_schema: bool = False
|
||||
redis_url: str = "redis://redis:6379/0"
|
||||
|
||||
# 渲染引擎选择:legacy=旧VideoComposeService,unified=新UnifiedRenderService
|
||||
render_engine: str = "legacy"
|
||||
|
||||
model_config = SettingsConfigDict(
|
||||
env_file=".env",
|
||||
|
||||
@@ -39,6 +39,10 @@ def _get_job_service():
|
||||
def compose_video(self, job_id: str, **kwargs):
|
||||
"""视频合成任务。
|
||||
|
||||
根据 RENDER_ENGINE 配置选择渲染引擎:
|
||||
- legacy: 旧 VideoComposeService(filter_complex 模式)
|
||||
- unified: 新 UnifiedRenderService(图层架构)
|
||||
|
||||
Args:
|
||||
job_id: JobService 中的任务 ID
|
||||
**kwargs: 来自 Job.payload 的额外参数(plan_id, output_path 等)
|
||||
@@ -56,66 +60,18 @@ def compose_video(self, job_id: str, **kwargs):
|
||||
job_service.fail_job(job_id, "Missing plan_id in job payload")
|
||||
return {"status": "error", "message": "Missing plan_id"}
|
||||
|
||||
# 标记为 running
|
||||
job_service.update_progress(job_id, progress=10.0, current_stage="初始化合成环境")
|
||||
# 判断使用哪个渲染引擎
|
||||
# 优先级:Redis Feature Flag(白名单 > 百分比) > 环境变量默认
|
||||
from video_processing.render_engine_resolver import get_render_engine_resolver
|
||||
|
||||
# 延迟导入 VideoComposeService
|
||||
from apps.api.app.services.video_compose_service import VideoComposeService
|
||||
resolver = get_render_engine_resolver()
|
||||
user_id = job.created_by_user_id or None
|
||||
engine = resolver.get_engine(user_id=user_id)
|
||||
|
||||
compose_svc = VideoComposeService(db)
|
||||
|
||||
# 校验合成条件
|
||||
job_service.update_progress(job_id, progress=20.0, current_stage="校验合成条件")
|
||||
validation = compose_svc.validate_compose(plan_id)
|
||||
if not validation.valid:
|
||||
error_msg = "; ".join(validation.errors)
|
||||
job_service.fail_job(job_id, f"合成校验失败: {error_msg}")
|
||||
return {"status": "error", "message": error_msg}
|
||||
|
||||
# 构建合成命令
|
||||
job_service.update_progress(job_id, progress=30.0, current_stage="构建 FFmpeg 命令")
|
||||
_output_dir = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
|
||||
output_path = os.path.join(_output_dir, f"{job_id}.mp4")
|
||||
compose_cmd = compose_svc.build_compose_command(plan_id, output_path)
|
||||
|
||||
# 执行 FFmpeg
|
||||
job_service.update_progress(job_id, progress=50.0, current_stage="正在执行视频合成")
|
||||
logger.info("Executing FFmpeg for job %s, plan %s", job_id, plan_id)
|
||||
|
||||
try:
|
||||
subprocess.run(
|
||||
compose_cmd.command,
|
||||
check=True,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
timeout=3600,
|
||||
)
|
||||
except subprocess.CalledProcessError as e:
|
||||
job_service.fail_job(job_id, f"FFmpeg 执行失败: {e.stderr[:500]}")
|
||||
raise
|
||||
|
||||
# 上传结果
|
||||
job_service.update_progress(job_id, progress=80.0, current_stage="上传合成结果")
|
||||
storage_key = f"rendered/{plan_id}/{job_id}.mp4"
|
||||
|
||||
from worker_app.tasks.edit_plan_generation import _upload_to_oss
|
||||
|
||||
output_url = _upload_to_oss(Path(output_path), storage_key)
|
||||
|
||||
# 更新 Job 状态为完成
|
||||
result_data = {
|
||||
"plan_id": plan_id,
|
||||
"output_path": output_path,
|
||||
"storage_key": storage_key,
|
||||
"output_url": output_url or "",
|
||||
"estimated_duration": compose_cmd.estimated_duration,
|
||||
"clip_count": len(compose_cmd.clip_chains),
|
||||
}
|
||||
job_service.complete_job(job_id, result=result_data)
|
||||
|
||||
logger.info("视频合成完成: job_id=%s, plan_id=%s", job_id, plan_id)
|
||||
return {"status": "completed", "job_id": job_id, "result": result_data}
|
||||
if engine == "unified":
|
||||
return _compose_with_unified_engine(self, job_service, job, plan_id, db)
|
||||
else:
|
||||
return _compose_with_legacy_engine(self, job_service, job, plan_id, db)
|
||||
|
||||
except self.retry_exc as exc:
|
||||
logger.warning("视频合成重试中: job_id=%s, exc=%s", job_id, exc)
|
||||
@@ -129,11 +85,145 @@ def compose_video(self, job_id: str, **kwargs):
|
||||
raise self.retry(exc=exc, countdown=60)
|
||||
finally:
|
||||
db.close()
|
||||
# 清理临时文件
|
||||
|
||||
|
||||
def _compose_with_legacy_engine(task, job_service, job, plan_id: str, db) -> dict:
|
||||
"""旧引擎渲染路径(VideoComposeService)。"""
|
||||
job_id = job.id
|
||||
|
||||
# 标记为 running
|
||||
job_service.update_progress(job_id, progress=10.0, current_stage="初始化合成环境")
|
||||
|
||||
# 延迟导入 VideoComposeService
|
||||
from apps.api.app.services.video_compose_service import VideoComposeService
|
||||
|
||||
compose_svc = VideoComposeService(db)
|
||||
|
||||
# 校验合成条件
|
||||
job_service.update_progress(job_id, progress=20.0, current_stage="校验合成条件")
|
||||
validation = compose_svc.validate_compose(plan_id)
|
||||
if not validation.valid:
|
||||
error_msg = "; ".join(validation.errors)
|
||||
job_service.fail_job(job_id, f"合成校验失败: {error_msg}")
|
||||
return {"status": "error", "message": error_msg}
|
||||
|
||||
# 构建合成命令
|
||||
job_service.update_progress(job_id, progress=30.0, current_stage="构建 FFmpeg 命令")
|
||||
_output_dir = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
|
||||
output_path = os.path.join(_output_dir, f"{job_id}.mp4")
|
||||
compose_cmd = compose_svc.build_compose_command(plan_id, output_path)
|
||||
|
||||
# 执行 FFmpeg
|
||||
job_service.update_progress(job_id, progress=50.0, current_stage="正在执行视频合成")
|
||||
logger.info("Executing FFmpeg for job %s, plan %s", job_id, plan_id)
|
||||
|
||||
try:
|
||||
subprocess.run(
|
||||
compose_cmd.command,
|
||||
check=True,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
timeout=3600,
|
||||
)
|
||||
except subprocess.CalledProcessError as e:
|
||||
job_service.fail_job(job_id, f"FFmpeg 执行失败: {e.stderr[:500]}")
|
||||
raise
|
||||
|
||||
# 上传结果
|
||||
job_service.update_progress(job_id, progress=80.0, current_stage="上传合成结果")
|
||||
storage_key = f"rendered/{plan_id}/{job_id}.mp4"
|
||||
|
||||
from worker_app.tasks.edit_plan_generation import _upload_to_oss
|
||||
|
||||
output_url = _upload_to_oss(Path(output_path), storage_key)
|
||||
|
||||
# 更新 Job 状态为完成
|
||||
result_data = {
|
||||
"plan_id": plan_id,
|
||||
"output_path": output_path,
|
||||
"storage_key": storage_key,
|
||||
"output_url": output_url or "",
|
||||
"estimated_duration": compose_cmd.estimated_duration,
|
||||
"clip_count": len(compose_cmd.clip_chains),
|
||||
"engine": "legacy",
|
||||
}
|
||||
job_service.complete_job(job_id, result=result_data)
|
||||
|
||||
logger.info("视频合成完成(legacy): job_id=%s, plan_id=%s", job_id, plan_id)
|
||||
return {"status": "completed", "job_id": job_id, "result": result_data}
|
||||
|
||||
|
||||
def _compose_with_unified_engine(task, job_service, job, plan_id: str, db) -> dict:
|
||||
"""新引擎渲染路径(UnifiedRenderService + RenderAdapter)。"""
|
||||
job_id = job.id
|
||||
|
||||
# 标记为 running
|
||||
job_service.update_progress(job_id, progress=10.0, current_stage="初始化统一渲染引擎")
|
||||
|
||||
from video_processing.render_adapter import RenderAdapter
|
||||
|
||||
adapter = RenderAdapter(db)
|
||||
|
||||
# 校验合成条件
|
||||
job_service.update_progress(job_id, progress=15.0, current_stage="校验合成条件")
|
||||
valid, errors, warnings, ready_count, total_count = adapter.validate_plan(plan_id)
|
||||
if not valid:
|
||||
error_msg = "; ".join(errors)
|
||||
job_service.fail_job(job_id, f"合成校验失败: {error_msg}")
|
||||
return {"status": "error", "message": error_msg}
|
||||
|
||||
# 进度回调
|
||||
def progress_cb(progress: float, stage: str) -> None:
|
||||
try:
|
||||
_output_dir = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
|
||||
output_path = os.path.join(_output_dir, f"{job_id}.mp4")
|
||||
if Path(output_path).exists():
|
||||
Path(output_path).unlink()
|
||||
except Exception as e:
|
||||
logger.warning(f"Operation failed in apps/worker/worker_app/tasks/compose_video.py: {e}", exc_info=True)
|
||||
job_service.update_progress(job_id, progress=progress, current_stage=stage)
|
||||
except Exception:
|
||||
logger.exception("更新进度失败")
|
||||
|
||||
# 执行渲染
|
||||
job_service.update_progress(job_id, progress=20.0, current_stage="开始渲染")
|
||||
logger.info("统一渲染引擎开始: job_id=%s plan_id=%s", job_id, plan_id)
|
||||
|
||||
result = adapter.render_plan(
|
||||
plan_id=plan_id,
|
||||
job_id=job_id,
|
||||
progress_cb=progress_cb,
|
||||
)
|
||||
|
||||
if not result.success:
|
||||
job_service.fail_job(job_id, f"渲染失败: {result.error_message}")
|
||||
raise RuntimeError(result.error_message)
|
||||
|
||||
# 更新 Job 状态为完成
|
||||
result_data = {
|
||||
"plan_id": plan_id,
|
||||
"output_path": str(result.output_path) if result.output_path else "",
|
||||
"storage_key": f"rendered/{plan_id}/{job_id}.mp4",
|
||||
"output_url": result.output_url,
|
||||
"estimated_duration": result.duration,
|
||||
"clip_count": result.clip_count,
|
||||
"engine": "unified",
|
||||
"width": result.width,
|
||||
"height": result.height,
|
||||
"file_size": result.file_size,
|
||||
}
|
||||
job_service.complete_job(job_id, result=result_data)
|
||||
|
||||
logger.info(
|
||||
"视频合成完成(unified): job_id=%s plan_id=%s duration=%.2fs",
|
||||
job_id,
|
||||
plan_id,
|
||||
result.duration,
|
||||
)
|
||||
return {"status": "completed", "job_id": job_id, "result": result_data}
|
||||
|
||||
|
||||
def _cleanup_output(job_id: str) -> None:
|
||||
"""清理临时输出文件。"""
|
||||
try:
|
||||
_output_dir = os.environ.get("VIDEO_OUTPUT_DIR", os.path.join(tempfile.gettempdir(), "video_output"))
|
||||
output_path = os.path.join(_output_dir, f"{job_id}.mp4")
|
||||
if Path(output_path).exists():
|
||||
Path(output_path).unlink()
|
||||
except Exception as e:
|
||||
logger.warning(f"清理输出文件失败: {e}", exc_info=True)
|
||||
|
||||
Regular → Executable
+325
-101
@@ -1,13 +1,18 @@
|
||||
"""剪辑计划渲染任务 — Phase 8 任务 2.05.
|
||||
"""剪辑计划渲染任务 — 支持 Feature Flag 灰度.
|
||||
|
||||
Celery 任务 worker.render_edit_plan:
|
||||
1. 加载 EditPlan + EditPlanClips
|
||||
2. 下载各片段素材
|
||||
3. 使用 UnifiedRenderService 按时间线+图层渲染
|
||||
2. 根据 Feature Flag 选择渲染引擎(legacy / unified)
|
||||
3. 下载各片段素材 + 渲染
|
||||
4. 上传渲染结果到 OSS
|
||||
5. 创建 GeneratedVideo 记录 + 查重
|
||||
6. 更新 EditPlan / EditPlanClip 状态
|
||||
7. 更新 GenerationTask 进度
|
||||
|
||||
渲染引擎灰度:
|
||||
- 走 Feature Flag (render_engine) 控制
|
||||
- legacy: VideoComposeService + FFmpeg filter_complex
|
||||
- unified: UnifiedRenderService 图层架构
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -63,14 +68,268 @@ def _get_repos():
|
||||
# ── Celery Task ───────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _resolve_render_engine(user_id: str) -> str:
|
||||
"""根据 Feature Flag 决定使用哪个渲染引擎。
|
||||
|
||||
Returns:
|
||||
"legacy" 或 "unified"
|
||||
"""
|
||||
try:
|
||||
from video_processing.render_engine_resolver import get_render_engine_resolver
|
||||
|
||||
resolver = get_render_engine_resolver()
|
||||
return resolver.get_engine(user_id=user_id)
|
||||
except Exception as exc:
|
||||
logger.warning("获取渲染引擎配置失败,fallback 到 legacy: %s", exc)
|
||||
return "legacy"
|
||||
|
||||
|
||||
def _mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, error_msg: str):
|
||||
"""统一的计划失败标记工具。"""
|
||||
plan = plan_repo.get(plan_id)
|
||||
if plan and plan.status.value == "rendering":
|
||||
plan.mark_failed()
|
||||
plan_repo.update(plan)
|
||||
if generation_task_id:
|
||||
gen_task = gen_task_repo.get(generation_task_id)
|
||||
if gen_task and gen_task.status.value != "failed":
|
||||
gen_task.status = "failed"
|
||||
gen_task.error_message = error_msg
|
||||
gen_task.completed_at = datetime.now(timezone.utc)
|
||||
gen_task_repo.update(gen_task)
|
||||
|
||||
|
||||
def _finalize_render_success(
|
||||
plan,
|
||||
plan_repo,
|
||||
clip_repo,
|
||||
gen_task_repo,
|
||||
db,
|
||||
plan_id: str,
|
||||
output_url: str,
|
||||
storage_key: str,
|
||||
duration: float,
|
||||
file_size: int,
|
||||
width: int,
|
||||
height: int,
|
||||
rendered_clip_ids: list[str],
|
||||
failed_clip_ids: list[str],
|
||||
generation_task_id: str,
|
||||
output_path: Path,
|
||||
engine: str,
|
||||
) -> dict:
|
||||
"""渲染成功后的统一收尾:查重 + 更新状态 + 返回结果。"""
|
||||
# 创建 GeneratedVideo 记录 + 查重
|
||||
project_id = plan.project_id or ""
|
||||
batch_id = plan.config.get("batch_id", "")
|
||||
mode = plan.config.get("mode", "edit_plan")
|
||||
if generation_task_id and project_id:
|
||||
try:
|
||||
create_video_record_and_dedup(
|
||||
generation_task_id=generation_task_id,
|
||||
project_id=project_id,
|
||||
batch_id=batch_id,
|
||||
file_url=output_url or "",
|
||||
file_size=file_size,
|
||||
duration=duration,
|
||||
video_path=str(output_path),
|
||||
mode=mode,
|
||||
session=db,
|
||||
width=width,
|
||||
height=height,
|
||||
fps=OUTPUT_FPS,
|
||||
)
|
||||
except Exception as dedup_err:
|
||||
logger.warning("查重失败(不影响渲染结果): %s", dedup_err)
|
||||
|
||||
# 更新片段状态为 rendered
|
||||
for clip_id in rendered_clip_ids:
|
||||
clip = clip_repo.get(clip_id)
|
||||
if clip and clip.status.value == "ready":
|
||||
clip.mark_rendered()
|
||||
clip_repo.update(clip)
|
||||
|
||||
# 更新 EditPlan 状态为 completed
|
||||
plan.config["rendered_url"] = output_url or ""
|
||||
plan.config["rendered_storage_key"] = storage_key
|
||||
plan.mark_completed()
|
||||
plan_repo.update(plan)
|
||||
|
||||
# 更新 GenerationTask 状态为 completed
|
||||
if generation_task_id:
|
||||
gen_task = gen_task_repo.get(generation_task_id)
|
||||
if gen_task:
|
||||
gen_task.status = "completed"
|
||||
gen_task.progress = 100.0
|
||||
gen_task.result_count = len(rendered_clip_ids)
|
||||
gen_task.completed_at = datetime.now(timezone.utc)
|
||||
gen_task_repo.update(gen_task)
|
||||
|
||||
logger.info(
|
||||
"剪辑计划渲染完成: plan_id=%s engine=%s rendered=%d failed=%d duration=%.1fs",
|
||||
plan_id,
|
||||
engine,
|
||||
len(rendered_clip_ids),
|
||||
len(failed_clip_ids),
|
||||
duration,
|
||||
)
|
||||
|
||||
return {
|
||||
"status": "completed",
|
||||
"plan_id": plan_id,
|
||||
"rendered_count": len(rendered_clip_ids),
|
||||
"failed_count": len(failed_clip_ids),
|
||||
"output_url": output_url,
|
||||
"duration": duration,
|
||||
}
|
||||
|
||||
|
||||
def _render_with_unified(
|
||||
plan,
|
||||
clips,
|
||||
asset_path_map: dict[str, Path],
|
||||
tmpdir_path: Path,
|
||||
rendered_clip_ids: list[str],
|
||||
plan_id: str,
|
||||
generation_task_id: str,
|
||||
plan_repo,
|
||||
clip_repo,
|
||||
gen_task_repo,
|
||||
db,
|
||||
) -> dict:
|
||||
"""统一渲染引擎路径(UnifiedRenderService 图层架构)。"""
|
||||
render_service = UnifiedRenderService(
|
||||
plan=plan,
|
||||
clips=clips,
|
||||
asset_path_map=asset_path_map,
|
||||
work_dir=tmpdir_path,
|
||||
output_width=OUTPUT_WIDTH,
|
||||
output_height=OUTPUT_HEIGHT,
|
||||
output_fps=int(OUTPUT_FPS),
|
||||
)
|
||||
|
||||
try:
|
||||
render_result = render_service.render()
|
||||
except Exception as render_err:
|
||||
logger.error("渲染失败(unified): %s — %s", plan_id, render_err)
|
||||
_mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, f"渲染失败: {render_err}")
|
||||
return {"status": "error", "message": f"渲染失败: {render_err}"}
|
||||
|
||||
output_path = render_result.output_path
|
||||
|
||||
# 上传到 OSS
|
||||
storage_key = f"rendered/{plan_id}/output.mp4"
|
||||
output_url = upload_to_oss(output_path, storage_key)
|
||||
|
||||
failed_clip_ids: list[str] = []
|
||||
return _finalize_render_success(
|
||||
plan=plan,
|
||||
plan_repo=plan_repo,
|
||||
clip_repo=clip_repo,
|
||||
gen_task_repo=gen_task_repo,
|
||||
db=db,
|
||||
plan_id=plan_id,
|
||||
output_url=output_url or "",
|
||||
storage_key=storage_key,
|
||||
duration=render_result.duration,
|
||||
file_size=render_result.file_size,
|
||||
width=render_result.width,
|
||||
height=render_result.height,
|
||||
rendered_clip_ids=rendered_clip_ids,
|
||||
failed_clip_ids=failed_clip_ids,
|
||||
generation_task_id=generation_task_id,
|
||||
output_path=output_path,
|
||||
engine="unified",
|
||||
)
|
||||
|
||||
|
||||
def _render_with_legacy(
|
||||
plan,
|
||||
clips,
|
||||
rendered_clip_ids: list[str],
|
||||
failed_clip_ids: list[str],
|
||||
tmpdir_path: Path,
|
||||
plan_id: str,
|
||||
generation_task_id: str,
|
||||
plan_repo,
|
||||
clip_repo,
|
||||
gen_task_repo,
|
||||
db,
|
||||
) -> dict:
|
||||
"""旧引擎路径(VideoComposeService + FFmpeg filter_complex)。"""
|
||||
import os
|
||||
import subprocess
|
||||
|
||||
from apps.api.app.services.video_compose_service import VideoComposeService
|
||||
|
||||
compose_svc = VideoComposeService(db)
|
||||
|
||||
# 校验合成条件
|
||||
validation = compose_svc.validate_compose(plan_id)
|
||||
if not validation.valid:
|
||||
error_msg = "; ".join(validation.errors)
|
||||
logger.error("合成校验失败(legacy): %s — %s", plan_id, error_msg)
|
||||
_mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, f"合成校验失败: {error_msg}")
|
||||
return {"status": "error", "message": error_msg}
|
||||
|
||||
# 构建 FFmpeg 命令
|
||||
output_dir = os.environ.get("VIDEO_OUTPUT_DIR", str(tmpdir_path))
|
||||
output_path = Path(output_dir) / f"{plan_id}.mp4"
|
||||
compose_cmd = compose_svc.build_compose_command(plan_id, str(output_path))
|
||||
|
||||
logger.info("执行 FFmpeg (legacy): plan_id=%s", plan_id)
|
||||
try:
|
||||
subprocess.run(
|
||||
compose_cmd.command,
|
||||
check=True,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
timeout=3600,
|
||||
)
|
||||
except subprocess.CalledProcessError as e:
|
||||
error_msg = f"FFmpeg 执行失败: {e.stderr[:500]}"
|
||||
logger.error("FFmpeg 执行失败(legacy): %s — %s", plan_id, error_msg)
|
||||
_mark_plan_failed(plan_repo, plan_id, gen_task_repo, generation_task_id, error_msg)
|
||||
return {"status": "error", "message": error_msg}
|
||||
|
||||
# 获取文件大小
|
||||
file_size = output_path.stat().st_size if output_path.exists() else 0
|
||||
duration = compose_cmd.estimated_duration or 0.0
|
||||
|
||||
# 上传到 OSS
|
||||
storage_key = f"rendered/{plan_id}/output.mp4"
|
||||
output_url = upload_to_oss(output_path, storage_key)
|
||||
|
||||
return _finalize_render_success(
|
||||
plan=plan,
|
||||
plan_repo=plan_repo,
|
||||
clip_repo=clip_repo,
|
||||
gen_task_repo=gen_task_repo,
|
||||
db=db,
|
||||
plan_id=plan_id,
|
||||
output_url=output_url or "",
|
||||
storage_key=storage_key,
|
||||
duration=duration,
|
||||
file_size=file_size,
|
||||
width=OUTPUT_WIDTH,
|
||||
height=OUTPUT_HEIGHT,
|
||||
rendered_clip_ids=rendered_clip_ids,
|
||||
failed_clip_ids=failed_clip_ids,
|
||||
generation_task_id=generation_task_id,
|
||||
output_path=output_path,
|
||||
engine="legacy",
|
||||
)
|
||||
|
||||
|
||||
@celery_app.task(name="worker.render_edit_plan", bind=True, max_retries=2)
|
||||
def render_edit_plan(self, plan_id: str) -> dict:
|
||||
"""渲染剪辑计划
|
||||
|
||||
流程:
|
||||
1. 加载 EditPlan + EditPlanClips
|
||||
2. 下载各片段素材到临时目录,构建 asset_path_map
|
||||
3. 使用 UnifiedRenderService 按时间线+图层渲染
|
||||
2. 根据 Feature Flag 选择渲染引擎(legacy / unified)
|
||||
3. 下载素材 + 渲染
|
||||
4. 上传渲染结果到 OSS
|
||||
5. 创建 GeneratedVideo 记录 + 查重
|
||||
6. 更新 EditPlan → completed, EditPlanClips → rendered
|
||||
@@ -79,6 +338,7 @@ def render_edit_plan(self, plan_id: str) -> dict:
|
||||
logger.info("开始渲染剪辑计划: plan_id=%s", plan_id)
|
||||
|
||||
generation_task_id = ""
|
||||
engine = "legacy"
|
||||
|
||||
for repos in _get_repos():
|
||||
plan_repo, clip_repo, gen_task_repo, db = repos
|
||||
@@ -93,7 +353,12 @@ def render_edit_plan(self, plan_id: str) -> dict:
|
||||
# 获取 generation_task_id(提前读取,确保 except 块可用)
|
||||
generation_task_id = plan.config.get("generation_task_id", "")
|
||||
|
||||
# 2. 加载片段列表(按 order 排序)
|
||||
# 2. 选择渲染引擎(Feature Flag 灰度控制)
|
||||
user_id = plan.created_by_user_id or ""
|
||||
engine = _resolve_render_engine(user_id)
|
||||
logger.info("剪辑计划渲染引擎: plan_id=%s engine=%s user_id=%s", plan_id, engine, user_id)
|
||||
|
||||
# 3. 加载片段列表(按 order 排序)
|
||||
clips = clip_repo.list_by_plan(plan_id, skip=0, limit=10000)
|
||||
if not clips:
|
||||
logger.warning("剪辑计划没有片段: %s", plan_id)
|
||||
@@ -116,6 +381,15 @@ def render_edit_plan(self, plan_id: str) -> dict:
|
||||
rendered_clip_ids: list[str] = []
|
||||
failed_clip_ids: list[str] = []
|
||||
|
||||
# 预先批量查询所有素材的 storage_key(file_url)
|
||||
from packages.adapters.sqlalchemy_impl.models import AssetModel
|
||||
|
||||
clip_asset_ids = [c.asset_id for c in clips if c.asset_id]
|
||||
asset_storage_map: dict[str, str] = {}
|
||||
if clip_asset_ids:
|
||||
assets = db.query(AssetModel).filter(AssetModel.id.in_(clip_asset_ids)).all()
|
||||
asset_storage_map = {a.id: a.file_url for a in assets if a.file_url}
|
||||
|
||||
for clip in clips:
|
||||
if not clip.asset_id:
|
||||
# 没有素材的片段跳过,标记为失败
|
||||
@@ -129,10 +403,22 @@ def render_edit_plan(self, plan_id: str) -> dict:
|
||||
rendered_clip_ids.append(clip.id)
|
||||
continue
|
||||
|
||||
storage_key = asset_storage_map.get(clip.asset_id)
|
||||
if not storage_key:
|
||||
logger.warning(
|
||||
"片段素材无 storage_key,跳过: clip_id=%s asset_id=%s",
|
||||
clip.id,
|
||||
clip.asset_id,
|
||||
)
|
||||
clip.mark_failed()
|
||||
clip_repo.update(clip)
|
||||
failed_clip_ids.append(clip.id)
|
||||
continue
|
||||
|
||||
# 下载素材
|
||||
ext = Path(clip.asset_id).suffix or ".mp4"
|
||||
ext = Path(storage_key).suffix or ".mp4"
|
||||
local_path = tmpdir_path / f"clip_{clip.order:04d}{ext}"
|
||||
if download_asset(clip.asset_id, local_path):
|
||||
if download_asset(storage_key, local_path):
|
||||
asset_path_map[clip.asset_id] = local_path
|
||||
rendered_clip_ids.append(clip.id)
|
||||
else:
|
||||
@@ -153,100 +439,38 @@ def render_edit_plan(self, plan_id: str) -> dict:
|
||||
gen_task_repo.update(gen_task)
|
||||
return {"status": "error", "message": "所有片段素材下载失败"}
|
||||
|
||||
# 4. 使用 UnifiedRenderService 渲染
|
||||
render_service = UnifiedRenderService(
|
||||
plan=plan,
|
||||
clips=clips,
|
||||
asset_path_map=asset_path_map,
|
||||
work_dir=tmpdir_path,
|
||||
output_width=OUTPUT_WIDTH,
|
||||
output_height=OUTPUT_HEIGHT,
|
||||
output_fps=int(OUTPUT_FPS),
|
||||
)
|
||||
# 4. 根据引擎选择渲染方式
|
||||
if engine == "unified":
|
||||
result = _render_with_unified(
|
||||
plan=plan,
|
||||
clips=clips,
|
||||
asset_path_map=asset_path_map,
|
||||
tmpdir_path=tmpdir_path,
|
||||
rendered_clip_ids=rendered_clip_ids,
|
||||
plan_id=plan_id,
|
||||
generation_task_id=generation_task_id,
|
||||
plan_repo=plan_repo,
|
||||
clip_repo=clip_repo,
|
||||
gen_task_repo=gen_task_repo,
|
||||
db=db,
|
||||
)
|
||||
else:
|
||||
result = _render_with_legacy(
|
||||
plan=plan,
|
||||
clips=clips,
|
||||
rendered_clip_ids=rendered_clip_ids,
|
||||
failed_clip_ids=failed_clip_ids,
|
||||
tmpdir_path=tmpdir_path,
|
||||
plan_id=plan_id,
|
||||
generation_task_id=generation_task_id,
|
||||
plan_repo=plan_repo,
|
||||
clip_repo=clip_repo,
|
||||
gen_task_repo=gen_task_repo,
|
||||
db=db,
|
||||
)
|
||||
|
||||
try:
|
||||
render_result = render_service.render()
|
||||
except Exception as render_err:
|
||||
logger.error("渲染失败: %s — %s", plan_id, render_err)
|
||||
plan.mark_failed()
|
||||
plan_repo.update(plan)
|
||||
if generation_task_id:
|
||||
gen_task = gen_task_repo.get(generation_task_id)
|
||||
if gen_task:
|
||||
gen_task.status = "failed"
|
||||
gen_task.error_message = f"渲染失败: {render_err}"
|
||||
gen_task.completed_at = datetime.now(timezone.utc)
|
||||
gen_task_repo.update(gen_task)
|
||||
return {"status": "error", "message": f"渲染失败: {render_err}"}
|
||||
|
||||
output_path = render_result.output_path
|
||||
|
||||
# 5. 上传到 OSS
|
||||
storage_key = f"rendered/{plan_id}/output.mp4"
|
||||
output_url = upload_to_oss(output_path, storage_key)
|
||||
|
||||
# 6. 创建 GeneratedVideo 记录 + 查重
|
||||
project_id = plan.project_id or ""
|
||||
batch_id = plan.config.get("batch_id", "")
|
||||
mode = plan.config.get("mode", "edit_plan")
|
||||
if generation_task_id and project_id:
|
||||
try:
|
||||
create_video_record_and_dedup(
|
||||
generation_task_id=generation_task_id,
|
||||
project_id=project_id,
|
||||
batch_id=batch_id,
|
||||
file_url=output_url or "",
|
||||
file_size=render_result.file_size,
|
||||
duration=render_result.duration,
|
||||
video_path=str(output_path),
|
||||
mode=mode,
|
||||
session=db,
|
||||
width=render_result.width,
|
||||
height=render_result.height,
|
||||
fps=OUTPUT_FPS,
|
||||
)
|
||||
except Exception as dedup_err:
|
||||
logger.warning("查重失败(不影响渲染结果): %s", dedup_err)
|
||||
|
||||
# 7. 更新片段状态为 rendered
|
||||
for clip_id in rendered_clip_ids:
|
||||
clip = clip_repo.get(clip_id)
|
||||
if clip and clip.status.value == "ready":
|
||||
clip.mark_rendered()
|
||||
clip_repo.update(clip)
|
||||
|
||||
# 8. 更新 EditPlan 状态为 completed
|
||||
plan.config["rendered_url"] = output_url or ""
|
||||
plan.config["rendered_storage_key"] = storage_key
|
||||
plan.mark_completed()
|
||||
plan_repo.update(plan)
|
||||
|
||||
# 9. 更新 GenerationTask 状态为 completed
|
||||
if generation_task_id:
|
||||
gen_task = gen_task_repo.get(generation_task_id)
|
||||
if gen_task:
|
||||
gen_task.status = "completed"
|
||||
gen_task.progress = 100.0
|
||||
gen_task.result_count = len(rendered_clip_ids)
|
||||
gen_task.completed_at = datetime.now(timezone.utc)
|
||||
gen_task_repo.update(gen_task)
|
||||
|
||||
logger.info(
|
||||
"剪辑计划渲染完成: plan_id=%s rendered=%d failed=%d duration=%.1fs",
|
||||
plan_id,
|
||||
len(rendered_clip_ids),
|
||||
len(failed_clip_ids),
|
||||
render_result.duration,
|
||||
)
|
||||
|
||||
return {
|
||||
"status": "completed",
|
||||
"plan_id": plan_id,
|
||||
"rendered_count": len(rendered_clip_ids),
|
||||
"failed_count": len(failed_clip_ids),
|
||||
"output_url": output_url,
|
||||
"duration": render_result.duration,
|
||||
}
|
||||
result["engine"] = engine
|
||||
return result
|
||||
|
||||
except Exception as exc:
|
||||
logger.exception("渲染剪辑计划异常: %s", plan_id)
|
||||
|
||||
Executable → Regular
+227
-41
@@ -113,6 +113,7 @@ from video_processing.oss_helpers import (
|
||||
get_signed_download_url,
|
||||
upload_to_oss,
|
||||
)
|
||||
from video_processing.render_engine_resolver import ENGINE_LEGACY, ENGINE_UNIFIED
|
||||
from video_processing.unified_render_service import UnifiedRenderService
|
||||
|
||||
# ── 虚拟 Plan / Clip(内存中构建,不写数据库) ────────────────────────────────
|
||||
@@ -124,6 +125,7 @@ class _VirtualPlan:
|
||||
|
||||
id: str
|
||||
name: str = ""
|
||||
config: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -161,13 +163,15 @@ def _build_plan_and_clips_from_task(
|
||||
"""
|
||||
plan = _VirtualPlan(id=task_id, name=f"Generated-{task_id[:8]}")
|
||||
|
||||
# 为每个下载路径生成合成 asset_id
|
||||
# 为每个下载路径生成合成 asset_id,并预探测素材时长
|
||||
asset_path_map: dict[str, Path] = {}
|
||||
path_to_asset_id: dict[Path, str] = {}
|
||||
path_duration: dict[Path, float] = {}
|
||||
for i, p in enumerate(downloaded_paths):
|
||||
asset_id = f"gen_{task_id[:8]}_{i:03d}{p.suffix or '.mp4'}"
|
||||
asset_path_map[asset_id] = p
|
||||
path_to_asset_id[p] = asset_id
|
||||
path_duration[p] = probe_duration(p)
|
||||
|
||||
clips: list[_VirtualClip] = []
|
||||
n = len(downloaded_paths)
|
||||
@@ -183,6 +187,7 @@ def _build_plan_and_clips_from_task(
|
||||
clip_type=clip_type,
|
||||
order=i,
|
||||
asset_id=path_to_asset_id[p],
|
||||
duration=path_duration[p],
|
||||
)
|
||||
)
|
||||
elif mode == "voice_over":
|
||||
@@ -195,6 +200,7 @@ def _build_plan_and_clips_from_task(
|
||||
clip_type="main",
|
||||
order=i,
|
||||
asset_id=path_to_asset_id[p],
|
||||
duration=path_duration[p],
|
||||
config={"role": "b_roll"},
|
||||
)
|
||||
)
|
||||
@@ -214,6 +220,7 @@ def _build_plan_and_clips_from_task(
|
||||
clip_type=clip_type,
|
||||
order=i,
|
||||
asset_id=path_to_asset_id[p],
|
||||
duration=path_duration[p],
|
||||
)
|
||||
)
|
||||
else:
|
||||
@@ -226,6 +233,7 @@ def _build_plan_and_clips_from_task(
|
||||
clip_type="main",
|
||||
order=i,
|
||||
asset_id=path_to_asset_id[p],
|
||||
duration=path_duration[p],
|
||||
)
|
||||
)
|
||||
|
||||
@@ -380,31 +388,37 @@ def _download_library_assets(
|
||||
|
||||
session = SessionLocal()
|
||||
try:
|
||||
# 构建查询:根据模式选择不同的过滤条件
|
||||
# 构建查询
|
||||
query = session.query(AssetModel).filter(
|
||||
AssetModel.status == "ready",
|
||||
AssetModel.file_type.in_(["video", "video/mp4", "video/quicktime"]),
|
||||
)
|
||||
|
||||
if asset_library_id:
|
||||
# 素材库模式
|
||||
query = query.filter(AssetModel.asset_library_id == asset_library_id)
|
||||
if asset_ids:
|
||||
# 明确指定了 asset_ids:直接按 ID 查,不预先按 library/project 过滤
|
||||
# 避免项目级素材或跨库素材因为 library_id 不匹配而查不到
|
||||
# 归属安全由后面的归属校验保证
|
||||
query = query.filter(AssetModel.id.in_(asset_ids))
|
||||
logger.info(
|
||||
"下载素材库视频: asset_library_id=%s asset_ids=%s",
|
||||
asset_library_id,
|
||||
asset_ids or "all",
|
||||
"下载指定素材: asset_ids=%d 个, asset_library_id=%s, project_id=%s",
|
||||
len(asset_ids),
|
||||
asset_library_id or "none",
|
||||
project_id or "none",
|
||||
)
|
||||
else:
|
||||
# 项目级模式
|
||||
query = query.filter(AssetModel.project_id == project_id)
|
||||
logger.info(
|
||||
"下载项目级视频: project_id=%s asset_ids=%s",
|
||||
project_id,
|
||||
asset_ids or "all",
|
||||
)
|
||||
|
||||
if asset_ids:
|
||||
query = query.filter(AssetModel.id.in_(asset_ids))
|
||||
# 未指定 asset_ids:按 library 或 project 下载全部 ready 视频
|
||||
if asset_library_id:
|
||||
query = query.filter(AssetModel.asset_library_id == asset_library_id)
|
||||
logger.info(
|
||||
"下载素材库全部视频: asset_library_id=%s",
|
||||
asset_library_id,
|
||||
)
|
||||
else:
|
||||
query = query.filter(AssetModel.project_id == project_id)
|
||||
logger.info(
|
||||
"下载项目全部视频: project_id=%s",
|
||||
project_id,
|
||||
)
|
||||
|
||||
assets = query.order_by(AssetModel.created_at).all()
|
||||
|
||||
@@ -421,13 +435,15 @@ def _download_library_assets(
|
||||
if missing_ids:
|
||||
raise ValueError(f"素材不存在: asset_ids={sorted(missing_ids)}")
|
||||
for asset in assets:
|
||||
# 校验素材库归属(只要传了 asset_library_id 就校验)
|
||||
if asset_library_id and asset.asset_library_id != asset_library_id:
|
||||
raise ValueError(
|
||||
f"素材不属于指定素材库: asset_id={asset.id}, "
|
||||
f"expected_asset_library_id={asset_library_id}, "
|
||||
f"actual_asset_library_id={asset.asset_library_id}"
|
||||
)
|
||||
if not asset_library_id and project_id and asset.project_id != project_id:
|
||||
# 校验项目归属(只要传了 project_id 就校验)
|
||||
if project_id and asset.project_id != project_id:
|
||||
raise ValueError(
|
||||
f"素材不属于指定项目: asset_id={asset.id}, "
|
||||
f"expected_project_id={project_id}, "
|
||||
@@ -558,6 +574,148 @@ def _validate_template_exists(template_id: str) -> None:
|
||||
session.close()
|
||||
|
||||
|
||||
# ── 渲染引擎选择 ─────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _resolve_render_engine(user_id: str) -> str:
|
||||
"""根据 Feature Flag 决定使用哪个渲染引擎。
|
||||
|
||||
Returns:
|
||||
"legacy" 或 "unified"
|
||||
"""
|
||||
try:
|
||||
from video_processing.render_engine_resolver import get_render_engine_resolver
|
||||
|
||||
resolver = get_render_engine_resolver()
|
||||
return resolver.get_engine(user_id=user_id)
|
||||
except Exception as exc:
|
||||
logger.warning("获取渲染引擎配置失败,fallback 到 unified: %s", exc)
|
||||
return ENGINE_UNIFIED
|
||||
|
||||
|
||||
# ── 旧引擎渲染(FFmpeg filter_complex) ────────────────────────────────────────
|
||||
|
||||
|
||||
def _render_with_legacy_engine(
|
||||
task_id: str,
|
||||
virtual_clips: list[_VirtualClip],
|
||||
asset_path_map: dict[str, Path],
|
||||
work_dir: Path,
|
||||
output_path: Path,
|
||||
) -> tuple[float, int]:
|
||||
"""旧引擎渲染路径:手动构建 FFmpeg filter_complex 命令。
|
||||
|
||||
说明:generate_video 任务使用虚拟 clips(无 EditPlan 数据库记录),
|
||||
因此无法直接复用 VideoComposeService。这里手动构建等价的 filter_complex
|
||||
命令,与旧引擎行为一致(scale → crop → setpts → trim → setpts,
|
||||
无 fps 归一化,保持原帧率)。
|
||||
|
||||
支持模式:one_take / pip / voice_over / voice_pip
|
||||
- 所有模式统一走 concat 滤镜(与旧引擎多片段逻辑一致)
|
||||
|
||||
Returns:
|
||||
(duration_seconds, file_size_bytes)
|
||||
"""
|
||||
import subprocess
|
||||
|
||||
main_clips = [
|
||||
c
|
||||
for c in virtual_clips
|
||||
if c.clip_type in ("main", "b_roll", "background")
|
||||
or (c.clip_type == "main" and c.config.get("role") == "b_roll")
|
||||
]
|
||||
if not main_clips:
|
||||
main_clips = virtual_clips[:1]
|
||||
|
||||
input_args: list[str] = []
|
||||
video_filters: list[str] = []
|
||||
audio_filters: list[str] = []
|
||||
|
||||
for i, clip in enumerate(main_clips):
|
||||
local_path = asset_path_map.get(clip.asset_id)
|
||||
if not local_path:
|
||||
continue
|
||||
input_args.extend(["-i", str(local_path)])
|
||||
|
||||
duration = clip.duration or 0.0
|
||||
|
||||
# 视频滤镜:scale → crop → setpts → trim → setpts(与旧引擎一致)
|
||||
vf = (
|
||||
f"[{i}:v]"
|
||||
f"scale={OUTPUT_WIDTH}:{OUTPUT_HEIGHT}:force_original_aspect_ratio=increase,"
|
||||
f"crop={OUTPUT_WIDTH}:{OUTPUT_HEIGHT},"
|
||||
f"setpts=PTS-STARTPTS,"
|
||||
f"trim=0:{duration:.3f},"
|
||||
f"setpts=PTS-STARTPTS"
|
||||
f"[v{i}]"
|
||||
)
|
||||
video_filters.append(vf)
|
||||
|
||||
# 音频滤镜:atrim → asetpts
|
||||
af = f"[{i}:a]atrim=0:{duration:.3f},asetpts=PTS-STARTPTS[a{i}]"
|
||||
audio_filters.append(af)
|
||||
|
||||
n = len(main_clips)
|
||||
|
||||
if n == 1:
|
||||
video_label = "[v0]"
|
||||
audio_label = "[a0]"
|
||||
else:
|
||||
# concat 视频
|
||||
v_inputs = "".join(f"[v{i}]" for i in range(n))
|
||||
video_filters.append(f"{v_inputs}concat=n={n}:v=1:a=0[outv]")
|
||||
# concat 音频
|
||||
a_inputs = "".join(f"[a{i}]" for i in range(n))
|
||||
audio_filters.append(f"{a_inputs}concat=n={n}:v=0:a=1[outa]")
|
||||
video_label = "[outv]"
|
||||
audio_label = "[outa]"
|
||||
|
||||
# 组装 filter_complex
|
||||
fc_parts = video_filters + audio_filters
|
||||
filter_complex = ";".join(fc_parts)
|
||||
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
*input_args,
|
||||
"-filter_complex",
|
||||
filter_complex,
|
||||
"-map",
|
||||
video_label,
|
||||
"-map",
|
||||
audio_label,
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-crf",
|
||||
"23",
|
||||
"-preset",
|
||||
"medium",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-b:a",
|
||||
"192k",
|
||||
"-movflags",
|
||||
"+faststart",
|
||||
str(output_path),
|
||||
]
|
||||
|
||||
logger.info("[task_id=%s] [渲染] legacy 引擎 FFmpeg 开始: clips=%d", task_id, n)
|
||||
try:
|
||||
run_ffmpeg(command)
|
||||
except subprocess.CalledProcessError as e:
|
||||
logger.error(
|
||||
"[task_id=%s] [渲染] legacy 引擎 FFmpeg 失败: %s\nfilter_complex: %s",
|
||||
task_id,
|
||||
e,
|
||||
filter_complex[:500],
|
||||
)
|
||||
raise
|
||||
|
||||
file_size = output_path.stat().st_size if output_path.exists() else 0
|
||||
duration = probe_duration(output_path)
|
||||
return duration, file_size
|
||||
|
||||
|
||||
# ── Celery Task ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@@ -715,31 +873,59 @@ def generate_video(self, task_id: str) -> dict:
|
||||
)
|
||||
_flush_logs(task_id, gen_task)
|
||||
|
||||
# 使用 UnifiedRenderService 渲染
|
||||
logger.info("[task_id=%s] [渲染] FFmpeg 渲染开始", task_id)
|
||||
# 3. 根据 Feature Flag 选择渲染引擎
|
||||
user_id = getattr(gen_task, "created_by_user_id", "") if gen_task else ""
|
||||
engine = _resolve_render_engine(user_id) if user_id else ENGINE_UNIFIED
|
||||
logger.info("[task_id=%s] [渲染] 引擎选择: %s (user_id=%s)", task_id, engine, user_id)
|
||||
|
||||
render_start = time.monotonic()
|
||||
render_service = UnifiedRenderService(
|
||||
plan=virtual_plan,
|
||||
clips=virtual_clips,
|
||||
asset_path_map=asset_path_map,
|
||||
work_dir=temp_path,
|
||||
output_width=OUTPUT_WIDTH,
|
||||
output_height=OUTPUT_HEIGHT,
|
||||
output_fps=int(OUTPUT_FPS),
|
||||
)
|
||||
render_result = render_service.render()
|
||||
render_elapsed = time.monotonic() - render_start
|
||||
logger.info(
|
||||
"[task_id=%s] [渲染] FFmpeg 渲染完成: 耗时=%.1fs",
|
||||
task_id,
|
||||
render_elapsed,
|
||||
)
|
||||
render_output_path = temp_path / f"rendered-{task_id}.mp4"
|
||||
|
||||
if engine == ENGINE_LEGACY:
|
||||
# 旧引擎:filter_complex + concat(保持原帧率,无 fps 归一化)
|
||||
render_duration, render_file_size = _render_with_legacy_engine(
|
||||
task_id=task_id,
|
||||
virtual_clips=virtual_clips,
|
||||
asset_path_map=asset_path_map,
|
||||
work_dir=temp_path,
|
||||
output_path=render_output_path,
|
||||
)
|
||||
render_elapsed = time.monotonic() - render_start
|
||||
logger.info(
|
||||
"[task_id=%s] [渲染] legacy 引擎完成: 耗时=%.1fs, 时长=%.2fs",
|
||||
task_id,
|
||||
render_elapsed,
|
||||
render_duration,
|
||||
)
|
||||
else:
|
||||
# 新引擎:UnifiedRenderService 图层架构
|
||||
logger.info("[task_id=%s] [渲染] unified 引擎 FFmpeg 渲染开始", task_id)
|
||||
render_service = UnifiedRenderService(
|
||||
plan=virtual_plan,
|
||||
clips=virtual_clips,
|
||||
asset_path_map=asset_path_map,
|
||||
work_dir=temp_path,
|
||||
output_width=OUTPUT_WIDTH,
|
||||
output_height=OUTPUT_HEIGHT,
|
||||
output_fps=int(OUTPUT_FPS),
|
||||
)
|
||||
render_result = render_service.render()
|
||||
render_output_path = render_result.output_path
|
||||
render_duration = render_result.duration
|
||||
render_file_size = render_result.file_size
|
||||
render_elapsed = time.monotonic() - render_start
|
||||
logger.info(
|
||||
"[task_id=%s] [渲染] unified 引擎完成: 耗时=%.1fs",
|
||||
task_id,
|
||||
render_elapsed,
|
||||
)
|
||||
|
||||
if gen_task:
|
||||
gen_task.append_log(
|
||||
"渲染",
|
||||
f"FFmpeg 渲染完成, 耗时={render_elapsed:.1f}s",
|
||||
f"引擎={engine}, 耗时={render_elapsed:.1f}s",
|
||||
duration=round(render_elapsed, 2),
|
||||
engine=engine,
|
||||
)
|
||||
_flush_logs(task_id, gen_task)
|
||||
|
||||
@@ -747,14 +933,14 @@ def generate_video(self, task_id: str) -> dict:
|
||||
if audio_path:
|
||||
final_path = temp_path / f"final-{task_id}.mp4"
|
||||
try:
|
||||
_mux_audio_track(render_result.output_path, audio_path, final_path)
|
||||
_mux_audio_track(render_output_path, audio_path, final_path)
|
||||
# 混音成功,使用混音后的文件
|
||||
output_path = final_path
|
||||
except Exception as mux_err:
|
||||
logger.warning("[task_id=%s] [混音] 音频混合失败,使用无音频版本: %s", task_id, mux_err)
|
||||
output_path = render_result.output_path
|
||||
output_path = render_output_path
|
||||
else:
|
||||
output_path = render_result.output_path
|
||||
output_path = render_output_path
|
||||
|
||||
file_size = output_path.stat().st_size
|
||||
duration = probe_duration(output_path)
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
"""Packages root."""
|
||||
@@ -1 +0,0 @@
|
||||
"""Adapters package for external implementations."""
|
||||
Regular → Executable
+16
-1
@@ -1,3 +1,9 @@
|
||||
from packages.adapters.redis.feature_flag_store import (
|
||||
FeatureFlagConfig,
|
||||
FeatureFlagStore,
|
||||
InMemoryFeatureFlagStore,
|
||||
RedisFeatureFlagStore,
|
||||
)
|
||||
from packages.adapters.redis.session_store import (
|
||||
NoopSessionStore,
|
||||
RedisConfig,
|
||||
@@ -5,4 +11,13 @@ from packages.adapters.redis.session_store import (
|
||||
get_session_store,
|
||||
)
|
||||
|
||||
__all__ = ["NoopSessionStore", "RedisConfig", "SessionStore", "get_session_store"]
|
||||
__all__ = [
|
||||
"FeatureFlagConfig",
|
||||
"FeatureFlagStore",
|
||||
"InMemoryFeatureFlagStore",
|
||||
"NoopSessionStore",
|
||||
"RedisConfig",
|
||||
"RedisFeatureFlagStore",
|
||||
"SessionStore",
|
||||
"get_session_store",
|
||||
]
|
||||
|
||||
+259
@@ -0,0 +1,259 @@
|
||||
"""Feature Flag 存储实现。
|
||||
|
||||
支持两种后端:
|
||||
- RedisFeatureFlagStore:生产环境使用,支持多实例共享、热更新
|
||||
- InMemoryFeatureFlagStore:测试/开发环境使用,纯内存
|
||||
|
||||
支持的 Flag 类型:
|
||||
- 全局开关(enabled: bool)
|
||||
- 白名单(whitelist: Set[str],如 user_id 列表)
|
||||
- 百分比切流(percentage: 0-100,基于标识符哈希取模)
|
||||
|
||||
判定优先级:白名单 > 百分比 > 全局开关
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Optional, Set
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# Redis key 前缀
|
||||
FEATURE_FLAG_REDIS_PREFIX = "feature_flag:"
|
||||
|
||||
|
||||
@dataclass
|
||||
class FeatureFlagConfig:
|
||||
"""单个 Feature Flag 的配置。"""
|
||||
|
||||
name: str
|
||||
enabled: bool = False
|
||||
percentage: int = 0 # 0-100
|
||||
whitelist: Set[str] = field(default_factory=set)
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"name": self.name,
|
||||
"enabled": self.enabled,
|
||||
"percentage": self.percentage,
|
||||
"whitelist": sorted(self.whitelist),
|
||||
}
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict) -> "FeatureFlagConfig":
|
||||
return cls(
|
||||
name=data["name"],
|
||||
enabled=bool(data.get("enabled", False)),
|
||||
percentage=int(data.get("percentage", 0)),
|
||||
whitelist=set(data.get("whitelist", [])),
|
||||
)
|
||||
|
||||
def is_active(self, identifier: Optional[str] = None) -> bool:
|
||||
"""判断当前 flag 是否激活。
|
||||
|
||||
判定优先级:
|
||||
1. 全局关闭 → False
|
||||
2. 白名单匹配 → True
|
||||
3. 百分比命中 → True
|
||||
4. 其他 → False
|
||||
|
||||
Args:
|
||||
identifier: 用于白名单匹配和百分比哈希的标识符(如 user_id)。
|
||||
传 None 时只看全局开关 + 百分比(百分比用随机值)。
|
||||
"""
|
||||
if not self.enabled:
|
||||
return False
|
||||
|
||||
# 白名单:精确匹配
|
||||
if identifier and identifier in self.whitelist:
|
||||
return True
|
||||
|
||||
# 百分比:0 直接 False,100 直接 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,6 +1,6 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import JSON, Boolean, Column, DateTime, Float, Integer, String, Text, UniqueConstraint, create_engine
|
||||
from sqlalchemy import JSON, Boolean, Column, DateTime, Float, Integer, String, Text, UniqueConstraint
|
||||
from sqlalchemy.orm import declarative_base
|
||||
|
||||
Base = declarative_base()
|
||||
|
||||
@@ -27,7 +27,6 @@ from .generated_videos import (
|
||||
from .generation_tasks import (
|
||||
CreateGenerationTaskCommand,
|
||||
CreateGenerationTaskUseCase,
|
||||
GetGenerationTaskUseCase,
|
||||
)
|
||||
from .ingest_jobs import SubmitIngestJobCommand, SubmitIngestJobUseCase
|
||||
from .jobs import (
|
||||
|
||||
@@ -12,10 +12,9 @@ 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, TokenType
|
||||
from packages.application.auth.jwt_service import JWTConfig, JWTService
|
||||
|
||||
|
||||
class JWTHandler:
|
||||
|
||||
@@ -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, "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
|
||||
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
|
||||
_jwt_service_instance = JWTService(JWTConfig(**kw))
|
||||
return _jwt_service_instance
|
||||
|
||||
|
||||
@@ -293,7 +293,7 @@ class LogoutUseCase:
|
||||
try:
|
||||
if request.logout_all_devices:
|
||||
# 删除所有设备的 session
|
||||
count = self.session_store.delete_all_user_sessions(request.user_id)
|
||||
self.session_store.delete_all_user_sessions(request.user_id)
|
||||
return True, None
|
||||
else:
|
||||
# 删除当前 session
|
||||
|
||||
@@ -85,8 +85,6 @@ 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, timedelta, timezone
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
from uuid import uuid4
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
"""
|
||||
|
||||
from math import ceil
|
||||
from typing import Generic, List, Optional, TypeVar
|
||||
from typing import Generic, List, TypeVar
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
@@ -7,9 +7,7 @@ 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
|
||||
|
||||
@@ -9,7 +9,6 @@ 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 +0,0 @@
|
||||
"""TTS Job application layer."""
|
||||
@@ -148,7 +148,6 @@ 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:
|
||||
"""合成单个分段并放入队列。"""
|
||||
|
||||
@@ -24,7 +24,7 @@ from packages.application.cosyvoice_service import (
|
||||
CosyVoiceError,
|
||||
CosyVoiceService,
|
||||
)
|
||||
from packages.application.tts_job.audio_merger import AudioMergeError, AudioMerger
|
||||
from packages.application.tts_job.audio_merger import 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,7 +2,6 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from typing import List, Optional
|
||||
|
||||
from packages.domain.voice_clone_profile import VoiceCloneProfile
|
||||
|
||||
@@ -10,7 +10,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any, Optional
|
||||
from typing import Optional
|
||||
|
||||
from packages.application.cosyvoice_service import (
|
||||
CosyVoiceAuthError,
|
||||
@@ -21,9 +21,8 @@ from packages.application.voice_clone.use_cases import (
|
||||
CreateVoiceCloneUseCase,
|
||||
RetryVoiceCloneUseCase,
|
||||
VoiceCloneNotFoundError,
|
||||
VoiceCloneNotRetryableError,
|
||||
)
|
||||
from packages.domain.voice_clone_profile import VoiceCloneProfile, VoiceCloneStatus
|
||||
from packages.domain.voice_clone_profile import VoiceCloneProfile
|
||||
from packages.ports.voice_clone_profile_repository import VoiceCloneProfileRepository
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -13,7 +13,6 @@ else:
|
||||
pass
|
||||
|
||||
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@ from __future__ import annotations
|
||||
|
||||
import copy
|
||||
from enum import Enum
|
||||
from typing import List, Optional
|
||||
from typing import Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
Regular → Executable
+36
@@ -134,6 +134,25 @@ class AssetStatus(StrEnum):
|
||||
PROCESSING = "processing"
|
||||
ERROR = "error"
|
||||
|
||||
@classmethod
|
||||
def _missing_(cls, value: object) -> "AssetStatus":
|
||||
"""兼容历史数据,避免枚举转换失败导致500。
|
||||
|
||||
- uploaded → READY(早期版本用 uploaded 表示上传完成)
|
||||
- 其他未知值 → READY(兜底,不阻塞业务)
|
||||
"""
|
||||
if isinstance(value, str):
|
||||
normalized = value.strip().lower()
|
||||
if normalized in ("uploaded", "success", "ok", "done", "complete"):
|
||||
return cls.READY
|
||||
if normalized in ("upload", "uploading_start", "upload_start"):
|
||||
return cls.UPLOADING
|
||||
if normalized in ("failed", "fail", "err"):
|
||||
return cls.ERROR
|
||||
if normalized in ("process", "processing", "running", "run"):
|
||||
return cls.PROCESSING
|
||||
return cls.READY
|
||||
|
||||
|
||||
class ClassificationStatus(StrEnum):
|
||||
PENDING = "pending"
|
||||
@@ -141,6 +160,23 @@ class ClassificationStatus(StrEnum):
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
|
||||
@classmethod
|
||||
def _missing_(cls, value: object) -> "ClassificationStatus":
|
||||
"""兼容历史数据,避免枚举转换失败导致500。
|
||||
|
||||
- done → COMPLETED(早期版本用 done 表示完成)
|
||||
- 其他未知值 → PENDING(兜底,不阻塞业务)
|
||||
"""
|
||||
if isinstance(value, str):
|
||||
normalized = value.strip().lower()
|
||||
if normalized in ("done", "success", "finished", "complete"):
|
||||
return cls.COMPLETED
|
||||
if normalized in ("fail", "error", "err"):
|
||||
return cls.FAILED
|
||||
if normalized in ("process", "processing", "running", "run"):
|
||||
return cls.PROCESSING
|
||||
return cls.PENDING
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class Asset:
|
||||
|
||||
@@ -24,7 +24,7 @@ from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Dict, Optional, Set
|
||||
from typing import Dict, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -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, Set
|
||||
from typing import Any, Callable, Dict, List, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
from packages.domain import AssetLibrary, AssetLibraryKind
|
||||
from packages.domain import AssetLibrary
|
||||
|
||||
|
||||
class AssetLibraryRepository(ABC):
|
||||
|
||||
@@ -30,7 +30,8 @@ fi
|
||||
# ---- Registry 配置 ----
|
||||
REGISTRY="${REGISTRY:-git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas}"
|
||||
CACHE_REGISTRY="${CACHE_REGISTRY:-git.xiaoxiajianji.com/xiaoxia/xiaoxia-saas}"
|
||||
CACHE_TAG="${CACHE_TAG:-release}"
|
||||
# 主缓存 tag:develop 分支构建时写入,所有分支读取
|
||||
CACHE_TAG_PRIMARY="${CACHE_TAG:-develop}"
|
||||
|
||||
API_IMAGE="xiaoxia-saas-api:$VERSION"
|
||||
WORKER_IMAGE="xiaoxia-saas-worker:$VERSION"
|
||||
@@ -45,6 +46,7 @@ REGISTRY_WEB="${REGISTRY}/xiaoxia-saas-web:$VERSION"
|
||||
|
||||
USE_CACHE=0
|
||||
USE_PUSH=0
|
||||
CACHE_WRITE=0
|
||||
|
||||
# 检查 buildx 和 Registry 认证
|
||||
if docker buildx version >/dev/null 2>&1; then
|
||||
@@ -54,7 +56,10 @@ if docker buildx version >/dev/null 2>&1; then
|
||||
docker buildx use default 2>/dev/null || true
|
||||
fi
|
||||
|
||||
echo "=== Building API image ==="
|
||||
# ---- 缓存读写策略(按分支隔离)----
|
||||
# 默认只读不写,防止 feature 分支污染主缓存
|
||||
# 只有 develop/main 分支才写回缓存
|
||||
BRANCH_NAME="${GITHUB_REF_NAME:-${CI_COMMIT_BRANCH:-unknown}}"
|
||||
if [ "$USE_CACHE" -eq 1 ]; then
|
||||
docker buildx build \
|
||||
--build-arg APP_VERSION="$VERSION" \
|
||||
@@ -68,6 +73,52 @@ else
|
||||
docker build --pull=false --build-arg APP_VERSION="$VERSION" -f infra/docker/api.Dockerfile -t "$API_IMAGE" -t "$API_LATEST" .
|
||||
fi
|
||||
|
||||
build_with_cache() {
|
||||
# usage: build_with_cache <image_name> <dockerfile> <extra_args...>
|
||||
IMG_NAME="$1"
|
||||
DOCKERFILE="$2"
|
||||
shift 2
|
||||
EXTRA_ARGS="$*"
|
||||
|
||||
CACHE_FROM="type=registry,ref=${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY},ignore-error=true"
|
||||
|
||||
if [ "$CACHE_WRITE" -eq 1 ]; then
|
||||
CACHE_TO="type=registry,ref=${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY},mode=max"
|
||||
echo " cache: read+write from ${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY}"
|
||||
else
|
||||
CACHE_TO=""
|
||||
echo " cache: read-only from ${CACHE_REGISTRY}/${IMG_NAME}-cache:${CACHE_TAG_PRIMARY}"
|
||||
fi
|
||||
|
||||
if [ "$USE_CACHE" -eq 1 ]; then
|
||||
if [ -n "$CACHE_TO" ]; then
|
||||
docker buildx build \
|
||||
$EXTRA_ARGS \
|
||||
--cache-from "$CACHE_FROM" \
|
||||
--cache-to "$CACHE_TO" \
|
||||
-f "$DOCKERFILE" \
|
||||
-t "$IMG_NAME:$VERSION" \
|
||||
--load \
|
||||
.
|
||||
else
|
||||
docker buildx build \
|
||||
$EXTRA_ARGS \
|
||||
--cache-from "$CACHE_FROM" \
|
||||
-f "$DOCKERFILE" \
|
||||
-t "$IMG_NAME:$VERSION" \
|
||||
--load \
|
||||
.
|
||||
fi
|
||||
else
|
||||
docker build --pull=false $EXTRA_ARGS -f "$DOCKERFILE" -t "$IMG_NAME:$VERSION" .
|
||||
fi
|
||||
}
|
||||
|
||||
echo "=== Building API image ==="
|
||||
build_with_cache "api" "infra/docker/api.Dockerfile" \
|
||||
"--build-arg APP_VERSION=$VERSION"
|
||||
docker tag "$API_IMAGE" "$API_LATEST"
|
||||
|
||||
echo "=== Building Worker image ==="
|
||||
if [ "$USE_CACHE" -eq 1 ]; then
|
||||
docker buildx build \
|
||||
@@ -83,9 +134,16 @@ else
|
||||
fi
|
||||
|
||||
echo "=== Building Web image (with buildx cache) ==="
|
||||
# 先构建前端产物
|
||||
# 先构建前端产物(使用持久化 npm 缓存卷)
|
||||
NPM_CACHE_VOLUME="xiaoxia-npm-cache"
|
||||
if ! docker volume inspect "$NPM_CACHE_VOLUME" >/dev/null 2>&1; then
|
||||
docker volume create "$NPM_CACHE_VOLUME" >/dev/null
|
||||
echo " Created npm cache volume: $NPM_CACHE_VOLUME"
|
||||
fi
|
||||
|
||||
docker run --rm \
|
||||
-v "$PWD:/workspace" \
|
||||
-v "$NPM_CACHE_VOLUME:/workspace/apps/web/node_modules" \
|
||||
-w /workspace/apps/web \
|
||||
docker.m.daocloud.io/library/node:20 \
|
||||
sh -lc "npm ci && npm run build"
|
||||
|
||||
@@ -247,3 +247,4 @@ def main() -> int:
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
|
||||
|
||||
@@ -0,0 +1,256 @@
|
||||
#!/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
|
||||
@@ -0,0 +1,219 @@
|
||||
#!/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
|
||||
@@ -0,0 +1,83 @@
|
||||
#!/usr/bin/env python3
|
||||
"""发送 CI 成功通知到飞书/项目群 webhook。"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import urllib.request
|
||||
|
||||
|
||||
def main() -> int:
|
||||
webhook = os.environ.get("CI_NOTIFY_WEBHOOK", "")
|
||||
if not webhook:
|
||||
print("未配置 CI_NOTIFY_WEBHOOK,跳过成功通知")
|
||||
print("如需启用,请在仓库 Settings -> Secrets and variables -> Actions 中添加 CI_NOTIFY_WEBHOOK")
|
||||
return 0
|
||||
|
||||
success_job = os.environ.get("SUCCESS_JOB", "Unknown Job")
|
||||
branch = os.environ.get("GITHUB_REF_NAME", "unknown")
|
||||
commit = os.environ.get("GITHUB_SHA", "unknown")[:8]
|
||||
actor = os.environ.get("GITHUB_ACTOR", "unknown")
|
||||
run_id = os.environ.get("GITHUB_RUN_ID", "unknown")
|
||||
repo = os.environ.get("GITHUB_REPOSITORY", "unknown")
|
||||
run_url = f"https://git.xiaoxiajianji.com/{repo}/actions/runs/{run_id}"
|
||||
|
||||
payload = {
|
||||
"msg_type": "interactive",
|
||||
"card": {
|
||||
"header": {
|
||||
"title": {
|
||||
"tag": "plain_text",
|
||||
"content": "✅ CI 构建成功",
|
||||
},
|
||||
"status": "green",
|
||||
},
|
||||
"elements": [
|
||||
{
|
||||
"tag": "div",
|
||||
"text": {
|
||||
"tag": "lark_md",
|
||||
"content": (
|
||||
f"**任务**: {success_job}\n"
|
||||
f"**分支**: {branch}\n"
|
||||
f"**提交**: {commit}\n"
|
||||
f"**提交者**: {actor}\n"
|
||||
f"**Run ID**: {run_id}"
|
||||
),
|
||||
},
|
||||
},
|
||||
{
|
||||
"tag": "action",
|
||||
"actions": [
|
||||
{
|
||||
"tag": "button",
|
||||
"text": {"tag": "plain_text", "content": "查看构建详情"},
|
||||
"url": run_url,
|
||||
"type": "primary",
|
||||
}
|
||||
],
|
||||
},
|
||||
],
|
||||
},
|
||||
}
|
||||
|
||||
data = json.dumps(payload).encode("utf-8")
|
||||
req = urllib.request.Request(
|
||||
webhook,
|
||||
data=data,
|
||||
headers={"Content-Type": "application/json"},
|
||||
method="POST",
|
||||
)
|
||||
try:
|
||||
with urllib.request.urlopen(req, timeout=10) as resp:
|
||||
resp.read()
|
||||
print("成功通知已发送")
|
||||
except Exception as e:
|
||||
print(f"成功通知发送失败: {e}", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -0,0 +1,136 @@
|
||||
# 灰度对比测试工具
|
||||
|
||||
用于统一渲染引擎灰度发布期间的新旧引擎对比验证。
|
||||
|
||||
## 能力
|
||||
|
||||
- **像素对比**:基于 FFmpeg SSIM + PSNR 双指标,评估视频画质差异
|
||||
- **音频对比**:基于差值音频 RMS,评估音频波形差异
|
||||
- **批量对比**:10个预设场景覆盖 P0/P1/P2 优先级
|
||||
- **HTML 报告**:可视化对比结果,包含画质、音频、性能三维度
|
||||
- **两种切换方式**:支持 engine 参数直传 或 Feature Flag 白名单切换
|
||||
|
||||
## 目录结构
|
||||
|
||||
```
|
||||
tests/render_compare/
|
||||
├── __init__.py # 包导出
|
||||
├── README.md # 本文档
|
||||
├── video_diff.py # 视频像素对比(SSIM + PSNR)
|
||||
├── audio_diff.py # 音频对比(差值 RMS)
|
||||
├── scenarios.py # 预定义对比场景(10个)
|
||||
└── runner.py # 批量对比执行器 + HTML 报告生成
|
||||
```
|
||||
|
||||
## 快速开始
|
||||
|
||||
### 环境要求
|
||||
|
||||
- FFmpeg 4.4+(需带 ssim 和 psnr 滤镜)
|
||||
- Python 3.10+
|
||||
- httpx(API 调用)
|
||||
|
||||
### 配置环境变量
|
||||
|
||||
```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 Flag(flag-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 命令是否正确
|
||||
@@ -0,0 +1,26 @@
|
||||
"""灰度对比测试工具包.
|
||||
|
||||
用于新旧渲染引擎的批量对比测试,包含:
|
||||
- video_diff: 视频像素对比(SSIM + PSNR)
|
||||
- audio_diff: 音频对比(差值 RMS)
|
||||
- scenarios: 预定义对比场景
|
||||
- runner: 批量对比执行器 + HTML 报告
|
||||
"""
|
||||
|
||||
from .audio_diff import AudioDiffResult, compute_audio_diff, extract_audio, probe_duration, probe_has_audio
|
||||
from .scenarios import SCENARIOS, CompareScenario, get_scenarios_by_priority
|
||||
from .video_diff import VideoDiffResult, compute_video_diff, save_diff_frame
|
||||
|
||||
__all__ = [
|
||||
"VideoDiffResult",
|
||||
"compute_video_diff",
|
||||
"save_diff_frame",
|
||||
"AudioDiffResult",
|
||||
"compute_audio_diff",
|
||||
"extract_audio",
|
||||
"probe_has_audio",
|
||||
"probe_duration",
|
||||
"SCENARIOS",
|
||||
"CompareScenario",
|
||||
"get_scenarios_by_priority",
|
||||
]
|
||||
@@ -0,0 +1,322 @@
|
||||
"""音频对比工具 — 基于 FFmpeg 的音频质量对比.
|
||||
|
||||
使用以下指标评估两段音频的相似度:
|
||||
1. 波形差异(RMS 差值)
|
||||
2. 频谱相似度(FFT 分帧比较)
|
||||
3. 时长差异
|
||||
|
||||
对比方式:
|
||||
- 直接对两个音频做 `ametadata=select='gt(scene\\,0.3)'` 过于复杂
|
||||
- 简化方案:用 `amerge` + `astats` 计算差值音频的 RMS
|
||||
|
||||
更精确的方案(已实现):
|
||||
- 将两轨音频做差(amix=0:weights='1 -1' → 实际上用 pan 更简单)
|
||||
- 对差值音频做 astats,获取差值的 RMS、峰值等指标
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
import shutil
|
||||
import subprocess # nosec B404
|
||||
from dataclasses import asdict, dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
FFMPEG_BIN: str = shutil.which("ffmpeg") or "ffmpeg"
|
||||
FFPROBE_BIN: str = shutil.which("ffprobe") or "ffprobe"
|
||||
|
||||
|
||||
@dataclass
|
||||
class AudioDiffResult:
|
||||
"""音频对比结果."""
|
||||
|
||||
audio_a: str
|
||||
audio_b: str
|
||||
duration_a: float
|
||||
duration_b: float
|
||||
duration_diff: float
|
||||
sample_rate_match: bool
|
||||
channels_match: bool
|
||||
diff_rms_db: float # 差值音频的 RMS(dB,越低越相似)
|
||||
diff_peak_db: float # 差值音频的峰值(dB,越低越相似)
|
||||
similarity_score: float # 综合相似度评分 [0, 1],1 = 完全一致
|
||||
passed: bool
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
def probe_duration(file_path: str) -> float:
|
||||
"""探测文件时长(秒),失败返回 0."""
|
||||
try:
|
||||
result = subprocess.run( # nosec B603
|
||||
[
|
||||
FFPROBE_BIN,
|
||||
"-v",
|
||||
"error",
|
||||
"-show_entries",
|
||||
"format=duration",
|
||||
"-of",
|
||||
"default=noprint_wrappers=1:nokey=1",
|
||||
str(file_path),
|
||||
],
|
||||
check=True,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
timeout=10,
|
||||
)
|
||||
return round(float(result.stdout.strip()), 3)
|
||||
except Exception:
|
||||
return 0.0
|
||||
|
||||
|
||||
def probe_has_audio(file_path: str | Path) -> bool:
|
||||
"""探测文件是否包含音频流."""
|
||||
try:
|
||||
result = subprocess.run( # nosec B603
|
||||
[
|
||||
FFPROBE_BIN,
|
||||
"-v",
|
||||
"error",
|
||||
"-select_streams",
|
||||
"a:0",
|
||||
"-show_entries",
|
||||
"stream=codec_type",
|
||||
"-of",
|
||||
"default=noprint_wrappers=1:nokey=1",
|
||||
str(file_path),
|
||||
],
|
||||
check=True,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
timeout=10,
|
||||
)
|
||||
return result.stdout.strip() == "audio"
|
||||
except Exception:
|
||||
return False # 探测失败保守返回 False,避免误判有音频
|
||||
|
||||
|
||||
def compute_audio_diff(
|
||||
audio_a: str | Path,
|
||||
audio_b: str | Path,
|
||||
*,
|
||||
similarity_threshold: float = 0.90,
|
||||
duration_tolerance: float = 0.1,
|
||||
) -> AudioDiffResult:
|
||||
"""计算两段音频的差异.
|
||||
|
||||
方案:用 pan 滤镜将两轨相减,对差值音频做 astats 分析。
|
||||
|
||||
Args:
|
||||
audio_a: 音频A(基线)
|
||||
audio_b: 音频B(对比)
|
||||
similarity_threshold: 相似度合格阈值
|
||||
duration_tolerance: 时长容忍度(秒)
|
||||
|
||||
Returns:
|
||||
AudioDiffResult 对比结果
|
||||
"""
|
||||
dur_a = probe_duration(str(audio_a))
|
||||
dur_b = probe_duration(str(audio_b))
|
||||
duration_diff = abs(dur_a - dur_b)
|
||||
|
||||
# 获取音频元信息
|
||||
info_a = _probe_audio_info(str(audio_a))
|
||||
info_b = _probe_audio_info(str(audio_b))
|
||||
|
||||
sample_rate_match = info_a["sample_rate"] == info_b["sample_rate"]
|
||||
channels_match = info_a["channels"] == info_b["channels"]
|
||||
|
||||
# 相减后分析差值
|
||||
# 取较短时长做对比
|
||||
min_dur = min(dur_a, dur_b)
|
||||
if min_dur <= 0:
|
||||
return AudioDiffResult(
|
||||
audio_a=str(audio_a),
|
||||
audio_b=str(audio_b),
|
||||
duration_a=dur_a,
|
||||
duration_b=dur_b,
|
||||
duration_diff=duration_diff,
|
||||
sample_rate_match=sample_rate_match,
|
||||
channels_match=channels_match,
|
||||
diff_rms_db=-999.0,
|
||||
diff_peak_db=-999.0,
|
||||
similarity_score=0.0,
|
||||
passed=False,
|
||||
)
|
||||
|
||||
# 做差值音频:a - b
|
||||
# 注意:amix 会自动按输入数归一化音量(除以N),
|
||||
# 所以 a + (-1)*b 经过 amix=inputs=2 后整体音量会减半(-6dB)。
|
||||
# 加 volume=2 补偿回来,确保差值 RMS 反映真实差异幅度。
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-i",
|
||||
str(audio_a),
|
||||
"-i",
|
||||
str(audio_b),
|
||||
"-filter_complex",
|
||||
# 第2轨反相 → amix混合 → volume=2补偿amix的自动缩放
|
||||
"[1:a]volume=-1[inv];[0:a][inv]amix=inputs=2:duration=shortest:dropout_transition=0,volume=2[diff]",
|
||||
"-map",
|
||||
"[diff]",
|
||||
"-f",
|
||||
"null",
|
||||
"-af",
|
||||
"astats=metadata=1:reset=0",
|
||||
"-",
|
||||
]
|
||||
|
||||
try:
|
||||
result = subprocess.run( # nosec B603
|
||||
command,
|
||||
check=True,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
timeout=120,
|
||||
)
|
||||
stderr = result.stderr or ""
|
||||
except subprocess.CalledProcessError as e:
|
||||
# 如果音频格式不兼容,返回失败
|
||||
return AudioDiffResult(
|
||||
audio_a=str(audio_a),
|
||||
audio_b=str(audio_b),
|
||||
duration_a=dur_a,
|
||||
duration_b=dur_b,
|
||||
duration_diff=duration_diff,
|
||||
sample_rate_match=sample_rate_match,
|
||||
channels_match=channels_match,
|
||||
diff_rms_db=999.0,
|
||||
diff_peak_db=999.0,
|
||||
similarity_score=0.0,
|
||||
passed=False,
|
||||
)
|
||||
|
||||
diff_rms_db, diff_peak_db = _parse_astats(stderr)
|
||||
|
||||
# 相似度评分:基于差值 RMS
|
||||
# 差值 RMS -60dB → 相似度 ~1.0(几乎无声差)
|
||||
# 差值 RMS -20dB → 相似度 ~0.5(有明显差异)
|
||||
# 差值 RMS 0dB → 相似度 ~0.0(完全相反)
|
||||
if diff_rms_db <= -60:
|
||||
similarity_score = 1.0
|
||||
elif diff_rms_db >= 0:
|
||||
similarity_score = 0.0
|
||||
else:
|
||||
# 线性映射:-60dB → 1.0, 0dB → 0.0
|
||||
similarity_score = max(0.0, min(1.0, 1.0 + diff_rms_db / 60.0))
|
||||
|
||||
passed = (
|
||||
duration_diff <= duration_tolerance
|
||||
and sample_rate_match
|
||||
and channels_match
|
||||
and similarity_score >= similarity_threshold
|
||||
)
|
||||
|
||||
return AudioDiffResult(
|
||||
audio_a=str(audio_a),
|
||||
audio_b=str(audio_b),
|
||||
duration_a=round(dur_a, 3),
|
||||
duration_b=round(dur_b, 3),
|
||||
duration_diff=round(duration_diff, 3),
|
||||
sample_rate_match=sample_rate_match,
|
||||
channels_match=channels_match,
|
||||
diff_rms_db=round(diff_rms_db, 2),
|
||||
diff_peak_db=round(diff_peak_db, 2),
|
||||
similarity_score=round(similarity_score, 4),
|
||||
passed=passed,
|
||||
)
|
||||
|
||||
|
||||
def _probe_audio_info(file_path: str) -> dict[str, int]:
|
||||
"""探测音频元信息."""
|
||||
try:
|
||||
result = subprocess.run( # nosec B603
|
||||
[
|
||||
FFPROBE_BIN,
|
||||
"-v",
|
||||
"error",
|
||||
"-select_streams",
|
||||
"a:0",
|
||||
"-show_entries",
|
||||
"stream=sample_rate,channels",
|
||||
"-of",
|
||||
"json",
|
||||
file_path,
|
||||
],
|
||||
check=True,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
timeout=10,
|
||||
)
|
||||
|
||||
info = json.loads(result.stdout)
|
||||
stream = info.get("streams", [{}])[0]
|
||||
return {
|
||||
"sample_rate": int(stream.get("sample_rate", 44100)),
|
||||
"channels": int(stream.get("channels", 2)),
|
||||
}
|
||||
except Exception:
|
||||
return {"sample_rate": 0, "channels": 0}
|
||||
|
||||
|
||||
def _parse_astats(stderr: str) -> tuple[float, float]:
|
||||
"""从 astats 输出中解析 RMS 和峰值.
|
||||
|
||||
astats 输出格式(在 stderr 中):
|
||||
[Parsed_astats_1 @ 0x...] Channel: 1
|
||||
[Parsed_astats_1 @ 0x...] ...
|
||||
[Parsed_astats_1 @ 0x...] Overall
|
||||
[Parsed_astats_1 @ 0x...] DC offset: 0.000000
|
||||
[Parsed_astats_1 @ 0x...] Min level: -0.123456
|
||||
[Parsed_astats_1 @ 0x...] Max level: 0.789012
|
||||
[Parsed_astats_1 @ 0x...] Peak level dB: -2.01
|
||||
[Parsed_astats_1 @ 0x...] RMS level dB: -10.56
|
||||
...
|
||||
"""
|
||||
lines = stderr.split("\n")
|
||||
rms_db = -999.0
|
||||
peak_db = -999.0
|
||||
|
||||
for line in lines:
|
||||
# 找 Overall 部分的统计(双声道时取整体值)
|
||||
rms_match = re.search(r"RMS level dB:\s*(-?\d+\.?\d*)", line)
|
||||
peak_match = re.search(r"Peak level dB:\s*(-?\d+\.?\d*)", line)
|
||||
if rms_match:
|
||||
rms_db = float(rms_match.group(1))
|
||||
if peak_match:
|
||||
peak_db = float(peak_match.group(1))
|
||||
|
||||
return rms_db, peak_db
|
||||
|
||||
|
||||
def extract_audio(video_path: str | Path, output_path: str | Path) -> Path:
|
||||
"""从视频中提取音频(AAC 格式).
|
||||
|
||||
Args:
|
||||
video_path: 视频文件路径
|
||||
output_path: 输出音频路径
|
||||
|
||||
Returns:
|
||||
输出音频文件路径
|
||||
"""
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-y",
|
||||
"-i",
|
||||
str(video_path),
|
||||
"-vn",
|
||||
"-acodec",
|
||||
"aac",
|
||||
"-b:a",
|
||||
"128k",
|
||||
str(output_path),
|
||||
]
|
||||
subprocess.run(command, check=True, capture_output=True, timeout=120) # nosec B603
|
||||
return Path(output_path)
|
||||
@@ -0,0 +1,628 @@
|
||||
"""灰度对比测试 Runner — 新旧引擎批量对比 + 报告生成.
|
||||
|
||||
使用方法:
|
||||
# 配置环境变量
|
||||
export STAGING_API_URL=https://api.staging.example.com
|
||||
export STAGING_API_KEY=your_key
|
||||
|
||||
# 运行全部 P0 场景
|
||||
python -m tests.render_compare.runner --priority P0 --output ./report/
|
||||
|
||||
# 只跑指定场景
|
||||
python -m tests.render_compare.runner --scenario simple_pass_through,subtitle_rendering
|
||||
|
||||
对比流程:
|
||||
1. 对每个场景,分别提交到 legacy 和 unified 引擎(通过 Feature Flag 白名单/百分比控制)
|
||||
- 方式A:通过内部 API 临时切换 flag(需要 admin key)
|
||||
- 方式B:提交任务时指定 engine 参数(如果 API 支持)
|
||||
2. 等待任务完成,下载输出视频
|
||||
3. 像素对比(SSIM + PSNR)+ 音频对比(差值RMS)
|
||||
4. 生成 HTML 对比报告
|
||||
|
||||
注意:默认假设 API 支持 `engine` 参数来指定渲染引擎。
|
||||
如果不支持,需要先通过内部 API 切换 Feature Flag,然后提交任务。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
# 确保项目根目录在 path 中
|
||||
sys.path.insert(0, str(Path(__file__).parent.parent.parent))
|
||||
|
||||
from .audio_diff import AudioDiffResult, compute_audio_diff
|
||||
from .scenarios import SCENARIOS, CompareScenario, get_scenarios_by_priority
|
||||
from .video_diff import VideoDiffResult, compute_video_diff
|
||||
|
||||
|
||||
@dataclass
|
||||
class ScenarioResult:
|
||||
"""单个场景的对比结果."""
|
||||
|
||||
scenario: CompareScenario
|
||||
legacy_task_id: str = ""
|
||||
unified_task_id: str = ""
|
||||
legacy_video_path: str = ""
|
||||
unified_video_path: str = ""
|
||||
legacy_duration_sec: float = 0.0
|
||||
unified_duration_sec: float = 0.0
|
||||
video_diff: VideoDiffResult | None = None
|
||||
audio_diff: AudioDiffResult | None = None
|
||||
legacy_success: bool = False
|
||||
unified_success: bool = False
|
||||
error: str = ""
|
||||
|
||||
@property
|
||||
def passed(self) -> bool:
|
||||
if not (self.legacy_success and self.unified_success):
|
||||
return False
|
||||
if self.video_diff and not self.video_diff.passed:
|
||||
return False
|
||||
if self.audio_diff and not self.audio_diff.passed:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
class StagingAPI:
|
||||
"""Staging 环境 API 客户端."""
|
||||
|
||||
def __init__(self, base_url: str, api_key: str, internal_api_key: str = ""):
|
||||
self.base_url = base_url.rstrip("/")
|
||||
self.api_key = api_key
|
||||
self.internal_api_key = internal_api_key
|
||||
self.client = httpx.Client(timeout=30.0)
|
||||
|
||||
def _headers(self, internal: bool = False) -> dict[str, str]:
|
||||
headers = {"Authorization": f"Bearer {self.api_key}"}
|
||||
if internal and self.internal_api_key:
|
||||
headers["X-API-Key"] = self.internal_api_key
|
||||
return headers
|
||||
|
||||
def submit_render_task(self, plan_payload: dict[str, Any], engine: str = "") -> str:
|
||||
"""提交渲染任务,返回 task_id.
|
||||
|
||||
Args:
|
||||
plan_payload: EditPlan payload
|
||||
engine: 可选,指定引擎("legacy" / "unified")
|
||||
|
||||
Returns:
|
||||
task_id
|
||||
"""
|
||||
url = f"{self.base_url}/api/v1/render/compose"
|
||||
payload = dict(plan_payload)
|
||||
if engine:
|
||||
payload["engine"] = engine
|
||||
resp = self.client.post(url, json=payload, headers=self._headers())
|
||||
resp.raise_for_status()
|
||||
data = resp.json()
|
||||
return data.get("task_id") or data.get("id", "")
|
||||
|
||||
def get_task_status(self, task_id: str) -> dict[str, Any]:
|
||||
"""获取任务状态."""
|
||||
url = f"{self.base_url}/api/v1/tasks/{task_id}"
|
||||
resp = self.client.get(url, headers=self._headers())
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
|
||||
def wait_for_task(self, task_id: str, timeout: float = 300.0, poll_interval: float = 3.0) -> dict[str, Any]:
|
||||
"""等待任务完成.
|
||||
|
||||
Returns:
|
||||
最终任务状态
|
||||
|
||||
Raises:
|
||||
TimeoutError: 超时
|
||||
"""
|
||||
start = time.time()
|
||||
while time.time() - start < timeout:
|
||||
status = self.get_task_status(task_id)
|
||||
state = status.get("status", "")
|
||||
if state in ("completed", "success", "done", "failed", "error"):
|
||||
return status
|
||||
time.sleep(poll_interval)
|
||||
raise TimeoutError(f"Task {task_id} timed out after {timeout}s")
|
||||
|
||||
def set_feature_flag(self, flag_name: str, enabled: bool, percentage: int = 0, whitelist: list[str] | None = None):
|
||||
"""通过内部 API 设置 Feature Flag.
|
||||
|
||||
用于不支持 engine 参数的场景,切换全局灰度比例。
|
||||
"""
|
||||
if not self.internal_api_key:
|
||||
raise ValueError("internal_api_key is required for feature flag operations")
|
||||
url = f"{self.base_url}/api/v1/internal/feature-flags/{flag_name}"
|
||||
body: dict[str, Any] = {"enabled": enabled, "percentage": percentage}
|
||||
if whitelist is not None:
|
||||
body["whitelist"] = whitelist
|
||||
resp = self.client.put(url, json=body, headers=self._headers(internal=True))
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
|
||||
def get_feature_flag(self, flag_name: str) -> dict[str, Any]:
|
||||
"""获取 Feature Flag 配置."""
|
||||
if not self.internal_api_key:
|
||||
raise ValueError("internal_api_key is required")
|
||||
url = f"{self.base_url}/api/v1/internal/feature-flags/{flag_name}"
|
||||
resp = self.client.get(url, headers=self._headers(internal=True))
|
||||
resp.raise_for_status()
|
||||
return resp.json()
|
||||
|
||||
def download_video(self, video_url: str, output_path: str | Path) -> Path:
|
||||
"""下载视频文件."""
|
||||
output_path = Path(output_path)
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
with self.client.stream("GET", video_url, timeout=60.0) as resp:
|
||||
resp.raise_for_status()
|
||||
with open(output_path, "wb") as f:
|
||||
for chunk in resp.iter_bytes():
|
||||
f.write(chunk)
|
||||
return output_path
|
||||
|
||||
|
||||
class CompareRunner:
|
||||
"""新旧引擎对比 Runner."""
|
||||
|
||||
# 全局默认阈值(唯一真实来源,所有入口统一引用)
|
||||
DEFAULT_SSIM_THRESHOLD: float = 0.95
|
||||
DEFAULT_PSNR_THRESHOLD: float = 28.0
|
||||
DEFAULT_AUDIO_SIMILARITY_THRESHOLD: float = 0.90
|
||||
DEFAULT_DURATION_TOLERANCE: float = 0.1
|
||||
DEFAULT_TASK_TIMEOUT: float = 300.0
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
api: StagingAPI,
|
||||
output_dir: Path,
|
||||
*,
|
||||
ssim_threshold: float | None = None,
|
||||
psnr_threshold: float | None = None,
|
||||
audio_similarity_threshold: float | None = None,
|
||||
task_timeout: float | None = None,
|
||||
flag_mode: bool = False, # 是否使用 Feature Flag 方式切换引擎
|
||||
duration_tolerance: float | None = None,
|
||||
):
|
||||
self.api = api
|
||||
self.output_dir = output_dir
|
||||
self.ssim_threshold = ssim_threshold if ssim_threshold is not None else self.DEFAULT_SSIM_THRESHOLD
|
||||
self.psnr_threshold = psnr_threshold if psnr_threshold is not None else self.DEFAULT_PSNR_THRESHOLD
|
||||
self.audio_similarity_threshold = (
|
||||
audio_similarity_threshold
|
||||
if audio_similarity_threshold is not None
|
||||
else self.DEFAULT_AUDIO_SIMILARITY_THRESHOLD
|
||||
)
|
||||
self.duration_tolerance = (
|
||||
duration_tolerance if duration_tolerance is not None else self.DEFAULT_DURATION_TOLERANCE
|
||||
)
|
||||
self.task_timeout = task_timeout if task_timeout is not None else self.DEFAULT_TASK_TIMEOUT
|
||||
self.flag_mode = flag_mode
|
||||
self.results: list[ScenarioResult] = []
|
||||
# flag_mode 下保存原始配置,测试结束后恢复(防污染线上)
|
||||
self._original_flag_config: dict[str, Any] | None = None
|
||||
|
||||
def run_scenario(self, scenario: CompareScenario) -> ScenarioResult:
|
||||
"""运行单个场景对比."""
|
||||
print(f"\n{'='*60}")
|
||||
print(f"[{scenario.priority}] {scenario.id}: {scenario.name}")
|
||||
print(f" {scenario.description}")
|
||||
|
||||
result = ScenarioResult(scenario=scenario)
|
||||
scenario_dir = self.output_dir / scenario.id
|
||||
scenario_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
try:
|
||||
# 1. 提交两个引擎的任务
|
||||
legacy_task_id = self._submit_with_engine(scenario, "legacy")
|
||||
unified_task_id = self._submit_with_engine(scenario, "unified")
|
||||
result.legacy_task_id = legacy_task_id
|
||||
result.unified_task_id = unified_task_id
|
||||
print(f" legacy task: {legacy_task_id}")
|
||||
print(f" unified task: {unified_task_id}")
|
||||
|
||||
# 2. 等待完成
|
||||
print(" waiting for legacy...", end="", flush=True)
|
||||
legacy_status = self.api.wait_for_task(legacy_task_id, timeout=self.task_timeout)
|
||||
result.legacy_success = legacy_status.get("status") in ("completed", "success", "done")
|
||||
legacy_video_url = legacy_status.get("output_url", "") or legacy_status.get("video_url", "")
|
||||
print(f" {'✅' if result.legacy_success else '❌'} ({legacy_status.get('duration_sec', '?')}s)")
|
||||
|
||||
print(" waiting for unified...", end="", flush=True)
|
||||
unified_status = self.api.wait_for_task(unified_task_id, timeout=self.task_timeout)
|
||||
result.unified_success = unified_status.get("status") in ("completed", "success", "done")
|
||||
unified_video_url = unified_status.get("output_url", "") or unified_status.get("video_url", "")
|
||||
print(f" {'✅' if result.unified_success else '❌'} ({unified_status.get('duration_sec', '?')}s)")
|
||||
|
||||
result.legacy_duration_sec = float(legacy_status.get("duration_sec", 0))
|
||||
result.unified_duration_sec = float(unified_status.get("duration_sec", 0))
|
||||
|
||||
if not (result.legacy_success and result.unified_success):
|
||||
result.error = f"Legacy success={result.legacy_success}, Unified success={result.unified_success}"
|
||||
print(" ⚠️ 任务未全部成功,跳过对比")
|
||||
return result
|
||||
|
||||
# 3. 下载视频
|
||||
print(" downloading...", end="", flush=True)
|
||||
legacy_path = self.api.download_video(legacy_video_url, scenario_dir / "legacy.mp4")
|
||||
unified_path = self.api.download_video(unified_video_url, scenario_dir / "unified.mp4")
|
||||
result.legacy_video_path = str(legacy_path)
|
||||
result.unified_video_path = str(unified_path)
|
||||
print(" ✅")
|
||||
|
||||
# 4. 像素对比
|
||||
print(" computing video diff...", end="", flush=True)
|
||||
result.video_diff = compute_video_diff(
|
||||
legacy_path,
|
||||
unified_path,
|
||||
ssim_threshold=self.ssim_threshold,
|
||||
psnr_threshold=self.psnr_threshold,
|
||||
duration_tolerance=self.duration_tolerance,
|
||||
)
|
||||
print(
|
||||
f" SSIM={result.video_diff.avg_ssim:.4f} PSNR={result.video_diff.avg_psnr:.2f}dB {'✅' if result.video_diff.passed else '❌'}"
|
||||
)
|
||||
|
||||
# 5. 音频对比(仅当都有音频时)
|
||||
from .audio_diff import probe_has_audio
|
||||
|
||||
legacy_has_audio = probe_has_audio(legacy_path)
|
||||
unified_has_audio = probe_has_audio(unified_path)
|
||||
|
||||
if legacy_has_audio and unified_has_audio:
|
||||
print(" computing audio diff...", end="", flush=True)
|
||||
result.audio_diff = compute_audio_diff(
|
||||
legacy_path,
|
||||
unified_path,
|
||||
similarity_threshold=self.audio_similarity_threshold,
|
||||
)
|
||||
print(
|
||||
f" similarity={result.audio_diff.similarity_score:.4f} {'✅' if result.audio_diff.passed else '❌'}"
|
||||
)
|
||||
elif legacy_has_audio != unified_has_audio:
|
||||
result.error = f"音频不一致: legacy_has_audio={legacy_has_audio}, unified_has_audio={unified_has_audio}"
|
||||
print(f" ⚠️ 音频不一致: legacy={legacy_has_audio}, unified={unified_has_audio}")
|
||||
else:
|
||||
print(" audio: both silent (skip)")
|
||||
|
||||
except Exception as e:
|
||||
result.error = str(e)
|
||||
print(f" ❌ 错误: {e}")
|
||||
|
||||
self.results.append(result)
|
||||
return result
|
||||
|
||||
def _submit_with_engine(self, scenario: CompareScenario, engine: str) -> str:
|
||||
"""提交指定引擎的任务.
|
||||
|
||||
如果 flag_mode=True,通过 Feature Flag 切换,否则通过 engine 参数。
|
||||
"""
|
||||
if self.flag_mode:
|
||||
# 先设置 flag(用白名单方式,确保只有当前测试用户命中)
|
||||
percentage = 0 if engine == "legacy" else 100
|
||||
self.api.set_feature_flag("render_engine", enabled=True, percentage=percentage)
|
||||
time.sleep(1) # 给 worker 一点时间刷新配置
|
||||
return self.api.submit_render_task(scenario.plan_payload)
|
||||
else:
|
||||
return self.api.submit_render_task(scenario.plan_payload, engine=engine)
|
||||
|
||||
def run_all(self, scenarios: list[CompareScenario]) -> list[ScenarioResult]:
|
||||
"""运行所有场景.
|
||||
|
||||
flag_mode=True 时,测试开始前保存原始 Feature Flag 配置,
|
||||
结束后(无论成功失败)自动恢复,避免污染线上环境。
|
||||
"""
|
||||
print(f"\n灰度对比测试开始 - {len(scenarios)} 个场景")
|
||||
print(f"输出目录: {self.output_dir}")
|
||||
print(f"视频阈值: SSIM>={self.ssim_threshold}, PSNR>={self.psnr_threshold}dB")
|
||||
print(f"音频阈值: similarity>={self.audio_similarity_threshold}")
|
||||
|
||||
# flag_mode:保存原始配置,测试结束后恢复(防污染)
|
||||
if self.flag_mode:
|
||||
try:
|
||||
self._original_flag_config = self.api.get_feature_flag("render_engine")
|
||||
print(f" [flag_mode] 已保存原始配置: {self._original_flag_config}")
|
||||
except Exception as e:
|
||||
print(f" ⚠️ [flag_mode] 保存原始配置失败: {e}")
|
||||
print(" 为避免污染线上,将中止测试。请检查 internal_api_key 配置。")
|
||||
return self.results
|
||||
|
||||
try:
|
||||
for i, scenario in enumerate(scenarios):
|
||||
print(f"\n进度: {i+1}/{len(scenarios)}")
|
||||
self.run_scenario(scenario)
|
||||
finally:
|
||||
# 始终恢复原始 flag 配置
|
||||
if self.flag_mode and self._original_flag_config:
|
||||
try:
|
||||
orig = self._original_flag_config
|
||||
self.api.set_feature_flag(
|
||||
"render_engine",
|
||||
enabled=orig.get("enabled", False),
|
||||
percentage=orig.get("percentage", 0),
|
||||
whitelist=orig.get("whitelist"),
|
||||
)
|
||||
print("\n[flag_mode] ✅ 已恢复原始 Feature Flag 配置")
|
||||
except Exception as e:
|
||||
print(f"\n[flag_mode] ❌ 恢复 Feature Flag 失败: {e}")
|
||||
print(" 请手动检查并恢复 render_engine flag 配置!")
|
||||
|
||||
return self.results
|
||||
|
||||
def summary(self) -> dict[str, Any]:
|
||||
"""生成汇总统计."""
|
||||
total = len(self.results)
|
||||
passed = sum(1 for r in self.results if r.passed)
|
||||
failed = total - passed
|
||||
|
||||
# 性能对比
|
||||
perf_diffs = []
|
||||
for r in self.results:
|
||||
if r.legacy_success and r.unified_success and r.legacy_duration_sec > 0:
|
||||
diff_pct = (r.unified_duration_sec - r.legacy_duration_sec) / r.legacy_duration_sec * 100
|
||||
perf_diffs.append(diff_pct)
|
||||
avg_perf_diff = sum(perf_diffs) / len(perf_diffs) if perf_diffs else 0.0
|
||||
|
||||
return {
|
||||
"total": total,
|
||||
"passed": passed,
|
||||
"failed": failed,
|
||||
"pass_rate": f"{passed/total*100:.1f}%" if total > 0 else "0%",
|
||||
"avg_perf_diff_pct": round(avg_perf_diff, 2),
|
||||
"scenarios": [self._result_to_dict(r) for r in self.results],
|
||||
"timestamp": datetime.now().isoformat(),
|
||||
"ssim_threshold": self.ssim_threshold,
|
||||
"psnr_threshold": self.psnr_threshold,
|
||||
"audio_threshold": self.audio_similarity_threshold,
|
||||
}
|
||||
|
||||
def _result_to_dict(self, r: ScenarioResult) -> dict[str, Any]:
|
||||
return {
|
||||
"id": r.scenario.id,
|
||||
"name": r.scenario.name,
|
||||
"priority": r.scenario.priority,
|
||||
"passed": r.passed,
|
||||
"legacy_success": r.legacy_success,
|
||||
"unified_success": r.unified_success,
|
||||
"legacy_duration_sec": r.legacy_duration_sec,
|
||||
"unified_duration_sec": r.unified_duration_sec,
|
||||
"video_diff": r.video_diff.to_dict() if r.video_diff else None,
|
||||
"audio_diff": r.audio_diff.to_dict() if r.audio_diff else None,
|
||||
"error": r.error,
|
||||
}
|
||||
|
||||
|
||||
def generate_html_report(summary: dict[str, Any], output_path: Path):
|
||||
"""生成 HTML 对比报告."""
|
||||
scenarios = summary["scenarios"]
|
||||
|
||||
# 按通过/失败分组
|
||||
passed_list = [s for s in scenarios if s["passed"]]
|
||||
failed_list = [s for s in scenarios if not s["passed"]]
|
||||
|
||||
# 构建场景卡片
|
||||
scenario_cards = ""
|
||||
for s in scenarios:
|
||||
status_class = "pass" if s["passed"] else "fail"
|
||||
status_text = "✅ 通过" if s["passed"] else "❌ 失败"
|
||||
|
||||
vdiff = s.get("video_diff") or {}
|
||||
adiff = s.get("audio_diff") or {}
|
||||
|
||||
video_info = ""
|
||||
if vdiff:
|
||||
video_info = f"""
|
||||
<div class="metric-row">
|
||||
<span>SSIM:</span>
|
||||
<span class="{'good' if vdiff.get('avg_ssim', 0) >= 0.95 else 'warn'}">{vdiff.get('avg_ssim', 0):.4f}</span>
|
||||
</div>
|
||||
<div class="metric-row">
|
||||
<span>PSNR:</span>
|
||||
<span>{vdiff.get('avg_psnr', 0):.2f} dB</span>
|
||||
</div>
|
||||
<div class="metric-row">
|
||||
<span>时长差:</span>
|
||||
<span>{vdiff.get('duration_diff', 0):.3f}s</span>
|
||||
</div>
|
||||
"""
|
||||
|
||||
audio_info = ""
|
||||
if adiff:
|
||||
audio_info = f"""
|
||||
<div class="metric-row">
|
||||
<span>音频相似度:</span>
|
||||
<span class="{'good' if adiff.get('similarity_score', 0) >= 0.9 else 'warn'}">{adiff.get('similarity_score', 0):.4f}</span>
|
||||
</div>
|
||||
<div class="metric-row">
|
||||
<span>差值 RMS:</span>
|
||||
<span>{adiff.get('diff_rms_db', 0):.2f} dB</span>
|
||||
</div>
|
||||
"""
|
||||
|
||||
perf_info = ""
|
||||
if s["legacy_duration_sec"] and s["unified_duration_sec"]:
|
||||
diff = s["unified_duration_sec"] - s["legacy_duration_sec"]
|
||||
pct = diff / s["legacy_duration_sec"] * 100 if s["legacy_duration_sec"] else 0
|
||||
trend = "🔴" if pct > 10 else ("🟡" if pct > 0 else "🟢")
|
||||
perf_info = f"""
|
||||
<div class="perf-row">
|
||||
<span>Legacy: {s['legacy_duration_sec']:.2f}s</span>
|
||||
<span>Unified: {s['unified_duration_sec']:.2f}s</span>
|
||||
<span>{trend} {pct:+.1f}%</span>
|
||||
</div>
|
||||
"""
|
||||
|
||||
error_info = f'<div class="error-box">{s["error"]}</div>' if s["error"] else ""
|
||||
|
||||
scenario_cards += f"""
|
||||
<div class="card {status_class}">
|
||||
<div class="card-header">
|
||||
<span class="badge">{s['priority']}</span>
|
||||
<span class="scenario-name">{s['name']}</span>
|
||||
<span class="status {status_class}">{status_text}</span>
|
||||
</div>
|
||||
<div class="card-body">
|
||||
<div class="grid-2">
|
||||
<div>
|
||||
<h4>视频质量</h4>
|
||||
{video_info or '<p class="muted">无数据</p>'}
|
||||
</div>
|
||||
<div>
|
||||
<h4>音频质量</h4>
|
||||
{audio_info or '<p class="muted">无音频或跳过</p>'}
|
||||
</div>
|
||||
</div>
|
||||
<div>
|
||||
<h4>性能对比</h4>
|
||||
{perf_info or '<p class="muted">无数据</p>'}
|
||||
</div>
|
||||
{error_info}
|
||||
</div>
|
||||
</div>
|
||||
"""
|
||||
|
||||
html = f"""<!DOCTYPE html>
|
||||
<html lang="zh-CN">
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1.0">
|
||||
<title>统一渲染引擎灰度对比报告</title>
|
||||
<style>
|
||||
* {{ box-sizing: border-box; margin: 0; padding: 0; }}
|
||||
body {{ font-family: -apple-system, BlinkMacSystemFont, 'Segoe UI', Roboto, sans-serif; background: #f5f5f5; color: #333; padding: 20px; }}
|
||||
.container {{ max-width: 1200px; margin: 0 auto; }}
|
||||
h1 {{ margin-bottom: 20px; font-size: 24px; }}
|
||||
.summary {{ background: white; border-radius: 12px; padding: 24px; margin-bottom: 24px; display: flex; gap: 32px; flex-wrap: wrap; }}
|
||||
.summary-item {{ text-align: center; }}
|
||||
.summary-item .value {{ font-size: 32px; font-weight: bold; margin-bottom: 4px; }}
|
||||
.summary-item .label {{ color: #666; font-size: 14px; }}
|
||||
.pass .value {{ color: #10b981; }}
|
||||
.fail .value {{ color: #ef4444; }}
|
||||
.card {{ background: white; border-radius: 12px; margin-bottom: 16px; overflow: hidden; border-left: 4px solid #10b981; }}
|
||||
.card.fail {{ border-left-color: #ef4444; }}
|
||||
.card-header {{ padding: 16px 20px; background: #fafafa; display: flex; align-items: center; gap: 12px; border-bottom: 1px solid #eee; }}
|
||||
.badge {{ background: #e5e7eb; color: #374151; padding: 2px 8px; border-radius: 4px; font-size: 12px; font-weight: 600; }}
|
||||
.scenario-name {{ flex: 1; font-weight: 600; }}
|
||||
.status {{ font-weight: 600; }}
|
||||
.status.pass {{ color: #10b981; }}
|
||||
.status.fail {{ color: #ef4444; }}
|
||||
.card-body {{ padding: 20px; }}
|
||||
.grid-2 {{ display: grid; grid-template-columns: 1fr 1fr; gap: 24px; margin-bottom: 16px; }}
|
||||
h4 {{ margin-bottom: 12px; color: #374151; font-size: 14px; }}
|
||||
.metric-row {{ display: flex; justify-content: space-between; padding: 6px 0; font-size: 14px; }}
|
||||
.metric-row .good {{ color: #10b981; font-weight: 600; }}
|
||||
.metric-row .warn {{ color: #f59e0b; font-weight: 600; }}
|
||||
.perf-row {{ display: flex; gap: 24px; padding: 8px 0; font-size: 14px; background: #f9fafb; padding: 12px; border-radius: 8px; }}
|
||||
.error-box {{ background: #fef2f2; color: #dc2626; padding: 12px; border-radius: 8px; margin-top: 12px; font-size: 13px; }}
|
||||
.muted {{ color: #9ca3af; font-size: 14px; }}
|
||||
.timestamp {{ text-align: center; color: #9ca3af; font-size: 12px; margin-top: 24px; }}
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<div class="container">
|
||||
<h1>🎬 统一渲染引擎灰度对比报告</h1>
|
||||
<div class="summary">
|
||||
<div class="summary-item">
|
||||
<div class="value">{summary['total']}</div>
|
||||
<div class="label">总场景数</div>
|
||||
</div>
|
||||
<div class="summary-item pass">
|
||||
<div class="value">{summary['passed']}</div>
|
||||
<div class="label">通过</div>
|
||||
</div>
|
||||
<div class="summary-item fail">
|
||||
<div class="value">{summary['failed']}</div>
|
||||
<div class="label">失败</div>
|
||||
</div>
|
||||
<div class="summary-item">
|
||||
<div class="value">{summary['pass_rate']}</div>
|
||||
<div class="label">通过率</div>
|
||||
</div>
|
||||
<div class="summary-item">
|
||||
<div class="value {'good' if summary['avg_perf_diff_pct'] <= 0 else 'warn'}" style="font-size: 24px; color: {'#10b981' if summary['avg_perf_diff_pct'] <= 0 else '#f59e0b'}">{summary['avg_perf_diff_pct']:+.1f}%</div>
|
||||
<div class="label">平均性能差异</div>
|
||||
</div>
|
||||
</div>
|
||||
{scenario_cards}
|
||||
<div class="timestamp">生成时间: {summary['timestamp']}</div>
|
||||
</div>
|
||||
</body>
|
||||
</html>"""
|
||||
|
||||
output_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
output_path.write_text(html, encoding="utf-8")
|
||||
return output_path
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="统一渲染引擎灰度对比测试")
|
||||
parser.add_argument("--priority", default="P0", choices=["P0", "P1", "P2"], help="最低优先级")
|
||||
parser.add_argument("--scenarios", default="", help="指定场景ID,逗号分隔")
|
||||
parser.add_argument("--output", default="./gray_compare_report", help="输出目录")
|
||||
parser.add_argument("--ssim-threshold", type=float, default=None, help="SSIM阈值(默认0.95)")
|
||||
parser.add_argument("--psnr-threshold", type=float, default=None, help="PSNR阈值(dB)(默认28.0)")
|
||||
parser.add_argument("--audio-threshold", type=float, default=None, help="音频相似度阈值(默认0.90)")
|
||||
parser.add_argument("--flag-mode", action="store_true", help="使用Feature Flag方式切换引擎")
|
||||
parser.add_argument("--task-timeout", type=float, default=300.0, help="单任务超时时间(秒)")
|
||||
args = parser.parse_args()
|
||||
|
||||
base_url = os.environ.get("STAGING_API_URL", "")
|
||||
api_key = os.environ.get("STAGING_API_KEY", "")
|
||||
internal_key = os.environ.get("STAGING_INTERNAL_API_KEY", "")
|
||||
|
||||
if not base_url or not api_key:
|
||||
print("❌ 请设置环境变量 STAGING_API_URL 和 STAGING_API_KEY")
|
||||
sys.exit(1)
|
||||
|
||||
# 选择场景
|
||||
if args.scenarios:
|
||||
scenario_ids = [s.strip() for s in args.scenarios.split(",")]
|
||||
selected = [s for s in SCENARIOS if s.id in scenario_ids]
|
||||
if not selected:
|
||||
print(f"❌ 未找到匹配的场景: {scenario_ids}")
|
||||
print(f"可用场景: {[s.id for s in SCENARIOS]}")
|
||||
sys.exit(1)
|
||||
else:
|
||||
selected = get_scenarios_by_priority(args.priority)
|
||||
|
||||
output_dir = Path(args.output).resolve()
|
||||
output_dir.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
api = StagingAPI(base_url, api_key, internal_key)
|
||||
runner = CompareRunner(
|
||||
api,
|
||||
output_dir,
|
||||
ssim_threshold=args.ssim_threshold,
|
||||
psnr_threshold=args.psnr_threshold,
|
||||
audio_similarity_threshold=args.audio_threshold,
|
||||
flag_mode=args.flag_mode,
|
||||
task_timeout=args.task_timeout,
|
||||
)
|
||||
|
||||
runner.run_all(selected)
|
||||
|
||||
# 生成报告
|
||||
summary = runner.summary()
|
||||
|
||||
# JSON 报告
|
||||
json_path = output_dir / "report.json"
|
||||
json_path.write_text(json.dumps(summary, indent=2, ensure_ascii=False), encoding="utf-8")
|
||||
|
||||
# HTML 报告
|
||||
html_path = output_dir / "report.html"
|
||||
generate_html_report(summary, html_path)
|
||||
|
||||
print(f"\n{'='*60}")
|
||||
print(f"对比完成: {summary['passed']}/{summary['total']} 通过 ({summary['pass_rate']})")
|
||||
print(f"报告: {html_path}")
|
||||
print(f"JSON: {json_path}")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,257 @@
|
||||
"""灰度对比测试场景定义 — 覆盖典型渲染场景.
|
||||
|
||||
每个场景对应一个 EditPlan,用于新旧引擎对比。
|
||||
覆盖场景:
|
||||
1. 简单直通(单clip无特效)
|
||||
2. 多clip转场(fade + slide)
|
||||
3. 画中画(main + overlay)
|
||||
4. 字幕渲染(ASS字幕)
|
||||
5. 独立音频轨(主视频 + BGM)
|
||||
6. 多图层混合(main + broll + overlay + audio)
|
||||
7. 背景图片 + 主视频(图片背景无音频)
|
||||
8. 无音频视频(纯画面,验证无音轨防御)
|
||||
9. 长视频(10+ clip,压力测试)
|
||||
10. 分辨率非标(竖屏9:16,验证scale策略)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass
|
||||
class CompareScenario:
|
||||
"""对比测试场景."""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
description: str
|
||||
priority: str # P0 / P1 / P2
|
||||
plan_payload: dict[str, Any] # EditPlan JSON payload(提交给 API 的数据)
|
||||
expected: dict[str, Any] = field(default_factory=dict) # 预期结果
|
||||
|
||||
|
||||
SCENARIOS: list[CompareScenario] = [
|
||||
CompareScenario(
|
||||
id="simple_pass_through",
|
||||
name="简单直通",
|
||||
description="单主clip,无转场无特效,验证直通优化路径",
|
||||
priority="P0",
|
||||
plan_payload={
|
||||
"width": 1280,
|
||||
"height": 720,
|
||||
"fps": 25,
|
||||
"clips": [
|
||||
{
|
||||
"clip_type": "main",
|
||||
"asset_id": "sample_5s.mp4",
|
||||
"duration": 5.0,
|
||||
"order": 0,
|
||||
}
|
||||
],
|
||||
},
|
||||
),
|
||||
CompareScenario(
|
||||
id="multi_clip_transition",
|
||||
name="多clip转场",
|
||||
description="3个clip,fade + 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]
|
||||
@@ -0,0 +1,283 @@
|
||||
"""视频对比工具 — 基于 FFmpeg 的像素级质量对比.
|
||||
|
||||
使用 SSIM + PSNR 双指标评估两个视频的相似度:
|
||||
- SSIM (Structural Similarity): 结构相似性,范围 [0, 1],越接近 1 越相似
|
||||
- PSNR (Peak Signal-to-Noise Ratio): 峰值信噪比,单位 dB,越高越好
|
||||
|
||||
灰度验收标准:
|
||||
- 平均 SSIM >= 0.95 → 视觉上几乎无差异(P0 场景必达)
|
||||
- 最低 SSIM >= 0.90 → 最严重帧差异可接受
|
||||
- 平均 PSNR >= 28dB → 质量达标
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import re
|
||||
import shutil
|
||||
import subprocess # nosec B404
|
||||
from dataclasses import asdict, dataclass
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
FFMPEG_BIN: str = shutil.which("ffmpeg") or "ffmpeg"
|
||||
FFPROBE_BIN: str = shutil.which("ffprobe") or "ffprobe"
|
||||
|
||||
|
||||
@dataclass
|
||||
class VideoDiffResult:
|
||||
"""视频对比结果."""
|
||||
|
||||
video_a: str
|
||||
video_b: str
|
||||
width: int
|
||||
height: int
|
||||
duration_a: float
|
||||
duration_b: float
|
||||
avg_ssim: float
|
||||
min_ssim: float
|
||||
avg_psnr: float # dB
|
||||
min_psnr: float
|
||||
frame_count: int
|
||||
duration_diff: float # 时长差(秒)
|
||||
resolution_match: bool
|
||||
passed: bool # 是否通过阈值
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return asdict(self)
|
||||
|
||||
|
||||
def probe_video_info(video_path: str) -> dict[str, Any]:
|
||||
"""获取视频信息(宽、高、时长、fps)."""
|
||||
try:
|
||||
result = subprocess.run( # nosec B603
|
||||
[
|
||||
FFPROBE_BIN,
|
||||
"-v",
|
||||
"error",
|
||||
"-select_streams",
|
||||
"v:0",
|
||||
"-show_entries",
|
||||
"stream=width,height,r_frame_rate,duration",
|
||||
"-show_entries",
|
||||
"format=duration",
|
||||
"-of",
|
||||
"json",
|
||||
video_path,
|
||||
],
|
||||
check=True,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
timeout=10,
|
||||
)
|
||||
info = json.loads(result.stdout)
|
||||
stream = info.get("streams", [{}])[0]
|
||||
fmt = info.get("format", {})
|
||||
|
||||
width = int(stream.get("width", 1280))
|
||||
height = int(stream.get("height", 720))
|
||||
fps_str = stream.get("r_frame_rate", "25/1")
|
||||
if "/" in fps_str:
|
||||
num, den = fps_str.split("/")
|
||||
fps = float(num) / float(den) if float(den) > 0 else 25.0
|
||||
else:
|
||||
fps = float(fps_str) if fps_str else 25.0
|
||||
duration = float(fmt.get("duration", 0)) or float(stream.get("duration", 0))
|
||||
|
||||
return {"width": width, "height": height, "duration": duration, "fps": round(fps, 2)}
|
||||
except Exception:
|
||||
return {"width": 1280, "height": 720, "duration": 0.0, "fps": 25.0}
|
||||
|
||||
|
||||
def compute_video_diff(
|
||||
video_a: str | Path,
|
||||
video_b: str | Path,
|
||||
*,
|
||||
ssim_threshold: float = 0.95,
|
||||
psnr_threshold: float = 28.0,
|
||||
duration_tolerance: float = 0.1,
|
||||
) -> VideoDiffResult:
|
||||
"""计算两个视频的像素差异.
|
||||
|
||||
使用 FFmpeg ssim + psnr 滤镜一次性计算两个指标。
|
||||
|
||||
Args:
|
||||
video_a: 视频A路径(基线)
|
||||
video_b: 视频B路径(对比)
|
||||
ssim_threshold: SSIM 合格阈值(默认 0.90)
|
||||
psnr_threshold: PSNR 合格阈值(默认 25dB)
|
||||
duration_tolerance: 时长容忍度(秒,默认 0.1s)
|
||||
|
||||
Returns:
|
||||
VideoDiffResult 对比结果
|
||||
|
||||
Raises:
|
||||
subprocess.CalledProcessError: FFmpeg 执行失败
|
||||
"""
|
||||
info_a = probe_video_info(str(video_a))
|
||||
info_b = probe_video_info(str(video_b))
|
||||
|
||||
duration_diff = abs(info_a["duration"] - info_b["duration"])
|
||||
resolution_match = info_a["width"] == info_b["width"] and info_a["height"] == info_b["height"]
|
||||
|
||||
# ssim 和 psnr 的 stats_file 都输出到 stdout
|
||||
# 用行格式区分:SSIM 行含 "All:",PSNR 行含 "psnr_avg:"
|
||||
command = [
|
||||
FFMPEG_BIN,
|
||||
"-i",
|
||||
str(video_a),
|
||||
"-i",
|
||||
str(video_b),
|
||||
"-lavfi",
|
||||
"[0:v][1:v]ssim=stats_file=-[out1];[0:v][1:v]psnr=stats_file=-[out2]",
|
||||
"-f",
|
||||
"null",
|
||||
"-",
|
||||
]
|
||||
|
||||
result = subprocess.run( # nosec B603
|
||||
command,
|
||||
check=True,
|
||||
stdout=subprocess.PIPE,
|
||||
stderr=subprocess.PIPE,
|
||||
text=True,
|
||||
timeout=300,
|
||||
)
|
||||
|
||||
# 逐帧统计在 stdout(stats_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)
|
||||
Executable
+72
@@ -0,0 +1,72 @@
|
||||
"""AssetStatus 枚举兼容性测试。
|
||||
|
||||
验证历史脏数据(如 'uploaded')不会导致枚举转换失败。
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.entities import AssetStatus
|
||||
|
||||
|
||||
class TestAssetStatusNormalValues:
|
||||
"""正常值应该正确映射。"""
|
||||
|
||||
def test_uploading(self):
|
||||
assert AssetStatus("uploading") == AssetStatus.UPLOADING
|
||||
|
||||
def test_ready(self):
|
||||
assert AssetStatus("ready") == AssetStatus.READY
|
||||
|
||||
def test_processing(self):
|
||||
assert AssetStatus("processing") == AssetStatus.PROCESSING
|
||||
|
||||
def test_error(self):
|
||||
assert AssetStatus("error") == AssetStatus.ERROR
|
||||
|
||||
|
||||
class TestAssetStatusHistoricalValues:
|
||||
"""历史脏数据应该正确映射到对应状态,不抛异常。"""
|
||||
|
||||
@pytest.mark.parametrize("value", ["uploaded", "Uploaded", "UPLOADED", " uploaded "])
|
||||
def test_uploaded_maps_to_ready(self, value):
|
||||
"""生产环境发现的 'uploaded' 历史值应映射为 READY。"""
|
||||
assert AssetStatus(value) == AssetStatus.READY
|
||||
|
||||
@pytest.mark.parametrize("value", ["success", "ok", "done", "complete"])
|
||||
def test_other_ready_like_values_map_to_ready(self, value):
|
||||
assert AssetStatus(value) == AssetStatus.READY
|
||||
|
||||
@pytest.mark.parametrize("value", ["upload", "uploading_start", "upload_start"])
|
||||
def test_upload_like_values_map_to_uploading(self, value):
|
||||
assert AssetStatus(value) == AssetStatus.UPLOADING
|
||||
|
||||
@pytest.mark.parametrize("value", ["failed", "fail", "err"])
|
||||
def test_error_like_values_map_to_error(self, value):
|
||||
assert AssetStatus(value) == AssetStatus.ERROR
|
||||
|
||||
@pytest.mark.parametrize("value", ["process", "running", "run"])
|
||||
def test_processing_like_values_map_to_processing(self, value):
|
||||
assert AssetStatus(value) == AssetStatus.PROCESSING
|
||||
|
||||
|
||||
class TestAssetStatusFallback:
|
||||
"""完全未知的值兜底为 READY,不抛500。"""
|
||||
|
||||
@pytest.mark.parametrize("value", ["unknown", "foo_bar", ""])
|
||||
def test_unknown_value_falls_back_to_ready(self, value):
|
||||
assert AssetStatus(value) == AssetStatus.READY
|
||||
|
||||
def test_none_value_falls_back_to_ready(self):
|
||||
assert AssetStatus(None) == AssetStatus.READY # type: ignore[arg-type]
|
||||
|
||||
def test_int_value_falls_back_to_ready(self):
|
||||
assert AssetStatus(123) == AssetStatus.READY # type: ignore[arg-type]
|
||||
|
||||
|
||||
class TestAssetStatusStrValue:
|
||||
"""枚举值仍为字符串类型,不影响序列化。"""
|
||||
|
||||
def test_value_unchanged(self):
|
||||
assert AssetStatus.READY.value == "ready"
|
||||
assert AssetStatus.ERROR.value == "error"
|
||||
assert isinstance(AssetStatus.READY, str)
|
||||
+68
@@ -0,0 +1,68 @@
|
||||
"""ClassificationStatus 枚举兼容性测试。
|
||||
|
||||
验证历史脏数据(如 'done')不会导致枚举转换失败。
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.entities import ClassificationStatus
|
||||
|
||||
|
||||
class TestClassificationStatusNormalValues:
|
||||
"""正常值应该正确映射。"""
|
||||
|
||||
def test_pending(self):
|
||||
assert ClassificationStatus("pending") == ClassificationStatus.PENDING
|
||||
|
||||
def test_processing(self):
|
||||
assert ClassificationStatus("processing") == ClassificationStatus.PROCESSING
|
||||
|
||||
def test_completed(self):
|
||||
assert ClassificationStatus("completed") == ClassificationStatus.COMPLETED
|
||||
|
||||
def test_failed(self):
|
||||
assert ClassificationStatus("failed") == ClassificationStatus.FAILED
|
||||
|
||||
|
||||
class TestClassificationStatusHistoricalValues:
|
||||
"""历史脏数据应该正确映射到对应状态,不抛异常。"""
|
||||
|
||||
@pytest.mark.parametrize("value", ["done", "Done", "DONE", " done "])
|
||||
def test_done_maps_to_completed(self, value):
|
||||
"""生产环境发现的 'done' 历史值应映射为 COMPLETED。"""
|
||||
assert ClassificationStatus(value) == ClassificationStatus.COMPLETED
|
||||
|
||||
@pytest.mark.parametrize("value", ["success", "finished", "complete"])
|
||||
def test_other_done_like_values_map_to_completed(self, value):
|
||||
assert ClassificationStatus(value) == ClassificationStatus.COMPLETED
|
||||
|
||||
@pytest.mark.parametrize("value", ["fail", "error", "err"])
|
||||
def test_error_like_values_map_to_failed(self, value):
|
||||
assert ClassificationStatus(value) == ClassificationStatus.FAILED
|
||||
|
||||
@pytest.mark.parametrize("value", ["process", "running", "run"])
|
||||
def test_processing_like_values_map_to_processing(self, value):
|
||||
assert ClassificationStatus(value) == ClassificationStatus.PROCESSING
|
||||
|
||||
|
||||
class TestClassificationStatusFallback:
|
||||
"""完全未知的值兜底为 PENDING,不抛500。"""
|
||||
|
||||
@pytest.mark.parametrize("value", ["unknown", "foo_bar", ""])
|
||||
def test_unknown_value_falls_back_to_pending(self, value):
|
||||
assert ClassificationStatus(value) == ClassificationStatus.PENDING
|
||||
|
||||
def test_none_value_falls_back_to_pending(self):
|
||||
assert ClassificationStatus(None) == ClassificationStatus.PENDING # type: ignore[arg-type]
|
||||
|
||||
def test_int_value_falls_back_to_pending(self):
|
||||
assert ClassificationStatus(123) == ClassificationStatus.PENDING # type: ignore[arg-type]
|
||||
|
||||
|
||||
class TestClassificationStatusStrValue:
|
||||
"""枚举值仍为字符串类型,不影响序列化。"""
|
||||
|
||||
def test_value_unchanged(self):
|
||||
assert ClassificationStatus.COMPLETED.value == "completed"
|
||||
assert ClassificationStatus.PENDING.value == "pending"
|
||||
assert isinstance(ClassificationStatus.COMPLETED, str)
|
||||
Regular → Executable
+2
@@ -63,6 +63,8 @@ class StubEditPlan:
|
||||
template_id: str = "tmpl-001"
|
||||
status: Any = None
|
||||
config: dict = field(default_factory=dict)
|
||||
project_id: str = ""
|
||||
created_by_user_id: str = "user-001"
|
||||
|
||||
def mark_failed(self):
|
||||
self.status = _StubStatus("failed")
|
||||
|
||||
Executable
+424
@@ -0,0 +1,424 @@
|
||||
"""Feature Flag 单元测试。
|
||||
|
||||
测试 FeatureFlagConfig、InMemoryFeatureFlagStore、RenderEngineResolver 的核心逻辑。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import time
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.adapters.redis.feature_flag_store import (
|
||||
FeatureFlagConfig,
|
||||
InMemoryFeatureFlagStore,
|
||||
)
|
||||
|
||||
# ── FeatureFlagConfig 测试 ──────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestFeatureFlagConfig:
|
||||
"""FeatureFlagConfig 核心逻辑测试。"""
|
||||
|
||||
def test_default_disabled(self):
|
||||
"""默认配置为关闭状态。"""
|
||||
config = FeatureFlagConfig(name="test_flag")
|
||||
assert config.enabled is False
|
||||
assert config.percentage == 0
|
||||
assert config.whitelist == set()
|
||||
assert config.is_active() is False
|
||||
assert config.is_active("user1") is False
|
||||
|
||||
def test_global_enabled_100_percent(self):
|
||||
"""100% + 启用 = 全部命中。"""
|
||||
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=100)
|
||||
assert config.is_active() is True
|
||||
assert config.is_active("user1") is True
|
||||
assert config.is_active("any_user") is True
|
||||
|
||||
def test_global_enabled_0_percent_no_whitelist(self):
|
||||
"""启用但 0% 且无白名单 = 不命中。"""
|
||||
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=0)
|
||||
assert config.is_active() is False
|
||||
assert config.is_active("user1") is False
|
||||
|
||||
def test_whitelist_takes_priority(self):
|
||||
"""白名单优先级高于百分比。"""
|
||||
config = FeatureFlagConfig(
|
||||
name="test_flag",
|
||||
enabled=True,
|
||||
percentage=0,
|
||||
whitelist={"user1", "user2"},
|
||||
)
|
||||
assert config.is_active("user1") is True
|
||||
assert config.is_active("user2") is True
|
||||
assert config.is_active("user3") is False
|
||||
|
||||
def test_whitelist_with_percentage(self):
|
||||
"""白名单用户即使百分比为0也命中,非白名单按百分比。"""
|
||||
config = FeatureFlagConfig(
|
||||
name="test_flag",
|
||||
enabled=True,
|
||||
percentage=100, # 100% 所有人命中
|
||||
whitelist={"user1"},
|
||||
)
|
||||
assert config.is_active("user1") is True
|
||||
assert config.is_active("user999") is True # 100% 命中
|
||||
|
||||
def test_percentage_consistency_same_user(self):
|
||||
"""同一用户多次调用结果一致(哈希确定性)。"""
|
||||
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=50)
|
||||
results = [config.is_active("user_fixed") for _ in range(100)]
|
||||
assert all(r == results[0] for r in results)
|
||||
|
||||
def test_percentage_different_users_distributed(self):
|
||||
"""不同用户分布大致符合百分比(统计检验,宽松阈值)。"""
|
||||
config = FeatureFlagConfig(name="test_flag", enabled=True, percentage=50)
|
||||
active_count = sum(1 for i in range(1000) if config.is_active(f"user_{i}"))
|
||||
# 50% 上下浮动 10% 都算合理
|
||||
assert 400 <= active_count <= 600, f"Expected ~500, got {active_count}"
|
||||
|
||||
def test_percentage_boundary_0_and_100(self):
|
||||
"""0% 和 100% 的边界情况。"""
|
||||
config_0 = FeatureFlagConfig(name="test", enabled=True, percentage=0)
|
||||
config_100 = FeatureFlagConfig(name="test", enabled=True, percentage=100)
|
||||
|
||||
for i in range(100):
|
||||
assert config_0.is_active(f"user_{i}") is False
|
||||
assert config_100.is_active(f"user_{i}") is True
|
||||
|
||||
def test_disabled_ignores_all_other_settings(self):
|
||||
"""关闭时忽略白名单和百分比。"""
|
||||
config = FeatureFlagConfig(
|
||||
name="test_flag",
|
||||
enabled=False,
|
||||
percentage=100,
|
||||
whitelist={"user1"},
|
||||
)
|
||||
assert config.is_active("user1") is False
|
||||
assert config.is_active() is False
|
||||
|
||||
def test_none_identifier_with_percentage(self):
|
||||
"""无 identifier 时按随机比例(0% 和 100% 是确定的)。"""
|
||||
config_0 = FeatureFlagConfig(name="test", enabled=True, percentage=0)
|
||||
config_100 = FeatureFlagConfig(name="test", enabled=True, percentage=100)
|
||||
assert config_0.is_active(None) is False
|
||||
assert config_100.is_active(None) is True
|
||||
|
||||
def test_to_dict_and_from_dict(self):
|
||||
"""序列化和反序列化对称。"""
|
||||
original = FeatureFlagConfig(
|
||||
name="test_flag",
|
||||
enabled=True,
|
||||
percentage=30,
|
||||
whitelist={"user_a", "user_b", "user_c"},
|
||||
)
|
||||
data = original.to_dict()
|
||||
restored = FeatureFlagConfig.from_dict(data)
|
||||
assert restored.name == original.name
|
||||
assert restored.enabled == original.enabled
|
||||
assert restored.percentage == original.percentage
|
||||
assert restored.whitelist == original.whitelist
|
||||
|
||||
def test_from_dict_with_missing_fields(self):
|
||||
"""from_dict 缺失字段时使用默认值。"""
|
||||
config = FeatureFlagConfig.from_dict({"name": "minimal"})
|
||||
assert config.name == "minimal"
|
||||
assert config.enabled is False
|
||||
assert config.percentage == 0
|
||||
assert config.whitelist == set()
|
||||
|
||||
|
||||
# ── InMemoryFeatureFlagStore 测试 ───────────────────────────────────────────
|
||||
|
||||
|
||||
class TestInMemoryFeatureFlagStore:
|
||||
"""内存存储实现测试。"""
|
||||
|
||||
def test_get_nonexistent_returns_default(self):
|
||||
"""获取不存在的 flag 返回默认配置(关闭)。"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
config = store.get("nonexistent")
|
||||
assert config.name == "nonexistent"
|
||||
assert config.enabled is False
|
||||
|
||||
def test_set_and_get(self):
|
||||
"""设置后可以读取。"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
config = FeatureFlagConfig(name="test", enabled=True, percentage=50, whitelist={"u1"})
|
||||
store.set(config)
|
||||
|
||||
got = store.get("test")
|
||||
assert got.enabled is True
|
||||
assert got.percentage == 50
|
||||
assert got.whitelist == {"u1"}
|
||||
|
||||
def test_delete_existing(self):
|
||||
"""删除存在的 flag 返回 True。"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
store.set(FeatureFlagConfig(name="test", enabled=True))
|
||||
assert store.delete("test") is True
|
||||
assert store.get("test").enabled is False
|
||||
|
||||
def test_delete_nonexistent(self):
|
||||
"""删除不存在的 flag 返回 False。"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
assert store.delete("nonexistent") is False
|
||||
|
||||
def test_list_all(self):
|
||||
"""列出所有 flag。"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
store.set(FeatureFlagConfig(name="flag_a", enabled=True))
|
||||
store.set(FeatureFlagConfig(name="flag_b", percentage=10))
|
||||
|
||||
all_flags = store.list_all()
|
||||
assert len(all_flags) == 2
|
||||
assert "flag_a" in all_flags
|
||||
assert "flag_b" in all_flags
|
||||
assert all_flags["flag_a"].enabled is True
|
||||
|
||||
def test_is_active_convenience(self):
|
||||
"""is_active 便捷方法。"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
store.set(FeatureFlagConfig(name="render", enabled=True, percentage=0, whitelist={"vip_user"}))
|
||||
assert store.is_active("render", "vip_user") is True
|
||||
assert store.is_active("render", "normal_user") is False
|
||||
assert store.is_active("nonexistent") is False
|
||||
|
||||
|
||||
# ── RenderEngineResolver 测试 ───────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestRenderEngineResolver:
|
||||
"""渲染引擎选择器测试。"""
|
||||
|
||||
def test_default_legacy_when_flag_disabled(self):
|
||||
"""flag 关闭时使用默认引擎(legacy)。"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
resolver = self._make_resolver(store=store, default="legacy")
|
||||
assert resolver.get_engine() == "legacy"
|
||||
assert resolver.get_engine("user1") == "legacy"
|
||||
|
||||
def test_default_unified_when_flag_disabled(self):
|
||||
"""flag 关闭但默认值是 unified 时返回 unified。"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
resolver = self._make_resolver(store=store, default="unified")
|
||||
assert resolver.get_engine() == "unified"
|
||||
|
||||
def test_whitelist_user_uses_unified(self):
|
||||
"""白名单用户走新引擎。"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
store.set(
|
||||
FeatureFlagConfig(
|
||||
name="render_engine",
|
||||
enabled=True,
|
||||
percentage=0,
|
||||
whitelist={"beta_tester"},
|
||||
)
|
||||
)
|
||||
resolver = self._make_resolver(store=store, default="legacy")
|
||||
assert resolver.get_engine("beta_tester") == "unified"
|
||||
assert resolver.get_engine("normal_user") == "legacy"
|
||||
|
||||
def test_100_percent_all_unified(self):
|
||||
"""100% 时所有用户走新引擎。"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
store.set(FeatureFlagConfig(name="render_engine", enabled=True, percentage=100))
|
||||
resolver = self._make_resolver(store=store, default="legacy")
|
||||
for i in range(50):
|
||||
assert resolver.get_engine(f"user_{i}") == "unified"
|
||||
|
||||
def test_invalid_default_engine_fallback(self):
|
||||
"""无效默认值回退到 legacy。"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
resolver = self._make_resolver(store=store, default="invalid_value")
|
||||
assert resolver.get_engine() == "legacy"
|
||||
|
||||
def test_should_use_unified_helper(self):
|
||||
"""should_use_unified 便捷方法。"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
store.set(
|
||||
FeatureFlagConfig(
|
||||
name="render_engine",
|
||||
enabled=True,
|
||||
percentage=0,
|
||||
whitelist={"user_a"},
|
||||
)
|
||||
)
|
||||
resolver = self._make_resolver(store=store)
|
||||
assert resolver.should_use_unified("user_a") is True
|
||||
assert resolver.should_use_unified("user_b") is False
|
||||
|
||||
def test_config_snapshot(self):
|
||||
"""配置快照。"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
store.set(
|
||||
FeatureFlagConfig(
|
||||
name="render_engine",
|
||||
enabled=True,
|
||||
percentage=30,
|
||||
whitelist={"u1", "u2"},
|
||||
)
|
||||
)
|
||||
resolver = self._make_resolver(store=store)
|
||||
snapshot = resolver.get_config_snapshot()
|
||||
assert snapshot["flag_name"] == "render_engine"
|
||||
assert snapshot["enabled"] is True
|
||||
assert snapshot["percentage"] == 30
|
||||
assert snapshot["whitelist"] == ["u1", "u2"]
|
||||
|
||||
def test_set_flag_updates_config(self):
|
||||
"""通过 set_flag 修改后立即生效。"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
resolver = self._make_resolver(store=store, default="legacy")
|
||||
|
||||
# 初始:关闭
|
||||
assert resolver.get_engine("user1") == "legacy"
|
||||
|
||||
# 开启 100%
|
||||
resolver.set_flag(FeatureFlagConfig(name="render_engine", enabled=True, percentage=100))
|
||||
assert resolver.get_engine("user1") == "unified"
|
||||
|
||||
# 关闭
|
||||
resolver.set_flag(FeatureFlagConfig(name="render_engine", enabled=False))
|
||||
assert resolver.get_engine("user1") == "legacy"
|
||||
|
||||
def test_force_refresh(self):
|
||||
"""强制刷新不报错。"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
resolver = self._make_resolver(store=store)
|
||||
resolver.force_refresh() # 不抛异常即可
|
||||
|
||||
def test_does_not_affect_in_flight_tasks(self):
|
||||
"""
|
||||
热更新不影响在途任务验证:
|
||||
任务开始时确定引擎,中途配置变更不改变当前任务的引擎选择。
|
||||
(这是通过"每次调用 get_engine 时读取当前配置"来保证的,
|
||||
任务开始时调用一次拿到结果,之后不再变化)
|
||||
"""
|
||||
store = InMemoryFeatureFlagStore()
|
||||
store.set(FeatureFlagConfig(name="render_engine", enabled=True, percentage=100))
|
||||
resolver = self._make_resolver(store=store, default="legacy")
|
||||
|
||||
# 模拟任务开始时获取引擎
|
||||
engine_at_start = resolver.get_engine("user1")
|
||||
assert engine_at_start == "unified"
|
||||
|
||||
# 任务进行中关闭 flag
|
||||
store.set(FeatureFlagConfig(name="render_engine", enabled=False))
|
||||
resolver.force_refresh()
|
||||
|
||||
# 在途任务持有的 engine_at_start 仍然是 unified(不随配置变化)
|
||||
assert engine_at_start == "unified"
|
||||
# 新任务会拿到 legacy
|
||||
assert resolver.get_engine("user1") == "legacy"
|
||||
|
||||
# ── 辅助方法 ──
|
||||
|
||||
@staticmethod
|
||||
def _make_resolver(store=None, default="legacy"):
|
||||
from apps.worker.video_processing.render_engine_resolver import (
|
||||
RenderEngineResolver,
|
||||
)
|
||||
|
||||
return RenderEngineResolver(
|
||||
default_engine=default,
|
||||
store=store or InMemoryFeatureFlagStore(),
|
||||
refresh_interval=9999, # 测试时禁用自动刷新
|
||||
)
|
||||
|
||||
|
||||
# ── RedisFeatureFlagStore 降级测试(无 Redis 环境) ───────────────────────
|
||||
|
||||
|
||||
class TestRedisStoreDegradation:
|
||||
"""Redis 不可用时的降级行为测试。"""
|
||||
|
||||
def test_get_returns_default_when_redis_unavailable(self):
|
||||
"""Redis 连接失败时返回默认关闭配置,不抛异常。"""
|
||||
import importlib
|
||||
|
||||
from packages.adapters.redis import feature_flag_store as ff_module
|
||||
|
||||
# 模拟 redis 模块不存在的场景不好做,这里直接测试异常捕获逻辑
|
||||
store = ff_module.RedisFeatureFlagStore.__new__(ff_module.RedisFeatureFlagStore)
|
||||
store._redis = MagicMock()
|
||||
store._redis.hgetall.side_effect = ConnectionError("Redis down")
|
||||
store._key_prefix = ff_module.FEATURE_FLAG_REDIS_PREFIX
|
||||
store._cache = {}
|
||||
store._cache_ttl = 5.0
|
||||
import threading
|
||||
|
||||
store._lock = threading.Lock()
|
||||
|
||||
config = store.get("render_engine")
|
||||
assert config.enabled is False
|
||||
assert config.name == "render_engine"
|
||||
|
||||
def test_list_all_returns_empty_on_redis_error(self):
|
||||
"""Redis 错误时 list_all 返回空字典。"""
|
||||
import importlib
|
||||
|
||||
from packages.adapters.redis import feature_flag_store as ff_module
|
||||
|
||||
store = ff_module.RedisFeatureFlagStore.__new__(ff_module.RedisFeatureFlagStore)
|
||||
store._redis = MagicMock()
|
||||
store._redis.scan.side_effect = ConnectionError("Redis down")
|
||||
store._key_prefix = ff_module.FEATURE_FLAG_REDIS_PREFIX
|
||||
store._cache = {}
|
||||
store._cache_ttl = 5.0
|
||||
import threading
|
||||
|
||||
store._lock = threading.Lock()
|
||||
|
||||
result = store.list_all()
|
||||
assert result == {}
|
||||
|
||||
|
||||
class TestRedisStoreListAll:
|
||||
"""RedisFeatureFlagStore list_all 正常路径测试。"""
|
||||
|
||||
def _make_store(self):
|
||||
from packages.adapters.redis import feature_flag_store as ff_module
|
||||
|
||||
store = ff_module.RedisFeatureFlagStore.__new__(ff_module.RedisFeatureFlagStore)
|
||||
store._redis = MagicMock()
|
||||
store._key_prefix = ff_module.FEATURE_FLAG_REDIS_PREFIX
|
||||
store._cache = {}
|
||||
store._cache_ttl = 5.0
|
||||
import threading
|
||||
|
||||
store._lock = threading.Lock()
|
||||
return store
|
||||
|
||||
def test_list_all_scan_with_match_param(self):
|
||||
"""list_all 调用 redis.scan 时使用正确的 match 参数名。"""
|
||||
store = self._make_store()
|
||||
prefix = store._key_prefix
|
||||
|
||||
# 模拟 scan 返回 2 个 key,分 2 次游标
|
||||
store._redis.scan.side_effect = [
|
||||
(10, [f"{prefix}render_engine", f"{prefix}other_flag"]),
|
||||
(0, []),
|
||||
]
|
||||
# 模拟 hgetall 返回配置
|
||||
store._redis.hgetall.return_value = {
|
||||
b"enabled": b"true",
|
||||
b"percentage": b"50",
|
||||
b"whitelist": b'["user1","user2"]',
|
||||
}
|
||||
|
||||
result = store.list_all()
|
||||
|
||||
# 验证 scan 被调用了 2 次(游标遍历)
|
||||
assert store._redis.scan.call_count == 2
|
||||
# 验证参数名是 match(不是 match_pattern)
|
||||
first_call_kwargs = store._redis.scan.call_args_list[0][1]
|
||||
assert "match" in first_call_kwargs
|
||||
assert "match_pattern" not in first_call_kwargs
|
||||
assert first_call_kwargs["match"] == f"{prefix}*"
|
||||
# 验证返回了 2 个 flag
|
||||
assert len(result) == 2
|
||||
assert "render_engine" in result
|
||||
assert "other_flag" in result
|
||||
@@ -0,0 +1,93 @@
|
||||
"""FFmpeg 超时保护测试。
|
||||
|
||||
验证 run_ffmpeg / probe_video_info 的超时保护机制,
|
||||
防止 FFmpeg hang 住导致 worker 永久阻塞。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import subprocess
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from video_processing.ffmpeg_utils import (
|
||||
DEFAULT_FFMPEG_TIMEOUT,
|
||||
probe_video_info,
|
||||
run_ffmpeg,
|
||||
)
|
||||
|
||||
# ── run_ffmpeg 超时保护 ──────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestRunFFmpegTimeout:
|
||||
"""run_ffmpeg 超时保护测试。"""
|
||||
|
||||
def test_default_timeout_is_set(self):
|
||||
"""默认超时应为 1800 秒(30分钟)。"""
|
||||
assert DEFAULT_FFMPEG_TIMEOUT == 1800
|
||||
|
||||
def test_timeout_expired_is_raised(self):
|
||||
"""超时未完成时 TimeoutExpired 异常被传播。"""
|
||||
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
|
||||
mock_run.side_effect = subprocess.TimeoutExpired(cmd=["ffmpeg", "test"], timeout=1)
|
||||
with pytest.raises(subprocess.TimeoutExpired):
|
||||
run_ffmpeg(["ffmpeg", "test"])
|
||||
|
||||
def test_custom_timeout(self):
|
||||
"""支持自定义超时时间。"""
|
||||
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
|
||||
mock_run.side_effect = subprocess.TimeoutExpired(cmd=["ffmpeg"], timeout=5)
|
||||
with pytest.raises(subprocess.TimeoutExpired):
|
||||
run_ffmpeg(["ffmpeg", "test"], timeout=5)
|
||||
|
||||
def test_none_timeout_disables_protection(self):
|
||||
"""timeout=None 可以禁用超时保护(不推荐)。"""
|
||||
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
|
||||
mock_result = MagicMock()
|
||||
mock_result.stdout = ""
|
||||
mock_result.stderr = ""
|
||||
mock_run.return_value = mock_result
|
||||
run_ffmpeg(["ffmpeg", "test"], timeout=None)
|
||||
# 验证 timeout=None 被传递
|
||||
call_kwargs = mock_run.call_args.kwargs
|
||||
assert call_kwargs["timeout"] is None
|
||||
|
||||
def test_called_process_error_still_raised(self):
|
||||
"""超时异常不影响原有 CalledProcessError 的抛出。"""
|
||||
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
|
||||
mock_run.side_effect = subprocess.CalledProcessError(returncode=1, cmd=["ffmpeg"], stderr="error msg")
|
||||
with pytest.raises(subprocess.CalledProcessError):
|
||||
run_ffmpeg(["ffmpeg", "test"])
|
||||
|
||||
|
||||
# ── probe_video_info 超时保护 ────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestProbeVideoInfoTimeout:
|
||||
"""probe_video_info 超时保护测试。"""
|
||||
|
||||
def test_probe_uses_timeout(self):
|
||||
"""probe_video_info 调用 ffprobe 时应设置 timeout=15。"""
|
||||
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
|
||||
mock_run.side_effect = subprocess.TimeoutExpired(cmd=["ffprobe"], timeout=15)
|
||||
# 超时异常被捕获,返回默认值
|
||||
result = probe_video_info("/tmp/test.mp4")
|
||||
assert result["width"] == 1280 # DEFAULT_OUTPUT_WIDTH
|
||||
assert result["height"] == 720 # DEFAULT_OUTPUT_HEIGHT
|
||||
|
||||
def test_probe_success(self):
|
||||
"""正常情况应解析 ffprobe JSON 输出。"""
|
||||
fake_output = """
|
||||
{
|
||||
"streams": [{"width": 1920, "height": 1080, "codec_type": "video", "r_frame_rate": "30/1", "duration": "10.5"}],
|
||||
"format": {"duration": "10.5"}
|
||||
}
|
||||
"""
|
||||
with patch("video_processing.ffmpeg_utils.subprocess.run") as mock_run:
|
||||
mock_result = MagicMock()
|
||||
mock_result.stdout = fake_output
|
||||
mock_run.return_value = mock_result
|
||||
result = probe_video_info("/tmp/test.mp4")
|
||||
assert result["width"] == 1920
|
||||
assert result["height"] == 1080
|
||||
assert abs(result["duration"] - 10.5) < 0.01
|
||||
@@ -0,0 +1,338 @@
|
||||
"""generate_video 任务 Feature Flag 灰度引擎选择单元测试.
|
||||
|
||||
覆盖:
|
||||
- _resolve_render_engine 正常返回 unified / legacy
|
||||
- Feature Flag 不可用时 fallback 到 unified
|
||||
- 白名单 / 百分比 / 全局开关各场景
|
||||
- _render_with_legacy_engine 命令构建与输出验证
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
from types import ModuleType
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
|
||||
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
|
||||
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
|
||||
# ── Mock worker 模块以避免数据库连接 ──────────────────────────────────────────
|
||||
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "worker"))
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", ".."))
|
||||
|
||||
_mock_db_mod = ModuleType("worker_app.db")
|
||||
_mock_db_mod.SessionLocal = MagicMock()
|
||||
sys.modules.setdefault("worker_app.db", _mock_db_mod)
|
||||
|
||||
_mock_celery_mod = ModuleType("worker_app.celery_app")
|
||||
_mock_celery_app = MagicMock()
|
||||
_mock_celery_app.task = lambda **kwargs: lambda fn: fn
|
||||
_mock_celery_mod.celery_app = _mock_celery_app
|
||||
sys.modules.setdefault("worker_app.celery_app", _mock_celery_mod)
|
||||
|
||||
# Mock worker_app.core.config 避免 settings 加载
|
||||
_mock_config_mod = ModuleType("worker_app.core.config")
|
||||
_mock_settings = MagicMock()
|
||||
_mock_settings.redis_url = None
|
||||
_mock_settings.render_engine = "unified"
|
||||
_mock_config_mod.get_settings = lambda: _mock_settings
|
||||
sys.modules.setdefault("worker_app.core", ModuleType("worker_app.core"))
|
||||
sys.modules.setdefault("worker_app.core.config", _mock_config_mod)
|
||||
|
||||
|
||||
# ── 测试用数据类 ──────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class _TestClip:
|
||||
def __init__(self, asset_id, duration=30.0, clip_type="main", config=None, order=0):
|
||||
self.id = f"clip_{asset_id}"
|
||||
self.plan_id = "test-plan"
|
||||
self.clip_type = clip_type
|
||||
self.order = order
|
||||
self.asset_id = asset_id
|
||||
self.duration = duration
|
||||
self.config = config or {}
|
||||
self.start_time = 0.0
|
||||
self.transition_effect = "cut"
|
||||
|
||||
|
||||
# ── RenderEngineResolver 基础行为测试 ───────────────────────────────────────
|
||||
|
||||
|
||||
def test_resolver_unified_when_enabled_100_percent():
|
||||
"""flag 全局开启(percentage=100)时,返回 unified。"""
|
||||
from video_processing.render_engine_resolver import RenderEngineResolver
|
||||
|
||||
from packages.adapters.redis.feature_flag_store import (
|
||||
FeatureFlagConfig,
|
||||
InMemoryFeatureFlagStore,
|
||||
)
|
||||
|
||||
store = InMemoryFeatureFlagStore()
|
||||
store.set(FeatureFlagConfig(name="render_engine", enabled=True, percentage=100))
|
||||
resolver = RenderEngineResolver(default_engine="legacy", store=store)
|
||||
|
||||
assert resolver.get_engine(user_id="user-123") == "unified"
|
||||
|
||||
|
||||
def test_resolver_legacy_when_flag_disabled():
|
||||
"""flag 全局关闭时,返回默认引擎 legacy。"""
|
||||
from video_processing.render_engine_resolver import RenderEngineResolver
|
||||
|
||||
from packages.adapters.redis.feature_flag_store import (
|
||||
FeatureFlagConfig,
|
||||
InMemoryFeatureFlagStore,
|
||||
)
|
||||
|
||||
store = InMemoryFeatureFlagStore()
|
||||
store.set(FeatureFlagConfig(name="render_engine", enabled=False, percentage=100))
|
||||
resolver = RenderEngineResolver(default_engine="legacy", store=store)
|
||||
|
||||
assert resolver.get_engine(user_id="user-123") == "legacy"
|
||||
|
||||
|
||||
def test_resolver_whitelist_overrides_percentage_0():
|
||||
"""白名单用户即使 percentage=0 也走 unified。"""
|
||||
from video_processing.render_engine_resolver import RenderEngineResolver
|
||||
|
||||
from packages.adapters.redis.feature_flag_store import (
|
||||
FeatureFlagConfig,
|
||||
InMemoryFeatureFlagStore,
|
||||
)
|
||||
|
||||
store = InMemoryFeatureFlagStore()
|
||||
store.set(
|
||||
FeatureFlagConfig(
|
||||
name="render_engine",
|
||||
enabled=True,
|
||||
percentage=0,
|
||||
whitelist={"user-vip"},
|
||||
)
|
||||
)
|
||||
resolver = RenderEngineResolver(default_engine="legacy", store=store)
|
||||
|
||||
assert resolver.get_engine(user_id="user-vip") == "unified"
|
||||
assert resolver.get_engine(user_id="user-other") == "legacy"
|
||||
|
||||
|
||||
def test_resolver_percentage_0_all_legacy():
|
||||
"""percentage=0 且无白名单时,全部走 legacy。"""
|
||||
from video_processing.render_engine_resolver import RenderEngineResolver
|
||||
|
||||
from packages.adapters.redis.feature_flag_store import (
|
||||
FeatureFlagConfig,
|
||||
InMemoryFeatureFlagStore,
|
||||
)
|
||||
|
||||
store = InMemoryFeatureFlagStore()
|
||||
store.set(FeatureFlagConfig(name="render_engine", enabled=True, percentage=0))
|
||||
resolver = RenderEngineResolver(default_engine="legacy", store=store)
|
||||
|
||||
for i in range(50):
|
||||
assert resolver.get_engine(user_id=f"user-{i}") == "legacy"
|
||||
|
||||
|
||||
def test_resolver_default_unified_when_flag_off():
|
||||
"""默认引擎设为 unified 且 flag 关闭时,返回 unified。"""
|
||||
from video_processing.render_engine_resolver import RenderEngineResolver
|
||||
|
||||
from packages.adapters.redis.feature_flag_store import (
|
||||
FeatureFlagConfig,
|
||||
InMemoryFeatureFlagStore,
|
||||
)
|
||||
|
||||
store = InMemoryFeatureFlagStore()
|
||||
store.set(FeatureFlagConfig(name="render_engine", enabled=False, percentage=0))
|
||||
resolver = RenderEngineResolver(default_engine="unified", store=store)
|
||||
|
||||
assert resolver.get_engine(user_id="user-123") == "unified"
|
||||
|
||||
|
||||
# ── _render_with_legacy_engine 集成测试 ──────────────────────────────────────
|
||||
|
||||
|
||||
def test_legacy_engine_single_clip_keeps_original_fps():
|
||||
"""单 clip 场景:输出保持原帧率(不做 fps 归一化),分辨率缩放正确。"""
|
||||
import subprocess
|
||||
import tempfile
|
||||
|
||||
from video_processing.ffmpeg_utils import probe_video_info
|
||||
|
||||
from apps.worker.worker_app.tasks.generation import _render_with_legacy_engine
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
tmp_path = Path(tmpdir)
|
||||
input_path = tmp_path / "input.mp4"
|
||||
output_path = tmp_path / "output.mp4"
|
||||
|
||||
# 生成 1 秒 30fps 测试视频(带音频)
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
"color=c=red:s=640x360:d=1:r=30",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
"anullsrc=r=44100:cl=stereo:d=1",
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-shortest",
|
||||
str(input_path),
|
||||
],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
)
|
||||
|
||||
clip = _TestClip(asset_id="asset-1", duration=1.0)
|
||||
asset_path_map = {"asset-1": input_path}
|
||||
|
||||
duration, file_size = _render_with_legacy_engine(
|
||||
task_id="test-task",
|
||||
virtual_clips=[clip],
|
||||
asset_path_map=asset_path_map,
|
||||
work_dir=tmp_path,
|
||||
output_path=output_path,
|
||||
)
|
||||
|
||||
assert output_path.exists()
|
||||
assert file_size > 0
|
||||
assert duration > 0
|
||||
|
||||
# 旧引擎保持原帧率(30fps),不做 fps 归一化
|
||||
info = probe_video_info(str(output_path))
|
||||
assert abs(info.get("fps", 0) - 30.0) < 0.5
|
||||
assert info.get("width") == 1280
|
||||
assert info.get("height") == 720
|
||||
|
||||
|
||||
def test_legacy_engine_two_clips_concat_duration():
|
||||
"""多 clip 场景:concat 后时长为两片段之和。"""
|
||||
import subprocess
|
||||
import tempfile
|
||||
|
||||
from video_processing.ffmpeg_utils import probe_duration
|
||||
|
||||
from apps.worker.worker_app.tasks.generation import _render_with_legacy_engine
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
tmp_path = Path(tmpdir)
|
||||
input1 = tmp_path / "input1.mp4"
|
||||
input2 = tmp_path / "input2.mp4"
|
||||
output_path = tmp_path / "output.mp4"
|
||||
|
||||
for idx, inp in enumerate([input1, input2]):
|
||||
color = "red" if idx == 0 else "blue"
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
f"color=c={color}:s=640x360:d=1:r=30",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
"anullsrc=r=44100:cl=stereo:d=1",
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-shortest",
|
||||
str(inp),
|
||||
],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
)
|
||||
|
||||
clip1 = _TestClip(asset_id="asset-1", duration=1.0, clip_type="main", order=0)
|
||||
clip2 = _TestClip(asset_id="asset-2", duration=1.0, clip_type="main", order=1)
|
||||
asset_path_map = {"asset-1": input1, "asset-2": input2}
|
||||
|
||||
duration, file_size = _render_with_legacy_engine(
|
||||
task_id="test-task",
|
||||
virtual_clips=[clip1, clip2],
|
||||
asset_path_map=asset_path_map,
|
||||
work_dir=tmp_path,
|
||||
output_path=output_path,
|
||||
)
|
||||
|
||||
assert output_path.exists()
|
||||
assert file_size > 0
|
||||
assert abs(duration - 2.0) < 0.2
|
||||
|
||||
|
||||
def test_legacy_engine_broll_mode_supported():
|
||||
"""b_roll 类型的 clip 也被正确识别为主图层并渲染。"""
|
||||
import subprocess
|
||||
import tempfile
|
||||
|
||||
from apps.worker.worker_app.tasks.generation import _render_with_legacy_engine
|
||||
|
||||
with tempfile.TemporaryDirectory() as tmpdir:
|
||||
tmp_path = Path(tmpdir)
|
||||
input_path = tmp_path / "input.mp4"
|
||||
output_path = tmp_path / "output.mp4"
|
||||
|
||||
subprocess.run(
|
||||
[
|
||||
"ffmpeg",
|
||||
"-y",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
"color=c=green:s=640x360:d=1:r=30",
|
||||
"-f",
|
||||
"lavfi",
|
||||
"-i",
|
||||
"anullsrc=r=44100:cl=stereo:d=1",
|
||||
"-c:v",
|
||||
"libx264",
|
||||
"-pix_fmt",
|
||||
"yuv420p",
|
||||
"-c:a",
|
||||
"aac",
|
||||
"-shortest",
|
||||
str(input_path),
|
||||
],
|
||||
check=True,
|
||||
capture_output=True,
|
||||
)
|
||||
|
||||
clip = _TestClip(
|
||||
asset_id="asset-1",
|
||||
duration=1.0,
|
||||
clip_type="main",
|
||||
config={"role": "b_roll"},
|
||||
)
|
||||
asset_path_map = {"asset-1": input_path}
|
||||
|
||||
duration, file_size = _render_with_legacy_engine(
|
||||
task_id="test-task",
|
||||
virtual_clips=[clip],
|
||||
asset_path_map=asset_path_map,
|
||||
work_dir=tmp_path,
|
||||
output_path=output_path,
|
||||
)
|
||||
|
||||
assert output_path.exists()
|
||||
assert file_size > 0
|
||||
assert duration > 0
|
||||
@@ -0,0 +1,190 @@
|
||||
"""渲染结果内部下载接口单元测试。
|
||||
|
||||
测试 internal_render 路由的核心逻辑,mock 掉 repository 和 storage 依赖。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from app.api.routes.internal_render import (
|
||||
InternalRenderDownloadUrlResponse,
|
||||
InternalRenderTaskVideosResponse,
|
||||
_video_to_item,
|
||||
get_render_task_videos,
|
||||
get_render_video_download_url,
|
||||
)
|
||||
|
||||
# ── Helpers ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class MockVideo:
|
||||
"""模拟 GeneratedVideo 领域对象。"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
self.id = kwargs.get("id", "video-1")
|
||||
self.generation_task_id = kwargs.get("generation_task_id", "task-1")
|
||||
self.project_id = kwargs.get("project_id", "proj-1")
|
||||
self.name = kwargs.get("name", "test_video.mp4")
|
||||
self.file_url = kwargs.get("file_url", "videos/test/output.mp4")
|
||||
self.file_size = kwargs.get("file_size", 1024000)
|
||||
self.duration = kwargs.get("duration", 30.5)
|
||||
self.width = kwargs.get("width", 1080)
|
||||
self.height = kwargs.get("height", 1920)
|
||||
self.fps = kwargs.get("fps", 30.0)
|
||||
self.status = kwargs.get("status", "completed")
|
||||
|
||||
|
||||
# ── _video_to_item 测试 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestVideoToItem:
|
||||
"""测试视频对象转响应项。"""
|
||||
|
||||
def test_basic_conversion(self):
|
||||
video = MockVideo(id="v1", generation_task_id="t1", status="completed")
|
||||
item = _video_to_item(video, "https://oss.example.com/download?v1")
|
||||
assert item.video_id == "v1"
|
||||
assert item.generation_task_id == "t1"
|
||||
assert item.status == "completed"
|
||||
assert item.download_url == "https://oss.example.com/download?v1"
|
||||
|
||||
def test_missing_optional_fields(self):
|
||||
"""缺可选字段时返回 None。"""
|
||||
video = MockVideo()
|
||||
# 去掉可选字段
|
||||
del video.file_size
|
||||
del video.duration
|
||||
item = _video_to_item(video, "https://example.com/dl")
|
||||
assert item.file_size is None
|
||||
assert item.duration is None
|
||||
assert item.width == 1080 # 还在
|
||||
|
||||
|
||||
# ── 路由函数测试 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestGetRenderVideoDownloadUrl:
|
||||
"""测试单个视频下载URL接口。"""
|
||||
|
||||
def test_video_exists(self):
|
||||
video = MockVideo(id="v-abc", file_url="videos/abc/out.mp4")
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.get.return_value = video
|
||||
mock_storage = MagicMock()
|
||||
mock_storage.get_download_url.return_value = "https://oss.test/signed?v=abc"
|
||||
|
||||
result = get_render_video_download_url(
|
||||
video_id="v-abc",
|
||||
_=True,
|
||||
generated_video_repository=mock_repo,
|
||||
storage_service=mock_storage,
|
||||
)
|
||||
|
||||
assert isinstance(result, InternalRenderDownloadUrlResponse)
|
||||
assert result.video_id == "v-abc"
|
||||
assert result.download_url == "https://oss.test/signed?v=abc"
|
||||
mock_repo.get.assert_called_once_with("v-abc")
|
||||
mock_storage.get_download_url.assert_called_once()
|
||||
|
||||
def test_video_not_found_raises_404(self):
|
||||
from fastapi import HTTPException
|
||||
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.get.return_value = None
|
||||
mock_storage = MagicMock()
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
get_render_video_download_url(
|
||||
video_id="nonexistent",
|
||||
_=True,
|
||||
generated_video_repository=mock_repo,
|
||||
storage_service=mock_storage,
|
||||
)
|
||||
assert exc_info.value.status_code == 404
|
||||
|
||||
def test_download_url_long_expiry(self):
|
||||
"""过期时间应为 24 小时(86400s)。"""
|
||||
video = MockVideo(id="v1")
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.get.return_value = video
|
||||
mock_storage = MagicMock()
|
||||
mock_storage.get_download_url.return_value = "https://oss.test/signed"
|
||||
|
||||
get_render_video_download_url(
|
||||
video_id="v1",
|
||||
_=True,
|
||||
generated_video_repository=mock_repo,
|
||||
storage_service=mock_storage,
|
||||
)
|
||||
|
||||
# 验证 expires_seconds=86400
|
||||
call_kwargs = mock_storage.get_download_url.call_args
|
||||
assert call_kwargs.kwargs.get("expires_seconds") == 86400 or call_kwargs[1].get("expires_seconds") == 86400
|
||||
|
||||
|
||||
class TestGetRenderTaskVideos:
|
||||
"""测试任务视频列表接口。"""
|
||||
|
||||
def test_list_multiple_videos(self):
|
||||
videos = [
|
||||
MockVideo(id="v1", status="completed"),
|
||||
MockVideo(id="v2", status="completed"),
|
||||
MockVideo(id="v3", status="failed"),
|
||||
]
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_by_generation_task.return_value = videos
|
||||
mock_storage = MagicMock()
|
||||
mock_storage.get_download_url.return_value = "https://oss.test/signed"
|
||||
|
||||
result = get_render_task_videos(
|
||||
task_id="task-1",
|
||||
status=None,
|
||||
_=True,
|
||||
generated_video_repository=mock_repo,
|
||||
storage_service=mock_storage,
|
||||
)
|
||||
|
||||
assert isinstance(result, InternalRenderTaskVideosResponse)
|
||||
assert result.task_id == "task-1"
|
||||
assert result.count == 3
|
||||
assert len(result.videos) == 3
|
||||
|
||||
def test_filter_by_status(self):
|
||||
videos = [
|
||||
MockVideo(id="v1", status="completed"),
|
||||
MockVideo(id="v2", status="completed"),
|
||||
MockVideo(id="v3", status="failed"),
|
||||
]
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_by_generation_task.return_value = videos
|
||||
mock_storage = MagicMock()
|
||||
mock_storage.get_download_url.return_value = "https://oss.test/signed"
|
||||
|
||||
result = get_render_task_videos(
|
||||
task_id="task-1",
|
||||
status="completed",
|
||||
_=True,
|
||||
generated_video_repository=mock_repo,
|
||||
storage_service=mock_storage,
|
||||
)
|
||||
|
||||
assert result.count == 2
|
||||
assert all(v.status == "completed" for v in result.videos)
|
||||
|
||||
def test_empty_task(self):
|
||||
mock_repo = MagicMock()
|
||||
mock_repo.list_by_generation_task.return_value = []
|
||||
mock_storage = MagicMock()
|
||||
|
||||
result = get_render_task_videos(
|
||||
task_id="empty-task",
|
||||
status=None,
|
||||
_=True,
|
||||
generated_video_repository=mock_repo,
|
||||
storage_service=mock_storage,
|
||||
)
|
||||
|
||||
assert result.count == 0
|
||||
assert result.videos == []
|
||||
@@ -132,10 +132,8 @@ class TestDownloadLibraryAssets:
|
||||
session.query.return_value = query
|
||||
filter_result = MagicMock()
|
||||
query.filter.return_value = filter_result
|
||||
in_filter = MagicMock()
|
||||
filter_result.filter.return_value = in_filter
|
||||
id_filter = MagicMock()
|
||||
in_filter.filter.return_value = id_filter
|
||||
filter_result.filter.return_value = id_filter
|
||||
assets = [self._make_asset("a1", "video/a1.mp4")]
|
||||
id_filter.order_by.return_value.all.return_value = assets
|
||||
|
||||
|
||||
Executable
+236
@@ -0,0 +1,236 @@
|
||||
"""P0-staging:OSS 上传崩溃修复测试.
|
||||
|
||||
测试:
|
||||
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()
|
||||
Executable
+424
@@ -0,0 +1,424 @@
|
||||
"""RenderAdapter 单元测试 — Phase 2.
|
||||
|
||||
测试适配层的计划加载、素材下载、引擎调用、结果上传等逻辑。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from video_processing.render_adapter import RenderAdapter, RenderAdapterResult
|
||||
|
||||
# ── Fixtures ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakeClip:
|
||||
"""模拟 EditPlanClip。"""
|
||||
|
||||
id: str
|
||||
plan_id: str = "plan_001"
|
||||
clip_type: str = "main"
|
||||
order: int = 0
|
||||
asset_id: str = ""
|
||||
text_content: str = ""
|
||||
start_time: float = 0.0
|
||||
duration: float = 0.0
|
||||
transition_effect: str = "cut"
|
||||
status: str = "ready"
|
||||
config: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class FakePlan:
|
||||
"""模拟 EditPlan。"""
|
||||
|
||||
id: str = "plan_001"
|
||||
name: str = "测试计划"
|
||||
status: str = "editing"
|
||||
config: dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
|
||||
def _make_clip(
|
||||
clip_id: str,
|
||||
clip_type: str = "main",
|
||||
order: int = 0,
|
||||
asset_id: str | None = None,
|
||||
duration: float = 5.0,
|
||||
status: str = "ready",
|
||||
transition_effect: str = "cut",
|
||||
config: dict[str, Any] | None = None,
|
||||
) -> FakeClip:
|
||||
# asset_id 为 None 时生成默认值,为空字符串时保留空串
|
||||
if asset_id is None:
|
||||
asset_id = f"asset_{clip_id}.mp4"
|
||||
return FakeClip(
|
||||
id=clip_id,
|
||||
clip_type=clip_type,
|
||||
order=order,
|
||||
asset_id=asset_id,
|
||||
duration=duration,
|
||||
status=status,
|
||||
transition_effect=transition_effect,
|
||||
config=config or {},
|
||||
)
|
||||
|
||||
|
||||
def _make_adapter(
|
||||
plan: FakePlan | None = None,
|
||||
clips: list[FakeClip] | None = None,
|
||||
) -> tuple[RenderAdapter, MagicMock, MagicMock]:
|
||||
"""创建测试用的 RenderAdapter 及 mock repo。
|
||||
|
||||
Returns:
|
||||
(adapter, mock_plan_repo, mock_clip_repo)
|
||||
"""
|
||||
mock_db = MagicMock()
|
||||
adapter = RenderAdapter(mock_db)
|
||||
|
||||
# 替换内部 repo
|
||||
mock_plan_repo = MagicMock()
|
||||
mock_clip_repo = MagicMock()
|
||||
adapter._plan_repo = mock_plan_repo
|
||||
adapter._clip_repo = mock_clip_repo
|
||||
|
||||
# 设置默认返回
|
||||
if plan is not None:
|
||||
mock_plan_repo.get.return_value = plan
|
||||
if clips is not None:
|
||||
mock_clip_repo.list_by_plan.return_value = clips
|
||||
|
||||
return adapter, mock_plan_repo, mock_clip_repo
|
||||
|
||||
|
||||
# ── validate_plan 测试 ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestValidatePlan:
|
||||
def test_plan_not_found(self):
|
||||
"""计划不存在时校验失败。"""
|
||||
adapter, mock_plan_repo, _ = _make_adapter(plan=None)
|
||||
mock_plan_repo.get.return_value = None
|
||||
|
||||
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
|
||||
|
||||
assert not valid
|
||||
assert len(errors) == 1
|
||||
assert "不存在" in errors[0]
|
||||
assert ready_count == 0
|
||||
assert total_count == 0
|
||||
|
||||
def test_no_clips(self):
|
||||
"""没有任何片段时校验失败。"""
|
||||
plan = FakePlan(id="plan_001", status="editing")
|
||||
adapter, _, mock_clip_repo = _make_adapter(plan=plan, clips=[])
|
||||
|
||||
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
|
||||
|
||||
assert not valid
|
||||
assert any("没有任何片段" in e for e in errors)
|
||||
|
||||
def test_no_ready_clips(self):
|
||||
"""没有 ready 片段时校验失败。"""
|
||||
plan = FakePlan(id="plan_001", status="editing")
|
||||
clips = [
|
||||
_make_clip("c1", status="pending"),
|
||||
_make_clip("c2", status="pending"),
|
||||
]
|
||||
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
|
||||
|
||||
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
|
||||
|
||||
assert not valid
|
||||
assert any("没有就绪" in e for e in errors)
|
||||
assert ready_count == 0
|
||||
assert total_count == 2
|
||||
|
||||
def test_ready_clip_no_asset(self):
|
||||
"""ready 片段没有 asset_id 时报错。"""
|
||||
plan = FakePlan(id="plan_001", status="editing")
|
||||
clips = [
|
||||
_make_clip("c1", asset_id=""),
|
||||
]
|
||||
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
|
||||
|
||||
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
|
||||
|
||||
assert not valid
|
||||
assert any("没有分配素材" in e for e in errors)
|
||||
|
||||
def test_valid_plan(self):
|
||||
"""正常计划校验通过。"""
|
||||
plan = FakePlan(id="plan_001", status="editing")
|
||||
clips = [
|
||||
_make_clip("c1", order=0, duration=3.0),
|
||||
_make_clip("c2", order=1, duration=4.0),
|
||||
]
|
||||
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
|
||||
|
||||
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
|
||||
|
||||
assert valid
|
||||
assert len(errors) == 0
|
||||
assert ready_count == 2
|
||||
assert total_count == 2
|
||||
|
||||
def test_wrong_status(self):
|
||||
"""计划状态不正确时报错。"""
|
||||
plan = FakePlan(id="plan_001", status="draft")
|
||||
clips = [_make_clip("c1")]
|
||||
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
|
||||
|
||||
valid, errors, _, _, _ = adapter.validate_plan("plan_001")
|
||||
|
||||
assert not valid
|
||||
assert any("状态不正确" in e for e in errors)
|
||||
|
||||
def test_mixed_status_with_warnings(self):
|
||||
"""混合状态时有 pending/failed 警告。"""
|
||||
plan = FakePlan(id="plan_001", status="editing")
|
||||
clips = [
|
||||
_make_clip("c1", order=0, status="ready"),
|
||||
_make_clip("c2", order=1, status="pending"),
|
||||
_make_clip("c3", order=2, status="failed"),
|
||||
]
|
||||
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
|
||||
|
||||
valid, errors, warnings, ready_count, total_count = adapter.validate_plan("plan_001")
|
||||
|
||||
assert valid
|
||||
assert any("pending" in w for w in warnings)
|
||||
assert any("failed" in w for w in warnings)
|
||||
assert ready_count == 1
|
||||
assert total_count == 3
|
||||
|
||||
|
||||
# ── render_plan 测试 ─────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestRenderPlan:
|
||||
def test_plan_not_found(self):
|
||||
"""计划不存在时返回失败。"""
|
||||
adapter, mock_plan_repo, _ = _make_adapter(plan=None)
|
||||
mock_plan_repo.get.return_value = None
|
||||
|
||||
result = adapter.render_plan("plan_001")
|
||||
|
||||
assert not result.success
|
||||
assert "不存在" in result.error_message
|
||||
|
||||
def test_no_ready_clips(self):
|
||||
"""没有 ready 片段时返回失败。"""
|
||||
plan = FakePlan(id="plan_001", status="editing")
|
||||
clips = [_make_clip("c1", status="pending")]
|
||||
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
|
||||
|
||||
result = adapter.render_plan("plan_001")
|
||||
|
||||
assert not result.success
|
||||
assert "没有可渲染" in result.error_message
|
||||
assert result.clip_count == 0
|
||||
|
||||
@patch("video_processing.render_adapter.download_asset")
|
||||
def test_all_assets_download_fail(self, mock_download):
|
||||
"""所有素材下载失败时返回失败。"""
|
||||
mock_download.return_value = False
|
||||
|
||||
plan = FakePlan(id="plan_001", status="editing")
|
||||
clips = [_make_clip("c1", order=0, duration=5.0)]
|
||||
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
|
||||
|
||||
result = adapter.render_plan("plan_001")
|
||||
|
||||
assert not result.success
|
||||
assert "素材下载失败" in result.error_message
|
||||
|
||||
@patch("video_processing.render_adapter.upload_to_oss")
|
||||
@patch("video_processing.render_adapter.UnifiedRenderService")
|
||||
@patch("video_processing.render_adapter.download_asset")
|
||||
def test_successful_render(self, mock_download, mock_render_cls, mock_upload, tmp_path):
|
||||
"""完整渲染流程成功。"""
|
||||
|
||||
# 素材下载成功
|
||||
def _fake_download(asset_id, local_path):
|
||||
local_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
local_path.write_bytes(b"fake video data")
|
||||
return True
|
||||
|
||||
mock_download.side_effect = _fake_download
|
||||
|
||||
# 渲染成功
|
||||
mock_render = MagicMock()
|
||||
mock_render.render.return_value = MagicMock(
|
||||
output_path=tmp_path / "output.mp4",
|
||||
duration=10.0,
|
||||
file_size=102400,
|
||||
width=1280,
|
||||
height=720,
|
||||
)
|
||||
mock_render_cls.return_value = mock_render
|
||||
|
||||
# 上传成功
|
||||
mock_upload.return_value = "https://oss.example.com/rendered/plan_001/job_001.mp4"
|
||||
|
||||
plan = FakePlan(id="plan_001", status="editing")
|
||||
clips = [
|
||||
_make_clip("c1", order=0, duration=5.0),
|
||||
_make_clip("c2", order=1, duration=5.0),
|
||||
]
|
||||
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
|
||||
|
||||
result = adapter.render_plan(
|
||||
"plan_001",
|
||||
job_id="job_001",
|
||||
work_dir=tmp_path / "work",
|
||||
)
|
||||
|
||||
assert result.success
|
||||
assert result.output_url.startswith("https://")
|
||||
assert result.duration == 10.0
|
||||
assert result.width == 1280
|
||||
assert result.height == 720
|
||||
assert result.clip_count == 2
|
||||
|
||||
# 验证 UnifiedRenderService 被正确调用
|
||||
mock_render_cls.assert_called_once()
|
||||
call_kwargs = mock_render_cls.call_args
|
||||
assert call_kwargs.kwargs["plan"] is plan
|
||||
assert len(call_kwargs.kwargs["clips"]) == 2
|
||||
assert len(call_kwargs.kwargs["asset_path_map"]) == 2
|
||||
|
||||
@patch("video_processing.render_adapter.download_asset")
|
||||
def test_progress_callback(self, mock_download, tmp_path):
|
||||
"""进度回调被正确触发。"""
|
||||
|
||||
def _fake_download(asset_id, local_path):
|
||||
local_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
local_path.write_bytes(b"fake data")
|
||||
return True
|
||||
|
||||
mock_download.side_effect = _fake_download
|
||||
|
||||
# 模拟渲染异常,避免走到最后
|
||||
with patch("video_processing.render_adapter.UnifiedRenderService") as mock_render_cls:
|
||||
mock_render = MagicMock()
|
||||
mock_render.render.side_effect = RuntimeError("render error")
|
||||
mock_render_cls.return_value = mock_render
|
||||
|
||||
plan = FakePlan(id="plan_001", status="editing")
|
||||
clips = [_make_clip("c1", order=0, duration=5.0)]
|
||||
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
|
||||
|
||||
progress_values = []
|
||||
|
||||
def progress_cb(progress: float, stage: str) -> None:
|
||||
progress_values.append((progress, stage))
|
||||
|
||||
result = adapter.render_plan(
|
||||
"plan_001",
|
||||
work_dir=tmp_path / "work",
|
||||
progress_cb=progress_cb,
|
||||
)
|
||||
|
||||
# 即使渲染失败,前期进度也应该上报了
|
||||
assert len(progress_values) > 0
|
||||
# 第一个进度应该是加载计划
|
||||
assert progress_values[0][1] == "加载剪辑计划"
|
||||
|
||||
@patch("video_processing.render_adapter.download_asset")
|
||||
def test_partial_asset_download(self, mock_download, tmp_path):
|
||||
"""部分素材下载失败时,只使用成功的素材。"""
|
||||
download_results = [True, False, True] # 3个素材中2个成功
|
||||
|
||||
def _fake_download(asset_id, local_path):
|
||||
idx = hash(asset_id) % 3
|
||||
if download_results[idx]:
|
||||
local_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
local_path.write_bytes(b"fake data")
|
||||
return True
|
||||
return False
|
||||
|
||||
mock_download.side_effect = _fake_download
|
||||
|
||||
with patch("video_processing.render_adapter.UnifiedRenderService") as mock_render_cls:
|
||||
mock_render = MagicMock()
|
||||
mock_render.render.return_value = MagicMock(
|
||||
output_path=tmp_path / "out.mp4",
|
||||
duration=5.0,
|
||||
file_size=1024,
|
||||
width=1280,
|
||||
height=720,
|
||||
)
|
||||
mock_render_cls.return_value = mock_render
|
||||
|
||||
with patch("video_processing.render_adapter.upload_to_oss", return_value="https://example.com/out.mp4"):
|
||||
plan = FakePlan(id="plan_001", status="editing")
|
||||
clips = [
|
||||
_make_clip("c1", order=0, duration=3.0, asset_id="asset_001.mp4"),
|
||||
_make_clip("c2", order=1, duration=3.0, asset_id="asset_002.mp4"),
|
||||
_make_clip("c3", order=2, duration=3.0, asset_id="asset_003.mp4"),
|
||||
]
|
||||
adapter, _, _ = _make_adapter(plan=plan, clips=clips)
|
||||
|
||||
result = adapter.render_plan(
|
||||
"plan_001",
|
||||
work_dir=tmp_path / "work",
|
||||
)
|
||||
|
||||
# 至少有部分素材成功,渲染应该进行
|
||||
# (具体成功数量取决于 hash 结果,但至少1个成功就能渲染)
|
||||
assert result.success or "素材下载失败" in result.error_message
|
||||
|
||||
|
||||
# ── _download_assets 测试 ────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TestDownloadAssets:
|
||||
@patch("video_processing.render_adapter.download_asset")
|
||||
def test_all_download_success(self, mock_download, tmp_path):
|
||||
"""全部素材下载成功。"""
|
||||
mock_download.return_value = True
|
||||
|
||||
clips = [
|
||||
_make_clip("c1", order=0, asset_id="key1.mp4"),
|
||||
_make_clip("c2", order=1, asset_id="key2.mp4"),
|
||||
]
|
||||
|
||||
result = RenderAdapter._download_assets(clips, tmp_path)
|
||||
|
||||
assert len(result) == 2
|
||||
assert "key1.mp4" in result
|
||||
assert "key2.mp4" in result
|
||||
assert mock_download.call_count == 2
|
||||
|
||||
@patch("video_processing.render_adapter.download_asset")
|
||||
def test_empty_asset_id_skipped(self, mock_download, tmp_path):
|
||||
"""空 asset_id 的片段被跳过。"""
|
||||
clips = [
|
||||
_make_clip("c1", order=0, asset_id=""),
|
||||
_make_clip("c2", order=1, asset_id="key2.mp4"),
|
||||
]
|
||||
mock_download.return_value = True
|
||||
|
||||
result = RenderAdapter._download_assets(clips, tmp_path)
|
||||
|
||||
assert len(result) == 1
|
||||
assert "key2.mp4" in result
|
||||
assert mock_download.call_count == 1 # 只调用了一次下载
|
||||
|
||||
@patch("video_processing.render_adapter.download_asset")
|
||||
def test_all_download_fail(self, mock_download, tmp_path):
|
||||
"""全部下载失败返回空字典。"""
|
||||
mock_download.return_value = False
|
||||
|
||||
clips = [
|
||||
_make_clip("c1", order=0, asset_id="key1.mp4"),
|
||||
]
|
||||
|
||||
result = RenderAdapter._download_assets(clips, tmp_path)
|
||||
|
||||
assert len(result) == 0
|
||||
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user