feat: P3 配音模块标签体系 — 标签 CRUD + 素材打标 + tag_ids 筛选
CI/CD Pipeline / Deploy Staging (push) Has been skipped
CI/CD Pipeline / Staging E2E Tests (push) Has been skipped
CI/CD Pipeline / Build Production Runtime Images (push) Has been skipped
CI/CD Pipeline / Deploy Production (push) Has been skipped
CI/CD Pipeline / Production Browser E2E (push) Has been skipped
CI/CD Pipeline / Frontend Lint (push) Failing after 48h55m59s
CI/CD Pipeline / Validate Code Quality And Tests (push) Failing after 48h55m59s

- 新增 Tag 领域实体 + TagModel/AssetTagModel ORM 模型
- Alembic 迁移 030:tags 表 + asset_tags 关联表
- TagRepository 端口 + SQLAlchemy/InMemory 实现
- Asset.tag_ids 多对多关联替代原 JSON tags
- GET /tags / POST /tags / DELETE /tags/{tag_id} 标签 CRUD
- POST /assets/{asset_id}/tags 打标 / DELETE 取消标签
- GET /assets 新增 tag_ids 筛选参数(逗号分隔,取交集)
- 19 个单元测试全部通过
This commit is contained in:
灵应
2026-07-07 11:58:27 +08:00
parent d60a963b62
commit b0f2e4712a
21 changed files with 782 additions and 45 deletions
@@ -0,0 +1,65 @@
"""Add tags and asset_tags tables
Revision ID: 030
Revises: 029
Create Date: 2026-07-07
新增标签表和素材-标签关联表,支持规范化多对多标签管理。
"""
import sqlalchemy as sa
from alembic import op
revision = "030"
down_revision = "029"
branch_labels = None
depends_on = None
def _table_exists(table: str) -> bool:
conn = op.get_bind()
result = conn.execute(
sa.text("SELECT COUNT(*) FROM information_schema.tables WHERE table_name = :table"),
{"table": table},
)
return result.scalar() > 0
def upgrade() -> None:
if not _table_exists("tags"):
op.create_table(
"tags",
sa.Column("id", sa.String(36), primary_key=True),
sa.Column("user_id", sa.String(36), nullable=False),
sa.Column("name", sa.String(100), nullable=False),
sa.Column(
"created_at",
sa.DateTime(),
nullable=False,
server_default=sa.func.now(),
),
sa.UniqueConstraint("user_id", "name", name="uq_tags_user_name"),
)
op.create_index("ix_tags_user_id", "tags", ["user_id"])
if not _table_exists("asset_tags"):
op.create_table(
"asset_tags",
sa.Column("asset_id", sa.String(36), primary_key=True),
sa.Column("tag_id", sa.String(36), primary_key=True),
sa.Column(
"created_at",
sa.DateTime(),
nullable=False,
server_default=sa.func.now(),
),
)
op.create_index("ix_asset_tags_tag_id", "asset_tags", ["tag_id"])
def downgrade() -> None:
op.drop_index("ix_asset_tags_tag_id", table_name="asset_tags")
op.drop_table("asset_tags")
op.drop_index("ix_tags_user_id", table_name="tags")
op.drop_table("tags")
+6
View File
@@ -16,6 +16,7 @@ from app.api.routes.jobs import router as jobs_router
from app.api.routes.projects import router as projects_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.tags import router as tags_router
from app.api.routes.task_center import router as task_center_router
from app.api.routes.templates import router as templates_router
from app.api.routes.titles import router as titles_router
@@ -38,6 +39,11 @@ api_router.include_router(
prefix="/projects",
tags=["Project"],
)
api_router.include_router(
tags_router,
prefix="/tags",
tags=["Tag"],
)
api_router.include_router(
task_center_router,
tags=["TaskCenter"],
+59 -3
View File
@@ -7,6 +7,7 @@ from app.dependencies import (
get_asset_library_repository,
get_asset_repository,
get_project_repository,
get_tag_repository,
)
from app.schemas.asset import (
AssetResponse,
@@ -17,6 +18,7 @@ from app.schemas.asset import (
UpdateAssetRequest,
UpdateAssetReviewRequest,
)
from app.schemas.tag import TagAssetsRequest
from fastapi import APIRouter, Depends, HTTPException, Query
from packages.application import (
@@ -66,6 +68,7 @@ def _to_asset_response(item, storage_service=None) -> AssetResponse:
classification_status=item.classification_status.value,
quality_score=item.quality_score,
uploaded_by_user_id=item.uploaded_by_user_id,
tag_ids=getattr(item, "tag_ids", []),
)
@@ -86,6 +89,7 @@ def list_assets(
keyword: Optional[str] = Query(None, description="按名称模糊匹配"),
gender: Optional[str] = Query(None, description="按 metadata.gender 筛选"),
style: Optional[str] = Query(None, description="按 metadata.style 筛选"),
tag_ids: Optional[str] = Query(None, description="按标签 ID 筛选(逗号分隔,取交集)"),
skip: int = Query(0, ge=0),
limit: int = Query(100, ge=1, le=500),
authenticated_user: AuthenticatedUser = Depends(get_current_user),
@@ -98,11 +102,18 @@ def list_assets(
# kind → file_type 映射(voice 对应 audio
kind_to_file_type = {"video": "video", "voice": "audio", "image": "image"}
# 需要内存过滤的标志(keyword/gender/style 无法在 DB 层过滤
needs_memory_filter = bool(keyword or gender or style)
# 解析 tag_ids 参数(逗号分隔
filter_tag_ids: list[str] | None = None
if tag_ids:
filter_tag_ids = [t.strip() for t in tag_ids.split(",") if t.strip()]
if not filter_tag_ids:
filter_tag_ids = None
# 需要内存过滤的标志(keyword/gender/style/tag_ids 无法在 DB 层过滤)
needs_memory_filter = bool(keyword or gender or style or filter_tag_ids)
def _apply_memory_filters(items):
"""应用 keyword / gender / style 内存过滤。"""
"""应用 keyword / gender / style / tag_ids 内存过滤。"""
result = items
if keyword:
kw = keyword.lower()
@@ -111,6 +122,9 @@ def list_assets(
result = [i for i in result if (i.metadata or {}).get("gender") == gender]
if style:
result = [i for i in result if (i.metadata or {}).get("style") == style]
if filter_tag_ids:
tag_set = set(filter_tag_ids)
result = [i for i in result if tag_set.issubset(set(getattr(i, "tag_ids", [])))]
return result
# ── 优化路径:无内存过滤时,使用 DB 级分页 ──
@@ -336,6 +350,48 @@ def delete_asset(
asset_repository.delete(asset_id)
@router.post("/{asset_id}/tags", response_model=AssetResponse)
def tag_asset(
asset_id: str,
request: TagAssetsRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
asset_repository: Any = Depends(get_asset_repository),
project_repository: Any = Depends(get_project_repository),
tag_repository: Any = Depends(get_tag_repository),
) -> AssetResponse:
"""给素材打标签。"""
item = asset_repository.find_by_id(asset_id)
if item is None:
raise HTTPException(status_code=404, detail=f"Asset {asset_id} not found")
_check_project_access(item.project_id, authenticated_user.user.id, project_repository)
for tag_id in request.tag_ids:
tag = tag_repository.get(tag_id)
if tag is None:
raise HTTPException(status_code=404, detail=f"Tag {tag_id} not found")
if tag.user_id != authenticated_user.user.id:
raise HTTPException(status_code=403, detail=f"无权使用标签 {tag_id}")
item.add_tag(tag_id)
updated = asset_repository.update(item)
return _to_asset_response(updated)
@router.delete("/{asset_id}/tags/{tag_id}", status_code=204)
def untag_asset(
asset_id: str,
tag_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
asset_repository: Any = Depends(get_asset_repository),
project_repository: Any = Depends(get_project_repository),
) -> None:
"""取消素材的标签。"""
item = asset_repository.find_by_id(asset_id)
if item is None:
raise HTTPException(status_code=404, detail=f"Asset {asset_id} not found")
_check_project_access(item.project_id, authenticated_user.user.id, project_repository)
item.remove_tag(tag_id)
asset_repository.update(item)
@router.post("", response_model=AssetResponse)
def create_asset(
request: CreateAssetRequest,
+67
View File
@@ -0,0 +1,67 @@
"""标签 CRUD 路由。"""
import logging
from typing import Any
from app.auth import AuthenticatedUser, get_current_user
from app.dependencies import get_tag_repository
from app.schemas.tag import (
CreateTagRequest,
ListTagsResponse,
TagResponse,
)
from fastapi import APIRouter, Depends, HTTPException
from packages.domain import Tag
logger = logging.getLogger(__name__)
router = APIRouter()
@router.get("", response_model=ListTagsResponse)
def list_tags(
skip: int = 0,
limit: int = 100,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
tag_repository: Any = Depends(get_tag_repository),
) -> ListTagsResponse:
"""列出当前用户的标签。"""
user_id = authenticated_user.user.id
items = tag_repository.list_by_user(user_id, skip=skip, limit=limit)
total = tag_repository.count_by_user(user_id)
return ListTagsResponse(
items=[TagResponse(id=t.id, name=t.name, created_at=t.created_at) for t in items],
total=total,
)
@router.post("", response_model=TagResponse, status_code=201)
def create_tag(
request: CreateTagRequest,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
tag_repository: Any = Depends(get_tag_repository),
) -> TagResponse:
"""创建标签(同用户同名去重,返回 409)。"""
user_id = authenticated_user.user.id
existing = tag_repository.find_by_name(user_id, request.name)
if existing:
raise HTTPException(status_code=409, detail="标签名称已存在")
tag = Tag.create(user_id=user_id, name=request.name)
created = tag_repository.create(tag)
return TagResponse(id=created.id, name=created.name, created_at=created.created_at)
@router.delete("/{tag_id}", status_code=204)
def delete_tag(
tag_id: str,
authenticated_user: AuthenticatedUser = Depends(get_current_user),
tag_repository: Any = Depends(get_tag_repository),
) -> None:
"""删除标签(同时清理素材关联)。"""
tag = tag_repository.get(tag_id)
if tag is None:
raise HTTPException(status_code=404, detail="标签不存在")
if tag.user_id != authenticated_user.user.id:
raise HTTPException(status_code=403, detail="无权删除该标签")
tag_repository.delete(tag_id)
+9
View File
@@ -39,6 +39,7 @@ from packages.adapters.sqlalchemy_impl.project_repository import (
SQLAlchemyProjectRepository,
)
from packages.adapters.sqlalchemy_impl.session import build_session_factory
from packages.adapters.sqlalchemy_impl.tag_repository import SQLAlchemyTagRepository
from packages.adapters.sqlalchemy_impl.title_library_repository import (
SQLAlchemyTitleLibraryRepository,
)
@@ -58,6 +59,7 @@ from packages.ports.generation_task_repository import GenerationTaskRepository
from packages.ports.ingest_job_repository import IngestJobRepository
from packages.ports.job_repository import JobRepository
from packages.ports.project_repository import ProjectRepository
from packages.ports.tag_repository import TagRepository
from packages.ports.title_library_repository import TitleLibraryRepository
from packages.ports.user_repository import UserRepository
from packages.ports.voice_clone_profile_repository import VoiceCloneProfileRepository
@@ -138,6 +140,13 @@ def get_project_repository(
return SQLAlchemyProjectRepository(session)
def get_tag_repository(
session: Session = Depends(get_db_session),
) -> TagRepository:
"""Provide the SQLAlchemy tag repository implementation."""
return SQLAlchemyTagRepository(session)
def get_user_repository(
session: Session = Depends(get_db_session),
) -> UserRepository:
+1
View File
@@ -51,6 +51,7 @@ class AssetResponse(BaseModel):
classification_status: str
quality_score: float | None = None
uploaded_by_user_id: str
tag_ids: list[str] = Field(default_factory=list)
class BatchDeleteRequest(BaseModel):
+24
View File
@@ -0,0 +1,24 @@
"""标签相关 Schema。"""
from datetime import datetime
from pydantic import BaseModel, Field
class CreateTagRequest(BaseModel):
name: str = Field(..., min_length=1, max_length=100)
class TagResponse(BaseModel):
id: str
name: str
created_at: datetime
class ListTagsResponse(BaseModel):
items: list[TagResponse]
total: int = Field(default=0, ge=0)
class TagAssetsRequest(BaseModel):
tag_ids: list[str] = Field(..., min_length=1, max_length=50)
@@ -51,3 +51,28 @@ class InMemoryAssetRepository:
del self._assets[aid]
count += 1
return count
def find_by_project(
self,
project_id: str,
skip: int = 0,
limit: int = 100,
) -> list[Asset]:
items = [a for a in self._assets.values() if a.project_id == project_id]
return items[skip : skip + limit]
def find_by_id(self, asset_id: str) -> Asset | None:
return self._assets.get(asset_id)
def find_by_tag_ids(
self,
tag_ids: list[str],
skip: int = 0,
limit: int = 100,
) -> list[Asset]:
"""查找包含所有指定标签的素材。"""
if not tag_ids:
return []
tag_set = set(tag_ids)
items = [a for a in self._assets.values() if tag_set.issubset(set(a.tag_ids))]
return items[skip : skip + limit]
@@ -0,0 +1,40 @@
"""标签 InMemory 仓储实现。"""
from packages.domain import Tag
class InMemoryTagRepository:
def __init__(self):
self._tags: dict[str, Tag] = {}
def create(self, tag: Tag) -> Tag:
self._tags[tag.id] = tag
return tag
def get(self, tag_id: str) -> Tag | None:
return self._tags.get(tag_id)
def find_by_name(self, user_id: str, name: str) -> Tag | None:
for tag in self._tags.values():
if tag.user_id == user_id and tag.name == name:
return tag
return None
def list_by_user(
self,
user_id: str,
skip: int = 0,
limit: int = 100,
) -> list[Tag]:
tags = [tag for tag in self._tags.values() if tag.user_id == user_id]
tags.sort(key=lambda t: t.created_at, reverse=True)
return tags[skip : skip + limit]
def count_by_user(self, user_id: str) -> int:
return sum(1 for tag in self._tags.values() if tag.user_id == user_id)
def delete(self, tag_id: str) -> bool:
if tag_id in self._tags:
del self._tags[tag_id]
return True
return False
@@ -3,7 +3,7 @@ from datetime import datetime, timezone
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import AssetModel
from packages.adapters.sqlalchemy_impl.models import AssetModel, AssetTagModel
from packages.domain import Asset, AssetStatus, ClassificationStatus
@@ -87,6 +87,8 @@ class SQLAlchemyAssetRepository:
updated_at=now,
)
self.session.add(model)
self.session.flush()
self._sync_asset_tags(asset.id, asset.tag_ids)
self.session.commit()
return asset
@@ -109,6 +111,8 @@ class SQLAlchemyAssetRepository:
model.quality_score = asset.quality_score
model.uploaded_by_user_id = asset.uploaded_by_user_id or model.uploaded_by_user_id
model.updated_at = datetime.now(timezone.utc)
self.session.flush()
self._sync_asset_tags(asset.id, asset.tag_ids)
self.session.commit()
return asset
@@ -203,6 +207,11 @@ class SQLAlchemyAssetRepository:
"audio": "audio/mpeg",
"image": "image/jpeg",
}.get(mime_type, mime_type)
# 查询关联的 tag_ids
tag_ids = [
row.tag_id
for row in self.session.query(AssetTagModel.tag_id).filter(AssetTagModel.asset_id == model.id).all()
]
return Asset(
id=model.id,
project_id=model.project_id,
@@ -222,6 +231,39 @@ class SQLAlchemyAssetRepository:
quality_score=model.quality_score,
uploaded_by_user_id=model.uploaded_by_user_id,
metadata=metadata,
tag_ids=tag_ids,
created_at=model.created_at,
updated_at=model.updated_at,
)
def _sync_asset_tags(self, asset_id: str, tag_ids: list[str]) -> None:
"""同步素材-标签关联表(全量替换)。"""
self.session.query(AssetTagModel).filter(AssetTagModel.asset_id == asset_id).delete(synchronize_session=False)
for tag_id in tag_ids:
self.session.add(AssetTagModel(asset_id=asset_id, tag_id=tag_id))
def find_by_tag_ids(
self,
tag_ids: list[str],
skip: int = 0,
limit: int = 100,
) -> list[Asset]:
"""查找包含所有指定标签的素材。"""
if not tag_ids:
return []
from sqlalchemy import func
# 找出同时拥有所有指定 tag_id 的 asset_id
tag_set = set(tag_ids)
asset_ids = (
self.session.query(AssetTagModel.asset_id)
.filter(AssetTagModel.tag_id.in_(tag_set))
.group_by(AssetTagModel.asset_id)
.having(func.count(AssetTagModel.tag_id) == len(tag_set))
.all()
)
ids = [row[0] for row in asset_ids]
if not ids:
return []
models = self.session.query(AssetModel).filter(AssetModel.id.in_(ids)).offset(skip).limit(limit).all()
return [self._to_domain(m) for m in models]
@@ -90,6 +90,29 @@ class AssetModel(Base):
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
class TagModel(Base):
"""标签 ORM 模型。"""
__tablename__ = "tags"
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))
__table_args__ = (UniqueConstraint("user_id", "name", name="uq_tags_user_name"),)
class AssetTagModel(Base):
"""素材-标签关联表 ORM 模型。"""
__tablename__ = "asset_tags"
asset_id = Column(String(36), primary_key=True)
tag_id = Column(String(36), primary_key=True)
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
class EditTemplateModel(Base):
"""Phase 8 剪辑模板 ORM 模型
@@ -0,0 +1,73 @@
"""标签 SQLAlchemy 仓储实现。"""
from sqlalchemy.orm import Session
from packages.adapters.sqlalchemy_impl.models import AssetTagModel, TagModel
from packages.domain import Tag
class SQLAlchemyTagRepository:
def __init__(self, session: Session):
self.session = session
def create(self, tag: Tag) -> Tag:
model = TagModel(
id=tag.id,
user_id=tag.user_id,
name=tag.name,
created_at=tag.created_at,
)
self.session.add(model)
self.session.commit()
return tag
def get(self, tag_id: str) -> Tag | None:
model = self.session.query(TagModel).filter(TagModel.id == tag_id).first()
if model is None:
return None
return self._to_domain(model)
def find_by_name(self, user_id: str, name: str) -> Tag | None:
model = self.session.query(TagModel).filter(TagModel.user_id == user_id, TagModel.name == name).first()
if model is None:
return None
return self._to_domain(model)
def list_by_user(
self,
user_id: str,
skip: int = 0,
limit: int = 100,
) -> list[Tag]:
models = (
self.session.query(TagModel)
.filter(TagModel.user_id == user_id)
.order_by(TagModel.created_at.desc())
.offset(skip)
.limit(limit)
.all()
)
return [self._to_domain(m) for m in models]
def count_by_user(self, user_id: str) -> int:
return self.session.query(TagModel).filter(TagModel.user_id == user_id).count()
def delete(self, tag_id: str) -> bool:
# 先清理关联表
self.session.query(AssetTagModel).filter(AssetTagModel.tag_id == tag_id).delete(synchronize_session=False)
model = self.session.query(TagModel).filter(TagModel.id == tag_id).first()
if model is None:
self.session.commit()
return False
self.session.delete(model)
self.session.commit()
return True
@staticmethod
def _to_domain(model: TagModel) -> Tag:
return Tag(
id=model.id,
user_id=model.user_id,
name=model.name,
created_at=model.created_at,
)
+2
View File
@@ -24,6 +24,7 @@ from .entities import (
from .generated_video import GeneratedVideo
from .generation_task import GenerationTask, GenerationTaskStatus
from .job import Job, JobStatus, JobType
from .tag import Tag
from .template_clip_config import ClipType, TemplateClipConfig, TransitionEffect
from .title_library import TitleLibraryItem
from .voice_library import VoiceLibraryItem
@@ -56,6 +57,7 @@ __all__ = [
"JobStatus",
"JobType",
"Project",
"Tag",
"TemplateClipConfig",
"TransitionEffect",
"User",
+14 -14
View File
@@ -162,7 +162,7 @@ class Asset:
quality_score: float | None = None
uploaded_by_user_id: str = ""
metadata: dict[str, Any] = field(default_factory=dict)
tags: list[str] = field(default_factory=list)
tag_ids: list[str] = field(default_factory=list)
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
@@ -214,23 +214,23 @@ class Asset:
quality_score=quality_score,
uploaded_by_user_id=uploaded_by_user_id.strip(),
metadata=metadata or {},
tags=[],
tag_ids=[],
)
def add_tag(self, tag: str) -> None:
"""添加标签。空标签会被忽略,自动去重。"""
clean_tag = tag.strip()
if not clean_tag:
raise ValueError("标签不能为空")
if clean_tag not in self.tags:
self.tags.append(clean_tag)
def add_tag(self, tag_id: str) -> None:
"""添加标签 ID。空 ID 会被忽略,自动去重。"""
clean_id = tag_id.strip()
if not clean_id:
raise ValueError("标签 ID 不能为空")
if clean_id not in self.tag_ids:
self.tag_ids.append(clean_id)
self.updated_at = datetime.now(timezone.utc)
def remove_tag(self, tag: str) -> None:
"""删除标签。如果标签不存在,不报错(幂等性)。"""
clean_tag = tag.strip()
if clean_tag in self.tags:
self.tags.remove(clean_tag)
def remove_tag(self, tag_id: str) -> None:
"""删除标签 ID。如果标签不存在,不报错(幂等性)。"""
clean_id = tag_id.strip()
if clean_id in self.tag_ids:
self.tag_ids.remove(clean_id)
self.updated_at = datetime.now(timezone.utc)
+24
View File
@@ -0,0 +1,24 @@
"""标签领域实体。"""
from dataclasses import dataclass, field
from datetime import datetime, timezone
from uuid import uuid4
@dataclass(slots=True)
class Tag:
id: str
user_id: str
name: str
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
@classmethod
def create(cls, user_id: str, name: str) -> "Tag":
clean_name = name.strip()
if not clean_name:
raise ValueError("标签名称不能为空")
return cls(
id=uuid4().hex,
user_id=user_id,
name=clean_name,
)
+2
View File
@@ -4,6 +4,7 @@ from .asset_library_repository import AssetLibraryRepository
from .asset_repository import AssetRepository
from .ingest_job_repository import IngestJobRepository
from .project_repository import ProjectRepository
from .tag_repository import TagRepository
from .title_library_repository import TitleLibraryRepository
from .voice_library_repository import VoiceLibraryRepository
@@ -12,6 +13,7 @@ __all__ = [
"AssetRepository",
"IngestJobRepository",
"ProjectRepository",
"TagRepository",
"TitleLibraryRepository",
"VoiceLibraryRepository",
]
+10
View File
@@ -83,3 +83,13 @@ class AssetRepository(ABC):
) -> list[Asset]:
"""按筛选条件搜索候选素材,按质量分降序排列。"""
pass
@abstractmethod
def find_by_tag_ids(
self,
tag_ids: list[str],
skip: int = 0,
limit: int = 100,
) -> list[Asset]:
"""查找包含所有指定标签的素材。"""
pass
+36
View File
@@ -0,0 +1,36 @@
"""标签仓储接口定义。"""
from abc import ABC, abstractmethod
from packages.domain import Tag
class TagRepository(ABC):
@abstractmethod
def create(self, tag: Tag) -> Tag:
pass
@abstractmethod
def get(self, tag_id: str) -> Tag | None:
pass
@abstractmethod
def find_by_name(self, user_id: str, name: str) -> Tag | None:
pass
@abstractmethod
def list_by_user(
self,
user_id: str,
skip: int = 0,
limit: int = 100,
) -> list[Tag]:
pass
@abstractmethod
def count_by_user(self, user_id: str) -> int:
pass
@abstractmethod
def delete(self, tag_id: str) -> bool:
pass
+27 -27
View File
@@ -4,7 +4,7 @@ from packages.domain import Asset
def test_add_tag_to_asset():
"""测试添加标签到 Asset。"""
"""测试添加标签 ID 到 Asset。"""
asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
@@ -13,16 +13,16 @@ def test_add_tag_to_asset():
mime_type="video/mp4",
)
asset.add_tag("风景")
asset.add_tag("自然")
asset.add_tag("tag-1")
asset.add_tag("tag-2")
assert len(asset.tags) == 2
assert "风景" in asset.tags
assert "自然" in asset.tags
assert len(asset.tag_ids) == 2
assert "tag-1" in asset.tag_ids
assert "tag-2" in asset.tag_ids
def test_add_duplicate_tag_should_ignore():
"""测试添加重复标签应自动去重。"""
"""测试添加重复标签 ID 应自动去重。"""
asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
@@ -31,15 +31,15 @@ def test_add_duplicate_tag_should_ignore():
mime_type="video/mp4",
)
asset.add_tag("风景")
asset.add_tag("风景") # 重复
asset.add_tag("tag-1")
asset.add_tag("tag-1") # 重复
assert len(asset.tags) == 1
assert asset.tags.count("风景") == 1
assert len(asset.tag_ids) == 1
assert asset.tag_ids.count("tag-1") == 1
def test_add_empty_tag_should_fail():
"""测试添加空标签应失败。"""
"""测试添加空标签 ID 应失败。"""
asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
@@ -48,15 +48,15 @@ def test_add_empty_tag_should_fail():
mime_type="video/mp4",
)
with pytest.raises(ValueError, match="标签不能为空"):
with pytest.raises(ValueError, match="标签 ID 不能为空"):
asset.add_tag("")
with pytest.raises(ValueError, match="标签不能为空"):
with pytest.raises(ValueError, match="标签 ID 不能为空"):
asset.add_tag(" ") # 仅空格
def test_remove_tag_from_asset():
"""测试从 Asset 删除标签。"""
"""测试从 Asset 删除标签 ID"""
asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
@@ -65,18 +65,18 @@ def test_remove_tag_from_asset():
mime_type="video/mp4",
)
asset.add_tag("风景")
asset.add_tag("自然")
asset.add_tag("tag-1")
asset.add_tag("tag-2")
asset.remove_tag("风景")
asset.remove_tag("tag-1")
assert len(asset.tags) == 1
assert "风景" not in asset.tags
assert "自然" in asset.tags
assert len(asset.tag_ids) == 1
assert "tag-1" not in asset.tag_ids
assert "tag-2" in asset.tag_ids
def test_remove_nonexistent_tag_should_be_idempotent():
"""测试删除不存在的标签应幂等(不报错)。"""
"""测试删除不存在的标签 ID 应幂等(不报错)。"""
asset = Asset.create(
project_id="proj-1",
library_id="lib-1",
@@ -85,10 +85,10 @@ def test_remove_nonexistent_tag_should_be_idempotent():
mime_type="video/mp4",
)
asset.add_tag("风景")
asset.add_tag("tag-1")
# 删除不存在的标签,不应报错
asset.remove_tag("不存在的标签")
# 删除不存在的标签 ID,不应报错
asset.remove_tag("nonexistent-tag")
assert len(asset.tags) == 1
assert "风景" in asset.tags
assert len(asset.tag_ids) == 1
assert "tag-1" in asset.tag_ids
+130
View File
@@ -0,0 +1,130 @@
"""素材打标/取消标签单元测试。"""
import pytest
from packages.adapters.in_memory.asset_repository import InMemoryAssetRepository
from packages.adapters.in_memory.tag_repository import InMemoryTagRepository
from packages.domain import Asset, Tag
@pytest.fixture
def asset_repo():
return InMemoryAssetRepository()
@pytest.fixture
def tag_repo():
return InMemoryTagRepository()
def _create_asset(asset_repo, **kwargs):
defaults = dict(
project_id="proj-1",
library_id="lib-1",
name="video.mp4",
storage_key="uploads/abc/video.mp4",
mime_type="video/mp4",
)
defaults.update(kwargs)
asset = Asset.create(**defaults)
return asset_repo.create(asset)
def test_tag_asset(asset_repo, tag_repo):
"""测试给素材打标签。"""
asset = _create_asset(asset_repo)
tag = tag_repo.create(Tag.create(user_id="user-1", name="风景"))
asset.add_tag(tag.id)
asset_repo.update(asset)
loaded = asset_repo.get(asset.id)
assert tag.id in loaded.tag_ids
def test_untag_asset(asset_repo, tag_repo):
"""测试取消素材标签。"""
asset = _create_asset(asset_repo)
tag = tag_repo.create(Tag.create(user_id="user-1", name="风景"))
asset.add_tag(tag.id)
asset_repo.update(asset)
asset.remove_tag(tag.id)
asset_repo.update(asset)
loaded = asset_repo.get(asset.id)
assert tag.id not in loaded.tag_ids
def test_tag_multiple_assets(asset_repo, tag_repo):
"""测试同一标签打给多个素材。"""
a1 = _create_asset(asset_repo, name="a.mp4")
a2 = _create_asset(asset_repo, name="b.mp4")
tag = tag_repo.create(Tag.create(user_id="user-1", name="风景"))
a1.add_tag(tag.id)
a2.add_tag(tag.id)
asset_repo.update(a1)
asset_repo.update(a2)
assert tag.id in asset_repo.get(a1.id).tag_ids
assert tag.id in asset_repo.get(a2.id).tag_ids
def test_find_by_tag_ids(asset_repo, tag_repo):
"""测试按标签 ID 筛选素材。"""
tag1 = tag_repo.create(Tag.create(user_id="user-1", name="风景"))
tag2 = tag_repo.create(Tag.create(user_id="user-1", name="自然"))
a1 = _create_asset(asset_repo, name="a.mp4")
a1.add_tag(tag1.id)
a1.add_tag(tag2.id)
asset_repo.update(a1)
a2 = _create_asset(asset_repo, name="b.mp4")
a2.add_tag(tag1.id)
asset_repo.update(a2)
a3 = _create_asset(asset_repo, name="c.mp4")
# 无标签
# 按 tag1 筛选 → a1, a2
result = asset_repo.find_by_tag_ids([tag1.id])
ids = {a.id for a in result}
assert ids == {a1.id, a2.id}
# 按 tag1 + tag2 筛选(交集)→ a1
result = asset_repo.find_by_tag_ids([tag1.id, tag2.id])
ids = {a.id for a in result}
assert ids == {a1.id}
# 空 tag_ids → 空结果
assert asset_repo.find_by_tag_ids([]) == []
def test_delete_tag_cleans_associations(asset_repo, tag_repo):
"""测试删除标签后素材的 tag_ids 不受影响(关联表由仓储层清理)。"""
asset = _create_asset(asset_repo)
tag = tag_repo.create(Tag.create(user_id="user-1", name="风景"))
asset.add_tag(tag.id)
asset_repo.update(asset)
# 删除标签
tag_repo.delete(tag.id)
assert tag_repo.get(tag.id) is None
# 素材的 tag_ids 在内存中仍有,但重新加载后 InMemory 不感知关联表
# 实际 SQLAlchemy 实现中 _sync_asset_tags 会在 update 时清理
def test_duplicate_tag_id_ignored(asset_repo, tag_repo):
"""测试重复打同一标签自动去重。"""
asset = _create_asset(asset_repo)
tag = tag_repo.create(Tag.create(user_id="user-1", name="风景"))
asset.add_tag(tag.id)
asset.add_tag(tag.id) # 重复
assert asset.tag_ids.count(tag.id) == 1
+102
View File
@@ -0,0 +1,102 @@
"""标签 CRUD 单元测试(使用 InMemoryTagRepository)。"""
import pytest
from packages.adapters.in_memory.tag_repository import InMemoryTagRepository
from packages.domain import Tag
@pytest.fixture
def tag_repo():
return InMemoryTagRepository()
def test_create_tag(tag_repo):
"""测试创建标签。"""
tag = Tag.create(user_id="user-1", name="风景")
created = tag_repo.create(tag)
assert created.id == tag.id
assert created.user_id == "user-1"
assert created.name == "风景"
def test_create_tag_strips_whitespace(tag_repo):
"""测试创建标签时自动去除首尾空格。"""
tag = Tag.create(user_id="user-1", name=" 风景 ")
assert tag.name == "风景"
def test_create_tag_empty_name_raises():
"""测试空名称抛出 ValueError。"""
with pytest.raises(ValueError, match="标签名称不能为空"):
Tag.create(user_id="user-1", name="")
with pytest.raises(ValueError, match="标签名称不能为空"):
Tag.create(user_id="user-1", name=" ")
def test_get_tag(tag_repo):
"""测试按 ID 获取标签。"""
tag = Tag.create(user_id="user-1", name="风景")
tag_repo.create(tag)
found = tag_repo.get(tag.id)
assert found is not None
assert found.name == "风景"
assert tag_repo.get("nonexistent") is None
def test_find_by_name(tag_repo):
"""测试按用户 ID + 名称查找标签。"""
tag = Tag.create(user_id="user-1", name="风景")
tag_repo.create(tag)
found = tag_repo.find_by_name("user-1", "风景")
assert found is not None
assert found.id == tag.id
# 不同用户同名标签不冲突
assert tag_repo.find_by_name("user-2", "风景") is None
# 不存在的名称
assert tag_repo.find_by_name("user-1", "不存在") is None
def test_list_by_user(tag_repo):
"""测试按用户列出标签(分页)。"""
for i in range(5):
tag_repo.create(Tag.create(user_id="user-1", name=f"标签{i}"))
# 另一个用户的标签
tag_repo.create(Tag.create(user_id="user-2", name="其他用户标签"))
items = tag_repo.list_by_user("user-1")
assert len(items) == 5
# 分页
items_page = tag_repo.list_by_user("user-1", skip=2, limit=2)
assert len(items_page) == 2
def test_count_by_user(tag_repo):
"""测试按用户统计标签数量。"""
for i in range(3):
tag_repo.create(Tag.create(user_id="user-1", name=f"标签{i}"))
tag_repo.create(Tag.create(user_id="user-2", name="其他"))
assert tag_repo.count_by_user("user-1") == 3
assert tag_repo.count_by_user("user-2") == 1
assert tag_repo.count_by_user("user-3") == 0
def test_delete_tag(tag_repo):
"""测试删除标签。"""
tag = Tag.create(user_id="user-1", name="风景")
tag_repo.create(tag)
assert tag_repo.delete(tag.id) is True
assert tag_repo.get(tag.id) is None
# 重复删除返回 False
assert tag_repo.delete(tag.id) is False