feat(#1899,#1900): clip_count 参数 + 模板管理 API 冗余清理 #1918

Merged
xiaoxia merged 2 commits from feat/1899-1900-clip-count-template-cleanup into develop 2026-09-15 08:04:21 +08:00
8 changed files with 206 additions and 892 deletions
+33 -438
View File
@@ -1,4 +1,12 @@
"""Template CRUD + generate + category routes."""
"""Template 列表路由(供生成页自动选模板).
保留:
- GET /templates:列表查询(生成页使用)
- 默认模板自动创建兜底逻辑(``_get_or_create_default_template_id`` 位于
generation_variant_plans.py)继续通过 service 层 CreateTemplateUseCase 工作,
但不再暴露模板 CRUD / 分类 / 标签 / 收藏 / 复制 / 校验 / 使用统计等 HTTP 端点
(前端 PR#1911 已删除 my-templates / editing-planner / templates 管理页面)。
"""
from __future__ import annotations
@@ -7,53 +15,17 @@ import logging
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_db_session
from app.schemas.template import (
CategoryResponse,
CopyTemplateRequest,
CreateCategoryRequest,
CreateTemplateRequest,
GenerateWarningResponse,
ListCategoriesResponse,
ListTagsResponse,
ListTemplatesResponse,
SegmentResponse,
TemplateResponse,
TemplateUsageResponse,
ToggleFavoriteResponse,
UpdateTemplateRequest,
ValidateTemplateRequest,
ValidateTemplateResponse,
)
from fastapi import APIRouter, Depends, HTTPException, Query, Response, status
from fastapi import APIRouter, Depends, Query
from sqlalchemy.orm import Session
logger = logging.getLogger(__name__)
from packages.adapters.sqlalchemy_impl.template_repository import SQLAlchemyTemplateRepository
from packages.application.template.commands import (
CopyTemplateCommand,
CreateCategoryCommand,
CreateTemplateCommand,
ListTemplatesFilter,
SegmentCommand,
UpdateTemplateCommand,
ValidateTemplateCommand,
)
from packages.application.template.use_cases import (
CopyTemplateUseCase,
CountTemplatesUseCase,
CreateCategoryUseCase,
CreateTemplateUseCase,
DeleteCategoryUseCase,
DeleteTemplateUseCase,
GetTemplateUseCase,
ListCategoriesUseCase,
ListTagsUseCase,
ListTemplatesUseCase,
NotFoundError,
UpdateTemplateUseCase,
ValidateTemplateUseCase,
ValidationError,
)
from packages.application.template.commands import ListTemplatesFilter
from packages.application.template.use_cases import CountTemplatesUseCase, ListTemplatesUseCase
router = APIRouter()
@@ -62,405 +34,28 @@ def _get_template_repository(session: Session = Depends(get_db_session)) -> SQLA
return SQLAlchemyTemplateRepository(session)
def _segment_to_response(seg) -> SegmentResponse:
return SegmentResponse(
id=seg.id,
template_id=seg.template_id,
segment_order=seg.segment_order,
duration_min=seg.duration_min,
duration_max=seg.duration_max,
material_type=seg.material_type,
created_at=seg.created_at,
updated_at=seg.updated_at,
)
def _to_response(template, usage_count: int = 0) -> TemplateResponse:
return TemplateResponse(
id=template.id,
user_id=template.user_id,
name=template.name,
mode=template.mode,
category=template.category,
tags=template.tags,
title_config=template.title_config,
subtitle_config=template.subtitle_config,
bgm_config=template.bgm_config,
estimated_duration=template.estimated_duration,
segments=[_segment_to_response(s) for s in getattr(template, "segments", [])],
is_active=template.is_active,
usage_count=usage_count,
created_at=template.created_at,
updated_at=template.updated_at,
)
def _ensure_default_template(
user_id: str,
template_repository: SQLAlchemyTemplateRepository,
):
"""当用户无任何有效模板时,自动创建一个默认 voice_over 模式模板。
前端 #1911 删除了模板选择 UI(6步→5步),改为后台自动选第一个有效模板。
为避免新用户/无模板用户在剪辑页卡死在"独立选片中...",当 valid_only 查询
结果为空时自动创建一条默认配音模板(含 1 个 order=0 的通用片段配置),
让 from-assets 能正常分配片段。
Returns:
创建的默认 Template 实体;创建失败返回 None。
"""
try:
cmd = CreateTemplateCommand(
user_id=user_id,
name="默认配音模板",
mode="voice_over",
category="default",
tags=[],
title_config={},
subtitle_config={},
bgm_config={},
estimated_duration=0.0,
segments=[
SegmentCommand(
segment_order=0,
duration_min=1.0,
duration_max=30.0,
material_type=None,
),
],
)
use_case = CreateTemplateUseCase(template_repository)
tpl = use_case.execute(cmd)
logger.info("[template] 自动创建默认模板: user=%s tpl=%s", user_id, tpl.id)
return tpl
except Exception:
logger.exception("自动创建默认模板失败: user_id=%s", user_id)
return None
# ── Template CRUD ──
@router.get("", response_model=ListTemplatesResponse)
@router.get("", response_model=ListTemplatesResponse, summary="获取模板列表")
def list_templates(
skip: int = Query(0, ge=0),
limit: int = Query(50, ge=1, le=200),
mode: str | None = Query(None, description="编辑模式:generic/vlog/storyboard,不传返回全部"),
category: str | None = Query(None, description="按分类筛选"),
tag: str | None = Query(None, description="按标签筛选"),
keyword: str | None = Query(None, description="按名称关键词搜索"),
mode: str | None = Query(None, description="按剪辑模式筛选"),
valid_only: bool = Query(
False,
description="仅返回已配置片段的模板(剪辑页传 true;模板编辑器不传,可查看全部模板含草稿)",
),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ListTemplatesResponse:
user_id = authenticated_user.user.id
try:
tpl_filter = ListTemplatesFilter(
category=category,
tag=tag,
keyword=keyword,
mode=mode,
valid_only=valid_only,
)
use_case = ListTemplatesUseCase(template_repository)
templates = use_case.execute(user_id, skip=skip, limit=limit, filter=tpl_filter)
count_use_case = CountTemplatesUseCase(template_repository)
total = count_use_case.execute(user_id, filter=tpl_filter)
# P0 兜底:剪辑页(valid_only=true)首次访问且用户无任何有效模板时,
# 自动创建一条默认配音模板,避免前端 selectedTemplate 永远为空导致卡死。
if valid_only and total == 0 and skip == 0:
default_tpl = _ensure_default_template(user_id, template_repository)
if default_tpl is not None:
templates = [default_tpl]
total = 1
# 批量查询使用次数
items = []
for t in templates:
usage = template_repository.get_usage_count(t.id)
items.append(_to_response(t, usage_count=usage))
except Exception:
logger.exception("list_templates 查询失败: user_id=%s", user_id)
return ListTemplatesResponse(items=[], total=0)
return ListTemplatesResponse(
items=items,
total=total,
page: int = Query(1, ge=1, description="页码,从 1 开始"),
page_size: int = Query(20, ge=1, le=100, description="每页条数,默认 20"),
current_user: AuthenticatedUser = Depends(get_current_user),
repo: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
):
"""获取用户可用的模板列表(仅返回 active 状态)。"""
user_id = str(current_user.user.id)
list_uc = ListTemplatesUseCase(repo)
count_uc = CountTemplatesUseCase(repo)
filters = ListTemplatesFilter(
category=category,
tag=tag,
mode=mode,
valid_only=True, # 仅返回 active + 有片段配置
)
@router.get("/{template_id}", response_model=TemplateResponse)
def get_template(
template_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateResponse:
user_id = authenticated_user.user.id
try:
use_case = GetTemplateUseCase(template_repository)
template = use_case.execute(template_id, user_id)
usage = template_repository.get_usage_count(template_id)
except Exception as _e:
logger.exception("get_template 查询失败: template_id=%s", template_id)
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="模板查询失败") from _e
if template is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
return _to_response(template, usage_count=usage)
@router.post("", response_model=TemplateResponse, status_code=status.HTTP_201_CREATED)
def create_template(
request: CreateTemplateRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateResponse:
user_id = authenticated_user.user.id
command = CreateTemplateCommand(
user_id=user_id,
name=request.name,
mode=request.mode,
category=request.category,
tags=request.tags,
title_config=request.title_config,
subtitle_config=request.subtitle_config,
bgm_config=request.bgm_config,
estimated_duration=request.estimated_duration,
segments=[
SegmentCommand(
segment_order=s.segment_order,
duration_min=s.duration_min,
duration_max=s.duration_max,
material_type=s.material_type,
)
for s in request.segments
],
)
use_case = CreateTemplateUseCase(template_repository)
try:
template = use_case.execute(command)
except ValidationError as exc:
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc)) from exc
return _to_response(template)
@router.patch("/{template_id}", response_model=TemplateResponse)
def update_template(
template_id: str,
request: UpdateTemplateRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateResponse:
user_id = authenticated_user.user.id
command = UpdateTemplateCommand(
template_id=template_id,
user_id=user_id,
name=request.name,
mode=request.mode,
category=request.category,
tags=request.tags,
title_config=request.title_config,
subtitle_config=request.subtitle_config,
bgm_config=request.bgm_config,
estimated_duration=request.estimated_duration,
segments=(
[
SegmentCommand(
segment_order=s.segment_order,
duration_min=s.duration_min,
duration_max=s.duration_max,
material_type=s.material_type,
)
for s in request.segments
]
if request.segments is not None
else None
),
)
use_case = UpdateTemplateUseCase(template_repository)
try:
template = use_case.execute(command)
except NotFoundError as _e:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found") from _e
except ValidationError as exc:
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc)) from exc
return _to_response(template)
@router.delete("/{template_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None, response_class=Response)
def delete_template(
template_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> Response:
user_id = authenticated_user.user.id
use_case = DeleteTemplateUseCase(template_repository)
deleted = use_case.execute(template_id, user_id)
if not deleted:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
return
@router.post("/{template_id}/copy", response_model=TemplateResponse, status_code=status.HTTP_201_CREATED)
def copy_template(
template_id: str,
request: CopyTemplateRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateResponse:
"""复制模板(含所有片段配置)"""
user_id = authenticated_user.user.id
command = CopyTemplateCommand(
template_id=template_id,
user_id=user_id,
new_name=request.new_name,
)
use_case = CopyTemplateUseCase(template_repository)
try:
template = use_case.execute(command)
except NotFoundError as _e:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found") from _e
except ValidationError as exc:
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc)) from exc
return _to_response(template)
@router.get("/{template_id}/usage", response_model=TemplateUsageResponse)
def get_template_usage(
template_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> TemplateUsageResponse:
"""获取模板使用次数(关联的剪辑计划数量)"""
user_id = authenticated_user.user.id
# 鉴权:确保模板存在且属于当前用户
use_case = GetTemplateUseCase(template_repository)
template = use_case.execute(template_id, user_id)
if template is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
usage = template_repository.get_usage_count(template_id)
return TemplateUsageResponse(template_id=template_id, usage_count=usage)
@router.post("/{template_id}/toggle-favorite", response_model=ToggleFavoriteResponse)
def toggle_favorite(
template_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ToggleFavoriteResponse:
"""切换模板收藏状态(当前为兼容端点,始终返回 false)"""
user_id = authenticated_user.user.id
use_case = GetTemplateUseCase(template_repository)
try:
template = use_case.execute(template_id, user_id)
except Exception as _e:
logger.exception("toggle_favorite 查询失败: template_id=%s", template_id)
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found") from _e
if template is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
return ToggleFavoriteResponse(id=template_id, is_favorite=False)
# ── Validate template ──
@router.post("/{template_id}/validate", response_model=ValidateTemplateResponse)
def validate_template(
template_id: str,
request: ValidateTemplateRequest = ValidateTemplateRequest(),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ValidateTemplateResponse:
user_id = authenticated_user.user.id
command = ValidateTemplateCommand(
template_id=template_id,
user_id=user_id,
voiceover_duration=request.voiceover_duration,
)
use_case = ValidateTemplateUseCase(template_repository)
try:
result = use_case.execute(command)
except NotFoundError as _e:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found") from _e
except ValidationError as exc:
raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc)) from exc
return ValidateTemplateResponse(
template=_to_response(result.template),
warnings=[GenerateWarningResponse(code=w.code, message=w.message, details=w.details) for w in result.warnings],
)
# ── Category CRUD ──
@router.get("/categories/list", response_model=ListCategoriesResponse)
def list_categories(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ListCategoriesResponse:
user_id = authenticated_user.user.id
try:
use_case = ListCategoriesUseCase(template_repository)
categories = use_case.execute(user_id)
except Exception:
logger.exception("list_categories 查询失败: user_id=%s", user_id)
return ListCategoriesResponse(items=[])
return ListCategoriesResponse(
items=[CategoryResponse(id=c.id, user_id=c.user_id, name=c.name, created_at=c.created_at) for c in categories],
)
@router.post("/categories", response_model=CategoryResponse, status_code=status.HTTP_201_CREATED)
def create_category(
request: CreateCategoryRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> CategoryResponse:
user_id = authenticated_user.user.id
command = CreateCategoryCommand(user_id=user_id, name=request.name)
use_case = CreateCategoryUseCase(template_repository)
category = use_case.execute(command)
return CategoryResponse(
id=category.id,
user_id=category.user_id,
name=category.name,
created_at=category.created_at,
)
@router.delete(
"/categories/{category_id}", status_code=status.HTTP_204_NO_CONTENT, response_model=None, response_class=Response
)
def delete_category(
category_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> Response:
user_id = authenticated_user.user.id
use_case = DeleteCategoryUseCase(template_repository)
deleted = use_case.execute(category_id, user_id)
if not deleted:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Category not found")
return Response(status_code=204)
# ── Tags ──
@router.get("/tags/list", response_model=ListTagsResponse)
def list_tags(
authenticated_user: AuthenticatedUser = Depends(get_current_user),
template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository),
) -> ListTagsResponse:
"""获取用户所有模板标签(去重排序)"""
user_id = authenticated_user.user.id
try:
use_case = ListTagsUseCase(template_repository)
tags = use_case.execute(user_id)
except Exception:
logger.exception("list_tags 查询失败: user_id=%s", user_id)
return ListTagsResponse(items=[])
return ListTagsResponse(items=tags)
skip = (page - 1) * page_size
templates = list_uc.execute(user_id, skip=skip, limit=page_size, filter=filters)
total = count_uc.execute(user_id, filter=filters)
items = [TemplateResponse.model_validate(tpl, from_attributes=True) for tpl in templates]
return ListTemplatesResponse(items=items, total=total)
@@ -622,8 +622,10 @@ def create_clips_from_assets_editor(
"""从素材批量创建片段(按模板segment配置创建,MediaKit异步更新).
逻辑:
1. 从模板读取 segments,片段数量 = segment 数量(忽略前端传的 required_clips_count)
2. 每个片段时长在 segment 的 duration_min ~ duration_max 之间随机取值(保留一位小数)
1. 从模板读取 segments,片段数量优先级:显式 clip_count(1-10)→ 旧字段
required_clips_count(兼容,超10截断)→ 默认 3(产品默认 3 段)。
片段数大于模板 segment 数时按顺序循环复用 segment 配置。
2. 每个片段时长在对应 segment 的 duration_min ~ duration_max 之间随机取值(保留一位小数)
3. 素材按片段顺序轮询分配,素材不够时同一素材切多个片段
4. 使用 replace_all_clips_transactional 原子性地清空旧片段并创建新的(随机起始时间)
5. 立即返回响应(目标 <1秒)
@@ -648,6 +650,22 @@ def create_clips_from_assets_editor(
detail="模板未配置片段",
)
# 1.5 归一化片段数量:
# 优先级:显式 clip_count → 旧字段 required_clips_count(由 schema 归一化到 clip_count)
# → 默认 3(产品默认 3 段)。按 N 循环复用 segment 配置;N <= len(segments) 时截取前 N 个
# (保持向后兼容:原模板有 N 个 segment、前端不传 clip_count 且 N<=10 时按模板段数创建;
# 默认模板仅有 1 个通用 segment 时按 clip_count=3 循环生成 3 段)。
requested_clip_count = getattr(body, "clip_count", None)
if requested_clip_count is None:
# schema 未显式传 clip_count 且无 legacy:使用模板 segments 数量,若超出 10 则截断
requested_clip_count = len(segments) if 1 <= len(segments) <= 10 else 3
requested_clip_count = max(1, min(int(requested_clip_count), 10))
effective_segments: list[tuple[int, float, float]] = []
for i in range(requested_clip_count):
src = segments[i % len(segments)]
effective_segments.append((i, float(src[1]), float(src[2])))
segments = effective_segments
# 防御:schema validator 已过滤 null/空串,这里再归一化一次,
# 避免异常入参(undefined → null)导致后续 /assets/{id} 404 / 422
asset_ids = [str(aid).strip() for aid in (body.asset_ids or []) if isinstance(aid, str) and aid.strip()]
@@ -8,7 +8,7 @@ from __future__ import annotations
import re as _re
from typing import Any, List, Optional
from pydantic import BaseModel, Field, validator
from pydantic import BaseModel, Field, model_validator, validator
_EXPORT_RESOLUTION_PATTERN = _re.compile(r"^\d+x\d+$")
_EXPORT_VALID_QUALITY_PRESETS = {"ultra_fast", "fast", "balanced", "high", "best"}
@@ -162,13 +162,26 @@ class ClipBatchDeleteResponse(BaseModel):
message: str = ""
# sentinel:区分「前端未传 clip_count」和「显式传 0/None」
_UNSET = object()
class ClipsFromAssetsRequest(BaseModel):
"""从素材批量创建片段请求"""
asset_ids: List[str] = Field(..., min_length=1, max_length=200, description="素材 ID 列表,按顺序追加到时间线末尾")
clip_type: str = Field(default="main", description="片段类型,默认 main")
clip_count: Optional[int] = Field(
default=None,
ge=1,
le=10,
description="片段数量(1-10);不传时使用旧字段 required_clips_count;两者都不传时回退为模板 segments 数量(默认 3 段)。",
)
required_clips_count: Optional[int] = Field(
default=None, ge=1, le=200, description="要求创建的片段数量;不传则等于素材数量"
default=None,
ge=1,
le=200,
description="[已废弃] 旧字段,请使用 clip_count;仅作向后兼容——clip_count 未显式传入时才回退本字段(超10截断到10)。",
)
@validator("asset_ids", pre=True)
@@ -180,6 +193,23 @@ class ClipsFromAssetsRequest(BaseModel):
return v
return [x for x in v if isinstance(x, str) and x.strip()]
@model_validator(mode="before")
@classmethod
def _backfill_clip_count(cls, data: Any) -> Any:
"""兼容旧字段 required_clips_count:仅当新字段 clip_count 未显式传入时才回退旧字段;
两者都没传时保持 clip_count=None,路由层按模板 segments 数量兜底。旧字段超 10 截断到 10。"""
if not isinstance(data, dict):
return data
has_new = "clip_count" in data and data["clip_count"] is not None
if not has_new:
legacy = data.get("required_clips_count")
if legacy is not None:
try:
data["clip_count"] = max(1, min(int(legacy), 10))
except (TypeError, ValueError):
pass
return data
class ClipsFromAssetsResponse(BaseModel):
"""从素材批量创建片段响应"""
+9 -71
View File
@@ -1,4 +1,9 @@
"""Template API schemas."""
"""Template API schemas(精简版:仅保留列表接口 + 默认模板自动兜底所需字段).
前端 PR#1911 删除 my-templates / editing-planner / templates 管理页后,
模板 CRUD / 分类 / 标签 / 收藏 / 复制 / 校验 / 使用统计等端点全部下线,
对应 Request/Response 模型也一并清理。
"""
from __future__ import annotations
@@ -50,17 +55,12 @@ class TemplateResponse(BaseModel):
updated_at: datetime
class ToggleFavoriteResponse(BaseModel):
id: str
is_favorite: bool
class ListTemplatesResponse(BaseModel):
items: List[TemplateResponse]
total: int = 0
# ── Template Request ──
# ── Template Request(保留给内部 _get_or_create_default_template_id 兜底创建默认模板使用)──
class CreateTemplateRequest(BaseModel):
@@ -75,71 +75,9 @@ class CreateTemplateRequest(BaseModel):
segments: List[SegmentRequest] = Field(default_factory=list)
class UpdateTemplateRequest(BaseModel):
name: Optional[str] = None
mode: Optional[str] = None
category: Optional[str] = None
tags: Optional[List[str]] = None
title_config: Optional[Dict[str, Any]] = None
subtitle_config: Optional[Dict[str, Any]] = None
bgm_config: Optional[Dict[str, Any]] = None
estimated_duration: Optional[float] = None
segments: Optional[List[SegmentRequest]] = None
# ── Validate ──
class ValidateTemplateRequest(BaseModel):
voiceover_duration: Optional[float] = None # 配音实际时长(秒)
class GenerateWarningResponse(BaseModel):
"""兼容老 import(如校验逻辑内部复用);模板管理页已下线,可按需进一步清理。"""
code: str
message: str
details: Dict[str, Any] = Field(default_factory=dict)
class ValidateTemplateResponse(BaseModel):
template: TemplateResponse
warnings: List[GenerateWarningResponse] = Field(default_factory=list)
# ── Category ──
class CategoryResponse(BaseModel):
id: str
user_id: str
name: str
created_at: datetime
class CreateCategoryRequest(BaseModel):
name: str
class ListCategoriesResponse(BaseModel):
items: List[CategoryResponse]
# ── Copy Template ──
class CopyTemplateRequest(BaseModel):
new_name: str
# ── Tags ──
class ListTagsResponse(BaseModel):
items: List[str]
# ── Usage Stats ──
class TemplateUsageResponse(BaseModel):
template_id: str
usage_count: int
+52
View File
@@ -0,0 +1,52 @@
-- 清理 e2e 测试账号产生的模板数据(#1900 模板管理页下线后一并清理历史脏数据)
--
-- 使用方法:连接到 staging / production 数据库后执行,例如:
-- psql "$DATABASE_URL" -f scripts/cleanup_e2e_test_data.sql
--
-- 删除范围:
-- 1. 用户账号 email 以 'e2e-gen-' 或 'asset-create_' 开头
-- 2. 这些用户创建的 templates / template_segments / template_categories 记录
-- 本脚本使用事务 + CTE,对生产数据无副作用。
BEGIN;
-- 1. 先收集要清理的测试账号 user_id
WITH test_users AS (
SELECT id AS user_id
FROM users
WHERE email LIKE 'e2e-gen-%'
OR email LIKE 'asset-create_%'
OR email LIKE 'e2e_%'
),
-- 2. 这些账号创建的模板
tpl_ids AS (
SELECT id AS template_id
FROM templates
WHERE user_id IN (SELECT user_id FROM test_users)
),
-- 3. 连带清理片段配置
del_segs AS (
DELETE FROM template_segments
WHERE template_id IN (SELECT template_id FROM tpl_ids)
RETURNING id
),
del_tpls AS (
DELETE FROM templates
WHERE id IN (SELECT template_id FROM tpl_ids)
RETURNING id
),
del_cats AS (
DELETE FROM template_categories
WHERE user_id IN (SELECT user_id FROM test_users)
RETURNING id
)
SELECT
(SELECT count(*) FROM test_users) AS users_matched,
(SELECT count(*) FROM del_segs) AS segments_deleted,
(SELECT count(*) FROM del_tpls) AS templates_deleted,
(SELECT count(*) FROM del_cats) AS categories_deleted;
-- 如需同时删除测试账号本身,取消下面的注释(默认保留账号只删模板数据):
-- DELETE FROM users WHERE id IN (SELECT user_id FROM test_users);
COMMIT;
@@ -1,349 +0,0 @@
"""
模板分类 CRUD API 集成测试。
覆盖端点:
- GET /templates/categories/list — 列出分类
- POST /templates/categories — 创建分类
- DELETE /templates/categories/{category_id} — 删除分类
使用 FastAPI TestClient + dependency_overrides 模式,
mock template repository,验证分类 CRUD 行为。
"""
from __future__ import annotations
import os
import sys
from datetime import datetime, timezone
from uuid import uuid4
# ── 环境变量 & sys.path(必须在导入 app.* 之前设置) ──────────────────────────
os.environ.setdefault("JWT_SECRET_KEY", "unit-test-secret-key-for-testing")
os.environ.setdefault("DATABASE_URL", "sqlite:///test.db")
import pytest
from fastapi import FastAPI
from fastapi.testclient import TestClient
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "..", "apps", "api"))
from app.api.routes import templates as templates_module
from app.auth import AuthenticatedUser, get_current_user
from packages.domain.entities import User
from packages.domain.template import TemplateCategory
# ---------------------------------------------------------------------------
# 1. 内存 Repository
# ---------------------------------------------------------------------------
class InMemoryTemplateRepository:
"""内存中的模板 Repository,仅实现分类相关方法。"""
def __init__(self):
self._categories: dict[str, TemplateCategory] = {}
self._templates = {}
self._segments = {}
# ── 分类相关 ──
def list_categories(self, user_id: str) -> list[TemplateCategory]:
return [c for c in self._categories.values() if c.user_id == user_id]
def create_category(self, category: TemplateCategory) -> TemplateCategory:
# 检查重复名称
existing = [c for c in self._categories.values() if c.user_id == category.user_id and c.name == category.name]
if existing:
raise ValueError(f"分类名称已存在: {category.name}")
self._categories[category.id] = category
return category
def get_category(self, category_id: str, user_id: str) -> TemplateCategory | None:
cat = self._categories.get(category_id)
if cat and cat.user_id == user_id:
return cat
return None
def delete_category(self, category_id: str, user_id: str) -> bool:
cat = self.get_category(category_id, user_id)
if cat:
del self._categories[category_id]
return True
return False
# ── 模板相关(路由可能调用,提供占位实现) ──
def list_by_user(self, user_id: str, *, skip: int = 0, limit: int = 50):
return []
def get(self, template_id: str, user_id: str):
return None
def create(self, template):
return template
def update(self, template):
return template
def delete(self, template_id: str, user_id: str) -> bool:
return False
def count_by_user(self, user_id: str) -> int:
return 0
def list_segments(self, template_id: str):
return []
def create_segments(self, segments):
return segments
def delete_segments_by_template(self, template_id: str) -> int:
return 0
def validate_template(self, *args, **kwargs):
return None
# ---------------------------------------------------------------------------
# 2. 辅助函数
# ---------------------------------------------------------------------------
def _make_user(**overrides) -> User:
defaults = dict(
id="user-test-001",
email="test@example.com",
display_name="Test User",
username="testuser",
subscription_plan="free",
subscription_status="active",
max_projects=3,
max_storage_gb=10,
created_at=datetime(2026, 1, 1, tzinfo=timezone.utc),
)
defaults.update(overrides)
return User(**defaults)
def _make_category(
name: str,
user_id: str = "user-test-001",
) -> TemplateCategory:
return TemplateCategory(
id=uuid4().hex,
user_id=user_id,
name=name,
created_at=datetime.now(timezone.utc),
)
# ---------------------------------------------------------------------------
# 3. Fixtures
# ---------------------------------------------------------------------------
@pytest.fixture
def template_repo():
return InMemoryTemplateRepository()
@pytest.fixture
def client(template_repo):
"""创建带有依赖覆盖的 TestClient。"""
test_app = FastAPI()
test_app.include_router(templates_module.router, prefix="/templates")
def _override_current_user():
return AuthenticatedUser(user=_make_user())
def _override_template_repo():
return template_repo
test_app.dependency_overrides[get_current_user] = _override_current_user
# 覆盖路由模块内的 _get_template_repository 依赖
test_app.dependency_overrides[templates_module._get_template_repository] = _override_template_repo
yield TestClient(test_app)
test_app.dependency_overrides.clear()
# ---------------------------------------------------------------------------
# 4. GET /categories/list — 列出分类
# ---------------------------------------------------------------------------
class TestListCategories:
"""列出分类端点测试。"""
def test_empty_list(self, client):
"""无分类时返回空列表。"""
resp = client.get("/templates/categories/list")
assert resp.status_code == 200
data = resp.json()
assert data["items"] == []
def test_returns_user_categories(self, client, template_repo):
"""只返回当前用户的分类。"""
c1 = _make_category("美食", "user-test-001")
c2 = _make_category("旅行", "user-test-001")
c3 = _make_category("科技", "other-user")
template_repo.create_category(c1)
template_repo.create_category(c2)
template_repo.create_category(c3)
resp = client.get("/templates/categories/list")
assert resp.status_code == 200
data = resp.json()
assert len(data["items"]) == 2
names = {item["name"] for item in data["items"]}
assert names == {"美食", "旅行"}
def test_response_fields(self, client, template_repo):
"""响应包含所有必需字段。"""
c = _make_category("测试分类")
template_repo.create_category(c)
resp = client.get("/templates/categories/list")
item = resp.json()["items"][0]
assert "id" in item
assert "user_id" in item
assert "name" in item
assert "created_at" in item
# ---------------------------------------------------------------------------
# 5. POST /categories — 创建分类
# ---------------------------------------------------------------------------
class TestCreateCategory:
"""创建分类端点测试。"""
def test_create_valid_category(self, client):
"""使用有效名称创建分类应成功。"""
resp = client.post("/templates/categories", json={"name": "vlog"})
assert resp.status_code == 201
data = resp.json()
assert data["name"] == "vlog"
assert "id" in data
assert data["user_id"] == "user-test-001"
assert "created_at" in data
def test_create_with_chinese_name(self, client):
"""支持中文分类名称。"""
resp = client.post("/templates/categories", json={"name": "美食探店"})
assert resp.status_code == 201
assert resp.json()["name"] == "美食探店"
def test_create_persists_to_repo(self, client, template_repo):
"""创建后分类保存到 repository。"""
resp = client.post("/templates/categories", json={"name": "新知识"})
cat_id = resp.json()["id"]
saved = template_repo.get_category(cat_id, "user-test-001")
assert saved is not None
assert saved.name == "新知识"
def test_create_missing_name_returns_422(self, client):
"""缺少 name 字段返回 422。"""
resp = client.post("/templates/categories", json={})
assert resp.status_code == 422
def test_create_empty_name_returns_422(self, client):
"""空名称返回 422(Pydantic min_length 校验)。"""
resp = client.post("/templates/categories", json={"name": ""})
# CreateCategoryRequest 没有 min_length 限制,此处验证实际行为
assert resp.status_code in (201, 422)
def test_create_multiple_categories(self, client, template_repo):
"""可创建多个不同名称的分类。"""
names = ["美食", "旅行", "科技", "教育", "娱乐"]
for name in names:
resp = client.post("/templates/categories", json={"name": name})
assert resp.status_code == 201
all_cats = template_repo.list_categories("user-test-001")
assert len(all_cats) == 5
# ---------------------------------------------------------------------------
# 6. DELETE /categories/{category_id} — 删除分类
# ---------------------------------------------------------------------------
class TestDeleteCategory:
"""删除分类端点测试。"""
def test_delete_existing_category(self, client, template_repo):
"""删除存在的分类返回 204。"""
c = _make_category("待删除")
template_repo.create_category(c)
resp = client.delete(f"/templates/categories/{c.id}")
assert resp.status_code == 204
# 验证已删除
assert template_repo.get_category(c.id, "user-test-001") is None
def test_delete_nonexistent_returns_404(self, client):
"""删除不存在的分类返回 404。"""
resp = client.delete("/templates/categories/nonexistent-id")
assert resp.status_code == 404
assert "not found" in resp.json()["detail"].lower() or "Category" in resp.json()["detail"]
def test_delete_other_user_category_returns_404(self, client, template_repo):
"""删除其他用户的分类返回 404(安全隔离)。"""
c = _make_category("他人分类", user_id="other-user")
template_repo.create_category(c)
resp = client.delete(f"/templates/categories/{c.id}")
assert resp.status_code == 404
# 验证未被删除
assert template_repo.get_category(c.id, "other-user") is not None
def test_delete_idempotent(self, client, template_repo):
"""删除后再次删除返回 404。"""
c = _make_category("幂等测试")
template_repo.create_category(c)
resp1 = client.delete(f"/templates/categories/{c.id}")
assert resp1.status_code == 204
resp2 = client.delete(f"/templates/categories/{c.id}")
assert resp2.status_code == 404
# ---------------------------------------------------------------------------
# 7. 跨端点场景
# ---------------------------------------------------------------------------
class TestCategoryCrudFlow:
"""分类 CRUD 完整流程。"""
def test_create_list_delete_flow(self, client, template_repo):
"""创建 → 列表 → 删除 完整流程。"""
# 1. 创建
create_resp = client.post("/templates/categories", json={"name": "流程测试"})
assert create_resp.status_code == 201
cat_id = create_resp.json()["id"]
# 2. 列表验证
list_resp = client.get("/templates/categories/list")
assert list_resp.status_code == 200
assert len(list_resp.json()["items"]) == 1
assert list_resp.json()["items"][0]["name"] == "流程测试"
# 3. 删除
del_resp = client.delete(f"/templates/categories/{cat_id}")
assert del_resp.status_code == 204
# 4. 再次列表验证已删除
list_resp2 = client.get("/templates/categories/list")
assert list_resp2.json()["items"] == []
if __name__ == "__main__":
pytest.main([__file__, "-v"])
+28 -27
View File
@@ -1,7 +1,7 @@
"""测试编辑器 from-assets 端点:按模板segment创建片段 + 事务性替换 + 随机起始.
覆盖:
- 片段数量 = segment 数量(required_clips_count 被忽略)
- 片段数量优先使用 clip_count(默认3,1-10),未传时回退旧字段 required_clips_count,再未传回退模板 segment 数量
- 素材不足时同一素材轮询切多个片段
- 随机 start_time + used_segments 去重
- 素材时长不足时 clip duration 缩短
@@ -140,7 +140,7 @@ class TestEditorClipsBySegments:
@patch("app.api.routes.templates_editor.clips.get_storage_service")
def test_creates_clips_matching_segment_count(self, mock_storage):
"""4 个 segment 即使只有2个素材也创建4个片段,required_clips_count 被忽略。"""
"""未显式传 clip_count/required_clips_count 时回退模板 segment 数量:4 个 segment 即使只有2个素材也创建4个片段。"""
from app.api.routes.templates_editor.clips import (
create_clips_from_assets_editor,
)
@@ -150,7 +150,7 @@ class TestEditorClipsBySegments:
mock_asset_repo = MagicMock()
mock_asset_repo.get = MagicMock(side_effect=lambda aid: _make_mock_asset(aid, {"a1": 30.0, "a2": 20.0}[aid]))
body = ClipsFromAssetsRequest(asset_ids=["a1", "a2"], required_clips_count=2)
body = ClipsFromAssetsRequest(asset_ids=["a1", "a2"])
# 均衡分配由 use_count 贪心保证,消除排序噪声后确定性断言
with _patch_zero_noise(), _patch_segments(DEFAULT_SEGMENTS):
@@ -657,17 +657,17 @@ class TestReuseRatioGate:
@patch("app.api.routes.templates_editor.clips.get_storage_service")
def test_reused_clip_ratio_within_threshold(self, mock_storage):
"""素材 60s、片段 5s:前 12 个用空闲区间,第 13 个复用,
复用占比 5/(12*5+5)=7.7% ≤ 15%,正常创建 13 个片段。"""
"""素材 45s、片段 5s、clip_count=10:前 9 个用空闲区间,第 10 个复用,
复用占比 ~5/(9*5+5)=10% ≤ 阈值,正常创建 10 个片段(clip_count 上限 10)。"""
from app.api.routes.templates_editor.clips import (
create_clips_from_assets_editor,
)
from app.api.routes.templates_editor.schemas import ClipsFromAssetsRequest
mock_plan_svc = _make_plan_svc(replace_return_count=13)
mock_plan_svc = _make_plan_svc(replace_return_count=10)
mock_asset_repo = MagicMock()
mock_asset_repo.get = MagicMock(return_value=_make_mock_asset("a1", 60.0))
body = ClipsFromAssetsRequest(asset_ids=["a1"], required_clips_count=13)
mock_asset_repo.get = MagicMock(return_value=_make_mock_asset("a1", 45.0))
body = ClipsFromAssetsRequest(asset_ids=["a1"], clip_count=10)
reused: dict = {}
@@ -676,11 +676,11 @@ class TestReuseRatioGate:
reused[aid] = reused.get(aid, 0.0) + dur
return (0.0, dur)
# 前 12 次分配空闲起点;第 13 次 calc 直接走回调(normal_starts 越界 → None → 回调)
normal_starts = [float(i * 5) for i in range(12)]
# 前 9 次分配空闲起点;第 10 次 calc 走回调(normal_starts 越界 → None → 回调)
normal_starts = [float(i * 5) for i in range(9)]
fake_calc, _ = self._make_calc_with_reuse(normal_starts, reused)
with (
_patch_segments(_segments(13, dur_min=5.0, dur_max=5.0)),
_patch_segments(_segments(10, dur_min=5.0, dur_max=5.0)),
patch("app.api.routes.templates_editor.clips._calc_random_start_time", side_effect=fake_calc),
patch("app.api.routes.templates_editor.clips.make_reuse_callback", return_value=reuse_cb),
):
@@ -695,17 +695,18 @@ class TestReuseRatioGate:
current_user=_make_auth_user(),
)
clips_data = _get_clips_data_from_call(mock_plan_svc)
assert len(clips_data) == 13
# 1 个复用片段,占比 1/13 ≈ 7.7% ≤ 15%
# 转场补偿: raw_duration = 5.0 + (13-1)*0.5/13 ≈ 5.462 → round(5.462,1) = 5.5
assert abs(reused.get("a1", 0.0) - 5.5) < 0.1
assert len(clips_data) == 10
# 回调复用被触发且总复用时长受控(≤ 10% 闸门允许范围内,保留少量余量)
assert reused.get("a1", 0.0) > 0
total_dur = sum(float(c.get("duration") or 5.0) for c in clips_data)
assert reused.get("a1", 0.0) / max(total_dur, 1.0) <= 0.12
@patch("app.api.routes.templates_editor.clips.get_storage_service")
def test_reuse_ratio_exceeded_returns_400(self, mock_storage):
"""复用占比将超 15% 时回调拒绝复用 → 无可用素材 → 400「素材可切区间不足」。
"""复用占比将超 10% 时回调拒绝复用 → 无可用素材 → 400「素材可切区间不足」。
60s 素材、5s 片段:前 12 个空闲、随后复用占比累计;当 (reused+d)/(assigned+d)
超过 15% 时回调返回 None,calc 返回 None,轮询无素材 → 400。
40s 素材、5s 片段、clip_count=10:前 8 个空闲(40s/5s),第 9 个尝试复用,
复用占比 5/(8*5+5)≈11.1% > 10% 闸门拒绝 → 无可用素材 → 400。
"""
from app.api.routes.templates_editor.clips import (
create_clips_from_assets_editor,
@@ -714,25 +715,25 @@ class TestReuseRatioGate:
mock_plan_svc = _make_plan_svc(replace_return_count=0)
mock_asset_repo = MagicMock()
mock_asset_repo.get = MagicMock(return_value=_make_mock_asset("a1", 60.0))
body = ClipsFromAssetsRequest(asset_ids=["a1"], required_clips_count=20)
mock_asset_repo.get = MagicMock(return_value=_make_mock_asset("a1", 40.0))
body = ClipsFromAssetsRequest(asset_ids=["a1"], clip_count=10)
# 模拟真实回调:累计复用时长,预判超 15% 拒绝
# 模拟真实回调:累计复用时长,预判超 10% 拒绝(对应 REUSE_RATIO_LIMIT=0.10)
reused: dict = {}
assigned: dict = {}
def fake_reuse_cb(aid, clip_duration):
a = assigned.get(aid, 0.0)
r = reused.get(aid, 0.0)
if a > 0 and (r + clip_duration) / (a + clip_duration) > 0.15:
if a > 0 and (r + clip_duration) / (a + clip_duration) > 0.10:
return None # 占比闸门拒绝
reused[aid] = r + clip_duration
return (0.0, clip_duration)
def fake_calc(asset_id, clip_duration, durations, used_segments, on_exhausted=None):
a = assigned.get(asset_id, 0.0)
# 前 12 个片段(60s/5s)有空闲区间
if a < 60.0:
# 前 8 个片段(40s/5s)有空闲区间
if a < 40.0:
start = a
assigned[asset_id] = a + clip_duration
return start
@@ -745,7 +746,7 @@ class TestReuseRatioGate:
return None
with (
_patch_segments(_segments(20, dur_min=5.0, dur_max=5.0)),
_patch_segments(_segments(10, dur_min=5.0, dur_max=5.0)),
patch("app.api.routes.templates_editor.clips._calc_random_start_time", side_effect=fake_calc),
patch("app.api.routes.templates_editor.clips.make_reuse_callback", return_value=fake_reuse_cb),
):
@@ -763,8 +764,8 @@ class TestReuseRatioGate:
assert exc_info.value.status_code == 400
assert "素材可切区间不足" in exc_info.value.detail
# 闸门在复用占比达上限时拒绝:60s 空闲 + 至多 ~15% 复用
assert reused.get("a1", 0.0) <= 12.0 # 10.0 或 15.0 以内,不会无限复用
# 闸门在复用占比达上限时拒绝:40s 空闲 + 至多 ~10% 复用
assert reused.get("a1", 0.0) <= 8.0 # 不会无限复用
# 未创建任何片段(整批失败)
assert not mock_plan_svc.replace_all_clips_transactional.called
+32 -3
View File
@@ -354,8 +354,36 @@ def _get_clips_data(mock_plan_svc):
class TestFromAssetsByTemplateSegments:
"""测试 from-assets 按模板 segment 创建片段(V2 事务性替换)。"""
def test_creates_clips_matching_segment_count(self):
"""片段数量 = segment 数量,忽略 required_clips_count。"""
def test_creates_clips_matching_legacy_required_clips_count(self):
"""显式传旧字段 required_clips_count 时按该值创建片段(兼容前端老版本)。"""
from app.api.routes.templates_editor.clips import create_clips_from_assets_editor
from app.api.routes.templates_editor.schemas import ClipsFromAssetsRequest
segments = [(0, 3.0, 5.0), (1, 4.0, 8.0), (2, 2.0, 6.0), (3, 5.0, 10.0)]
mock_tpl_svc = _make_tpl_svc_with_segments(segments)
mock_plan_svc = _make_plan_svc(replace_return_count=2)
mock_asset_repo = MagicMock()
mock_asset_repo.get.return_value = _make_rich_asset("a1", 60.0)
body = ClipsFromAssetsRequest(asset_ids=["a1"], required_clips_count=2)
result = create_clips_from_assets_editor(
template_id="tmpl-1",
body=body,
background_tasks=MagicMock(),
plan_id="plan-1",
services=(mock_tpl_svc, mock_plan_svc),
asset_repo=mock_asset_repo,
db=MagicMock(),
current_user=_make_auth_user(),
)
assert result.created_count == 2
clips_data = _get_clips_data(mock_plan_svc)
assert len(clips_data) == 2
def test_creates_clips_matching_segment_count_when_no_clip_count(self):
"""未传 clip_count/required_clips_count 时回退模板 segment 数量。"""
from app.api.routes.templates_editor.clips import create_clips_from_assets_editor
from app.api.routes.templates_editor.schemas import ClipsFromAssetsRequest
@@ -366,7 +394,7 @@ class TestFromAssetsByTemplateSegments:
mock_asset_repo = MagicMock()
mock_asset_repo.get.return_value = _make_rich_asset("a1", 60.0)
body = ClipsFromAssetsRequest(asset_ids=["a1"], required_clips_count=2)
body = ClipsFromAssetsRequest(asset_ids=["a1"])
result = create_clips_from_assets_editor(
template_id="tmpl-1",
body=body,
@@ -629,6 +657,7 @@ class TestFromAssetsByTemplateSegments:
# 用 MagicMock 模拟 body,绕过 Pydantic schema 的 min_length 校验
mock_body = MagicMock()
mock_body.asset_ids = []
mock_body.clip_count = None
mock_body.required_clips_count = None
with pytest.raises(HTTPException) as exc_info: