diff --git a/alembic/versions/030_add_tags_and_asset_tags.py b/alembic/versions/030_add_tags_and_asset_tags.py new file mode 100644 index 000000000..6dbfa212b --- /dev/null +++ b/alembic/versions/030_add_tags_and_asset_tags.py @@ -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") diff --git a/apps/api/app/api/router.py b/apps/api/app/api/router.py index 714181055..875833cda 100755 --- a/apps/api/app/api/router.py +++ b/apps/api/app/api/router.py @@ -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"], diff --git a/apps/api/app/api/routes/assets.py b/apps/api/app/api/routes/assets.py index 3725b3fda..f351699f7 100644 --- a/apps/api/app/api/routes/assets.py +++ b/apps/api/app/api/routes/assets.py @@ -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, diff --git a/apps/api/app/api/routes/tags.py b/apps/api/app/api/routes/tags.py new file mode 100644 index 000000000..7de3e278b --- /dev/null +++ b/apps/api/app/api/routes/tags.py @@ -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) diff --git a/apps/api/app/dependencies.py b/apps/api/app/dependencies.py index 5d38872f7..9c42a59b3 100755 --- a/apps/api/app/dependencies.py +++ b/apps/api/app/dependencies.py @@ -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: diff --git a/apps/api/app/schemas/asset.py b/apps/api/app/schemas/asset.py index ff0cefa53..1c686772f 100644 --- a/apps/api/app/schemas/asset.py +++ b/apps/api/app/schemas/asset.py @@ -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): diff --git a/apps/api/app/schemas/tag.py b/apps/api/app/schemas/tag.py new file mode 100644 index 000000000..70c8e2eb5 --- /dev/null +++ b/apps/api/app/schemas/tag.py @@ -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) diff --git a/packages/adapters/in_memory/asset_repository.py b/packages/adapters/in_memory/asset_repository.py index e194dc48d..6a3a64043 100755 --- a/packages/adapters/in_memory/asset_repository.py +++ b/packages/adapters/in_memory/asset_repository.py @@ -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] diff --git a/packages/adapters/in_memory/tag_repository.py b/packages/adapters/in_memory/tag_repository.py new file mode 100644 index 000000000..451b8848f --- /dev/null +++ b/packages/adapters/in_memory/tag_repository.py @@ -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 diff --git a/packages/adapters/sqlalchemy_impl/asset_repository.py b/packages/adapters/sqlalchemy_impl/asset_repository.py index 7fc53ff7e..ab6083043 100644 --- a/packages/adapters/sqlalchemy_impl/asset_repository.py +++ b/packages/adapters/sqlalchemy_impl/asset_repository.py @@ -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] diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index 982875e00..27a07513e 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -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 模型 diff --git a/packages/adapters/sqlalchemy_impl/tag_repository.py b/packages/adapters/sqlalchemy_impl/tag_repository.py new file mode 100644 index 000000000..4353484ca --- /dev/null +++ b/packages/adapters/sqlalchemy_impl/tag_repository.py @@ -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, + ) diff --git a/packages/domain/__init__.py b/packages/domain/__init__.py index e0ce5822d..7d2992d69 100755 --- a/packages/domain/__init__.py +++ b/packages/domain/__init__.py @@ -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", diff --git a/packages/domain/entities.py b/packages/domain/entities.py index ef568686c..b31897d15 100644 --- a/packages/domain/entities.py +++ b/packages/domain/entities.py @@ -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) diff --git a/packages/domain/tag.py b/packages/domain/tag.py new file mode 100644 index 000000000..b83b1e8b6 --- /dev/null +++ b/packages/domain/tag.py @@ -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, + ) diff --git a/packages/ports/__init__.py b/packages/ports/__init__.py index 5d915d17b..a4461e159 100644 --- a/packages/ports/__init__.py +++ b/packages/ports/__init__.py @@ -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", ] diff --git a/packages/ports/asset_repository.py b/packages/ports/asset_repository.py index 2432d0eb0..9f96d1f1c 100644 --- a/packages/ports/asset_repository.py +++ b/packages/ports/asset_repository.py @@ -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 diff --git a/packages/ports/tag_repository.py b/packages/ports/tag_repository.py new file mode 100644 index 000000000..901d2b12d --- /dev/null +++ b/packages/ports/tag_repository.py @@ -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 diff --git a/tests/integration/test_asset_tags.py b/tests/integration/test_asset_tags.py index 212af3c8a..bc3165203 100644 --- a/tests/integration/test_asset_tags.py +++ b/tests/integration/test_asset_tags.py @@ -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 diff --git a/tests/unit/test_asset_tagging.py b/tests/unit/test_asset_tagging.py new file mode 100644 index 000000000..2b36b40a4 --- /dev/null +++ b/tests/unit/test_asset_tagging.py @@ -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 diff --git a/tests/unit/test_tag_crud.py b/tests/unit/test_tag_crud.py new file mode 100644 index 000000000..894bd944a --- /dev/null +++ b/tests/unit/test_tag_crud.py @@ -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