Compare commits
4 Commits
| Author | SHA1 | Date | |
|---|---|---|---|
| c424295588 | |||
| 7c41c56284 | |||
| 80491386bc | |||
| 4f6461d4a1 |
Regular → Executable
+94
-6
@@ -8,13 +8,16 @@ 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,
|
||||
@@ -27,19 +30,25 @@ 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,
|
||||
GetTemplateUsageUseCase,
|
||||
GetTemplateUseCase,
|
||||
ListCategoriesUseCase,
|
||||
ListTagsUseCase,
|
||||
ListTemplatesUseCase,
|
||||
NotFoundError,
|
||||
UpdateTemplateUseCase,
|
||||
@@ -67,7 +76,7 @@ def _segment_to_response(seg) -> SegmentResponse:
|
||||
)
|
||||
|
||||
|
||||
def _to_response(template) -> TemplateResponse:
|
||||
def _to_response(template, usage_count: int = 0) -> TemplateResponse:
|
||||
return TemplateResponse(
|
||||
id=template.id,
|
||||
user_id=template.user_id,
|
||||
@@ -81,6 +90,7 @@ def _to_response(template) -> TemplateResponse:
|
||||
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,
|
||||
)
|
||||
@@ -93,19 +103,36 @@ def _to_response(template) -> TemplateResponse:
|
||||
def list_templates(
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(50, ge=1, le=200),
|
||||
category: str | None = Query(None, description="按分类筛选"),
|
||||
tag: str | None = Query(None, description="按标签筛选"),
|
||||
keyword: str | None = Query(None, description="按名称关键词搜索"),
|
||||
mode: str | None = Query(None, description="按剪辑模式筛选"),
|
||||
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,
|
||||
)
|
||||
use_case = ListTemplatesUseCase(template_repository)
|
||||
templates = use_case.execute(user_id, skip=skip, limit=limit)
|
||||
total = template_repository.count_by_user(user_id)
|
||||
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)
|
||||
|
||||
# 批量查询使用次数
|
||||
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=[_to_response(t) for t in templates],
|
||||
items=items,
|
||||
total=total,
|
||||
)
|
||||
|
||||
@@ -120,12 +147,13 @@ def get_template(
|
||||
try:
|
||||
use_case = GetTemplateUseCase(template_repository)
|
||||
template = use_case.execute(template_id, user_id)
|
||||
usage = template_repository.get_usage_count(template_id)
|
||||
except Exception:
|
||||
logger.exception("get_template 查询失败: template_id=%s", template_id)
|
||||
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail="模板查询失败")
|
||||
if template is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Template not found")
|
||||
return _to_response(template)
|
||||
return _to_response(template, usage_count=usage)
|
||||
|
||||
|
||||
@router.post("", response_model=TemplateResponse, status_code=status.HTTP_201_CREATED)
|
||||
@@ -220,6 +248,47 @@ def delete_template(
|
||||
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:
|
||||
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.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,
|
||||
@@ -318,4 +387,23 @@ def delete_category(
|
||||
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
|
||||
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)
|
||||
|
||||
Regular → Executable
+23
@@ -45,6 +45,7 @@ class TemplateResponse(BaseModel):
|
||||
segments: List[SegmentResponse] = Field(default_factory=list)
|
||||
is_active: bool = True
|
||||
is_favorite: bool = False
|
||||
usage_count: int = 0
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
@@ -120,3 +121,25 @@ class CreateCategoryRequest(BaseModel):
|
||||
|
||||
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
|
||||
|
||||
Regular → Executable
+118
-18
@@ -2,11 +2,14 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from typing import List, Optional
|
||||
|
||||
from sqlalchemy import func
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import (
|
||||
EditPlanModel,
|
||||
TemplateCategoryModel,
|
||||
TemplateModel,
|
||||
TemplateSegmentModel,
|
||||
@@ -28,18 +31,26 @@ class SQLAlchemyTemplateRepository:
|
||||
*,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
category: Optional[str] = None,
|
||||
tag: Optional[str] = None,
|
||||
keyword: Optional[str] = None,
|
||||
mode: Optional[str] = None,
|
||||
) -> List[Template]:
|
||||
models = (
|
||||
self.session.query(TemplateModel)
|
||||
.filter(
|
||||
TemplateModel.user_id == user_id,
|
||||
TemplateModel.is_active.is_(True),
|
||||
)
|
||||
.order_by(TemplateModel.created_at.desc())
|
||||
.offset(skip)
|
||||
.limit(limit)
|
||||
.all()
|
||||
query = self.session.query(TemplateModel).filter(
|
||||
TemplateModel.user_id == user_id,
|
||||
TemplateModel.is_active.is_(True),
|
||||
)
|
||||
if category:
|
||||
query = query.filter(TemplateModel.category == category)
|
||||
if mode:
|
||||
query = query.filter(TemplateModel.mode == mode)
|
||||
if keyword:
|
||||
like_pattern = f"%{keyword}%"
|
||||
query = query.filter(TemplateModel.name.like(like_pattern))
|
||||
if tag:
|
||||
# JSON 数组包含指定标签(MySQL JSON_CONTAINS / SQLite json_each 兼容写法用 LIKE)
|
||||
query = query.filter(TemplateModel.tags.like(f'%"{tag}"%'))
|
||||
models = query.order_by(TemplateModel.created_at.desc()).offset(skip).limit(limit).all()
|
||||
templates = [self._model_to_entity(m) for m in models]
|
||||
# 批量加载所有 segments,避免 N+1 查询
|
||||
if templates:
|
||||
@@ -142,15 +153,77 @@ class SQLAlchemyTemplateRepository:
|
||||
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.is_(True),
|
||||
)
|
||||
.count()
|
||||
def count_by_user(
|
||||
self,
|
||||
user_id: str,
|
||||
*,
|
||||
category: Optional[str] = None,
|
||||
tag: Optional[str] = None,
|
||||
keyword: Optional[str] = None,
|
||||
mode: Optional[str] = None,
|
||||
) -> int:
|
||||
query = self.session.query(TemplateModel).filter(
|
||||
TemplateModel.user_id == user_id,
|
||||
TemplateModel.is_active.is_(True),
|
||||
)
|
||||
if category:
|
||||
query = query.filter(TemplateModel.category == category)
|
||||
if mode:
|
||||
query = query.filter(TemplateModel.mode == mode)
|
||||
if keyword:
|
||||
query = query.filter(TemplateModel.name.like(f"%{keyword}%"))
|
||||
if tag:
|
||||
query = query.filter(TemplateModel.tags.like(f'%"{tag}"%'))
|
||||
return query.count()
|
||||
|
||||
def copy_template(self, template_id: str, user_id: str, new_name: str) -> Template:
|
||||
"""复制模板(含所有 segments)。"""
|
||||
source = self.get(template_id, user_id)
|
||||
if source is None:
|
||||
raise ValueError(f"Template {template_id} not found")
|
||||
|
||||
new_id = str(uuid.uuid4())
|
||||
new_template = Template(
|
||||
id=new_id,
|
||||
user_id=user_id,
|
||||
name=new_name,
|
||||
mode=source.mode,
|
||||
category=source.category,
|
||||
tags=list(source.tags),
|
||||
title_config=dict(source.title_config),
|
||||
subtitle_config=dict(source.subtitle_config),
|
||||
bgm_config=dict(source.bgm_config),
|
||||
estimated_duration=source.estimated_duration,
|
||||
is_active=True,
|
||||
)
|
||||
created = self.create(new_template)
|
||||
|
||||
# 复制 segments
|
||||
new_segments: List[TemplateSegment] = []
|
||||
for seg in source.segments:
|
||||
new_seg = TemplateSegment(
|
||||
id=str(uuid.uuid4()),
|
||||
template_id=new_id,
|
||||
segment_order=seg.segment_order,
|
||||
duration_min=seg.duration_min,
|
||||
duration_max=seg.duration_max,
|
||||
material_type=seg.material_type,
|
||||
)
|
||||
new_segments.append(new_seg)
|
||||
model = TemplateSegmentModel(
|
||||
id=new_seg.id,
|
||||
template_id=new_seg.template_id,
|
||||
segment_order=new_seg.segment_order,
|
||||
duration_min=new_seg.duration_min,
|
||||
duration_max=new_seg.duration_max,
|
||||
material_type=new_seg.material_type,
|
||||
)
|
||||
self.session.add(model)
|
||||
if new_segments:
|
||||
self.session.commit()
|
||||
|
||||
created.segments = new_segments
|
||||
return created
|
||||
|
||||
# ── Segments ──
|
||||
|
||||
@@ -234,6 +307,33 @@ class SQLAlchemyTemplateRepository:
|
||||
self.session.commit()
|
||||
return True
|
||||
|
||||
# ── Tags ──
|
||||
|
||||
def list_tags(self, user_id: str) -> List[str]:
|
||||
"""获取用户所有模板的标签(去重)。"""
|
||||
models = (
|
||||
self.session.query(TemplateModel)
|
||||
.filter(
|
||||
TemplateModel.user_id == user_id,
|
||||
TemplateModel.is_active.is_(True),
|
||||
TemplateModel.tags.isnot(None),
|
||||
)
|
||||
.all()
|
||||
)
|
||||
tags_set: set[str] = set()
|
||||
for m in models:
|
||||
if m.tags:
|
||||
for t in m.tags:
|
||||
if t:
|
||||
tags_set.add(t)
|
||||
return sorted(tags_set)
|
||||
|
||||
# ── Usage Stats ──
|
||||
|
||||
def get_usage_count(self, template_id: str) -> int:
|
||||
"""获取模板被使用的次数(关联的剪辑计划数量)。"""
|
||||
return self.session.query(EditPlanModel).filter(EditPlanModel.template_id == template_id).count()
|
||||
|
||||
# ── Mapping helpers ──
|
||||
|
||||
@staticmethod
|
||||
|
||||
Regular → Executable
+15
@@ -49,6 +49,21 @@ class CreateCategoryCommand:
|
||||
name: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class CopyTemplateCommand:
|
||||
template_id: str
|
||||
user_id: str
|
||||
new_name: str
|
||||
|
||||
|
||||
@dataclass
|
||||
class ListTemplatesFilter:
|
||||
category: Optional[str] = None
|
||||
tag: Optional[str] = None
|
||||
keyword: Optional[str] = None
|
||||
mode: Optional[str] = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class ValidateTemplateCommand:
|
||||
template_id: str
|
||||
|
||||
Regular → Executable
+74
-1
@@ -7,8 +7,10 @@ from dataclasses import dataclass, field
|
||||
from typing import List, Optional
|
||||
|
||||
from packages.application.template.commands import (
|
||||
CopyTemplateCommand,
|
||||
CreateCategoryCommand,
|
||||
CreateTemplateCommand,
|
||||
ListTemplatesFilter,
|
||||
UpdateTemplateCommand,
|
||||
ValidateTemplateCommand,
|
||||
)
|
||||
@@ -102,8 +104,40 @@ class ListTemplatesUseCase:
|
||||
*,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
filter: Optional[ListTemplatesFilter] = None,
|
||||
) -> List[Template]:
|
||||
return self.repository.list_by_user(user_id, skip=skip, limit=limit)
|
||||
if filter is None:
|
||||
return self.repository.list_by_user(user_id, skip=skip, limit=limit)
|
||||
return self.repository.list_by_user(
|
||||
user_id,
|
||||
skip=skip,
|
||||
limit=limit,
|
||||
category=filter.category,
|
||||
tag=filter.tag,
|
||||
keyword=filter.keyword,
|
||||
mode=filter.mode,
|
||||
)
|
||||
|
||||
|
||||
class CountTemplatesUseCase:
|
||||
def __init__(self, repository: TemplateRepositoryPort) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def execute(
|
||||
self,
|
||||
user_id: str,
|
||||
*,
|
||||
filter: Optional[ListTemplatesFilter] = None,
|
||||
) -> int:
|
||||
if filter is None:
|
||||
return self.repository.count_by_user(user_id)
|
||||
return self.repository.count_by_user(
|
||||
user_id,
|
||||
category=filter.category,
|
||||
tag=filter.tag,
|
||||
keyword=filter.keyword,
|
||||
mode=filter.mode,
|
||||
)
|
||||
|
||||
|
||||
class GetTemplateUseCase:
|
||||
@@ -175,6 +209,23 @@ class DeleteTemplateUseCase:
|
||||
return self.repository.delete(template_id, user_id)
|
||||
|
||||
|
||||
class CopyTemplateUseCase:
|
||||
def __init__(self, repository: TemplateRepositoryPort) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def execute(self, command: CopyTemplateCommand) -> 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 not command.new_name or not command.new_name.strip():
|
||||
raise ValidationError("新模板名称不能为空")
|
||||
return self.repository.copy_template(
|
||||
command.template_id,
|
||||
command.user_id,
|
||||
command.new_name.strip(),
|
||||
)
|
||||
|
||||
|
||||
# ── Validate template ──
|
||||
|
||||
|
||||
@@ -258,3 +309,25 @@ class DeleteCategoryUseCase:
|
||||
|
||||
def execute(self, category_id: str, user_id: str) -> bool:
|
||||
return self.repository.delete_category(category_id, user_id)
|
||||
|
||||
|
||||
# ── Tags ──
|
||||
|
||||
|
||||
class ListTagsUseCase:
|
||||
def __init__(self, repository: TemplateRepositoryPort) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def execute(self, user_id: str) -> List[str]:
|
||||
return self.repository.list_tags(user_id)
|
||||
|
||||
|
||||
# ── Usage Stats ──
|
||||
|
||||
|
||||
class GetTemplateUsageUseCase:
|
||||
def __init__(self, repository: TemplateRepositoryPort) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def execute(self, template_id: str) -> int:
|
||||
return self.repository.get_usage_count(template_id)
|
||||
|
||||
Regular → Executable
+23
-2
@@ -8,12 +8,31 @@ 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 list_by_user(
|
||||
self,
|
||||
user_id: str,
|
||||
*,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
category: Optional[str] = None,
|
||||
tag: Optional[str] = None,
|
||||
keyword: Optional[str] = None,
|
||||
mode: Optional[str] = None,
|
||||
) -> 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 count_by_user(
|
||||
self,
|
||||
user_id: str,
|
||||
*,
|
||||
category: Optional[str] = None,
|
||||
tag: Optional[str] = None,
|
||||
keyword: Optional[str] = None,
|
||||
mode: Optional[str] = None,
|
||||
) -> int: ...
|
||||
def copy_template(self, template_id: str, user_id: str, new_name: str) -> Template: ...
|
||||
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: ...
|
||||
@@ -21,3 +40,5 @@ class TemplateRepositoryPort(Protocol):
|
||||
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: ...
|
||||
def list_tags(self, user_id: str) -> List[str]: ...
|
||||
def get_usage_count(self, template_id: str) -> int: ...
|
||||
|
||||
@@ -15,3 +15,4 @@ exclude =
|
||||
per-file-ignores =
|
||||
*/__init__.py:F401,F403,F405
|
||||
tests/*:E402,F401,F841
|
||||
packages/ports/*:E301,E704
|
||||
|
||||
Regular → Executable
+231
@@ -7,18 +7,24 @@ from unittest.mock import Mock
|
||||
import pytest
|
||||
|
||||
from packages.application.template.commands import (
|
||||
CopyTemplateCommand,
|
||||
CreateCategoryCommand,
|
||||
CreateTemplateCommand,
|
||||
ListTemplatesFilter,
|
||||
SegmentCommand,
|
||||
UpdateTemplateCommand,
|
||||
ValidateTemplateCommand,
|
||||
)
|
||||
from packages.application.template.use_cases import (
|
||||
CopyTemplateUseCase,
|
||||
CountTemplatesUseCase,
|
||||
CreateCategoryUseCase,
|
||||
CreateTemplateUseCase,
|
||||
DeleteTemplateUseCase,
|
||||
GetTemplateUsageUseCase,
|
||||
GetTemplateUseCase,
|
||||
ListCategoriesUseCase,
|
||||
ListTagsUseCase,
|
||||
ListTemplatesUseCase,
|
||||
NotFoundError,
|
||||
UpdateTemplateUseCase,
|
||||
@@ -44,6 +50,9 @@ def _make_repo():
|
||||
repo.create_category = Mock()
|
||||
repo.get_category = Mock(return_value=None)
|
||||
repo.delete_category = Mock(return_value=False)
|
||||
repo.copy_template = Mock()
|
||||
repo.list_tags = Mock(return_value=[])
|
||||
repo.get_usage_count = Mock(return_value=0)
|
||||
return repo
|
||||
|
||||
|
||||
@@ -442,3 +451,225 @@ class TestGetTemplateUseCase:
|
||||
result = use_case.execute("nonexistent", "user-001")
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
# ── CopyTemplateUseCase ──
|
||||
|
||||
|
||||
class TestCopyTemplateUseCase:
|
||||
@pytest.fixture
|
||||
def repo(self):
|
||||
repo = _make_repo()
|
||||
source = _make_template(
|
||||
id="tmpl-src",
|
||||
name="源模板",
|
||||
segments=[
|
||||
TemplateSegment(
|
||||
id="seg-1",
|
||||
template_id="tmpl-src",
|
||||
segment_order=0,
|
||||
duration_min=5.0,
|
||||
duration_max=10.0,
|
||||
material_type=None,
|
||||
),
|
||||
],
|
||||
)
|
||||
repo.get = Mock(return_value=source)
|
||||
|
||||
def _copy_side_effect(template_id, user_id, new_name):
|
||||
return _make_template(
|
||||
id="tmpl-copied",
|
||||
user_id=user_id,
|
||||
name=new_name,
|
||||
segments=[
|
||||
TemplateSegment(
|
||||
id="seg-copied",
|
||||
template_id="tmpl-copied",
|
||||
segment_order=0,
|
||||
duration_min=5.0,
|
||||
duration_max=10.0,
|
||||
material_type=None,
|
||||
)
|
||||
],
|
||||
)
|
||||
|
||||
repo.copy_template = Mock(side_effect=_copy_side_effect)
|
||||
return repo
|
||||
|
||||
@pytest.fixture
|
||||
def use_case(self, repo):
|
||||
return CopyTemplateUseCase(repo)
|
||||
|
||||
def test_copy_success(self, use_case, repo):
|
||||
command = CopyTemplateCommand(
|
||||
template_id="tmpl-src",
|
||||
user_id="user-001",
|
||||
new_name="复制的模板",
|
||||
)
|
||||
result = use_case.execute(command)
|
||||
|
||||
assert result.id == "tmpl-copied"
|
||||
assert result.name == "复制的模板"
|
||||
assert len(result.segments) == 1
|
||||
repo.copy_template.assert_called_once_with(
|
||||
"tmpl-src",
|
||||
"user-001",
|
||||
"复制的模板",
|
||||
)
|
||||
|
||||
def test_copy_not_found_raises(self, use_case, repo):
|
||||
repo.get = Mock(return_value=None)
|
||||
command = CopyTemplateCommand(
|
||||
template_id="tmpl-nonexist",
|
||||
user_id="user-001",
|
||||
new_name="新名字",
|
||||
)
|
||||
with pytest.raises(NotFoundError):
|
||||
use_case.execute(command)
|
||||
|
||||
def test_copy_empty_name_raises(self, use_case, repo):
|
||||
command = CopyTemplateCommand(
|
||||
template_id="tmpl-src",
|
||||
user_id="user-001",
|
||||
new_name=" ",
|
||||
)
|
||||
with pytest.raises(ValidationError):
|
||||
use_case.execute(command)
|
||||
|
||||
|
||||
# ── ListTemplatesUseCase (filter) ──
|
||||
|
||||
|
||||
class TestListTemplatesUseCaseWithFilter:
|
||||
def test_list_with_category_filter(self):
|
||||
repo = _make_repo()
|
||||
repo.list_by_user = Mock(return_value=[])
|
||||
use_case = ListTemplatesUseCase(repo)
|
||||
f = ListTemplatesFilter(category="vlog")
|
||||
|
||||
use_case.execute("user-001", skip=0, limit=10, filter=f)
|
||||
|
||||
repo.list_by_user.assert_called_once()
|
||||
call_kwargs = repo.list_by_user.call_args
|
||||
assert call_kwargs[1]["category"] == "vlog"
|
||||
|
||||
def test_list_with_tag_filter(self):
|
||||
repo = _make_repo()
|
||||
repo.list_by_user = Mock(return_value=[])
|
||||
use_case = ListTemplatesUseCase(repo)
|
||||
f = ListTemplatesFilter(tag="热门")
|
||||
|
||||
use_case.execute("user-001", filter=f)
|
||||
|
||||
call_kwargs = repo.list_by_user.call_args
|
||||
assert call_kwargs[1]["tag"] == "热门"
|
||||
|
||||
def test_list_with_keyword_filter(self):
|
||||
repo = _make_repo()
|
||||
repo.list_by_user = Mock(return_value=[])
|
||||
use_case = ListTemplatesUseCase(repo)
|
||||
f = ListTemplatesFilter(keyword="vlog")
|
||||
|
||||
use_case.execute("user-001", filter=f)
|
||||
|
||||
call_kwargs = repo.list_by_user.call_args
|
||||
assert call_kwargs[1]["keyword"] == "vlog"
|
||||
|
||||
def test_list_with_mode_filter(self):
|
||||
repo = _make_repo()
|
||||
repo.list_by_user = Mock(return_value=[])
|
||||
use_case = ListTemplatesUseCase(repo)
|
||||
f = ListTemplatesFilter(mode="one_take")
|
||||
|
||||
use_case.execute("user-001", filter=f)
|
||||
|
||||
call_kwargs = repo.list_by_user.call_args
|
||||
assert call_kwargs[1]["mode"] == "one_take"
|
||||
|
||||
def test_list_without_filter_uses_defaults(self):
|
||||
repo = _make_repo()
|
||||
repo.list_by_user = Mock(return_value=[])
|
||||
use_case = ListTemplatesUseCase(repo)
|
||||
|
||||
use_case.execute("user-001", skip=0, limit=50)
|
||||
|
||||
call_args = repo.list_by_user.call_args
|
||||
assert call_args[0][0] == "user-001"
|
||||
assert call_args[1]["skip"] == 0
|
||||
assert call_args[1]["limit"] == 50
|
||||
|
||||
|
||||
# ── CountTemplatesUseCase ──
|
||||
|
||||
|
||||
class TestCountTemplatesUseCase:
|
||||
def test_count_without_filter(self):
|
||||
repo = _make_repo()
|
||||
repo.count_by_user = Mock(return_value=5)
|
||||
use_case = CountTemplatesUseCase(repo)
|
||||
|
||||
result = use_case.execute("user-001")
|
||||
|
||||
assert result == 5
|
||||
repo.count_by_user.assert_called_once_with("user-001")
|
||||
|
||||
def test_count_with_filter(self):
|
||||
repo = _make_repo()
|
||||
repo.count_by_user = Mock(return_value=2)
|
||||
use_case = CountTemplatesUseCase(repo)
|
||||
f = ListTemplatesFilter(category="vlog", tag="热门")
|
||||
|
||||
result = use_case.execute("user-001", filter=f)
|
||||
|
||||
assert result == 2
|
||||
call_kwargs = repo.count_by_user.call_args
|
||||
assert call_kwargs[1]["category"] == "vlog"
|
||||
assert call_kwargs[1]["tag"] == "热门"
|
||||
|
||||
|
||||
# ── ListTagsUseCase ──
|
||||
|
||||
|
||||
class TestListTagsUseCase:
|
||||
def test_list_tags_returns_sorted(self):
|
||||
repo = _make_repo()
|
||||
repo.list_tags = Mock(return_value=["vlog", "热门", "教程"])
|
||||
use_case = ListTagsUseCase(repo)
|
||||
|
||||
result = use_case.execute("user-001")
|
||||
|
||||
assert result == ["vlog", "热门", "教程"]
|
||||
repo.list_tags.assert_called_once_with("user-001")
|
||||
|
||||
def test_list_tags_empty(self):
|
||||
repo = _make_repo()
|
||||
repo.list_tags = Mock(return_value=[])
|
||||
use_case = ListTagsUseCase(repo)
|
||||
|
||||
result = use_case.execute("user-001")
|
||||
|
||||
assert result == []
|
||||
|
||||
|
||||
# ── GetTemplateUsageUseCase ──
|
||||
|
||||
|
||||
class TestGetTemplateUsageUseCase:
|
||||
def test_get_usage_count(self):
|
||||
repo = _make_repo()
|
||||
repo.get_usage_count = Mock(return_value=3)
|
||||
use_case = GetTemplateUsageUseCase(repo)
|
||||
|
||||
result = use_case.execute("tmpl-001")
|
||||
|
||||
assert result == 3
|
||||
repo.get_usage_count.assert_called_once_with("tmpl-001")
|
||||
|
||||
def test_get_usage_zero(self):
|
||||
repo = _make_repo()
|
||||
repo.get_usage_count = Mock(return_value=0)
|
||||
use_case = GetTemplateUsageUseCase(repo)
|
||||
|
||||
result = use_case.execute("tmpl-001")
|
||||
|
||||
assert result == 0
|
||||
|
||||
Reference in New Issue
Block a user