Compare commits
1 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| f1597749f1 |
@@ -1,4 +1,4 @@
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime
|
||||
|
||||
import psycopg2
|
||||
import redis
|
||||
@@ -13,7 +13,7 @@ router = APIRouter(tags=["Health"])
|
||||
async def health_check():
|
||||
return {
|
||||
"status": "healthy",
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"timestamp": datetime.utcnow().isoformat(),
|
||||
"version": settings.APP_VERSION,
|
||||
}
|
||||
|
||||
@@ -33,7 +33,7 @@ async def startup_check():
|
||||
all_ready = all(check["status"] == "healthy" for check in checks.values())
|
||||
response = {
|
||||
"status": "started" if all_ready else "starting",
|
||||
"timestamp": datetime.now(timezone.utc).isoformat(),
|
||||
"timestamp": datetime.utcnow().isoformat(),
|
||||
"checks": checks,
|
||||
}
|
||||
if not all_ready:
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,65 +0,0 @@
|
||||
"""模板编辑器 API 路由包.
|
||||
|
||||
将原来 2560 行的 templates_editor.py 巨无霸拆分为 12 个模块:
|
||||
- schemas.py: 所有 Pydantic model
|
||||
- dependencies.py: 依赖注入
|
||||
- _utils.py: 工具函数
|
||||
- _fallback.py: 自动兜底逻辑
|
||||
- draft.py: 草稿管理(详情/更新/发布/版本/回滚)
|
||||
- clips.py: 片段管理(CRUD/分割/合并/重排/批量删除/从素材创建)
|
||||
- adjustments.py: 片段调整(速度/音量/裁剪/批量调速)
|
||||
- bgm.py: BGM 管理
|
||||
- effects.py: 转场 + 滤镜
|
||||
- export.py: 导出配置
|
||||
- cover.py: 封面管理 + AI 生成封面
|
||||
- subtitles.py: 字幕管理
|
||||
- ai_features.py: AI 推荐
|
||||
- generation.py: 生成(触发/进度/记录)
|
||||
- timeline.py: 时间线
|
||||
|
||||
挂载路径: /api/v1/templates/{template_id}/editor/
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
# 向后兼容:测试和其他模块可能直接从 templates_editor 导入这些符号
|
||||
from app.auth import get_current_user # noqa: F401
|
||||
from app.dependencies import get_db_session # noqa: F401
|
||||
from fastapi import APIRouter
|
||||
|
||||
from .adjustments import router as adjustments_router
|
||||
from .ai_features import router as ai_features_router
|
||||
from .bgm import router as bgm_router
|
||||
from .clips import router as clips_router
|
||||
from .cover import router as cover_router
|
||||
from .dependencies import get_draft_plan_id, get_editor_services # noqa: F401
|
||||
from .draft import router as draft_router
|
||||
from .effects import router as effects_router
|
||||
from .export import router as export_router
|
||||
from .generation import router as generation_router
|
||||
from .subtitles import router as subtitles_router
|
||||
from .timeline import router as timeline_router
|
||||
|
||||
# 主 router,所有子路由都合并到这里
|
||||
router = APIRouter(tags=["Template Editor"])
|
||||
|
||||
# 合并所有子模块的路由(不用 include_router 是因为子路由有空路径 "")
|
||||
_sub_routers = [
|
||||
draft_router,
|
||||
clips_router,
|
||||
adjustments_router,
|
||||
bgm_router,
|
||||
effects_router,
|
||||
export_router,
|
||||
cover_router,
|
||||
subtitles_router,
|
||||
ai_features_router,
|
||||
generation_router,
|
||||
timeline_router,
|
||||
]
|
||||
|
||||
for sub in _sub_routers:
|
||||
for route in sub.routes:
|
||||
router.routes.append(route)
|
||||
|
||||
__all__ = ["router"]
|
||||
@@ -1,163 +0,0 @@
|
||||
"""模板编辑器自动兜底逻辑.
|
||||
|
||||
generate_editor_draft 触发生成前的自动修复流程:
|
||||
1. draft → editing 状态迁移
|
||||
2. 无片段时从模板复制片段配置
|
||||
3. 为无素材片段分配指定素材
|
||||
4. 项目有素材库时自动选素材
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import random
|
||||
from typing import Any
|
||||
|
||||
from app.services.edit_plan_service import EditPlanService
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.template_clip_config_repository import (
|
||||
SQLAlchemyTemplateClipConfigRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.template_repository import (
|
||||
SQLAlchemyTemplateRepository,
|
||||
)
|
||||
from packages.domain.edit_plan import EditPlanStatus
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _auto_fallback_draft_to_editing(
|
||||
svc: EditPlanService, plan_id: str, plan_check
|
||||
) -> None:
|
||||
"""自动兜底 1: draft → editing"""
|
||||
if plan_check.status == EditPlanStatus.DRAFT:
|
||||
logger.info("模板编辑器自动兜底: plan=%s draft→editing", plan_id)
|
||||
svc.transition_status(plan_id, EditPlanStatus.EDITING)
|
||||
|
||||
|
||||
def _auto_fallback_copy_template_clips(
|
||||
svc: EditPlanService, plan_id: str, plan_check, db: Session
|
||||
) -> None:
|
||||
"""自动兜底 2: 无片段 + 有 template_id → 从模板复制片段配置"""
|
||||
existing_clips = svc.count_clips(plan_id)
|
||||
if existing_clips == 0 and plan_check.template_id:
|
||||
logger.info(
|
||||
"模板编辑器自动兜底: plan=%s 无片段,从模板 %s 复制片段配置",
|
||||
plan_id,
|
||||
plan_check.template_id,
|
||||
)
|
||||
clip_config_repo = SQLAlchemyTemplateClipConfigRepository(db)
|
||||
configs = clip_config_repo.list_by_template(plan_check.template_id)
|
||||
if configs:
|
||||
for cfg in configs:
|
||||
svc.create_clip(
|
||||
plan_id=plan_id,
|
||||
clip_type=cfg.clip_type.value
|
||||
if hasattr(cfg.clip_type, "value")
|
||||
else cfg.clip_type,
|
||||
order=cfg.order,
|
||||
template_clip_config_id=cfg.id,
|
||||
duration=cfg.default_duration,
|
||||
transition_effect=cfg.transition_effect.value
|
||||
if hasattr(cfg.transition_effect, "value")
|
||||
else cfg.transition_effect,
|
||||
)
|
||||
logger.info(
|
||||
"模板编辑器自动兜底: plan=%s 从 template_clip_configs 复制了 %d 个片段",
|
||||
plan_id,
|
||||
len(configs),
|
||||
)
|
||||
else:
|
||||
tpl_repo = SQLAlchemyTemplateRepository(db)
|
||||
segments = tpl_repo.list_segments(plan_check.template_id)
|
||||
for seg in segments:
|
||||
avg_duration = (seg.duration_min + seg.duration_max) / 2
|
||||
svc.create_clip(
|
||||
plan_id=plan_id,
|
||||
clip_type="main",
|
||||
order=seg.segment_order,
|
||||
duration=avg_duration,
|
||||
config={
|
||||
"material_type": seg.material_type or "",
|
||||
"template_segment_id": seg.id,
|
||||
},
|
||||
)
|
||||
logger.info(
|
||||
"模板编辑器自动兜底: plan=%s 从旧模板 segments 复制了 %d 个片段",
|
||||
plan_id,
|
||||
len(segments),
|
||||
)
|
||||
|
||||
|
||||
def _auto_fallback_assign_assets(
|
||||
svc: EditPlanService, plan_id: str, plan_check
|
||||
) -> list:
|
||||
"""自动兜底 3: 为没有素材的片段分配素材。返回剩余无素材片段列表。"""
|
||||
all_clips = svc.list_clips(plan_id)
|
||||
clips_without_asset = [c for c in all_clips if not c.asset_id]
|
||||
config_asset_ids = (plan_check.config or {}).get("asset_ids", [])
|
||||
|
||||
if clips_without_asset and config_asset_ids:
|
||||
logger.info(
|
||||
"模板编辑器自动兜底3: plan=%s 为 %d 个无素材片段分配 %d 个指定素材",
|
||||
plan_id,
|
||||
len(clips_without_asset),
|
||||
len(config_asset_ids),
|
||||
)
|
||||
for i, clip in enumerate(clips_without_asset):
|
||||
asset_idx = i % len(config_asset_ids)
|
||||
svc.assign_asset(clip.id, config_asset_ids[asset_idx])
|
||||
logger.info("模板编辑器自动兜底3: plan=%s 素材分配完成", plan_id)
|
||||
clips_without_asset = []
|
||||
|
||||
return clips_without_asset
|
||||
|
||||
|
||||
def _auto_fallback_auto_material_mode(
|
||||
svc: EditPlanService,
|
||||
plan_id: str,
|
||||
plan_check,
|
||||
clips_without_asset: list,
|
||||
asset_library_repo: Any,
|
||||
asset_repo: Any,
|
||||
) -> None:
|
||||
"""自动兜底 4: 项目有视频素材库时自动选素材"""
|
||||
if not clips_without_asset:
|
||||
return
|
||||
if not plan_check.project_id:
|
||||
return
|
||||
|
||||
logger.info(
|
||||
"模板编辑器自动兜底4: plan=%s 自动选素材分配给 %d 个无素材片段",
|
||||
plan_id,
|
||||
len(clips_without_asset),
|
||||
)
|
||||
libs = asset_library_repo.find_by_project(plan_check.project_id)
|
||||
video_lib = None
|
||||
for lib in libs:
|
||||
lib_kind = lib.kind.value if hasattr(lib.kind, "value") else lib.kind
|
||||
if lib_kind == "video":
|
||||
video_lib = lib
|
||||
break
|
||||
|
||||
if video_lib:
|
||||
assets = asset_repo.find_by_library(video_lib.id)
|
||||
ready_videos = [
|
||||
a
|
||||
for a in assets
|
||||
if (a.status.value if hasattr(a.status, "value") else a.status) == "ready"
|
||||
and a.mime_type
|
||||
and a.mime_type.startswith("video")
|
||||
]
|
||||
if ready_videos:
|
||||
random.shuffle(ready_videos)
|
||||
for i, clip in enumerate(clips_without_asset):
|
||||
asset = ready_videos[i % len(ready_videos)]
|
||||
svc.assign_asset(clip.id, asset.id)
|
||||
logger.info(
|
||||
"模板编辑器自动兜底4: plan=%s 从素材库 %s 分配了 %d 个素材",
|
||||
plan_id,
|
||||
video_lib.name,
|
||||
len(ready_videos),
|
||||
)
|
||||
@@ -1,109 +0,0 @@
|
||||
"""模板编辑器内部工具函数.
|
||||
|
||||
纯函数,不依赖请求上下文。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from .schemas import ClipAdjustResponse
|
||||
|
||||
# 时间线场景颜色映射
|
||||
_CLIP_TYPE_COLORS = {
|
||||
"intro": "#6366f1",
|
||||
"title": "#6366f1",
|
||||
"product": "#818cf8",
|
||||
"showcase": "#10b981",
|
||||
"scene": "#10b981",
|
||||
"subtitle": "#f59e0b",
|
||||
"text": "#f59e0b",
|
||||
"cta": "#ef4444",
|
||||
"outro": "#ef4444",
|
||||
"voiceover": "#8b5cf6",
|
||||
"transition": "#64748b",
|
||||
}
|
||||
_DEFAULT_COLOR = "#6366f1"
|
||||
|
||||
|
||||
def _format_time(seconds: float) -> str:
|
||||
"""秒数格式化为 m:ss"""
|
||||
m = int(seconds) // 60
|
||||
s = int(seconds) % 60
|
||||
return f"{m}:{s:02d}"
|
||||
|
||||
|
||||
def _clip_type_to_scene_label(clip_type: str, text_content: str) -> str:
|
||||
"""片段类型转时间线场景标签"""
|
||||
type_labels = {
|
||||
"intro": "开场",
|
||||
"title": "标题",
|
||||
"product": "产品展示",
|
||||
"showcase": "场景展示",
|
||||
"scene": "场景",
|
||||
"subtitle": "字幕",
|
||||
"text": "文字",
|
||||
"cta": "结尾 CTA",
|
||||
"outro": "结尾",
|
||||
"voiceover": "配音",
|
||||
"transition": "转场",
|
||||
}
|
||||
label = type_labels.get(clip_type, clip_type or "片段")
|
||||
if text_content:
|
||||
short = text_content[:20].strip()
|
||||
if short:
|
||||
return f"{label} - {short}"
|
||||
return label
|
||||
|
||||
|
||||
# ── 片段调整相关工具 ────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
def _get_clip_config(clip) -> dict:
|
||||
"""安全获取 clip.config"""
|
||||
config = getattr(clip, "config", {}) or {}
|
||||
if not isinstance(config, dict):
|
||||
config = {}
|
||||
return config
|
||||
|
||||
|
||||
def _get_adjust_volume(clip) -> float:
|
||||
"""获取片段音量"""
|
||||
config = _get_clip_config(clip)
|
||||
return float(config.get("volume", 1.0))
|
||||
|
||||
|
||||
def _get_adjust_trim(clip) -> tuple[float, float]:
|
||||
"""获取片段裁剪起止"""
|
||||
config = _get_clip_config(clip)
|
||||
trim_start = float(config.get("trim_start", 0.0))
|
||||
trim_end = float(config.get("trim_end", 0.0))
|
||||
return trim_start, trim_end
|
||||
|
||||
|
||||
def _build_adjust_response(clip) -> ClipAdjustResponse:
|
||||
"""构造片段调整响应"""
|
||||
trim_start, trim_end = _get_adjust_trim(clip)
|
||||
return ClipAdjustResponse(
|
||||
clip_id=clip.id,
|
||||
speed=clip.playback_speed,
|
||||
volume=_get_adjust_volume(clip),
|
||||
trim_start=trim_start,
|
||||
trim_end=trim_end,
|
||||
duration=clip.duration,
|
||||
)
|
||||
|
||||
|
||||
def _validate_trim(trim_start: float, trim_end: float, total_duration: float) -> None:
|
||||
"""校验裁剪时长合法性"""
|
||||
if trim_start + trim_end >= total_duration:
|
||||
raise ValueError(
|
||||
f"裁剪总时长({trim_start + trim_end:.2f}s)不能大于等于片段总时长({total_duration:.2f}s)"
|
||||
)
|
||||
|
||||
|
||||
def _clip_value(value: Any) -> str:
|
||||
"""获取枚举/字符串值的统一方法"""
|
||||
if hasattr(value, "value"):
|
||||
return value.value
|
||||
return str(value)
|
||||
@@ -1,167 +0,0 @@
|
||||
"""片段调整路由.
|
||||
|
||||
端点:
|
||||
- PUT /clips/{clip_id}/speed 调速
|
||||
- PUT /clips/{clip_id}/volume 调音量
|
||||
- PUT /clips/{clip_id}/trim 裁剪
|
||||
- PUT /clips/{clip_id}/adjustments 统一调整
|
||||
- POST /clips/batch-speed 批量调速
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.services.edit_plan_service import EditPlanService
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
|
||||
from ._utils import _build_adjust_response, _get_adjust_trim, _get_clip_config, _validate_trim
|
||||
from .dependencies import get_draft_plan_id, get_editor_services
|
||||
from .schemas import (
|
||||
BatchSpeedRequest,
|
||||
BatchSpeedResponse,
|
||||
ClipAdjustmentsRequest,
|
||||
ClipAdjustResponse,
|
||||
SpeedAdjustRequest,
|
||||
TrimAdjustRequest,
|
||||
VolumeAdjustRequest,
|
||||
)
|
||||
|
||||
router = APIRouter(tags=["Template Editor"])
|
||||
|
||||
|
||||
@router.put("/clips/{clip_id}/speed", response_model=ClipAdjustResponse)
|
||||
def adjust_editor_clip_speed(
|
||||
template_id: str,
|
||||
clip_id: str,
|
||||
body: SpeedAdjustRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> ClipAdjustResponse:
|
||||
"""调整片段播放速度"""
|
||||
_, plan_svc = services
|
||||
clip = plan_svc.get_clip(clip_id)
|
||||
if not clip:
|
||||
raise HTTPException(status_code=404, detail="片段不存在")
|
||||
|
||||
updated = plan_svc.update_clip(clip_id, playback_speed=body.speed)
|
||||
return _build_adjust_response(updated)
|
||||
|
||||
|
||||
@router.put("/clips/{clip_id}/volume", response_model=ClipAdjustResponse)
|
||||
def adjust_editor_clip_volume(
|
||||
template_id: str,
|
||||
clip_id: str,
|
||||
body: VolumeAdjustRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> ClipAdjustResponse:
|
||||
"""调整片段音量"""
|
||||
_, plan_svc = services
|
||||
clip = plan_svc.get_clip(clip_id)
|
||||
if not clip:
|
||||
raise HTTPException(status_code=404, detail="片段不存在")
|
||||
|
||||
config = dict(_get_clip_config(clip))
|
||||
config["volume"] = body.volume
|
||||
updated = plan_svc.update_clip(clip_id, config=config)
|
||||
return _build_adjust_response(updated)
|
||||
|
||||
|
||||
@router.put("/clips/{clip_id}/trim", response_model=ClipAdjustResponse)
|
||||
def adjust_editor_clip_trim(
|
||||
template_id: str,
|
||||
clip_id: str,
|
||||
body: TrimAdjustRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> ClipAdjustResponse:
|
||||
"""裁剪片段(trim in/out)"""
|
||||
_, plan_svc = services
|
||||
clip = plan_svc.get_clip(clip_id)
|
||||
if not clip:
|
||||
raise HTTPException(status_code=404, detail="片段不存在")
|
||||
|
||||
try:
|
||||
_validate_trim(body.trim_start, body.trim_end, clip.duration)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e)) from e
|
||||
|
||||
config = dict(_get_clip_config(clip))
|
||||
config["trim_start"] = body.trim_start
|
||||
config["trim_end"] = body.trim_end
|
||||
updated = plan_svc.update_clip(clip_id, config=config)
|
||||
return _build_adjust_response(updated)
|
||||
|
||||
|
||||
@router.put("/clips/{clip_id}/adjustments", response_model=ClipAdjustResponse)
|
||||
def adjust_editor_clip_all(
|
||||
template_id: str,
|
||||
clip_id: str,
|
||||
body: ClipAdjustmentsRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> ClipAdjustResponse:
|
||||
"""统一调整片段的 speed / volume / trim"""
|
||||
_, plan_svc = services
|
||||
clip = plan_svc.get_clip(clip_id)
|
||||
if not clip:
|
||||
raise HTTPException(status_code=404, detail="片段不存在")
|
||||
|
||||
update_kwargs: dict[str, Any] = {}
|
||||
config_updates: dict[str, Any] = {}
|
||||
|
||||
if body.speed is not None:
|
||||
update_kwargs["playback_speed"] = body.speed
|
||||
if body.volume is not None:
|
||||
config_updates["volume"] = body.volume
|
||||
if body.trim_start is not None:
|
||||
config_updates["trim_start"] = body.trim_start
|
||||
if body.trim_end is not None:
|
||||
config_updates["trim_end"] = body.trim_end
|
||||
|
||||
current_trim_start, current_trim_end = _get_adjust_trim(clip)
|
||||
new_trim_start = body.trim_start if body.trim_start is not None else current_trim_start
|
||||
new_trim_end = body.trim_end if body.trim_end is not None else current_trim_end
|
||||
|
||||
if body.trim_start is not None or body.trim_end is not None:
|
||||
try:
|
||||
_validate_trim(new_trim_start, new_trim_end, clip.duration)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e)) from e
|
||||
|
||||
if config_updates:
|
||||
config = dict(_get_clip_config(clip))
|
||||
config.update(config_updates)
|
||||
update_kwargs["config"] = config
|
||||
|
||||
if not update_kwargs:
|
||||
return _build_adjust_response(clip)
|
||||
|
||||
updated = plan_svc.update_clip(clip_id, **update_kwargs)
|
||||
return _build_adjust_response(updated)
|
||||
|
||||
|
||||
@router.post("/clips/batch-speed", response_model=BatchSpeedResponse)
|
||||
def batch_adjust_editor_speed(
|
||||
template_id: str,
|
||||
body: BatchSpeedRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> BatchSpeedResponse:
|
||||
"""批量调整草稿内所有片段的播放速度"""
|
||||
_, plan_svc = services
|
||||
clips = plan_svc.list_clips(plan_id, limit=500, skip=0)
|
||||
count = 0
|
||||
for clip in clips:
|
||||
plan_svc.update_clip(clip.id, playback_speed=body.speed)
|
||||
count += 1
|
||||
|
||||
return BatchSpeedResponse(updated_count=count, plan_id=plan_id)
|
||||
@@ -1,121 +0,0 @@
|
||||
"""AI 功能路由.
|
||||
|
||||
端点:
|
||||
- POST /ai-recommend AI 推荐片段方案
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import get_db_session
|
||||
from app.services.edit_plan_service import EditPlanService
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.domain.config_schemas import normalize_plan_config
|
||||
|
||||
from .dependencies import get_draft_plan_id, get_editor_services
|
||||
from .schemas import AIRecommendRequest, AIRecommendResponse
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(tags=["Template Editor"])
|
||||
|
||||
|
||||
@router.post("/ai-recommend", response_model=AIRecommendResponse)
|
||||
def editor_ai_recommend(
|
||||
template_id: str,
|
||||
body: AIRecommendRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
db: Session = Depends(get_db_session),
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> AIRecommendResponse:
|
||||
"""AI 推荐片段方案"""
|
||||
_, plan_svc = services
|
||||
plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
|
||||
plan_status = plan.status.value if hasattr(plan.status, "value") else plan.status
|
||||
if plan_status not in ("draft", "editing"):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="当前草稿状态不支持AI推荐,请先编辑后再试",
|
||||
)
|
||||
|
||||
from packages.shared.ai_service import run_ai_recommend
|
||||
|
||||
result = run_ai_recommend(
|
||||
plan_id=plan_id,
|
||||
template_id=plan.template_id,
|
||||
asset_ids=body.asset_ids,
|
||||
editing_mode=body.editing_mode,
|
||||
target_duration=body.target_duration,
|
||||
)
|
||||
|
||||
try:
|
||||
plan_svc.delete_all_clips(plan_id)
|
||||
|
||||
for clip_data in result["clips"]:
|
||||
plan_svc.create_clip(
|
||||
plan_id=plan_id,
|
||||
clip_type=clip_data["clip_type"],
|
||||
order=clip_data["order"],
|
||||
text_content=clip_data.get("text_content", ""),
|
||||
duration=clip_data["duration"],
|
||||
transition_effect=clip_data.get("transition_effect", "cut"),
|
||||
asset_id=clip_data.get("asset_id", ""),
|
||||
start_time=clip_data.get("start_time", 0.0),
|
||||
config=clip_data.get("config", {}),
|
||||
)
|
||||
|
||||
normalized_config = normalize_plan_config(result.get("config", {}))
|
||||
plan_svc.update_plan(
|
||||
plan_id,
|
||||
config=normalized_config,
|
||||
total_duration=result["total_duration"],
|
||||
)
|
||||
except Exception as _e:
|
||||
logger.exception(
|
||||
"模板编辑器AI推荐写入失败: template_id=%s plan_id=%s",
|
||||
template_id,
|
||||
plan_id,
|
||||
)
|
||||
try:
|
||||
db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="AI推荐结果保存失败,请稍后重试",
|
||||
) from _e
|
||||
|
||||
logger.info(
|
||||
"模板编辑器AI推荐: template_id=%s plan_id=%s clips=%d duration=%.1f by user=%s",
|
||||
template_id,
|
||||
plan_id,
|
||||
len(result["clips"]),
|
||||
result["total_duration"],
|
||||
current_user.user.id,
|
||||
)
|
||||
|
||||
return AIRecommendResponse(
|
||||
plan_id=plan_id,
|
||||
clips=[
|
||||
{
|
||||
"clip_type": c["clip_type"],
|
||||
"order": c["order"],
|
||||
"text_content": c.get("text_content", ""),
|
||||
"duration": c["duration"],
|
||||
"transition_effect": c.get("transition_effect", "cut"),
|
||||
"asset_id": c.get("asset_id", ""),
|
||||
"start_time": c.get("start_time", 0.0),
|
||||
"config": c.get("config", {}),
|
||||
}
|
||||
for c in result["clips"]
|
||||
],
|
||||
config=normalized_config,
|
||||
total_duration=result["total_duration"],
|
||||
confidence=result["confidence"],
|
||||
)
|
||||
@@ -1,133 +0,0 @@
|
||||
"""BGM 管理路由.
|
||||
|
||||
端点:
|
||||
- GET /bgm 获取 BGM 配置
|
||||
- PUT /bgm 更新 BGM 配置
|
||||
- GET /bgm/presets 预设 BGM 列表
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.services.edit_plan_service import EditPlanService
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
|
||||
from .dependencies import get_draft_plan_id, get_editor_services
|
||||
from .schemas import BGMConfigUpdateRequest
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(tags=["Template Editor"])
|
||||
|
||||
|
||||
@router.get("/bgm", response_model=dict[str, Any])
|
||||
def get_editor_bgm(
|
||||
template_id: str,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""获取草稿的 BGM 配置"""
|
||||
_, plan_svc = services
|
||||
plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
config = plan.config or {}
|
||||
return {
|
||||
"plan_id": plan.id,
|
||||
"bgm": config.get("bgm", {}),
|
||||
}
|
||||
|
||||
|
||||
@router.put("/bgm", response_model=dict[str, Any])
|
||||
def update_editor_bgm(
|
||||
template_id: str,
|
||||
body: BGMConfigUpdateRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""更新草稿的 BGM 配置"""
|
||||
_, plan_svc = services
|
||||
plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
|
||||
config = dict(plan.config) if plan.config else {}
|
||||
current_bgm = dict(config.get("bgm", {}))
|
||||
update_data = body.model_dump(exclude_none=True)
|
||||
current_bgm.update(update_data)
|
||||
|
||||
if current_bgm.get("enabled"):
|
||||
has_source = any(
|
||||
current_bgm.get(key)
|
||||
for key in ("asset_id", "preset_id", "audio_url")
|
||||
if current_bgm.get(key)
|
||||
)
|
||||
if not has_source:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="启用 BGM 时需要指定素材来源(asset_id / preset_id / audio_url)",
|
||||
)
|
||||
|
||||
config["bgm"] = current_bgm
|
||||
updated_plan = plan_svc.update_plan_config(plan_id, config)
|
||||
|
||||
logger.info(
|
||||
"模板编辑器更新BGM: template_id=%s plan_id=%s enabled=%s by user=%s",
|
||||
template_id,
|
||||
plan_id,
|
||||
current_bgm.get("enabled", False),
|
||||
current_user.user.id,
|
||||
)
|
||||
|
||||
return {
|
||||
"plan_id": updated_plan.id,
|
||||
"bgm": current_bgm,
|
||||
}
|
||||
|
||||
|
||||
@router.get("/bgm/presets", response_model=dict[str, Any])
|
||||
def list_editor_bgm_presets(
|
||||
style: str | None = Query(default=None, description="按风格筛选"),
|
||||
keyword: str | None = Query(default=None, description="关键词搜索"),
|
||||
skip: int = Query(default=0, ge=0, description="分页偏移"),
|
||||
limit: int = Query(default=50, ge=1, le=200, description="每页数量"),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""获取预设 BGM 列表"""
|
||||
from packages.domain.preset_bgm import (
|
||||
BGM_STYLES,
|
||||
PRESET_BGM_LIBRARY,
|
||||
list_preset_bgm_by_style,
|
||||
search_preset_bgm,
|
||||
)
|
||||
|
||||
bgm_list = PRESET_BGM_LIBRARY
|
||||
if keyword:
|
||||
bgm_list = search_preset_bgm(keyword)
|
||||
elif style:
|
||||
bgm_list = list_preset_bgm_by_style(style)
|
||||
|
||||
total = len(bgm_list)
|
||||
paged = bgm_list[skip : skip + limit]
|
||||
|
||||
return {
|
||||
"total": total,
|
||||
"skip": skip,
|
||||
"limit": limit,
|
||||
"styles": BGM_STYLES,
|
||||
"items": [
|
||||
{
|
||||
"id": bgm.id,
|
||||
"name": bgm.name,
|
||||
"style": bgm.style,
|
||||
"style_label": BGM_STYLES.get(bgm.style, bgm.style),
|
||||
"duration": bgm.duration,
|
||||
"artist": bgm.artist,
|
||||
"description": bgm.description,
|
||||
"tags": bgm.tags,
|
||||
"audio_url": bgm.audio_url,
|
||||
}
|
||||
for bgm in paged
|
||||
],
|
||||
}
|
||||
@@ -1,318 +0,0 @@
|
||||
"""片段管理路由.
|
||||
|
||||
端点:
|
||||
- GET /clips 片段列表
|
||||
- POST /clips 创建片段
|
||||
- GET /clips/{clip_id} 片段详情
|
||||
- PUT /clips/{clip_id} 更新片段
|
||||
- DELETE /clips/{clip_id} 删除片段
|
||||
- POST /clips/{clip_id}/split 分割片段
|
||||
- POST /clips/merge 合并片段
|
||||
- POST /clips/reorder 重排片段
|
||||
- POST /clips/batch-delete 批量删除
|
||||
- POST /clips/from-assets 从素材创建片段
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.services.edit_plan_service import EditPlanService
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
|
||||
from .dependencies import get_draft_plan_id, get_editor_services
|
||||
from .schemas import (
|
||||
ClipBatchDeleteRequest,
|
||||
ClipBatchDeleteResponse,
|
||||
ClipReorderRequest,
|
||||
ClipReorderResponse,
|
||||
ClipsFromAssetsRequest,
|
||||
ClipsFromAssetsResponse,
|
||||
EditorClipCreateRequest,
|
||||
EditorClipListResponse,
|
||||
EditorClipResponse,
|
||||
EditorClipUpdateRequest,
|
||||
MergeClipsRequest,
|
||||
SplitClipRequest,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(tags=["Template Editor"])
|
||||
|
||||
|
||||
def _clip_to_response(clip) -> EditorClipResponse:
|
||||
"""统一构造片段响应"""
|
||||
return EditorClipResponse(
|
||||
id=clip.id,
|
||||
plan_id=clip.plan_id,
|
||||
clip_type=clip.clip_type.value
|
||||
if hasattr(clip.clip_type, "value")
|
||||
else str(clip.clip_type),
|
||||
order=clip.order,
|
||||
duration=clip.duration,
|
||||
text_content=clip.text_content or "",
|
||||
transition_effect=clip.transition_effect.value
|
||||
if hasattr(clip.transition_effect, "value")
|
||||
else str(clip.transition_effect),
|
||||
playback_speed=clip.playback_speed or 1.0,
|
||||
config=clip.config or {},
|
||||
)
|
||||
|
||||
|
||||
@router.get("/clips", response_model=EditorClipListResponse)
|
||||
def list_draft_clips(
|
||||
template_id: str,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
skip: int = Query(default=0, ge=0),
|
||||
limit: int = Query(default=100, ge=1, le=500),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""获取草稿的片段列表"""
|
||||
_, plan_svc = services
|
||||
clips = plan_svc.list_clips(plan_id, skip=skip, limit=limit)
|
||||
total = plan_svc.count_clips(plan_id)
|
||||
return EditorClipListResponse(
|
||||
items=[_clip_to_response(c) for c in clips],
|
||||
total=total,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/clips", response_model=EditorClipResponse, status_code=status.HTTP_201_CREATED)
|
||||
def create_draft_clip(
|
||||
template_id: str,
|
||||
req: EditorClipCreateRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""在草稿中创建新片段"""
|
||||
_, plan_svc = services
|
||||
try:
|
||||
clip = plan_svc.create_clip(
|
||||
plan_id,
|
||||
clip_type=req.clip_type,
|
||||
order=req.order,
|
||||
duration=req.duration,
|
||||
text_content=req.text_content,
|
||||
transition_effect=req.transition_effect,
|
||||
config=req.config,
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
return _clip_to_response(clip)
|
||||
|
||||
|
||||
@router.put("/clips/{clip_id}", response_model=EditorClipResponse)
|
||||
def update_draft_clip(
|
||||
template_id: str,
|
||||
clip_id: str,
|
||||
req: EditorClipUpdateRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""更新草稿中的片段"""
|
||||
_, plan_svc = services
|
||||
try:
|
||||
clip = plan_svc.update_clip(
|
||||
clip_id,
|
||||
order=req.order,
|
||||
duration=req.duration,
|
||||
text_content=req.text_content,
|
||||
transition_effect=req.transition_effect,
|
||||
playback_speed=req.playback_speed,
|
||||
config=req.config,
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
return _clip_to_response(clip)
|
||||
|
||||
|
||||
@router.delete("/clips/{clip_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
def delete_draft_clip(
|
||||
template_id: str,
|
||||
clip_id: str,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""删除草稿中的片段"""
|
||||
_, plan_svc = services
|
||||
success = plan_svc.delete_clip(clip_id)
|
||||
if not success:
|
||||
raise HTTPException(status_code=404, detail="片段不存在")
|
||||
return None
|
||||
|
||||
|
||||
@router.get("/clips/{clip_id}", response_model=EditorClipResponse)
|
||||
def get_draft_clip_detail(
|
||||
template_id: str,
|
||||
clip_id: str,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""获取草稿中的片段详情"""
|
||||
_, plan_svc = services
|
||||
clip = plan_svc.get_clip(clip_id)
|
||||
if clip is None:
|
||||
raise HTTPException(status_code=404, detail="片段不存在")
|
||||
if clip.plan_id != plan_id:
|
||||
raise HTTPException(status_code=404, detail="片段不存在")
|
||||
return _clip_to_response(clip)
|
||||
|
||||
|
||||
@router.post("/clips/{clip_id}/split", response_model=dict[str, Any], status_code=status.HTTP_200_OK)
|
||||
def split_draft_clip(
|
||||
template_id: str,
|
||||
clip_id: str,
|
||||
body: SplitClipRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""将一个片段从指定时间点分割为两个片段"""
|
||||
_, plan_svc = services
|
||||
clip = plan_svc.get_clip(clip_id)
|
||||
if clip is None or clip.plan_id != plan_id:
|
||||
raise HTTPException(status_code=404, detail="片段不存在")
|
||||
try:
|
||||
result = plan_svc.split_clip(clip_id, body.split_time)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)
|
||||
) from exc
|
||||
left = result["left_clip"]
|
||||
right = result["right_clip"]
|
||||
return {
|
||||
"left_clip": {
|
||||
"id": left.id,
|
||||
"plan_id": left.plan_id,
|
||||
"clip_type": left.clip_type,
|
||||
"order": left.order,
|
||||
"duration": left.duration,
|
||||
"start_time": left.start_time,
|
||||
},
|
||||
"right_clip": {
|
||||
"id": right.id,
|
||||
"plan_id": right.plan_id,
|
||||
"clip_type": right.clip_type,
|
||||
"order": right.order,
|
||||
"duration": right.duration,
|
||||
"start_time": right.start_time,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
@router.post("/clips/merge", response_model=dict[str, Any], status_code=status.HTTP_200_OK)
|
||||
def merge_draft_clips(
|
||||
template_id: str,
|
||||
body: MergeClipsRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""将多个连续的同类型片段合并为一个片段"""
|
||||
_, plan_svc = services
|
||||
for cid in body.clip_ids:
|
||||
clip = plan_svc.get_clip(cid)
|
||||
if clip is None or clip.plan_id != plan_id:
|
||||
raise HTTPException(status_code=404, detail=f"片段不存在: {cid}")
|
||||
try:
|
||||
merged = plan_svc.merge_clips(body.clip_ids)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)
|
||||
) from exc
|
||||
return {
|
||||
"id": merged.id,
|
||||
"plan_id": merged.plan_id,
|
||||
"clip_type": merged.clip_type,
|
||||
"order": merged.order,
|
||||
"duration": merged.duration,
|
||||
"text_content": merged.text_content,
|
||||
}
|
||||
|
||||
|
||||
@router.post("/clips/reorder", response_model=ClipReorderResponse)
|
||||
def reorder_editor_clips(
|
||||
template_id: str,
|
||||
body: ClipReorderRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> ClipReorderResponse:
|
||||
"""批量重排片段顺序"""
|
||||
_, plan_svc = services
|
||||
count = 0
|
||||
for item in body.items:
|
||||
try:
|
||||
plan_svc.update_clip(item.clip_id, order=item.new_order)
|
||||
count += 1
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
return ClipReorderResponse(updated_count=count, plan_id=plan_id)
|
||||
|
||||
|
||||
@router.post("/clips/batch-delete", response_model=ClipBatchDeleteResponse)
|
||||
def batch_delete_editor_clips(
|
||||
template_id: str,
|
||||
body: ClipBatchDeleteRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> ClipBatchDeleteResponse:
|
||||
"""批量删除片段"""
|
||||
_, plan_svc = services
|
||||
deleted = 0
|
||||
for clip_id in body.clip_ids:
|
||||
if plan_svc.delete_clip(clip_id):
|
||||
deleted += 1
|
||||
|
||||
return ClipBatchDeleteResponse(deleted_count=deleted, plan_id=plan_id)
|
||||
|
||||
|
||||
@router.post("/clips/from-assets", response_model=ClipsFromAssetsResponse)
|
||||
def create_clips_from_assets_editor(
|
||||
template_id: str,
|
||||
body: ClipsFromAssetsRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> ClipsFromAssetsResponse:
|
||||
"""从素材批量创建片段"""
|
||||
_, plan_svc = services
|
||||
clips = []
|
||||
for i, asset_id in enumerate(body.asset_ids):
|
||||
try:
|
||||
clip = plan_svc.create_clip(
|
||||
plan_id,
|
||||
clip_type="main",
|
||||
order=body.start_order + i if hasattr(body, "start_order") else i,
|
||||
duration=5.0,
|
||||
asset_id=asset_id,
|
||||
)
|
||||
clips.append(clip)
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
logger.info(
|
||||
"模板编辑器从素材创建片段: template_id=%s plan_id=%s count=%d by user=%s",
|
||||
template_id,
|
||||
plan_id,
|
||||
len(clips),
|
||||
current_user.user.id,
|
||||
)
|
||||
|
||||
return ClipsFromAssetsResponse(
|
||||
created_count=len(clips),
|
||||
plan_id=plan_id,
|
||||
clip_ids=[c.id for c in clips],
|
||||
)
|
||||
@@ -1,209 +0,0 @@
|
||||
"""封面管理路由.
|
||||
|
||||
端点:
|
||||
- GET /cover 封面配置
|
||||
- PUT /cover 更新封面
|
||||
- POST /cover/extract 抽帧生成封面
|
||||
- POST /cover/smart 智能选帧
|
||||
- POST /generate-cover AI 生成封面
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.services.edit_plan_service import EditPlanService
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from packages.domain.config_schemas import normalize_plan_config
|
||||
|
||||
from .dependencies import get_draft_plan_id, get_editor_services
|
||||
from .schemas import (
|
||||
CoverConfigResponse,
|
||||
CoverExtractRequest,
|
||||
CoverGenerateResponse,
|
||||
CoverSmartRequest,
|
||||
CoverUpdateRequest,
|
||||
GenerateCoverRequest,
|
||||
GenerateCoverResponse,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(tags=["Template Editor"])
|
||||
|
||||
|
||||
@router.get("/cover", response_model=CoverConfigResponse)
|
||||
def get_editor_cover(
|
||||
template_id: str,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> CoverConfigResponse:
|
||||
"""获取草稿封面配置"""
|
||||
_, plan_svc = services
|
||||
plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
config = plan.config or {}
|
||||
cover_config = config.get("cover", {})
|
||||
|
||||
return CoverConfigResponse(
|
||||
type=cover_config.get("cover_type", "auto"),
|
||||
image_url=cover_config.get("cover_image_url", ""),
|
||||
frame_time=cover_config.get("frame_time", 0.0),
|
||||
)
|
||||
|
||||
|
||||
@router.put("/cover", response_model=CoverConfigResponse)
|
||||
def update_editor_cover(
|
||||
template_id: str,
|
||||
body: CoverUpdateRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> CoverConfigResponse:
|
||||
"""更新草稿封面配置"""
|
||||
_, plan_svc = services
|
||||
plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
|
||||
config = dict(plan.config) if plan.config else {}
|
||||
current_cover = dict(config.get("cover", {}))
|
||||
update_data = body.model_dump(exclude_none=True)
|
||||
current_cover.update(update_data)
|
||||
|
||||
config["cover"] = current_cover
|
||||
normalized = normalize_plan_config(config)
|
||||
plan_svc.update_plan_config(plan_id, {"cover": normalized["cover"]})
|
||||
|
||||
return CoverConfigResponse(
|
||||
type=current_cover.get("cover_type", "auto"),
|
||||
image_url=current_cover.get("cover_image_url", ""),
|
||||
frame_time=current_cover.get("frame_time", 0.0),
|
||||
)
|
||||
|
||||
|
||||
@router.post("/cover/extract", response_model=CoverGenerateResponse)
|
||||
def extract_editor_cover(
|
||||
template_id: str,
|
||||
body: CoverExtractRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> CoverGenerateResponse:
|
||||
"""从指定片段抽帧生成封面"""
|
||||
_, plan_svc = services
|
||||
plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
|
||||
clip = plan_svc.get_clip(body.clip_id)
|
||||
if not clip or clip.plan_id != plan_id:
|
||||
raise HTTPException(status_code=400, detail="片段不存在或不属于当前草稿")
|
||||
|
||||
cover_url = f"cover/extract/{plan_id}_{body.clip_id}_{body.frame_time}.jpg"
|
||||
|
||||
config = dict(plan.config) if plan.config else {}
|
||||
cover_config = dict(config.get("cover", {}))
|
||||
cover_config.update(
|
||||
{
|
||||
"cover_type": "extract",
|
||||
"cover_image_url": cover_url,
|
||||
"clip_id": body.clip_id,
|
||||
"frame_time": body.frame_time,
|
||||
}
|
||||
)
|
||||
config["cover"] = cover_config
|
||||
normalized = normalize_plan_config(config)
|
||||
plan_svc.update_plan_config(plan_id, {"cover": normalized["cover"]})
|
||||
|
||||
logger.info(
|
||||
"模板编辑器封面抽帧: template_id=%s plan_id=%s clip_id=%s by user=%s",
|
||||
template_id,
|
||||
plan_id,
|
||||
body.clip_id,
|
||||
current_user.user.id,
|
||||
)
|
||||
|
||||
return CoverGenerateResponse(
|
||||
type="extract",
|
||||
image_url=cover_url,
|
||||
frame_time=body.frame_time,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/cover/smart", response_model=CoverGenerateResponse)
|
||||
def smart_editor_cover(
|
||||
template_id: str,
|
||||
body: CoverSmartRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> CoverGenerateResponse:
|
||||
"""智能选帧生成封面"""
|
||||
_, plan_svc = services
|
||||
plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
|
||||
cover_url = f"cover/smart/{plan_id}_smart.jpg"
|
||||
strategy = getattr(body, "strategy", "auto")
|
||||
|
||||
config = dict(plan.config) if plan.config else {}
|
||||
cover_config = dict(config.get("cover", {}))
|
||||
cover_config.update(
|
||||
{
|
||||
"cover_type": "smart",
|
||||
"cover_image_url": cover_url,
|
||||
"strategy": strategy,
|
||||
}
|
||||
)
|
||||
config["cover"] = cover_config
|
||||
normalized = normalize_plan_config(config)
|
||||
plan_svc.update_plan_config(plan_id, {"cover": normalized["cover"]})
|
||||
|
||||
logger.info(
|
||||
"模板编辑器智能封面: template_id=%s plan_id=%s strategy=%s by user=%s",
|
||||
template_id,
|
||||
plan_id,
|
||||
strategy,
|
||||
current_user.user.id,
|
||||
)
|
||||
|
||||
return CoverGenerateResponse(
|
||||
type="smart",
|
||||
image_url=cover_url,
|
||||
frame_time=None,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/generate-cover", response_model=GenerateCoverResponse)
|
||||
def editor_generate_cover(
|
||||
template_id: str,
|
||||
body: GenerateCoverRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> GenerateCoverResponse:
|
||||
"""AI 生成封面"""
|
||||
_, plan_svc = services
|
||||
plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
|
||||
from packages.shared.ai_service import run_generate_cover
|
||||
|
||||
cover_data = run_generate_cover(
|
||||
plan_id=plan_id,
|
||||
asset_ids=body.asset_ids,
|
||||
cover_type=body.cover_type,
|
||||
frame_time=body.frame_time,
|
||||
)
|
||||
|
||||
current_config = dict(plan.config) if plan.config else {}
|
||||
current_config["cover"] = cover_data
|
||||
normalized = normalize_plan_config(current_config)
|
||||
plan_svc.update_plan_config(plan_id, {"cover": normalized["cover"]})
|
||||
|
||||
logger.info(
|
||||
"模板编辑器封面生成: template_id=%s plan_id=%s type=%s by user=%s",
|
||||
template_id,
|
||||
plan_id,
|
||||
body.cover_type,
|
||||
current_user.user.id,
|
||||
)
|
||||
|
||||
return GenerateCoverResponse(plan_id=plan_id, cover=cover_data)
|
||||
@@ -1,141 +0,0 @@
|
||||
"""模板编辑器依赖注入.
|
||||
|
||||
核心依赖:
|
||||
- get_editor_services: 获取模板+计划服务
|
||||
- get_draft_plan_id: 根据 template_id 获取或创建草稿,返回 plan_id
|
||||
- _check_queue_limits: 生成队列限流检查
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.task_enqueue import GLOBAL_PENDING_LIMIT, USER_PENDING_LIMIT
|
||||
from app.dependencies import get_db_session
|
||||
from app.services.edit_plan_service import EditPlanService
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
from fastapi import Depends, HTTPException, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.template_repository import (
|
||||
SQLAlchemyTemplateRepository,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def get_editor_services(
|
||||
db: Session = Depends(get_db_session),
|
||||
) -> tuple[EditTemplateService, EditPlanService]:
|
||||
"""获取模板编辑器所需的两个服务"""
|
||||
return EditTemplateService(db), EditPlanService(db)
|
||||
|
||||
|
||||
def get_draft_plan_id(
|
||||
template_id: str,
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
db: Session = Depends(get_db_session),
|
||||
) -> str:
|
||||
"""路径依赖:根据 template_id 获取或创建草稿,返回 plan_id.
|
||||
|
||||
这是模板编辑器路由的核心依赖——所有编辑器端点都先经过这里,
|
||||
确保 template_id → plan_id 的映射始终存在。
|
||||
|
||||
兼容策略:优先从新模板系统(edit_templates 表)查找,
|
||||
若不存在则回退到旧模板系统(templates 表),确保用户自建模板可用。
|
||||
"""
|
||||
tpl_svc, plan_svc = services
|
||||
user_id = str(current_user.user.id)
|
||||
|
||||
# 1. 草稿已存在 → 直接返回
|
||||
draft = tpl_svc.get_template_draft(template_id)
|
||||
if draft is not None:
|
||||
return draft.id
|
||||
|
||||
# 2. 新系统有模板 → 用新服务创建草稿
|
||||
if tpl_svc.get_template(template_id) is not None:
|
||||
draft = tpl_svc.create_template_draft(template_id, user_id=user_id)
|
||||
return draft.id
|
||||
|
||||
# 3. 回退到旧模板系统(templates 表)
|
||||
old_repo = SQLAlchemyTemplateRepository(db)
|
||||
old_template = old_repo.get(template_id, user_id=user_id)
|
||||
if old_template is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="模板不存在")
|
||||
|
||||
# 4. 基于旧模板创建草稿计划
|
||||
from app.services.plan_generator_service import PlanGeneratorService
|
||||
|
||||
from packages.domain.edit_template import EditTemplate, EditTemplateStatus
|
||||
from packages.domain.template_clip_config import ClipType, TemplateClipConfig
|
||||
|
||||
# 构造伪 EditTemplate 对象(只填 generate_from_template 需要的字段)
|
||||
pseudo_template = EditTemplate(
|
||||
id=old_template.id,
|
||||
name=old_template.name,
|
||||
editing_mode=old_template.mode,
|
||||
status=EditTemplateStatus.ACTIVE,
|
||||
)
|
||||
|
||||
# 将旧模板 segments 转换为 clip_configs
|
||||
clip_configs: list[TemplateClipConfig] = []
|
||||
for seg in old_template.segments or []:
|
||||
clip_configs.append(
|
||||
TemplateClipConfig(
|
||||
id=f"seg_{seg.id}",
|
||||
template_id=old_template.id,
|
||||
clip_type=ClipType.MAIN,
|
||||
order=seg.segment_order,
|
||||
min_duration=seg.duration_min,
|
||||
max_duration=seg.duration_max,
|
||||
)
|
||||
)
|
||||
|
||||
generator = PlanGeneratorService(db)
|
||||
result = generator.generate_from_template(
|
||||
template=pseudo_template,
|
||||
clip_configs=clip_configs,
|
||||
asset_ids=[],
|
||||
created_by_user_id=user_id,
|
||||
name=f"{old_template.name} - 草稿",
|
||||
)
|
||||
plan = result["plan"]
|
||||
|
||||
# 标记为模板草稿(后续可复用 tpl_svc.get_template_draft 的查找逻辑)
|
||||
plan_svc.update_plan_config(plan.id, {"is_template_draft": True})
|
||||
|
||||
logger.info(
|
||||
"旧模板自动创建草稿: template_id=%s draft_plan_id=%s user_id=%s",
|
||||
template_id,
|
||||
plan.id,
|
||||
user_id,
|
||||
)
|
||||
return plan.id
|
||||
|
||||
|
||||
def _check_queue_limits(gen_task_repo, user_id: str) -> None:
|
||||
"""队列限流预检查"""
|
||||
try:
|
||||
has_count = (
|
||||
hasattr(gen_task_repo, "count_pending_by_user")
|
||||
and hasattr(gen_task_repo, "count_pending_total")
|
||||
)
|
||||
if has_count:
|
||||
user_pending = gen_task_repo.count_pending_by_user(user_id)
|
||||
global_pending = gen_task_repo.count_pending_total()
|
||||
if user_pending >= USER_PENDING_LIMIT:
|
||||
raise HTTPException(
|
||||
status_code=429,
|
||||
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
|
||||
)
|
||||
if global_pending >= GLOBAL_PENDING_LIMIT:
|
||||
raise HTTPException(
|
||||
status_code=503,
|
||||
detail="系统繁忙,请稍后再试",
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.warning("[模板编辑器队列限流] 检查失败,跳过: %s", e)
|
||||
@@ -1,164 +0,0 @@
|
||||
"""草稿管理路由.
|
||||
|
||||
端点:
|
||||
- GET / 获取草稿详情
|
||||
- PUT / 更新草稿
|
||||
- POST /publish 发布草稿到模板
|
||||
- GET /versions 模板版本历史
|
||||
- POST /rollback 回滚到指定版本
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.services.edit_plan_service import EditPlanService
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
|
||||
from .dependencies import get_draft_plan_id, get_editor_services
|
||||
from .schemas import (
|
||||
EditorDraftResponse,
|
||||
EditorPublishResponse,
|
||||
EditorRollbackRequest,
|
||||
EditorRollbackResponse,
|
||||
EditorTemplateVersionItem,
|
||||
EditorUpdateRequest,
|
||||
EditorVersionListResponse,
|
||||
)
|
||||
|
||||
router = APIRouter(tags=["Template Editor"])
|
||||
|
||||
|
||||
@router.get("", response_model=EditorDraftResponse)
|
||||
def get_editor_draft(
|
||||
template_id: str,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""获取模板编辑器草稿详情
|
||||
|
||||
首次访问时自动创建草稿。
|
||||
"""
|
||||
_, plan_svc = services
|
||||
plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
clips = plan_svc.list_clips(plan_id)
|
||||
return EditorDraftResponse(
|
||||
plan_id=plan.id,
|
||||
template_id=plan.template_id,
|
||||
name=plan.name,
|
||||
status=plan.status.value if hasattr(plan.status, "value") else str(plan.status),
|
||||
config=plan.config or {},
|
||||
total_duration=plan.total_duration,
|
||||
clip_count=len(clips),
|
||||
)
|
||||
|
||||
|
||||
@router.put("", response_model=EditorDraftResponse)
|
||||
def update_editor_draft(
|
||||
template_id: str,
|
||||
req: EditorUpdateRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""更新模板编辑器草稿"""
|
||||
_, plan_svc = services
|
||||
plan = plan_svc.update_plan(
|
||||
plan_id,
|
||||
name=req.name,
|
||||
config=req.config,
|
||||
total_duration=req.total_duration,
|
||||
)
|
||||
clips = plan_svc.list_clips(plan_id)
|
||||
return EditorDraftResponse(
|
||||
plan_id=plan.id,
|
||||
template_id=plan.template_id,
|
||||
name=plan.name,
|
||||
status=plan.status.value if hasattr(plan.status, "value") else str(plan.status),
|
||||
config=plan.config or {},
|
||||
total_duration=plan.total_duration,
|
||||
clip_count=len(clips),
|
||||
)
|
||||
|
||||
|
||||
@router.post("/publish", response_model=EditorPublishResponse, status_code=status.HTTP_200_OK)
|
||||
def publish_draft_to_template(
|
||||
template_id: str,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""将草稿发布(同步)到正式模板
|
||||
|
||||
草稿的 config 和 clips 会同步覆盖到模板,事务保证一致性。
|
||||
"""
|
||||
tpl_svc, plan_svc = services
|
||||
try:
|
||||
tpl = tpl_svc.publish_template_from_draft(template_id, plan_id)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
clips = plan_svc.list_clips(plan_id)
|
||||
return EditorPublishResponse(
|
||||
template_id=tpl.id,
|
||||
status="published",
|
||||
clip_count=len(clips),
|
||||
version=tpl.version,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/versions", response_model=EditorVersionListResponse)
|
||||
def list_template_versions(
|
||||
template_id: str,
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
limit: int = Query(default=50, ge=1, le=200),
|
||||
):
|
||||
"""查询模板发布版本历史"""
|
||||
tpl_svc, _ = services
|
||||
versions = tpl_svc.list_template_versions(template_id, limit=limit)
|
||||
items = [
|
||||
EditorTemplateVersionItem(
|
||||
version=v.version,
|
||||
name=v.name,
|
||||
editing_mode=v.editing_mode,
|
||||
clip_count=len(v.clip_configs),
|
||||
change_note=v.change_note,
|
||||
published_by=v.published_by,
|
||||
created_at=(
|
||||
v.created_at.isoformat()
|
||||
if hasattr(v.created_at, "isoformat")
|
||||
else str(v.created_at)
|
||||
),
|
||||
)
|
||||
for v in versions
|
||||
]
|
||||
return EditorVersionListResponse(versions=items, total=len(items))
|
||||
|
||||
|
||||
@router.post("/rollback", response_model=EditorRollbackResponse, status_code=status.HTTP_200_OK)
|
||||
def rollback_template(
|
||||
template_id: str,
|
||||
request: EditorRollbackRequest,
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""回滚模板到指定历史版本
|
||||
|
||||
回滚本身也是一次发布,版本号会 +1,可以再次回滚。
|
||||
"""
|
||||
tpl_svc, _ = services
|
||||
try:
|
||||
tpl = tpl_svc.rollback_to_version(template_id, request.version)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
clip_configs = tpl_svc.list_clip_configs(template_id)
|
||||
return EditorRollbackResponse(
|
||||
template_id=tpl.id,
|
||||
status="rolled_back",
|
||||
rollback_to_version=request.version,
|
||||
new_version=tpl.version,
|
||||
clip_count=len(clip_configs),
|
||||
)
|
||||
@@ -1,195 +0,0 @@
|
||||
"""转场 & 滤镜路由.
|
||||
|
||||
端点:
|
||||
- GET /transition-presets 转场预设列表
|
||||
- PUT /clips/{clip_id}/transition 单片段转场
|
||||
- POST /transitions/batch 批量转场
|
||||
- GET /filter-presets 滤镜预设列表
|
||||
- GET /filter 滤镜配置
|
||||
- PUT /filter 更新滤镜
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.services.edit_plan_service import EditPlanService
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from packages.domain.config_schemas import normalize_plan_config
|
||||
|
||||
from .dependencies import get_draft_plan_id, get_editor_services
|
||||
from .schemas import (
|
||||
BatchTransitionRequest,
|
||||
BatchTransitionResponse,
|
||||
ClipTransitionResponse,
|
||||
FilterConfigResponse,
|
||||
FilterPresetListResponse,
|
||||
FilterUpdateRequest,
|
||||
TransitionPresetListResponse,
|
||||
TransitionUpdateRequest,
|
||||
)
|
||||
|
||||
router = APIRouter(tags=["Template Editor"])
|
||||
|
||||
|
||||
# ── 转场 ────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get("/transition-presets", response_model=TransitionPresetListResponse)
|
||||
def list_editor_transition_presets(
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> TransitionPresetListResponse:
|
||||
"""获取转场预设列表"""
|
||||
from packages.domain.transition_presets import TRANSITION_PRESETS
|
||||
|
||||
items = [
|
||||
{
|
||||
"id": p["id"],
|
||||
"name": p["name"],
|
||||
"category": p.get("category", "通用"),
|
||||
"duration": p.get("default_duration", 0.5),
|
||||
"description": p.get("description", ""),
|
||||
}
|
||||
for p in TRANSITION_PRESETS
|
||||
]
|
||||
return TransitionPresetListResponse(items=items, total=len(items))
|
||||
|
||||
|
||||
@router.put("/clips/{clip_id}/transition", response_model=ClipTransitionResponse)
|
||||
def update_editor_clip_transition(
|
||||
template_id: str,
|
||||
clip_id: str,
|
||||
body: TransitionUpdateRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> ClipTransitionResponse:
|
||||
"""设置单个片段的转场效果"""
|
||||
_, plan_svc = services
|
||||
try:
|
||||
clip = plan_svc.update_clip(
|
||||
clip_id,
|
||||
transition_effect=body.effect,
|
||||
transition_duration=body.duration,
|
||||
)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(status_code=400, detail=str(exc)) from exc
|
||||
|
||||
return ClipTransitionResponse(
|
||||
clip_id=clip.id,
|
||||
effect=clip.transition_effect.value
|
||||
if hasattr(clip.transition_effect, "value")
|
||||
else clip.transition_effect,
|
||||
duration=clip.transition_duration or 0.5,
|
||||
)
|
||||
|
||||
|
||||
@router.post("/transitions/batch", response_model=BatchTransitionResponse)
|
||||
def batch_update_editor_transitions(
|
||||
template_id: str,
|
||||
body: BatchTransitionRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> BatchTransitionResponse:
|
||||
"""批量设置所有片段的转场效果"""
|
||||
_, plan_svc = services
|
||||
clips = plan_svc.list_clips(plan_id, limit=500)
|
||||
updated = 0
|
||||
for clip in clips:
|
||||
if clip.order > 0: # 第一个片段不加转场
|
||||
try:
|
||||
plan_svc.update_clip(
|
||||
clip.id,
|
||||
transition_effect=body.effect,
|
||||
transition_duration=body.duration,
|
||||
)
|
||||
updated += 1
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
return BatchTransitionResponse(
|
||||
updated_count=updated,
|
||||
plan_id=plan_id,
|
||||
)
|
||||
|
||||
|
||||
# ── 滤镜 ────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
@router.get("/filter-presets", response_model=FilterPresetListResponse)
|
||||
def list_editor_filter_presets(
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> FilterPresetListResponse:
|
||||
"""获取滤镜预设列表"""
|
||||
from packages.domain.filter_presets import FILTER_PRESETS
|
||||
|
||||
items = [
|
||||
{
|
||||
"id": p["id"],
|
||||
"name": p["name"],
|
||||
"category": p.get("category", "通用"),
|
||||
"thumbnail": p.get("thumbnail", ""),
|
||||
"description": p.get("description", ""),
|
||||
}
|
||||
for p in FILTER_PRESETS
|
||||
]
|
||||
return FilterPresetListResponse(items=items, total=len(items))
|
||||
|
||||
|
||||
@router.get("/filter", response_model=FilterConfigResponse)
|
||||
def get_editor_filter(
|
||||
template_id: str,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> FilterConfigResponse:
|
||||
"""获取草稿的全局滤镜配置"""
|
||||
_, plan_svc = services
|
||||
plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
config = plan.config or {}
|
||||
filter_config = config.get("filter", {})
|
||||
|
||||
return FilterConfigResponse(
|
||||
plan_id=plan.id,
|
||||
enabled=filter_config.get("enabled", False),
|
||||
preset_id=filter_config.get("preset_id", ""),
|
||||
intensity=filter_config.get("intensity", 1.0),
|
||||
brightness=filter_config.get("brightness", 0.0),
|
||||
contrast=filter_config.get("contrast", 1.0),
|
||||
saturation=filter_config.get("saturation", 1.0),
|
||||
warmth=filter_config.get("warmth", 0.0),
|
||||
)
|
||||
|
||||
|
||||
@router.put("/filter", response_model=FilterConfigResponse)
|
||||
def update_editor_filter(
|
||||
template_id: str,
|
||||
body: FilterUpdateRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> FilterConfigResponse:
|
||||
"""更新草稿的全局滤镜配置"""
|
||||
_, plan_svc = services
|
||||
plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
|
||||
config = dict(plan.config) if plan.config else {}
|
||||
current_filter = dict(config.get("filter", {}))
|
||||
update_data = body.model_dump(exclude_none=True)
|
||||
current_filter.update(update_data)
|
||||
|
||||
config["filter"] = current_filter
|
||||
updated_plan = plan_svc.update_plan_config(plan_id, normalize_plan_config(config))
|
||||
|
||||
return FilterConfigResponse(
|
||||
plan_id=updated_plan.id,
|
||||
enabled=current_filter.get("enabled", False),
|
||||
preset_id=current_filter.get("preset_id", ""),
|
||||
intensity=current_filter.get("intensity", 1.0),
|
||||
brightness=current_filter.get("brightness", 0.0),
|
||||
contrast=current_filter.get("contrast", 1.0),
|
||||
saturation=current_filter.get("saturation", 1.0),
|
||||
warmth=current_filter.get("warmth", 0.0),
|
||||
)
|
||||
@@ -1,106 +0,0 @@
|
||||
"""导出配置路由.
|
||||
|
||||
端点:
|
||||
- GET /export-presets 导出预设列表
|
||||
- GET /export 导出配置
|
||||
- PUT /export 更新导出配置
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.services.edit_plan_service import EditPlanService
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
from packages.domain.config_schemas import normalize_plan_config
|
||||
|
||||
from .dependencies import get_draft_plan_id, get_editor_services
|
||||
from .schemas import ExportConfigResponse, ExportPresetListResponse, ExportUpdateRequest
|
||||
|
||||
router = APIRouter(tags=["Template Editor"])
|
||||
|
||||
|
||||
@router.get("/export-presets", response_model=ExportPresetListResponse)
|
||||
def list_editor_export_presets(
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> ExportPresetListResponse:
|
||||
"""获取导出预设列表"""
|
||||
from packages.domain.export_presets import EXPORT_PRESETS
|
||||
|
||||
items = [
|
||||
{
|
||||
"id": p["id"],
|
||||
"name": p["name"],
|
||||
"resolution": p.get("resolution", "1080p"),
|
||||
"fps": p.get("fps", 30),
|
||||
"video_bitrate": p.get("bitrate", ""),
|
||||
"audio_bitrate": p.get("audio_bitrate", 128),
|
||||
"format": p.get("format", "mp4"),
|
||||
"quality_preset": p.get("quality_preset", "balanced"),
|
||||
"description": p.get("description", ""),
|
||||
"size_hint": p.get("size_hint", ""),
|
||||
}
|
||||
for p in EXPORT_PRESETS
|
||||
]
|
||||
return ExportPresetListResponse(items=items, total=len(items))
|
||||
|
||||
|
||||
@router.get("/export", response_model=ExportConfigResponse)
|
||||
def get_editor_export(
|
||||
template_id: str,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> ExportConfigResponse:
|
||||
"""获取草稿的导出配置"""
|
||||
_, plan_svc = services
|
||||
plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
config = plan.config or {}
|
||||
export_config = config.get("export", {})
|
||||
|
||||
return ExportConfigResponse(
|
||||
plan_id=plan.id,
|
||||
resolution=export_config.get("resolution", "1080p"),
|
||||
fps=export_config.get("fps", 30),
|
||||
video_bitrate=export_config.get("video_bitrate", 4000),
|
||||
audio_bitrate=export_config.get("audio_bitrate", 128),
|
||||
format=export_config.get("format", "mp4"),
|
||||
quality_preset=export_config.get("quality_preset", "balanced"),
|
||||
watermark_enabled=export_config.get("watermark_enabled", True),
|
||||
watermark_text=export_config.get("watermark_text", ""),
|
||||
)
|
||||
|
||||
|
||||
@router.put("/export", response_model=ExportConfigResponse)
|
||||
def update_editor_export(
|
||||
template_id: str,
|
||||
body: ExportUpdateRequest,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> ExportConfigResponse:
|
||||
"""更新草稿的导出配置"""
|
||||
_, plan_svc = services
|
||||
plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
|
||||
config = dict(plan.config) if plan.config else {}
|
||||
current_export = dict(config.get("export", {}))
|
||||
update_data = body.model_dump(exclude_none=True)
|
||||
current_export.update(update_data)
|
||||
|
||||
config["export"] = current_export
|
||||
updated_plan = plan_svc.update_plan_config(plan_id, normalize_plan_config(config))
|
||||
updated_export = (updated_plan.config or {}).get("export", {})
|
||||
|
||||
return ExportConfigResponse(
|
||||
plan_id=updated_plan.id,
|
||||
resolution=updated_export.get("resolution", "1080p"),
|
||||
fps=updated_export.get("fps", 30),
|
||||
video_bitrate=updated_export.get("video_bitrate", 4000),
|
||||
audio_bitrate=updated_export.get("audio_bitrate", 128),
|
||||
format=updated_export.get("format", "mp4"),
|
||||
quality_preset=updated_export.get("quality_preset", "balanced"),
|
||||
watermark_enabled=updated_export.get("watermark_enabled", True),
|
||||
watermark_text=updated_export.get("watermark_text", ""),
|
||||
)
|
||||
@@ -1,250 +0,0 @@
|
||||
"""草稿生成路由.
|
||||
|
||||
端点:
|
||||
- POST /generate 触发生成
|
||||
- GET /generation-status 生成进度
|
||||
- GET /generations 生成记录列表
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import Any
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.core.celery_app import celery_app
|
||||
from app.core.storage import OSSStorageService, get_storage_service
|
||||
from app.dependencies import (
|
||||
get_asset_library_repository,
|
||||
get_asset_repository,
|
||||
get_db_session,
|
||||
)
|
||||
from app.schemas.generation_task import GenerationTaskResponse
|
||||
from app.services.edit_plan_service import EditPlanService
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
|
||||
SQLAlchemyGenerationTaskRepository,
|
||||
)
|
||||
from packages.application.generation_tasks import (
|
||||
CreateGenerationTaskCommand,
|
||||
CreateGenerationTaskUseCase,
|
||||
)
|
||||
from packages.domain.edit_plan import EditPlanStatus
|
||||
|
||||
from ._fallback import (
|
||||
_auto_fallback_assign_assets,
|
||||
_auto_fallback_auto_material_mode,
|
||||
_auto_fallback_copy_template_clips,
|
||||
_auto_fallback_draft_to_editing,
|
||||
)
|
||||
from .dependencies import _check_queue_limits, get_draft_plan_id, get_editor_services
|
||||
from .schemas import (
|
||||
ClipStatusItem,
|
||||
EditPlanGenerateResponse,
|
||||
EditPlanGenerationsResponse,
|
||||
EditPlanGenerationStatusResponse,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
router = APIRouter(tags=["Template Editor"])
|
||||
|
||||
|
||||
@router.post("/generate", response_model=EditPlanGenerateResponse)
|
||||
def generate_editor_draft(
|
||||
template_id: str,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
db: Session = Depends(get_db_session),
|
||||
current_user: AuthenticatedUser = Depends(get_current_user),
|
||||
asset_library_repo: Any = Depends(get_asset_library_repository),
|
||||
asset_repo: Any = Depends(get_asset_repository),
|
||||
) -> EditPlanGenerateResponse:
|
||||
"""触发模板草稿渲染生成"""
|
||||
_, plan_svc = services
|
||||
plan_check = plan_svc.get_plan_or_raise(plan_id)
|
||||
|
||||
# 自动兜底流程
|
||||
_auto_fallback_draft_to_editing(plan_svc, plan_id, plan_check)
|
||||
_auto_fallback_copy_template_clips(plan_svc, plan_id, plan_check, db)
|
||||
clips_without_asset = _auto_fallback_assign_assets(plan_svc, plan_id, plan_check)
|
||||
_auto_fallback_auto_material_mode(
|
||||
plan_svc, plan_id, plan_check, clips_without_asset, asset_library_repo, asset_repo
|
||||
)
|
||||
|
||||
# 检查是否可生成
|
||||
try:
|
||||
can_gen, reason = plan_svc.can_generate(plan_id)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)
|
||||
) from exc
|
||||
if not can_gen:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST, detail=reason
|
||||
)
|
||||
|
||||
try:
|
||||
clip_count = plan_svc.mark_clips_ready(plan_id)
|
||||
|
||||
gen_task_repo = SQLAlchemyGenerationTaskRepository(db)
|
||||
user_id = current_user.user.id
|
||||
_check_queue_limits(gen_task_repo, user_id)
|
||||
|
||||
gen_task_use_case = CreateGenerationTaskUseCase(gen_task_repo)
|
||||
plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
config_asset_ids = (plan.config or {}).get("asset_ids", [])
|
||||
gen_task = gen_task_use_case.execute(
|
||||
CreateGenerationTaskCommand(
|
||||
project_id=plan.project_id or "",
|
||||
template_id=plan.template_id,
|
||||
created_by_user_id=current_user.user.id,
|
||||
source_edit_plan_id=plan_id,
|
||||
asset_ids=list(config_asset_ids) if config_asset_ids else [],
|
||||
),
|
||||
)
|
||||
|
||||
plan_svc.update_plan_config(plan_id, {"generation_task_id": gen_task.id})
|
||||
plan_svc.transition_status(plan_id, EditPlanStatus.RENDERING)
|
||||
celery_app.send_task("worker.render_edit_plan", args=[plan_id])
|
||||
|
||||
updated_plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
|
||||
logger.info(
|
||||
"模板编辑器触发生成: template_id=%s plan_id=%s gen_task_id=%s clips=%d by user=%s",
|
||||
template_id,
|
||||
plan_id,
|
||||
gen_task.id,
|
||||
clip_count,
|
||||
current_user.user.id,
|
||||
)
|
||||
|
||||
return EditPlanGenerateResponse(
|
||||
plan_id=plan_id,
|
||||
plan_status=updated_plan.status.value
|
||||
if hasattr(updated_plan.status, "value")
|
||||
else updated_plan.status,
|
||||
generation_task_id=gen_task.id,
|
||||
clip_count=clip_count,
|
||||
)
|
||||
except HTTPException:
|
||||
raise
|
||||
except Exception as _e:
|
||||
logger.exception(
|
||||
"模板编辑器触发生成失败: template_id=%s plan_id=%s",
|
||||
template_id,
|
||||
plan_id,
|
||||
)
|
||||
try:
|
||||
plan_svc.transition_status(plan_id, EditPlanStatus.FAILED)
|
||||
except Exception:
|
||||
pass
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail="生成失败,请稍后重试",
|
||||
) from _e
|
||||
|
||||
|
||||
@router.get("/generation-status", response_model=EditPlanGenerationStatusResponse)
|
||||
def get_editor_generation_status(
|
||||
template_id: str,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
storage_service: OSSStorageService = Depends(get_storage_service),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> EditPlanGenerationStatusResponse:
|
||||
"""查询草稿生成进度"""
|
||||
_, plan_svc = services
|
||||
try:
|
||||
gen_status = plan_svc.get_generation_status(plan_id)
|
||||
except ValueError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_404_NOT_FOUND, detail=str(exc)
|
||||
) from exc
|
||||
|
||||
plan = gen_status["plan"]
|
||||
clips = gen_status["clips"]
|
||||
|
||||
clip_items = [
|
||||
ClipStatusItem(
|
||||
clip_id=c.id,
|
||||
clip_type=c.clip_type,
|
||||
order=c.order,
|
||||
status=c.status.value if hasattr(c.status, "value") else c.status,
|
||||
asset_id=c.asset_id or "",
|
||||
text_content=c.text_content or "",
|
||||
duration=c.duration,
|
||||
)
|
||||
for c in clips
|
||||
]
|
||||
|
||||
raw_video_url = (plan.config or {}).get("rendered_url", "")
|
||||
video_url = ""
|
||||
if raw_video_url:
|
||||
try:
|
||||
video_url = storage_service.get_download_url(
|
||||
raw_video_url, expires_seconds=86400
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
"生成视频签名URL失败: template_id=%s error=%s", template_id, e
|
||||
)
|
||||
video_url = raw_video_url
|
||||
|
||||
progress = gen_status.get("progress", 0.0)
|
||||
error_message = gen_status.get("error_message", "")
|
||||
gen_task_status = gen_status.get("generation_task_status")
|
||||
plan_status_val = (
|
||||
plan.status.value if hasattr(plan.status, "value") else plan.status
|
||||
)
|
||||
if plan_status_val == "completed" and progress < 100:
|
||||
progress = 100.0
|
||||
|
||||
return EditPlanGenerationStatusResponse(
|
||||
plan_id=plan_id,
|
||||
plan_status=plan_status_val,
|
||||
generation_task_id=gen_status["generation_task_id"],
|
||||
generation_task_status=gen_task_status,
|
||||
progress=progress,
|
||||
video_url=video_url,
|
||||
error_message=error_message,
|
||||
clips=clip_items,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/generations", response_model=EditPlanGenerationsResponse)
|
||||
def list_editor_generations(
|
||||
template_id: str,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
db: Session = Depends(get_db_session),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> EditPlanGenerationsResponse:
|
||||
"""查询草稿关联的生成记录列表"""
|
||||
_, plan_svc = services
|
||||
plan_svc.get_plan_or_raise(plan_id)
|
||||
|
||||
gen_task_repo = SQLAlchemyGenerationTaskRepository(db)
|
||||
tasks = gen_task_repo.list_by_source_edit_plan(plan_id)
|
||||
items = [
|
||||
GenerationTaskResponse(
|
||||
id=t.id,
|
||||
project_id=t.project_id,
|
||||
asset_library_id=t.asset_library_id,
|
||||
strategy_id=t.strategy_id,
|
||||
voice_library_id=t.voice_library_id,
|
||||
template_id=t.template_id,
|
||||
asset_ids=t.asset_ids,
|
||||
title_ids=t.title_ids,
|
||||
voice_ids=t.voice_ids,
|
||||
source_edit_plan_id=t.source_edit_plan_id or "",
|
||||
status=t.status.value if hasattr(t.status, "value") else t.status,
|
||||
progress=t.progress,
|
||||
result_count=t.result_count,
|
||||
error_message=t.error_message,
|
||||
)
|
||||
for t in tasks
|
||||
]
|
||||
return EditPlanGenerationsResponse(items=items, total=len(items))
|
||||
@@ -1,622 +0,0 @@
|
||||
"""模板编辑器所有 Pydantic Schema 定义.
|
||||
|
||||
集中管理,避免在路由文件里散落 40+ 个 model。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re as _re
|
||||
from typing import Any, List, Optional
|
||||
|
||||
from app.schemas.generation_task import GenerationTaskResponse
|
||||
from pydantic import BaseModel, Field, validator
|
||||
|
||||
_EXPORT_RESOLUTION_PATTERN = _re.compile(r"^\d+x\d+$")
|
||||
_EXPORT_VALID_QUALITY_PRESETS = {"ultra_fast", "fast", "balanced", "high", "best"}
|
||||
_EXPORT_VALID_FORMATS = {"mp4", "mov"}
|
||||
|
||||
|
||||
# ── 生成状态相关 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class ClipStatusItem(BaseModel):
|
||||
"""片段生成状态"""
|
||||
|
||||
clip_id: str
|
||||
clip_type: str
|
||||
order: int
|
||||
status: str
|
||||
asset_id: str
|
||||
text_content: str
|
||||
duration: float
|
||||
|
||||
|
||||
class EditPlanGenerationStatusResponse(BaseModel):
|
||||
"""剪辑计划生成进度响应体"""
|
||||
|
||||
plan_id: str
|
||||
plan_status: str
|
||||
generation_task_id: Optional[str] = None
|
||||
generation_task_status: Optional[str] = None
|
||||
progress: float = 0.0
|
||||
video_url: str = ""
|
||||
error_message: str = ""
|
||||
clips: List[ClipStatusItem]
|
||||
|
||||
|
||||
class EditPlanGenerateResponse(BaseModel):
|
||||
"""剪辑计划触发生成响应体"""
|
||||
|
||||
plan_id: str
|
||||
plan_status: str
|
||||
generation_task_id: str
|
||||
clip_count: int
|
||||
|
||||
|
||||
class EditPlanGenerationsResponse(BaseModel):
|
||||
"""剪辑计划关联的生成记录列表响应体"""
|
||||
|
||||
items: List[GenerationTaskResponse]
|
||||
total: int
|
||||
|
||||
|
||||
# ── AI 推荐 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class AIRecommendRequest(BaseModel):
|
||||
"""AI 推荐片段方案请求体"""
|
||||
|
||||
asset_ids: List[str] = Field(default_factory=list, description="素材 ID 列表")
|
||||
editing_mode: str = Field(
|
||||
default="one_take", description="剪辑模式: one_take / pip / voice_over / voice_pip"
|
||||
)
|
||||
target_duration: float = Field(
|
||||
default=30.0, ge=1.0, le=600.0, description="目标时长(秒)"
|
||||
)
|
||||
|
||||
|
||||
class AIRecommendClipItem(BaseModel):
|
||||
"""AI 推荐的单个片段"""
|
||||
|
||||
clip_type: str = Field(..., description="片段类型: intro / showcase / title / subtitle / cta / outro")
|
||||
order: int = Field(..., ge=0, description="片段顺序")
|
||||
text_content: str = Field(default="", description="文字内容")
|
||||
duration: float = Field(..., ge=0.0, description="片段时长(秒)")
|
||||
transition_effect: str = Field(default="cut", description="转场效果")
|
||||
transition_duration: float = Field(default=0.0, ge=0.0, description="转场时长(秒),0 表示使用默认值")
|
||||
asset_id: str = Field(default="", description="关联素材 ID")
|
||||
start_time: float = Field(default=0.0, ge=0.0, description="素材截取起始时间(秒)")
|
||||
config: dict[str, Any] = Field(default_factory=dict, description="片段额外配置")
|
||||
|
||||
|
||||
class AIRecommendResponse(BaseModel):
|
||||
"""AI 推荐片段方案响应体"""
|
||||
|
||||
plan_id: str = Field(..., description="剪辑计划 ID")
|
||||
clips: List[AIRecommendClipItem] = Field(..., description="推荐的片段列表")
|
||||
config: dict[str, Any] = Field(..., description="推荐的 plan config(cover/title/subtitle/bgm)")
|
||||
total_duration: float = Field(..., ge=0.0, description="推荐方案总时长(秒)")
|
||||
confidence: float = Field(..., ge=0.0, le=1.0, description="AI 推荐置信度 (0~1)")
|
||||
|
||||
|
||||
# ── 封面生成 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class GenerateCoverRequest(BaseModel):
|
||||
"""AI 封面生成请求体"""
|
||||
|
||||
asset_ids: List[str] = Field(default_factory=list, description="素材 ID 列表(确定视频来源)")
|
||||
cover_type: str = Field(
|
||||
default="ai_frame",
|
||||
description="封面类型: ai_frame / manual / upload / ai_regenerate",
|
||||
)
|
||||
frame_time: Optional[float] = Field(
|
||||
default=None,
|
||||
ge=0.0,
|
||||
description="手动选帧时间点(秒),仅 cover_type=manual 时有效",
|
||||
)
|
||||
|
||||
|
||||
class GenerateCoverResponse(BaseModel):
|
||||
"""AI 封面生成响应体"""
|
||||
|
||||
plan_id: str = Field(..., description="剪辑计划 ID")
|
||||
cover: dict[str, Any] = Field(..., description="封面数据(type / image_url / frame_time 等)")
|
||||
|
||||
|
||||
# ── BGM ────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class BGMConfigUpdateRequest(BaseModel):
|
||||
"""更新BGM配置请求体"""
|
||||
|
||||
enabled: Optional[bool] = Field(default=None, description="是否启用 BGM")
|
||||
source: Optional[str] = Field(default=None, description="BGM 来源: library/upload/ai_recommend")
|
||||
asset_id: Optional[str] = Field(default=None, max_length=64, description="BGM 素材 ID")
|
||||
preset_id: Optional[str] = Field(default=None, max_length=64, description="预设 BGM ID")
|
||||
audio_url: Optional[str] = Field(default=None, max_length=500, description="BGM 音频 URL")
|
||||
volume: Optional[float] = Field(default=None, ge=0.0, le=1.0, description="音量 (0.0 ~ 1.0)")
|
||||
fade_in: Optional[float] = Field(default=None, ge=0.0, le=30.0, description="淡入时长(秒)")
|
||||
fade_out: Optional[float] = Field(default=None, ge=0.0, le=30.0, description="淡出时长(秒)")
|
||||
loop_enabled: Optional[bool] = Field(default=None, description="是否循环播放")
|
||||
sidechain_enabled: Optional[bool] = Field(default=None, description="是否启用人声闪避")
|
||||
sidechain_ratio: Optional[float] = Field(default=None, ge=0.0, le=1.0, description="闪避音量降低比例")
|
||||
|
||||
|
||||
# ── 片段调整 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class SpeedAdjustRequest(BaseModel):
|
||||
"""调速请求"""
|
||||
|
||||
speed: float = Field(..., ge=0.25, le=4.0, description="播放速度 0.25~4.0")
|
||||
|
||||
|
||||
class VolumeAdjustRequest(BaseModel):
|
||||
"""音量调节请求"""
|
||||
|
||||
volume: float = Field(..., ge=0.0, le=2.0, description="音量倍率 0~2.0(1.0=原音量)")
|
||||
|
||||
|
||||
class TrimAdjustRequest(BaseModel):
|
||||
"""裁剪请求"""
|
||||
|
||||
trim_start: float = Field(0.0, ge=0.0, description="开头裁剪秒数")
|
||||
trim_end: float = Field(0.0, ge=0.0, description="结尾裁剪秒数")
|
||||
|
||||
|
||||
class ClipAdjustmentsRequest(BaseModel):
|
||||
"""统一调整请求"""
|
||||
|
||||
speed: Optional[float] = Field(default=None, ge=0.25, le=4.0)
|
||||
volume: Optional[float] = Field(default=None, ge=0.0, le=2.0)
|
||||
trim_start: Optional[float] = Field(default=None, ge=0.0)
|
||||
trim_end: Optional[float] = Field(default=None, ge=0.0)
|
||||
|
||||
|
||||
class BatchSpeedRequest(BaseModel):
|
||||
"""批量调速请求"""
|
||||
|
||||
speed: float = Field(..., ge=0.25, le=4.0, description="播放速度")
|
||||
|
||||
|
||||
class ClipAdjustResponse(BaseModel):
|
||||
"""片段调整响应"""
|
||||
|
||||
clip_id: str
|
||||
speed: float
|
||||
volume: float
|
||||
trim_start: float
|
||||
trim_end: float
|
||||
duration: float
|
||||
|
||||
|
||||
class BatchSpeedResponse(BaseModel):
|
||||
"""批量调速响应"""
|
||||
|
||||
updated_count: int
|
||||
plan_id: str
|
||||
|
||||
|
||||
# ── 片段批量操作 ────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class ClipReorderItem(BaseModel):
|
||||
"""重排序条目"""
|
||||
|
||||
clip_id: str
|
||||
new_order: int = Field(..., ge=0, description="新的排序序号")
|
||||
|
||||
|
||||
class ClipReorderRequest(BaseModel):
|
||||
"""片段重排序请求"""
|
||||
|
||||
items: List[ClipReorderItem] = Field(..., min_length=1, max_length=500, description="重排序条目列表")
|
||||
|
||||
|
||||
class ClipReorderResponse(BaseModel):
|
||||
"""片段重排序响应"""
|
||||
|
||||
success: bool = True
|
||||
updated_count: int
|
||||
message: str = ""
|
||||
|
||||
|
||||
class ClipBatchDeleteRequest(BaseModel):
|
||||
"""批量删除片段请求"""
|
||||
|
||||
clip_ids: List[str] = Field(..., min_length=1, max_length=500, description="要删除的片段ID列表")
|
||||
|
||||
|
||||
class ClipBatchDeleteResponse(BaseModel):
|
||||
"""批量删除片段响应"""
|
||||
|
||||
success: bool = True
|
||||
deleted_count: int
|
||||
message: str = ""
|
||||
|
||||
|
||||
class ClipsFromAssetsRequest(BaseModel):
|
||||
"""从素材批量创建片段请求"""
|
||||
|
||||
asset_ids: List[str] = Field(
|
||||
..., min_length=1, max_length=200, description="素材 ID 列表,按顺序追加到时间线末尾"
|
||||
)
|
||||
clip_type: str = Field(default="main", description="片段类型,默认 main")
|
||||
|
||||
|
||||
class ClipsFromAssetsResponse(BaseModel):
|
||||
"""从素材批量创建片段响应"""
|
||||
|
||||
success: bool = True
|
||||
created_count: int
|
||||
message: str = ""
|
||||
clip_ids: List[str] = Field(default_factory=list, description="创建的片段ID列表")
|
||||
|
||||
|
||||
# ── 封面配置 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class CoverConfigResponse(BaseModel):
|
||||
"""封面配置响应"""
|
||||
|
||||
type: str = Field(..., description="封面类型: ai_frame / manual / upload")
|
||||
image_url: str = Field(default="", description="封面图片 URL")
|
||||
frame_time: Optional[float] = Field(default=None, description="抽帧时间点(秒)")
|
||||
|
||||
|
||||
class CoverUpdateRequest(BaseModel):
|
||||
"""更新封面配置请求"""
|
||||
|
||||
type: Optional[str] = Field(default=None, description="封面类型")
|
||||
image_url: Optional[str] = Field(default=None, description="封面图片 URL")
|
||||
frame_time: Optional[float] = Field(default=None, ge=0.0, description="抽帧时间点(秒)")
|
||||
|
||||
|
||||
class CoverExtractRequest(BaseModel):
|
||||
"""从片段抽帧生成封面请求"""
|
||||
|
||||
clip_id: str = Field(..., description="片段 ID")
|
||||
frame_time: float = Field(1.0, ge=0.0, description="抽帧时间点(秒)")
|
||||
|
||||
|
||||
class CoverSmartRequest(BaseModel):
|
||||
"""智能选帧请求"""
|
||||
|
||||
clip_id: Optional[str] = Field(default=None, description="指定片段 ID(不传则用第一个视频片段)")
|
||||
|
||||
|
||||
class CoverGenerateResponse(BaseModel):
|
||||
"""封面生成响应"""
|
||||
|
||||
type: str = Field(..., description="封面类型")
|
||||
image_url: str = Field(..., description="封面图片 URL")
|
||||
frame_time: Optional[float] = Field(default=None, description="抽帧时间点(秒)")
|
||||
|
||||
|
||||
# ── 导出配置 ────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class ExportConfigResponse(BaseModel):
|
||||
"""导出配置响应"""
|
||||
|
||||
resolution: str
|
||||
fps: int
|
||||
video_bitrate: int
|
||||
audio_bitrate: int
|
||||
format: str
|
||||
quality_preset: str
|
||||
watermark_enabled: bool
|
||||
watermark_text: str
|
||||
|
||||
|
||||
class ExportUpdateRequest(BaseModel):
|
||||
"""更新导出配置请求"""
|
||||
|
||||
resolution: Optional[str] = None
|
||||
fps: Optional[int] = Field(default=None, ge=15, le=60)
|
||||
video_bitrate: Optional[int] = Field(default=None, ge=1000, le=20000)
|
||||
audio_bitrate: Optional[int] = Field(default=None, ge=64, le=320)
|
||||
format: Optional[str] = None
|
||||
quality_preset: Optional[str] = None
|
||||
watermark_enabled: Optional[bool] = None
|
||||
watermark_text: Optional[str] = None
|
||||
|
||||
@validator("resolution")
|
||||
def validate_resolution(cls, v):
|
||||
if v is None:
|
||||
return v
|
||||
if not _EXPORT_RESOLUTION_PATTERN.match(v):
|
||||
raise ValueError("分辨率格式错误,应为 宽x高,如 1080x1920")
|
||||
w, h = v.split("x")
|
||||
if int(w) < 100 or int(h) < 100:
|
||||
raise ValueError("分辨率数值过小")
|
||||
if int(w) > 4096 or int(h) > 4096:
|
||||
raise ValueError("分辨率数值过大,最大 4096x4096")
|
||||
return v
|
||||
|
||||
@validator("format")
|
||||
def validate_format(cls, v):
|
||||
if v is None:
|
||||
return v
|
||||
if v not in _EXPORT_VALID_FORMATS:
|
||||
raise ValueError(f"无效格式: {v},支持: {_EXPORT_VALID_FORMATS}")
|
||||
return v
|
||||
|
||||
@validator("quality_preset")
|
||||
def validate_quality_preset(cls, v):
|
||||
if v is None:
|
||||
return v
|
||||
if v not in _EXPORT_VALID_QUALITY_PRESETS:
|
||||
raise ValueError(f"无效质量预设: {v},支持: {_EXPORT_VALID_QUALITY_PRESETS}")
|
||||
return v
|
||||
|
||||
|
||||
class ExportPresetItem(BaseModel):
|
||||
"""导出预设条目"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
resolution: str
|
||||
fps: int
|
||||
video_bitrate: int
|
||||
audio_bitrate: int
|
||||
format: str
|
||||
quality_preset: str
|
||||
description: str
|
||||
size_hint: str
|
||||
|
||||
|
||||
class ExportPresetListResponse(BaseModel):
|
||||
"""导出预设列表响应"""
|
||||
|
||||
items: List[ExportPresetItem]
|
||||
total: int
|
||||
|
||||
|
||||
# ── 滤镜 ────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class FilterPresetResponse(BaseModel):
|
||||
"""滤镜预设响应"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
category: str
|
||||
description: str
|
||||
tags: List[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class FilterConfigResponse(BaseModel):
|
||||
"""滤镜配置响应"""
|
||||
|
||||
enabled: bool
|
||||
preset_id: str
|
||||
intensity: int
|
||||
brightness: float
|
||||
contrast: float
|
||||
saturation: float
|
||||
warmth: float
|
||||
|
||||
|
||||
class FilterUpdateRequest(BaseModel):
|
||||
"""更新滤镜配置请求"""
|
||||
|
||||
enabled: Optional[bool] = None
|
||||
preset_id: Optional[str] = None
|
||||
intensity: Optional[int] = Field(default=None, ge=0, le=100)
|
||||
brightness: Optional[float] = Field(default=None, ge=-1.0, le=1.0)
|
||||
contrast: Optional[float] = Field(default=None, ge=0.0, le=2.0)
|
||||
saturation: Optional[float] = Field(default=None, ge=0.0, le=3.0)
|
||||
warmth: Optional[float] = Field(default=None, ge=-1.0, le=1.0)
|
||||
|
||||
|
||||
class FilterPresetListResponse(BaseModel):
|
||||
"""滤镜预设列表响应"""
|
||||
|
||||
items: List[FilterPresetResponse]
|
||||
total: int
|
||||
|
||||
|
||||
# ── 转场 ────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class TransitionPresetResponse(BaseModel):
|
||||
"""转场预设响应"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
category: str
|
||||
description: str
|
||||
tags: List[str] = Field(default_factory=list)
|
||||
default_duration: float
|
||||
min_duration: float
|
||||
max_duration: float
|
||||
|
||||
|
||||
class TransitionUpdateRequest(BaseModel):
|
||||
"""更新转场请求"""
|
||||
|
||||
effect: str = Field(..., description="转场效果 ID")
|
||||
duration: Optional[float] = Field(default=None, ge=0.0, description="转场时长(秒)")
|
||||
|
||||
|
||||
class BatchTransitionRequest(BaseModel):
|
||||
"""批量设置转场请求"""
|
||||
|
||||
effect: str = Field(..., description="转场效果 ID")
|
||||
duration: Optional[float] = Field(default=None, ge=0.0, description="转场时长(秒)")
|
||||
apply_to: str = Field(
|
||||
default="all",
|
||||
description="应用范围: all=所有片段, except_first=除第一个外, except_last=除最后一个, middle=中间片段",
|
||||
)
|
||||
|
||||
|
||||
class ClipTransitionResponse(BaseModel):
|
||||
"""片段转场信息响应"""
|
||||
|
||||
clip_id: str
|
||||
effect: str
|
||||
duration: float
|
||||
|
||||
|
||||
class BatchTransitionResponse(BaseModel):
|
||||
"""批量转场响应"""
|
||||
|
||||
updated_count: int
|
||||
plan_id: str
|
||||
|
||||
|
||||
class TransitionPresetListResponse(BaseModel):
|
||||
"""转场预设列表响应"""
|
||||
|
||||
items: List[TransitionPresetResponse]
|
||||
total: int
|
||||
|
||||
|
||||
# ── 编辑器草稿 & 片段 ───────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class EditorDraftResponse(BaseModel):
|
||||
"""模板编辑器草稿详情响应"""
|
||||
|
||||
plan_id: str
|
||||
template_id: str
|
||||
name: str
|
||||
status: str
|
||||
config: dict[str, Any]
|
||||
total_duration: float
|
||||
clip_count: int
|
||||
is_draft: bool = True
|
||||
|
||||
|
||||
class EditorUpdateRequest(BaseModel):
|
||||
"""更新草稿请求"""
|
||||
|
||||
name: Optional[str] = Field(default=None, min_length=1, max_length=200)
|
||||
config: Optional[dict[str, Any]] = Field(default=None)
|
||||
total_duration: Optional[float] = Field(default=None, ge=0.0)
|
||||
|
||||
|
||||
class EditorClipResponse(BaseModel):
|
||||
"""片段响应"""
|
||||
|
||||
id: str
|
||||
plan_id: str
|
||||
clip_type: str
|
||||
order: int
|
||||
duration: float
|
||||
text_content: str = ""
|
||||
transition_effect: str = "cut"
|
||||
playback_speed: float = 1.0
|
||||
config: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class EditorClipListResponse(BaseModel):
|
||||
"""片段列表响应"""
|
||||
|
||||
items: List[EditorClipResponse]
|
||||
total: int
|
||||
|
||||
|
||||
class EditorClipCreateRequest(BaseModel):
|
||||
"""创建片段请求"""
|
||||
|
||||
clip_type: str = Field(..., min_length=1, max_length=32)
|
||||
order: int = Field(..., ge=0)
|
||||
duration: float = Field(..., gt=0.0)
|
||||
text_content: str = Field(default="", max_length=2000)
|
||||
transition_effect: str = Field(default="cut", max_length=32)
|
||||
config: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
|
||||
class EditorClipUpdateRequest(BaseModel):
|
||||
"""更新片段请求"""
|
||||
|
||||
order: Optional[int] = Field(default=None, ge=0)
|
||||
duration: Optional[float] = Field(default=None, gt=0.0)
|
||||
text_content: Optional[str] = Field(default=None, max_length=2000)
|
||||
transition_effect: Optional[str] = Field(default=None, max_length=32)
|
||||
playback_speed: Optional[float] = Field(default=None, gt=0.0)
|
||||
config: Optional[dict[str, Any]] = None
|
||||
|
||||
|
||||
class EditorPublishResponse(BaseModel):
|
||||
"""发布草稿响应"""
|
||||
|
||||
template_id: str
|
||||
status: str = "published"
|
||||
clip_count: int
|
||||
version: int = 1
|
||||
|
||||
|
||||
class EditorTemplateVersionItem(BaseModel):
|
||||
"""模板版本历史条目"""
|
||||
|
||||
version: int
|
||||
name: str
|
||||
editing_mode: str
|
||||
clip_count: int
|
||||
change_note: str
|
||||
published_by: str
|
||||
created_at: str
|
||||
|
||||
|
||||
class EditorVersionListResponse(BaseModel):
|
||||
"""模板版本列表响应"""
|
||||
|
||||
versions: list[EditorTemplateVersionItem]
|
||||
total: int
|
||||
|
||||
|
||||
class EditorRollbackRequest(BaseModel):
|
||||
"""回滚请求体"""
|
||||
|
||||
version: int
|
||||
|
||||
|
||||
class EditorRollbackResponse(BaseModel):
|
||||
"""回滚响应"""
|
||||
|
||||
template_id: str
|
||||
status: str = "rolled_back"
|
||||
rollback_to_version: int
|
||||
new_version: int
|
||||
clip_count: int
|
||||
|
||||
|
||||
# ── 片段分割与合并 ──────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class SplitClipRequest(BaseModel):
|
||||
"""分割片段请求体"""
|
||||
|
||||
split_time: float = Field(..., gt=0, description="分割点(秒,相对于片段起始)")
|
||||
|
||||
|
||||
class MergeClipsRequest(BaseModel):
|
||||
"""合并片段请求体"""
|
||||
|
||||
clip_ids: list[str] = Field(..., min_length=2, description="要合并的片段 ID 列表")
|
||||
|
||||
|
||||
# ── 时间线 ──────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
class EditorTimelineSceneResponse(BaseModel):
|
||||
"""时间线场景"""
|
||||
|
||||
scene: str
|
||||
time: str
|
||||
duration: float
|
||||
color: str
|
||||
clip_id: str = ""
|
||||
clip_type: str = ""
|
||||
|
||||
|
||||
class EditorTimelineResponse(BaseModel):
|
||||
"""时间线响应"""
|
||||
|
||||
plan_id: str
|
||||
total_duration: float
|
||||
scenes: List[EditorTimelineSceneResponse]
|
||||
@@ -1,173 +0,0 @@
|
||||
"""字幕管理路由.
|
||||
|
||||
端点:
|
||||
- GET /clips/{clip_id}/subtitles 字幕列表
|
||||
- POST /clips/{clip_id}/subtitles 新增字幕
|
||||
- PUT /clips/{clip_id}/subtitles/{subtitle_id} 更新字幕
|
||||
- DELETE /clips/{clip_id}/subtitles/{subtitle_id} 删除字幕
|
||||
- PUT /clips/{clip_id}/subtitles 批量更新字幕(全量替换)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.services.edit_plan_service import EditPlanService
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
|
||||
from .dependencies import get_draft_plan_id, get_editor_services
|
||||
|
||||
router = APIRouter(tags=["Template Editor"])
|
||||
|
||||
|
||||
def _get_clip_subtitles(plan_svc: EditPlanService, clip_id: str, plan_id: str) -> list[dict[str, Any]]:
|
||||
"""获取片段字幕列表,统一校验"""
|
||||
clip = plan_svc.get_clip(clip_id)
|
||||
if not clip:
|
||||
raise HTTPException(status_code=404, detail="片段不存在")
|
||||
if clip.plan_id != plan_id:
|
||||
raise HTTPException(status_code=404, detail="片段不存在")
|
||||
config = clip.config or {}
|
||||
subtitles = config.get("subtitles", [])
|
||||
if not isinstance(subtitles, list):
|
||||
subtitles = []
|
||||
return subtitles
|
||||
|
||||
|
||||
@router.get("/clips/{clip_id}/subtitles", response_model=list[dict[str, Any]])
|
||||
def get_editor_clip_subtitles(
|
||||
template_id: str,
|
||||
clip_id: str,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> list[dict[str, Any]]:
|
||||
"""获取片段的字幕列表"""
|
||||
_, plan_svc = services
|
||||
return _get_clip_subtitles(plan_svc, clip_id, plan_id)
|
||||
|
||||
|
||||
@router.post("/clips/{clip_id}/subtitles", response_model=dict[str, Any])
|
||||
def create_editor_clip_subtitle(
|
||||
template_id: str,
|
||||
clip_id: str,
|
||||
body: dict[str, Any],
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> dict[str, Any]:
|
||||
"""新增片段字幕"""
|
||||
_, plan_svc = services
|
||||
clip = plan_svc.get_clip(clip_id)
|
||||
if not clip:
|
||||
raise HTTPException(status_code=404, detail="片段不存在")
|
||||
|
||||
config = dict(clip.config) if clip.config else {}
|
||||
subtitles = config.get("subtitles", [])
|
||||
if not isinstance(subtitles, list):
|
||||
subtitles = []
|
||||
|
||||
new_id = f"sub_{len(subtitles) + 1}"
|
||||
new_subtitle = {
|
||||
"id": body.get("id", new_id),
|
||||
"start_time": body.get("start_time", 0.0),
|
||||
"end_time": body.get("end_time", 0.0),
|
||||
"text": body.get("text", ""),
|
||||
"style": body.get("style", {}),
|
||||
}
|
||||
subtitles.append(new_subtitle)
|
||||
config["subtitles"] = subtitles
|
||||
|
||||
plan_svc.update_clip(clip_id, config=config)
|
||||
return new_subtitle
|
||||
|
||||
|
||||
@router.put("/clips/{clip_id}/subtitles/{subtitle_id}", response_model=dict[str, Any])
|
||||
def update_editor_clip_subtitle(
|
||||
template_id: str,
|
||||
clip_id: str,
|
||||
subtitle_id: str,
|
||||
body: dict[str, Any],
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> dict[str, Any]:
|
||||
"""更新片段字幕"""
|
||||
_, plan_svc = services
|
||||
clip = plan_svc.get_clip(clip_id)
|
||||
if not clip:
|
||||
raise HTTPException(status_code=404, detail="片段不存在")
|
||||
|
||||
config = dict(clip.config) if clip.config else {}
|
||||
subtitles = config.get("subtitles", [])
|
||||
if not isinstance(subtitles, list):
|
||||
subtitles = []
|
||||
|
||||
found = False
|
||||
for i, sub in enumerate(subtitles):
|
||||
if sub.get("id") == subtitle_id:
|
||||
subtitles[i].update(body)
|
||||
found = True
|
||||
break
|
||||
|
||||
if not found:
|
||||
raise HTTPException(status_code=404, detail="字幕不存在")
|
||||
|
||||
config["subtitles"] = subtitles
|
||||
plan_svc.update_clip(clip_id, config=config)
|
||||
return subtitles[i]
|
||||
|
||||
|
||||
@router.delete(
|
||||
"/clips/{clip_id}/subtitles/{subtitle_id}",
|
||||
status_code=status.HTTP_204_NO_CONTENT,
|
||||
)
|
||||
def delete_editor_clip_subtitle(
|
||||
template_id: str,
|
||||
clip_id: str,
|
||||
subtitle_id: str,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
):
|
||||
"""删除片段字幕"""
|
||||
_, plan_svc = services
|
||||
clip = plan_svc.get_clip(clip_id)
|
||||
if not clip:
|
||||
raise HTTPException(status_code=404, detail="片段不存在")
|
||||
|
||||
config = dict(clip.config) if clip.config else {}
|
||||
subtitles = config.get("subtitles", [])
|
||||
if not isinstance(subtitles, list):
|
||||
subtitles = []
|
||||
|
||||
new_subtitles = [s for s in subtitles if s.get("id") != subtitle_id]
|
||||
if len(new_subtitles) == len(subtitles):
|
||||
raise HTTPException(status_code=404, detail="字幕不存在")
|
||||
|
||||
config["subtitles"] = new_subtitles
|
||||
plan_svc.update_clip(clip_id, config=config)
|
||||
return None
|
||||
|
||||
|
||||
@router.put("/clips/{clip_id}/subtitles", response_model=list[dict[str, Any]])
|
||||
def batch_update_editor_clip_subtitles(
|
||||
template_id: str,
|
||||
clip_id: str,
|
||||
body: list[dict[str, Any]],
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> list[dict[str, Any]]:
|
||||
"""批量更新片段字幕(全量替换)"""
|
||||
_, plan_svc = services
|
||||
clip = plan_svc.get_clip(clip_id)
|
||||
if not clip:
|
||||
raise HTTPException(status_code=404, detail="片段不存在")
|
||||
|
||||
config = dict(clip.config) if clip.config else {}
|
||||
config["subtitles"] = body
|
||||
plan_svc.update_clip(clip_id, config=config)
|
||||
return body
|
||||
@@ -1,61 +0,0 @@
|
||||
"""时间线路由.
|
||||
|
||||
端点:
|
||||
- GET /timeline 时间线场景数据
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.services.edit_plan_service import EditPlanService
|
||||
from app.services.edit_template_service import EditTemplateService
|
||||
from fastapi import APIRouter, Depends
|
||||
|
||||
from ._utils import _CLIP_TYPE_COLORS, _DEFAULT_COLOR, _clip_type_to_scene_label, _format_time
|
||||
from .dependencies import get_draft_plan_id, get_editor_services
|
||||
from .schemas import EditorTimelineResponse, EditorTimelineSceneResponse
|
||||
|
||||
router = APIRouter(tags=["Template Editor"])
|
||||
|
||||
|
||||
@router.get("/timeline", response_model=EditorTimelineResponse)
|
||||
def get_editor_timeline(
|
||||
template_id: str,
|
||||
plan_id: str = Depends(get_draft_plan_id),
|
||||
services: tuple[EditTemplateService, EditPlanService] = Depends(get_editor_services),
|
||||
_: AuthenticatedUser = Depends(get_current_user),
|
||||
) -> EditorTimelineResponse:
|
||||
"""获取草稿的时间线场景数据"""
|
||||
_, plan_svc = services
|
||||
plan = plan_svc.get_plan_or_raise(plan_id)
|
||||
clips = plan_svc.list_clips(plan_id=plan_id, skip=0, limit=200)
|
||||
clips.sort(key=lambda c: c.order)
|
||||
|
||||
scenes = []
|
||||
current_time = 0.0
|
||||
|
||||
for clip in clips:
|
||||
start = current_time
|
||||
end = start + clip.duration
|
||||
color = _CLIP_TYPE_COLORS.get(clip.clip_type, _DEFAULT_COLOR)
|
||||
scene_label = _clip_type_to_scene_label(clip.clip_type, clip.text_content)
|
||||
|
||||
scenes.append(
|
||||
EditorTimelineSceneResponse(
|
||||
scene=scene_label,
|
||||
time=f"{_format_time(start)} - {_format_time(end)}",
|
||||
duration=clip.duration,
|
||||
color=color,
|
||||
clip_id=clip.id,
|
||||
clip_type=clip.clip_type,
|
||||
)
|
||||
)
|
||||
current_time = end
|
||||
|
||||
total_duration = sum(s.duration for s in scenes) or plan.total_duration
|
||||
|
||||
return EditorTimelineResponse(
|
||||
plan_id=plan_id,
|
||||
total_duration=total_duration,
|
||||
scenes=scenes,
|
||||
)
|
||||
@@ -6,6 +6,8 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Dict, List
|
||||
|
||||
from packages.shared.ai_service import (
|
||||
_call_ai_cover_service,
|
||||
_call_ai_recommend_service,
|
||||
|
||||
Regular → Executable
+8
-11
@@ -4,7 +4,6 @@ Redis Session 存储
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional
|
||||
|
||||
@@ -13,8 +12,6 @@ from redis import Redis
|
||||
|
||||
from packages.domain.auth.session_store import SessionStorePort
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class RedisConfig:
|
||||
"""Redis 配置"""
|
||||
@@ -150,7 +147,7 @@ class SessionStore(SessionStorePort):
|
||||
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis session save failed: {e}")
|
||||
print(f"Failed to save session: {e}")
|
||||
return False
|
||||
|
||||
def get_session(self, session_id: str) -> Optional[dict]:
|
||||
@@ -171,7 +168,7 @@ class SessionStore(SessionStorePort):
|
||||
return json.loads(data)
|
||||
return None
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis session get failed: {e}")
|
||||
print(f"Failed to get session: {e}")
|
||||
return None
|
||||
|
||||
def get_session_by_refresh_token(self, refresh_token: str) -> Optional[dict]:
|
||||
@@ -195,7 +192,7 @@ class SessionStore(SessionStorePort):
|
||||
# 再获取完整的 session 数据
|
||||
return self.get_session(session_id)
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis session get by refresh failed: {e}")
|
||||
print(f"Failed to get session by refresh_token: {e}")
|
||||
return None
|
||||
|
||||
def get_refresh_token(self, session_id: str) -> Optional[str]:
|
||||
@@ -212,7 +209,7 @@ class SessionStore(SessionStorePort):
|
||||
refresh_token_key = self._refresh_token_key(session_id)
|
||||
return self.redis.get(refresh_token_key)
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis refresh token get failed: {e}")
|
||||
print(f"Failed to get refresh_token: {e}")
|
||||
return None
|
||||
|
||||
def update_last_active(self, session_id: str) -> bool:
|
||||
@@ -241,7 +238,7 @@ class SessionStore(SessionStorePort):
|
||||
|
||||
return False
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis session update active failed: {e}")
|
||||
print(f"Failed to update last active: {e}")
|
||||
return False
|
||||
|
||||
def delete_session(self, session_id: str) -> bool:
|
||||
@@ -281,7 +278,7 @@ class SessionStore(SessionStorePort):
|
||||
|
||||
return True
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis session delete failed: {e}")
|
||||
print(f"Failed to delete session: {e}")
|
||||
return False
|
||||
|
||||
def get_user_sessions(self, user_id: str) -> list[dict]:
|
||||
@@ -306,7 +303,7 @@ class SessionStore(SessionStorePort):
|
||||
|
||||
return sessions
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis user sessions list failed: {e}")
|
||||
print(f"Failed to get user sessions: {e}")
|
||||
return []
|
||||
|
||||
def delete_all_user_sessions(self, user_id: str) -> int:
|
||||
@@ -333,7 +330,7 @@ class SessionStore(SessionStorePort):
|
||||
|
||||
return count
|
||||
except Exception as e:
|
||||
logger.warning(f"Redis user sessions delete failed: {e}")
|
||||
print(f"Failed to delete all user sessions: {e}")
|
||||
return 0
|
||||
|
||||
def session_exists(self, session_id: str) -> bool:
|
||||
|
||||
Regular → Executable
+3
-3
@@ -1,6 +1,6 @@
|
||||
"""JWT Token 生成、验证、解析服务"""
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from datetime import datetime, timedelta
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
import jwt
|
||||
@@ -88,7 +88,7 @@ class JWTService(JWTServicePort):
|
||||
Returns:
|
||||
JWT Token 字符串
|
||||
"""
|
||||
now = datetime.now(timezone.utc)
|
||||
now = datetime.utcnow()
|
||||
expire = now + timedelta(minutes=self.config.ACCESS_TOKEN_EXPIRE_MINUTES) # noqa: E501
|
||||
|
||||
payload = {
|
||||
@@ -115,7 +115,7 @@ class JWTService(JWTServicePort):
|
||||
Returns:
|
||||
JWT Token 字符串
|
||||
"""
|
||||
now = datetime.now(timezone.utc)
|
||||
now = datetime.utcnow()
|
||||
expire = now + timedelta(days=self.config.REFRESH_TOKEN_EXPIRE_DAYS)
|
||||
|
||||
payload = {
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
使用 bcrypt 安全存储密码
|
||||
"""
|
||||
|
||||
from typing import Optional
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import bcrypt
|
||||
|
||||
|
||||
Regular → Executable
+2
-5
@@ -2,7 +2,6 @@
|
||||
密码重置 Use Case
|
||||
"""
|
||||
|
||||
import logging
|
||||
import secrets
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional
|
||||
@@ -10,8 +9,6 @@ from typing import Optional
|
||||
from packages.adapters.smtp import get_email_service
|
||||
from packages.application.auth.password_hasher import password_hasher, password_validator
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class RequestPasswordResetRequest:
|
||||
"""请求密码重置"""
|
||||
@@ -85,10 +82,10 @@ class RequestPasswordResetUseCase:
|
||||
)
|
||||
|
||||
if not success:
|
||||
logger.warning(f"Password reset email failed: {error}")
|
||||
print(f"Failed to send password reset email: {error}")
|
||||
# 不返回错误,避免暴露用户存在性
|
||||
except Exception as e:
|
||||
logger.error(f"Email service error: {e}", exc_info=True)
|
||||
print(f"Email service error: {e}")
|
||||
|
||||
return True, None
|
||||
|
||||
|
||||
Regular → Executable
+2
-5
@@ -2,7 +2,6 @@
|
||||
用户注册 Use Case
|
||||
"""
|
||||
|
||||
import logging
|
||||
import secrets
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
@@ -12,8 +11,6 @@ from packages.adapters.smtp import get_email_service
|
||||
from packages.application.auth.password_hasher import password_hasher, password_validator
|
||||
from packages.domain.entities import User
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class RegisterUserRequest:
|
||||
"""注册请求"""
|
||||
@@ -139,9 +136,9 @@ class RegisterUserUseCase:
|
||||
email_sent = success
|
||||
|
||||
if not success:
|
||||
logger.warning(f"Verification email failed: {error}")
|
||||
print(f"Failed to send verification email: {error}")
|
||||
except Exception as e:
|
||||
logger.error(f"Email service error: {e}", exc_info=True)
|
||||
print(f"Email service error: {e}")
|
||||
|
||||
# 10. 返回响应(即使邮件发送失败,用户也已创建)
|
||||
return (
|
||||
|
||||
@@ -11,11 +11,14 @@ from packages.application.recipe.commands import (
|
||||
CreateRecipeCommand,
|
||||
UpdateRecipeCommand,
|
||||
)
|
||||
from packages.domain.exceptions import NotFoundError
|
||||
from packages.domain.recipe import Recipe, RecipeItem
|
||||
from packages.infrastructure.feature_flags import FeatureScope, feature_flags
|
||||
|
||||
|
||||
class NotFoundError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class FeatureDisabledError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
@@ -15,10 +15,20 @@ from packages.application.template.commands import (
|
||||
ValidateTemplateCommand,
|
||||
)
|
||||
from packages.domain.editing_mode import EditingMode
|
||||
from packages.domain.exceptions import NotFoundError, ValidationError
|
||||
from packages.domain.template import Template, TemplateCategory, TemplateSegment
|
||||
from packages.ports.template_repository import TemplateRepositoryPort
|
||||
|
||||
|
||||
class NotFoundError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class ValidationError(Exception):
|
||||
"""业务规则校验失败."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
VALID_MODES = {m.value for m in EditingMode}
|
||||
VALID_MATERIAL_TYPES = {"人物", "场景"}
|
||||
|
||||
|
||||
@@ -12,7 +12,6 @@ from packages.application.title_library.commands import (
|
||||
PickTitleCommand,
|
||||
UpdateTitleLibraryCommand,
|
||||
)
|
||||
from packages.domain.exceptions import NotFoundError, QuotaExceededError
|
||||
from packages.domain.quota import QuotaDimension, quota_checker
|
||||
from packages.domain.title_library import TitleLibraryItem
|
||||
|
||||
@@ -163,3 +162,15 @@ class PickTitleUseCase:
|
||||
|
||||
# 随机选一个
|
||||
return random.choice(pool)
|
||||
|
||||
|
||||
class QuotaExceededError(Exception):
|
||||
def __init__(self, dimension: str, limit: float, used: float) -> None:
|
||||
self.dimension = dimension
|
||||
self.limit = limit
|
||||
self.used = used
|
||||
super().__init__(f"Quota exceeded for {dimension}: {used}/{limit}")
|
||||
|
||||
|
||||
class NotFoundError(Exception):
|
||||
pass
|
||||
|
||||
@@ -1,7 +0,0 @@
|
||||
"""TTS Job 模块公共异常定义。"""
|
||||
|
||||
|
||||
class TTSJobNotFoundError(Exception):
|
||||
"""TTS 任务未找到。"""
|
||||
|
||||
pass
|
||||
@@ -4,11 +4,16 @@ from __future__ import annotations
|
||||
|
||||
from typing import List, Optional
|
||||
|
||||
from packages.application.tts_job.exceptions import TTSJobNotFoundError
|
||||
from packages.domain.tts_job import TTSJob
|
||||
from packages.ports.tts_job_repository import TTSJobRepository
|
||||
|
||||
|
||||
class TTSJobNotFoundError(Exception):
|
||||
"""TTS 任务未找到。"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class CreateTTSJobUseCase:
|
||||
"""创建 TTS 合成任务。"""
|
||||
|
||||
|
||||
@@ -19,7 +19,6 @@ from typing import Optional
|
||||
|
||||
from packages.application.cosyvoice_service import CosyVoiceAuthError, CosyVoiceError, CosyVoiceService
|
||||
from packages.application.tts_job.audio_merger import AudioMerger
|
||||
from packages.application.tts_job.exceptions import TTSJobNotFoundError
|
||||
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
|
||||
@@ -44,6 +43,12 @@ class TTSWorkflowError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class TTSJobNotFoundError(Exception):
|
||||
"""TTS 任务未找到。"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class TTSWorkflowService:
|
||||
"""TTS 合成工作流编排服务。
|
||||
|
||||
|
||||
@@ -10,13 +10,18 @@ from packages.application.video_share.commands import (
|
||||
CreateShareCommand,
|
||||
UpdateShareCommand,
|
||||
)
|
||||
from packages.domain.exceptions import NotFoundError
|
||||
from packages.domain.generated_video import GeneratedVideo
|
||||
from packages.domain.video_share import VideoShare
|
||||
from packages.ports.generated_video_repository import GeneratedVideoRepository
|
||||
from packages.ports.video_share_repository import VideoShareRepositoryPort
|
||||
|
||||
|
||||
class NotFoundError(Exception):
|
||||
"""分享记录不存在."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class VideoNotFoundError(Exception):
|
||||
"""视频不存在."""
|
||||
|
||||
|
||||
@@ -10,7 +10,6 @@ from packages.application.voice_library.commands import (
|
||||
CreateVoiceLibraryCommand,
|
||||
UpdateVoiceLibraryCommand,
|
||||
)
|
||||
from packages.domain.exceptions import NotFoundError, QuotaExceededError
|
||||
from packages.domain.quota import QuotaDimension, quota_checker
|
||||
from packages.domain.voice_library import VoiceLibraryItem
|
||||
|
||||
@@ -118,3 +117,15 @@ class DeleteVoiceLibraryUseCase:
|
||||
|
||||
def execute(self, voice_id: str, user_id: str) -> bool:
|
||||
return self.repository.delete(voice_id, user_id)
|
||||
|
||||
|
||||
class QuotaExceededError(Exception):
|
||||
def __init__(self, dimension: str, limit: float, used: float) -> None:
|
||||
self.dimension = dimension
|
||||
self.limit = limit
|
||||
self.used = used
|
||||
super().__init__(f"Quota exceeded for {dimension}: {used}/{limit}")
|
||||
|
||||
|
||||
class NotFoundError(Exception):
|
||||
pass
|
||||
|
||||
@@ -19,7 +19,6 @@ from uuid import uuid4
|
||||
class AssetLibraryKind(StrEnum):
|
||||
VIDEO = "video"
|
||||
VOICE = "voice"
|
||||
IMAGE = "image"
|
||||
|
||||
|
||||
class IngestJobStatus(StrEnum):
|
||||
@@ -35,27 +34,6 @@ class ClassificationJobStatus(StrEnum):
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
|
||||
@classmethod
|
||||
def _missing_(cls, value: object) -> "ClassificationJobStatus":
|
||||
"""兼容历史数据,避免枚举转换失败导致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
|
||||
|
||||
|
||||
# 向后兼容别名
|
||||
ClassificationStatus = ClassificationJobStatus
|
||||
|
||||
|
||||
class AssetClassification(StrEnum):
|
||||
"""Asset classification categories."""
|
||||
|
||||
@@ -16,12 +16,18 @@ else:
|
||||
from typing import Any
|
||||
from uuid import uuid4
|
||||
|
||||
# 枚举统一从 classification 模块导入,消除重复定义
|
||||
from packages.domain.classification import (
|
||||
AssetLibraryKind,
|
||||
ClassificationStatus,
|
||||
IngestJobStatus,
|
||||
)
|
||||
|
||||
class AssetLibraryKind(StrEnum):
|
||||
VIDEO = "video"
|
||||
VOICE = "voice"
|
||||
IMAGE = "image"
|
||||
|
||||
|
||||
class IngestJobStatus(StrEnum):
|
||||
PENDING = "pending"
|
||||
PROCESSING = "processing"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
@@ -154,6 +160,30 @@ class AssetStatus(StrEnum):
|
||||
return cls.READY
|
||||
|
||||
|
||||
class ClassificationStatus(StrEnum):
|
||||
PENDING = "pending"
|
||||
PROCESSING = "processing"
|
||||
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:
|
||||
id: str
|
||||
|
||||
@@ -1,35 +0,0 @@
|
||||
"""
|
||||
领域层通用异常定义。
|
||||
|
||||
所有应用层共享的通用异常在此统一定义,
|
||||
消除各 use_cases 模块中重复的异常类。
|
||||
领域特定的异常仍在各模块内定义。
|
||||
"""
|
||||
|
||||
|
||||
class DomainError(Exception):
|
||||
"""领域层异常基类。"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class NotFoundError(DomainError):
|
||||
"""资源不存在。"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class ValidationError(DomainError):
|
||||
"""业务规则校验失败。"""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class QuotaExceededError(DomainError):
|
||||
"""配额超限。"""
|
||||
|
||||
def __init__(self, dimension: str, limit: float, used: float) -> None:
|
||||
self.dimension = dimension
|
||||
self.limit = limit
|
||||
self.used = used
|
||||
super().__init__(f"Quota exceeded for {dimension}: {used}/{limit}")
|
||||
@@ -1,157 +0,0 @@
|
||||
"""素材库 UseCase 单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.asset_libraries import (
|
||||
CreateAssetLibraryCommand,
|
||||
CreateAssetLibraryUseCase,
|
||||
ListAssetLibrariesUseCase,
|
||||
)
|
||||
from packages.domain import AssetLibrary, AssetLibraryKind
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_repo():
|
||||
repo = MagicMock()
|
||||
return repo
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_library():
|
||||
lib = AssetLibrary.create(
|
||||
project_id="proj_456",
|
||||
name="测试视频库",
|
||||
kind=AssetLibraryKind.VIDEO,
|
||||
)
|
||||
lib.id = "lib_123"
|
||||
return lib
|
||||
|
||||
|
||||
class TestListAssetLibrariesUseCase:
|
||||
"""ListAssetLibrariesUseCase 测试"""
|
||||
|
||||
def test_list_returns_repo_results(self, mock_repo, sample_library):
|
||||
"""正常返回 repository 的查询结果"""
|
||||
mock_repo.find_by_project.return_value = [sample_library]
|
||||
use_case = ListAssetLibrariesUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("proj_456")
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0].id == "lib_123"
|
||||
mock_repo.find_by_project.assert_called_once_with("proj_456")
|
||||
|
||||
def test_empty_project_raises_value_error(self, mock_repo):
|
||||
"""空 project_id 抛出 ValueError"""
|
||||
use_case = ListAssetLibrariesUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
use_case.execute("")
|
||||
|
||||
mock_repo.find_by_project.assert_not_called()
|
||||
|
||||
def test_whitespace_project_raises_value_error(self, mock_repo):
|
||||
"""纯空格 project_id 也抛出 ValueError"""
|
||||
use_case = ListAssetLibrariesUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(ValueError, match="project_id 不能为空"):
|
||||
use_case.execute(" ")
|
||||
|
||||
mock_repo.find_by_project.assert_not_called()
|
||||
|
||||
def test_project_id_stripped_before_query(self, mock_repo, sample_library):
|
||||
"""project_id 会被 strip 后再查询"""
|
||||
mock_repo.find_by_project.return_value = [sample_library]
|
||||
use_case = ListAssetLibrariesUseCase(mock_repo)
|
||||
|
||||
use_case.execute(" proj_456 ")
|
||||
|
||||
mock_repo.find_by_project.assert_called_once_with("proj_456")
|
||||
|
||||
def test_empty_list(self, mock_repo):
|
||||
"""项目没有素材库时返回空列表"""
|
||||
mock_repo.find_by_project.return_value = []
|
||||
use_case = ListAssetLibrariesUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("proj_456")
|
||||
|
||||
assert result == []
|
||||
|
||||
|
||||
class TestCreateAssetLibraryUseCase:
|
||||
"""CreateAssetLibraryUseCase 测试"""
|
||||
|
||||
def test_create_success(self, mock_repo, sample_library):
|
||||
"""创建成功返回 AssetLibrary"""
|
||||
mock_repo.create.return_value = sample_library
|
||||
use_case = CreateAssetLibraryUseCase(mock_repo)
|
||||
|
||||
command = CreateAssetLibraryCommand(
|
||||
project_id="proj_456",
|
||||
name="新素材库",
|
||||
kind=AssetLibraryKind.IMAGE,
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.id == "lib_123"
|
||||
mock_repo.create.assert_called_once()
|
||||
# 验证传入 repository 的是一个 AssetLibrary 对象
|
||||
created = mock_repo.create.call_args[0][0]
|
||||
assert isinstance(created, AssetLibrary)
|
||||
assert created.project_id == "proj_456"
|
||||
assert created.name == "新素材库"
|
||||
assert created.kind == AssetLibraryKind.IMAGE
|
||||
|
||||
def test_create_with_video_kind(self, mock_repo):
|
||||
"""创建视频类型素材库"""
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateAssetLibraryUseCase(mock_repo)
|
||||
|
||||
command = CreateAssetLibraryCommand(
|
||||
project_id="proj_1",
|
||||
name="视频库",
|
||||
kind=AssetLibraryKind.VIDEO,
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.kind == AssetLibraryKind.VIDEO
|
||||
assert result.name == "视频库"
|
||||
|
||||
def test_create_with_voice_kind(self, mock_repo):
|
||||
"""创建音色类型素材库"""
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateAssetLibraryUseCase(mock_repo)
|
||||
|
||||
command = CreateAssetLibraryCommand(
|
||||
project_id="proj_1",
|
||||
name="音色库",
|
||||
kind=AssetLibraryKind.VOICE,
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.kind == AssetLibraryKind.VOICE
|
||||
|
||||
|
||||
class TestCreateAssetLibraryCommand:
|
||||
"""CreateAssetLibraryCommand 数据类测试"""
|
||||
|
||||
def test_command_fields(self):
|
||||
"""命令对象字段正确"""
|
||||
cmd = CreateAssetLibraryCommand(
|
||||
project_id="proj_1",
|
||||
name="test",
|
||||
kind=AssetLibraryKind.VIDEO,
|
||||
)
|
||||
assert cmd.project_id == "proj_1"
|
||||
assert cmd.name == "test"
|
||||
assert cmd.kind == AssetLibraryKind.VIDEO
|
||||
|
||||
def test_command_is_dataclass(self):
|
||||
"""命令是 dataclass"""
|
||||
from dataclasses import is_dataclass
|
||||
|
||||
assert is_dataclass(CreateAssetLibraryCommand)
|
||||
@@ -1,105 +0,0 @@
|
||||
"""BGM 配置工具函数单元测试."""
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.bgm_utils import merge_bgm_config
|
||||
|
||||
|
||||
class TestMergeBgmConfig:
|
||||
"""merge_bgm_config 测试"""
|
||||
|
||||
def test_user_bgm_empty_returns_template_copy(self):
|
||||
"""用户配置为空时,返回模板配置的拷贝"""
|
||||
template = {"enabled": True, "volume": 0.5, "asset_id": "tpl_123"}
|
||||
result = merge_bgm_config(template, {})
|
||||
assert result == template
|
||||
assert result is not template
|
||||
|
||||
def test_user_bgm_none_returns_template_copy(self):
|
||||
"""用户配置为 None 时,返回模板配置的拷贝"""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
result = merge_bgm_config(template, None) # type: ignore
|
||||
assert result == template
|
||||
|
||||
def test_template_bgm_empty_returns_user_copy(self):
|
||||
"""模板配置为空时,返回用户配置的拷贝"""
|
||||
user = {"enabled": False, "volume": 0.8, "asset_id": "user_456"}
|
||||
result = merge_bgm_config({}, user)
|
||||
assert result == user
|
||||
assert result is not user
|
||||
|
||||
def test_template_bgm_none_returns_user_copy(self):
|
||||
"""模板配置为 None 时,返回用户配置的拷贝"""
|
||||
user = {"enabled": False, "volume": 0.8}
|
||||
result = merge_bgm_config(None, user) # type: ignore
|
||||
assert result == user
|
||||
|
||||
def test_user_fields_override_template(self):
|
||||
"""用户显式指定的字段覆盖模板对应字段"""
|
||||
template = {
|
||||
"enabled": True,
|
||||
"volume": 0.5,
|
||||
"asset_id": "tpl_123",
|
||||
"fade_in": 1.0,
|
||||
}
|
||||
user = {
|
||||
"volume": 0.8,
|
||||
"asset_id": "user_456",
|
||||
}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["volume"] == 0.8
|
||||
assert result["asset_id"] == "user_456"
|
||||
assert result["fade_in"] == 1.0 # 模板值保留
|
||||
|
||||
def test_enabled_not_in_user_preserves_template_enabled(self):
|
||||
"""enabled 特殊处理:用户没传 enabled 时保留模板的 enabled 值"""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
user = {"volume": 0.8} # 没传 enabled
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["enabled"] is True # 保留模板的
|
||||
assert result["volume"] == 0.8 # 用户指定的覆盖
|
||||
|
||||
def test_enabled_in_user_overrides_template(self):
|
||||
"""用户传了 enabled 时覆盖模板的 enabled"""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
user = {"enabled": False, "volume": 0.8}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["enabled"] is False
|
||||
assert result["volume"] == 0.8
|
||||
|
||||
def test_user_adds_new_fields(self):
|
||||
"""用户配置中的新字段会被添加到结果中"""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
user = {"sidechain_enabled": True, "sidechain_ratio": 0.6}
|
||||
result = merge_bgm_config(template, user)
|
||||
assert result["enabled"] is True
|
||||
assert result["volume"] == 0.5
|
||||
assert result["sidechain_enabled"] is True
|
||||
assert result["sidechain_ratio"] == 0.6
|
||||
|
||||
def test_both_empty_returns_empty_dict(self):
|
||||
"""两者都为空时返回空字典"""
|
||||
result = merge_bgm_config({}, {})
|
||||
assert result == {}
|
||||
|
||||
def test_nested_dict_shallow_merge(self):
|
||||
"""嵌套字典是浅合并(当前设计)"""
|
||||
template = {"enabled": True, "config": {"eq": True, "compression": False}}
|
||||
user = {"config": {"compression": True, "reverb": 0.5}}
|
||||
result = merge_bgm_config(template, user)
|
||||
# 浅合并:整个 config 被用户值覆盖
|
||||
assert result["config"] == {"compression": True, "reverb": 0.5}
|
||||
|
||||
def test_does_not_mutate_template(self):
|
||||
"""不修改原始模板配置"""
|
||||
template = {"enabled": True, "volume": 0.5}
|
||||
original = dict(template)
|
||||
merge_bgm_config(template, {"volume": 0.8})
|
||||
assert template == original
|
||||
|
||||
def test_does_not_mutate_user(self):
|
||||
"""不修改原始用户配置"""
|
||||
user = {"volume": 0.8}
|
||||
original = dict(user)
|
||||
merge_bgm_config({"enabled": True}, user)
|
||||
assert user == original
|
||||
@@ -1,131 +0,0 @@
|
||||
"""领域层通用异常单元测试."""
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.domain.exceptions import (
|
||||
DomainError,
|
||||
NotFoundError,
|
||||
QuotaExceededError,
|
||||
ValidationError,
|
||||
)
|
||||
|
||||
|
||||
class TestDomainError:
|
||||
"""DomainError 基类测试"""
|
||||
|
||||
def test_is_exception(self):
|
||||
"""DomainError 是 Exception 的子类"""
|
||||
assert issubclass(DomainError, Exception)
|
||||
|
||||
def test_can_raise_and_catch(self):
|
||||
"""可以抛出和捕获"""
|
||||
with pytest.raises(DomainError):
|
||||
raise DomainError("something went wrong")
|
||||
|
||||
def test_message(self):
|
||||
"""异常消息正确"""
|
||||
err = DomainError("test message")
|
||||
assert str(err) == "test message"
|
||||
|
||||
|
||||
class TestNotFoundError:
|
||||
"""NotFoundError 测试"""
|
||||
|
||||
def test_is_domain_error(self):
|
||||
"""NotFoundError 继承自 DomainError"""
|
||||
assert issubclass(NotFoundError, DomainError)
|
||||
|
||||
def test_can_raise_as_domain_error(self):
|
||||
"""可以作为 DomainError 捕获"""
|
||||
with pytest.raises(DomainError):
|
||||
raise NotFoundError("resource not found")
|
||||
|
||||
def test_default_message(self):
|
||||
"""无参构造"""
|
||||
err = NotFoundError()
|
||||
assert isinstance(err, NotFoundError)
|
||||
|
||||
def test_custom_message(self):
|
||||
"""自定义消息"""
|
||||
err = NotFoundError("user 123 not found")
|
||||
assert str(err) == "user 123 not found"
|
||||
|
||||
|
||||
class TestValidationError:
|
||||
"""ValidationError 测试"""
|
||||
|
||||
def test_is_domain_error(self):
|
||||
"""ValidationError 继承自 DomainError"""
|
||||
assert issubclass(ValidationError, DomainError)
|
||||
|
||||
def test_can_raise_as_domain_error(self):
|
||||
"""可以作为 DomainError 捕获"""
|
||||
with pytest.raises(DomainError):
|
||||
raise ValidationError("invalid input")
|
||||
|
||||
def test_custom_message(self):
|
||||
"""自定义消息"""
|
||||
err = ValidationError("duration must be positive")
|
||||
assert str(err) == "duration must be positive"
|
||||
|
||||
|
||||
class TestQuotaExceededError:
|
||||
"""QuotaExceededError 测试"""
|
||||
|
||||
def test_is_domain_error(self):
|
||||
"""QuotaExceededError 继承自 DomainError"""
|
||||
assert issubclass(QuotaExceededError, DomainError)
|
||||
|
||||
def test_can_raise_as_domain_error(self):
|
||||
"""可以作为 DomainError 捕获"""
|
||||
with pytest.raises(DomainError):
|
||||
raise QuotaExceededError("storage", 1024.0, 2048.0)
|
||||
|
||||
def test_stores_dimension_limit_used(self):
|
||||
"""保存 dimension、limit、used 属性"""
|
||||
err = QuotaExceededError("storage_mb", 1024.0, 1500.0)
|
||||
assert err.dimension == "storage_mb"
|
||||
assert err.limit == 1024.0
|
||||
assert err.used == 1500.0
|
||||
|
||||
def test_error_message_format(self):
|
||||
"""异常消息格式正确"""
|
||||
err = QuotaExceededError("storage_mb", 1024.0, 1500.0)
|
||||
msg = str(err)
|
||||
assert "storage_mb" in msg
|
||||
assert "1500.0" in msg
|
||||
assert "1024.0" in msg
|
||||
assert "Quota exceeded" in msg
|
||||
|
||||
def test_integer_values(self):
|
||||
"""整数值也能正常工作"""
|
||||
err = QuotaExceededError("projects", 10, 15)
|
||||
assert err.dimension == "projects"
|
||||
assert err.limit == 10
|
||||
assert err.used == 15
|
||||
|
||||
def test_zero_limit(self):
|
||||
"""限制为 0 时也能正常工作"""
|
||||
err = QuotaExceededError("custom_templates", 0, 1)
|
||||
assert err.limit == 0
|
||||
assert err.used == 1
|
||||
|
||||
|
||||
class TestExceptionHierarchy:
|
||||
"""异常继承关系测试"""
|
||||
|
||||
def test_all_are_domain_errors(self):
|
||||
"""所有异常都可以作为 DomainError 捕获"""
|
||||
errors = [
|
||||
NotFoundError(),
|
||||
ValidationError("bad"),
|
||||
QuotaExceededError("x", 10.0, 20.0),
|
||||
]
|
||||
for err in errors:
|
||||
assert isinstance(err, DomainError)
|
||||
|
||||
def test_distinct_types(self):
|
||||
"""不同异常类型可以区分"""
|
||||
assert not issubclass(NotFoundError, ValidationError)
|
||||
assert not issubclass(ValidationError, QuotaExceededError)
|
||||
assert not issubclass(NotFoundError, QuotaExceededError)
|
||||
@@ -1,270 +0,0 @@
|
||||
"""TTS Job UseCase 单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.tts_job.exceptions import TTSJobNotFoundError
|
||||
from packages.application.tts_job.use_cases import (
|
||||
CreateTTSJobUseCase,
|
||||
DeleteTTSJobUseCase,
|
||||
GetTTSJobStatusUseCase,
|
||||
GetTTSJobUseCase,
|
||||
ListTTSJobsUseCase,
|
||||
)
|
||||
from packages.domain.tts_job import TTSJob
|
||||
|
||||
|
||||
def _make_job(
|
||||
id: str = "job1",
|
||||
user_id: str = "user_1",
|
||||
status: str = "pending",
|
||||
) -> TTSJob:
|
||||
j = TTSJob.create(
|
||||
user_id=user_id,
|
||||
input_text="你好世界",
|
||||
voice_id="voice_001",
|
||||
voice_model="cosyvoice",
|
||||
project_id="proj_1",
|
||||
voice_clone_profile_id="",
|
||||
sample_rate=22050,
|
||||
format="mp3",
|
||||
max_retries=3,
|
||||
)
|
||||
j.id = id
|
||||
object.__setattr__(j, "status", status)
|
||||
return j
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_repo():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_job():
|
||||
return _make_job()
|
||||
|
||||
|
||||
class TestCreateTTSJobUseCase:
|
||||
"""CreateTTSJobUseCase 测试"""
|
||||
|
||||
def test_create_success(self, mock_repo, sample_job):
|
||||
"""创建成功"""
|
||||
mock_repo.create.return_value = sample_job
|
||||
use_case = CreateTTSJobUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute(
|
||||
"user_1",
|
||||
"你好世界",
|
||||
voice_id="voice_001",
|
||||
voice_model="cosyvoice",
|
||||
project_id="proj_1",
|
||||
)
|
||||
|
||||
assert result.id == "job1"
|
||||
assert result.status == "pending"
|
||||
mock_repo.create.assert_called_once()
|
||||
created = mock_repo.create.call_args[0][0]
|
||||
assert isinstance(created, TTSJob)
|
||||
assert created.input_text == "你好世界"
|
||||
|
||||
def test_create_with_defaults(self, mock_repo):
|
||||
"""使用默认参数创建"""
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateTTSJobUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("user_1", "测试文本")
|
||||
|
||||
assert result.user_id == "user_1"
|
||||
assert result.input_text == "测试文本"
|
||||
assert result.sample_rate == 22050
|
||||
assert result.format == "mp3"
|
||||
assert result.max_retries == 3
|
||||
|
||||
def test_create_with_clone_profile(self, mock_repo):
|
||||
"""创建时带音色克隆ID"""
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateTTSJobUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute(
|
||||
"user_1", "文本", voice_clone_profile_id="vc_123"
|
||||
)
|
||||
|
||||
assert result.voice_clone_profile_id == "vc_123"
|
||||
|
||||
def test_create_with_metadata(self, mock_repo):
|
||||
"""创建时带metadata"""
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateTTSJobUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute(
|
||||
"user_1", "text", metadata={"source": "api", "priority": "high"}
|
||||
)
|
||||
|
||||
assert result.metadata == {"source": "api", "priority": "high"}
|
||||
|
||||
|
||||
class TestListTTSJobsUseCase:
|
||||
"""ListTTSJobsUseCase 测试"""
|
||||
|
||||
def test_list_returns_items_and_total(self, mock_repo, sample_job):
|
||||
"""返回 (items, total) 元组"""
|
||||
mock_repo.list_by_user.return_value = [sample_job]
|
||||
mock_repo.count_by_user.return_value = 1
|
||||
use_case = ListTTSJobsUseCase(mock_repo)
|
||||
|
||||
items, total = use_case.execute("user_1")
|
||||
|
||||
assert len(items) == 1
|
||||
assert total == 1
|
||||
mock_repo.list_by_user.assert_called_once_with(
|
||||
"user_1", status=None, limit=50, offset=0
|
||||
)
|
||||
mock_repo.count_by_user.assert_called_once_with("user_1", status=None)
|
||||
|
||||
def test_list_with_status(self, mock_repo, sample_job):
|
||||
"""按状态过滤"""
|
||||
mock_repo.list_by_user.return_value = [sample_job]
|
||||
mock_repo.count_by_user.return_value = 5
|
||||
use_case = ListTTSJobsUseCase(mock_repo)
|
||||
|
||||
use_case.execute("user_1", status="completed")
|
||||
|
||||
mock_repo.list_by_user.assert_called_once_with(
|
||||
"user_1", status="completed", limit=50, offset=0
|
||||
)
|
||||
mock_repo.count_by_user.assert_called_once_with("user_1", status="completed")
|
||||
|
||||
def test_list_with_pagination(self, mock_repo, sample_job):
|
||||
"""带分页参数"""
|
||||
mock_repo.list_by_user.return_value = [sample_job]
|
||||
mock_repo.count_by_user.return_value = 100
|
||||
use_case = ListTTSJobsUseCase(mock_repo)
|
||||
|
||||
use_case.execute("user_1", skip=20, limit=10)
|
||||
|
||||
mock_repo.list_by_user.assert_called_once_with(
|
||||
"user_1", status=None, limit=10, offset=20
|
||||
)
|
||||
|
||||
def test_empty_list(self, mock_repo):
|
||||
"""空列表"""
|
||||
mock_repo.list_by_user.return_value = []
|
||||
mock_repo.count_by_user.return_value = 0
|
||||
use_case = ListTTSJobsUseCase(mock_repo)
|
||||
|
||||
items, total = use_case.execute("user_1")
|
||||
|
||||
assert items == []
|
||||
assert total == 0
|
||||
|
||||
|
||||
class TestGetTTSJobUseCase:
|
||||
"""GetTTSJobUseCase 测试"""
|
||||
|
||||
def test_get_existing(self, mock_repo, sample_job):
|
||||
"""获取存在的任务"""
|
||||
mock_repo.get.return_value = sample_job
|
||||
use_case = GetTTSJobUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("job1", "user_1")
|
||||
|
||||
assert result.id == "job1"
|
||||
mock_repo.get.assert_called_once_with("job1")
|
||||
|
||||
def test_get_nonexistent_raises(self, mock_repo):
|
||||
"""不存在抛出 TTSJobNotFoundError"""
|
||||
mock_repo.get.return_value = None
|
||||
use_case = GetTTSJobUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(TTSJobNotFoundError, match="not found"):
|
||||
use_case.execute("noexist", "user_1")
|
||||
|
||||
def test_get_wrong_user_raises(self, mock_repo, sample_job):
|
||||
"""非本人任务抛出"""
|
||||
sample_job.user_id = "other_user"
|
||||
mock_repo.get.return_value = sample_job
|
||||
use_case = GetTTSJobUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(TTSJobNotFoundError):
|
||||
use_case.execute("job1", "user_1")
|
||||
|
||||
|
||||
class TestGetTTSJobStatusUseCase:
|
||||
"""GetTTSJobStatusUseCase 测试"""
|
||||
|
||||
def test_get_status_pending(self, mock_repo, sample_job):
|
||||
"""获取 pending 状态"""
|
||||
mock_repo.get.return_value = sample_job
|
||||
use_case = GetTTSJobStatusUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("job1", "user_1")
|
||||
|
||||
assert result.status == "pending"
|
||||
|
||||
def test_get_status_completed(self, mock_repo):
|
||||
"""获取 completed 状态"""
|
||||
job = _make_job(status="completed")
|
||||
mock_repo.get.return_value = job
|
||||
use_case = GetTTSJobStatusUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("job1", "user_1")
|
||||
|
||||
assert result.status == "completed"
|
||||
|
||||
def test_get_status_not_found(self, mock_repo):
|
||||
"""不存在抛出"""
|
||||
mock_repo.get.return_value = None
|
||||
use_case = GetTTSJobStatusUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(TTSJobNotFoundError):
|
||||
use_case.execute("noexist", "user_1")
|
||||
|
||||
def test_get_status_wrong_user(self, mock_repo, sample_job):
|
||||
"""非本人抛出"""
|
||||
sample_job.user_id = "other"
|
||||
mock_repo.get.return_value = sample_job
|
||||
use_case = GetTTSJobStatusUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(TTSJobNotFoundError):
|
||||
use_case.execute("job1", "user_1")
|
||||
|
||||
|
||||
class TestDeleteTTSJobUseCase:
|
||||
"""DeleteTTSJobUseCase 测试"""
|
||||
|
||||
def test_delete_success(self, mock_repo, sample_job):
|
||||
"""删除成功"""
|
||||
mock_repo.get.return_value = sample_job
|
||||
mock_repo.delete.return_value = True
|
||||
use_case = DeleteTTSJobUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("job1", "user_1")
|
||||
|
||||
assert result is True
|
||||
mock_repo.get.assert_called_once_with("job1")
|
||||
mock_repo.delete.assert_called_once_with("job1")
|
||||
|
||||
def test_delete_nonexistent_returns_false(self, mock_repo):
|
||||
"""不存在返回 False"""
|
||||
mock_repo.get.return_value = None
|
||||
use_case = DeleteTTSJobUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("noexist", "user_1")
|
||||
|
||||
assert result is False
|
||||
mock_repo.delete.assert_not_called()
|
||||
|
||||
def test_delete_wrong_user_returns_false(self, mock_repo, sample_job):
|
||||
"""非本人返回 False"""
|
||||
sample_job.user_id = "other"
|
||||
mock_repo.get.return_value = sample_job
|
||||
use_case = DeleteTTSJobUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("job1", "user_1")
|
||||
|
||||
assert result is False
|
||||
mock_repo.delete.assert_not_called()
|
||||
@@ -1,300 +0,0 @@
|
||||
"""音色克隆 Voice Clone UseCase 单元测试."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.voice_clone.use_cases import (
|
||||
CreateVoiceCloneUseCase,
|
||||
DeleteVoiceCloneUseCase,
|
||||
GetVoiceCloneStatusUseCase,
|
||||
GetVoiceCloneUseCase,
|
||||
ListVoiceClonesUseCase,
|
||||
RetryVoiceCloneUseCase,
|
||||
VoiceCloneNotFoundError,
|
||||
VoiceCloneNotRetryableError,
|
||||
)
|
||||
from packages.domain.voice_clone_profile import VoiceCloneProfile
|
||||
|
||||
|
||||
def _make_profile(
|
||||
id: str = "vc1",
|
||||
user_id: str = "user_1",
|
||||
name: str = "我的音色",
|
||||
status: str = "completed",
|
||||
) -> VoiceCloneProfile:
|
||||
p = VoiceCloneProfile.create(
|
||||
user_id=user_id,
|
||||
name=name,
|
||||
description="测试音色",
|
||||
source_audio_url="https://oss.example.com/source.wav",
|
||||
voice_model="cosyvoice",
|
||||
language="zh-CN",
|
||||
gender="female",
|
||||
max_retries=3,
|
||||
)
|
||||
p.id = id
|
||||
# 直接设状态绕过状态机
|
||||
object.__setattr__(p, "status", status)
|
||||
return p
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_repo():
|
||||
return MagicMock()
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_profile():
|
||||
return _make_profile()
|
||||
|
||||
|
||||
class TestCreateVoiceCloneUseCase:
|
||||
"""CreateVoiceCloneUseCase 测试"""
|
||||
|
||||
def test_create_success(self, mock_repo, sample_profile):
|
||||
"""创建成功"""
|
||||
mock_repo.create.return_value = sample_profile
|
||||
use_case = CreateVoiceCloneUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute(
|
||||
"user_1",
|
||||
"新音色",
|
||||
description="自定义音色",
|
||||
source_audio_url="https://oss.example.com/src.wav",
|
||||
voice_model="cosyvoice",
|
||||
language="zh-CN",
|
||||
gender="female",
|
||||
)
|
||||
|
||||
assert result.id == "vc1"
|
||||
assert result.name == "我的音色"
|
||||
mock_repo.create.assert_called_once()
|
||||
created = mock_repo.create.call_args[0][0]
|
||||
assert isinstance(created, VoiceCloneProfile)
|
||||
|
||||
def test_create_with_defaults(self, mock_repo):
|
||||
"""使用默认参数创建"""
|
||||
mock_repo.create.side_effect = lambda x: x
|
||||
use_case = CreateVoiceCloneUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("user_1", "极简音色")
|
||||
|
||||
assert result.user_id == "user_1"
|
||||
assert result.name == "极简音色"
|
||||
assert result.language == "zh-CN"
|
||||
assert result.gender == "unknown"
|
||||
assert result.max_retries == 3
|
||||
|
||||
|
||||
class TestListVoiceClonesUseCase:
|
||||
"""ListVoiceClonesUseCase 测试"""
|
||||
|
||||
def test_list_returns_items_and_total(self, mock_repo, sample_profile):
|
||||
"""返回 (items, total) 元组"""
|
||||
mock_repo.list_by_user.return_value = [sample_profile]
|
||||
mock_repo.count_by_user.return_value = 1
|
||||
use_case = ListVoiceClonesUseCase(mock_repo)
|
||||
|
||||
items, total = use_case.execute("user_1")
|
||||
|
||||
assert len(items) == 1
|
||||
assert total == 1
|
||||
mock_repo.list_by_user.assert_called_once_with(
|
||||
"user_1", status=None, limit=50, offset=0
|
||||
)
|
||||
mock_repo.count_by_user.assert_called_once_with("user_1", status=None)
|
||||
|
||||
def test_list_with_status(self, mock_repo, sample_profile):
|
||||
"""按状态过滤"""
|
||||
mock_repo.list_by_user.return_value = [sample_profile]
|
||||
mock_repo.count_by_user.return_value = 1
|
||||
use_case = ListVoiceClonesUseCase(mock_repo)
|
||||
|
||||
use_case.execute("user_1", status="processing")
|
||||
|
||||
mock_repo.list_by_user.assert_called_once_with(
|
||||
"user_1", status="processing", limit=50, offset=0
|
||||
)
|
||||
mock_repo.count_by_user.assert_called_once_with("user_1", status="processing")
|
||||
|
||||
def test_list_with_pagination(self, mock_repo, sample_profile):
|
||||
"""带分页参数"""
|
||||
mock_repo.list_by_user.return_value = [sample_profile]
|
||||
mock_repo.count_by_user.return_value = 10
|
||||
use_case = ListVoiceClonesUseCase(mock_repo)
|
||||
|
||||
use_case.execute("user_1", skip=5, limit=10)
|
||||
|
||||
mock_repo.list_by_user.assert_called_once_with(
|
||||
"user_1", status=None, limit=10, offset=5
|
||||
)
|
||||
|
||||
def test_empty_list(self, mock_repo):
|
||||
"""空列表"""
|
||||
mock_repo.list_by_user.return_value = []
|
||||
mock_repo.count_by_user.return_value = 0
|
||||
use_case = ListVoiceClonesUseCase(mock_repo)
|
||||
|
||||
items, total = use_case.execute("user_1")
|
||||
|
||||
assert items == []
|
||||
assert total == 0
|
||||
|
||||
|
||||
class TestGetVoiceCloneUseCase:
|
||||
"""GetVoiceCloneUseCase 测试"""
|
||||
|
||||
def test_get_existing(self, mock_repo, sample_profile):
|
||||
"""获取存在的音色克隆"""
|
||||
mock_repo.get.return_value = sample_profile
|
||||
use_case = GetVoiceCloneUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("vc1", "user_1")
|
||||
|
||||
assert result.id == "vc1"
|
||||
mock_repo.get.assert_called_once_with("vc1")
|
||||
|
||||
def test_get_nonexistent_raises(self, mock_repo):
|
||||
"""获取不存在的抛出 VoiceCloneNotFoundError"""
|
||||
mock_repo.get.return_value = None
|
||||
use_case = GetVoiceCloneUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(VoiceCloneNotFoundError, match="not found"):
|
||||
use_case.execute("noexist", "user_1")
|
||||
|
||||
def test_get_wrong_user_raises(self, mock_repo, sample_profile):
|
||||
"""非本人的音色克隆抛出 NotFoundError"""
|
||||
sample_profile.user_id = "other_user"
|
||||
mock_repo.get.return_value = sample_profile
|
||||
use_case = GetVoiceCloneUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(VoiceCloneNotFoundError):
|
||||
use_case.execute("vc1", "user_1")
|
||||
|
||||
|
||||
class TestGetVoiceCloneStatusUseCase:
|
||||
"""GetVoiceCloneStatusUseCase 测试"""
|
||||
|
||||
def test_get_status_completed(self, mock_repo, sample_profile):
|
||||
"""获取 completed 状态"""
|
||||
mock_repo.get.return_value = sample_profile
|
||||
use_case = GetVoiceCloneStatusUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("vc1", "user_1")
|
||||
|
||||
assert result.status == "completed"
|
||||
|
||||
def test_get_status_not_found(self, mock_repo):
|
||||
"""不存在抛出"""
|
||||
mock_repo.get.return_value = None
|
||||
use_case = GetVoiceCloneStatusUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(VoiceCloneNotFoundError):
|
||||
use_case.execute("noexist", "user_1")
|
||||
|
||||
def test_get_status_wrong_user(self, mock_repo, sample_profile):
|
||||
"""非本人抛出"""
|
||||
sample_profile.user_id = "other"
|
||||
mock_repo.get.return_value = sample_profile
|
||||
use_case = GetVoiceCloneStatusUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(VoiceCloneNotFoundError):
|
||||
use_case.execute("vc1", "user_1")
|
||||
|
||||
|
||||
class TestDeleteVoiceCloneUseCase:
|
||||
"""DeleteVoiceCloneUseCase 测试"""
|
||||
|
||||
def test_delete_success(self, mock_repo, sample_profile):
|
||||
"""删除成功"""
|
||||
mock_repo.get.return_value = sample_profile
|
||||
mock_repo.delete.return_value = True
|
||||
use_case = DeleteVoiceCloneUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("vc1", "user_1")
|
||||
|
||||
assert result is True
|
||||
mock_repo.get.assert_called_once_with("vc1")
|
||||
mock_repo.delete.assert_called_once_with("vc1")
|
||||
|
||||
def test_delete_nonexistent_returns_false(self, mock_repo):
|
||||
"""不存在返回 False"""
|
||||
mock_repo.get.return_value = None
|
||||
use_case = DeleteVoiceCloneUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("noexist", "user_1")
|
||||
|
||||
assert result is False
|
||||
mock_repo.delete.assert_not_called()
|
||||
|
||||
def test_delete_wrong_user_returns_false(self, mock_repo, sample_profile):
|
||||
"""非本人返回 False"""
|
||||
sample_profile.user_id = "other"
|
||||
mock_repo.get.return_value = sample_profile
|
||||
use_case = DeleteVoiceCloneUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("vc1", "user_1")
|
||||
|
||||
assert result is False
|
||||
mock_repo.delete.assert_not_called()
|
||||
|
||||
|
||||
class TestRetryVoiceCloneUseCase:
|
||||
"""RetryVoiceCloneUseCase 测试"""
|
||||
|
||||
def test_retry_failed_profile(self, mock_repo):
|
||||
"""失败的音色克隆可以重试"""
|
||||
profile = _make_profile(status="failed")
|
||||
mock_repo.get.return_value = profile
|
||||
mock_repo.update.side_effect = lambda x: x
|
||||
use_case = RetryVoiceCloneUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("vc1", "user_1")
|
||||
|
||||
assert result is not None
|
||||
assert result.status == "pending"
|
||||
assert result.retry_count >= 1
|
||||
mock_repo.update.assert_called_once()
|
||||
|
||||
def test_retry_not_found(self, mock_repo):
|
||||
"""不存在抛出"""
|
||||
mock_repo.get.return_value = None
|
||||
use_case = RetryVoiceCloneUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(VoiceCloneNotFoundError):
|
||||
use_case.execute("noexist", "user_1")
|
||||
|
||||
mock_repo.update.assert_not_called()
|
||||
|
||||
def test_retry_wrong_user(self, mock_repo, sample_profile):
|
||||
"""非本人抛出"""
|
||||
sample_profile.user_id = "other"
|
||||
mock_repo.get.return_value = sample_profile
|
||||
use_case = RetryVoiceCloneUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(VoiceCloneNotFoundError):
|
||||
use_case.execute("vc1", "user_1")
|
||||
|
||||
mock_repo.update.assert_not_called()
|
||||
|
||||
def test_retry_completed_not_retryable(self, mock_repo, sample_profile):
|
||||
"""completed 状态不可重试"""
|
||||
mock_repo.get.return_value = sample_profile
|
||||
use_case = RetryVoiceCloneUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(VoiceCloneNotRetryableError):
|
||||
use_case.execute("vc1", "user_1")
|
||||
|
||||
mock_repo.update.assert_not_called()
|
||||
|
||||
def test_retry_processing_not_retryable(self, mock_repo):
|
||||
"""processing 状态不可重试"""
|
||||
profile = _make_profile(status="processing")
|
||||
mock_repo.get.return_value = profile
|
||||
use_case = RetryVoiceCloneUseCase(mock_repo)
|
||||
|
||||
with pytest.raises(VoiceCloneNotRetryableError):
|
||||
use_case.execute("vc1", "user_1")
|
||||
Executable → Regular
+630
-212
@@ -1,8 +1,14 @@
|
||||
"""音色库 UseCase 单元测试."""
|
||||
"""
|
||||
配音库(Voice Library)Use Case 回归测试
|
||||
|
||||
from __future__ import annotations
|
||||
测试目标:
|
||||
1. CreateVoiceLibraryUseCase - 创建配音库条目,验证 voice_id 字段映射正确(PR#74 P0 bug 修复)
|
||||
2. UpdateVoiceLibraryUseCase - 更新配音库条目,验证 voice_id 字段映射正确
|
||||
3. 配额逻辑覆盖 - free=10, basic=100, premium=100
|
||||
4. 边界条件与异常场景
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
|
||||
@@ -15,259 +21,671 @@ from packages.application.voice_library.use_cases import (
|
||||
DeleteVoiceLibraryUseCase,
|
||||
GetVoiceLibraryUseCase,
|
||||
ListVoiceLibraryUseCase,
|
||||
NotFoundError,
|
||||
QuotaExceededError,
|
||||
UpdateVoiceLibraryUseCase,
|
||||
)
|
||||
from packages.domain.exceptions import NotFoundError, QuotaExceededError
|
||||
from packages.domain.voice_library import VoiceLibraryItem
|
||||
|
||||
|
||||
def _make_item(id: str = "v1", name: str = "测试音色", user_id: str = "user_1") -> VoiceLibraryItem:
|
||||
return VoiceLibraryItem(
|
||||
id=id,
|
||||
user_id=user_id,
|
||||
name=name,
|
||||
text="示例文本",
|
||||
voice_provider="cosyvoice",
|
||||
voice_id="voice_001",
|
||||
voice_name="温柔女声",
|
||||
audio_url="https://oss.example.com/voice.mp3",
|
||||
duration=10.5,
|
||||
file_size=102400,
|
||||
status="ready",
|
||||
project_id="",
|
||||
tags=["温柔", "女声"],
|
||||
metadata_={},
|
||||
)
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_repo():
|
||||
return MagicMock()
|
||||
"""创建 Mock 仓储"""
|
||||
repo = Mock()
|
||||
repo.count_by_user = Mock(return_value=0)
|
||||
repo.create = Mock(side_effect=lambda item: item)
|
||||
repo.update = Mock(side_effect=lambda item: item)
|
||||
repo.get = Mock(return_value=None)
|
||||
repo.delete = Mock(return_value=True)
|
||||
repo.list_by_user = Mock(return_value=[])
|
||||
return repo
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_item():
|
||||
return _make_item()
|
||||
def create_use_case(mock_repo):
|
||||
return CreateVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def update_use_case(mock_repo):
|
||||
return UpdateVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_create_command():
|
||||
"""标准创建命令"""
|
||||
return CreateVoiceLibraryCommand(
|
||||
user_id="user-001",
|
||||
name="测试配音",
|
||||
text="你好世界",
|
||||
voice_provider="aliyun",
|
||||
voice_id="voice-abc-123",
|
||||
voice_name="小云",
|
||||
audio_url="https://oss.example.com/audio/abc.wav",
|
||||
duration=3.5,
|
||||
file_size=56000,
|
||||
status="completed",
|
||||
project_id="proj-001",
|
||||
tags=["测试", "中文"],
|
||||
metadata_={"source": "unit_test"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def existing_voice_item():
|
||||
"""模拟已存在的配音条目"""
|
||||
return VoiceLibraryItem(
|
||||
id="existing-voice-001",
|
||||
user_id="user-001",
|
||||
name="旧配音",
|
||||
text="旧文本",
|
||||
voice_provider="old_provider",
|
||||
voice_id="old-voice-id",
|
||||
voice_name="旧声音",
|
||||
audio_url="https://oss.example.com/old.wav",
|
||||
duration=1.0,
|
||||
file_size=16000,
|
||||
status="completed",
|
||||
project_id="proj-001",
|
||||
tags=["旧"],
|
||||
metadata_={},
|
||||
)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 1. CreateVoiceLibraryUseCase 测试
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestCreateVoiceLibraryUseCase:
|
||||
"""配音库创建 UseCase 测试"""
|
||||
|
||||
def test_create_success_all_fields(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""测试创建成功 - 所有字段完整传入"""
|
||||
result = create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert result is not None
|
||||
assert result.user_id == "user-001"
|
||||
assert result.name == "测试配音"
|
||||
assert result.text == "你好世界"
|
||||
assert result.voice_provider == "aliyun"
|
||||
assert result.voice_name == "小云"
|
||||
assert result.audio_url == "https://oss.example.com/audio/abc.wav"
|
||||
assert result.duration == 3.5
|
||||
assert result.file_size == 56000
|
||||
assert result.status == "completed"
|
||||
assert result.project_id == "proj-001"
|
||||
assert result.tags == ["测试", "中文"]
|
||||
assert result.metadata_ == {"source": "unit_test"}
|
||||
|
||||
mock_repo.count_by_user.assert_called_once_with("user-001")
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_create_voice_id_field_mapping(self, create_use_case, mock_repo):
|
||||
"""
|
||||
【P0 回归】验证 voice_id 字段映射正确
|
||||
|
||||
PR#74 修复了 command.id 被错误使用的问题。
|
||||
此测试确保 CreateVoiceLibraryCommand 中的 voice_id 字段
|
||||
被正确传递到 VoiceLibraryItem 的 voice_id 属性上,
|
||||
而非被其他字段(如 item 自身的 id)覆盖。
|
||||
"""
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id="user-001",
|
||||
name="voice_id 回归测试",
|
||||
voice_id="specific-voice-id-xyz",
|
||||
voice_provider="azure",
|
||||
voice_name="Azure Xiaoxiao",
|
||||
)
|
||||
|
||||
result = create_use_case.execute(command, plan_name="free")
|
||||
|
||||
# 核心断言:voice_id 必须来自 command.voice_id
|
||||
assert result.voice_id == "specific-voice-id-xyz", "voice_id 应来自 command.voice_id,而非其他字段"
|
||||
# 同时确保 item 自身生成的 id 与 voice_id 不同
|
||||
assert result.id != "specific-voice-id-xyz", "item.id(UUID)不应与 voice_id 混淆"
|
||||
|
||||
def test_create_voice_id_empty_string(self, create_use_case, mock_repo):
|
||||
"""测试 voice_id 为空字符串的合法场景"""
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id="user-001",
|
||||
name="无 voice_id 配音",
|
||||
voice_id="",
|
||||
voice_provider="custom",
|
||||
)
|
||||
|
||||
result = create_use_case.execute(command, plan_name="free")
|
||||
|
||||
assert result.voice_id == ""
|
||||
|
||||
def test_create_default_values(self, create_use_case, mock_repo):
|
||||
"""测试默认值填充"""
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id="user-001",
|
||||
name="最小化创建",
|
||||
)
|
||||
|
||||
result = create_use_case.execute(command, plan_name="free")
|
||||
|
||||
assert result.text == ""
|
||||
assert result.voice_provider == ""
|
||||
assert result.voice_id == ""
|
||||
assert result.voice_name == ""
|
||||
assert result.audio_url == ""
|
||||
assert result.duration == 0
|
||||
assert result.file_size == 0
|
||||
assert result.status == "completed"
|
||||
assert result.project_id is None
|
||||
assert result.tags == []
|
||||
assert result.metadata_ == {}
|
||||
|
||||
def test_create_generates_uuid(self, create_use_case, mock_repo):
|
||||
"""测试创建时自动生成 UUID 作为 id"""
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id="user-001",
|
||||
name="UUID 测试",
|
||||
)
|
||||
|
||||
result = create_use_case.execute(command, plan_name="free")
|
||||
|
||||
assert result.id is not None
|
||||
assert len(result.id) == 32 # uuid4().hex 长度为 32
|
||||
assert result.id.isalnum()
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 2. 配额逻辑测试(Create 时的配额检查)
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestCreateVoiceLibraryQuota:
|
||||
"""配音库创建配额检查测试"""
|
||||
|
||||
def test_quota_free_plan_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""free 套餐(上限10),当前 5 个,允许创建"""
|
||||
mock_repo.count_by_user.return_value = 5
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_free_plan_at_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""free 套餐(上限10),当前 10 个,拒绝创建"""
|
||||
mock_repo.count_by_user.return_value = 10
|
||||
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert exc_info.value.dimension == "max_voiceovers"
|
||||
assert exc_info.value.limit == 10
|
||||
assert exc_info.value.used == 10
|
||||
mock_repo.create.assert_not_called()
|
||||
|
||||
def test_quota_free_plan_over_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""free 套餐(上限10),当前 15 个,拒绝创建"""
|
||||
mock_repo.count_by_user.return_value = 15
|
||||
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert exc_info.value.dimension == "max_voiceovers"
|
||||
assert exc_info.value.limit == 10
|
||||
assert exc_info.value.used == 15
|
||||
|
||||
def test_quota_free_plan_just_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""free 套餐(上限10),当前 9 个,允许创建(边界)"""
|
||||
mock_repo.count_by_user.return_value = 9
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_basic_plan_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""basic 套餐(上限100),当前 50 个,允许创建"""
|
||||
mock_repo.count_by_user.return_value = 50
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="basic")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_basic_plan_at_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""basic 套餐(上限100),当前 100 个,拒绝创建"""
|
||||
mock_repo.count_by_user.return_value = 100
|
||||
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
create_use_case.execute(sample_create_command, plan_name="basic")
|
||||
|
||||
assert exc_info.value.dimension == "max_voiceovers"
|
||||
assert exc_info.value.limit == 100
|
||||
assert exc_info.value.used == 100
|
||||
|
||||
def test_quota_basic_plan_just_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""basic 套餐(上限100),当前 99 个,允许创建(边界)"""
|
||||
mock_repo.count_by_user.return_value = 99
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="basic")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_premium_plan_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""premium 套餐(上限100),当前 50 个,允许创建"""
|
||||
mock_repo.count_by_user.return_value = 50
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="premium")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_premium_plan_at_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""premium 套餐(上限100),当前 100 个,拒绝创建"""
|
||||
mock_repo.count_by_user.return_value = 100
|
||||
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
create_use_case.execute(sample_create_command, plan_name="premium")
|
||||
|
||||
assert exc_info.value.dimension == "max_voiceovers"
|
||||
assert exc_info.value.limit == 100
|
||||
assert exc_info.value.used == 100
|
||||
|
||||
def test_quota_premium_plan_just_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""premium 套餐(上限100),当前 99 个,允许创建(边界)"""
|
||||
mock_repo.count_by_user.return_value = 99
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="premium")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_zero_usage(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""新用户零使用量,所有套餐均可创建"""
|
||||
mock_repo.count_by_user.return_value = 0
|
||||
|
||||
for plan in ["free", "basic", "premium"]:
|
||||
mock_repo.create.reset_mock()
|
||||
mock_repo.count_by_user.reset_mock()
|
||||
mock_repo.count_by_user.return_value = 0
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name=plan)
|
||||
assert result is not None, f"{plan} 套餐零使用量应允许创建"
|
||||
|
||||
def test_quota_unknown_plan_defaults_to_zero(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""未知套餐名默认配额为 0,即使 0 使用量也无法创建"""
|
||||
mock_repo.count_by_user.return_value = 0
|
||||
|
||||
with pytest.raises(QuotaExceededError):
|
||||
create_use_case.execute(sample_create_command, plan_name="unknown_plan")
|
||||
|
||||
def test_quota_exceeded_error_attributes(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""QuotaExceededError 异常属性完整性"""
|
||||
mock_repo.count_by_user.return_value = 10
|
||||
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
err = exc_info.value
|
||||
assert hasattr(err, "dimension")
|
||||
assert hasattr(err, "limit")
|
||||
assert hasattr(err, "used")
|
||||
assert "max_voiceovers" in str(err)
|
||||
assert "10" in str(err)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 3. UpdateVoiceLibraryUseCase 测试
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestUpdateVoiceLibraryUseCase:
|
||||
"""配音库更新 UseCase 测试"""
|
||||
|
||||
def test_update_success_all_fields(self, update_use_case, mock_repo, existing_voice_item):
|
||||
"""测试全字段更新成功"""
|
||||
mock_repo.get.return_value = existing_voice_item
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="existing-voice-001",
|
||||
user_id="user-001",
|
||||
name="更新后的名称",
|
||||
text="更新后的文本",
|
||||
voice_provider="new_provider",
|
||||
voice_id="new-voice-id-456",
|
||||
voice_name="新声音",
|
||||
audio_url="https://oss.example.com/new.wav",
|
||||
duration=5.0,
|
||||
file_size=80000,
|
||||
status="processing",
|
||||
tags=["新标签"],
|
||||
metadata_={"updated": True},
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
assert result.name == "更新后的名称"
|
||||
assert result.text == "更新后的文本"
|
||||
assert result.voice_provider == "new_provider"
|
||||
assert result.voice_name == "新声音"
|
||||
assert result.audio_url == "https://oss.example.com/new.wav"
|
||||
assert result.duration == 5.0
|
||||
assert result.file_size == 80000
|
||||
assert result.status == "processing"
|
||||
assert result.tags == ["新标签"]
|
||||
assert result.metadata_ == {"updated": True}
|
||||
|
||||
mock_repo.update.assert_called_once()
|
||||
|
||||
def test_update_voice_id_field_mapping(self, update_use_case, mock_repo, existing_voice_item):
|
||||
"""
|
||||
【P0 回归】验证 update 时 voice_id 字段映射正确
|
||||
|
||||
PR#74 修复了 API 路由层将 command.id 错误传给 voice_id 的 bug。
|
||||
此测试确保 UpdateVoiceLibraryCommand 中 voice_id 字段
|
||||
被正确写入 VoiceLibraryItem.voice_id,而非被 item.id 覆盖。
|
||||
"""
|
||||
mock_repo.get.return_value = existing_voice_item
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="existing-voice-001",
|
||||
user_id="user-001",
|
||||
voice_id="completely-different-voice-id",
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
# 核心断言:voice_id 应被更新为新值
|
||||
assert result.voice_id == "completely-different-voice-id", "voice_id 应被更新为 command.voice_id 的值"
|
||||
# item 自身的 id 保持不变
|
||||
assert result.id == "existing-voice-001"
|
||||
|
||||
def test_update_partial_only_voice_id(self, update_use_case, mock_repo, existing_voice_item):
|
||||
"""测试仅更新 voice_id 一个字段"""
|
||||
mock_repo.get.return_value = existing_voice_item
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="existing-voice-001",
|
||||
user_id="user-001",
|
||||
voice_id="only-voice-id-changed",
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
assert result.voice_id == "only-voice-id-changed"
|
||||
# 其他字段保持不变
|
||||
assert result.name == "旧配音"
|
||||
assert result.text == "旧文本"
|
||||
assert result.voice_provider == "old_provider"
|
||||
assert result.voice_name == "旧声音"
|
||||
assert result.audio_url == "https://oss.example.com/old.wav"
|
||||
assert result.duration == 1.0
|
||||
assert result.file_size == 16000
|
||||
|
||||
def test_update_partial_only_name(self, update_use_case, mock_repo, existing_voice_item):
|
||||
"""测试仅更新 name"""
|
||||
mock_repo.get.return_value = existing_voice_item
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="existing-voice-001",
|
||||
user_id="user-001",
|
||||
name="仅改名",
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
assert result.name == "仅改名"
|
||||
assert result.voice_id == "old-voice-id" # voice_id 不变
|
||||
|
||||
def test_update_not_found(self, update_use_case, mock_repo):
|
||||
"""测试更新不存在的条目"""
|
||||
mock_repo.get.return_value = None
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="nonexistent-id",
|
||||
user_id="user-001",
|
||||
name="不存在",
|
||||
)
|
||||
|
||||
with pytest.raises(NotFoundError, match="nonexistent-id"):
|
||||
update_use_case.execute(command)
|
||||
|
||||
mock_repo.update.assert_not_called()
|
||||
|
||||
def test_update_wrong_user(self, update_use_case, mock_repo):
|
||||
"""测试用户隔离 - 不能更新其他用户的条目"""
|
||||
mock_repo.get.return_value = None # repo 返回 None 表示找不到(不同 user_id)
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="existing-voice-001",
|
||||
user_id="other-user-999",
|
||||
name="恶意修改",
|
||||
)
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
update_use_case.execute(command)
|
||||
|
||||
def test_update_none_fields_not_changed(self, update_use_case, mock_repo, existing_voice_item):
|
||||
"""测试 None 字段不覆盖原有值"""
|
||||
mock_repo.get.return_value = existing_voice_item
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="existing-voice-001",
|
||||
user_id="user-001",
|
||||
# 所有可选字段保持 None
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
# 所有字段应保持不变
|
||||
assert result.name == "旧配音"
|
||||
assert result.text == "旧文本"
|
||||
assert result.voice_id == "old-voice-id"
|
||||
assert result.voice_provider == "old_provider"
|
||||
assert result.voice_name == "旧声音"
|
||||
assert result.audio_url == "https://oss.example.com/old.wav"
|
||||
assert result.duration == 1.0
|
||||
assert result.file_size == 16000
|
||||
assert result.status == "completed"
|
||||
|
||||
def test_update_voice_id_empty_string(self, update_use_case, mock_repo, existing_voice_item):
|
||||
"""测试 voice_id 更新为空字符串(合法场景:清除 voice_id)"""
|
||||
mock_repo.get.return_value = existing_voice_item
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="existing-voice-001",
|
||||
user_id="user-001",
|
||||
voice_id="",
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
assert result.voice_id == ""
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 4. DeleteVoiceLibraryUseCase 测试
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestDeleteVoiceLibraryUseCase:
|
||||
"""配音库删除 UseCase 测试"""
|
||||
|
||||
def test_delete_success(self, mock_repo):
|
||||
"""测试删除成功"""
|
||||
mock_repo.delete.return_value = True
|
||||
use_case = DeleteVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("voice-001", "user-001")
|
||||
|
||||
assert result is True
|
||||
mock_repo.delete.assert_called_once_with("voice-001", "user-001")
|
||||
|
||||
def test_delete_not_found(self, mock_repo):
|
||||
"""测试删除不存在的条目"""
|
||||
mock_repo.delete.return_value = False
|
||||
use_case = DeleteVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("nonexistent", "user-001")
|
||||
|
||||
assert result is False
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 5. GetVoiceLibraryUseCase 测试
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestGetVoiceLibraryUseCase:
|
||||
"""配音库查询 UseCase 测试"""
|
||||
|
||||
def test_get_existing(self, mock_repo):
|
||||
"""测试查询存在的条目"""
|
||||
expected = VoiceLibraryItem(
|
||||
id="v-001",
|
||||
user_id="user-001",
|
||||
name="测试",
|
||||
voice_id="voice-xyz",
|
||||
)
|
||||
mock_repo.get.return_value = expected
|
||||
use_case = GetVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("v-001", "user-001")
|
||||
|
||||
assert result is not None
|
||||
assert result.id == "v-001"
|
||||
assert result.voice_id == "voice-xyz"
|
||||
mock_repo.get.assert_called_once_with("v-001", "user-001")
|
||||
|
||||
def test_get_not_found(self, mock_repo):
|
||||
"""测试查询不存在的条目"""
|
||||
mock_repo.get.return_value = None
|
||||
use_case = GetVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("nonexistent", "user-001")
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 6. ListVoiceLibraryUseCase 测试
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestListVoiceLibraryUseCase:
|
||||
"""ListVoiceLibraryUseCase 测试"""
|
||||
"""配音库列表 UseCase 测试"""
|
||||
|
||||
def test_list_returns_items_and_total(self, mock_repo, sample_item):
|
||||
"""返回 (items, total_count) 元组"""
|
||||
mock_repo.list_by_user.return_value = [sample_item]
|
||||
mock_repo.count_by_user.return_value = 1
|
||||
use_case = ListVoiceLibraryUseCase(mock_repo)
|
||||
def test_list_default(self, mock_repo):
|
||||
"""测试默认列表查询"""
|
||||
items = [
|
||||
VoiceLibraryItem(id="v1", user_id="user-001", name="A"),
|
||||
VoiceLibraryItem(id="v2", user_id="user-001", name="B"),
|
||||
]
|
||||
mock_repo.list_by_user.return_value = items
|
||||
mock_repo.count_by_user.return_value = 2
|
||||
use_case = ListVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
items, total = use_case.execute("user_1")
|
||||
result_items, total = use_case.execute("user-001")
|
||||
|
||||
assert len(items) == 1
|
||||
assert items[0].id == "v1"
|
||||
assert total == 1
|
||||
mock_repo.list_by_user.assert_called_once_with(
|
||||
"user_1", status=None, skip=0, limit=50
|
||||
)
|
||||
mock_repo.count_by_user.assert_called_once_with("user_1")
|
||||
assert len(result_items) == 2
|
||||
assert total == 2
|
||||
mock_repo.list_by_user.assert_called_once_with("user-001", status=None, skip=0, limit=50)
|
||||
|
||||
def test_list_with_status_filter(self, mock_repo, sample_item):
|
||||
"""按状态过滤"""
|
||||
mock_repo.list_by_user.return_value = [sample_item]
|
||||
mock_repo.count_by_user.return_value = 1
|
||||
use_case = ListVoiceLibraryUseCase(mock_repo)
|
||||
def test_list_with_status_filter(self, mock_repo):
|
||||
"""测试按状态筛选"""
|
||||
mock_repo.list_by_user.return_value = []
|
||||
use_case = ListVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
items, total = use_case.execute("user_1", status="ready")
|
||||
use_case.execute("user-001", status="completed", skip=10, limit=20)
|
||||
|
||||
assert total == 1
|
||||
mock_repo.list_by_user.assert_called_once_with(
|
||||
"user_1", status="ready", skip=0, limit=50
|
||||
)
|
||||
mock_repo.count_by_user.assert_called_once_with("user_1", status="ready")
|
||||
mock_repo.list_by_user.assert_called_once_with("user-001", status="completed", skip=10, limit=20)
|
||||
|
||||
def test_list_with_pagination(self, mock_repo, sample_item):
|
||||
"""带分页参数"""
|
||||
mock_repo.list_by_user.return_value = [sample_item]
|
||||
mock_repo.count_by_user.return_value = 10
|
||||
use_case = ListVoiceLibraryUseCase(mock_repo)
|
||||
|
||||
use_case.execute("user_1", skip=10, limit=20)
|
||||
|
||||
mock_repo.list_by_user.assert_called_once_with(
|
||||
"user_1", status=None, skip=10, limit=20
|
||||
)
|
||||
|
||||
def test_empty_list(self, mock_repo):
|
||||
"""空列表"""
|
||||
def test_list_empty(self, mock_repo):
|
||||
"""测试空列表"""
|
||||
mock_repo.list_by_user.return_value = []
|
||||
mock_repo.count_by_user.return_value = 0
|
||||
use_case = ListVoiceLibraryUseCase(mock_repo)
|
||||
use_case = ListVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
items, total = use_case.execute("user_1")
|
||||
items, total = use_case.execute("user-001")
|
||||
|
||||
assert items == []
|
||||
assert total == 0
|
||||
|
||||
|
||||
class TestGetVoiceLibraryUseCase:
|
||||
"""GetVoiceLibraryUseCase 测试"""
|
||||
|
||||
def test_get_existing(self, mock_repo, sample_item):
|
||||
"""获取存在的音色"""
|
||||
mock_repo.get.return_value = sample_item
|
||||
use_case = GetVoiceLibraryUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("v1", "user_1")
|
||||
|
||||
assert result is not None
|
||||
assert result.id == "v1"
|
||||
mock_repo.get.assert_called_once_with("v1", "user_1")
|
||||
|
||||
def test_get_nonexistent_returns_none(self, mock_repo):
|
||||
"""获取不存在的返回 None"""
|
||||
mock_repo.get.return_value = None
|
||||
use_case = GetVoiceLibraryUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("noexist", "user_1")
|
||||
|
||||
assert result is None
|
||||
# ===========================================================================
|
||||
# 7. voice_id 与 id 字段隔离专项回归测试
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestCreateVoiceLibraryUseCase:
|
||||
"""CreateVoiceLibraryUseCase 测试"""
|
||||
class TestVoiceIdFieldIsolation:
|
||||
"""
|
||||
PR#74 P0 Bug 回归:voice_id 与 item.id 字段隔离
|
||||
|
||||
def test_create_success(self, mock_repo, sample_item):
|
||||
"""创建成功"""
|
||||
mock_repo.count_by_user.return_value = 0
|
||||
mock_repo.create.return_value = sample_item
|
||||
use_case = CreateVoiceLibraryUseCase(mock_repo)
|
||||
原 bug:API 路由层误将 command.id(item 主键)用作 voice_id,
|
||||
导致 voice_id 字段值错误。本测试类从 UseCase 层验证
|
||||
这两个字段在整个 CRUD 生命周期中互不干扰。
|
||||
"""
|
||||
|
||||
def test_create_id_and_voice_id_are_independent(self, create_use_case, mock_repo):
|
||||
"""创建时 id 自动生成,voice_id 来自 command"""
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id="user_1",
|
||||
name="新音色",
|
||||
text="你好",
|
||||
voice_provider="cosyvoice",
|
||||
voice_id="v_new",
|
||||
voice_name="新音色名",
|
||||
audio_url="https://oss.example.com/new.mp3",
|
||||
duration=5.0,
|
||||
file_size=51200,
|
||||
status="processing",
|
||||
project_id="",
|
||||
tags=[],
|
||||
metadata_={},
|
||||
user_id="user-001",
|
||||
name="隔离测试",
|
||||
voice_id="tts-voice-001",
|
||||
voice_provider="openai",
|
||||
)
|
||||
result = use_case.execute(command, plan_name="free")
|
||||
|
||||
assert result.id == "v1"
|
||||
mock_repo.count_by_user.assert_called_once_with("user_1")
|
||||
mock_repo.create.assert_called_once()
|
||||
created = mock_repo.create.call_args[0][0]
|
||||
assert isinstance(created, VoiceLibraryItem)
|
||||
assert created.name == "新音色"
|
||||
result = create_use_case.execute(command, plan_name="free")
|
||||
|
||||
def test_create_quota_exceeded(self, mock_repo):
|
||||
"""超过配额抛出 QuotaExceededError"""
|
||||
mock_repo.count_by_user.return_value = 9999
|
||||
use_case = CreateVoiceLibraryUseCase(mock_repo)
|
||||
assert result.id != result.voice_id, "id 和 voice_id 应为不同值"
|
||||
assert result.voice_id == "tts-voice-001"
|
||||
assert len(result.id) == 32 # UUID hex
|
||||
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id="user_1",
|
||||
name="超限音色",
|
||||
text="text",
|
||||
voice_provider="cosyvoice",
|
||||
voice_id="v",
|
||||
voice_name="v",
|
||||
audio_url="url",
|
||||
duration=1.0,
|
||||
file_size=100,
|
||||
status="ready",
|
||||
project_id="",
|
||||
tags=[],
|
||||
metadata_={},
|
||||
def test_update_voice_id_does_not_change_id(self, update_use_case, mock_repo):
|
||||
"""更新 voice_id 不影响 item 主键 id"""
|
||||
existing = VoiceLibraryItem(
|
||||
id="stable-id-001",
|
||||
user_id="user-001",
|
||||
name="测试",
|
||||
voice_id="old-voice",
|
||||
)
|
||||
with pytest.raises(QuotaExceededError):
|
||||
use_case.execute(command, plan_name="free")
|
||||
|
||||
mock_repo.create.assert_not_called()
|
||||
|
||||
|
||||
class TestUpdateVoiceLibraryUseCase:
|
||||
"""UpdateVoiceLibraryUseCase 测试"""
|
||||
|
||||
def test_update_name(self, mock_repo, sample_item):
|
||||
"""更新名称"""
|
||||
mock_repo.get.return_value = sample_item
|
||||
mock_repo.update.side_effect = lambda x: x
|
||||
use_case = UpdateVoiceLibraryUseCase(mock_repo)
|
||||
|
||||
command = UpdateVoiceLibraryCommand(id="v1", user_id="user_1", name="新名字")
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.name == "新名字"
|
||||
# 其他字段不变
|
||||
assert result.voice_name == "温柔女声"
|
||||
mock_repo.get.assert_called_once_with("v1", "user_1")
|
||||
mock_repo.update.assert_called_once()
|
||||
|
||||
def test_update_status(self, mock_repo, sample_item):
|
||||
"""更新状态"""
|
||||
mock_repo.get.return_value = sample_item
|
||||
mock_repo.update.side_effect = lambda x: x
|
||||
use_case = UpdateVoiceLibraryUseCase(mock_repo)
|
||||
|
||||
command = UpdateVoiceLibraryCommand(id="v1", user_id="user_1", status="failed")
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.status == "failed"
|
||||
|
||||
def test_update_multiple_fields(self, mock_repo, sample_item):
|
||||
"""同时更新多个字段"""
|
||||
mock_repo.get.return_value = sample_item
|
||||
mock_repo.update.side_effect = lambda x: x
|
||||
use_case = UpdateVoiceLibraryUseCase(mock_repo)
|
||||
mock_repo.get.return_value = existing
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="v1",
|
||||
user_id="user_1",
|
||||
name="更新后",
|
||||
duration=15.0,
|
||||
tags=["新标签"],
|
||||
id="stable-id-001",
|
||||
user_id="user-001",
|
||||
voice_id="new-voice-999",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.name == "更新后"
|
||||
assert result.duration == 15.0
|
||||
assert result.tags == ["新标签"]
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
def test_update_nonexistent_raises(self, mock_repo):
|
||||
"""更新不存在的抛出 NotFoundError"""
|
||||
mock_repo.get.return_value = None
|
||||
use_case = UpdateVoiceLibraryUseCase(mock_repo)
|
||||
assert result.id == "stable-id-001", "item 主键 id 不应改变"
|
||||
assert result.voice_id == "new-voice-999", "voice_id 应被更新"
|
||||
|
||||
command = UpdateVoiceLibraryCommand(id="noexist", user_id="user_1", name="x")
|
||||
with pytest.raises(NotFoundError, match="not found"):
|
||||
use_case.execute(command)
|
||||
def test_create_then_update_voice_id_preserves_id(self, create_use_case, update_use_case, mock_repo):
|
||||
"""创建后再更新 voice_id,id 始终不变"""
|
||||
# 创建
|
||||
create_cmd = CreateVoiceLibraryCommand(
|
||||
user_id="user-001",
|
||||
name="生命周期测试",
|
||||
voice_id="initial-voice",
|
||||
)
|
||||
created = create_use_case.execute(create_cmd, plan_name="free")
|
||||
original_id = created.id
|
||||
|
||||
mock_repo.update.assert_not_called()
|
||||
# 更新
|
||||
mock_repo.get.return_value = created
|
||||
update_cmd = UpdateVoiceLibraryCommand(
|
||||
id=original_id,
|
||||
user_id="user-001",
|
||||
voice_id="updated-voice",
|
||||
)
|
||||
updated = update_use_case.execute(update_cmd)
|
||||
|
||||
|
||||
class TestDeleteVoiceLibraryUseCase:
|
||||
"""DeleteVoiceLibraryUseCase 测试"""
|
||||
|
||||
def test_delete_success(self, mock_repo):
|
||||
"""删除成功"""
|
||||
mock_repo.delete.return_value = True
|
||||
use_case = DeleteVoiceLibraryUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("v1", "user_1")
|
||||
|
||||
assert result is True
|
||||
mock_repo.delete.assert_called_once_with("v1", "user_1")
|
||||
|
||||
def test_delete_nonexistent_returns_false(self, mock_repo):
|
||||
"""删除不存在的返回 False"""
|
||||
mock_repo.delete.return_value = False
|
||||
use_case = DeleteVoiceLibraryUseCase(mock_repo)
|
||||
|
||||
result = use_case.execute("noexist", "user_1")
|
||||
|
||||
assert result is False
|
||||
assert updated.id == original_id, "经过创建和更新,id 应保持一致"
|
||||
assert updated.voice_id == "updated-voice"
|
||||
assert updated.voice_id != "initial-voice"
|
||||
|
||||
Reference in New Issue
Block a user