From 62c8d1cdff0c63e65592521fc90e4c096df948ac Mon Sep 17 00:00:00 2001 From: Audit Bot Date: Mon, 29 Jun 2026 14:14:58 +0800 Subject: [PATCH 1/3] =?UTF-8?q?feat:=20=E5=89=AA=E8=BE=91=E8=AE=A1?= =?UTF-8?q?=E5=88=92=E7=BC=96=E8=BE=91=E5=99=A8=E5=90=8E=E7=AB=AF=20?= =?UTF-8?q?=E2=80=94=20=E6=A8=A1=E6=9D=BF=20CRUD=20+=20=E5=88=86=E7=B1=BB?= =?UTF-8?q?=20+=20=E7=94=9F=E6=88=90=E6=A0=A1=E9=AA=8C?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - domain: Template, TemplateSegment, TemplateCategory 实体 - ports: TemplateRepositoryPort Protocol - adapters: SQLAlchemyTemplateRepository + 3 个 ORM Model - application: CRUD use cases + GenerateFromTemplateUseCase - 业务规则: one_take=1片段, voice_over_mix=每段需material_type, 配音±30%警告 - schemas: Pydantic request/response models - routes: /api/v1/templates CRUD + /generate + /categories - alembic: 014_add_template_tables (templates/template_segments/template_categories) - tests: 21 个单元测试全部通过 --- alembic/versions/014_add_template_tables.py | 87 ++++ apps/api/app/api/router.py | 6 + apps/api/app/api/routes/templates.py | 289 +++++++++++++ apps/api/app/schemas/template.py | 111 +++++ packages/adapters/sqlalchemy_impl/models.py | 40 ++ .../sqlalchemy_impl/template_repository.py | 258 +++++++++++ packages/application/template/commands.py | 55 +++ packages/application/template/use_cases.py | 255 +++++++++++ packages/domain/template.py | 47 ++ packages/ports/template_repository.py | 22 + tests/unit/test_template_use_cases.py | 409 ++++++++++++++++++ 11 files changed, 1579 insertions(+) create mode 100644 alembic/versions/014_add_template_tables.py create mode 100644 apps/api/app/api/routes/templates.py create mode 100644 apps/api/app/schemas/template.py create mode 100644 packages/adapters/sqlalchemy_impl/template_repository.py create mode 100644 packages/application/template/commands.py create mode 100644 packages/application/template/use_cases.py create mode 100644 packages/domain/template.py create mode 100644 packages/ports/template_repository.py create mode 100644 tests/unit/test_template_use_cases.py diff --git a/alembic/versions/014_add_template_tables.py b/alembic/versions/014_add_template_tables.py new file mode 100644 index 000000000..3ae2ee4f8 --- /dev/null +++ b/alembic/versions/014_add_template_tables.py @@ -0,0 +1,87 @@ +"""Phase 3 - 剪辑计划模板:templates + template_segments + template_categories + +Revision ID: 014 +Revises: 013 +Create Date: 2026-06-29 + +This migration creates three new tables: +1. templates — 剪辑计划模板主表 +2. template_segments — 模板片段表 +3. template_categories — 模板分类表 +""" +from alembic import op +import sqlalchemy as sa + +# revision identifiers +revision = "014" +down_revision = "013" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + conn = op.get_bind() + + # ── 1. Create templates table ── + conn.execute(sa.text(""" + CREATE TABLE IF NOT EXISTS templates ( + id VARCHAR(36) PRIMARY KEY, + user_id VARCHAR(36) NOT NULL, + name VARCHAR(200) NOT NULL, + mode VARCHAR(30) NOT NULL, + category VARCHAR(100) NOT NULL DEFAULT '', + tags JSONB NOT NULL DEFAULT '[]', + title_config JSONB NOT NULL DEFAULT '{}', + subtitle_config JSONB NOT NULL DEFAULT '{}', + bgm_config JSONB NOT NULL DEFAULT '{}', + estimated_duration FLOAT NOT NULL DEFAULT 0.0, + is_active BOOLEAN NOT NULL DEFAULT TRUE, + created_at TIMESTAMP NOT NULL DEFAULT NOW(), + updated_at TIMESTAMP NOT NULL DEFAULT NOW() + ) + """)) + conn.execute(sa.text( + "CREATE INDEX IF NOT EXISTS ix_templates_user_id ON templates(user_id)" + )) + conn.execute(sa.text( + "CREATE INDEX IF NOT EXISTS ix_templates_mode ON templates(mode)" + )) + + # ── 2. Create template_segments table ── + conn.execute(sa.text(""" + CREATE TABLE IF NOT EXISTS template_segments ( + id VARCHAR(36) PRIMARY KEY, + template_id VARCHAR(36) NOT NULL, + segment_order INTEGER NOT NULL, + duration_min FLOAT NOT NULL, + duration_max FLOAT NOT NULL, + material_type VARCHAR(20), + created_at TIMESTAMP NOT NULL DEFAULT NOW(), + updated_at TIMESTAMP NOT NULL DEFAULT NOW() + ) + """)) + conn.execute(sa.text( + "CREATE INDEX IF NOT EXISTS ix_template_segments_template_id " + "ON template_segments(template_id)" + )) + + # ── 3. Create template_categories table ── + conn.execute(sa.text(""" + CREATE TABLE IF NOT EXISTS template_categories ( + id VARCHAR(36) PRIMARY KEY, + user_id VARCHAR(36) NOT NULL, + name VARCHAR(100) NOT NULL, + created_at TIMESTAMP NOT NULL DEFAULT NOW() + ) + """)) + conn.execute(sa.text( + "CREATE INDEX IF NOT EXISTS ix_template_categories_user_id " + "ON template_categories(user_id)" + )) + + +def downgrade() -> None: + conn = op.get_bind() + conn.execute(sa.text("DROP TABLE IF EXISTS template_categories")) + conn.execute(sa.text("DROP TABLE IF EXISTS template_segments")) + conn.execute(sa.text("DROP TABLE IF EXISTS templates")) diff --git a/apps/api/app/api/router.py b/apps/api/app/api/router.py index 93b94fd24..712580c93 100644 --- a/apps/api/app/api/router.py +++ b/apps/api/app/api/router.py @@ -8,6 +8,7 @@ from app.api.routes.duplication import router as duplication_router from app.api.routes.generated_videos import router as generated_videos_router from app.api.routes.recipes import router as recipes_router from app.api.routes.subscription import router as subscription_router +from app.api.routes.templates import router as templates_router from app.api.routes.titles import router as titles_router from app.api.routes.voices import router as voices_router from app.api.routes.generation_tasks import router as generation_tasks_router @@ -104,3 +105,8 @@ api_router.include_router( prefix="/recipes", tags=["Recipe"], ) +api_router.include_router( + templates_router, + prefix="/templates", + tags=["Template"], +) diff --git a/apps/api/app/api/routes/templates.py b/apps/api/app/api/routes/templates.py new file mode 100644 index 000000000..af4884876 --- /dev/null +++ b/apps/api/app/api/routes/templates.py @@ -0,0 +1,289 @@ +"""Template CRUD + generate + category routes.""" +from __future__ import annotations + +from typing import List + +from fastapi import APIRouter, Depends, HTTPException, Query, Response, status +from sqlalchemy.orm import Session + +from app.auth import AuthenticatedUser, get_current_user +from app.dependencies import get_db_session +from app.schemas.template import ( + CategoryResponse, + CreateCategoryRequest, + CreateTemplateRequest, + GenerateFromTemplateRequest, + GenerateFromTemplateResponse, + GenerateWarningResponse, + ListCategoriesResponse, + ListTemplatesResponse, + SegmentResponse, + TemplateResponse, + UpdateTemplateRequest, +) +from packages.adapters.sqlalchemy_impl.template_repository import SQLAlchemyTemplateRepository +from packages.application.template.commands import ( + CreateCategoryCommand, + CreateTemplateCommand, + GenerateFromTemplateCommand, + SegmentCommand, + UpdateTemplateCommand, +) +from packages.application.template.use_cases import ( + CreateCategoryUseCase, + CreateTemplateUseCase, + DeleteCategoryUseCase, + DeleteTemplateUseCase, + GenerateFromTemplateUseCase, + GetTemplateUseCase, + ListCategoriesUseCase, + ListTemplatesUseCase, + NotFoundError, + UpdateTemplateUseCase, + ValidationError, +) + +router = APIRouter() + + +def _get_template_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyTemplateRepository: + 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) -> 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, + created_at=template.created_at, + updated_at=template.updated_at, + ) + + +# ── Template CRUD ── + + +@router.get("", response_model=ListTemplatesResponse) +def list_templates( + skip: int = Query(0, ge=0), + limit: int = Query(50, ge=1, le=200), + authenticated_user: AuthenticatedUser = Depends(get_current_user), + template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository), +) -> ListTemplatesResponse: + user_id = authenticated_user.user.id + use_case = ListTemplatesUseCase(template_repository) + templates = use_case.execute(user_id, skip=skip, limit=limit) + total = template_repository.count_by_user(user_id) + return ListTemplatesResponse( + items=[_to_response(t) for t in templates], + total=total, + ) + + +@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 + 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") + return _to_response(template) + + +@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)) + 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: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found") + except ValidationError as exc: + raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc)) + return _to_response(template) + + +@router.delete("/{template_id}", status_code=status.HTTP_204_NO_CONTENT, 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 Response(status_code=204) + + +# ── Generate from template ── + + +@router.post("/{template_id}/generate", response_model=GenerateFromTemplateResponse) +def generate_from_template( + template_id: str, + request: GenerateFromTemplateRequest = GenerateFromTemplateRequest(), + authenticated_user: AuthenticatedUser = Depends(get_current_user), + template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository), +) -> GenerateFromTemplateResponse: + user_id = authenticated_user.user.id + command = GenerateFromTemplateCommand( + template_id=template_id, + user_id=user_id, + voiceover_duration=request.voiceover_duration, + ) + use_case = GenerateFromTemplateUseCase(template_repository) + try: + result = use_case.execute(command) + except NotFoundError: + raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found") + except ValidationError as exc: + raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc)) + + return GenerateFromTemplateResponse( + 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 + use_case = ListCategoriesUseCase(template_repository) + categories = use_case.execute(user_id) + 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_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) diff --git a/apps/api/app/schemas/template.py b/apps/api/app/schemas/template.py new file mode 100644 index 000000000..5f3bc1635 --- /dev/null +++ b/apps/api/app/schemas/template.py @@ -0,0 +1,111 @@ +"""Template API schemas.""" +from __future__ import annotations + +from datetime import datetime +from typing import Any, Dict, List, Optional + +from pydantic import BaseModel, Field + + +# ── Segment ── + +class SegmentResponse(BaseModel): + id: str + template_id: str + segment_order: int + duration_min: float + duration_max: float + material_type: Optional[str] = None + created_at: datetime + updated_at: datetime + + +class SegmentRequest(BaseModel): + segment_order: int + duration_min: float + duration_max: float + material_type: Optional[str] = None + + +# ── Template Response ── + +class TemplateResponse(BaseModel): + id: str + user_id: str + name: str + mode: str + category: str = "" + tags: List[str] = Field(default_factory=list) + title_config: Dict[str, Any] = Field(default_factory=dict) + subtitle_config: Dict[str, Any] = Field(default_factory=dict) + bgm_config: Dict[str, Any] = Field(default_factory=dict) + estimated_duration: float = 0.0 + segments: List[SegmentResponse] = Field(default_factory=list) + is_active: bool = True + created_at: datetime + updated_at: datetime + + +class ListTemplatesResponse(BaseModel): + items: List[TemplateResponse] + total: int = 0 + + +# ── Template Request ── + +class CreateTemplateRequest(BaseModel): + name: str + mode: str + category: str = "" + tags: List[str] = Field(default_factory=list) + title_config: Dict[str, Any] = Field(default_factory=dict) + subtitle_config: Dict[str, Any] = Field(default_factory=dict) + bgm_config: Dict[str, Any] = Field(default_factory=dict) + estimated_duration: float = 0.0 + 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 + + +# ── Generate ── + +class GenerateFromTemplateRequest(BaseModel): + voiceover_duration: Optional[float] = None # 配音实际时长(秒) + + +class GenerateWarningResponse(BaseModel): + code: str + message: str + details: Dict[str, Any] = Field(default_factory=dict) + + +class GenerateFromTemplateResponse(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] diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index f9ded57a3..f5854e025 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -280,3 +280,43 @@ class RecipeItemModel(Base): position = Column(Integer, nullable=False, default=0) extra_meta = Column('metadata', JSON, nullable=False, default=dict) + +class TemplateModel(Base): + __tablename__ = "templates" + + id = Column(String(36), primary_key=True) + user_id = Column(String(36), nullable=False, index=True) + name = Column(String(200), nullable=False) + mode = Column(String(30), nullable=False, index=True) # pip / person_narration / one_take / voice_over_mix + category = Column(String(100), nullable=False, default="") + tags = Column(JSON, nullable=False, default=list) + title_config = Column(JSON, nullable=False, default=dict) + subtitle_config = Column(JSON, nullable=False, default=dict) + bgm_config = Column(JSON, nullable=False, default=dict) + estimated_duration = Column(Float, nullable=False, default=0.0) + is_active = Column(Boolean, nullable=False, default=True) + created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) + updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) + + +class TemplateSegmentModel(Base): + __tablename__ = "template_segments" + + id = Column(String(36), primary_key=True) + template_id = Column(String(36), nullable=False, index=True) + segment_order = Column(Integer, nullable=False) + duration_min = Column(Float, nullable=False) + duration_max = Column(Float, nullable=False) + material_type = Column(String(20), nullable=True) # 仅 voice_over_mix: 人物/场景 + created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) + updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) + + +class TemplateCategoryModel(Base): + __tablename__ = "template_categories" + + id = Column(String(36), primary_key=True) + user_id = Column(String(36), nullable=False, index=True) + name = Column(String(100), nullable=False) + created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) + diff --git a/packages/adapters/sqlalchemy_impl/template_repository.py b/packages/adapters/sqlalchemy_impl/template_repository.py new file mode 100644 index 000000000..f255a2af7 --- /dev/null +++ b/packages/adapters/sqlalchemy_impl/template_repository.py @@ -0,0 +1,258 @@ +"""SQLAlchemy implementation of TemplateRepository.""" +from __future__ import annotations + +from typing import List, Optional + +from sqlalchemy.orm import Session + +from packages.adapters.sqlalchemy_impl.models import ( + TemplateCategoryModel, + TemplateModel, + TemplateSegmentModel, +) +from packages.domain.template import Template, TemplateCategory, TemplateSegment + + +class SQLAlchemyTemplateRepository: + """SQLAlchemy 剪辑计划模板仓储.""" + + def __init__(self, session: Session) -> None: + self.session = session + + # ── Template CRUD ── + + def list_by_user( + self, + user_id: str, + *, + skip: int = 0, + limit: int = 50, + ) -> List[Template]: + models = ( + self.session.query(TemplateModel) + .filter( + TemplateModel.user_id == user_id, + TemplateModel.is_active == True, + ) + .order_by(TemplateModel.created_at.desc()) + .offset(skip) + .limit(limit) + .all() + ) + templates = [self._model_to_entity(m) for m in models] + for t in templates: + t.segments = self.list_segments(t.id) + return templates + + def get(self, template_id: str, user_id: str) -> Optional[Template]: + model = ( + self.session.query(TemplateModel) + .filter( + TemplateModel.id == template_id, + TemplateModel.user_id == user_id, + ) + .first() + ) + if model is None: + return None + template = self._model_to_entity(model) + template.segments = self.list_segments(template.id) + return template + + def create(self, template: Template) -> Template: + model = TemplateModel( + 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, + is_active=template.is_active, + ) + self.session.add(model) + self.session.commit() + self.session.refresh(model) + result = self._model_to_entity(model) + result.segments = template.segments + return result + + def update(self, template: Template) -> Template: + model = ( + self.session.query(TemplateModel) + .filter( + TemplateModel.id == template.id, + TemplateModel.user_id == template.user_id, + ) + .first() + ) + if model is None: + raise ValueError(f"Template {template.id} not found") + model.name = template.name + model.mode = template.mode + model.category = template.category + model.tags = template.tags + model.title_config = template.title_config + model.subtitle_config = template.subtitle_config + model.bgm_config = template.bgm_config + model.estimated_duration = template.estimated_duration + model.is_active = template.is_active + self.session.commit() + self.session.refresh(model) + result = self._model_to_entity(model) + result.segments = template.segments + return result + + def delete(self, template_id: str, user_id: str) -> bool: + model = ( + self.session.query(TemplateModel) + .filter( + TemplateModel.id == template_id, + TemplateModel.user_id == user_id, + ) + .first() + ) + if model is None: + return False + model.is_active = False + self.session.commit() + return True + + def count_by_user(self, user_id: str) -> int: + return ( + self.session.query(TemplateModel) + .filter( + TemplateModel.user_id == user_id, + TemplateModel.is_active == True, + ) + .count() + ) + + # ── Segments ── + + def list_segments(self, template_id: str) -> List[TemplateSegment]: + models = ( + self.session.query(TemplateSegmentModel) + .filter(TemplateSegmentModel.template_id == template_id) + .order_by(TemplateSegmentModel.segment_order) + .all() + ) + return [self._segment_model_to_entity(m) for m in models] + + def create_segments(self, segments: List[TemplateSegment]) -> List[TemplateSegment]: + for seg in segments: + model = TemplateSegmentModel( + 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, + ) + self.session.add(model) + self.session.commit() + return segments + + def delete_segments_by_template(self, template_id: str) -> int: + count = ( + self.session.query(TemplateSegmentModel) + .filter(TemplateSegmentModel.template_id == template_id) + .delete() + ) + self.session.commit() + return count + + # ── Categories ── + + def list_categories(self, user_id: str) -> List[TemplateCategory]: + models = ( + self.session.query(TemplateCategoryModel) + .filter(TemplateCategoryModel.user_id == user_id) + .order_by(TemplateCategoryModel.created_at) + .all() + ) + return [self._category_model_to_entity(m) for m in models] + + def create_category(self, category: TemplateCategory) -> TemplateCategory: + model = TemplateCategoryModel( + id=category.id, + user_id=category.user_id, + name=category.name, + ) + self.session.add(model) + self.session.commit() + self.session.refresh(model) + return self._category_model_to_entity(model) + + def get_category(self, category_id: str, user_id: str) -> Optional[TemplateCategory]: + model = ( + self.session.query(TemplateCategoryModel) + .filter( + TemplateCategoryModel.id == category_id, + TemplateCategoryModel.user_id == user_id, + ) + .first() + ) + if model is None: + return None + return self._category_model_to_entity(model) + + def delete_category(self, category_id: str, user_id: str) -> bool: + model = ( + self.session.query(TemplateCategoryModel) + .filter( + TemplateCategoryModel.id == category_id, + TemplateCategoryModel.user_id == user_id, + ) + .first() + ) + if model is None: + return False + self.session.delete(model) + self.session.commit() + return True + + # ── Mapping helpers ── + + @staticmethod + def _model_to_entity(model: TemplateModel) -> Template: + return Template( + id=model.id, + user_id=model.user_id, + name=model.name, + mode=model.mode, + category=model.category or "", + tags=model.tags or [], + title_config=model.title_config or {}, + subtitle_config=model.subtitle_config or {}, + bgm_config=model.bgm_config or {}, + estimated_duration=model.estimated_duration or 0.0, + is_active=model.is_active, + created_at=model.created_at, + updated_at=model.updated_at, + ) + + @staticmethod + def _segment_model_to_entity(model: TemplateSegmentModel) -> TemplateSegment: + return TemplateSegment( + id=model.id, + template_id=model.template_id, + segment_order=model.segment_order, + duration_min=model.duration_min, + duration_max=model.duration_max, + material_type=model.material_type, + created_at=model.created_at, + updated_at=model.updated_at, + ) + + @staticmethod + def _category_model_to_entity(model: TemplateCategoryModel) -> TemplateCategory: + return TemplateCategory( + id=model.id, + user_id=model.user_id, + name=model.name, + created_at=model.created_at, + ) diff --git a/packages/application/template/commands.py b/packages/application/template/commands.py new file mode 100644 index 000000000..693df8282 --- /dev/null +++ b/packages/application/template/commands.py @@ -0,0 +1,55 @@ +"""Template commands.""" +from __future__ import annotations + +from dataclasses import dataclass, field +from typing import List, Optional + + +@dataclass +class SegmentCommand: + segment_order: int + duration_min: float + duration_max: float + material_type: Optional[str] = None + + +@dataclass +class CreateTemplateCommand: + user_id: str + name: str + mode: str + category: str = "" + tags: List[str] = field(default_factory=list) + title_config: dict = field(default_factory=dict) + subtitle_config: dict = field(default_factory=dict) + bgm_config: dict = field(default_factory=dict) + estimated_duration: float = 0.0 + segments: List[SegmentCommand] = field(default_factory=list) + + +@dataclass +class UpdateTemplateCommand: + template_id: str + user_id: str + name: Optional[str] = None + mode: Optional[str] = None + category: Optional[str] = None + tags: Optional[List[str]] = None + title_config: Optional[dict] = None + subtitle_config: Optional[dict] = None + bgm_config: Optional[dict] = None + estimated_duration: Optional[float] = None + segments: Optional[List[SegmentCommand]] = None + + +@dataclass +class CreateCategoryCommand: + user_id: str + name: str + + +@dataclass +class GenerateFromTemplateCommand: + template_id: str + user_id: str + voiceover_duration: Optional[float] = None # 配音实际时长(用于偏差校验) diff --git a/packages/application/template/use_cases.py b/packages/application/template/use_cases.py new file mode 100644 index 000000000..ab790a1ec --- /dev/null +++ b/packages/application/template/use_cases.py @@ -0,0 +1,255 @@ +"""Template use cases.""" +from __future__ import annotations + +import uuid +from dataclasses import dataclass, field +from typing import List, Optional + +from packages.adapters.sqlalchemy_impl.template_repository import SQLAlchemyTemplateRepository +from packages.application.template.commands import ( + CreateCategoryCommand, + CreateTemplateCommand, + GenerateFromTemplateCommand, + UpdateTemplateCommand, +) +from packages.domain.template import Template, TemplateCategory, TemplateSegment + + +class NotFoundError(Exception): + pass + + +class ValidationError(Exception): + """业务规则校验失败.""" + pass + + +VALID_MODES = {"pip", "person_narration", "one_take", "voice_over_mix"} +VALID_MATERIAL_TYPES = {"人物", "场景"} + + +@dataclass +class GenerateWarning: + """生成时的警告信息.""" + code: str # voiceover_duration_mismatch / missing_material_type / ... + message: str + details: dict = field(default_factory=dict) + + +@dataclass +class GenerateResult: + """使用模板生成的结果.""" + template: Template + warnings: List[GenerateWarning] = field(default_factory=list) + + +# ── Template CRUD ── + + +class CreateTemplateUseCase: + def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + self.repository = repository + + def execute(self, command: CreateTemplateCommand) -> Template: + if command.mode not in VALID_MODES: + raise ValidationError(f"无效的剪辑模式: {command.mode},可选值: {VALID_MODES}") + + template_id = uuid.uuid4().hex + template = Template( + id=template_id, + user_id=command.user_id, + name=command.name, + mode=command.mode, + category=command.category, + tags=command.tags, + title_config=command.title_config, + subtitle_config=command.subtitle_config, + bgm_config=command.bgm_config, + estimated_duration=command.estimated_duration, + ) + template = self.repository.create(template) + + if command.segments: + segments = [ + TemplateSegment( + id=uuid.uuid4().hex, + template_id=template.id, + segment_order=seg.segment_order, + duration_min=seg.duration_min, + duration_max=seg.duration_max, + material_type=seg.material_type, + ) + for seg in command.segments + ] + self.repository.create_segments(segments) + template.segments = segments + + return template + + +class ListTemplatesUseCase: + def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + self.repository = repository + + def execute( + self, + user_id: str, + *, + skip: int = 0, + limit: int = 50, + ) -> List[Template]: + return self.repository.list_by_user(user_id, skip=skip, limit=limit) + + +class GetTemplateUseCase: + def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + self.repository = repository + + def execute(self, template_id: str, user_id: str) -> Optional[Template]: + return self.repository.get(template_id, user_id) + + +class UpdateTemplateUseCase: + def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + self.repository = repository + + def execute(self, command: UpdateTemplateCommand) -> Template: + existing = self.repository.get(command.template_id, command.user_id) + if existing is None: + raise NotFoundError(f"Template {command.template_id} not found") + + if command.mode is not None and command.mode not in VALID_MODES: + raise ValidationError(f"无效的剪辑模式: {command.mode}") + + if command.name is not None: + existing.name = command.name + if command.mode is not None: + existing.mode = command.mode + if command.category is not None: + existing.category = command.category + if command.tags is not None: + existing.tags = command.tags + if command.title_config is not None: + existing.title_config = command.title_config + if command.subtitle_config is not None: + existing.subtitle_config = command.subtitle_config + if command.bgm_config is not None: + existing.bgm_config = command.bgm_config + if command.estimated_duration is not None: + existing.estimated_duration = command.estimated_duration + + self.repository.update(existing) + + # Replace segments if provided + if command.segments is not None: + self.repository.delete_segments_by_template(existing.id) + segments = [ + TemplateSegment( + id=uuid.uuid4().hex, + template_id=existing.id, + segment_order=seg.segment_order, + duration_min=seg.duration_min, + duration_max=seg.duration_max, + material_type=seg.material_type, + ) + for seg in command.segments + ] + self.repository.create_segments(segments) + existing.segments = segments + else: + existing.segments = self.repository.list_segments(existing.id) + + return existing + + +class DeleteTemplateUseCase: + def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + self.repository = repository + + def execute(self, template_id: str, user_id: str) -> bool: + return self.repository.delete(template_id, user_id) + + +# ── Generate from template ── + + +class GenerateFromTemplateUseCase: + """使用模板生成剪辑计划 — 校验业务规则.""" + + def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + self.repository = repository + + def execute(self, command: GenerateFromTemplateCommand) -> GenerateResult: + template = self.repository.get(command.template_id, command.user_id) + if template is None: + raise NotFoundError(f"Template {command.template_id} not found") + + warnings: List[GenerateWarning] = [] + + # 业务规则 1: one_take 必须恰好 1 个片段 + if template.mode == "one_take": + if len(template.segments) != 1: + raise ValidationError( + f"一镜到底模式必须恰好有 1 个片段,当前有 {len(template.segments)} 个" + ) + + # 业务规则 2: voice_over_mix 每个片段必须有 material_type + if template.mode == "voice_over_mix": + for seg in template.segments: + if not seg.material_type or seg.material_type not in VALID_MATERIAL_TYPES: + raise ValidationError( + f"口播+混剪模式下每个片段必须指定 material_type(人物/场景)," + f"片段 {seg.segment_order} 的 material_type 无效: {seg.material_type}" + ) + + # 业务规则 3: 配音时长偏差 ±30% 警告 + if command.voiceover_duration is not None and template.estimated_duration > 0: + ratio = command.voiceover_duration / template.estimated_duration + if ratio < 0.7 or ratio > 1.3: + warnings.append(GenerateWarning( + code="voiceover_duration_mismatch", + message=( + f"配音时长 ({command.voiceover_duration:.1f}s) " + f"与预估时长 ({template.estimated_duration:.1f}s) " + f"偏差超过 ±30%,可能影响剪辑效果" + ), + details={ + "voiceover_duration": command.voiceover_duration, + "estimated_duration": template.estimated_duration, + "ratio": round(ratio, 3), + }, + )) + + return GenerateResult(template=template, warnings=warnings) + + +# ── Category CRUD ── + + +class CreateCategoryUseCase: + def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + self.repository = repository + + def execute(self, command: CreateCategoryCommand) -> TemplateCategory: + category = TemplateCategory( + id=uuid.uuid4().hex, + user_id=command.user_id, + name=command.name, + ) + return self.repository.create_category(category) + + +class ListCategoriesUseCase: + def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + self.repository = repository + + def execute(self, user_id: str) -> List[TemplateCategory]: + return self.repository.list_categories(user_id) + + +class DeleteCategoryUseCase: + def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + self.repository = repository + + def execute(self, category_id: str, user_id: str) -> bool: + return self.repository.delete_category(category_id, user_id) diff --git a/packages/domain/template.py b/packages/domain/template.py new file mode 100644 index 000000000..8d0620743 --- /dev/null +++ b/packages/domain/template.py @@ -0,0 +1,47 @@ +"""Template domain entities — 剪辑计划模板.""" +from __future__ import annotations + +from dataclasses import dataclass, field +from datetime import datetime, timezone +from typing import List, Optional + + +@dataclass +class TemplateSegment: + """模板中的单个片段.""" + id: str + template_id: str + segment_order: int + duration_min: float + duration_max: float + material_type: Optional[str] = None # 仅 voice_over_mix: 人物/场景; 其他模式 null + created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) + updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) + + +@dataclass +class Template: + """剪辑计划模板.""" + id: str + user_id: str + name: str + mode: str # pip / person_narration / one_take / voice_over_mix + category: str = "" + tags: List[str] = field(default_factory=list) + title_config: dict = field(default_factory=dict) + subtitle_config: dict = field(default_factory=dict) + bgm_config: dict = field(default_factory=dict) + estimated_duration: float = 0.0 + segments: List[TemplateSegment] = field(default_factory=list) + is_active: bool = True + created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) + updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) + + +@dataclass +class TemplateCategory: + """模板分类.""" + id: str + user_id: str + name: str + created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc)) diff --git a/packages/ports/template_repository.py b/packages/ports/template_repository.py new file mode 100644 index 000000000..a93a7d4e7 --- /dev/null +++ b/packages/ports/template_repository.py @@ -0,0 +1,22 @@ +"""Template repository port (Protocol).""" +from __future__ import annotations + +from typing import List, Optional, Protocol + +from packages.domain.template import Template, TemplateCategory, TemplateSegment + + +class TemplateRepositoryPort(Protocol): + def list_by_user(self, user_id: str, *, skip: int = 0, limit: int = 50) -> List[Template]: ... + def get(self, template_id: str, user_id: str) -> Optional[Template]: ... + def create(self, template: Template) -> Template: ... + def update(self, template: Template) -> Template: ... + def delete(self, template_id: str, user_id: str) -> bool: ... + def count_by_user(self, user_id: str) -> int: ... + def list_segments(self, template_id: str) -> List[TemplateSegment]: ... + def create_segments(self, segments: List[TemplateSegment]) -> List[TemplateSegment]: ... + def delete_segments_by_template(self, template_id: str) -> int: ... + def list_categories(self, user_id: str) -> List[TemplateCategory]: ... + def create_category(self, category: TemplateCategory) -> TemplateCategory: ... + def get_category(self, category_id: str, user_id: str) -> Optional[TemplateCategory]: ... + def delete_category(self, category_id: str, user_id: str) -> bool: ... diff --git a/tests/unit/test_template_use_cases.py b/tests/unit/test_template_use_cases.py new file mode 100644 index 000000000..bc1b4e6d0 --- /dev/null +++ b/tests/unit/test_template_use_cases.py @@ -0,0 +1,409 @@ +""" +Template Use Cases 单元测试 — 剪辑计划模板 CRUD + 业务规则校验 +""" +from unittest.mock import MagicMock, Mock + +import pytest + +from packages.application.template.commands import ( + CreateCategoryCommand, + CreateTemplateCommand, + GenerateFromTemplateCommand, + SegmentCommand, + UpdateTemplateCommand, +) +from packages.application.template.use_cases import ( + CreateCategoryUseCase, + CreateTemplateUseCase, + DeleteTemplateUseCase, + GenerateFromTemplateUseCase, + GetTemplateUseCase, + ListCategoriesUseCase, + ListTemplatesUseCase, + NotFoundError, + UpdateTemplateUseCase, + ValidationError, +) +from packages.domain.template import Template, TemplateCategory, TemplateSegment + + +def _make_repo(): + """创建一个 mock repository.""" + repo = Mock() + repo.list_by_user = Mock(return_value=[]) + repo.get = Mock(return_value=None) + repo.create = Mock() + repo.update = Mock() + repo.delete = Mock(return_value=False) + repo.count_by_user = Mock(return_value=0) + repo.list_segments = Mock(return_value=[]) + repo.create_segments = Mock() + repo.delete_segments_by_template = Mock(return_value=0) + repo.list_categories = Mock(return_value=[]) + repo.create_category = Mock() + repo.get_category = Mock(return_value=None) + repo.delete_category = Mock(return_value=False) + return repo + + +def _make_template(**kwargs) -> Template: + defaults = dict( + id="tmpl-001", + user_id="user-001", + name="测试模板", + mode="pip", + category="default", + tags=["test"], + title_config={"ai_auto_select": True}, + subtitle_config={"enabled": True}, + bgm_config={"enabled": False}, + estimated_duration=60.0, + segments=[], + ) + defaults.update(kwargs) + return Template(**defaults) + + +# ── CreateTemplateUseCase ── + + +class TestCreateTemplateUseCase: + @pytest.fixture + def repo(self): + return _make_repo() + + @pytest.fixture + def use_case(self, repo): + return CreateTemplateUseCase(repo) + + def test_create_basic_template(self, use_case, repo): + """创建基础模板(无片段).""" + repo.create.side_effect = lambda t: t # 返回传入的 template + + command = CreateTemplateCommand( + user_id="user-001", + name="画中画模板", + mode="pip", + category="vlog", + tags=["vlog", "pip"], + estimated_duration=90.0, + ) + result = use_case.execute(command) + + assert result.name == "画中画模板" + assert result.mode == "pip" + assert result.user_id == "user-001" + repo.create.assert_called_once() + + def test_create_with_segments(self, use_case, repo): + """创建模板并附带片段.""" + repo.create.side_effect = lambda t: t + repo.create_segments.side_effect = lambda segs: segs + + command = CreateTemplateCommand( + user_id="user-001", + name="口播混剪模板", + mode="voice_over_mix", + segments=[ + SegmentCommand(segment_order=1, duration_min=5, duration_max=15, material_type="人物"), + SegmentCommand(segment_order=2, duration_min=10, duration_max=30, material_type="场景"), + ], + ) + result = use_case.execute(command) + + assert len(result.segments) == 2 + assert result.segments[0].material_type == "人物" + repo.create_segments.assert_called_once() + + def test_create_invalid_mode_raises(self, use_case): + """无效剪辑模式应抛出 ValidationError.""" + command = CreateTemplateCommand( + user_id="user-001", + name="无效模板", + mode="invalid_mode", + ) + with pytest.raises(ValidationError, match="无效的剪辑模式"): + use_case.execute(command) + + +# ── UpdateTemplateUseCase ── + + +class TestUpdateTemplateUseCase: + @pytest.fixture + def repo(self): + return _make_repo() + + @pytest.fixture + def use_case(self, repo): + return UpdateTemplateUseCase(repo) + + def test_update_name(self, use_case, repo): + """更新模板名称.""" + existing = _make_template() + repo.get.return_value = existing + repo.update.side_effect = lambda t: t + + command = UpdateTemplateCommand( + template_id="tmpl-001", + user_id="user-001", + name="新名称", + ) + result = use_case.execute(command) + + assert result.name == "新名称" + repo.update.assert_called_once() + + def test_update_not_found_raises(self, use_case, repo): + """模板不存在时抛出 NotFoundError.""" + repo.get.return_value = None + + command = UpdateTemplateCommand( + template_id="nonexistent", + user_id="user-001", + name="新名称", + ) + with pytest.raises(NotFoundError): + use_case.execute(command) + + def test_update_invalid_mode_raises(self, use_case, repo): + """更新为无效模式时抛出 ValidationError.""" + existing = _make_template() + repo.get.return_value = existing + + command = UpdateTemplateCommand( + template_id="tmpl-001", + user_id="user-001", + mode="bad_mode", + ) + with pytest.raises(ValidationError, match="无效的剪辑模式"): + use_case.execute(command) + + def test_replace_segments(self, use_case, repo): + """替换片段列表.""" + existing = _make_template() + repo.get.return_value = existing + repo.update.side_effect = lambda t: t + repo.create_segments.side_effect = lambda segs: segs + + command = UpdateTemplateCommand( + template_id="tmpl-001", + user_id="user-001", + segments=[ + SegmentCommand(segment_order=1, duration_min=5, duration_max=20, material_type=None), + ], + ) + result = use_case.execute(command) + + repo.delete_segments_by_template.assert_called_once_with("tmpl-001") + repo.create_segments.assert_called_once() + assert len(result.segments) == 1 + + +# ── GenerateFromTemplateUseCase — 业务规则校验 ── + + +class TestGenerateFromTemplateUseCase: + @pytest.fixture + def repo(self): + return _make_repo() + + @pytest.fixture + def use_case(self, repo): + return GenerateFromTemplateUseCase(repo) + + def test_one_take_with_one_segment_ok(self, use_case, repo): + """一镜到底 + 恰好 1 个片段 → 通过.""" + seg = TemplateSegment( + id="seg-001", template_id="tmpl-001", segment_order=1, + duration_min=0, duration_max=60, + ) + template = _make_template(mode="one_take", segments=[seg]) + repo.get.return_value = template + + command = GenerateFromTemplateCommand( + template_id="tmpl-001", user_id="user-001", + ) + result = use_case.execute(command) + + assert result.template.mode == "one_take" + assert result.warnings == [] + + def test_one_take_with_two_segments_raises(self, use_case, repo): + """一镜到底 + 2 个片段 → ValidationError.""" + segs = [ + TemplateSegment(id=f"seg-{i}", template_id="tmpl-001", segment_order=i, + duration_min=0, duration_max=30) + for i in (1, 2) + ] + template = _make_template(mode="one_take", segments=segs) + repo.get.return_value = template + + command = GenerateFromTemplateCommand( + template_id="tmpl-001", user_id="user-001", + ) + with pytest.raises(ValidationError, match="一镜到底模式必须恰好有 1 个片段"): + use_case.execute(command) + + def test_voice_over_mix_all_segments_have_material_type_ok(self, use_case, repo): + """口播+混剪 + 所有片段都有 material_type → 通过.""" + segs = [ + TemplateSegment(id="seg-1", template_id="tmpl-001", segment_order=1, + duration_min=5, duration_max=15, material_type="人物"), + TemplateSegment(id="seg-2", template_id="tmpl-001", segment_order=2, + duration_min=10, duration_max=30, material_type="场景"), + ] + template = _make_template(mode="voice_over_mix", segments=segs) + repo.get.return_value = template + + command = GenerateFromTemplateCommand( + template_id="tmpl-001", user_id="user-001", + ) + result = use_case.execute(command) + assert result.warnings == [] + + def test_voice_over_mix_missing_material_type_raises(self, use_case, repo): + """口播+混剪 + 某片段缺少 material_type → ValidationError.""" + segs = [ + TemplateSegment(id="seg-1", template_id="tmpl-001", segment_order=1, + duration_min=5, duration_max=15, material_type="人物"), + TemplateSegment(id="seg-2", template_id="tmpl-001", segment_order=2, + duration_min=10, duration_max=30, material_type=None), # 缺失 + ] + template = _make_template(mode="voice_over_mix", segments=segs) + repo.get.return_value = template + + command = GenerateFromTemplateCommand( + template_id="tmpl-001", user_id="user-001", + ) + with pytest.raises(ValidationError, match="material_type"): + use_case.execute(command) + + def test_voiceover_duration_within_tolerance_no_warning(self, use_case, repo): + """配音时长在 ±30% 以内 → 无警告.""" + template = _make_template(estimated_duration=60.0) + repo.get.return_value = template + + command = GenerateFromTemplateCommand( + template_id="tmpl-001", user_id="user-001", + voiceover_duration=70.0, # 70/60 = 1.167, within ±30% + ) + result = use_case.execute(command) + assert result.warnings == [] + + def test_voiceover_duration_exceeds_tolerance_warning(self, use_case, repo): + """配音时长超过 ±30% → 警告.""" + template = _make_template(estimated_duration=60.0) + repo.get.return_value = template + + command = GenerateFromTemplateCommand( + template_id="tmpl-001", user_id="user-001", + voiceover_duration=100.0, # 100/60 = 1.667, exceeds +30% + ) + result = use_case.execute(command) + + assert len(result.warnings) == 1 + assert result.warnings[0].code == "voiceover_duration_mismatch" + + def test_voiceover_duration_too_short_warning(self, use_case, repo): + """配音时长过短(< 70%)→ 警告.""" + template = _make_template(estimated_duration=60.0) + repo.get.return_value = template + + command = GenerateFromTemplateCommand( + template_id="tmpl-001", user_id="user-001", + voiceover_duration=30.0, # 30/60 = 0.5, below -30% + ) + result = use_case.execute(command) + + assert len(result.warnings) == 1 + assert result.warnings[0].code == "voiceover_duration_mismatch" + + def test_template_not_found_raises(self, use_case, repo): + """模板不存在 → NotFoundError.""" + repo.get.return_value = None + + command = GenerateFromTemplateCommand( + template_id="nonexistent", user_id="user-001", + ) + with pytest.raises(NotFoundError): + use_case.execute(command) + + +# ── Category Use Cases ── + + +class TestCategoryUseCases: + @pytest.fixture + def repo(self): + return _make_repo() + + def test_create_category(self, repo): + repo.create_category.side_effect = lambda c: c + + use_case = CreateCategoryUseCase(repo) + command = CreateCategoryCommand(user_id="user-001", name="Vlog") + result = use_case.execute(command) + + assert result.name == "Vlog" + repo.create_category.assert_called_once() + + def test_list_categories(self, repo): + categories = [ + TemplateCategory(id="cat-1", user_id="user-001", name="Vlog"), + TemplateCategory(id="cat-2", user_id="user-001", name="教程"), + ] + repo.list_categories.return_value = categories + + use_case = ListCategoriesUseCase(repo) + result = use_case.execute("user-001") + + assert len(result) == 2 + assert result[0].name == "Vlog" + + def test_delete_category_not_found(self, repo): + repo.delete_category.return_value = False + + use_case = DeleteTemplateUseCase(repo) + result = use_case.execute("nonexistent", "user-001") + assert result is False + + +# ── ListTemplatesUseCase ── + + +class TestListTemplatesUseCase: + def test_list_returns_templates(self): + repo = _make_repo() + templates = [_make_template(id=f"t-{i}") for i in range(3)] + repo.list_by_user.return_value = templates + + use_case = ListTemplatesUseCase(repo) + result = use_case.execute("user-001", skip=0, limit=50) + + assert len(result) == 3 + repo.list_by_user.assert_called_once_with("user-001", skip=0, limit=50) + + +# ── GetTemplateUseCase ── + + +class TestGetTemplateUseCase: + def test_get_existing(self): + repo = _make_repo() + template = _make_template() + repo.get.return_value = template + + use_case = GetTemplateUseCase(repo) + result = use_case.execute("tmpl-001", "user-001") + + assert result.id == "tmpl-001" + + def test_get_nonexistent_returns_none(self): + repo = _make_repo() + repo.get.return_value = None + + use_case = GetTemplateUseCase(repo) + result = use_case.execute("nonexistent", "user-001") + + assert result is None -- 2.54.0 From 9b6a0b8b14ab31b3c58a4dadb844100acf14592d Mon Sep 17 00:00:00 2001 From: Audit Bot Date: Mon, 29 Jun 2026 14:44:24 +0800 Subject: [PATCH 2/3] =?UTF-8?q?fix:=20=E4=BF=AE=E5=A4=8D=E4=BB=A3=E7=A0=81?= =?UTF-8?q?=E5=AE=A1=E8=AE=A15=E4=B8=AA=E9=97=AE=E9=A2=98=20(P0-1,P0-2,P1-?= =?UTF-8?q?3,P1-4,P1-5)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit P0-1: 剪辑模式使用 EditingMode 枚举,移除硬编码字符串 - voice_over_mix → voice_over, person_narration → voice_pip - use_cases.py 导入 EditingMode 枚举替代硬编码 VALID_MODES P0-2: 补充 packages/application/template/__init__.py P1-3: Use Cases 依赖 TemplateRepositoryPort 而非 SQLAlchemyTemplateRepository P1-4: generate 端点重命名为 validate(只校验不生成) - GenerateFromTemplateUseCase → ValidateTemplateUseCase - POST /{id}/generate → POST /{id}/validate P1-5: 软删除模板时级联清理关联 segments,避免孤儿数据 测试全部通过 (21/21) --- apps/api/app/api/routes/templates.py | 26 +++++----- apps/api/app/schemas/template.py | 6 +-- .../sqlalchemy_impl/template_repository.py | 4 ++ packages/application/template/__init__.py | 0 packages/application/template/commands.py | 2 +- packages/application/template/use_cases.py | 47 ++++++++++--------- packages/domain/template.py | 2 +- tests/unit/test_template_use_cases.py | 40 ++++++++-------- 8 files changed, 66 insertions(+), 61 deletions(-) create mode 100644 packages/application/template/__init__.py diff --git a/apps/api/app/api/routes/templates.py b/apps/api/app/api/routes/templates.py index af4884876..2d6b6afae 100644 --- a/apps/api/app/api/routes/templates.py +++ b/apps/api/app/api/routes/templates.py @@ -12,34 +12,34 @@ from app.schemas.template import ( CategoryResponse, CreateCategoryRequest, CreateTemplateRequest, - GenerateFromTemplateRequest, - GenerateFromTemplateResponse, - GenerateWarningResponse, ListCategoriesResponse, ListTemplatesResponse, SegmentResponse, TemplateResponse, UpdateTemplateRequest, + ValidateTemplateRequest, + ValidateTemplateResponse, + GenerateWarningResponse, ) from packages.adapters.sqlalchemy_impl.template_repository import SQLAlchemyTemplateRepository from packages.application.template.commands import ( CreateCategoryCommand, CreateTemplateCommand, - GenerateFromTemplateCommand, SegmentCommand, UpdateTemplateCommand, + ValidateTemplateCommand, ) from packages.application.template.use_cases import ( CreateCategoryUseCase, CreateTemplateUseCase, DeleteCategoryUseCase, DeleteTemplateUseCase, - GenerateFromTemplateUseCase, GetTemplateUseCase, ListCategoriesUseCase, ListTemplatesUseCase, NotFoundError, UpdateTemplateUseCase, + ValidateTemplateUseCase, ValidationError, ) @@ -208,23 +208,23 @@ def delete_template( return Response(status_code=204) -# ── Generate from template ── +# ── Validate template ── -@router.post("/{template_id}/generate", response_model=GenerateFromTemplateResponse) -def generate_from_template( +@router.post("/{template_id}/validate", response_model=ValidateTemplateResponse) +def validate_template( template_id: str, - request: GenerateFromTemplateRequest = GenerateFromTemplateRequest(), + request: ValidateTemplateRequest = ValidateTemplateRequest(), authenticated_user: AuthenticatedUser = Depends(get_current_user), template_repository: SQLAlchemyTemplateRepository = Depends(_get_template_repository), -) -> GenerateFromTemplateResponse: +) -> ValidateTemplateResponse: user_id = authenticated_user.user.id - command = GenerateFromTemplateCommand( + command = ValidateTemplateCommand( template_id=template_id, user_id=user_id, voiceover_duration=request.voiceover_duration, ) - use_case = GenerateFromTemplateUseCase(template_repository) + use_case = ValidateTemplateUseCase(template_repository) try: result = use_case.execute(command) except NotFoundError: @@ -232,7 +232,7 @@ def generate_from_template( except ValidationError as exc: raise HTTPException(status_code=status.HTTP_422_UNPROCESSABLE_ENTITY, detail=str(exc)) - return GenerateFromTemplateResponse( + return ValidateTemplateResponse( template=_to_response(result.template), warnings=[ GenerateWarningResponse(code=w.code, message=w.message, details=w.details) diff --git a/apps/api/app/schemas/template.py b/apps/api/app/schemas/template.py index 5f3bc1635..0d7bea514 100644 --- a/apps/api/app/schemas/template.py +++ b/apps/api/app/schemas/template.py @@ -77,9 +77,9 @@ class UpdateTemplateRequest(BaseModel): segments: Optional[List[SegmentRequest]] = None -# ── Generate ── +# ── Validate ── -class GenerateFromTemplateRequest(BaseModel): +class ValidateTemplateRequest(BaseModel): voiceover_duration: Optional[float] = None # 配音实际时长(秒) @@ -89,7 +89,7 @@ class GenerateWarningResponse(BaseModel): details: Dict[str, Any] = Field(default_factory=dict) -class GenerateFromTemplateResponse(BaseModel): +class ValidateTemplateResponse(BaseModel): template: TemplateResponse warnings: List[GenerateWarningResponse] = Field(default_factory=list) diff --git a/packages/adapters/sqlalchemy_impl/template_repository.py b/packages/adapters/sqlalchemy_impl/template_repository.py index f255a2af7..38cabae5d 100644 --- a/packages/adapters/sqlalchemy_impl/template_repository.py +++ b/packages/adapters/sqlalchemy_impl/template_repository.py @@ -118,6 +118,10 @@ class SQLAlchemyTemplateRepository: if model is None: return False model.is_active = False + # 级联清理关联的 segments,避免孤儿数据 + self.session.query(TemplateSegmentModel).filter( + TemplateSegmentModel.template_id == template_id, + ).delete(synchronize_session=False) self.session.commit() return True diff --git a/packages/application/template/__init__.py b/packages/application/template/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/packages/application/template/commands.py b/packages/application/template/commands.py index 693df8282..0fdc71e8e 100644 --- a/packages/application/template/commands.py +++ b/packages/application/template/commands.py @@ -49,7 +49,7 @@ class CreateCategoryCommand: @dataclass -class GenerateFromTemplateCommand: +class ValidateTemplateCommand: template_id: str user_id: str voiceover_duration: Optional[float] = None # 配音实际时长(用于偏差校验) diff --git a/packages/application/template/use_cases.py b/packages/application/template/use_cases.py index ab790a1ec..a6d112c32 100644 --- a/packages/application/template/use_cases.py +++ b/packages/application/template/use_cases.py @@ -5,14 +5,15 @@ import uuid from dataclasses import dataclass, field from typing import List, Optional -from packages.adapters.sqlalchemy_impl.template_repository import SQLAlchemyTemplateRepository from packages.application.template.commands import ( CreateCategoryCommand, CreateTemplateCommand, - GenerateFromTemplateCommand, UpdateTemplateCommand, + ValidateTemplateCommand, ) +from packages.domain.editing_mode import EditingMode from packages.domain.template import Template, TemplateCategory, TemplateSegment +from packages.ports.template_repository import TemplateRepositoryPort class NotFoundError(Exception): @@ -24,7 +25,7 @@ class ValidationError(Exception): pass -VALID_MODES = {"pip", "person_narration", "one_take", "voice_over_mix"} +VALID_MODES = {m.value for m in EditingMode} VALID_MATERIAL_TYPES = {"人物", "场景"} @@ -37,8 +38,8 @@ class GenerateWarning: @dataclass -class GenerateResult: - """使用模板生成的结果.""" +class ValidateResult: + """模板校验结果.""" template: Template warnings: List[GenerateWarning] = field(default_factory=list) @@ -47,7 +48,7 @@ class GenerateResult: class CreateTemplateUseCase: - def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + def __init__(self, repository: TemplateRepositoryPort) -> None: self.repository = repository def execute(self, command: CreateTemplateCommand) -> Template: @@ -88,7 +89,7 @@ class CreateTemplateUseCase: class ListTemplatesUseCase: - def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + def __init__(self, repository: TemplateRepositoryPort) -> None: self.repository = repository def execute( @@ -102,7 +103,7 @@ class ListTemplatesUseCase: class GetTemplateUseCase: - def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + def __init__(self, repository: TemplateRepositoryPort) -> None: self.repository = repository def execute(self, template_id: str, user_id: str) -> Optional[Template]: @@ -110,7 +111,7 @@ class GetTemplateUseCase: class UpdateTemplateUseCase: - def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + def __init__(self, repository: TemplateRepositoryPort) -> None: self.repository = repository def execute(self, command: UpdateTemplateCommand) -> Template: @@ -163,23 +164,23 @@ class UpdateTemplateUseCase: class DeleteTemplateUseCase: - def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + def __init__(self, repository: TemplateRepositoryPort) -> None: self.repository = repository def execute(self, template_id: str, user_id: str) -> bool: return self.repository.delete(template_id, user_id) -# ── Generate from template ── +# ── Validate template ── -class GenerateFromTemplateUseCase: - """使用模板生成剪辑计划 — 校验业务规则.""" +class ValidateTemplateUseCase: + """校验模板业务规则.""" - def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + def __init__(self, repository: TemplateRepositoryPort) -> None: self.repository = repository - def execute(self, command: GenerateFromTemplateCommand) -> GenerateResult: + def execute(self, command: ValidateTemplateCommand) -> ValidateResult: template = self.repository.get(command.template_id, command.user_id) if template is None: raise NotFoundError(f"Template {command.template_id} not found") @@ -187,18 +188,18 @@ class GenerateFromTemplateUseCase: warnings: List[GenerateWarning] = [] # 业务规则 1: one_take 必须恰好 1 个片段 - if template.mode == "one_take": + if template.mode == EditingMode.ONE_TAKE.value: if len(template.segments) != 1: raise ValidationError( f"一镜到底模式必须恰好有 1 个片段,当前有 {len(template.segments)} 个" ) - # 业务规则 2: voice_over_mix 每个片段必须有 material_type - if template.mode == "voice_over_mix": + # 业务规则 2: voice_over 每个片段必须有 material_type + if template.mode == EditingMode.VOICE_OVER.value: for seg in template.segments: if not seg.material_type or seg.material_type not in VALID_MATERIAL_TYPES: raise ValidationError( - f"口播+混剪模式下每个片段必须指定 material_type(人物/场景)," + f"口播+B-roll模式下每个片段必须指定 material_type(人物/场景)," f"片段 {seg.segment_order} 的 material_type 无效: {seg.material_type}" ) @@ -220,14 +221,14 @@ class GenerateFromTemplateUseCase: }, )) - return GenerateResult(template=template, warnings=warnings) + return ValidateResult(template=template, warnings=warnings) # ── Category CRUD ── class CreateCategoryUseCase: - def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + def __init__(self, repository: TemplateRepositoryPort) -> None: self.repository = repository def execute(self, command: CreateCategoryCommand) -> TemplateCategory: @@ -240,7 +241,7 @@ class CreateCategoryUseCase: class ListCategoriesUseCase: - def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + def __init__(self, repository: TemplateRepositoryPort) -> None: self.repository = repository def execute(self, user_id: str) -> List[TemplateCategory]: @@ -248,7 +249,7 @@ class ListCategoriesUseCase: class DeleteCategoryUseCase: - def __init__(self, repository: SQLAlchemyTemplateRepository) -> None: + def __init__(self, repository: TemplateRepositoryPort) -> None: self.repository = repository def execute(self, category_id: str, user_id: str) -> bool: diff --git a/packages/domain/template.py b/packages/domain/template.py index 8d0620743..c9e27bbc6 100644 --- a/packages/domain/template.py +++ b/packages/domain/template.py @@ -25,7 +25,7 @@ class Template: id: str user_id: str name: str - mode: str # pip / person_narration / one_take / voice_over_mix + mode: str # EditingMode 枚举值: pip / voice_pip / one_take / voice_over category: str = "" tags: List[str] = field(default_factory=list) title_config: dict = field(default_factory=dict) diff --git a/tests/unit/test_template_use_cases.py b/tests/unit/test_template_use_cases.py index bc1b4e6d0..4585ffbe6 100644 --- a/tests/unit/test_template_use_cases.py +++ b/tests/unit/test_template_use_cases.py @@ -8,20 +8,20 @@ import pytest from packages.application.template.commands import ( CreateCategoryCommand, CreateTemplateCommand, - GenerateFromTemplateCommand, SegmentCommand, UpdateTemplateCommand, + ValidateTemplateCommand, ) from packages.application.template.use_cases import ( CreateCategoryUseCase, CreateTemplateUseCase, DeleteTemplateUseCase, - GenerateFromTemplateUseCase, GetTemplateUseCase, ListCategoriesUseCase, ListTemplatesUseCase, NotFoundError, UpdateTemplateUseCase, + ValidateTemplateUseCase, ValidationError, ) from packages.domain.template import Template, TemplateCategory, TemplateSegment @@ -103,7 +103,7 @@ class TestCreateTemplateUseCase: command = CreateTemplateCommand( user_id="user-001", name="口播混剪模板", - mode="voice_over_mix", + mode="voice_over", segments=[ SegmentCommand(segment_order=1, duration_min=5, duration_max=15, material_type="人物"), SegmentCommand(segment_order=2, duration_min=10, duration_max=30, material_type="场景"), @@ -200,17 +200,17 @@ class TestUpdateTemplateUseCase: assert len(result.segments) == 1 -# ── GenerateFromTemplateUseCase — 业务规则校验 ── +# ── ValidateTemplateUseCase — 业务规则校验 ── -class TestGenerateFromTemplateUseCase: +class TestValidateTemplateUseCase: @pytest.fixture def repo(self): return _make_repo() @pytest.fixture def use_case(self, repo): - return GenerateFromTemplateUseCase(repo) + return ValidateTemplateUseCase(repo) def test_one_take_with_one_segment_ok(self, use_case, repo): """一镜到底 + 恰好 1 个片段 → 通过.""" @@ -221,7 +221,7 @@ class TestGenerateFromTemplateUseCase: template = _make_template(mode="one_take", segments=[seg]) repo.get.return_value = template - command = GenerateFromTemplateCommand( + command = ValidateTemplateCommand( template_id="tmpl-001", user_id="user-001", ) result = use_case.execute(command) @@ -239,41 +239,41 @@ class TestGenerateFromTemplateUseCase: template = _make_template(mode="one_take", segments=segs) repo.get.return_value = template - command = GenerateFromTemplateCommand( + command = ValidateTemplateCommand( template_id="tmpl-001", user_id="user-001", ) with pytest.raises(ValidationError, match="一镜到底模式必须恰好有 1 个片段"): use_case.execute(command) - def test_voice_over_mix_all_segments_have_material_type_ok(self, use_case, repo): - """口播+混剪 + 所有片段都有 material_type → 通过.""" + def test_voice_over_all_segments_have_material_type_ok(self, use_case, repo): + """口播+B-roll + 所有片段都有 material_type → 通过.""" segs = [ TemplateSegment(id="seg-1", template_id="tmpl-001", segment_order=1, duration_min=5, duration_max=15, material_type="人物"), TemplateSegment(id="seg-2", template_id="tmpl-001", segment_order=2, duration_min=10, duration_max=30, material_type="场景"), ] - template = _make_template(mode="voice_over_mix", segments=segs) + template = _make_template(mode="voice_over", segments=segs) repo.get.return_value = template - command = GenerateFromTemplateCommand( + command = ValidateTemplateCommand( template_id="tmpl-001", user_id="user-001", ) result = use_case.execute(command) assert result.warnings == [] - def test_voice_over_mix_missing_material_type_raises(self, use_case, repo): - """口播+混剪 + 某片段缺少 material_type → ValidationError.""" + def test_voice_over_missing_material_type_raises(self, use_case, repo): + """口播+B-roll + 某片段缺少 material_type → ValidationError.""" segs = [ TemplateSegment(id="seg-1", template_id="tmpl-001", segment_order=1, duration_min=5, duration_max=15, material_type="人物"), TemplateSegment(id="seg-2", template_id="tmpl-001", segment_order=2, duration_min=10, duration_max=30, material_type=None), # 缺失 ] - template = _make_template(mode="voice_over_mix", segments=segs) + template = _make_template(mode="voice_over", segments=segs) repo.get.return_value = template - command = GenerateFromTemplateCommand( + command = ValidateTemplateCommand( template_id="tmpl-001", user_id="user-001", ) with pytest.raises(ValidationError, match="material_type"): @@ -284,7 +284,7 @@ class TestGenerateFromTemplateUseCase: template = _make_template(estimated_duration=60.0) repo.get.return_value = template - command = GenerateFromTemplateCommand( + command = ValidateTemplateCommand( template_id="tmpl-001", user_id="user-001", voiceover_duration=70.0, # 70/60 = 1.167, within ±30% ) @@ -296,7 +296,7 @@ class TestGenerateFromTemplateUseCase: template = _make_template(estimated_duration=60.0) repo.get.return_value = template - command = GenerateFromTemplateCommand( + command = ValidateTemplateCommand( template_id="tmpl-001", user_id="user-001", voiceover_duration=100.0, # 100/60 = 1.667, exceeds +30% ) @@ -310,7 +310,7 @@ class TestGenerateFromTemplateUseCase: template = _make_template(estimated_duration=60.0) repo.get.return_value = template - command = GenerateFromTemplateCommand( + command = ValidateTemplateCommand( template_id="tmpl-001", user_id="user-001", voiceover_duration=30.0, # 30/60 = 0.5, below -30% ) @@ -323,7 +323,7 @@ class TestGenerateFromTemplateUseCase: """模板不存在 → NotFoundError.""" repo.get.return_value = None - command = GenerateFromTemplateCommand( + command = ValidateTemplateCommand( template_id="nonexistent", user_id="user-001", ) with pytest.raises(NotFoundError): -- 2.54.0 From 78524a04eb69ab0d82b15e2da350cb2e671bc8f8 Mon Sep 17 00:00:00 2001 From: Audit Bot Date: Mon, 29 Jun 2026 14:53:21 +0800 Subject: [PATCH 3/3] =?UTF-8?q?fix:=20=E4=BF=AE=E5=A4=8D=E5=A4=8D=E5=AE=A1?= =?UTF-8?q?2=E4=B8=AA=E5=BB=BA=E8=AE=AE=20(models=E6=B3=A8=E9=87=8A=20+=20?= =?UTF-8?q?N+1=E6=9F=A5=E8=AF=A2=20+=20=E4=BA=8B=E5=8A=A1=E8=BE=B9?= =?UTF-8?q?=E7=95=8C)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - models.py: 注释中旧模式名对齐 EditingMode 枚举 - list_by_user: 批量加载 segments 避免 N+1 查询 - create(): flush 替代 commit,create_segments 统一提交事务 --- packages/adapters/sqlalchemy_impl/models.py | 4 +-- .../sqlalchemy_impl/template_repository.py | 22 +++++++++++++-- packages/application/template/use_cases.py | 28 +++++++++---------- 3 files changed, 35 insertions(+), 19 deletions(-) diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index f5854e025..7bfc5fab2 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -287,7 +287,7 @@ class TemplateModel(Base): id = Column(String(36), primary_key=True) user_id = Column(String(36), nullable=False, index=True) name = Column(String(200), nullable=False) - mode = Column(String(30), nullable=False, index=True) # pip / person_narration / one_take / voice_over_mix + mode = Column(String(30), nullable=False, index=True) # EditingMode 枚举值: pip / voice_pip / one_take / voice_over category = Column(String(100), nullable=False, default="") tags = Column(JSON, nullable=False, default=list) title_config = Column(JSON, nullable=False, default=dict) @@ -307,7 +307,7 @@ class TemplateSegmentModel(Base): segment_order = Column(Integer, nullable=False) duration_min = Column(Float, nullable=False) duration_max = Column(Float, nullable=False) - material_type = Column(String(20), nullable=True) # 仅 voice_over_mix: 人物/场景 + material_type = Column(String(20), nullable=True) # 仅 voice_over 模式: 人物/场景 created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) diff --git a/packages/adapters/sqlalchemy_impl/template_repository.py b/packages/adapters/sqlalchemy_impl/template_repository.py index 38cabae5d..a006692ff 100644 --- a/packages/adapters/sqlalchemy_impl/template_repository.py +++ b/packages/adapters/sqlalchemy_impl/template_repository.py @@ -40,8 +40,23 @@ class SQLAlchemyTemplateRepository: .all() ) templates = [self._model_to_entity(m) for m in models] - for t in templates: - t.segments = self.list_segments(t.id) + # 批量加载所有 segments,避免 N+1 查询 + if templates: + template_ids = [t.id for t in templates] + seg_models = ( + self.session.query(TemplateSegmentModel) + .filter(TemplateSegmentModel.template_id.in_(template_ids)) + .order_by(TemplateSegmentModel.segment_order) + .all() + ) + # 按 template_id 分组 + seg_map: dict[str, list] = {} + for sm in seg_models: + seg_map.setdefault(sm.template_id, []).append( + self._segment_model_to_entity(sm), + ) + for t in templates: + t.segments = seg_map.get(t.id, []) return templates def get(self, template_id: str, user_id: str) -> Optional[Template]: @@ -74,7 +89,8 @@ class SQLAlchemyTemplateRepository: is_active=template.is_active, ) self.session.add(model) - self.session.commit() + # flush 而非 commit,让 create + create_segments 在同一事务中提交 + self.session.flush() self.session.refresh(model) result = self._model_to_entity(model) result.segments = template.segments diff --git a/packages/application/template/use_cases.py b/packages/application/template/use_cases.py index a6d112c32..f43cb0122 100644 --- a/packages/application/template/use_cases.py +++ b/packages/application/template/use_cases.py @@ -70,20 +70,20 @@ class CreateTemplateUseCase: ) template = self.repository.create(template) - if command.segments: - segments = [ - TemplateSegment( - id=uuid.uuid4().hex, - template_id=template.id, - segment_order=seg.segment_order, - duration_min=seg.duration_min, - duration_max=seg.duration_max, - material_type=seg.material_type, - ) - for seg in command.segments - ] - self.repository.create_segments(segments) - template.segments = segments + # 始终调用 create_segments 以确保在同一事务中提交 + segments = [ + TemplateSegment( + id=uuid.uuid4().hex, + template_id=template.id, + segment_order=seg.segment_order, + duration_min=seg.duration_min, + duration_max=seg.duration_max, + material_type=seg.material_type, + ) + for seg in command.segments + ] + self.repository.create_segments(segments) + template.segments = segments return template -- 2.54.0