diff --git a/apps/api/app/api/routes/templates.py b/apps/api/app/api/routes/templates.py index 3f632d4d5..05bbd8355 100644 --- a/apps/api/app/api/routes/templates.py +++ b/apps/api/app/api/routes/templates.py @@ -106,6 +106,10 @@ def list_templates( 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: @@ -116,6 +120,7 @@ def list_templates( 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) diff --git a/packages/adapters/sqlalchemy_impl/template_repository.py b/packages/adapters/sqlalchemy_impl/template_repository.py index 275d6bc3d..a4a0a3f92 100755 --- a/packages/adapters/sqlalchemy_impl/template_repository.py +++ b/packages/adapters/sqlalchemy_impl/template_repository.py @@ -10,6 +10,7 @@ from __future__ import annotations import uuid from typing import List, Optional +from sqlalchemy import or_ from sqlalchemy.orm import Session from packages.adapters.sqlalchemy_impl.models import ( @@ -28,6 +29,20 @@ class SQLAlchemyTemplateRepository: def __init__(self, session: Session) -> None: self.session = session + def _filter_with_segment_configs(self, query): + """只保留在 template_clip_configs 或 template_segments 中存在片段配置的模板。 + + 两张表都没有记录的模板无法用于生成(from-assets 会 400), + 剪辑页选模板时应排除;模板编辑器不传 valid_only,仍可见全部模板。 + """ + has_clip_config = self.session.query(TemplateClipConfigModel.id).filter( + TemplateClipConfigModel.template_id == TemplateModel.id, + ) + has_segment = self.session.query(TemplateSegmentModel.id).filter( + TemplateSegmentModel.template_id == TemplateModel.id, + ) + return query.filter(or_(has_clip_config.exists(), has_segment.exists())) + # ── Template CRUD ── def list_by_user( @@ -40,11 +55,14 @@ class SQLAlchemyTemplateRepository: tag: Optional[str] = None, keyword: Optional[str] = None, mode: Optional[str] = None, + valid_only: bool = False, ) -> List[Template]: query = self.session.query(TemplateModel).filter( TemplateModel.user_id == user_id, TemplateModel.is_active.is_(True), ) + if valid_only: + query = self._filter_with_segment_configs(query) if category: query = query.filter(TemplateModel.category == category) if mode: @@ -173,11 +191,14 @@ class SQLAlchemyTemplateRepository: tag: Optional[str] = None, keyword: Optional[str] = None, mode: Optional[str] = None, + valid_only: bool = False, ) -> int: query = self.session.query(TemplateModel).filter( TemplateModel.user_id == user_id, TemplateModel.is_active.is_(True), ) + if valid_only: + query = self._filter_with_segment_configs(query) if category: query = query.filter(TemplateModel.category == category) if mode: diff --git a/packages/application/template/commands.py b/packages/application/template/commands.py index a7a07bb0f..90dbea846 100755 --- a/packages/application/template/commands.py +++ b/packages/application/template/commands.py @@ -62,6 +62,7 @@ class ListTemplatesFilter: tag: Optional[str] = None keyword: Optional[str] = None mode: Optional[str] = None + valid_only: bool = False @dataclass diff --git a/packages/application/template/use_cases.py b/packages/application/template/use_cases.py index 4b3f2202e..e4e537c2f 100755 --- a/packages/application/template/use_cases.py +++ b/packages/application/template/use_cases.py @@ -106,6 +106,7 @@ class ListTemplatesUseCase: tag=filter.tag, keyword=filter.keyword, mode=filter.mode, + valid_only=filter.valid_only, ) @@ -127,6 +128,7 @@ class CountTemplatesUseCase: tag=filter.tag, keyword=filter.keyword, mode=filter.mode, + valid_only=filter.valid_only, ) diff --git a/packages/ports/template_repository.py b/packages/ports/template_repository.py index 7b071ab63..8cd7f057b 100755 --- a/packages/ports/template_repository.py +++ b/packages/ports/template_repository.py @@ -18,6 +18,7 @@ class TemplateRepositoryPort(Protocol): tag: Optional[str] = None, keyword: Optional[str] = None, mode: Optional[str] = None, + valid_only: bool = False, ) -> List[Template]: ... def get(self, template_id: str, user_id: str) -> Optional[Template]: ... def create(self, template: Template) -> Template: ... @@ -31,6 +32,7 @@ class TemplateRepositoryPort(Protocol): tag: Optional[str] = None, keyword: Optional[str] = None, mode: Optional[str] = None, + valid_only: bool = False, ) -> int: ... def copy_template(self, template_id: str, user_id: str, new_name: str) -> Template: ... def list_segments(self, template_id: str) -> List[TemplateSegment]: ... diff --git a/tests/unit/test_template_use_cases.py b/tests/unit/test_template_use_cases.py index 7dc2eec95..0fb5aff48 100755 --- a/tests/unit/test_template_use_cases.py +++ b/tests/unit/test_template_use_cases.py @@ -188,6 +188,7 @@ class TestListTemplatesUseCase: tag="tag1", keyword="test", mode="one_take", + valid_only=False, ) def test_list_pagination(self): @@ -226,6 +227,7 @@ class TestCountTemplatesUseCase: tag="tag1", keyword="kw", mode="pip", + valid_only=False, ) diff --git a/tests/unit/test_unify_template_segments.py b/tests/unit/test_unify_template_segments.py index 1f6e2d4af..b043a4395 100644 --- a/tests/unit/test_unify_template_segments.py +++ b/tests/unit/test_unify_template_segments.py @@ -158,6 +158,52 @@ class TestListByUser: assert len(result[0].segments) == 1 assert result[0].segments[0].duration_min == 2.0 + def test_valid_only_filters_templates_without_segments(self, repo, session): + """#1769: valid_only=True 时排除两张片段表都没有记录的无效模板.""" + # 有效模板:有 clip_configs + valid_clip = _make_template(name="有效模板-clip_configs") + repo.create(valid_clip) + repo.create_segments([_make_segment(valid_clip.id, order=1)]) + # 有效模板:仅有旧表 template_segments 记录 + valid_old = _make_template(name="有效模板-old_segments") + repo.create(valid_old) + old = TemplateSegmentModel( + id=str(uuid.uuid4()), + template_id=valid_old.id, + segment_order=1, + duration_min=2.0, + duration_max=6.0, + ) + session.add(old) + session.commit() + # 无效模板:两张表都没有记录 + invalid = _make_template(name="无效模板-无片段") + repo.create(invalid) + + # 默认不过滤:编辑器视角能看到全部 3 个模板 + all_templates = repo.list_by_user("u1") + assert len(all_templates) == 3 + assert repo.count_by_user("u1") == 3 + + # valid_only=True:剪辑页视角只返回 2 个有效模板 + valid_templates = repo.list_by_user("u1", valid_only=True) + assert {t.name for t in valid_templates} == {"有效模板-clip_configs", "有效模板-old_segments"} + assert all(len(t.segments) > 0 for t in valid_templates) + assert repo.count_by_user("u1", valid_only=True) == 2 + + def test_valid_only_with_filters_and_pagination(self, repo, session): + """valid_only 与其他过滤/分页条件组合使用.""" + tpl = _make_template(name="口播模板", mode="voice_over") + repo.create(tpl) + repo.create_segments([_make_segment(tpl.id, order=1, material_type="人物")]) + _invalid = _make_template(name="口播无效模板", mode="voice_over") + repo.create(_invalid) + + result = repo.list_by_user("u1", mode="voice_over", valid_only=True) + assert len(result) == 1 + assert result[0].name == "口播模板" + assert repo.count_by_user("u1", mode="voice_over", valid_only=True) == 1 + class TestCopyTemplate: def test_copy_writes_to_clip_configs(self, repo, session):