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
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:
@@ -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")
|
||||
@@ -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"],
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
@@ -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:
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user