feat: Phase 1 - 核心重构(去Project层/标题库API/配音库API/清理废弃代码) #74
@@ -0,0 +1,144 @@
|
||||
"""Phase 1 - 核心重构:清理废弃表
|
||||
|
||||
Revision ID: 011
|
||||
Revises: 010
|
||||
Create Date: 2026-06-28
|
||||
|
||||
This migration:
|
||||
1. Drops 6 deprecated tables:
|
||||
- tasks (任务管理)
|
||||
- milestones (里程碑)
|
||||
- task_issues (任务问题)
|
||||
- project_titles (项目标题,已被 title_libraries 替代)
|
||||
- edit_plans (编辑计划)
|
||||
- edit_plan_clips (编辑计划片段)
|
||||
2. Removes edit_plan_id column from generation_tasks table
|
||||
"""
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
# revision identifiers
|
||||
revision = "011"
|
||||
down_revision = "010"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# ── 1. Drop deprecated tables ──
|
||||
|
||||
# Drop in reverse dependency order
|
||||
conn.execute(sa.text("DROP TABLE IF EXISTS task_issues"))
|
||||
conn.execute(sa.text("DROP TABLE IF EXISTS milestones"))
|
||||
conn.execute(sa.text("DROP TABLE IF EXISTS tasks"))
|
||||
conn.execute(sa.text("DROP TABLE IF EXISTS project_titles"))
|
||||
conn.execute(sa.text("DROP TABLE IF EXISTS edit_plan_clips"))
|
||||
conn.execute(sa.text("DROP TABLE IF EXISTS edit_plans"))
|
||||
|
||||
# ── 2. Remove edit_plan_id from generation_tasks ──
|
||||
|
||||
conn.execute(sa.text(
|
||||
"ALTER TABLE generation_tasks DROP COLUMN IF EXISTS edit_plan_id"
|
||||
))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
conn = op.get_bind()
|
||||
|
||||
# ── 1. Re-add edit_plan_id to generation_tasks ──
|
||||
|
||||
conn.execute(sa.text(
|
||||
"ALTER TABLE generation_tasks ADD COLUMN IF NOT EXISTS edit_plan_id VARCHAR(32)"
|
||||
))
|
||||
|
||||
# ── 2. Recreate deprecated tables (basic structure) ──
|
||||
|
||||
# Note: Full schema recreation is complex; this is a minimal downgrade
|
||||
conn.execute(sa.text("""
|
||||
CREATE TABLE IF NOT EXISTS edit_plans (
|
||||
id VARCHAR(32) PRIMARY KEY,
|
||||
project_id VARCHAR(32) NOT NULL,
|
||||
name VARCHAR(255) NOT NULL,
|
||||
description TEXT NOT NULL DEFAULT '',
|
||||
status VARCHAR(20) NOT NULL DEFAULT 'draft',
|
||||
created_by_user_id VARCHAR(32) NOT NULL DEFAULT '',
|
||||
metadata JSONB NOT NULL DEFAULT '{}',
|
||||
created_at TIMESTAMP NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT NOW()
|
||||
)
|
||||
"""))
|
||||
|
||||
conn.execute(sa.text("""
|
||||
CREATE TABLE IF NOT EXISTS edit_plan_clips (
|
||||
id VARCHAR(32) PRIMARY KEY,
|
||||
edit_plan_id VARCHAR(32) NOT NULL,
|
||||
asset_id VARCHAR(32) NOT NULL,
|
||||
order_index INTEGER NOT NULL DEFAULT 0,
|
||||
start_time FLOAT NOT NULL DEFAULT 0,
|
||||
end_time FLOAT NOT NULL DEFAULT 0,
|
||||
metadata JSONB NOT NULL DEFAULT '{}',
|
||||
created_at TIMESTAMP NOT NULL DEFAULT NOW()
|
||||
)
|
||||
"""))
|
||||
|
||||
conn.execute(sa.text("""
|
||||
CREATE TABLE IF NOT EXISTS project_titles (
|
||||
id VARCHAR(36) PRIMARY KEY,
|
||||
project_id VARCHAR(36) NOT NULL,
|
||||
text VARCHAR(500) NOT NULL,
|
||||
category VARCHAR(50) NOT NULL DEFAULT 'default',
|
||||
source VARCHAR(20) NOT NULL DEFAULT 'manual',
|
||||
tags JSONB NOT NULL DEFAULT '[]',
|
||||
favorite BOOLEAN NOT NULL DEFAULT FALSE,
|
||||
usage_count INTEGER NOT NULL DEFAULT 0,
|
||||
metadata JSONB NOT NULL DEFAULT '{}',
|
||||
created_at TIMESTAMP NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT NOW()
|
||||
)
|
||||
"""))
|
||||
|
||||
conn.execute(sa.text("""
|
||||
CREATE TABLE IF NOT EXISTS tasks (
|
||||
id VARCHAR(32) PRIMARY KEY,
|
||||
project_id VARCHAR(32) NOT NULL,
|
||||
title VARCHAR(255) NOT NULL,
|
||||
description TEXT NOT NULL DEFAULT '',
|
||||
status VARCHAR(20) NOT NULL DEFAULT 'pending',
|
||||
priority VARCHAR(20) NOT NULL DEFAULT 'medium',
|
||||
assigned_to_user_id VARCHAR(32) NOT NULL DEFAULT '',
|
||||
due_date TIMESTAMP,
|
||||
metadata JSONB NOT NULL DEFAULT '{}',
|
||||
created_at TIMESTAMP NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT NOW()
|
||||
)
|
||||
"""))
|
||||
|
||||
conn.execute(sa.text("""
|
||||
CREATE TABLE IF NOT EXISTS milestones (
|
||||
id VARCHAR(32) PRIMARY KEY,
|
||||
project_id VARCHAR(32) NOT NULL,
|
||||
name VARCHAR(255) NOT NULL,
|
||||
description TEXT NOT NULL DEFAULT '',
|
||||
due_date TIMESTAMP,
|
||||
status VARCHAR(20) NOT NULL DEFAULT 'pending',
|
||||
metadata JSONB NOT NULL DEFAULT '{}',
|
||||
created_at TIMESTAMP NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT NOW()
|
||||
)
|
||||
"""))
|
||||
|
||||
conn.execute(sa.text("""
|
||||
CREATE TABLE IF NOT EXISTS task_issues (
|
||||
id VARCHAR(32) PRIMARY KEY,
|
||||
task_id VARCHAR(32) NOT NULL,
|
||||
title VARCHAR(255) NOT NULL,
|
||||
description TEXT NOT NULL DEFAULT '',
|
||||
status VARCHAR(20) NOT NULL DEFAULT 'open',
|
||||
priority VARCHAR(20) NOT NULL DEFAULT 'medium',
|
||||
metadata JSONB NOT NULL DEFAULT '{}',
|
||||
created_at TIMESTAMP NOT NULL DEFAULT NOW(),
|
||||
updated_at TIMESTAMP NOT NULL DEFAULT NOW()
|
||||
)
|
||||
"""))
|
||||
+10
-13
@@ -4,13 +4,12 @@ from app.api.routes.assets import router as assets_router
|
||||
from app.api.routes.auth import router as auth_router
|
||||
from app.api.routes.chunked_upload import router as chunked_upload_router
|
||||
from app.api.routes.classification_jobs import router as classification_jobs_router
|
||||
from app.api.routes.edit_plans import router as edit_plans_router
|
||||
from app.api.routes.generated_videos import router as generated_videos_router
|
||||
from app.api.routes.titles import router as titles_router
|
||||
from app.api.routes.voices import router as voices_router
|
||||
from app.api.routes.generation_tasks import router as generation_tasks_router
|
||||
from app.api.routes.health import router as health_check_router
|
||||
from app.api.routes.ingest_jobs import router as ingest_jobs_router
|
||||
from app.api.routes.project_management import router as project_management_router
|
||||
from app.api.routes.project_titles import router as project_titles_router
|
||||
from app.api.routes.projects import router as projects_router
|
||||
from app.api.routes.task_center import router as task_center_router
|
||||
from app.api.routes.upload import router as upload_router
|
||||
@@ -29,13 +28,6 @@ api_router.include_router(
|
||||
prefix="/projects",
|
||||
tags=["Project"],
|
||||
)
|
||||
api_router.include_router(
|
||||
project_titles_router,
|
||||
tags=["TitleLibrary"],
|
||||
)
|
||||
api_router.include_router(
|
||||
edit_plans_router,
|
||||
)
|
||||
api_router.include_router(
|
||||
task_center_router,
|
||||
tags=["TaskCenter"],
|
||||
@@ -85,7 +77,12 @@ api_router.include_router(
|
||||
tags=["GeneratedVideo"],
|
||||
)
|
||||
api_router.include_router(
|
||||
project_management_router,
|
||||
prefix="/project-management",
|
||||
tags=["ProjectManagement"],
|
||||
titles_router,
|
||||
prefix="/titles",
|
||||
tags=["TitleLibrary"],
|
||||
)
|
||||
api_router.include_router(
|
||||
voices_router,
|
||||
prefix="/voices",
|
||||
tags=["VoiceLibrary"],
|
||||
)
|
||||
|
||||
@@ -7,7 +7,7 @@ from app.schemas.asset_library import (
|
||||
CreateAssetLibraryRequest,
|
||||
ListAssetLibrariesResponse,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
|
||||
from packages.application import (
|
||||
CreateAssetLibraryCommand,
|
||||
@@ -42,18 +42,30 @@ def _to_asset_library_response(item) -> AssetLibraryResponse:
|
||||
|
||||
@router.get("", response_model=ListAssetLibrariesResponse)
|
||||
def list_asset_libraries(
|
||||
project_id: str,
|
||||
project_id: str | None = Query(None),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
asset_library_repository: Any = Depends(get_asset_library_repository),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
) -> ListAssetLibrariesResponse:
|
||||
project = GetProjectUseCase(project_repository).execute(project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
|
||||
if not project.can_access(authenticated_user.id):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Access denied to project")
|
||||
user_id = authenticated_user.user.id
|
||||
use_case = ListAssetLibrariesUseCase(asset_library_repository)
|
||||
items = use_case.execute(project_id)
|
||||
|
||||
if project_id:
|
||||
# If project_id provided, check access and filter by project
|
||||
project = GetProjectUseCase(project_repository).execute(project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
|
||||
if not project.can_access(user_id):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Access denied to project")
|
||||
items = use_case.execute(project_id)
|
||||
else:
|
||||
# If no project_id, list all libraries from accessible projects
|
||||
accessible_projects = project_repository.find_accessible_projects(user_id)
|
||||
all_items = []
|
||||
for proj in accessible_projects:
|
||||
all_items.extend(use_case.execute(proj.id))
|
||||
items = all_items
|
||||
|
||||
return ListAssetLibrariesResponse(items=[_to_asset_library_response(item) for item in items])
|
||||
|
||||
|
||||
@@ -67,7 +79,7 @@ def create_asset_library(
|
||||
project = project_repository.find_by_id(request.project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
|
||||
if not project.can_access(authenticated_user.id):
|
||||
if not project.can_access(authenticated_user.user.id):
|
||||
raise HTTPException(status_code=status.HTTP_403_FORBIDDEN, detail="Access denied to project")
|
||||
use_case = CreateAssetLibraryUseCase(asset_library_repository)
|
||||
item = use_case.execute(
|
||||
|
||||
@@ -62,7 +62,7 @@ def list_assets(
|
||||
library = asset_library_repository.get(library_id)
|
||||
if library is None:
|
||||
raise HTTPException(status_code=404, detail=f"AssetLibrary {library_id} not found")
|
||||
_check_project_access(library.project_id, authenticated_user.id, project_repository)
|
||||
_check_project_access(library.project_id, authenticated_user.user.id, project_repository)
|
||||
use_case = ListAssetsUseCase(asset_repository)
|
||||
items = use_case.execute(library_id)
|
||||
return ListAssetsResponse(items=[_to_asset_response(item) for item in items])
|
||||
@@ -87,7 +87,7 @@ def update_asset_review_status(
|
||||
item = asset_repository.get(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.id, project_repository)
|
||||
_check_project_access(item.project_id, authenticated_user.user.id, project_repository)
|
||||
_apply_asset_review_status(item, request.review_status)
|
||||
updated = asset_repository.update(item)
|
||||
return _to_asset_response(updated)
|
||||
@@ -104,7 +104,7 @@ def create_asset(
|
||||
project = project_repository.find_by_id(request.project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=404, detail=f"Project {request.project_id} not found")
|
||||
if not project.can_access(authenticated_user.id):
|
||||
if not project.can_access(authenticated_user.user.id):
|
||||
raise HTTPException(status_code=403, detail="Access denied to project")
|
||||
|
||||
library = asset_library_repository.get(request.library_id)
|
||||
@@ -130,7 +130,7 @@ def create_asset(
|
||||
status=AssetStatus(request.status),
|
||||
classification_status=ClassificationStatus(request.classification_status),
|
||||
quality_score=request.quality_score,
|
||||
uploaded_by_user_id=authenticated_user.id,
|
||||
uploaded_by_user_id=authenticated_user.user.id,
|
||||
)
|
||||
)
|
||||
return _to_asset_response(item)
|
||||
|
||||
@@ -1,287 +0,0 @@
|
||||
from datetime import datetime, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import (
|
||||
get_asset_repository,
|
||||
get_db_session,
|
||||
get_project_repository,
|
||||
)
|
||||
from app.schemas.edit_plan import (
|
||||
AutoGenerateEditPlanRequest,
|
||||
CreateEditPlanRequest,
|
||||
EditPlanClipResponse,
|
||||
EditPlanResponse,
|
||||
EditTemplateResponse,
|
||||
)
|
||||
from packages.domain.edit_plan import EditingMode, SmartEditPlanGenerator
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import EditPlanClipModel, EditPlanModel, EditTemplateModel
|
||||
from packages.domain import AssetStatus
|
||||
|
||||
router = APIRouter(prefix="/projects/{project_id}/edit-plans", tags=["剪辑计划"])
|
||||
|
||||
|
||||
def _ensure_project(project_id: str, project_repository):
|
||||
project = project_repository.find_by_id(project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=404, detail="Project not found")
|
||||
return project
|
||||
|
||||
|
||||
def _default_template(session: Session, project_id: str, user_id: str) -> EditTemplateModel:
|
||||
template = (
|
||||
session.query(EditTemplateModel)
|
||||
.filter(
|
||||
EditTemplateModel.project_id == project_id,
|
||||
EditTemplateModel.is_active.is_(True),
|
||||
)
|
||||
.order_by(EditTemplateModel.created_at.asc())
|
||||
.first()
|
||||
)
|
||||
if template is not None:
|
||||
return template
|
||||
template = EditTemplateModel(
|
||||
id=uuid4().hex,
|
||||
project_id=project_id,
|
||||
name="基础节奏模板",
|
||||
description="自动选择可用视频素材,按上传顺序生成三段式剪辑计划。",
|
||||
target_duration=30,
|
||||
clip_count=3,
|
||||
created_by_user_id=user_id,
|
||||
)
|
||||
session.add(template)
|
||||
session.commit()
|
||||
return template
|
||||
|
||||
|
||||
def _to_template_response(template: EditTemplateModel) -> EditTemplateResponse:
|
||||
return EditTemplateResponse(
|
||||
id=template.id,
|
||||
project_id=template.project_id,
|
||||
name=template.name,
|
||||
description=template.description,
|
||||
target_duration=float(template.target_duration or 0),
|
||||
clip_count=int(template.clip_count or 0),
|
||||
is_active=bool(template.is_active),
|
||||
created_at=template.created_at,
|
||||
)
|
||||
|
||||
|
||||
def _to_plan_response(
|
||||
plan: EditPlanModel, clips: list[EditPlanClipModel], asset_names: dict[str, str]
|
||||
) -> EditPlanResponse:
|
||||
return EditPlanResponse(
|
||||
id=plan.id,
|
||||
project_id=plan.project_id,
|
||||
template_id=plan.template_id,
|
||||
asset_library_id=plan.asset_library_id,
|
||||
title_id=plan.title_id,
|
||||
status=plan.status,
|
||||
summary=plan.summary,
|
||||
editing_mode=plan.editing_mode,
|
||||
clips=[
|
||||
EditPlanClipResponse(
|
||||
id=clip.id,
|
||||
asset_id=clip.asset_id,
|
||||
asset_name=asset_names.get(clip.asset_id, clip.asset_id),
|
||||
sequence=clip.sequence,
|
||||
start_time=float(clip.start_time or 0),
|
||||
duration=float(clip.duration or 0),
|
||||
reason=clip.reason,
|
||||
layer=clip.layer,
|
||||
)
|
||||
for clip in clips
|
||||
],
|
||||
created_at=plan.created_at,
|
||||
updated_at=plan.updated_at,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/templates/", response_model=list[EditTemplateResponse])
|
||||
def list_edit_templates(
|
||||
project_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
project_repository=Depends(get_project_repository),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> list[EditTemplateResponse]:
|
||||
_ensure_project(project_id, project_repository)
|
||||
template = _default_template(session, project_id, authenticated_user.user.id)
|
||||
templates = (
|
||||
session.query(EditTemplateModel)
|
||||
.filter(EditTemplateModel.project_id == project_id, EditTemplateModel.is_active.is_(True))
|
||||
.all()
|
||||
)
|
||||
return [_to_template_response(item) for item in templates or [template]]
|
||||
|
||||
|
||||
@router.post("", response_model=EditPlanResponse)
|
||||
def create_edit_plan(
|
||||
project_id: str,
|
||||
request: CreateEditPlanRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
project_repository=Depends(get_project_repository),
|
||||
asset_repository=Depends(get_asset_repository),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> EditPlanResponse:
|
||||
_ensure_project(project_id, project_repository)
|
||||
template = (
|
||||
session.query(EditTemplateModel).filter(EditTemplateModel.id == request.template_id).first()
|
||||
if request.template_id
|
||||
else None
|
||||
)
|
||||
if template is None:
|
||||
template = _default_template(session, project_id, authenticated_user.user.id)
|
||||
assets = [
|
||||
asset
|
||||
for asset in asset_repository.list_by_library(request.asset_library_id)
|
||||
if asset.status == AssetStatus.READY and asset.mime_type.startswith("video/")
|
||||
]
|
||||
if not assets:
|
||||
raise HTTPException(status_code=422, detail="素材库暂无可用于剪辑计划的视频素材")
|
||||
selected = sorted(assets, key=lambda asset: (-(asset.quality_score or 0), asset.created_at))[
|
||||
: max(1, int(template.clip_count or 3))
|
||||
]
|
||||
plan = EditPlanModel(
|
||||
id=uuid4().hex,
|
||||
project_id=project_id,
|
||||
template_id=template.id,
|
||||
asset_library_id=request.asset_library_id,
|
||||
title_id=request.title_id,
|
||||
status="draft",
|
||||
summary=f"按《{template.name}》自动选择 {len(selected)} 段素材,预计生成约 {int(template.target_duration or 30)} 秒成片。",
|
||||
created_by_user_id=authenticated_user.user.id,
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
)
|
||||
session.add(plan)
|
||||
clips: list[EditPlanClipModel] = []
|
||||
clip_duration = max(1, float(template.target_duration or 30) / len(selected))
|
||||
for index, asset in enumerate(selected, start=1):
|
||||
clip = EditPlanClipModel(
|
||||
id=uuid4().hex,
|
||||
edit_plan_id=plan.id,
|
||||
asset_id=asset.id,
|
||||
sequence=index,
|
||||
start_time=0,
|
||||
duration=min(float(asset.duration or clip_duration), clip_duration),
|
||||
reason="优先选择已就绪、质量分较高的视频素材。",
|
||||
)
|
||||
session.add(clip)
|
||||
clips.append(clip)
|
||||
session.commit()
|
||||
return _to_plan_response(plan, clips, {asset.id: asset.name for asset in selected})
|
||||
|
||||
|
||||
@router.post("/auto-generate", response_model=EditPlanResponse)
|
||||
def auto_generate_edit_plan(
|
||||
project_id: str,
|
||||
request: AutoGenerateEditPlanRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
project_repository=Depends(get_project_repository),
|
||||
asset_repository=Depends(get_asset_repository),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> EditPlanResponse:
|
||||
"""
|
||||
智能生成剪辑计划
|
||||
|
||||
根据素材的分类结果和质量评分,自动编排剪辑计划。
|
||||
支持多种剪辑模式:
|
||||
- one_take: 按分类分组,组内按质量排序,顺序拼接
|
||||
- pip: 第一个高质量素材为主画面,其余为画中画
|
||||
- voice_over: person 类素材为主播口播,其余穿插为 B-roll
|
||||
- voice_pip: 结合 voice_over 和 pip,第一个高质量 person 素材为主画面
|
||||
"""
|
||||
_ensure_project(project_id, project_repository)
|
||||
|
||||
# 获取素材库中的所有素材
|
||||
assets = asset_repository.list_by_library(request.asset_library_id)
|
||||
|
||||
if not assets:
|
||||
raise HTTPException(status_code=422, detail="素材库中暂无素材")
|
||||
|
||||
# 使用智能生成器
|
||||
generator = SmartEditPlanGenerator(project_id, assets)
|
||||
|
||||
try:
|
||||
plan_result = generator.generate_plan(
|
||||
editing_mode=request.editing_mode,
|
||||
target_duration=request.target_duration
|
||||
)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
|
||||
if not plan_result.clips:
|
||||
raise HTTPException(status_code=422, detail="无符合条件的视频素材")
|
||||
|
||||
# 获取模板
|
||||
template = (
|
||||
session.query(EditTemplateModel).filter(EditTemplateModel.id == request.template_id).first()
|
||||
if request.template_id
|
||||
else None
|
||||
)
|
||||
if template is None:
|
||||
template = _default_template(session, project_id, authenticated_user.user.id)
|
||||
|
||||
# 创建剪辑计划
|
||||
plan = EditPlanModel(
|
||||
id=uuid4().hex,
|
||||
project_id=project_id,
|
||||
template_id=template.id,
|
||||
asset_library_id=request.asset_library_id,
|
||||
title_id=request.title_id or "",
|
||||
status="draft",
|
||||
editing_mode=request.editing_mode,
|
||||
summary=plan_result.summary,
|
||||
created_by_user_id=authenticated_user.user.id,
|
||||
updated_at=datetime.now(timezone.utc),
|
||||
)
|
||||
session.add(plan)
|
||||
|
||||
# 创建剪辑片段
|
||||
clips: list[EditPlanClipModel] = []
|
||||
asset_name_map = {asset.id: asset.name for asset in assets}
|
||||
|
||||
for clip_plan in plan_result.clips:
|
||||
clip = EditPlanClipModel(
|
||||
id=uuid4().hex,
|
||||
edit_plan_id=plan.id,
|
||||
asset_id=clip_plan.asset_id,
|
||||
sequence=clip_plan.sequence,
|
||||
start_time=clip_plan.start_time,
|
||||
duration=clip_plan.duration,
|
||||
reason=clip_plan.reason,
|
||||
layer=clip_plan.layer,
|
||||
)
|
||||
session.add(clip)
|
||||
clips.append(clip)
|
||||
|
||||
session.commit()
|
||||
|
||||
return _to_plan_response(plan, clips, asset_name_map)
|
||||
|
||||
|
||||
@router.get("/{plan_id}", response_model=EditPlanResponse)
|
||||
def get_edit_plan(
|
||||
project_id: str,
|
||||
plan_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
project_repository=Depends(get_project_repository),
|
||||
asset_repository=Depends(get_asset_repository),
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> EditPlanResponse:
|
||||
plan = (
|
||||
session.query(EditPlanModel).filter(EditPlanModel.id == plan_id, EditPlanModel.project_id == project_id).first()
|
||||
)
|
||||
if plan is None:
|
||||
raise HTTPException(status_code=404, detail="Edit plan not found")
|
||||
_ensure_project(project_id, project_repository)
|
||||
clips = (
|
||||
session.query(EditPlanClipModel)
|
||||
.filter(EditPlanClipModel.edit_plan_id == plan.id)
|
||||
.order_by(EditPlanClipModel.sequence.asc())
|
||||
.all()
|
||||
)
|
||||
assets = asset_repository.list_by_library(plan.asset_library_id)
|
||||
return _to_plan_response(plan, clips, {asset.id: asset.name for asset in assets})
|
||||
@@ -9,7 +9,7 @@ from app.schemas.generated_video import (
|
||||
ListGeneratedVideosResponse,
|
||||
UpdateGeneratedVideoReviewRequest,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||||
|
||||
from packages.application import (
|
||||
GetGeneratedVideoDownloadUrlUseCase,
|
||||
@@ -42,17 +42,29 @@ def _to_generated_video_response(item, download_url: str | None = None) -> Gener
|
||||
|
||||
@router.get("", response_model=ListGeneratedVideosResponse)
|
||||
def list_generated_videos(
|
||||
project_id: str,
|
||||
project_id: str | None = Query(None),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
generated_video_repository: Any = Depends(get_generated_video_repository),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
storage_service: OSSStorageService = Depends(get_storage_service),
|
||||
) -> ListGeneratedVideosResponse:
|
||||
project = project_repository.find_by_id(project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=404, detail=f"Project {project_id} not found")
|
||||
user_id = authenticated_user.user.id
|
||||
use_case = ListGeneratedVideosUseCase(generated_video_repository)
|
||||
items = use_case.execute(project_id)
|
||||
|
||||
if project_id:
|
||||
# If project_id provided, check access and filter by project
|
||||
project = project_repository.find_by_id(project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=404, detail=f"Project {project_id} not found")
|
||||
items = use_case.execute(project_id)
|
||||
else:
|
||||
# If no project_id, list all videos from accessible projects
|
||||
accessible_projects = project_repository.find_accessible_projects(user_id)
|
||||
all_items = []
|
||||
for proj in accessible_projects:
|
||||
all_items.extend(use_case.execute(proj.id))
|
||||
items = all_items
|
||||
|
||||
# Generate download URLs for each video
|
||||
responses = []
|
||||
for item in items:
|
||||
|
||||
@@ -8,7 +8,6 @@ from app.dependencies import (
|
||||
get_generated_video_repository,
|
||||
get_generation_task_repository,
|
||||
get_project_repository,
|
||||
get_project_title_repository,
|
||||
)
|
||||
from app.schemas.generated_video import (
|
||||
GeneratedVideoResponse,
|
||||
@@ -46,7 +45,6 @@ def _to_generation_task_response(task) -> GenerationTaskResponse:
|
||||
asset_library_id=task.asset_library_id,
|
||||
strategy_id=task.strategy_id,
|
||||
voice_library_id=task.voice_library_id,
|
||||
edit_plan_id=task.edit_plan_id,
|
||||
status=task.status,
|
||||
progress=task.progress,
|
||||
result_count=task.result_count,
|
||||
@@ -81,21 +79,6 @@ def _ensure_library_has_ready_video_assets(assets) -> None:
|
||||
)
|
||||
|
||||
|
||||
def _select_title_id(project_title_repository: Any, project_id: str) -> str:
|
||||
active_titles = project_title_repository.list_by_project(project_id, active_only=True)
|
||||
if not active_titles:
|
||||
return ""
|
||||
selected = sorted(
|
||||
active_titles,
|
||||
key=lambda title: (
|
||||
0 if getattr(title, "favorite", False) else 1,
|
||||
int(title.usage_count or 0),
|
||||
title.created_at,
|
||||
),
|
||||
)[0]
|
||||
return selected.id
|
||||
|
||||
|
||||
@router.post("/tasks/", response_model=GenerationTaskResponse)
|
||||
def create_generation_task(
|
||||
request: CreateGenerationTaskRequest,
|
||||
@@ -104,12 +87,11 @@ def create_generation_task(
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
asset_library_repository: Any = Depends(get_asset_library_repository),
|
||||
asset_repository: Any = Depends(get_asset_repository),
|
||||
project_title_repository: Any = Depends(get_project_title_repository),
|
||||
) -> GenerationTaskResponse:
|
||||
project = project_repository.find_by_id(request.project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=404, detail=f"Project {request.project_id} not found")
|
||||
if not project.can_access(authenticated_user.id):
|
||||
if not project.can_access(authenticated_user.user.id):
|
||||
raise HTTPException(status_code=403, detail="Access denied to project")
|
||||
|
||||
library = asset_library_repository.get(request.asset_library_id)
|
||||
@@ -124,10 +106,9 @@ def create_generation_task(
|
||||
CreateGenerationTaskCommand(
|
||||
project_id=request.project_id,
|
||||
asset_library_id=request.asset_library_id,
|
||||
strategy_id=request.strategy_id or _select_title_id(project_title_repository, request.project_id),
|
||||
strategy_id=request.strategy_id,
|
||||
voice_library_id=request.voice_library_id,
|
||||
edit_plan_id=request.edit_plan_id,
|
||||
created_by_user_id=authenticated_user.id,
|
||||
created_by_user_id=authenticated_user.user.id,
|
||||
)
|
||||
)
|
||||
celery_app.send_task("worker.generate_video", args=[task.id])
|
||||
@@ -145,7 +126,7 @@ def get_generation_task(
|
||||
task = use_case.execute(task_id)
|
||||
if task is None:
|
||||
raise HTTPException(status_code=404, detail=f"GenerationTask {task_id} not found")
|
||||
_check_project_access(task.project_id, authenticated_user.id, project_repository)
|
||||
_check_project_access(task.project_id, authenticated_user.user.id, project_repository)
|
||||
return _to_generation_task_response(task)
|
||||
|
||||
|
||||
@@ -160,7 +141,7 @@ def list_generation_results(
|
||||
task = generation_task_repository.get(task_id)
|
||||
if task is None:
|
||||
raise HTTPException(status_code=404, detail=f"GenerationTask {task_id} not found")
|
||||
_check_project_access(task.project_id, authenticated_user.id, project_repository)
|
||||
_check_project_access(task.project_id, authenticated_user.user.id, project_repository)
|
||||
use_case = ListGeneratedVideosByTaskUseCase(generated_video_repository)
|
||||
items = use_case.execute(task_id)
|
||||
return ListGeneratedVideosResponse(items=[_to_generated_video_response(item) for item in items])
|
||||
|
||||
@@ -1,466 +0,0 @@
|
||||
"""项目管理 API 路由"""
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Annotated
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from packages.adapters.sqlite_tracker.project_management_repositories import (
|
||||
SQLiteMilestoneRepository,
|
||||
SQLiteTaskIssueRepository,
|
||||
SQLiteTaskRepository,
|
||||
)
|
||||
from packages.application.get_task_detail_use_case import GetTaskDetailUseCase
|
||||
from packages.application.project_management_use_cases import (
|
||||
CreateMilestoneUseCase,
|
||||
CreateTaskIssueUseCase,
|
||||
CreateTaskUseCase,
|
||||
ListProjectMilestonesUseCase,
|
||||
ListProjectTasksUseCase,
|
||||
ListTaskIssuesUseCase,
|
||||
ResolveTaskIssueUseCase,
|
||||
UpdateTaskProgressUseCase,
|
||||
UpdateTaskStatusUseCase,
|
||||
)
|
||||
from packages.application.update_task_use_case import UpdateTaskUseCase
|
||||
from packages.domain import TaskPriority, TaskStatus
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
# 使用 SQLite tracker.db
|
||||
_task_repo = SQLiteTaskRepository()
|
||||
_milestone_repo = SQLiteMilestoneRepository()
|
||||
_issue_repo = SQLiteTaskIssueRepository()
|
||||
|
||||
|
||||
def get_task_repo():
|
||||
return _task_repo
|
||||
|
||||
|
||||
def get_milestone_repo():
|
||||
return _milestone_repo
|
||||
|
||||
|
||||
def get_issue_repo():
|
||||
return _issue_repo
|
||||
|
||||
|
||||
# ========== Request/Response Models ==========
|
||||
|
||||
|
||||
class CreateTaskRequest(BaseModel):
|
||||
project_id: str
|
||||
name: str
|
||||
description: str = ""
|
||||
priority: TaskPriority = TaskPriority.MEDIUM
|
||||
parent_task_id: str = ""
|
||||
assignee_user_id: str = ""
|
||||
|
||||
|
||||
class TaskResponse(BaseModel):
|
||||
id: str
|
||||
project_id: str
|
||||
name: str
|
||||
description: str
|
||||
status: TaskStatus
|
||||
priority: TaskPriority
|
||||
parent_task_id: str
|
||||
assignee_user_id: str
|
||||
progress: float
|
||||
planned_start_date: datetime | None
|
||||
planned_end_date: datetime | None
|
||||
actual_start_date: datetime | None
|
||||
actual_end_date: datetime | None
|
||||
tags: list[str]
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class UpdateTaskRequest(BaseModel):
|
||||
name: str | None = None
|
||||
description: str | None = None
|
||||
priority: str | None = None
|
||||
assignee_user_id: str | None = None
|
||||
|
||||
|
||||
class UpdateTaskStatusRequest(BaseModel):
|
||||
status: TaskStatus
|
||||
|
||||
|
||||
class UpdateTaskProgressRequest(BaseModel):
|
||||
progress: Annotated[float, Field(ge=0, le=100)]
|
||||
|
||||
|
||||
class CreateMilestoneRequest(BaseModel):
|
||||
project_id: str
|
||||
name: str
|
||||
description: str = ""
|
||||
|
||||
|
||||
class MilestoneResponse(BaseModel):
|
||||
id: str
|
||||
project_id: str
|
||||
name: str
|
||||
description: str
|
||||
target_date: datetime | None
|
||||
completed: bool
|
||||
completed_at: datetime | None
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class CreateTaskIssueRequest(BaseModel):
|
||||
task_id: str
|
||||
project_id: str
|
||||
title: str
|
||||
description: str = ""
|
||||
created_by_user_id: str = ""
|
||||
|
||||
|
||||
class TaskIssueResponse(BaseModel):
|
||||
id: str
|
||||
task_id: str
|
||||
project_id: str
|
||||
title: str
|
||||
description: str
|
||||
resolved: bool
|
||||
resolved_at: datetime | None
|
||||
created_by_user_id: str
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
# ========== Task Endpoints ==========
|
||||
|
||||
|
||||
@router.post("/tasks", response_model=TaskResponse)
|
||||
def create_task(
|
||||
req: CreateTaskRequest,
|
||||
task_repo=Depends(get_task_repo),
|
||||
):
|
||||
"""创建任务"""
|
||||
use_case = CreateTaskUseCase(task_repo)
|
||||
task = use_case.execute(
|
||||
project_id=req.project_id,
|
||||
name=req.name,
|
||||
description=req.description,
|
||||
priority=req.priority,
|
||||
parent_task_id=req.parent_task_id,
|
||||
assignee_user_id=req.assignee_user_id,
|
||||
)
|
||||
return TaskResponse(
|
||||
id=task.id,
|
||||
project_id=task.project_id,
|
||||
name=task.name,
|
||||
description=task.description,
|
||||
status=task.status,
|
||||
priority=task.priority,
|
||||
parent_task_id=task.parent_task_id,
|
||||
assignee_user_id=task.assignee_user_id,
|
||||
progress=task.progress,
|
||||
planned_start_date=task.planned_start_date,
|
||||
planned_end_date=task.planned_end_date,
|
||||
actual_start_date=task.actual_start_date,
|
||||
actual_end_date=task.actual_end_date,
|
||||
tags=task.tags,
|
||||
created_at=task.created_at,
|
||||
updated_at=task.updated_at,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/tasks", response_model=list[TaskResponse])
|
||||
def list_tasks(
|
||||
project_id: str,
|
||||
task_repo=Depends(get_task_repo),
|
||||
):
|
||||
"""获取项目任务列表"""
|
||||
use_case = ListProjectTasksUseCase(task_repo)
|
||||
tasks = use_case.execute(project_id)
|
||||
return [
|
||||
TaskResponse(
|
||||
id=t.id,
|
||||
project_id=t.project_id,
|
||||
name=t.name,
|
||||
description=t.description,
|
||||
status=t.status,
|
||||
priority=t.priority,
|
||||
parent_task_id=t.parent_task_id,
|
||||
assignee_user_id=t.assignee_user_id,
|
||||
progress=t.progress,
|
||||
planned_start_date=t.planned_start_date,
|
||||
planned_end_date=t.planned_end_date,
|
||||
actual_start_date=t.actual_start_date,
|
||||
actual_end_date=t.actual_end_date,
|
||||
tags=t.tags,
|
||||
created_at=t.created_at,
|
||||
updated_at=t.updated_at,
|
||||
)
|
||||
for t in tasks
|
||||
]
|
||||
|
||||
|
||||
@router.get("/tasks/{task_id}", response_model=TaskResponse)
|
||||
def get_task(
|
||||
task_id: str,
|
||||
task_repo=Depends(get_task_repo),
|
||||
):
|
||||
"""获取任务详情"""
|
||||
use_case = GetTaskDetailUseCase(task_repo)
|
||||
try:
|
||||
task = use_case.execute(task_id)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=404, detail=str(e))
|
||||
return TaskResponse(
|
||||
id=task.id,
|
||||
project_id=task.project_id,
|
||||
name=task.name,
|
||||
description=task.description,
|
||||
status=task.status,
|
||||
priority=task.priority,
|
||||
parent_task_id=task.parent_task_id,
|
||||
assignee_user_id=task.assignee_user_id,
|
||||
progress=task.progress,
|
||||
planned_start_date=task.planned_start_date,
|
||||
planned_end_date=task.planned_end_date,
|
||||
actual_start_date=task.actual_start_date,
|
||||
actual_end_date=task.actual_end_date,
|
||||
tags=task.tags,
|
||||
created_at=task.created_at,
|
||||
updated_at=task.updated_at,
|
||||
)
|
||||
|
||||
|
||||
@router.patch("/tasks/{task_id}", response_model=TaskResponse)
|
||||
def update_task(
|
||||
task_id: str,
|
||||
req: UpdateTaskRequest,
|
||||
task_repo=Depends(get_task_repo),
|
||||
):
|
||||
"""更新任务基本信息"""
|
||||
use_case = UpdateTaskUseCase(task_repo)
|
||||
try:
|
||||
task = use_case.execute(
|
||||
task_id=task_id,
|
||||
name=req.name,
|
||||
description=req.description,
|
||||
priority=req.priority,
|
||||
assignee_user_id=req.assignee_user_id,
|
||||
)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=404, detail=str(e))
|
||||
return TaskResponse(
|
||||
id=task.id,
|
||||
project_id=task.project_id,
|
||||
name=task.name,
|
||||
description=task.description,
|
||||
status=task.status,
|
||||
priority=task.priority,
|
||||
parent_task_id=task.parent_task_id,
|
||||
assignee_user_id=task.assignee_user_id,
|
||||
progress=task.progress,
|
||||
planned_start_date=task.planned_start_date,
|
||||
planned_end_date=task.planned_end_date,
|
||||
actual_start_date=task.actual_start_date,
|
||||
actual_end_date=task.actual_end_date,
|
||||
tags=task.tags,
|
||||
created_at=task.created_at,
|
||||
updated_at=task.updated_at,
|
||||
)
|
||||
|
||||
|
||||
@router.patch("/tasks/{task_id}/status", response_model=TaskResponse)
|
||||
def update_task_status(
|
||||
task_id: str,
|
||||
req: UpdateTaskStatusRequest,
|
||||
task_repo=Depends(get_task_repo),
|
||||
):
|
||||
"""更新任务状态"""
|
||||
use_case = UpdateTaskStatusUseCase(task_repo)
|
||||
try:
|
||||
task = use_case.execute(task_id, req.status)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=404, detail=str(e))
|
||||
return TaskResponse(
|
||||
id=task.id,
|
||||
project_id=task.project_id,
|
||||
name=task.name,
|
||||
description=task.description,
|
||||
status=task.status,
|
||||
priority=task.priority,
|
||||
parent_task_id=task.parent_task_id,
|
||||
assignee_user_id=task.assignee_user_id,
|
||||
progress=task.progress,
|
||||
planned_start_date=task.planned_start_date,
|
||||
planned_end_date=task.planned_end_date,
|
||||
actual_start_date=task.actual_start_date,
|
||||
actual_end_date=task.actual_end_date,
|
||||
tags=task.tags,
|
||||
created_at=task.created_at,
|
||||
updated_at=task.updated_at,
|
||||
)
|
||||
|
||||
|
||||
@router.patch("/tasks/{task_id}/progress", response_model=TaskResponse)
|
||||
def update_task_progress(
|
||||
task_id: str,
|
||||
req: UpdateTaskProgressRequest,
|
||||
task_repo=Depends(get_task_repo),
|
||||
):
|
||||
"""更新任务进度"""
|
||||
use_case = UpdateTaskProgressUseCase(task_repo)
|
||||
try:
|
||||
task = use_case.execute(task_id, req.progress)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=404, detail=str(e))
|
||||
return TaskResponse(
|
||||
id=task.id,
|
||||
project_id=task.project_id,
|
||||
name=task.name,
|
||||
description=task.description,
|
||||
status=task.status,
|
||||
priority=task.priority,
|
||||
parent_task_id=task.parent_task_id,
|
||||
assignee_user_id=task.assignee_user_id,
|
||||
progress=task.progress,
|
||||
planned_start_date=task.planned_start_date,
|
||||
planned_end_date=task.planned_end_date,
|
||||
actual_start_date=task.actual_start_date,
|
||||
actual_end_date=task.actual_end_date,
|
||||
tags=task.tags,
|
||||
created_at=task.created_at,
|
||||
updated_at=task.updated_at,
|
||||
)
|
||||
|
||||
|
||||
# ========== Milestone Endpoints ==========
|
||||
|
||||
|
||||
@router.post("/milestones", response_model=MilestoneResponse)
|
||||
def create_milestone(
|
||||
req: CreateMilestoneRequest,
|
||||
milestone_repo=Depends(get_milestone_repo),
|
||||
):
|
||||
"""创建里程碑"""
|
||||
use_case = CreateMilestoneUseCase(milestone_repo)
|
||||
milestone = use_case.execute(
|
||||
project_id=req.project_id,
|
||||
name=req.name,
|
||||
description=req.description,
|
||||
)
|
||||
return MilestoneResponse(
|
||||
id=milestone.id,
|
||||
project_id=milestone.project_id,
|
||||
name=milestone.name,
|
||||
description=milestone.description,
|
||||
target_date=milestone.target_date,
|
||||
completed=milestone.completed,
|
||||
completed_at=milestone.completed_at,
|
||||
created_at=milestone.created_at,
|
||||
updated_at=milestone.updated_at,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/milestones", response_model=list[MilestoneResponse])
|
||||
def list_milestones(
|
||||
project_id: str,
|
||||
milestone_repo=Depends(get_milestone_repo),
|
||||
):
|
||||
"""获取项目里程碑列表"""
|
||||
use_case = ListProjectMilestonesUseCase(milestone_repo)
|
||||
milestones = use_case.execute(project_id)
|
||||
return [
|
||||
MilestoneResponse(
|
||||
id=m.id,
|
||||
project_id=m.project_id,
|
||||
name=m.name,
|
||||
description=m.description,
|
||||
target_date=m.target_date,
|
||||
completed=m.completed,
|
||||
completed_at=m.completed_at,
|
||||
created_at=m.created_at,
|
||||
updated_at=m.updated_at,
|
||||
)
|
||||
for m in milestones
|
||||
]
|
||||
|
||||
|
||||
# ========== Task Issue Endpoints ==========
|
||||
|
||||
|
||||
@router.post("/issues", response_model=TaskIssueResponse)
|
||||
def create_issue(
|
||||
req: CreateTaskIssueRequest,
|
||||
issue_repo=Depends(get_issue_repo),
|
||||
):
|
||||
"""创建任务问题"""
|
||||
use_case = CreateTaskIssueUseCase(issue_repo)
|
||||
issue = use_case.execute(
|
||||
task_id=req.task_id,
|
||||
project_id=req.project_id,
|
||||
title=req.title,
|
||||
description=req.description,
|
||||
created_by_user_id=req.created_by_user_id,
|
||||
)
|
||||
return TaskIssueResponse(
|
||||
id=issue.id,
|
||||
task_id=issue.task_id,
|
||||
project_id=issue.project_id,
|
||||
title=issue.title,
|
||||
description=issue.description,
|
||||
resolved=issue.resolved,
|
||||
resolved_at=issue.resolved_at,
|
||||
created_by_user_id=issue.created_by_user_id,
|
||||
created_at=issue.created_at,
|
||||
updated_at=issue.updated_at,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/issues", response_model=list[TaskIssueResponse])
|
||||
def list_issues(
|
||||
task_id: str,
|
||||
issue_repo=Depends(get_issue_repo),
|
||||
):
|
||||
"""获取任务问题列表"""
|
||||
use_case = ListTaskIssuesUseCase(issue_repo)
|
||||
issues = use_case.execute(task_id)
|
||||
return [
|
||||
TaskIssueResponse(
|
||||
id=i.id,
|
||||
task_id=i.task_id,
|
||||
project_id=i.project_id,
|
||||
title=i.title,
|
||||
description=i.description,
|
||||
resolved=i.resolved,
|
||||
resolved_at=i.resolved_at,
|
||||
created_by_user_id=i.created_by_user_id,
|
||||
created_at=i.created_at,
|
||||
updated_at=i.updated_at,
|
||||
)
|
||||
for i in issues
|
||||
]
|
||||
|
||||
|
||||
@router.patch("/issues/{issue_id}/resolve", response_model=TaskIssueResponse)
|
||||
def resolve_issue(
|
||||
issue_id: str,
|
||||
issue_repo=Depends(get_issue_repo),
|
||||
):
|
||||
"""解决任务问题"""
|
||||
use_case = ResolveTaskIssueUseCase(issue_repo)
|
||||
try:
|
||||
issue = use_case.execute(issue_id)
|
||||
except ValueError as e:
|
||||
raise HTTPException(status_code=404, detail=str(e))
|
||||
return TaskIssueResponse(
|
||||
id=issue.id,
|
||||
task_id=issue.task_id,
|
||||
project_id=issue.project_id,
|
||||
title=issue.title,
|
||||
description=issue.description,
|
||||
resolved=issue.resolved,
|
||||
resolved_at=issue.resolved_at,
|
||||
created_by_user_id=issue.created_by_user_id,
|
||||
created_at=issue.created_at,
|
||||
updated_at=issue.updated_at,
|
||||
)
|
||||
@@ -1,91 +0,0 @@
|
||||
from typing import Any
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import (
|
||||
get_project_repository,
|
||||
get_project_title_repository,
|
||||
)
|
||||
from app.schemas.project_title import (
|
||||
CreateProjectTitleRequest,
|
||||
ListProjectTitlesResponse,
|
||||
ProjectTitleResponse,
|
||||
UpdateProjectTitleRequest,
|
||||
)
|
||||
from fastapi import APIRouter, Depends, HTTPException, status
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _to_response(item) -> ProjectTitleResponse:
|
||||
return ProjectTitleResponse(
|
||||
id=item.id,
|
||||
project_id=item.project_id,
|
||||
text=item.text,
|
||||
category=item.category,
|
||||
favorite=bool(getattr(item, "favorite", False)),
|
||||
usage_count=int(item.usage_count or 0),
|
||||
is_active=bool(item.is_active),
|
||||
created_at=item.created_at,
|
||||
updated_at=item.updated_at,
|
||||
)
|
||||
|
||||
|
||||
def _get_project_or_404(project_id: str, project_repository: Any):
|
||||
project = project_repository.find_by_id(project_id)
|
||||
if project is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project not found")
|
||||
return project
|
||||
|
||||
|
||||
@router.get("/projects/{project_id}/titles", response_model=ListProjectTitlesResponse)
|
||||
def list_project_titles(
|
||||
project_id: str,
|
||||
active_only: bool = False,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
title_repository: Any = Depends(get_project_title_repository),
|
||||
) -> ListProjectTitlesResponse:
|
||||
_get_project_or_404(project_id, project_repository)
|
||||
return ListProjectTitlesResponse(
|
||||
items=[_to_response(item) for item in title_repository.list_by_project(project_id, active_only)]
|
||||
)
|
||||
|
||||
|
||||
@router.post("/projects/{project_id}/titles", response_model=ProjectTitleResponse)
|
||||
def create_project_title(
|
||||
project_id: str,
|
||||
request: CreateProjectTitleRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
project_repository: Any = Depends(get_project_repository),
|
||||
title_repository: Any = Depends(get_project_title_repository),
|
||||
) -> ProjectTitleResponse:
|
||||
_get_project_or_404(project_id, project_repository)
|
||||
item = title_repository.create(
|
||||
project_id=project_id,
|
||||
text=request.text,
|
||||
category=request.category,
|
||||
favorite=request.favorite,
|
||||
created_by_user_id=authenticated_user.user.id,
|
||||
)
|
||||
return _to_response(item)
|
||||
|
||||
|
||||
@router.patch("/project-titles/{title_id}", response_model=ProjectTitleResponse)
|
||||
def update_project_title(
|
||||
title_id: str,
|
||||
request: UpdateProjectTitleRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
title_repository: Any = Depends(get_project_title_repository),
|
||||
) -> ProjectTitleResponse:
|
||||
item = title_repository.get(title_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Project title not found")
|
||||
if request.text is not None:
|
||||
item.text = request.text.strip()
|
||||
if request.category is not None:
|
||||
item.category = request.category
|
||||
if request.favorite is not None:
|
||||
item.favorite = request.favorite
|
||||
if request.is_active is not None:
|
||||
item.is_active = request.is_active
|
||||
return _to_response(title_repository.update(item))
|
||||
@@ -0,0 +1,155 @@
|
||||
"""Title library CRUD routes."""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import get_db_session, get_user_repository
|
||||
from app.schemas.title_library import (
|
||||
CreateTitleLibraryRequest,
|
||||
ListTitleLibraryResponse,
|
||||
TitleLibraryItemResponse,
|
||||
UpdateTitleLibraryRequest,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.title_library_repository import SQLAlchemyTitleLibraryRepository
|
||||
from packages.application.title_library.commands import CreateTitleLibraryCommand, UpdateTitleLibraryCommand
|
||||
from packages.application.title_library.use_cases import (
|
||||
CreateTitleLibraryUseCase,
|
||||
DeleteTitleLibraryUseCase,
|
||||
GetTitleLibraryUseCase,
|
||||
ListTitleLibraryUseCase,
|
||||
UpdateTitleLibraryUseCase,
|
||||
NotFoundError,
|
||||
QuotaExceededError,
|
||||
)
|
||||
from packages.ports.user_repository import UserRepository
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _get_title_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyTitleLibraryRepository:
|
||||
return SQLAlchemyTitleLibraryRepository(session)
|
||||
|
||||
|
||||
def _to_response(item) -> TitleLibraryItemResponse:
|
||||
return TitleLibraryItemResponse(
|
||||
id=item.id,
|
||||
user_id=item.user_id,
|
||||
name=item.name,
|
||||
text=item.text,
|
||||
category=item.category,
|
||||
description=item.description,
|
||||
tags=item.tags,
|
||||
usage_count=item.usage_count,
|
||||
is_active=item.is_active,
|
||||
created_at=item.created_at,
|
||||
updated_at=item.updated_at,
|
||||
)
|
||||
|
||||
|
||||
def _get_user_plan(user_id: str, user_repository: UserRepository) -> str:
|
||||
user = user_repository.get_by_id(user_id)
|
||||
if user is None:
|
||||
return "free"
|
||||
return getattr(user, "subscription_plan", "free") or "free"
|
||||
|
||||
|
||||
@router.get("/", response_model=ListTitleLibraryResponse)
|
||||
def list_titles(
|
||||
category: Optional[str] = Query(None),
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(50, ge=1, le=200),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
|
||||
) -> ListTitleLibraryResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
use_case = ListTitleLibraryUseCase(title_repository)
|
||||
items = use_case.execute(user_id, category=category, skip=skip, limit=limit)
|
||||
total = title_repository.count_by_user(user_id)
|
||||
return ListTitleLibraryResponse(
|
||||
items=[_to_response(i) for i in items],
|
||||
total=total,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/{title_id}", response_model=TitleLibraryItemResponse)
|
||||
def get_title(
|
||||
title_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
|
||||
) -> TitleLibraryItemResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
use_case = GetTitleLibraryUseCase(title_repository)
|
||||
item = use_case.execute(title_id, user_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Title not found")
|
||||
return _to_response(item)
|
||||
|
||||
|
||||
@router.post("/", response_model=TitleLibraryItemResponse, status_code=status.HTTP_201_CREATED)
|
||||
def create_title(
|
||||
request: CreateTitleLibraryRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
|
||||
user_repository: UserRepository = Depends(get_user_repository),
|
||||
) -> TitleLibraryItemResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
plan_name = _get_user_plan(user_id, user_repository)
|
||||
command = CreateTitleLibraryCommand(
|
||||
user_id=user_id,
|
||||
name=request.name,
|
||||
text=request.text,
|
||||
category=request.category,
|
||||
description=request.description,
|
||||
tags=request.tags,
|
||||
)
|
||||
use_case = CreateTitleLibraryUseCase(title_repository)
|
||||
try:
|
||||
item = use_case.execute(command, plan_name=plan_name)
|
||||
except QuotaExceededError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
|
||||
detail=f"标题库配额已满({exc.used}/{exc.limit}),请升级套餐",
|
||||
)
|
||||
return _to_response(item)
|
||||
|
||||
|
||||
@router.put("/{title_id}", response_model=TitleLibraryItemResponse)
|
||||
def update_title(
|
||||
title_id: str,
|
||||
request: UpdateTitleLibraryRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
|
||||
) -> TitleLibraryItemResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
command = UpdateTitleLibraryCommand(
|
||||
title_id=title_id,
|
||||
user_id=user_id,
|
||||
name=request.name,
|
||||
text=request.text,
|
||||
category=request.category,
|
||||
description=request.description,
|
||||
tags=request.tags,
|
||||
)
|
||||
use_case = UpdateTitleLibraryUseCase(title_repository)
|
||||
try:
|
||||
item = use_case.execute(command)
|
||||
except NotFoundError:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Title not found")
|
||||
return _to_response(item)
|
||||
|
||||
|
||||
@router.delete("/{title_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
def delete_title(
|
||||
title_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
title_repository: SQLAlchemyTitleLibraryRepository = Depends(_get_title_repository),
|
||||
) -> None:
|
||||
user_id = authenticated_user.user.id
|
||||
use_case = DeleteTitleLibraryUseCase(title_repository)
|
||||
deleted = use_case.execute(title_id, user_id)
|
||||
if not deleted:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Title not found")
|
||||
@@ -0,0 +1,170 @@
|
||||
"""Voice library CRUD routes."""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Query, status
|
||||
|
||||
from app.auth import AuthenticatedUser, get_current_user
|
||||
from app.dependencies import get_db_session, get_user_repository
|
||||
from app.schemas.voice_library import (
|
||||
CreateVoiceLibraryRequest,
|
||||
ListVoiceLibraryResponse,
|
||||
VoiceLibraryItemResponse,
|
||||
UpdateVoiceLibraryRequest,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.voice_library_repository import SQLAlchemyVoiceLibraryRepository
|
||||
from packages.application.voice_library.commands import CreateVoiceLibraryCommand, UpdateVoiceLibraryCommand
|
||||
from packages.application.voice_library.use_cases import (
|
||||
CreateVoiceLibraryUseCase,
|
||||
DeleteVoiceLibraryUseCase,
|
||||
GetVoiceLibraryUseCase,
|
||||
ListVoiceLibraryUseCase,
|
||||
UpdateVoiceLibraryUseCase,
|
||||
NotFoundError,
|
||||
QuotaExceededError,
|
||||
)
|
||||
from packages.ports.user_repository import UserRepository
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
def _get_voice_repository(session: Session = Depends(get_db_session)) -> SQLAlchemyVoiceLibraryRepository:
|
||||
return SQLAlchemyVoiceLibraryRepository(session)
|
||||
|
||||
|
||||
def _to_response(item) -> VoiceLibraryItemResponse:
|
||||
return VoiceLibraryItemResponse(
|
||||
id=item.id,
|
||||
user_id=item.user_id,
|
||||
name=item.name,
|
||||
text=item.text,
|
||||
voice_provider=item.voice_provider,
|
||||
voice_id=item.voice_id,
|
||||
voice_name=item.voice_name,
|
||||
audio_url=item.audio_url,
|
||||
duration=item.duration,
|
||||
file_size=item.file_size,
|
||||
status=item.status,
|
||||
project_id=item.project_id,
|
||||
tags=item.tags,
|
||||
created_at=item.created_at,
|
||||
updated_at=item.updated_at,
|
||||
)
|
||||
|
||||
|
||||
def _get_user_plan(user_id: str, user_repository: UserRepository) -> str:
|
||||
user = user_repository.get_by_id(user_id)
|
||||
if user is None:
|
||||
return "free"
|
||||
return getattr(user, "subscription_plan", "free") or "free"
|
||||
|
||||
|
||||
@router.get("/", response_model=ListVoiceLibraryResponse)
|
||||
def list_voices(
|
||||
status_filter: Optional[str] = Query(None, alias="status"),
|
||||
skip: int = Query(0, ge=0),
|
||||
limit: int = Query(50, ge=1, le=200),
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
voice_repository: SQLAlchemyVoiceLibraryRepository = Depends(_get_voice_repository),
|
||||
) -> ListVoiceLibraryResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
use_case = ListVoiceLibraryUseCase(voice_repository)
|
||||
items = use_case.execute(user_id, status=status_filter, skip=skip, limit=limit)
|
||||
total = voice_repository.count_by_user(user_id)
|
||||
return ListVoiceLibraryResponse(
|
||||
items=[_to_response(i) for i in items],
|
||||
total=total,
|
||||
)
|
||||
|
||||
|
||||
@router.get("/{voice_id}", response_model=VoiceLibraryItemResponse)
|
||||
def get_voice(
|
||||
voice_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
voice_repository: SQLAlchemyVoiceLibraryRepository = Depends(_get_voice_repository),
|
||||
) -> VoiceLibraryItemResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
use_case = GetVoiceLibraryUseCase(voice_repository)
|
||||
item = use_case.execute(voice_id, user_id)
|
||||
if item is None:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Voice not found")
|
||||
return _to_response(item)
|
||||
|
||||
|
||||
@router.post("/", response_model=VoiceLibraryItemResponse, status_code=status.HTTP_201_CREATED)
|
||||
def create_voice(
|
||||
request: CreateVoiceLibraryRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
voice_repository: SQLAlchemyVoiceLibraryRepository = Depends(_get_voice_repository),
|
||||
user_repository: UserRepository = Depends(get_user_repository),
|
||||
) -> VoiceLibraryItemResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
plan_name = _get_user_plan(user_id, user_repository)
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id=user_id,
|
||||
name=request.name,
|
||||
text=request.text,
|
||||
voice_provider=request.voice_provider,
|
||||
voice_id=request.voice_id,
|
||||
voice_name=request.voice_name,
|
||||
audio_url=request.audio_url,
|
||||
duration=request.duration,
|
||||
file_size=request.file_size,
|
||||
status=request.status,
|
||||
project_id=request.project_id,
|
||||
tags=request.tags,
|
||||
)
|
||||
use_case = CreateVoiceLibraryUseCase(voice_repository)
|
||||
try:
|
||||
item = use_case.execute(command, plan_name=plan_name)
|
||||
except QuotaExceededError as exc:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_429_TOO_MANY_REQUESTS,
|
||||
detail=f"配音库配额已满({exc.used}/{exc.limit}),请升级套餐",
|
||||
)
|
||||
return _to_response(item)
|
||||
|
||||
|
||||
@router.put("/{voice_id}", response_model=VoiceLibraryItemResponse)
|
||||
def update_voice(
|
||||
voice_id: str,
|
||||
request: UpdateVoiceLibraryRequest,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
voice_repository: SQLAlchemyVoiceLibraryRepository = Depends(_get_voice_repository),
|
||||
) -> VoiceLibraryItemResponse:
|
||||
user_id = authenticated_user.user.id
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id=voice_id,
|
||||
user_id=user_id,
|
||||
name=request.name,
|
||||
text=request.text,
|
||||
voice_provider=request.voice_provider,
|
||||
voice_id=request.voice_id,
|
||||
voice_name=request.voice_name,
|
||||
audio_url=request.audio_url,
|
||||
duration=request.duration,
|
||||
file_size=request.file_size,
|
||||
status=request.status,
|
||||
tags=request.tags,
|
||||
)
|
||||
use_case = UpdateVoiceLibraryUseCase(voice_repository)
|
||||
try:
|
||||
item = use_case.execute(command)
|
||||
except NotFoundError:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Voice not found")
|
||||
return _to_response(item)
|
||||
|
||||
|
||||
@router.delete("/{voice_id}", status_code=status.HTTP_204_NO_CONTENT)
|
||||
def delete_voice(
|
||||
voice_id: str,
|
||||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||||
voice_repository: SQLAlchemyVoiceLibraryRepository = Depends(_get_voice_repository),
|
||||
) -> None:
|
||||
user_id = authenticated_user.user.id
|
||||
use_case = DeleteVoiceLibraryUseCase(voice_repository)
|
||||
deleted = use_case.execute(voice_id, user_id)
|
||||
if not deleted:
|
||||
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Voice not found")
|
||||
@@ -27,15 +27,18 @@ from packages.adapters.sqlalchemy_impl.generated_video_repository import (
|
||||
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
|
||||
SQLAlchemyGenerationTaskRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.title_library_repository import (
|
||||
SQLAlchemyTitleLibraryRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.voice_library_repository import (
|
||||
SQLAlchemyVoiceLibraryRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.ingest_job_repository import (
|
||||
SQLAlchemyIngestJobRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.project_repository import (
|
||||
SQLAlchemyProjectRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.project_title_repository import (
|
||||
SQLAlchemyProjectTitleRepository,
|
||||
)
|
||||
from packages.adapters.sqlalchemy_impl.session import build_session_factory
|
||||
from packages.adapters.sqlalchemy_impl.user_repository import SQLAlchemyUserRepository
|
||||
from packages.ports.asset_repository import AssetRepository
|
||||
@@ -43,10 +46,11 @@ from packages.ports.asset_library_repository import AssetLibraryRepository
|
||||
from packages.ports.user_repository import UserRepository
|
||||
from packages.ports.classification_job_repository import ClassificationJobRepository
|
||||
from packages.ports.generation_task_repository import GenerationTaskRepository
|
||||
from packages.ports.title_library_repository import TitleLibraryRepository
|
||||
from packages.ports.voice_library_repository import VoiceLibraryRepository
|
||||
from packages.ports.generated_video_repository import GeneratedVideoRepository
|
||||
from packages.ports.ingest_job_repository import IngestJobRepository
|
||||
from packages.ports.project_repository import ProjectRepository
|
||||
from packages.ports.project_title_repository import ProjectTitleRepository
|
||||
|
||||
_engine, _SessionLocal = build_session_factory(settings.DATABASE_URL)
|
||||
|
||||
@@ -109,12 +113,6 @@ def get_project_repository(
|
||||
return SQLAlchemyProjectRepository(session)
|
||||
|
||||
|
||||
def get_project_title_repository(
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> SQLAlchemyProjectTitleRepository:
|
||||
"""Provide the SQLAlchemy project title repository implementation."""
|
||||
return SQLAlchemyProjectTitleRepository(session)
|
||||
|
||||
|
||||
def get_user_repository(
|
||||
session: Session = Depends(get_db_session),
|
||||
@@ -146,3 +144,16 @@ def get_auth_email_service() -> NoopEmailService | EmailService:
|
||||
),
|
||||
enabled=True,
|
||||
)
|
||||
|
||||
def get_title_library_repository(
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> SQLAlchemyTitleLibraryRepository:
|
||||
"""Provide the SQLAlchemy title library repository implementation."""
|
||||
return SQLAlchemyTitleLibraryRepository(session)
|
||||
|
||||
|
||||
def get_voice_library_repository(
|
||||
session: Session = Depends(get_db_session),
|
||||
) -> SQLAlchemyVoiceLibraryRepository:
|
||||
"""Provide the SQLAlchemy voice library repository implementation."""
|
||||
return SQLAlchemyVoiceLibraryRepository(session)
|
||||
|
||||
@@ -1,60 +0,0 @@
|
||||
from datetime import datetime
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class EditTemplateResponse(BaseModel):
|
||||
id: str
|
||||
project_id: str
|
||||
name: str
|
||||
description: str
|
||||
target_duration: float
|
||||
clip_count: int
|
||||
is_active: bool
|
||||
created_at: datetime | None = None
|
||||
|
||||
|
||||
class EditPlanClipResponse(BaseModel):
|
||||
id: str
|
||||
asset_id: str
|
||||
asset_name: str
|
||||
sequence: int
|
||||
start_time: float
|
||||
duration: float
|
||||
reason: str
|
||||
layer: str = "main" # main, pip, broll
|
||||
|
||||
|
||||
class EditPlanResponse(BaseModel):
|
||||
id: str
|
||||
project_id: str
|
||||
template_id: str
|
||||
asset_library_id: str
|
||||
title_id: str = ""
|
||||
status: str
|
||||
summary: str
|
||||
editing_mode: str | None = None # one_take, pip, voice_over, voice_pip
|
||||
clips: list[EditPlanClipResponse] = Field(default_factory=list)
|
||||
created_at: datetime | None = None
|
||||
updated_at: datetime | None = None
|
||||
|
||||
|
||||
class CreateEditPlanRequest(BaseModel):
|
||||
asset_library_id: str
|
||||
template_id: str = ""
|
||||
title_id: str = ""
|
||||
|
||||
|
||||
class AutoGenerateEditPlanRequest(BaseModel):
|
||||
"""智能生成剪辑计划请求"""
|
||||
asset_library_id: str
|
||||
editing_mode: str = Field(
|
||||
default="one_take",
|
||||
description="剪辑模式: one_take, pip, voice_over, voice_pip"
|
||||
)
|
||||
target_duration: float = Field(
|
||||
default=30.0,
|
||||
description="目标时长(秒)"
|
||||
)
|
||||
template_id: str = ""
|
||||
title_id: str = ""
|
||||
@@ -6,7 +6,6 @@ class CreateGenerationTaskRequest(BaseModel):
|
||||
asset_library_id: str = Field(..., min_length=1)
|
||||
strategy_id: str = ""
|
||||
voice_library_id: str = ""
|
||||
edit_plan_id: str = ""
|
||||
created_by_user_id: str = ""
|
||||
|
||||
|
||||
@@ -16,7 +15,6 @@ class GenerationTaskResponse(BaseModel):
|
||||
asset_library_id: str
|
||||
strategy_id: str
|
||||
voice_library_id: str
|
||||
edit_plan_id: str
|
||||
status: str
|
||||
progress: float
|
||||
result_count: int
|
||||
|
||||
@@ -1,35 +0,0 @@
|
||||
from datetime import datetime
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
TitleCategory = Literal["default", "marketing", "tutorial", "story", "promo"]
|
||||
|
||||
|
||||
class ProjectTitleResponse(BaseModel):
|
||||
id: str
|
||||
project_id: str
|
||||
text: str
|
||||
category: str
|
||||
favorite: bool
|
||||
usage_count: int
|
||||
is_active: bool
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class ListProjectTitlesResponse(BaseModel):
|
||||
items: list[ProjectTitleResponse]
|
||||
|
||||
|
||||
class CreateProjectTitleRequest(BaseModel):
|
||||
text: str = Field(min_length=1, max_length=200)
|
||||
category: TitleCategory = "default"
|
||||
favorite: bool = False
|
||||
|
||||
|
||||
class UpdateProjectTitleRequest(BaseModel):
|
||||
text: str | None = Field(default=None, min_length=1, max_length=200)
|
||||
category: TitleCategory | None = None
|
||||
favorite: bool | None = None
|
||||
is_active: bool | None = None
|
||||
@@ -0,0 +1,42 @@
|
||||
"""Title library Pydantic schemas."""
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import List, Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class TitleLibraryItemResponse(BaseModel):
|
||||
id: str
|
||||
user_id: str
|
||||
name: str
|
||||
text: str
|
||||
category: str = "default"
|
||||
description: str = ""
|
||||
tags: List[str] = Field(default_factory=list)
|
||||
usage_count: int = 0
|
||||
is_active: bool = True
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class ListTitleLibraryResponse(BaseModel):
|
||||
items: list[TitleLibraryItemResponse]
|
||||
total: int = 0
|
||||
|
||||
|
||||
class CreateTitleLibraryRequest(BaseModel):
|
||||
name: str = Field(..., min_length=1, max_length=255)
|
||||
text: str = Field(..., min_length=1, max_length=500)
|
||||
category: str = "default"
|
||||
description: str = ""
|
||||
tags: List[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class UpdateTitleLibraryRequest(BaseModel):
|
||||
name: Optional[str] = Field(None, min_length=1, max_length=255)
|
||||
text: Optional[str] = Field(None, min_length=1, max_length=500)
|
||||
category: Optional[str] = None
|
||||
description: Optional[str] = None
|
||||
tags: Optional[List[str]] = None
|
||||
@@ -0,0 +1,57 @@
|
||||
"""Voice library Pydantic schemas."""
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import List, Optional
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
|
||||
class VoiceLibraryItemResponse(BaseModel):
|
||||
id: str
|
||||
user_id: str
|
||||
name: str
|
||||
text: str = ""
|
||||
voice_provider: str = ""
|
||||
voice_id: str = ""
|
||||
voice_name: str = ""
|
||||
audio_url: str = ""
|
||||
duration: float = 0
|
||||
file_size: int = 0
|
||||
status: str = "completed"
|
||||
project_id: Optional[str] = None
|
||||
tags: List[str] = Field(default_factory=list)
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
|
||||
|
||||
class ListVoiceLibraryResponse(BaseModel):
|
||||
items: list[VoiceLibraryItemResponse]
|
||||
total: int = 0
|
||||
|
||||
|
||||
class CreateVoiceLibraryRequest(BaseModel):
|
||||
name: str = Field(..., min_length=1, max_length=255)
|
||||
text: str = ""
|
||||
voice_provider: str = ""
|
||||
voice_id: str = ""
|
||||
voice_name: str = ""
|
||||
audio_url: str = ""
|
||||
duration: float = 0
|
||||
file_size: int = 0
|
||||
status: str = "completed"
|
||||
project_id: Optional[str] = None
|
||||
tags: List[str] = Field(default_factory=list)
|
||||
|
||||
|
||||
class UpdateVoiceLibraryRequest(BaseModel):
|
||||
name: Optional[str] = Field(None, min_length=1, max_length=255)
|
||||
text: Optional[str] = None
|
||||
voice_provider: Optional[str] = None
|
||||
voice_id: Optional[str] = None
|
||||
voice_name: Optional[str] = None
|
||||
audio_url: Optional[str] = None
|
||||
duration: Optional[float] = None
|
||||
file_size: Optional[int] = None
|
||||
status: Optional[str] = None
|
||||
tags: Optional[List[str]] = None
|
||||
@@ -16,7 +16,7 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# 从 domain 层导入 EditingMode,避免重复定义
|
||||
from packages.domain.edit_plan import EditingMode
|
||||
from packages.domain.editing_mode import EditingMode
|
||||
|
||||
|
||||
class PIPPosition(StrEnum):
|
||||
|
||||
@@ -1,13 +1,13 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import ProjectTitleModel
|
||||
from packages.adapters.sqlalchemy_impl.models import TitleLibraryModel
|
||||
|
||||
|
||||
def mark_title_used_for_generation(db, task) -> None:
|
||||
if not task.strategy_id:
|
||||
return
|
||||
title = db.query(ProjectTitleModel).filter(ProjectTitleModel.id == task.strategy_id).first()
|
||||
if title is None or title.project_id != task.project_id:
|
||||
title = db.query(TitleLibraryModel).filter(TitleLibraryModel.id == task.strategy_id).first()
|
||||
if title is None:
|
||||
return
|
||||
title.usage_count = int(title.usage_count or 0) + 1
|
||||
title.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
@@ -1,358 +0,0 @@
|
||||
"""Smart Edit Plan Generator - 根据分类和质量评分智能编排剪辑计划"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
from collections import defaultdict
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
from typing import Any
|
||||
|
||||
from packages.domain import Asset, AssetClassification, AssetStatus
|
||||
from packages.domain.edit_plan import EditClipPlan, EditPlanResult, EditingMode
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _calculate_start_times(clips: list[EditClipPlan]) -> list[EditClipPlan]:
|
||||
"""
|
||||
计算时间轴,根据前面的片段时长累加 start_time
|
||||
|
||||
Args:
|
||||
clips: 已按 sequence 排序的片段列表
|
||||
|
||||
Returns:
|
||||
修正后的片段列表
|
||||
"""
|
||||
current_time = 0.0
|
||||
for clip in clips:
|
||||
clip.start_time = current_time
|
||||
current_time += clip.duration
|
||||
return clips
|
||||
|
||||
|
||||
class SmartEditPlanGenerator:
|
||||
"""
|
||||
智能剪辑计划生成器
|
||||
|
||||
根据素材的分类结果和质量评分,自动编排剪辑计划。
|
||||
支持多种剪辑模式:one_take, pip, voice_over, voice_pip
|
||||
"""
|
||||
|
||||
def __init__(self, project_id: str, assets: list[Asset]):
|
||||
self.project_id = project_id
|
||||
# 筛选已就绪的视频素材
|
||||
self.assets = [
|
||||
a for a in assets
|
||||
if a.status == AssetStatus.READY and a.mime_type.startswith("video/")
|
||||
]
|
||||
self.assets_by_classification: dict[str, list[Asset]] = defaultdict(list)
|
||||
|
||||
def _parse_classification(self, asset: Asset) -> str:
|
||||
"""解析素材的分类结果"""
|
||||
# 从 metadata 中获取分类
|
||||
classification = asset.metadata.get("classification", "")
|
||||
if not classification:
|
||||
# 尝试从 classification_result 字段获取
|
||||
classification = asset.metadata.get("classification_result", "")
|
||||
|
||||
# 如果是 JSON 字符串,解析它
|
||||
if classification and isinstance(classification, str):
|
||||
try:
|
||||
parsed = json.loads(classification)
|
||||
if isinstance(parsed, dict):
|
||||
classification = parsed.get("classification", "other")
|
||||
elif isinstance(parsed, str):
|
||||
classification = parsed
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
|
||||
# 验证分类值是否有效
|
||||
valid_classifications = [c.value for c in AssetClassification]
|
||||
if classification not in valid_classifications:
|
||||
classification = "other"
|
||||
|
||||
return classification
|
||||
|
||||
def _group_by_classification(self) -> None:
|
||||
"""按分类结果对素材分组"""
|
||||
for asset in self.assets:
|
||||
classification = self._parse_classification(asset)
|
||||
self.assets_by_classification[classification].append(asset)
|
||||
|
||||
def _sort_by_quality(self, assets: list[Asset]) -> list[Asset]:
|
||||
"""按质量评分排序,高分在前"""
|
||||
return sorted(
|
||||
assets,
|
||||
key=lambda a: (-(a.quality_score or 0), a.created_at)
|
||||
)
|
||||
|
||||
def _calculate_clip_duration(self, asset: Asset, target_duration: float, clip_count: int) -> float:
|
||||
"""计算单个片段的时长"""
|
||||
if asset.duration:
|
||||
# 如果素材时长超过平均时长,取平均时长
|
||||
avg_duration = target_duration / max(1, clip_count)
|
||||
return min(float(asset.duration), avg_duration)
|
||||
return target_duration / max(1, clip_count)
|
||||
|
||||
def _generate_one_take(self, target_duration: float = 30.0) -> EditPlanResult:
|
||||
"""
|
||||
One-Take 模式:按分类分组,组内按质量排序,顺序拼接
|
||||
"""
|
||||
self._group_by_classification()
|
||||
|
||||
clips: list[EditClipPlan] = []
|
||||
sequence = 1
|
||||
|
||||
# 按优先级排序分类:person > scenic > product > other
|
||||
priority_order = ["person", "scenic", "product", "animal", "food", "tech", "sport", "music", "other"]
|
||||
sorted_classifications = sorted(
|
||||
self.assets_by_classification.keys(),
|
||||
key=lambda c: priority_order.index(c) if c in priority_order else len(priority_order)
|
||||
)
|
||||
|
||||
for classification in sorted_classifications:
|
||||
sorted_assets = self._sort_by_quality(self.assets_by_classification[classification])
|
||||
for asset in sorted_assets:
|
||||
duration = self._calculate_clip_duration(
|
||||
asset, target_duration, len(self.assets)
|
||||
)
|
||||
clips.append(EditClipPlan(
|
||||
asset_id=asset.id,
|
||||
sequence=sequence,
|
||||
start_time=0,
|
||||
duration=duration,
|
||||
layer="main",
|
||||
reason=f"按分类 [{classification}] 排列,质量评分 {asset.quality_score or 0:.1f}"
|
||||
))
|
||||
sequence += 1
|
||||
|
||||
total_duration = sum(c.duration for c in clips)
|
||||
_calculate_start_times(clips)
|
||||
return EditPlanResult(
|
||||
project_id=self.project_id,
|
||||
editing_mode=EditingMode.ONE_TAKE,
|
||||
clips=clips,
|
||||
total_duration=total_duration,
|
||||
summary=f"One-Take 模式:按 {len(sorted_classifications)} 个分类分组,共 {len(clips)} 段素材"
|
||||
)
|
||||
|
||||
def _generate_pip(self, target_duration: float = 30.0) -> EditPlanResult:
|
||||
"""
|
||||
PIP 模式:第一个高质量素材为主画面,其余为画中画
|
||||
"""
|
||||
sorted_assets = self._sort_by_quality(self.assets)
|
||||
|
||||
if not sorted_assets:
|
||||
return EditPlanResult(
|
||||
project_id=self.project_id,
|
||||
editing_mode=EditingMode.PIP,
|
||||
clips=[],
|
||||
total_duration=0,
|
||||
summary="无素材可用"
|
||||
)
|
||||
|
||||
clips: list[EditClipPlan] = []
|
||||
sequence = 1
|
||||
|
||||
# 第一个高质量素材作为主画面
|
||||
main_asset = sorted_assets[0]
|
||||
main_duration = min(
|
||||
float(main_asset.duration) if main_asset.duration else target_duration,
|
||||
target_duration
|
||||
)
|
||||
clips.append(EditClipPlan(
|
||||
asset_id=main_asset.id,
|
||||
sequence=sequence,
|
||||
start_time=0,
|
||||
duration=main_duration,
|
||||
layer="main",
|
||||
reason=f"高质量主画面 (质量评分: {main_asset.quality_score or 0:.1f})"
|
||||
))
|
||||
sequence += 1
|
||||
|
||||
# 其余素材作为画中画
|
||||
for asset in sorted_assets[1:]:
|
||||
duration = self._calculate_clip_duration(asset, target_duration, len(sorted_assets))
|
||||
clips.append(EditClipPlan(
|
||||
asset_id=asset.id,
|
||||
sequence=sequence,
|
||||
start_time=0,
|
||||
duration=duration,
|
||||
layer="pip",
|
||||
reason=f"画中画素材 (质量评分: {asset.quality_score or 0:.1f})"
|
||||
))
|
||||
sequence += 1
|
||||
|
||||
total_duration = main_duration
|
||||
_calculate_start_times(clips)
|
||||
return EditPlanResult(
|
||||
project_id=self.project_id,
|
||||
editing_mode=EditingMode.PIP,
|
||||
clips=clips,
|
||||
total_duration=total_duration,
|
||||
summary=f"PIP 模式:1 个主画面 + {len(sorted_assets) - 1} 个画中画"
|
||||
)
|
||||
|
||||
def _generate_voiceover(self, target_duration: float = 30.0) -> EditPlanResult:
|
||||
"""
|
||||
Voiceover 模式:person 类素材为主播口播,其余穿插为 B-roll
|
||||
"""
|
||||
self._group_by_classification()
|
||||
|
||||
person_assets = self._sort_by_quality(
|
||||
self.assets_by_classification.get("person", [])
|
||||
)
|
||||
other_assets = self._sort_by_quality([
|
||||
a for assets in self.assets_by_classification.values()
|
||||
for a in assets
|
||||
if self._parse_classification(a) != "person"
|
||||
])
|
||||
|
||||
clips: list[EditClipPlan] = []
|
||||
sequence = 1
|
||||
|
||||
# 合并口播和 B-roll
|
||||
main_assets = person_assets if person_assets else other_assets
|
||||
broll_assets = [a for a in other_assets if a not in person_assets] if person_assets else []
|
||||
|
||||
# 优先使用 person 素材作为口播
|
||||
for i, asset in enumerate(main_assets):
|
||||
duration = self._calculate_clip_duration(asset, target_duration, len(main_assets))
|
||||
is_person = asset in person_assets
|
||||
clips.append(EditClipPlan(
|
||||
asset_id=asset.id,
|
||||
sequence=sequence,
|
||||
start_time=0,
|
||||
duration=duration,
|
||||
layer="main" if is_person else "broll",
|
||||
reason=f"{'主播口播' if is_person else 'B-roll'} (质量评分: {asset.quality_score or 0:.1f})"
|
||||
))
|
||||
sequence += 1
|
||||
|
||||
# 在口播之间穿插 B-roll
|
||||
if is_person and broll_assets and i < len(main_assets) - 1:
|
||||
broll_asset = broll_assets[i % len(broll_assets)]
|
||||
broll_duration = self._calculate_clip_duration(
|
||||
broll_asset, target_duration, len(main_assets) + len(broll_assets)
|
||||
)
|
||||
clips.append(EditClipPlan(
|
||||
asset_id=broll_asset.id,
|
||||
sequence=sequence,
|
||||
start_time=0,
|
||||
duration=broll_duration,
|
||||
layer="broll",
|
||||
reason=f"B-roll 穿插 (质量评分: {broll_asset.quality_score or 0:.1f})"
|
||||
))
|
||||
sequence += 1
|
||||
|
||||
total_duration = sum(c.duration for c in clips)
|
||||
person_count = len(person_assets)
|
||||
_calculate_start_times(clips)
|
||||
return EditPlanResult(
|
||||
project_id=self.project_id,
|
||||
editing_mode=EditingMode.VOICE_OVER,
|
||||
clips=clips,
|
||||
total_duration=total_duration,
|
||||
summary=f"Voiceover 模式:{person_count} 段口播 + {len(clips) - person_count} 段 B-roll"
|
||||
)
|
||||
|
||||
def _generate_voice_pip(self, target_duration: float = 30.0) -> EditPlanResult:
|
||||
"""
|
||||
Voice-PIP 模式:结合 voiceover 和 pip
|
||||
第一个高质量 person 素材为主画面,其余为 PIP B-roll
|
||||
"""
|
||||
self._group_by_classification()
|
||||
|
||||
person_assets = self._sort_by_quality(
|
||||
self.assets_by_classification.get("person", [])
|
||||
)
|
||||
other_assets = self._sort_by_quality([
|
||||
a for assets in self.assets_by_classification.values()
|
||||
for a in assets
|
||||
if self._parse_classification(a) != "person"
|
||||
])
|
||||
|
||||
clips: list[EditClipPlan] = []
|
||||
sequence = 1
|
||||
|
||||
# 主画面:优先使用高质量 person 素材
|
||||
main_asset = person_assets[0] if person_assets else (other_assets[0] if other_assets else None)
|
||||
if main_asset:
|
||||
main_duration = min(
|
||||
float(main_asset.duration) if main_asset.duration else target_duration,
|
||||
target_duration
|
||||
)
|
||||
is_person = main_asset in person_assets
|
||||
clips.append(EditClipPlan(
|
||||
asset_id=main_asset.id,
|
||||
sequence=sequence,
|
||||
start_time=0,
|
||||
duration=main_duration,
|
||||
layer="main",
|
||||
reason=f"{'主播口播' if is_person else '主画面'} (质量评分: {main_asset.quality_score or 0:.1f})"
|
||||
))
|
||||
sequence += 1
|
||||
|
||||
# PIP 素材
|
||||
pip_assets = [a for a in (person_assets[1:] + other_assets) if a != main_asset]
|
||||
for asset in pip_assets:
|
||||
duration = self._calculate_clip_duration(asset, target_duration, len(pip_assets) + 1)
|
||||
clips.append(EditClipPlan(
|
||||
asset_id=asset.id,
|
||||
sequence=sequence,
|
||||
start_time=0,
|
||||
duration=duration,
|
||||
layer="pip",
|
||||
reason=f"PIP 素材 (质量评分: {asset.quality_score or 0:.1f})"
|
||||
))
|
||||
sequence += 1
|
||||
|
||||
total_duration = sum(c.duration for c in clips)
|
||||
_calculate_start_times(clips)
|
||||
return EditPlanResult(
|
||||
project_id=self.project_id,
|
||||
editing_mode=EditingMode.VOICE_PIP,
|
||||
clips=clips,
|
||||
total_duration=total_duration,
|
||||
summary=f"Voice-PIP 模式:1 个主画面 + {len(pip_assets)} 个 PIP 素材"
|
||||
)
|
||||
|
||||
def generate_plan(
|
||||
self,
|
||||
editing_mode: str = "one_take",
|
||||
target_duration: float = 30.0
|
||||
) -> EditPlanResult:
|
||||
"""
|
||||
生成剪辑计划
|
||||
|
||||
Args:
|
||||
editing_mode: 剪辑模式 (one_take/pip/voice_over/voice_pip)
|
||||
target_duration: 目标时长(秒)
|
||||
|
||||
Returns:
|
||||
EditPlanResult: 编排好的剪辑计划
|
||||
"""
|
||||
logger.info(f"Generating edit plan for project {self.project_id} with mode {editing_mode}")
|
||||
|
||||
if not self.assets:
|
||||
logger.warning(f"No ready video assets found for project {self.project_id}")
|
||||
return EditPlanResult(
|
||||
project_id=self.project_id,
|
||||
editing_mode=EditingMode(editing_mode),
|
||||
clips=[],
|
||||
total_duration=0,
|
||||
summary="无素材可用"
|
||||
)
|
||||
|
||||
mode = EditingMode(editing_mode.lower())
|
||||
|
||||
if mode == EditingMode.ONE_TAKE:
|
||||
return self._generate_one_take(target_duration)
|
||||
elif mode == EditingMode.PIP:
|
||||
return self._generate_pip(target_duration)
|
||||
elif mode == EditingMode.VOICE_OVER:
|
||||
return self._generate_voiceover(target_duration)
|
||||
elif mode == EditingMode.VOICE_PIP:
|
||||
return self._generate_voice_pip(target_duration)
|
||||
else:
|
||||
raise ValueError(f"Unknown editing mode: {editing_mode}")
|
||||
@@ -219,7 +219,7 @@ def generate_video(self, task_id: str) -> dict:
|
||||
Returns:
|
||||
生成结果字典
|
||||
"""
|
||||
from packages.domain import GeneratedVideo, GenerationMode, GenerationTaskStatus
|
||||
from packages.domain import EditingMode, GeneratedVideo, GenerationTaskStatus
|
||||
from packages.adapters.sqlalchemy_impl.generation_task_repository import (
|
||||
SQLAlchemyGenerationTaskRepository,
|
||||
)
|
||||
@@ -240,9 +240,9 @@ def generate_video(self, task_id: str) -> dict:
|
||||
session.close()
|
||||
|
||||
try:
|
||||
editing_mode = GenerationMode(mode)
|
||||
editing_mode = EditingMode(mode)
|
||||
except ValueError:
|
||||
editing_mode = GenerationMode.ONE_TAKE
|
||||
editing_mode = EditingMode.ONE_TAKE
|
||||
|
||||
output_name = f"generated-{task_id}.mp4"
|
||||
storage_key = f"generated/projects/{project_id}/tasks/{task_id}/{output_name}"
|
||||
|
||||
@@ -1,92 +0,0 @@
|
||||
"""项目管理 In-Memory Repository 实现"""
|
||||
|
||||
from packages.domain import Milestone, Task, TaskIssue
|
||||
from packages.ports.project_management_repositories import (
|
||||
MilestoneRepository,
|
||||
TaskIssueRepository,
|
||||
TaskRepository,
|
||||
)
|
||||
|
||||
|
||||
class InMemoryTaskRepository(TaskRepository):
|
||||
"""任务 In-Memory 仓储实现"""
|
||||
|
||||
def __init__(self):
|
||||
self._store: dict[str, Task] = {}
|
||||
|
||||
def create(self, task: Task) -> Task:
|
||||
self._store[task.id] = task
|
||||
return task
|
||||
|
||||
def get_by_id(self, task_id: str) -> Task | None:
|
||||
return self._store.get(task_id)
|
||||
|
||||
def list_by_project(self, project_id: str) -> list[Task]:
|
||||
return [t for t in self._store.values() if t.project_id == project_id]
|
||||
|
||||
def list_by_parent(self, parent_task_id: str) -> list[Task]:
|
||||
return [t for t in self._store.values() if t.parent_task_id == parent_task_id]
|
||||
|
||||
def update(self, task: Task) -> Task:
|
||||
if task.id not in self._store:
|
||||
raise ValueError(f"Task {task.id} not found")
|
||||
self._store[task.id] = task
|
||||
return task
|
||||
|
||||
def delete(self, task_id: str) -> None:
|
||||
self._store.pop(task_id, None)
|
||||
|
||||
|
||||
class InMemoryMilestoneRepository(MilestoneRepository):
|
||||
"""里程碑 In-Memory 仓储实现"""
|
||||
|
||||
def __init__(self):
|
||||
self._store: dict[str, Milestone] = {}
|
||||
|
||||
def create(self, milestone: Milestone) -> Milestone:
|
||||
self._store[milestone.id] = milestone
|
||||
return milestone
|
||||
|
||||
def get_by_id(self, milestone_id: str) -> Milestone | None:
|
||||
return self._store.get(milestone_id)
|
||||
|
||||
def list_by_project(self, project_id: str) -> list[Milestone]:
|
||||
return [m for m in self._store.values() if m.project_id == project_id]
|
||||
|
||||
def update(self, milestone: Milestone) -> Milestone:
|
||||
if milestone.id not in self._store:
|
||||
raise ValueError(f"Milestone {milestone.id} not found")
|
||||
self._store[milestone.id] = milestone
|
||||
return milestone
|
||||
|
||||
def delete(self, milestone_id: str) -> None:
|
||||
self._store.pop(milestone_id, None)
|
||||
|
||||
|
||||
class InMemoryTaskIssueRepository(TaskIssueRepository):
|
||||
"""任务问题 In-Memory 仓储实现"""
|
||||
|
||||
def __init__(self):
|
||||
self._store: dict[str, TaskIssue] = {}
|
||||
|
||||
def create(self, issue: TaskIssue) -> TaskIssue:
|
||||
self._store[issue.id] = issue
|
||||
return issue
|
||||
|
||||
def get_by_id(self, issue_id: str) -> TaskIssue | None:
|
||||
return self._store.get(issue_id)
|
||||
|
||||
def list_by_task(self, task_id: str) -> list[TaskIssue]:
|
||||
return [i for i in self._store.values() if i.task_id == task_id]
|
||||
|
||||
def list_by_project(self, project_id: str) -> list[TaskIssue]:
|
||||
return [i for i in self._store.values() if i.project_id == project_id]
|
||||
|
||||
def update(self, issue: TaskIssue) -> TaskIssue:
|
||||
if issue.id not in self._store:
|
||||
raise ValueError(f"TaskIssue {issue.id} not found")
|
||||
self._store[issue.id] = issue
|
||||
return issue
|
||||
|
||||
def delete(self, issue_id: str) -> None:
|
||||
self._store.pop(issue_id, None)
|
||||
@@ -15,7 +15,6 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
strategy_id=task.strategy_id,
|
||||
asset_library_id=task.asset_library_id,
|
||||
voice_library_id=task.voice_library_id,
|
||||
edit_plan_id=task.edit_plan_id,
|
||||
status=task.status,
|
||||
progress=task.progress,
|
||||
result_count=task.result_count,
|
||||
@@ -39,7 +38,6 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
strategy_id=model.strategy_id,
|
||||
asset_library_id=model.asset_library_id,
|
||||
voice_library_id=model.voice_library_id,
|
||||
edit_plan_id=getattr(model, "edit_plan_id", "") or "",
|
||||
status=model.status,
|
||||
progress=model.progress,
|
||||
result_count=int(model.result_count or 0),
|
||||
@@ -59,7 +57,6 @@ class SQLAlchemyGenerationTaskRepository:
|
||||
if model is None:
|
||||
raise ValueError(f"GenerationTask {task.id} not found")
|
||||
model.voice_library_id = task.voice_library_id
|
||||
model.edit_plan_id = task.edit_plan_id
|
||||
model.status = task.status
|
||||
model.progress = task.progress
|
||||
model.result_count = task.result_count
|
||||
|
||||
@@ -87,20 +87,6 @@ class AssetModel(Base):
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
class ProjectTitleModel(Base):
|
||||
__tablename__ = "project_titles"
|
||||
|
||||
id = Column(String(36), primary_key=True)
|
||||
project_id = Column(String(36), nullable=False, index=True)
|
||||
text = Column(String(200), nullable=False)
|
||||
category = Column(String(50), nullable=False, default="default", index=True)
|
||||
favorite = Column(Boolean, nullable=False, default=False, index=True)
|
||||
usage_count = Column(Integer, nullable=False, default=0)
|
||||
is_active = Column(Boolean, nullable=False, default=True)
|
||||
created_by_user_id = Column(String(36), nullable=False)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
class EditTemplateModel(Base):
|
||||
__tablename__ = "edit_templates"
|
||||
@@ -118,33 +104,7 @@ class EditTemplateModel(Base):
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
class EditPlanModel(Base):
|
||||
__tablename__ = "edit_plans"
|
||||
|
||||
id = Column(String(32), primary_key=True)
|
||||
project_id = Column(String(32), nullable=False, index=True)
|
||||
template_id = Column(String(32), nullable=False, index=True)
|
||||
asset_library_id = Column(String(32), nullable=False, index=True)
|
||||
title_id = Column(String(32), nullable=False, default="")
|
||||
editing_mode = Column(String(20), nullable=True, default=None, index=True)
|
||||
status = Column(String(20), nullable=False, default="draft", index=True)
|
||||
summary = Column(Text, nullable=False, default="")
|
||||
created_by_user_id = Column(String(32), nullable=False, default="")
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
class EditPlanClipModel(Base):
|
||||
__tablename__ = "edit_plan_clips"
|
||||
|
||||
id = Column(String(32), primary_key=True)
|
||||
edit_plan_id = Column(String(32), nullable=False, index=True)
|
||||
asset_id = Column(String(32), nullable=False, index=True)
|
||||
sequence = Column(Integer, nullable=False)
|
||||
start_time = Column(Float, nullable=False, default=0)
|
||||
duration = Column(Float, nullable=False, default=0)
|
||||
reason = Column(Text, nullable=False, default="")
|
||||
layer = Column(String(20), nullable=False, default="main")
|
||||
|
||||
|
||||
class IngestJobModel(Base):
|
||||
@@ -183,7 +143,6 @@ class GenerationTaskModel(Base):
|
||||
strategy_id = Column(String(32), nullable=False, default="")
|
||||
asset_library_id = Column(String(32), nullable=False, index=True)
|
||||
voice_library_id = Column(String(32), nullable=False, default="")
|
||||
edit_plan_id = Column(String(32), nullable=False, default="", index=True)
|
||||
editing_mode = Column(String(20), nullable=False, default="one_take", index=True) # 剪辑模式: one_take, pip, voice_over, voice_pip
|
||||
status = Column(String(20), nullable=False, default="pending", index=True)
|
||||
progress = Column(Float, nullable=False, default=0.0)
|
||||
@@ -223,54 +182,8 @@ class GeneratedVideoModel(Base):
|
||||
duplicate_of = Column(String(32), nullable=True)
|
||||
|
||||
|
||||
class TaskModel(Base):
|
||||
__tablename__ = "tasks"
|
||||
|
||||
id = Column(String(32), primary_key=True)
|
||||
project_id = Column(String(32), nullable=False, index=True)
|
||||
name = Column(String(200), nullable=False)
|
||||
description = Column(Text, nullable=False, default="")
|
||||
status = Column(String(20), nullable=False, default="pending", index=True)
|
||||
priority = Column(String(20), nullable=False, default="medium")
|
||||
parent_task_id = Column(String(32), nullable=False, default="", index=True)
|
||||
assignee_user_id = Column(String(32), nullable=False, default="")
|
||||
progress = Column(Float, nullable=False, default=0.0)
|
||||
planned_start_date = Column(DateTime, nullable=True)
|
||||
planned_end_date = Column(DateTime, nullable=True)
|
||||
actual_start_date = Column(DateTime, nullable=True)
|
||||
actual_end_date = Column(DateTime, nullable=True)
|
||||
tags_json = Column(Text, nullable=False, default="[]")
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
class MilestoneModel(Base):
|
||||
__tablename__ = "milestones"
|
||||
|
||||
id = Column(String(32), primary_key=True)
|
||||
project_id = Column(String(32), nullable=False, index=True)
|
||||
name = Column(String(200), nullable=False)
|
||||
description = Column(Text, nullable=False, default="")
|
||||
target_date = Column(DateTime, nullable=True)
|
||||
completed = Column(Boolean, nullable=False, default=False)
|
||||
completed_at = Column(DateTime, nullable=True)
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
class TaskIssueModel(Base):
|
||||
__tablename__ = "task_issues"
|
||||
|
||||
id = Column(String(32), primary_key=True)
|
||||
task_id = Column(String(32), nullable=False, index=True)
|
||||
project_id = Column(String(32), nullable=False, index=True)
|
||||
title = Column(String(200), nullable=False)
|
||||
description = Column(Text, nullable=False, default="")
|
||||
resolved = Column(Boolean, nullable=False, default=False)
|
||||
resolved_at = Column(DateTime, nullable=True)
|
||||
created_by_user_id = Column(String(32), nullable=False, default="")
|
||||
created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
updated_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc))
|
||||
|
||||
class TitleLibraryModel(Base):
|
||||
__tablename__ = "title_libraries"
|
||||
|
||||
@@ -1,241 +0,0 @@
|
||||
"""项目管理 SQLAlchemy Repository 实现"""
|
||||
|
||||
import json
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.domain import Milestone, Task, TaskIssue
|
||||
from packages.ports.project_management_repositories import (
|
||||
MilestoneRepository,
|
||||
TaskIssueRepository,
|
||||
TaskRepository,
|
||||
)
|
||||
|
||||
from .models import MilestoneModel, TaskIssueModel, TaskModel
|
||||
|
||||
|
||||
class SQLAlchemyTaskRepository(TaskRepository):
|
||||
"""任务 SQLAlchemy 仓储实现"""
|
||||
|
||||
def __init__(self, session: Session):
|
||||
self._session = session
|
||||
|
||||
def create(self, task: Task) -> Task:
|
||||
model = TaskModel(
|
||||
id=task.id,
|
||||
project_id=task.project_id,
|
||||
name=task.name,
|
||||
description=task.description,
|
||||
status=task.status.value,
|
||||
priority=task.priority.value,
|
||||
parent_task_id=task.parent_task_id,
|
||||
assignee_user_id=task.assignee_user_id,
|
||||
progress=task.progress,
|
||||
planned_start_date=task.planned_start_date,
|
||||
planned_end_date=task.planned_end_date,
|
||||
actual_start_date=task.actual_start_date,
|
||||
actual_end_date=task.actual_end_date,
|
||||
tags_json=json.dumps(task.tags, ensure_ascii=False),
|
||||
created_at=task.created_at,
|
||||
updated_at=task.updated_at,
|
||||
)
|
||||
self._session.add(model)
|
||||
self._session.commit()
|
||||
return task
|
||||
|
||||
def get_by_id(self, task_id: str) -> Task | None:
|
||||
model = self._session.query(TaskModel).filter(TaskModel.id == task_id).first()
|
||||
if not model:
|
||||
return None
|
||||
return self._model_to_entity(model)
|
||||
|
||||
def list_by_project(self, project_id: str) -> list[Task]:
|
||||
models = self._session.query(TaskModel).filter(TaskModel.project_id == project_id).all()
|
||||
return [self._model_to_entity(m) for m in models]
|
||||
|
||||
def list_by_parent(self, parent_task_id: str) -> list[Task]:
|
||||
models = self._session.query(TaskModel).filter(TaskModel.parent_task_id == parent_task_id).all()
|
||||
return [self._model_to_entity(m) for m in models]
|
||||
|
||||
def update(self, task: Task) -> Task:
|
||||
model = self._session.query(TaskModel).filter(TaskModel.id == task.id).first()
|
||||
if not model:
|
||||
raise ValueError(f"Task {task.id} not found")
|
||||
|
||||
model.name = task.name
|
||||
model.description = task.description
|
||||
model.status = task.status.value
|
||||
model.priority = task.priority.value
|
||||
model.parent_task_id = task.parent_task_id
|
||||
model.assignee_user_id = task.assignee_user_id
|
||||
model.progress = task.progress
|
||||
model.planned_start_date = task.planned_start_date
|
||||
model.planned_end_date = task.planned_end_date
|
||||
model.actual_start_date = task.actual_start_date
|
||||
model.actual_end_date = task.actual_end_date
|
||||
model.tags_json = json.dumps(task.tags, ensure_ascii=False)
|
||||
model.updated_at = task.updated_at
|
||||
|
||||
self._session.commit()
|
||||
return task
|
||||
|
||||
def delete(self, task_id: str) -> None:
|
||||
self._session.query(TaskModel).filter(TaskModel.id == task_id).delete()
|
||||
self._session.commit()
|
||||
|
||||
def _model_to_entity(self, model: TaskModel) -> Task:
|
||||
from packages.domain.project_management import TaskPriority, TaskStatus
|
||||
|
||||
return Task(
|
||||
id=model.id,
|
||||
project_id=model.project_id,
|
||||
name=model.name,
|
||||
description=model.description,
|
||||
status=TaskStatus(model.status),
|
||||
priority=TaskPriority(model.priority),
|
||||
parent_task_id=model.parent_task_id,
|
||||
assignee_user_id=model.assignee_user_id,
|
||||
progress=model.progress,
|
||||
planned_start_date=model.planned_start_date,
|
||||
planned_end_date=model.planned_end_date,
|
||||
actual_start_date=model.actual_start_date,
|
||||
actual_end_date=model.actual_end_date,
|
||||
tags=json.loads(model.tags_json),
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
|
||||
|
||||
class SQLAlchemyMilestoneRepository(MilestoneRepository):
|
||||
"""里程碑 SQLAlchemy 仓储实现"""
|
||||
|
||||
def __init__(self, session: Session):
|
||||
self._session = session
|
||||
|
||||
def create(self, milestone: Milestone) -> Milestone:
|
||||
model = MilestoneModel(
|
||||
id=milestone.id,
|
||||
project_id=milestone.project_id,
|
||||
name=milestone.name,
|
||||
description=milestone.description,
|
||||
target_date=milestone.target_date,
|
||||
completed=milestone.completed,
|
||||
completed_at=milestone.completed_at,
|
||||
created_at=milestone.created_at,
|
||||
updated_at=milestone.updated_at,
|
||||
)
|
||||
self._session.add(model)
|
||||
self._session.commit()
|
||||
return milestone
|
||||
|
||||
def get_by_id(self, milestone_id: str) -> Milestone | None:
|
||||
model = self._session.query(MilestoneModel).filter(MilestoneModel.id == milestone_id).first()
|
||||
if not model:
|
||||
return None
|
||||
return self._model_to_entity(model)
|
||||
|
||||
def list_by_project(self, project_id: str) -> list[Milestone]:
|
||||
models = self._session.query(MilestoneModel).filter(MilestoneModel.project_id == project_id).all()
|
||||
return [self._model_to_entity(m) for m in models]
|
||||
|
||||
def update(self, milestone: Milestone) -> Milestone:
|
||||
model = self._session.query(MilestoneModel).filter(MilestoneModel.id == milestone.id).first()
|
||||
if not model:
|
||||
raise ValueError(f"Milestone {milestone.id} not found")
|
||||
|
||||
model.name = milestone.name
|
||||
model.description = milestone.description
|
||||
model.target_date = milestone.target_date
|
||||
model.completed = milestone.completed
|
||||
model.completed_at = milestone.completed_at
|
||||
model.updated_at = milestone.updated_at
|
||||
|
||||
self._session.commit()
|
||||
return milestone
|
||||
|
||||
def delete(self, milestone_id: str) -> None:
|
||||
self._session.query(MilestoneModel).filter(MilestoneModel.id == milestone_id).delete()
|
||||
self._session.commit()
|
||||
|
||||
def _model_to_entity(self, model: MilestoneModel) -> Milestone:
|
||||
return Milestone(
|
||||
id=model.id,
|
||||
project_id=model.project_id,
|
||||
name=model.name,
|
||||
description=model.description,
|
||||
target_date=model.target_date,
|
||||
completed=model.completed,
|
||||
completed_at=model.completed_at,
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
|
||||
|
||||
class SQLAlchemyTaskIssueRepository(TaskIssueRepository):
|
||||
"""任务问题 SQLAlchemy 仓储实现"""
|
||||
|
||||
def __init__(self, session: Session):
|
||||
self._session = session
|
||||
|
||||
def create(self, issue: TaskIssue) -> TaskIssue:
|
||||
model = TaskIssueModel(
|
||||
id=issue.id,
|
||||
task_id=issue.task_id,
|
||||
project_id=issue.project_id,
|
||||
title=issue.title,
|
||||
description=issue.description,
|
||||
resolved=issue.resolved,
|
||||
resolved_at=issue.resolved_at,
|
||||
created_by_user_id=issue.created_by_user_id,
|
||||
created_at=issue.created_at,
|
||||
updated_at=issue.updated_at,
|
||||
)
|
||||
self._session.add(model)
|
||||
self._session.commit()
|
||||
return issue
|
||||
|
||||
def get_by_id(self, issue_id: str) -> TaskIssue | None:
|
||||
model = self._session.query(TaskIssueModel).filter(TaskIssueModel.id == issue_id).first()
|
||||
if not model:
|
||||
return None
|
||||
return self._model_to_entity(model)
|
||||
|
||||
def list_by_task(self, task_id: str) -> list[TaskIssue]:
|
||||
models = self._session.query(TaskIssueModel).filter(TaskIssueModel.task_id == task_id).all()
|
||||
return [self._model_to_entity(m) for m in models]
|
||||
|
||||
def list_by_project(self, project_id: str) -> list[TaskIssue]:
|
||||
models = self._session.query(TaskIssueModel).filter(TaskIssueModel.project_id == project_id).all()
|
||||
return [self._model_to_entity(m) for m in models]
|
||||
|
||||
def update(self, issue: TaskIssue) -> TaskIssue:
|
||||
model = self._session.query(TaskIssueModel).filter(TaskIssueModel.id == issue.id).first()
|
||||
if not model:
|
||||
raise ValueError(f"TaskIssue {issue.id} not found")
|
||||
|
||||
model.title = issue.title
|
||||
model.description = issue.description
|
||||
model.resolved = issue.resolved
|
||||
model.resolved_at = issue.resolved_at
|
||||
model.updated_at = issue.updated_at
|
||||
|
||||
self._session.commit()
|
||||
return issue
|
||||
|
||||
def delete(self, issue_id: str) -> None:
|
||||
self._session.query(TaskIssueModel).filter(TaskIssueModel.id == issue_id).delete()
|
||||
self._session.commit()
|
||||
|
||||
def _model_to_entity(self, model: TaskIssueModel) -> TaskIssue:
|
||||
return TaskIssue(
|
||||
id=model.id,
|
||||
task_id=model.task_id,
|
||||
project_id=model.project_id,
|
||||
title=model.title,
|
||||
description=model.description,
|
||||
resolved=model.resolved,
|
||||
resolved_at=model.resolved_at,
|
||||
created_by_user_id=model.created_by_user_id,
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
@@ -1,52 +0,0 @@
|
||||
from datetime import datetime, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import ProjectTitleModel
|
||||
|
||||
|
||||
class SQLAlchemyProjectTitleRepository:
|
||||
def __init__(self, session):
|
||||
self.session = session
|
||||
|
||||
def list_by_project(self, project_id: str, active_only: bool = False) -> list[ProjectTitleModel]:
|
||||
query = self.session.query(ProjectTitleModel).filter(ProjectTitleModel.project_id == project_id)
|
||||
if active_only:
|
||||
query = query.filter(ProjectTitleModel.is_active.is_(True))
|
||||
return query.order_by(ProjectTitleModel.created_at.desc()).all()
|
||||
|
||||
def get(self, title_id: str) -> ProjectTitleModel | None:
|
||||
return self.session.query(ProjectTitleModel).filter(ProjectTitleModel.id == title_id).first()
|
||||
|
||||
def create(
|
||||
self,
|
||||
*,
|
||||
project_id: str,
|
||||
text: str,
|
||||
category: str,
|
||||
created_by_user_id: str,
|
||||
favorite: bool = False,
|
||||
) -> ProjectTitleModel:
|
||||
now = datetime.now(timezone.utc)
|
||||
item = ProjectTitleModel(
|
||||
id=uuid4().hex,
|
||||
project_id=project_id,
|
||||
text=text.strip(),
|
||||
category=category,
|
||||
favorite=favorite,
|
||||
usage_count=0,
|
||||
is_active=True,
|
||||
created_by_user_id=created_by_user_id,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
self.session.add(item)
|
||||
self.session.commit()
|
||||
self.session.refresh(item)
|
||||
return item
|
||||
|
||||
def update(self, item: ProjectTitleModel) -> ProjectTitleModel:
|
||||
item.updated_at = datetime.now(timezone.utc)
|
||||
self.session.add(item)
|
||||
self.session.commit()
|
||||
self.session.refresh(item)
|
||||
return item
|
||||
@@ -0,0 +1,114 @@
|
||||
"""SQLAlchemy implementation of TitleLibraryRepository."""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import List, Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import TitleLibraryModel
|
||||
from packages.domain.title_library import TitleLibraryItem
|
||||
|
||||
|
||||
class SQLAlchemyTitleLibraryRepository:
|
||||
"""SQLAlchemy 标题库仓储"""
|
||||
|
||||
def __init__(self, session: Session) -> None:
|
||||
self.session = session
|
||||
|
||||
def list_by_user(
|
||||
self,
|
||||
user_id: str,
|
||||
*,
|
||||
category: Optional[str] = None,
|
||||
is_active: bool = True,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
) -> List[TitleLibraryItem]:
|
||||
query = self.session.query(TitleLibraryModel).filter(
|
||||
TitleLibraryModel.user_id == user_id,
|
||||
TitleLibraryModel.is_active == is_active,
|
||||
)
|
||||
if category:
|
||||
query = query.filter(TitleLibraryModel.category == category)
|
||||
query = query.order_by(TitleLibraryModel.created_at.desc())
|
||||
models = query.offset(skip).limit(limit).all()
|
||||
return [self._model_to_entity(m) for m in models]
|
||||
|
||||
def get(self, title_id: str, user_id: str) -> Optional[TitleLibraryItem]:
|
||||
model = self.session.query(TitleLibraryModel).filter(
|
||||
TitleLibraryModel.id == title_id,
|
||||
TitleLibraryModel.user_id == user_id,
|
||||
).first()
|
||||
if model is None:
|
||||
return None
|
||||
return self._model_to_entity(model)
|
||||
|
||||
def create(self, item: TitleLibraryItem) -> TitleLibraryItem:
|
||||
model = TitleLibraryModel(
|
||||
id=item.id,
|
||||
user_id=item.user_id,
|
||||
name=item.name,
|
||||
description=item.description,
|
||||
category=item.category,
|
||||
text=item.text,
|
||||
tags=item.tags,
|
||||
usage_count=item.usage_count,
|
||||
is_active=item.is_active,
|
||||
metadata=item.metadata_,
|
||||
)
|
||||
self.session.add(model)
|
||||
self.session.commit()
|
||||
self.session.refresh(model)
|
||||
return self._model_to_entity(model)
|
||||
|
||||
def update(self, item: TitleLibraryItem) -> TitleLibraryItem:
|
||||
model = self.session.query(TitleLibraryModel).filter(
|
||||
TitleLibraryModel.id == item.id,
|
||||
TitleLibraryModel.user_id == item.user_id,
|
||||
).first()
|
||||
if model is None:
|
||||
raise ValueError(f"TitleLibraryItem {item.id} not found")
|
||||
model.name = item.name
|
||||
model.description = item.description
|
||||
model.category = item.category
|
||||
model.text = item.text
|
||||
model.tags = item.tags
|
||||
model.is_active = item.is_active
|
||||
model.metadata = item.metadata_
|
||||
self.session.commit()
|
||||
self.session.refresh(model)
|
||||
return self._model_to_entity(model)
|
||||
|
||||
def delete(self, title_id: str, user_id: str) -> bool:
|
||||
model = self.session.query(TitleLibraryModel).filter(
|
||||
TitleLibraryModel.id == title_id,
|
||||
TitleLibraryModel.user_id == user_id,
|
||||
).first()
|
||||
if model is None:
|
||||
return False
|
||||
model.is_active = False
|
||||
self.session.commit()
|
||||
return True
|
||||
|
||||
def count_by_user(self, user_id: str, is_active: bool = True) -> int:
|
||||
return self.session.query(TitleLibraryModel).filter(
|
||||
TitleLibraryModel.user_id == user_id,
|
||||
TitleLibraryModel.is_active == is_active,
|
||||
).count()
|
||||
|
||||
@staticmethod
|
||||
def _model_to_entity(model: TitleLibraryModel) -> TitleLibraryItem:
|
||||
return TitleLibraryItem(
|
||||
id=model.id,
|
||||
user_id=model.user_id,
|
||||
name=model.name,
|
||||
description=model.description,
|
||||
category=model.category,
|
||||
text=model.text,
|
||||
tags=model.tags or [],
|
||||
usage_count=model.usage_count or 0,
|
||||
is_active=model.is_active,
|
||||
metadata_=model.metadata or {},
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
@@ -0,0 +1,125 @@
|
||||
"""SQLAlchemy implementation of VoiceLibraryRepository."""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import List, Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import VoiceLibraryModel
|
||||
from packages.domain.voice_library import VoiceLibraryItem
|
||||
|
||||
|
||||
class SQLAlchemyVoiceLibraryRepository:
|
||||
"""SQLAlchemy 配音库仓储"""
|
||||
|
||||
def __init__(self, session: Session) -> None:
|
||||
self.session = session
|
||||
|
||||
def list_by_user(
|
||||
self,
|
||||
user_id: str,
|
||||
*,
|
||||
status: Optional[str] = None,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
) -> List[VoiceLibraryItem]:
|
||||
query = self.session.query(VoiceLibraryModel).filter(
|
||||
VoiceLibraryModel.user_id == user_id,
|
||||
)
|
||||
if status:
|
||||
query = query.filter(VoiceLibraryModel.status == status)
|
||||
query = query.order_by(VoiceLibraryModel.created_at.desc())
|
||||
models = query.offset(skip).limit(limit).all()
|
||||
return [self._model_to_entity(m) for m in models]
|
||||
|
||||
def get(self, voice_id: str, user_id: str) -> Optional[VoiceLibraryItem]:
|
||||
model = self.session.query(VoiceLibraryModel).filter(
|
||||
VoiceLibraryModel.id == voice_id,
|
||||
VoiceLibraryModel.user_id == user_id,
|
||||
).first()
|
||||
if model is None:
|
||||
return None
|
||||
return self._model_to_entity(model)
|
||||
|
||||
def create(self, item: VoiceLibraryItem) -> VoiceLibraryItem:
|
||||
model = VoiceLibraryModel(
|
||||
id=item.id,
|
||||
user_id=item.user_id,
|
||||
project_id=item.project_id or "",
|
||||
name=item.name,
|
||||
text=item.text,
|
||||
voice_provider=item.voice_provider,
|
||||
voice_id=item.voice_id,
|
||||
voice_name=item.voice_name,
|
||||
audio_url=item.audio_url,
|
||||
duration=item.duration,
|
||||
file_size=item.file_size,
|
||||
status=item.status,
|
||||
tags=item.tags,
|
||||
metadata=item.metadata_,
|
||||
)
|
||||
self.session.add(model)
|
||||
self.session.commit()
|
||||
self.session.refresh(model)
|
||||
return self._model_to_entity(model)
|
||||
|
||||
def update(self, item: VoiceLibraryItem) -> VoiceLibraryItem:
|
||||
model = self.session.query(VoiceLibraryModel).filter(
|
||||
VoiceLibraryModel.id == item.id,
|
||||
VoiceLibraryModel.user_id == item.user_id,
|
||||
).first()
|
||||
if model is None:
|
||||
raise ValueError(f"VoiceLibraryItem {item.id} not found")
|
||||
model.name = item.name
|
||||
model.text = item.text
|
||||
model.voice_provider = item.voice_provider
|
||||
model.voice_id = item.voice_id
|
||||
model.voice_name = item.voice_name
|
||||
model.audio_url = item.audio_url
|
||||
model.duration = item.duration
|
||||
model.file_size = item.file_size
|
||||
model.status = item.status
|
||||
model.tags = item.tags
|
||||
model.metadata = item.metadata_
|
||||
self.session.commit()
|
||||
self.session.refresh(model)
|
||||
return self._model_to_entity(model)
|
||||
|
||||
def delete(self, voice_id: str, user_id: str) -> bool:
|
||||
model = self.session.query(VoiceLibraryModel).filter(
|
||||
VoiceLibraryModel.id == voice_id,
|
||||
VoiceLibraryModel.user_id == user_id,
|
||||
).first()
|
||||
if model is None:
|
||||
return False
|
||||
# Soft delete by setting status to deleted
|
||||
model.status = "deleted"
|
||||
self.session.commit()
|
||||
return True
|
||||
|
||||
def count_by_user(self, user_id: str) -> int:
|
||||
return self.session.query(VoiceLibraryModel).filter(
|
||||
VoiceLibraryModel.user_id == user_id,
|
||||
VoiceLibraryModel.status != "deleted",
|
||||
).count()
|
||||
|
||||
@staticmethod
|
||||
def _model_to_entity(model: VoiceLibraryModel) -> VoiceLibraryItem:
|
||||
return VoiceLibraryItem(
|
||||
id=model.id,
|
||||
user_id=model.user_id,
|
||||
name=model.name,
|
||||
text=model.text,
|
||||
voice_provider=model.voice_provider,
|
||||
voice_id=model.voice_id,
|
||||
voice_name=model.voice_name,
|
||||
audio_url=model.audio_url,
|
||||
duration=model.duration or 0,
|
||||
file_size=model.file_size or 0,
|
||||
status=model.status,
|
||||
project_id=model.project_id if model.project_id else None,
|
||||
tags=model.tags or [],
|
||||
metadata_=model.metadata or {},
|
||||
created_at=model.created_at,
|
||||
updated_at=model.updated_at,
|
||||
)
|
||||
@@ -1,10 +1,5 @@
|
||||
"""SQLite Tracker Adapter"""
|
||||
|
||||
from .project_management_repositories import (
|
||||
SQLiteMilestoneRepository,
|
||||
SQLiteTaskIssueRepository,
|
||||
SQLiteTaskRepository,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"SQLiteTaskRepository",
|
||||
|
||||
@@ -1,201 +0,0 @@
|
||||
"""SQLite 实现的项目管理 Repository"""
|
||||
|
||||
import sqlite3
|
||||
from datetime import datetime
|
||||
from typing import List, Optional
|
||||
|
||||
from packages.domain.project_management import (
|
||||
Milestone,
|
||||
Task,
|
||||
TaskIssue,
|
||||
TaskPriority,
|
||||
TaskStatus,
|
||||
)
|
||||
|
||||
DB_PATH = "tracker.db"
|
||||
|
||||
|
||||
class SQLiteTaskRepository:
|
||||
"""基于 SQLite 的任务仓储"""
|
||||
|
||||
def get_by_id(self, task_id: str) -> Optional[Task]:
|
||||
conn = sqlite3.connect(DB_PATH)
|
||||
conn.row_factory = sqlite3.Row
|
||||
cursor = conn.cursor()
|
||||
|
||||
cursor.execute("SELECT * FROM tasks WHERE id = ?", (task_id,))
|
||||
row = cursor.fetchone()
|
||||
conn.close()
|
||||
|
||||
if not row:
|
||||
return None
|
||||
|
||||
return Task(
|
||||
id=str(row["id"]),
|
||||
name=row["name"],
|
||||
description=row["description"] or "",
|
||||
status=TaskStatus(row["status"]) if row["status"] else TaskStatus.PENDING,
|
||||
priority=(TaskPriority(row["priority"]) if row["priority"] else TaskPriority.MEDIUM),
|
||||
progress=0, # tracker.db 没有 progress 字段
|
||||
project_id=row["phase"] or "xiaoxia-saas",
|
||||
assignee_user_id=row["assigned_to"] or "",
|
||||
created_at=(datetime.fromisoformat(row["created_at"]) if row["created_at"] else datetime.now()),
|
||||
updated_at=datetime.now(),
|
||||
)
|
||||
|
||||
def list_by_project(self, project_id: str, skip: int = 0, limit: int = 100) -> List[Task]:
|
||||
conn = sqlite3.connect(DB_PATH)
|
||||
conn.row_factory = sqlite3.Row
|
||||
cursor = conn.cursor()
|
||||
|
||||
# 返回所有任务(忽略 project_id 过滤,因为 tracker.db 使用 phase)
|
||||
cursor.execute(
|
||||
"""
|
||||
SELECT * FROM tasks
|
||||
ORDER BY created_at DESC
|
||||
LIMIT ? OFFSET ?
|
||||
""",
|
||||
(limit, skip),
|
||||
)
|
||||
|
||||
rows = cursor.fetchall()
|
||||
conn.close()
|
||||
|
||||
tasks = []
|
||||
for row in rows:
|
||||
tasks.append(
|
||||
Task(
|
||||
id=str(row["id"]),
|
||||
name=row["name"],
|
||||
description=row["description"] or "",
|
||||
status=(TaskStatus(row["status"]) if row["status"] else TaskStatus.PENDING),
|
||||
priority=(TaskPriority(row["priority"]) if row["priority"] else TaskPriority.MEDIUM),
|
||||
progress=0,
|
||||
project_id=row["phase"] or "xiaoxia-saas",
|
||||
assignee_user_id=row["assigned_to"] or "",
|
||||
created_at=(datetime.fromisoformat(row["created_at"]) if row["created_at"] else datetime.now()),
|
||||
updated_at=datetime.now(),
|
||||
)
|
||||
)
|
||||
|
||||
return tasks
|
||||
|
||||
def save(self, task: Task) -> Task:
|
||||
conn = sqlite3.connect(DB_PATH)
|
||||
cursor = conn.cursor()
|
||||
|
||||
if task.id and task.id.isdigit():
|
||||
# 更新现有任务
|
||||
cursor.execute(
|
||||
"""
|
||||
UPDATE tasks
|
||||
SET name = ?, description = ?, status = ?, priority = ?, assigned_to = ?
|
||||
WHERE id = ?
|
||||
""",
|
||||
(
|
||||
task.name,
|
||||
task.description,
|
||||
task.status.value,
|
||||
task.priority.value,
|
||||
task.assignee_user_id,
|
||||
task.id,
|
||||
),
|
||||
)
|
||||
else:
|
||||
# 创建新任务
|
||||
cursor.execute(
|
||||
"""
|
||||
INSERT INTO tasks (name, description, status, phase, priority, assigned_to, created_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
task.name,
|
||||
task.description,
|
||||
task.status.value,
|
||||
task.project_id,
|
||||
task.priority.value,
|
||||
task.assignee_user_id,
|
||||
datetime.now().isoformat(),
|
||||
),
|
||||
)
|
||||
task.id = str(cursor.lastrowid)
|
||||
|
||||
conn.commit()
|
||||
conn.close()
|
||||
return task
|
||||
|
||||
|
||||
class SQLiteMilestoneRepository:
|
||||
"""基于 SQLite 的里程碑仓储"""
|
||||
|
||||
def list_by_project(self, project_id: str) -> List[Milestone]:
|
||||
conn = sqlite3.connect(DB_PATH)
|
||||
conn.row_factory = sqlite3.Row
|
||||
cursor = conn.cursor()
|
||||
|
||||
cursor.execute("SELECT * FROM milestones ORDER BY start_date")
|
||||
rows = cursor.fetchall()
|
||||
conn.close()
|
||||
|
||||
milestones = []
|
||||
for row in rows:
|
||||
milestones.append(
|
||||
Milestone(
|
||||
id=str(row["id"]),
|
||||
name=row["name"],
|
||||
description=row["description"] or "",
|
||||
target_date=row["end_date"] or "",
|
||||
project_id=row["phase"] or "xiaoxia-saas",
|
||||
created_at=(datetime.fromisoformat(row["start_date"]) if row["start_date"] else datetime.now()),
|
||||
)
|
||||
)
|
||||
|
||||
return milestones
|
||||
|
||||
def save(self, milestone: Milestone) -> Milestone:
|
||||
conn = sqlite3.connect(DB_PATH)
|
||||
cursor = conn.cursor()
|
||||
|
||||
if milestone.id and milestone.id.isdigit():
|
||||
cursor.execute(
|
||||
"""
|
||||
UPDATE milestones
|
||||
SET name = ?, description = ?, end_date = ?
|
||||
WHERE id = ?
|
||||
""",
|
||||
(
|
||||
milestone.name,
|
||||
milestone.description,
|
||||
milestone.target_date,
|
||||
milestone.id,
|
||||
),
|
||||
)
|
||||
else:
|
||||
cursor.execute(
|
||||
"""
|
||||
INSERT INTO milestones (name, description, phase, start_date, end_date)
|
||||
VALUES (?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
milestone.name,
|
||||
milestone.description,
|
||||
milestone.project_id,
|
||||
datetime.now().isoformat(),
|
||||
milestone.target_date,
|
||||
),
|
||||
)
|
||||
milestone.id = str(cursor.lastrowid)
|
||||
|
||||
conn.commit()
|
||||
conn.close()
|
||||
return milestone
|
||||
|
||||
|
||||
class SQLiteTaskIssueRepository:
|
||||
"""空实现 - tracker.db 没有 issues 表"""
|
||||
|
||||
def list_by_task(self, task_id: str) -> List[TaskIssue]:
|
||||
return []
|
||||
|
||||
def save(self, issue: TaskIssue) -> TaskIssue:
|
||||
return issue
|
||||
@@ -1,17 +0,0 @@
|
||||
"""获取单个任务详情用例"""
|
||||
|
||||
from packages.domain import Task
|
||||
from packages.ports import TaskRepository
|
||||
|
||||
|
||||
class GetTaskDetailUseCase:
|
||||
"""获取任务详情用例"""
|
||||
|
||||
def __init__(self, task_repo: TaskRepository):
|
||||
self.task_repo = task_repo
|
||||
|
||||
def execute(self, task_id: str) -> Task:
|
||||
task = self.task_repo.get_by_id(task_id)
|
||||
if not task:
|
||||
raise ValueError(f"Task {task_id} not found")
|
||||
return task
|
||||
@@ -1,146 +0,0 @@
|
||||
"""项目管理 Use Cases"""
|
||||
|
||||
from packages.domain import Milestone, Task, TaskIssue, TaskPriority, TaskStatus
|
||||
from packages.ports import MilestoneRepository, TaskIssueRepository, TaskRepository
|
||||
|
||||
|
||||
class CreateTaskUseCase:
|
||||
"""创建任务用例"""
|
||||
|
||||
def __init__(self, task_repo: TaskRepository):
|
||||
self.task_repo = task_repo
|
||||
|
||||
def execute(
|
||||
self,
|
||||
project_id: str,
|
||||
name: str,
|
||||
description: str = "",
|
||||
priority: TaskPriority = TaskPriority.MEDIUM,
|
||||
parent_task_id: str = "",
|
||||
assignee_user_id: str = "",
|
||||
) -> Task:
|
||||
task = Task.create(
|
||||
project_id=project_id,
|
||||
name=name,
|
||||
description=description,
|
||||
priority=priority,
|
||||
parent_task_id=parent_task_id,
|
||||
assignee_user_id=assignee_user_id,
|
||||
)
|
||||
return self.task_repo.create(task)
|
||||
|
||||
|
||||
class ListProjectTasksUseCase:
|
||||
"""获取项目任务列表用例"""
|
||||
|
||||
def __init__(self, task_repo: TaskRepository):
|
||||
self.task_repo = task_repo
|
||||
|
||||
def execute(self, project_id: str) -> list[Task]:
|
||||
return self.task_repo.list_by_project(project_id)
|
||||
|
||||
|
||||
class UpdateTaskStatusUseCase:
|
||||
"""更新任务状态用例"""
|
||||
|
||||
def __init__(self, task_repo: TaskRepository):
|
||||
self.task_repo = task_repo
|
||||
|
||||
def execute(self, task_id: str, new_status: TaskStatus) -> Task:
|
||||
task = self.task_repo.get_by_id(task_id)
|
||||
if not task:
|
||||
raise ValueError(f"Task {task_id} not found")
|
||||
task.update_status(new_status)
|
||||
return self.task_repo.update(task)
|
||||
|
||||
|
||||
class UpdateTaskProgressUseCase:
|
||||
"""更新任务进度用例"""
|
||||
|
||||
def __init__(self, task_repo: TaskRepository):
|
||||
self.task_repo = task_repo
|
||||
|
||||
def execute(self, task_id: str, progress: float) -> Task:
|
||||
task = self.task_repo.get_by_id(task_id)
|
||||
if not task:
|
||||
raise ValueError(f"Task {task_id} not found")
|
||||
task.update_progress(progress)
|
||||
return self.task_repo.update(task)
|
||||
|
||||
|
||||
class CreateMilestoneUseCase:
|
||||
"""创建里程碑用例"""
|
||||
|
||||
def __init__(self, milestone_repo: MilestoneRepository):
|
||||
self.milestone_repo = milestone_repo
|
||||
|
||||
def execute(
|
||||
self,
|
||||
project_id: str,
|
||||
name: str,
|
||||
description: str = "",
|
||||
) -> Milestone:
|
||||
milestone = Milestone.create(
|
||||
project_id=project_id,
|
||||
name=name,
|
||||
description=description,
|
||||
)
|
||||
return self.milestone_repo.create(milestone)
|
||||
|
||||
|
||||
class ListProjectMilestonesUseCase:
|
||||
"""获取项目里程碑列表用例"""
|
||||
|
||||
def __init__(self, milestone_repo: MilestoneRepository):
|
||||
self.milestone_repo = milestone_repo
|
||||
|
||||
def execute(self, project_id: str) -> list[Milestone]:
|
||||
return self.milestone_repo.list_by_project(project_id)
|
||||
|
||||
|
||||
class CreateTaskIssueUseCase:
|
||||
"""创建任务问题用例"""
|
||||
|
||||
def __init__(self, issue_repo: TaskIssueRepository):
|
||||
self.issue_repo = issue_repo
|
||||
|
||||
def execute(
|
||||
self,
|
||||
task_id: str,
|
||||
project_id: str,
|
||||
title: str,
|
||||
description: str = "",
|
||||
created_by_user_id: str = "",
|
||||
) -> TaskIssue:
|
||||
issue = TaskIssue.create(
|
||||
task_id=task_id,
|
||||
project_id=project_id,
|
||||
title=title,
|
||||
description=description,
|
||||
created_by_user_id=created_by_user_id,
|
||||
)
|
||||
return self.issue_repo.create(issue)
|
||||
|
||||
|
||||
class ListTaskIssuesUseCase:
|
||||
"""获取任务问题列表用例"""
|
||||
|
||||
def __init__(self, issue_repo: TaskIssueRepository):
|
||||
self.issue_repo = issue_repo
|
||||
|
||||
def execute(self, task_id: str) -> list[TaskIssue]:
|
||||
return self.issue_repo.list_by_task(task_id)
|
||||
|
||||
|
||||
class ResolveTaskIssueUseCase:
|
||||
"""解决任务问题用例"""
|
||||
|
||||
def __init__(self, issue_repo: TaskIssueRepository):
|
||||
self.issue_repo = issue_repo
|
||||
|
||||
def execute(self, issue_id: str) -> TaskIssue:
|
||||
issue = self.issue_repo.get_by_id(issue_id)
|
||||
if not issue:
|
||||
raise ValueError(f"TaskIssue {issue_id} not found")
|
||||
issue.mark_resolved()
|
||||
return self.issue_repo.update(issue)
|
||||
@@ -0,0 +1,20 @@
|
||||
"""Title library application module."""
|
||||
from packages.application.title_library.use_cases import (
|
||||
CreateTitleLibraryUseCase,
|
||||
DeleteTitleLibraryUseCase,
|
||||
GetTitleLibraryUseCase,
|
||||
ListTitleLibraryUseCase,
|
||||
UpdateTitleLibraryUseCase,
|
||||
QuotaExceededError,
|
||||
NotFoundError,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"CreateTitleLibraryUseCase",
|
||||
"DeleteTitleLibraryUseCase",
|
||||
"GetTitleLibraryUseCase",
|
||||
"ListTitleLibraryUseCase",
|
||||
"UpdateTitleLibraryUseCase",
|
||||
"QuotaExceededError",
|
||||
"NotFoundError",
|
||||
]
|
||||
@@ -0,0 +1,29 @@
|
||||
"""Title library commands."""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Optional
|
||||
|
||||
|
||||
@dataclass
|
||||
class CreateTitleLibraryCommand:
|
||||
user_id: str
|
||||
name: str
|
||||
text: str
|
||||
category: str = "default"
|
||||
description: str = ""
|
||||
tags: List[str] = field(default_factory=list)
|
||||
metadata_: dict = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class UpdateTitleLibraryCommand:
|
||||
title_id: str
|
||||
user_id: str
|
||||
name: Optional[str] = None
|
||||
text: Optional[str] = None
|
||||
category: Optional[str] = None
|
||||
description: Optional[str] = None
|
||||
tags: Optional[List[str]] = None
|
||||
is_active: Optional[bool] = None
|
||||
metadata_: Optional[dict] = None
|
||||
@@ -0,0 +1,111 @@
|
||||
"""Title library use cases."""
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from typing import List, Optional
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.title_library_repository import SQLAlchemyTitleLibraryRepository
|
||||
from packages.application.title_library.commands import (
|
||||
CreateTitleLibraryCommand,
|
||||
UpdateTitleLibraryCommand,
|
||||
)
|
||||
from packages.domain.quota import QuotaDimension, quota_checker
|
||||
from packages.domain.title_library import TitleLibraryItem
|
||||
|
||||
|
||||
class ListTitleLibraryUseCase:
|
||||
def __init__(self, repository: SQLAlchemyTitleLibraryRepository) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def execute(
|
||||
self,
|
||||
user_id: str,
|
||||
*,
|
||||
category: Optional[str] = None,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
) -> List[TitleLibraryItem]:
|
||||
return self.repository.list_by_user(user_id, category=category, skip=skip, limit=limit)
|
||||
|
||||
|
||||
class GetTitleLibraryUseCase:
|
||||
def __init__(self, repository: SQLAlchemyTitleLibraryRepository) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def execute(self, title_id: str, user_id: str) -> Optional[TitleLibraryItem]:
|
||||
return self.repository.get(title_id, user_id)
|
||||
|
||||
|
||||
class CreateTitleLibraryUseCase:
|
||||
def __init__(self, repository: SQLAlchemyTitleLibraryRepository) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def execute(self, command: CreateTitleLibraryCommand, plan_name: str = "free") -> TitleLibraryItem:
|
||||
# Quota check
|
||||
current_count = self.repository.count_by_user(command.user_id)
|
||||
result = quota_checker.check(plan_name, QuotaDimension.MAX_TITLES.value, current_count)
|
||||
if not result.allowed:
|
||||
raise QuotaExceededError(
|
||||
dimension=QuotaDimension.MAX_TITLES.value,
|
||||
limit=result.limit,
|
||||
used=result.used,
|
||||
)
|
||||
|
||||
item = TitleLibraryItem(
|
||||
id=uuid.uuid4().hex,
|
||||
user_id=command.user_id,
|
||||
name=command.name,
|
||||
text=command.text,
|
||||
category=command.category,
|
||||
description=command.description,
|
||||
tags=command.tags,
|
||||
metadata_=command.metadata_,
|
||||
)
|
||||
return self.repository.create(item)
|
||||
|
||||
|
||||
class UpdateTitleLibraryUseCase:
|
||||
def __init__(self, repository: SQLAlchemyTitleLibraryRepository) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def execute(self, command: UpdateTitleLibraryCommand) -> TitleLibraryItem:
|
||||
existing = self.repository.get(command.title_id, command.user_id)
|
||||
if existing is None:
|
||||
raise NotFoundError(f"Title {command.title_id} not found")
|
||||
|
||||
if command.name is not None:
|
||||
existing.name = command.name
|
||||
if command.text is not None:
|
||||
existing.text = command.text
|
||||
if command.category is not None:
|
||||
existing.category = command.category
|
||||
if command.description is not None:
|
||||
existing.description = command.description
|
||||
if command.tags is not None:
|
||||
existing.tags = command.tags
|
||||
if command.is_active is not None:
|
||||
existing.is_active = command.is_active
|
||||
if command.metadata_ is not None:
|
||||
existing.metadata_ = command.metadata_
|
||||
|
||||
return self.repository.update(existing)
|
||||
|
||||
|
||||
class DeleteTitleLibraryUseCase:
|
||||
def __init__(self, repository: SQLAlchemyTitleLibraryRepository) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def execute(self, title_id: str, user_id: str) -> bool:
|
||||
return self.repository.delete(title_id, user_id)
|
||||
|
||||
|
||||
class QuotaExceededError(Exception):
|
||||
def __init__(self, dimension: str, limit: float, used: float) -> None:
|
||||
self.dimension = dimension
|
||||
self.limit = limit
|
||||
self.used = used
|
||||
super().__init__(f"Quota exceeded for {dimension}: {used}/{limit}")
|
||||
|
||||
|
||||
class NotFoundError(Exception):
|
||||
pass
|
||||
@@ -1,35 +0,0 @@
|
||||
"""更新任务基本信息用例"""
|
||||
|
||||
from packages.domain import Task
|
||||
from packages.ports import TaskRepository
|
||||
|
||||
|
||||
class UpdateTaskUseCase:
|
||||
"""更新任务基本信息"""
|
||||
|
||||
def __init__(self, task_repo: TaskRepository):
|
||||
self.task_repo = task_repo
|
||||
|
||||
def execute(
|
||||
self,
|
||||
task_id: str,
|
||||
name: str | None = None,
|
||||
description: str | None = None,
|
||||
priority: str | None = None,
|
||||
assignee_user_id: str | None = None,
|
||||
) -> Task:
|
||||
task = self.task_repo.get_by_id(task_id)
|
||||
if not task:
|
||||
raise ValueError(f"Task {task_id} not found")
|
||||
|
||||
if name is not None:
|
||||
task.name = name
|
||||
if description is not None:
|
||||
task.description = description
|
||||
if priority is not None:
|
||||
task.priority = priority
|
||||
if assignee_user_id is not None:
|
||||
task.assignee_user_id = assignee_user_id
|
||||
|
||||
self.task_repo.update(task)
|
||||
return task
|
||||
@@ -0,0 +1,20 @@
|
||||
"""Voice library application module."""
|
||||
from packages.application.voice_library.use_cases import (
|
||||
CreateVoiceLibraryUseCase,
|
||||
DeleteVoiceLibraryUseCase,
|
||||
GetVoiceLibraryUseCase,
|
||||
ListVoiceLibraryUseCase,
|
||||
UpdateVoiceLibraryUseCase,
|
||||
QuotaExceededError,
|
||||
NotFoundError,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"CreateVoiceLibraryUseCase",
|
||||
"DeleteVoiceLibraryUseCase",
|
||||
"GetVoiceLibraryUseCase",
|
||||
"ListVoiceLibraryUseCase",
|
||||
"UpdateVoiceLibraryUseCase",
|
||||
"QuotaExceededError",
|
||||
"NotFoundError",
|
||||
]
|
||||
@@ -0,0 +1,39 @@
|
||||
"""Voice library commands."""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import List, Optional
|
||||
|
||||
|
||||
@dataclass
|
||||
class CreateVoiceLibraryCommand:
|
||||
user_id: str
|
||||
name: str
|
||||
text: str = ""
|
||||
voice_provider: str = ""
|
||||
voice_id: str = ""
|
||||
voice_name: str = ""
|
||||
audio_url: str = ""
|
||||
duration: float = 0
|
||||
file_size: int = 0
|
||||
status: str = "completed"
|
||||
project_id: Optional[str] = None
|
||||
tags: List[str] = field(default_factory=list)
|
||||
metadata_: dict = field(default_factory=dict)
|
||||
|
||||
|
||||
@dataclass
|
||||
class UpdateVoiceLibraryCommand:
|
||||
id: str # ID of the voice library item to update
|
||||
user_id: str
|
||||
name: Optional[str] = None
|
||||
text: Optional[str] = None
|
||||
voice_provider: Optional[str] = None
|
||||
voice_id: Optional[str] = None
|
||||
voice_name: Optional[str] = None
|
||||
audio_url: Optional[str] = None
|
||||
duration: Optional[float] = None
|
||||
file_size: Optional[int] = None
|
||||
status: Optional[str] = None
|
||||
tags: Optional[List[str]] = None
|
||||
metadata_: Optional[dict] = None
|
||||
@@ -0,0 +1,125 @@
|
||||
"""Voice library use cases."""
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from typing import List, Optional
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.voice_library_repository import SQLAlchemyVoiceLibraryRepository
|
||||
from packages.application.voice_library.commands import (
|
||||
CreateVoiceLibraryCommand,
|
||||
UpdateVoiceLibraryCommand,
|
||||
)
|
||||
from packages.domain.quota import QuotaDimension, quota_checker
|
||||
from packages.domain.voice_library import VoiceLibraryItem
|
||||
|
||||
|
||||
class ListVoiceLibraryUseCase:
|
||||
def __init__(self, repository: SQLAlchemyVoiceLibraryRepository) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def execute(
|
||||
self,
|
||||
user_id: str,
|
||||
*,
|
||||
status: Optional[str] = None,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
) -> List[VoiceLibraryItem]:
|
||||
return self.repository.list_by_user(user_id, status=status, skip=skip, limit=limit)
|
||||
|
||||
|
||||
class GetVoiceLibraryUseCase:
|
||||
def __init__(self, repository: SQLAlchemyVoiceLibraryRepository) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def execute(self, voice_id: str, user_id: str) -> Optional[VoiceLibraryItem]:
|
||||
return self.repository.get(voice_id, user_id)
|
||||
|
||||
|
||||
class CreateVoiceLibraryUseCase:
|
||||
def __init__(self, repository: SQLAlchemyVoiceLibraryRepository) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def execute(self, command: CreateVoiceLibraryCommand, plan_name: str = "free") -> VoiceLibraryItem:
|
||||
# Quota check
|
||||
current_count = self.repository.count_by_user(command.user_id)
|
||||
result = quota_checker.check(plan_name, QuotaDimension.MAX_VOICEOVERS.value, current_count)
|
||||
if not result.allowed:
|
||||
raise QuotaExceededError(
|
||||
dimension=QuotaDimension.MAX_VOICEOVERS.value,
|
||||
limit=result.limit,
|
||||
used=result.used,
|
||||
)
|
||||
|
||||
item = VoiceLibraryItem(
|
||||
id=uuid.uuid4().hex,
|
||||
user_id=command.user_id,
|
||||
name=command.name,
|
||||
text=command.text,
|
||||
voice_provider=command.voice_provider,
|
||||
voice_id=command.voice_id,
|
||||
voice_name=command.voice_name,
|
||||
audio_url=command.audio_url,
|
||||
duration=command.duration,
|
||||
file_size=command.file_size,
|
||||
status=command.status,
|
||||
project_id=command.project_id,
|
||||
tags=command.tags,
|
||||
metadata_=command.metadata_,
|
||||
)
|
||||
return self.repository.create(item)
|
||||
|
||||
|
||||
class UpdateVoiceLibraryUseCase:
|
||||
def __init__(self, repository: SQLAlchemyVoiceLibraryRepository) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def execute(self, command: UpdateVoiceLibraryCommand) -> VoiceLibraryItem:
|
||||
existing = self.repository.get(command.id, command.user_id)
|
||||
if existing is None:
|
||||
raise NotFoundError(f"Voice {command.id} not found")
|
||||
|
||||
if command.name is not None:
|
||||
existing.name = command.name
|
||||
if command.text is not None:
|
||||
existing.text = command.text
|
||||
if command.voice_provider is not None:
|
||||
existing.voice_provider = command.voice_provider
|
||||
if command.voice_id is not None:
|
||||
existing.voice_id = command.voice_id
|
||||
if command.voice_name is not None:
|
||||
existing.voice_name = command.voice_name
|
||||
if command.audio_url is not None:
|
||||
existing.audio_url = command.audio_url
|
||||
if command.duration is not None:
|
||||
existing.duration = command.duration
|
||||
if command.file_size is not None:
|
||||
existing.file_size = command.file_size
|
||||
if command.status is not None:
|
||||
existing.status = command.status
|
||||
if command.tags is not None:
|
||||
existing.tags = command.tags
|
||||
if command.metadata_ is not None:
|
||||
existing.metadata_ = command.metadata_
|
||||
|
||||
return self.repository.update(existing)
|
||||
|
||||
|
||||
class DeleteVoiceLibraryUseCase:
|
||||
def __init__(self, repository: SQLAlchemyVoiceLibraryRepository) -> None:
|
||||
self.repository = repository
|
||||
|
||||
def execute(self, voice_id: str, user_id: str) -> bool:
|
||||
return self.repository.delete(voice_id, user_id)
|
||||
|
||||
|
||||
class QuotaExceededError(Exception):
|
||||
def __init__(self, dimension: str, limit: float, used: float) -> None:
|
||||
self.dimension = dimension
|
||||
self.limit = limit
|
||||
self.used = used
|
||||
super().__init__(f"Quota exceeded for {dimension}: {used}/{limit}")
|
||||
|
||||
|
||||
class NotFoundError(Exception):
|
||||
pass
|
||||
@@ -5,11 +5,7 @@ from .classification import (
|
||||
ClassificationJob,
|
||||
ClassificationJobStatus,
|
||||
)
|
||||
from .edit_plan import (
|
||||
EditClipPlan,
|
||||
EditPlanResult,
|
||||
EditingMode,
|
||||
)
|
||||
from .editing_mode import EditingMode
|
||||
from .entities import (
|
||||
Asset,
|
||||
AssetLibrary,
|
||||
@@ -23,7 +19,8 @@ from .entities import (
|
||||
)
|
||||
from .generated_video import GeneratedVideo
|
||||
from .generation_task import GenerationTask, GenerationTaskStatus
|
||||
from .project_management import Milestone, Task, TaskIssue, TaskPriority, TaskStatus
|
||||
from .title_library import TitleLibraryItem
|
||||
from .voice_library import VoiceLibraryItem
|
||||
|
||||
__all__ = [
|
||||
"Asset",
|
||||
@@ -34,19 +31,14 @@ __all__ = [
|
||||
"ClassificationJob",
|
||||
"ClassificationJobStatus",
|
||||
"ClassificationStatus",
|
||||
"EditClipPlan",
|
||||
"EditPlanResult",
|
||||
"EditingMode",
|
||||
"GeneratedVideo",
|
||||
"GenerationTask",
|
||||
"GenerationTaskStatus",
|
||||
"IngestJob",
|
||||
"IngestJobStatus",
|
||||
"Milestone",
|
||||
"Project",
|
||||
"Task",
|
||||
"TaskIssue",
|
||||
"TaskPriority",
|
||||
"TaskStatus",
|
||||
"User",
|
||||
"TitleLibraryItem",
|
||||
"VoiceLibraryItem",
|
||||
]
|
||||
|
||||
@@ -1,379 +0,0 @@
|
||||
"""Edit Plan domain models - shared between API and Worker."""
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import StrEnum
|
||||
|
||||
|
||||
class EditingMode(StrEnum):
|
||||
"""剪辑模式枚举"""
|
||||
ONE_TAKE = "one_take" # 顺序拼接模式
|
||||
PIP = "pip" # 画中画模式
|
||||
VOICE_OVER = "voice_over" # 口播+B-roll模式
|
||||
VOICE_PIP = "voice_pip" # 口播+画中画组合模式
|
||||
|
||||
|
||||
@dataclass
|
||||
class EditClipPlan:
|
||||
"""单个剪辑片段的编排计划"""
|
||||
asset_id: str
|
||||
sequence: int
|
||||
start_time: float = 0.0
|
||||
duration: float = 0.0
|
||||
layer: str = "main" # main, pip, broll
|
||||
reason: str = ""
|
||||
|
||||
|
||||
@dataclass
|
||||
class EditPlanResult:
|
||||
"""完整剪辑计划结果"""
|
||||
project_id: str
|
||||
editing_mode: EditingMode
|
||||
clips: list[EditClipPlan]
|
||||
total_duration: float
|
||||
summary: str
|
||||
|
||||
|
||||
|
||||
|
||||
import json as _json
|
||||
import logging as _logging
|
||||
from collections import defaultdict as _defaultdict
|
||||
from packages.domain.entities import Asset, AssetStatus
|
||||
from packages.domain.classification import AssetClassification
|
||||
|
||||
_logger = _logging.getLogger(__name__)
|
||||
|
||||
def _calculate_start_times(clips):
|
||||
current_time = 0.0
|
||||
for clip in clips:
|
||||
clip.start_time = current_time
|
||||
current_time += clip.duration
|
||||
return clips
|
||||
|
||||
|
||||
class SmartEditPlanGenerator:
|
||||
"""
|
||||
智能剪辑计划生成器
|
||||
|
||||
根据素材的分类结果和质量评分,自动编排剪辑计划。
|
||||
支持多种剪辑模式:one_take, pip, voice_over, voice_pip
|
||||
"""
|
||||
|
||||
def __init__(self, project_id: str, assets: list[Asset]):
|
||||
self.project_id = project_id
|
||||
# 筛选已就绪的视频素材
|
||||
self.assets = [
|
||||
a for a in assets
|
||||
if a.status == AssetStatus.READY and a.mime_type.startswith("video/")
|
||||
]
|
||||
self.assets_by_classification: dict[str, list[Asset]] = defaultdict(list)
|
||||
|
||||
def _parse_classification(self, asset: Asset) -> str:
|
||||
"""解析素材的分类结果"""
|
||||
# 从 metadata 中获取分类
|
||||
classification = asset.metadata.get("classification", "")
|
||||
if not classification:
|
||||
# 尝试从 classification_result 字段获取
|
||||
classification = asset.metadata.get("classification_result", "")
|
||||
|
||||
# 如果是 JSON 字符串,解析它
|
||||
if classification and isinstance(classification, str):
|
||||
try:
|
||||
parsed = json.loads(classification)
|
||||
if isinstance(parsed, dict):
|
||||
classification = parsed.get("classification", "other")
|
||||
elif isinstance(parsed, str):
|
||||
classification = parsed
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
|
||||
# 验证分类值是否有效
|
||||
valid_classifications = [c.value for c in AssetClassification]
|
||||
if classification not in valid_classifications:
|
||||
classification = "other"
|
||||
|
||||
return classification
|
||||
|
||||
def _group_by_classification(self) -> None:
|
||||
"""按分类结果对素材分组"""
|
||||
for asset in self.assets:
|
||||
classification = self._parse_classification(asset)
|
||||
self.assets_by_classification[classification].append(asset)
|
||||
|
||||
def _sort_by_quality(self, assets: list[Asset]) -> list[Asset]:
|
||||
"""按质量评分排序,高分在前"""
|
||||
return sorted(
|
||||
assets,
|
||||
key=lambda a: (-(a.quality_score or 0), a.created_at)
|
||||
)
|
||||
|
||||
def _calculate_clip_duration(self, asset: Asset, target_duration: float, clip_count: int) -> float:
|
||||
"""计算单个片段的时长"""
|
||||
if asset.duration:
|
||||
# 如果素材时长超过平均时长,取平均时长
|
||||
avg_duration = target_duration / max(1, clip_count)
|
||||
return min(float(asset.duration), avg_duration)
|
||||
return target_duration / max(1, clip_count)
|
||||
|
||||
def _generate_one_take(self, target_duration: float = 30.0) -> EditPlanResult:
|
||||
"""
|
||||
One-Take 模式:按分类分组,组内按质量排序,顺序拼接
|
||||
"""
|
||||
self._group_by_classification()
|
||||
|
||||
clips: list[EditClipPlan] = []
|
||||
sequence = 1
|
||||
|
||||
# 按优先级排序分类:person > scenic > product > other
|
||||
priority_order = ["person", "scenic", "product", "animal", "food", "tech", "sport", "music", "other"]
|
||||
sorted_classifications = sorted(
|
||||
self.assets_by_classification.keys(),
|
||||
key=lambda c: priority_order.index(c) if c in priority_order else len(priority_order)
|
||||
)
|
||||
|
||||
for classification in sorted_classifications:
|
||||
sorted_assets = self._sort_by_quality(self.assets_by_classification[classification])
|
||||
for asset in sorted_assets:
|
||||
duration = self._calculate_clip_duration(
|
||||
asset, target_duration, len(self.assets)
|
||||
)
|
||||
clips.append(EditClipPlan(
|
||||
asset_id=asset.id,
|
||||
sequence=sequence,
|
||||
start_time=0,
|
||||
duration=duration,
|
||||
layer="main",
|
||||
reason=f"按分类 [{classification}] 排列,质量评分 {asset.quality_score or 0:.1f}"
|
||||
))
|
||||
sequence += 1
|
||||
|
||||
total_duration = sum(c.duration for c in clips)
|
||||
_calculate_start_times(clips)
|
||||
return EditPlanResult(
|
||||
project_id=self.project_id,
|
||||
editing_mode=EditingMode.ONE_TAKE,
|
||||
clips=clips,
|
||||
total_duration=total_duration,
|
||||
summary=f"One-Take 模式:按 {len(sorted_classifications)} 个分类分组,共 {len(clips)} 段素材"
|
||||
)
|
||||
|
||||
def _generate_pip(self, target_duration: float = 30.0) -> EditPlanResult:
|
||||
"""
|
||||
PIP 模式:第一个高质量素材为主画面,其余为画中画
|
||||
"""
|
||||
sorted_assets = self._sort_by_quality(self.assets)
|
||||
|
||||
if not sorted_assets:
|
||||
return EditPlanResult(
|
||||
project_id=self.project_id,
|
||||
editing_mode=EditingMode.PIP,
|
||||
clips=[],
|
||||
total_duration=0,
|
||||
summary="无素材可用"
|
||||
)
|
||||
|
||||
clips: list[EditClipPlan] = []
|
||||
sequence = 1
|
||||
|
||||
# 第一个高质量素材作为主画面
|
||||
main_asset = sorted_assets[0]
|
||||
main_duration = min(
|
||||
float(main_asset.duration) if main_asset.duration else target_duration,
|
||||
target_duration
|
||||
)
|
||||
clips.append(EditClipPlan(
|
||||
asset_id=main_asset.id,
|
||||
sequence=sequence,
|
||||
start_time=0,
|
||||
duration=main_duration,
|
||||
layer="main",
|
||||
reason=f"高质量主画面 (质量评分: {main_asset.quality_score or 0:.1f})"
|
||||
))
|
||||
sequence += 1
|
||||
|
||||
# 其余素材作为画中画
|
||||
for asset in sorted_assets[1:]:
|
||||
duration = self._calculate_clip_duration(asset, target_duration, len(sorted_assets))
|
||||
clips.append(EditClipPlan(
|
||||
asset_id=asset.id,
|
||||
sequence=sequence,
|
||||
start_time=0,
|
||||
duration=duration,
|
||||
layer="pip",
|
||||
reason=f"画中画素材 (质量评分: {asset.quality_score or 0:.1f})"
|
||||
))
|
||||
sequence += 1
|
||||
|
||||
total_duration = main_duration
|
||||
_calculate_start_times(clips)
|
||||
return EditPlanResult(
|
||||
project_id=self.project_id,
|
||||
editing_mode=EditingMode.PIP,
|
||||
clips=clips,
|
||||
total_duration=total_duration,
|
||||
summary=f"PIP 模式:1 个主画面 + {len(sorted_assets) - 1} 个画中画"
|
||||
)
|
||||
|
||||
def _generate_voiceover(self, target_duration: float = 30.0) -> EditPlanResult:
|
||||
"""
|
||||
Voiceover 模式:person 类素材为主播口播,其余穿插为 B-roll
|
||||
"""
|
||||
self._group_by_classification()
|
||||
|
||||
person_assets = self._sort_by_quality(
|
||||
self.assets_by_classification.get("person", [])
|
||||
)
|
||||
other_assets = self._sort_by_quality([
|
||||
a for assets in self.assets_by_classification.values()
|
||||
for a in assets
|
||||
if self._parse_classification(a) != "person"
|
||||
])
|
||||
|
||||
clips: list[EditClipPlan] = []
|
||||
sequence = 1
|
||||
|
||||
# 合并口播和 B-roll
|
||||
main_assets = person_assets if person_assets else other_assets
|
||||
broll_assets = [a for a in other_assets if a not in person_assets] if person_assets else []
|
||||
|
||||
# 优先使用 person 素材作为口播
|
||||
for i, asset in enumerate(main_assets):
|
||||
duration = self._calculate_clip_duration(asset, target_duration, len(main_assets))
|
||||
is_person = asset in person_assets
|
||||
clips.append(EditClipPlan(
|
||||
asset_id=asset.id,
|
||||
sequence=sequence,
|
||||
start_time=0,
|
||||
duration=duration,
|
||||
layer="main" if is_person else "broll",
|
||||
reason=f"{'主播口播' if is_person else 'B-roll'} (质量评分: {asset.quality_score or 0:.1f})"
|
||||
))
|
||||
sequence += 1
|
||||
|
||||
# 在口播之间穿插 B-roll
|
||||
if is_person and broll_assets and i < len(main_assets) - 1:
|
||||
broll_asset = broll_assets[i % len(broll_assets)]
|
||||
broll_duration = self._calculate_clip_duration(
|
||||
broll_asset, target_duration, len(main_assets) + len(broll_assets)
|
||||
)
|
||||
clips.append(EditClipPlan(
|
||||
asset_id=broll_asset.id,
|
||||
sequence=sequence,
|
||||
start_time=0,
|
||||
duration=broll_duration,
|
||||
layer="broll",
|
||||
reason=f"B-roll 穿插 (质量评分: {broll_asset.quality_score or 0:.1f})"
|
||||
))
|
||||
sequence += 1
|
||||
|
||||
total_duration = sum(c.duration for c in clips)
|
||||
person_count = len(person_assets)
|
||||
_calculate_start_times(clips)
|
||||
return EditPlanResult(
|
||||
project_id=self.project_id,
|
||||
editing_mode=EditingMode.VOICE_OVER,
|
||||
clips=clips,
|
||||
total_duration=total_duration,
|
||||
summary=f"Voiceover 模式:{person_count} 段口播 + {len(clips) - person_count} 段 B-roll"
|
||||
)
|
||||
|
||||
def _generate_voice_pip(self, target_duration: float = 30.0) -> EditPlanResult:
|
||||
"""
|
||||
Voice-PIP 模式:结合 voiceover 和 pip
|
||||
第一个高质量 person 素材为主画面,其余为 PIP B-roll
|
||||
"""
|
||||
self._group_by_classification()
|
||||
|
||||
person_assets = self._sort_by_quality(
|
||||
self.assets_by_classification.get("person", [])
|
||||
)
|
||||
other_assets = self._sort_by_quality([
|
||||
a for assets in self.assets_by_classification.values()
|
||||
for a in assets
|
||||
if self._parse_classification(a) != "person"
|
||||
])
|
||||
|
||||
clips: list[EditClipPlan] = []
|
||||
sequence = 1
|
||||
|
||||
# 主画面:优先使用高质量 person 素材
|
||||
main_asset = person_assets[0] if person_assets else (other_assets[0] if other_assets else None)
|
||||
if main_asset:
|
||||
main_duration = min(
|
||||
float(main_asset.duration) if main_asset.duration else target_duration,
|
||||
target_duration
|
||||
)
|
||||
is_person = main_asset in person_assets
|
||||
clips.append(EditClipPlan(
|
||||
asset_id=main_asset.id,
|
||||
sequence=sequence,
|
||||
start_time=0,
|
||||
duration=main_duration,
|
||||
layer="main",
|
||||
reason=f"{'主播口播' if is_person else '主画面'} (质量评分: {main_asset.quality_score or 0:.1f})"
|
||||
))
|
||||
sequence += 1
|
||||
|
||||
# PIP 素材
|
||||
pip_assets = [a for a in (person_assets[1:] + other_assets) if a != main_asset]
|
||||
for asset in pip_assets:
|
||||
duration = self._calculate_clip_duration(asset, target_duration, len(pip_assets) + 1)
|
||||
clips.append(EditClipPlan(
|
||||
asset_id=asset.id,
|
||||
sequence=sequence,
|
||||
start_time=0,
|
||||
duration=duration,
|
||||
layer="pip",
|
||||
reason=f"PIP 素材 (质量评分: {asset.quality_score or 0:.1f})"
|
||||
))
|
||||
sequence += 1
|
||||
|
||||
total_duration = sum(c.duration for c in clips)
|
||||
_calculate_start_times(clips)
|
||||
return EditPlanResult(
|
||||
project_id=self.project_id,
|
||||
editing_mode=EditingMode.VOICE_PIP,
|
||||
clips=clips,
|
||||
total_duration=total_duration,
|
||||
summary=f"Voice-PIP 模式:1 个主画面 + {len(pip_assets)} 个 PIP 素材"
|
||||
)
|
||||
|
||||
def generate_plan(
|
||||
self,
|
||||
editing_mode: str = "one_take",
|
||||
target_duration: float = 30.0
|
||||
) -> EditPlanResult:
|
||||
"""
|
||||
生成剪辑计划
|
||||
|
||||
Args:
|
||||
editing_mode: 剪辑模式 (one_take/pip/voice_over/voice_pip)
|
||||
target_duration: 目标时长(秒)
|
||||
|
||||
Returns:
|
||||
EditPlanResult: 编排好的剪辑计划
|
||||
"""
|
||||
logger.info(f"Generating edit plan for project {self.project_id} with mode {editing_mode}")
|
||||
|
||||
if not self.assets:
|
||||
logger.warning(f"No ready video assets found for project {self.project_id}")
|
||||
return EditPlanResult(
|
||||
project_id=self.project_id,
|
||||
editing_mode=EditingMode(editing_mode),
|
||||
clips=[],
|
||||
total_duration=0,
|
||||
summary="无素材可用"
|
||||
)
|
||||
|
||||
mode = EditingMode(editing_mode.lower())
|
||||
|
||||
if mode == EditingMode.ONE_TAKE:
|
||||
return self._generate_one_take(target_duration)
|
||||
elif mode == EditingMode.PIP:
|
||||
return self._generate_pip(target_duration)
|
||||
elif mode == EditingMode.VOICE_OVER:
|
||||
return self._generate_voiceover(target_duration)
|
||||
elif mode == EditingMode.VOICE_PIP:
|
||||
return self._generate_voice_pip(target_duration)
|
||||
else:
|
||||
raise ValueError(f"Unknown editing mode: {editing_mode}")
|
||||
@@ -0,0 +1,11 @@
|
||||
"""Editing mode enum — extracted from edit_plan for independent use."""
|
||||
|
||||
from enum import StrEnum
|
||||
|
||||
|
||||
class EditingMode(StrEnum):
|
||||
"""剪辑模式枚举"""
|
||||
ONE_TAKE = "one_take" # 顺序拼接模式
|
||||
PIP = "pip" # 画中画模式
|
||||
VOICE_OVER = "voice_over" # 口播+B-roll模式
|
||||
VOICE_PIP = "voice_pip" # 口播+画中画组合模式
|
||||
@@ -317,76 +317,6 @@ class EditTemplate:
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class EditPlan:
|
||||
id: str
|
||||
project_id: str
|
||||
template_id: str
|
||||
asset_library_id: str
|
||||
title_id: str = ""
|
||||
status: str = "draft"
|
||||
summary: str = ""
|
||||
created_by_user_id: str = ""
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ProjectTitle:
|
||||
id: str
|
||||
project_id: str
|
||||
text: str
|
||||
category: str = "default"
|
||||
favorite: bool = False
|
||||
usage_count: int = 0
|
||||
is_active: bool = True
|
||||
created_by_user_id: str = ""
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class Task:
|
||||
id: str
|
||||
project_id: str
|
||||
name: str
|
||||
description: str = ""
|
||||
status: str = "pending"
|
||||
priority: str = "medium"
|
||||
parent_task_id: str = ""
|
||||
assignee_user_id: str = ""
|
||||
progress: float = 0.0
|
||||
planned_start_date: datetime | None = None
|
||||
planned_end_date: datetime | None = None
|
||||
actual_start_date: datetime | None = None
|
||||
actual_end_date: datetime | None = None
|
||||
tags: 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))
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class Milestone:
|
||||
id: str
|
||||
project_id: str
|
||||
name: str
|
||||
description: str = ""
|
||||
target_date: datetime | None = None
|
||||
completed: bool = False
|
||||
completed_at: datetime | None = None
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class TaskIssue:
|
||||
id: str
|
||||
task_id: str
|
||||
project_id: str
|
||||
title: str
|
||||
description: str = ""
|
||||
resolved: bool = False
|
||||
resolved_at: datetime | None = None
|
||||
created_by_user_id: str = ""
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
@@ -21,7 +21,6 @@ class GenerationTask:
|
||||
asset_library_id: str
|
||||
strategy_id: str = ""
|
||||
voice_library_id: str = ""
|
||||
edit_plan_id: str = ""
|
||||
status: GenerationTaskStatus = GenerationTaskStatus.PENDING
|
||||
progress: float = 0.0
|
||||
result_count: int = 0
|
||||
@@ -39,7 +38,6 @@ class GenerationTask:
|
||||
*,
|
||||
strategy_id: str = "",
|
||||
voice_library_id: str = "",
|
||||
edit_plan_id: str = "",
|
||||
created_by_user_id: str = "",
|
||||
) -> "GenerationTask":
|
||||
if not project_id.strip():
|
||||
@@ -52,6 +50,5 @@ class GenerationTask:
|
||||
asset_library_id=asset_library_id.strip(),
|
||||
strategy_id=strategy_id.strip(),
|
||||
voice_library_id=voice_library_id.strip(),
|
||||
edit_plan_id=edit_plan_id.strip(),
|
||||
created_by_user_id=created_by_user_id.strip(),
|
||||
)
|
||||
|
||||
@@ -1,232 +0,0 @@
|
||||
"""项目管理领域对象:任务、里程碑、项目阶段"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from enum import StrEnum
|
||||
from uuid import uuid4
|
||||
|
||||
|
||||
class TaskStatus(StrEnum):
|
||||
"""任务状态"""
|
||||
|
||||
PENDING = "pending" # 待开始
|
||||
IN_PROGRESS = "in_progress" # 进行中
|
||||
BLOCKED = "blocked" # 阻塞
|
||||
COMPLETED = "completed" # 已完成
|
||||
CANCELLED = "cancelled" # 已取消
|
||||
|
||||
|
||||
class TaskPriority(StrEnum):
|
||||
"""任务优先级"""
|
||||
|
||||
LOW = "low"
|
||||
MEDIUM = "medium"
|
||||
HIGH = "high"
|
||||
URGENT = "urgent"
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class Task:
|
||||
"""任务实体"""
|
||||
|
||||
id: str
|
||||
project_id: str
|
||||
name: str
|
||||
description: str = ""
|
||||
status: TaskStatus = TaskStatus.PENDING
|
||||
priority: TaskPriority = TaskPriority.MEDIUM
|
||||
parent_task_id: str = "" # 父任务ID(支持子任务层级)
|
||||
assignee_user_id: str = "" # 负责人
|
||||
progress: float = 0.0 # 进度 0-100
|
||||
planned_start_date: datetime | None = None
|
||||
planned_end_date: datetime | None = None
|
||||
actual_start_date: datetime | None = None
|
||||
actual_end_date: datetime | None = None
|
||||
tags: 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))
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
cls,
|
||||
project_id: str,
|
||||
name: str,
|
||||
description: str = "",
|
||||
priority: TaskPriority = TaskPriority.MEDIUM,
|
||||
parent_task_id: str = "",
|
||||
assignee_user_id: str = "",
|
||||
planned_start_date: datetime | None = None,
|
||||
planned_end_date: datetime | None = None,
|
||||
) -> "Task":
|
||||
"""创建任务"""
|
||||
clean_name = name.strip()
|
||||
if not clean_name:
|
||||
raise ValueError("任务名称不能为空")
|
||||
if not project_id.strip():
|
||||
raise ValueError("project_id 不能为空")
|
||||
|
||||
return cls(
|
||||
id=uuid4().hex,
|
||||
project_id=project_id.strip(),
|
||||
name=clean_name,
|
||||
description=description.strip(),
|
||||
priority=priority,
|
||||
parent_task_id=parent_task_id.strip(),
|
||||
assignee_user_id=assignee_user_id.strip(),
|
||||
planned_start_date=planned_start_date,
|
||||
planned_end_date=planned_end_date,
|
||||
)
|
||||
|
||||
def update_status(self, new_status: TaskStatus) -> None:
|
||||
"""更新任务状态"""
|
||||
self.status = new_status
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
# 自动设置实际开始/结束时间
|
||||
if new_status == TaskStatus.IN_PROGRESS and self.actual_start_date is None:
|
||||
self.actual_start_date = datetime.now(timezone.utc)
|
||||
elif new_status == TaskStatus.COMPLETED and self.actual_end_date is None:
|
||||
self.actual_end_date = datetime.now(timezone.utc)
|
||||
self.progress = 100.0
|
||||
|
||||
def update_progress(self, progress: float) -> None:
|
||||
"""更新任务进度"""
|
||||
if not 0 <= progress <= 100:
|
||||
raise ValueError("进度必须在 0-100 之间")
|
||||
self.progress = progress
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
# 自动更新状态
|
||||
if progress > 0 and self.status == TaskStatus.PENDING:
|
||||
self.status = TaskStatus.IN_PROGRESS
|
||||
if progress == 100 and self.status != TaskStatus.COMPLETED:
|
||||
self.status = TaskStatus.COMPLETED
|
||||
if self.actual_end_date is None:
|
||||
self.actual_end_date = datetime.now(timezone.utc)
|
||||
|
||||
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)
|
||||
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)
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class Milestone:
|
||||
"""里程碑实体"""
|
||||
|
||||
id: str
|
||||
project_id: str
|
||||
name: str
|
||||
description: str = ""
|
||||
target_date: datetime | None = None
|
||||
completed: bool = False
|
||||
completed_at: datetime | None = None
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
cls,
|
||||
project_id: str,
|
||||
name: str,
|
||||
description: str = "",
|
||||
target_date: datetime | None = None,
|
||||
) -> "Milestone":
|
||||
"""创建里程碑"""
|
||||
clean_name = name.strip()
|
||||
if not clean_name:
|
||||
raise ValueError("里程碑名称不能为空")
|
||||
if not project_id.strip():
|
||||
raise ValueError("project_id 不能为空")
|
||||
|
||||
return cls(
|
||||
id=uuid4().hex,
|
||||
project_id=project_id.strip(),
|
||||
name=clean_name,
|
||||
description=description.strip(),
|
||||
target_date=target_date,
|
||||
)
|
||||
|
||||
def mark_completed(self) -> None:
|
||||
"""标记为已完成"""
|
||||
if not self.completed:
|
||||
self.completed = True
|
||||
self.completed_at = datetime.now(timezone.utc)
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def reopen(self) -> None:
|
||||
"""重新打开里程碑"""
|
||||
if self.completed:
|
||||
self.completed = False
|
||||
self.completed_at = None
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class TaskIssue:
|
||||
"""任务问题/卡点实体"""
|
||||
|
||||
id: str
|
||||
task_id: str
|
||||
project_id: str
|
||||
title: str
|
||||
description: str = ""
|
||||
resolved: bool = False
|
||||
resolved_at: datetime | None = None
|
||||
created_by_user_id: str = ""
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
cls,
|
||||
task_id: str,
|
||||
project_id: str,
|
||||
title: str,
|
||||
description: str = "",
|
||||
created_by_user_id: str = "",
|
||||
) -> "TaskIssue":
|
||||
"""创建任务问题"""
|
||||
clean_title = title.strip()
|
||||
if not clean_title:
|
||||
raise ValueError("问题标题不能为空")
|
||||
if not task_id.strip():
|
||||
raise ValueError("task_id 不能为空")
|
||||
if not project_id.strip():
|
||||
raise ValueError("project_id 不能为空")
|
||||
|
||||
return cls(
|
||||
id=uuid4().hex,
|
||||
task_id=task_id.strip(),
|
||||
project_id=project_id.strip(),
|
||||
title=clean_title,
|
||||
description=description.strip(),
|
||||
created_by_user_id=created_by_user_id.strip(),
|
||||
)
|
||||
|
||||
def mark_resolved(self) -> None:
|
||||
"""标记为已解决"""
|
||||
if not self.resolved:
|
||||
self.resolved = True
|
||||
self.resolved_at = datetime.now(timezone.utc)
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
|
||||
def reopen(self) -> None:
|
||||
"""重新打开问题"""
|
||||
if self.resolved:
|
||||
self.resolved = False
|
||||
self.resolved_at = None
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
@@ -0,0 +1,23 @@
|
||||
"""Title library domain entity."""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from typing import List
|
||||
|
||||
|
||||
@dataclass
|
||||
class TitleLibraryItem:
|
||||
"""标题库条目"""
|
||||
id: str
|
||||
user_id: str
|
||||
name: str
|
||||
text: str
|
||||
category: str = "default"
|
||||
description: str = ""
|
||||
tags: List[str] = field(default_factory=list)
|
||||
usage_count: int = 0
|
||||
is_active: bool = True
|
||||
metadata_: dict = field(default_factory=dict)
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
@@ -0,0 +1,27 @@
|
||||
"""Voice library domain entity."""
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from typing import List, Optional
|
||||
|
||||
|
||||
@dataclass
|
||||
class VoiceLibraryItem:
|
||||
"""配音库条目"""
|
||||
id: str
|
||||
user_id: str
|
||||
name: str
|
||||
text: str = ""
|
||||
voice_provider: str = ""
|
||||
voice_id: str = ""
|
||||
voice_name: str = ""
|
||||
audio_url: str = ""
|
||||
duration: float = 0
|
||||
file_size: int = 0
|
||||
status: str = "completed"
|
||||
project_id: Optional[str] = None
|
||||
tags: List[str] = field(default_factory=list)
|
||||
metadata_: dict = field(default_factory=dict)
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
updated_at: datetime = field(default_factory=lambda: datetime.now(timezone.utc))
|
||||
@@ -3,19 +3,15 @@
|
||||
from .asset_library_repository import AssetLibraryRepository
|
||||
from .asset_repository import AssetRepository
|
||||
from .ingest_job_repository import IngestJobRepository
|
||||
from .project_management_repositories import (
|
||||
MilestoneRepository,
|
||||
TaskIssueRepository,
|
||||
TaskRepository,
|
||||
)
|
||||
from .project_repository import ProjectRepository
|
||||
from .title_library_repository import TitleLibraryRepository
|
||||
from .voice_library_repository import VoiceLibraryRepository
|
||||
|
||||
__all__ = [
|
||||
"AssetLibraryRepository",
|
||||
"AssetRepository",
|
||||
"IngestJobRepository",
|
||||
"MilestoneRepository",
|
||||
"ProjectRepository",
|
||||
"TaskIssueRepository",
|
||||
"TaskRepository",
|
||||
"TitleLibraryRepository",
|
||||
"VoiceLibraryRepository",
|
||||
]
|
||||
|
||||
@@ -1,102 +0,0 @@
|
||||
"""项目管理 Repository 接口定义"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
|
||||
from packages.domain import Milestone, Task, TaskIssue
|
||||
|
||||
|
||||
class TaskRepository(ABC):
|
||||
"""任务仓储接口"""
|
||||
|
||||
@abstractmethod
|
||||
def create(self, task: Task) -> Task:
|
||||
"""创建任务"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_by_id(self, task_id: str) -> Task | None:
|
||||
"""根据ID获取任务"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def list_by_project(self, project_id: str) -> list[Task]:
|
||||
"""获取项目下的所有任务"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def list_by_parent(self, parent_task_id: str) -> list[Task]:
|
||||
"""获取子任务列表"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def update(self, task: Task) -> Task:
|
||||
"""更新任务"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def delete(self, task_id: str) -> None:
|
||||
"""删除任务"""
|
||||
pass
|
||||
|
||||
|
||||
class MilestoneRepository(ABC):
|
||||
"""里程碑仓储接口"""
|
||||
|
||||
@abstractmethod
|
||||
def create(self, milestone: Milestone) -> Milestone:
|
||||
"""创建里程碑"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_by_id(self, milestone_id: str) -> Milestone | None:
|
||||
"""根据ID获取里程碑"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def list_by_project(self, project_id: str) -> list[Milestone]:
|
||||
"""获取项目下的所有里程碑"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def update(self, milestone: Milestone) -> Milestone:
|
||||
"""更新里程碑"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def delete(self, milestone_id: str) -> None:
|
||||
"""删除里程碑"""
|
||||
pass
|
||||
|
||||
|
||||
class TaskIssueRepository(ABC):
|
||||
"""任务问题仓储接口"""
|
||||
|
||||
@abstractmethod
|
||||
def create(self, issue: TaskIssue) -> TaskIssue:
|
||||
"""创建任务问题"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def get_by_id(self, issue_id: str) -> TaskIssue | None:
|
||||
"""根据ID获取任务问题"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def list_by_task(self, task_id: str) -> list[TaskIssue]:
|
||||
"""获取任务下的所有问题"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def list_by_project(self, project_id: str) -> list[TaskIssue]:
|
||||
"""获取项目下的所有问题"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def update(self, issue: TaskIssue) -> TaskIssue:
|
||||
"""更新任务问题"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def delete(self, issue_id: str) -> None:
|
||||
"""删除任务问题"""
|
||||
pass
|
||||
@@ -1,9 +0,0 @@
|
||||
"""Port interface for project title repository."""
|
||||
from __future__ import annotations
|
||||
from typing import Protocol, Any
|
||||
|
||||
class ProjectTitleRepository(Protocol):
|
||||
def list_by_project(self, project_id: str, active_only: bool = False) -> list[Any]: ...
|
||||
def get(self, title_id: str) -> Any | None: ...
|
||||
def create(self, *, project_id: str, text: str, category: str, created_by_user_id: str, favorite: bool = False) -> Any: ...
|
||||
def update(self, item: Any) -> Any: ...
|
||||
@@ -0,0 +1,36 @@
|
||||
"""Title library repository port."""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import List, Optional, Protocol
|
||||
|
||||
from packages.domain.title_library import TitleLibraryItem
|
||||
|
||||
|
||||
class TitleLibraryRepository(Protocol):
|
||||
"""标题库仓储接口"""
|
||||
|
||||
def list_by_user(
|
||||
self,
|
||||
user_id: str,
|
||||
*,
|
||||
category: Optional[str] = None,
|
||||
is_active: bool = True,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
) -> List[TitleLibraryItem]:
|
||||
...
|
||||
|
||||
def get(self, title_id: str, user_id: str) -> Optional[TitleLibraryItem]:
|
||||
...
|
||||
|
||||
def create(self, item: TitleLibraryItem) -> TitleLibraryItem:
|
||||
...
|
||||
|
||||
def update(self, item: TitleLibraryItem) -> TitleLibraryItem:
|
||||
...
|
||||
|
||||
def delete(self, title_id: str, user_id: str) -> bool:
|
||||
...
|
||||
|
||||
def count_by_user(self, user_id: str, is_active: bool = True) -> int:
|
||||
...
|
||||
@@ -0,0 +1,35 @@
|
||||
"""Voice library repository port."""
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import List, Optional, Protocol
|
||||
|
||||
from packages.domain.voice_library import VoiceLibraryItem
|
||||
|
||||
|
||||
class VoiceLibraryRepository(Protocol):
|
||||
"""配音库仓储接口"""
|
||||
|
||||
def list_by_user(
|
||||
self,
|
||||
user_id: str,
|
||||
*,
|
||||
status: Optional[str] = None,
|
||||
skip: int = 0,
|
||||
limit: int = 50,
|
||||
) -> List[VoiceLibraryItem]:
|
||||
...
|
||||
|
||||
def get(self, voice_id: str, user_id: str) -> Optional[VoiceLibraryItem]:
|
||||
...
|
||||
|
||||
def create(self, item: VoiceLibraryItem) -> VoiceLibraryItem:
|
||||
...
|
||||
|
||||
def update(self, item: VoiceLibraryItem) -> VoiceLibraryItem:
|
||||
...
|
||||
|
||||
def delete(self, voice_id: str, user_id: str) -> bool:
|
||||
...
|
||||
|
||||
def count_by_user(self, user_id: str) -> int:
|
||||
...
|
||||
@@ -3,7 +3,6 @@ from pathlib import Path
|
||||
ALLOWED_API_ADAPTER_IMPORTS = {
|
||||
Path("apps/api/app/dependencies.py"),
|
||||
Path("apps/api/app/db.py"),
|
||||
Path("apps/api/app/api/routes/edit_plans.py"),
|
||||
}
|
||||
|
||||
|
||||
|
||||
@@ -1,71 +0,0 @@
|
||||
import sys
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "worker"))
|
||||
|
||||
from app.api.routes.generation_tasks import _select_title_id
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
from worker_app.core.title_usage import mark_title_used_for_generation
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import Base, ProjectTitleModel
|
||||
from packages.adapters.sqlalchemy_impl.project_title_repository import SQLAlchemyProjectTitleRepository
|
||||
|
||||
|
||||
def _session_and_repository():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session = sessionmaker(bind=engine)()
|
||||
return session, SQLAlchemyProjectTitleRepository(session)
|
||||
|
||||
|
||||
def test_select_title_prefers_favorite_then_lowest_usage():
|
||||
_, repository = _session_and_repository()
|
||||
normal = repository.create(
|
||||
project_id="project-1",
|
||||
text="普通标题",
|
||||
category="default",
|
||||
created_by_user_id="user-1",
|
||||
)
|
||||
favorite = repository.create(
|
||||
project_id="project-1",
|
||||
text="常用标题",
|
||||
category="default",
|
||||
favorite=True,
|
||||
created_by_user_id="user-1",
|
||||
)
|
||||
normal.usage_count = 0
|
||||
favorite.usage_count = 10
|
||||
repository.update(normal)
|
||||
repository.update(favorite)
|
||||
|
||||
assert _select_title_id(repository, "project-1") == favorite.id
|
||||
|
||||
|
||||
def test_mark_title_used_after_generation_completion():
|
||||
session, _ = _session_and_repository()
|
||||
now = datetime.now(timezone.utc)
|
||||
title = ProjectTitleModel(
|
||||
id="title-1",
|
||||
project_id="project-1",
|
||||
text="生成标题",
|
||||
category="default",
|
||||
favorite=False,
|
||||
usage_count=2,
|
||||
is_active=True,
|
||||
created_by_user_id="user-1",
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
)
|
||||
session.add(title)
|
||||
session.commit()
|
||||
|
||||
mark_title_used_for_generation(
|
||||
session,
|
||||
)
|
||||
|
||||
updated = session.query(ProjectTitleModel).filter(ProjectTitleModel.id == "title-1").first()
|
||||
assert updated.usage_count == 3
|
||||
@@ -1,50 +0,0 @@
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parents[2] / "apps" / "api"))
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from packages.adapters.sqlalchemy_impl.models import Base
|
||||
from packages.adapters.sqlalchemy_impl.project_title_repository import SQLAlchemyProjectTitleRepository
|
||||
|
||||
|
||||
def _repository():
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session = sessionmaker(bind=engine)()
|
||||
return SQLAlchemyProjectTitleRepository(session)
|
||||
|
||||
|
||||
def test_project_title_repository_creates_and_lists_titles():
|
||||
repository = _repository()
|
||||
|
||||
title = repository.create(
|
||||
project_id="project-1",
|
||||
text=" 3 分钟看懂产品亮点 ",
|
||||
category="marketing",
|
||||
favorite=True,
|
||||
created_by_user_id="user-1",
|
||||
)
|
||||
|
||||
assert title.text == "3 分钟看懂产品亮点"
|
||||
assert title.category == "marketing"
|
||||
assert title.favorite is True
|
||||
assert title.usage_count == 0
|
||||
assert title.is_active is True
|
||||
assert repository.list_by_project("project-1") == [title]
|
||||
|
||||
|
||||
def test_project_title_repository_filters_inactive_titles():
|
||||
repository = _repository()
|
||||
title = repository.create(
|
||||
project_id="project-1",
|
||||
text="停用标题",
|
||||
category="default",
|
||||
created_by_user_id="user-1",
|
||||
)
|
||||
title.is_active = False
|
||||
repository.update(title)
|
||||
|
||||
assert repository.list_by_project("project-1", active_only=True) == []
|
||||
@@ -0,0 +1,488 @@
|
||||
"""
|
||||
标题库(Title Library)Use Case 回归测试
|
||||
|
||||
测试目标:
|
||||
1. CreateTitleLibraryUseCase - 创建标题库条目
|
||||
2. UpdateTitleLibraryUseCase - 更新标题库条目
|
||||
3. 配额逻辑覆盖 - titles: free=50, basic=500, premium=500
|
||||
4. 边界条件与异常场景
|
||||
"""
|
||||
|
||||
from unittest.mock import Mock
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.title_library.commands import (
|
||||
CreateTitleLibraryCommand,
|
||||
UpdateTitleLibraryCommand,
|
||||
)
|
||||
from packages.application.title_library.use_cases import (
|
||||
CreateTitleLibraryUseCase,
|
||||
UpdateTitleLibraryUseCase,
|
||||
DeleteTitleLibraryUseCase,
|
||||
GetTitleLibraryUseCase,
|
||||
ListTitleLibraryUseCase,
|
||||
QuotaExceededError,
|
||||
NotFoundError,
|
||||
)
|
||||
from packages.domain.title_library import TitleLibraryItem
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.fixture
|
||||
def mock_repo():
|
||||
"""创建 Mock 仓储"""
|
||||
repo = Mock()
|
||||
repo.count_by_user = Mock(return_value=0)
|
||||
repo.create = Mock(side_effect=lambda item: item)
|
||||
repo.update = Mock(side_effect=lambda item: item)
|
||||
repo.get = Mock(return_value=None)
|
||||
repo.delete = Mock(return_value=True)
|
||||
repo.list_by_user = Mock(return_value=[])
|
||||
return repo
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def create_use_case(mock_repo):
|
||||
return CreateTitleLibraryUseCase(repository=mock_repo)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def update_use_case(mock_repo):
|
||||
return UpdateTitleLibraryUseCase(repository=mock_repo)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_create_command():
|
||||
"""标准创建命令"""
|
||||
return CreateTitleLibraryCommand(
|
||||
user_id="user-001",
|
||||
name="测试标题",
|
||||
text="这是一个测试标题文本",
|
||||
category="新闻",
|
||||
description="用于测试的标题",
|
||||
tags=["测试", "新闻"],
|
||||
metadata_={"source": "unit_test"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def existing_title_item():
|
||||
"""模拟已存在的标题条目"""
|
||||
return TitleLibraryItem(
|
||||
id="existing-title-001",
|
||||
user_id="user-001",
|
||||
name="旧标题",
|
||||
text="旧文本",
|
||||
category="旧分类",
|
||||
description="旧描述",
|
||||
tags=["旧"],
|
||||
is_active=True,
|
||||
metadata_={},
|
||||
)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 1. CreateTitleLibraryUseCase 测试
|
||||
# ===========================================================================
|
||||
|
||||
class TestCreateTitleLibraryUseCase:
|
||||
"""标题库创建 UseCase 测试"""
|
||||
|
||||
def test_create_success_all_fields(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""测试创建成功 - 所有字段完整传入"""
|
||||
result = create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert result is not None
|
||||
assert result.user_id == "user-001"
|
||||
assert result.name == "测试标题"
|
||||
assert result.text == "这是一个测试标题文本"
|
||||
assert result.category == "新闻"
|
||||
assert result.description == "用于测试的标题"
|
||||
assert result.tags == ["测试", "新闻"]
|
||||
assert result.metadata_ == {"source": "unit_test"}
|
||||
|
||||
mock_repo.count_by_user.assert_called_once_with("user-001")
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_create_generates_uuid(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""测试创建时自动生成 UUID 作为 id"""
|
||||
result = create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert result.id is not None
|
||||
assert len(result.id) == 32 # uuid4().hex 长度为 32
|
||||
assert result.id.isalnum()
|
||||
|
||||
def test_create_default_values(self, create_use_case, mock_repo):
|
||||
"""测试默认值填充"""
|
||||
command = CreateTitleLibraryCommand(
|
||||
user_id="user-001",
|
||||
name="最小化创建",
|
||||
text="文本",
|
||||
)
|
||||
|
||||
result = create_use_case.execute(command, plan_name="free")
|
||||
|
||||
assert result.category == "default"
|
||||
assert result.description == ""
|
||||
assert result.tags == []
|
||||
assert result.metadata_ == {}
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 2. 配额逻辑测试(titles: free=50, basic=500, premium=500)
|
||||
# ===========================================================================
|
||||
|
||||
class TestCreateTitleLibraryQuota:
|
||||
"""标题库创建配额检查测试"""
|
||||
|
||||
def test_quota_free_plan_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""free 套餐(上限50),当前 25 个,允许创建"""
|
||||
mock_repo.count_by_user.return_value = 25
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_free_plan_at_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""free 套餐(上限50),当前 50 个,拒绝创建"""
|
||||
mock_repo.count_by_user.return_value = 50
|
||||
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert exc_info.value.dimension == "max_titles"
|
||||
assert exc_info.value.limit == 50
|
||||
assert exc_info.value.used == 50
|
||||
mock_repo.create.assert_not_called()
|
||||
|
||||
def test_quota_free_plan_just_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""free 套餐(上限50),当前 49 个,允许创建(边界)"""
|
||||
mock_repo.count_by_user.return_value = 49
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_free_plan_over_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""free 套餐(上限50),当前 60 个,拒绝创建"""
|
||||
mock_repo.count_by_user.return_value = 60
|
||||
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert exc_info.value.dimension == "max_titles"
|
||||
assert exc_info.value.limit == 50
|
||||
assert exc_info.value.used == 60
|
||||
|
||||
def test_quota_basic_plan_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""basic 套餐(上限500),当前 200 个,允许创建"""
|
||||
mock_repo.count_by_user.return_value = 200
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="basic")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_basic_plan_at_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""basic 套餐(上限500),当前 500 个,拒绝创建"""
|
||||
mock_repo.count_by_user.return_value = 500
|
||||
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
create_use_case.execute(sample_create_command, plan_name="basic")
|
||||
|
||||
assert exc_info.value.dimension == "max_titles"
|
||||
assert exc_info.value.limit == 500
|
||||
assert exc_info.value.used == 500
|
||||
|
||||
def test_quota_basic_plan_just_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""basic 套餐(上限500),当前 499 个,允许创建(边界)"""
|
||||
mock_repo.count_by_user.return_value = 499
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="basic")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_premium_plan_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""premium 套餐(上限500),当前 250 个,允许创建"""
|
||||
mock_repo.count_by_user.return_value = 250
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="premium")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_premium_plan_at_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""premium 套餐(上限500),当前 500 个,拒绝创建"""
|
||||
mock_repo.count_by_user.return_value = 500
|
||||
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
create_use_case.execute(sample_create_command, plan_name="premium")
|
||||
|
||||
assert exc_info.value.dimension == "max_titles"
|
||||
assert exc_info.value.limit == 500
|
||||
assert exc_info.value.used == 500
|
||||
|
||||
def test_quota_premium_plan_just_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""premium 套餐(上限500),当前 499 个,允许创建(边界)"""
|
||||
mock_repo.count_by_user.return_value = 499
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="premium")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_zero_usage_all_plans(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""新用户零使用量,所有套餐均可创建"""
|
||||
mock_repo.count_by_user.return_value = 0
|
||||
|
||||
for plan in ["free", "basic", "premium"]:
|
||||
mock_repo.create.reset_mock()
|
||||
mock_repo.count_by_user.reset_mock()
|
||||
mock_repo.count_by_user.return_value = 0
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name=plan)
|
||||
assert result is not None, f"{plan} 套餐零使用量应允许创建"
|
||||
|
||||
def test_quota_unknown_plan_defaults_to_zero(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""未知套餐名默认配额为 0,无法创建"""
|
||||
mock_repo.count_by_user.return_value = 0
|
||||
|
||||
with pytest.raises(QuotaExceededError):
|
||||
create_use_case.execute(sample_create_command, plan_name="unknown_plan")
|
||||
|
||||
def test_quota_exceeded_error_attributes(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""QuotaExceededError 异常属性完整性"""
|
||||
mock_repo.count_by_user.return_value = 50
|
||||
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
err = exc_info.value
|
||||
assert hasattr(err, "dimension")
|
||||
assert hasattr(err, "limit")
|
||||
assert hasattr(err, "used")
|
||||
assert "max_titles" in str(err)
|
||||
assert "50" in str(err)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 3. UpdateTitleLibraryUseCase 测试
|
||||
# ===========================================================================
|
||||
|
||||
class TestUpdateTitleLibraryUseCase:
|
||||
"""标题库更新 UseCase 测试"""
|
||||
|
||||
def test_update_success_all_fields(self, update_use_case, mock_repo, existing_title_item):
|
||||
"""测试全字段更新成功"""
|
||||
mock_repo.get.return_value = existing_title_item
|
||||
|
||||
command = UpdateTitleLibraryCommand(
|
||||
title_id="existing-title-001",
|
||||
user_id="user-001",
|
||||
name="更新后标题",
|
||||
text="更新后文本",
|
||||
category="新分类",
|
||||
description="新描述",
|
||||
tags=["新标签"],
|
||||
is_active=False,
|
||||
metadata_={"updated": True},
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
assert result.name == "更新后标题"
|
||||
assert result.text == "更新后文本"
|
||||
assert result.category == "新分类"
|
||||
assert result.description == "新描述"
|
||||
assert result.tags == ["新标签"]
|
||||
assert result.is_active is False
|
||||
assert result.metadata_ == {"updated": True}
|
||||
|
||||
mock_repo.update.assert_called_once()
|
||||
|
||||
def test_update_partial_only_name(self, update_use_case, mock_repo, existing_title_item):
|
||||
"""测试仅更新 name"""
|
||||
mock_repo.get.return_value = existing_title_item
|
||||
|
||||
command = UpdateTitleLibraryCommand(
|
||||
title_id="existing-title-001",
|
||||
user_id="user-001",
|
||||
name="仅改名",
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
assert result.name == "仅改名"
|
||||
# 其他字段保持不变
|
||||
assert result.text == "旧文本"
|
||||
assert result.category == "旧分类"
|
||||
assert result.description == "旧描述"
|
||||
|
||||
def test_update_partial_only_is_active(self, update_use_case, mock_repo, existing_title_item):
|
||||
"""测试仅更新 is_active(软删除/恢复)"""
|
||||
mock_repo.get.return_value = existing_title_item
|
||||
|
||||
command = UpdateTitleLibraryCommand(
|
||||
title_id="existing-title-001",
|
||||
user_id="user-001",
|
||||
is_active=False,
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
assert result.is_active is False
|
||||
assert result.name == "旧标题" # 其他字段不变
|
||||
|
||||
def test_update_not_found(self, update_use_case, mock_repo):
|
||||
"""测试更新不存在的条目"""
|
||||
mock_repo.get.return_value = None
|
||||
|
||||
command = UpdateTitleLibraryCommand(
|
||||
title_id="nonexistent-id",
|
||||
user_id="user-001",
|
||||
name="不存在",
|
||||
)
|
||||
|
||||
with pytest.raises(NotFoundError, match="nonexistent-id"):
|
||||
update_use_case.execute(command)
|
||||
|
||||
mock_repo.update.assert_not_called()
|
||||
|
||||
def test_update_wrong_user(self, update_use_case, mock_repo):
|
||||
"""测试用户隔离"""
|
||||
mock_repo.get.return_value = None
|
||||
|
||||
command = UpdateTitleLibraryCommand(
|
||||
title_id="existing-title-001",
|
||||
user_id="other-user-999",
|
||||
name="恶意修改",
|
||||
)
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
update_use_case.execute(command)
|
||||
|
||||
def test_update_none_fields_not_changed(self, update_use_case, mock_repo, existing_title_item):
|
||||
"""测试 None 字段不覆盖原有值"""
|
||||
mock_repo.get.return_value = existing_title_item
|
||||
|
||||
command = UpdateTitleLibraryCommand(
|
||||
title_id="existing-title-001",
|
||||
user_id="user-001",
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
assert result.name == "旧标题"
|
||||
assert result.text == "旧文本"
|
||||
assert result.category == "旧分类"
|
||||
assert result.is_active is True
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 4. DeleteTitleLibraryUseCase 测试
|
||||
# ===========================================================================
|
||||
|
||||
class TestDeleteTitleLibraryUseCase:
|
||||
"""标题库删除 UseCase 测试"""
|
||||
|
||||
def test_delete_success(self, mock_repo):
|
||||
"""测试删除成功"""
|
||||
mock_repo.delete.return_value = True
|
||||
use_case = DeleteTitleLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("title-001", "user-001")
|
||||
|
||||
assert result is True
|
||||
mock_repo.delete.assert_called_once_with("title-001", "user-001")
|
||||
|
||||
def test_delete_not_found(self, mock_repo):
|
||||
"""测试删除不存在的条目"""
|
||||
mock_repo.delete.return_value = False
|
||||
use_case = DeleteTitleLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("nonexistent", "user-001")
|
||||
|
||||
assert result is False
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 5. GetTitleLibraryUseCase 测试
|
||||
# ===========================================================================
|
||||
|
||||
class TestGetTitleLibraryUseCase:
|
||||
"""标题库查询 UseCase 测试"""
|
||||
|
||||
def test_get_existing(self, mock_repo):
|
||||
"""测试查询存在的条目"""
|
||||
expected = TitleLibraryItem(
|
||||
id="t-001",
|
||||
user_id="user-001",
|
||||
name="测试",
|
||||
text="文本",
|
||||
)
|
||||
mock_repo.get.return_value = expected
|
||||
use_case = GetTitleLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("t-001", "user-001")
|
||||
|
||||
assert result is not None
|
||||
assert result.id == "t-001"
|
||||
mock_repo.get.assert_called_once_with("t-001", "user-001")
|
||||
|
||||
def test_get_not_found(self, mock_repo):
|
||||
"""测试查询不存在的条目"""
|
||||
mock_repo.get.return_value = None
|
||||
use_case = GetTitleLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("nonexistent", "user-001")
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 6. ListTitleLibraryUseCase 测试
|
||||
# ===========================================================================
|
||||
|
||||
class TestListTitleLibraryUseCase:
|
||||
"""标题库列表 UseCase 测试"""
|
||||
|
||||
def test_list_default(self, mock_repo):
|
||||
"""测试默认列表查询"""
|
||||
items = [
|
||||
TitleLibraryItem(id="t1", user_id="user-001", name="A", text="a"),
|
||||
TitleLibraryItem(id="t2", user_id="user-001", name="B", text="b"),
|
||||
]
|
||||
mock_repo.list_by_user.return_value = items
|
||||
use_case = ListTitleLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("user-001")
|
||||
|
||||
assert len(result) == 2
|
||||
mock_repo.list_by_user.assert_called_once_with("user-001", category=None, skip=0, limit=50)
|
||||
|
||||
def test_list_with_category_filter(self, mock_repo):
|
||||
"""测试按分类筛选"""
|
||||
mock_repo.list_by_user.return_value = []
|
||||
use_case = ListTitleLibraryUseCase(repository=mock_repo)
|
||||
|
||||
use_case.execute("user-001", category="新闻", skip=5, limit=10)
|
||||
|
||||
mock_repo.list_by_user.assert_called_once_with(
|
||||
"user-001", category="新闻", skip=5, limit=10
|
||||
)
|
||||
|
||||
def test_list_empty(self, mock_repo):
|
||||
"""测试空列表"""
|
||||
mock_repo.list_by_user.return_value = []
|
||||
use_case = ListTitleLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("user-001")
|
||||
|
||||
assert result == []
|
||||
@@ -0,0 +1,685 @@
|
||||
"""
|
||||
配音库(Voice Library)Use Case 回归测试
|
||||
|
||||
测试目标:
|
||||
1. CreateVoiceLibraryUseCase - 创建配音库条目,验证 voice_id 字段映射正确(PR#74 P0 bug 修复)
|
||||
2. UpdateVoiceLibraryUseCase - 更新配音库条目,验证 voice_id 字段映射正确
|
||||
3. 配额逻辑覆盖 - free=10, basic=100, premium=100
|
||||
4. 边界条件与异常场景
|
||||
"""
|
||||
|
||||
from unittest.mock import Mock, MagicMock, call
|
||||
|
||||
import pytest
|
||||
|
||||
from packages.application.voice_library.commands import (
|
||||
CreateVoiceLibraryCommand,
|
||||
UpdateVoiceLibraryCommand,
|
||||
)
|
||||
from packages.application.voice_library.use_cases import (
|
||||
CreateVoiceLibraryUseCase,
|
||||
UpdateVoiceLibraryUseCase,
|
||||
DeleteVoiceLibraryUseCase,
|
||||
GetVoiceLibraryUseCase,
|
||||
ListVoiceLibraryUseCase,
|
||||
QuotaExceededError,
|
||||
NotFoundError,
|
||||
)
|
||||
from packages.domain.voice_library import VoiceLibraryItem
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fixtures
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@pytest.fixture
|
||||
def mock_repo():
|
||||
"""创建 Mock 仓储"""
|
||||
repo = Mock()
|
||||
repo.count_by_user = Mock(return_value=0)
|
||||
repo.create = Mock(side_effect=lambda item: item)
|
||||
repo.update = Mock(side_effect=lambda item: item)
|
||||
repo.get = Mock(return_value=None)
|
||||
repo.delete = Mock(return_value=True)
|
||||
repo.list_by_user = Mock(return_value=[])
|
||||
return repo
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def create_use_case(mock_repo):
|
||||
return CreateVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def update_use_case(mock_repo):
|
||||
return UpdateVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def sample_create_command():
|
||||
"""标准创建命令"""
|
||||
return CreateVoiceLibraryCommand(
|
||||
user_id="user-001",
|
||||
name="测试配音",
|
||||
text="你好世界",
|
||||
voice_provider="aliyun",
|
||||
voice_id="voice-abc-123",
|
||||
voice_name="小云",
|
||||
audio_url="https://oss.example.com/audio/abc.wav",
|
||||
duration=3.5,
|
||||
file_size=56000,
|
||||
status="completed",
|
||||
project_id="proj-001",
|
||||
tags=["测试", "中文"],
|
||||
metadata_={"source": "unit_test"},
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def existing_voice_item():
|
||||
"""模拟已存在的配音条目"""
|
||||
return VoiceLibraryItem(
|
||||
id="existing-voice-001",
|
||||
user_id="user-001",
|
||||
name="旧配音",
|
||||
text="旧文本",
|
||||
voice_provider="old_provider",
|
||||
voice_id="old-voice-id",
|
||||
voice_name="旧声音",
|
||||
audio_url="https://oss.example.com/old.wav",
|
||||
duration=1.0,
|
||||
file_size=16000,
|
||||
status="completed",
|
||||
project_id="proj-001",
|
||||
tags=["旧"],
|
||||
metadata_={},
|
||||
)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 1. CreateVoiceLibraryUseCase 测试
|
||||
# ===========================================================================
|
||||
|
||||
class TestCreateVoiceLibraryUseCase:
|
||||
"""配音库创建 UseCase 测试"""
|
||||
|
||||
def test_create_success_all_fields(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""测试创建成功 - 所有字段完整传入"""
|
||||
result = create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert result is not None
|
||||
assert result.user_id == "user-001"
|
||||
assert result.name == "测试配音"
|
||||
assert result.text == "你好世界"
|
||||
assert result.voice_provider == "aliyun"
|
||||
assert result.voice_name == "小云"
|
||||
assert result.audio_url == "https://oss.example.com/audio/abc.wav"
|
||||
assert result.duration == 3.5
|
||||
assert result.file_size == 56000
|
||||
assert result.status == "completed"
|
||||
assert result.project_id == "proj-001"
|
||||
assert result.tags == ["测试", "中文"]
|
||||
assert result.metadata_ == {"source": "unit_test"}
|
||||
|
||||
mock_repo.count_by_user.assert_called_once_with("user-001")
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_create_voice_id_field_mapping(self, create_use_case, mock_repo):
|
||||
"""
|
||||
【P0 回归】验证 voice_id 字段映射正确
|
||||
|
||||
PR#74 修复了 command.id 被错误使用的问题。
|
||||
此测试确保 CreateVoiceLibraryCommand 中的 voice_id 字段
|
||||
被正确传递到 VoiceLibraryItem 的 voice_id 属性上,
|
||||
而非被其他字段(如 item 自身的 id)覆盖。
|
||||
"""
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id="user-001",
|
||||
name="voice_id 回归测试",
|
||||
voice_id="specific-voice-id-xyz",
|
||||
voice_provider="azure",
|
||||
voice_name="Azure Xiaoxiao",
|
||||
)
|
||||
|
||||
result = create_use_case.execute(command, plan_name="free")
|
||||
|
||||
# 核心断言:voice_id 必须来自 command.voice_id
|
||||
assert result.voice_id == "specific-voice-id-xyz", \
|
||||
"voice_id 应来自 command.voice_id,而非其他字段"
|
||||
# 同时确保 item 自身生成的 id 与 voice_id 不同
|
||||
assert result.id != "specific-voice-id-xyz", \
|
||||
"item.id(UUID)不应与 voice_id 混淆"
|
||||
|
||||
def test_create_voice_id_empty_string(self, create_use_case, mock_repo):
|
||||
"""测试 voice_id 为空字符串的合法场景"""
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id="user-001",
|
||||
name="无 voice_id 配音",
|
||||
voice_id="",
|
||||
voice_provider="custom",
|
||||
)
|
||||
|
||||
result = create_use_case.execute(command, plan_name="free")
|
||||
|
||||
assert result.voice_id == ""
|
||||
|
||||
def test_create_default_values(self, create_use_case, mock_repo):
|
||||
"""测试默认值填充"""
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id="user-001",
|
||||
name="最小化创建",
|
||||
)
|
||||
|
||||
result = create_use_case.execute(command, plan_name="free")
|
||||
|
||||
assert result.text == ""
|
||||
assert result.voice_provider == ""
|
||||
assert result.voice_id == ""
|
||||
assert result.voice_name == ""
|
||||
assert result.audio_url == ""
|
||||
assert result.duration == 0
|
||||
assert result.file_size == 0
|
||||
assert result.status == "completed"
|
||||
assert result.project_id is None
|
||||
assert result.tags == []
|
||||
assert result.metadata_ == {}
|
||||
|
||||
def test_create_generates_uuid(self, create_use_case, mock_repo):
|
||||
"""测试创建时自动生成 UUID 作为 id"""
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id="user-001",
|
||||
name="UUID 测试",
|
||||
)
|
||||
|
||||
result = create_use_case.execute(command, plan_name="free")
|
||||
|
||||
assert result.id is not None
|
||||
assert len(result.id) == 32 # uuid4().hex 长度为 32
|
||||
assert result.id.isalnum()
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 2. 配额逻辑测试(Create 时的配额检查)
|
||||
# ===========================================================================
|
||||
|
||||
class TestCreateVoiceLibraryQuota:
|
||||
"""配音库创建配额检查测试"""
|
||||
|
||||
def test_quota_free_plan_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""free 套餐(上限10),当前 5 个,允许创建"""
|
||||
mock_repo.count_by_user.return_value = 5
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_free_plan_at_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""free 套餐(上限10),当前 10 个,拒绝创建"""
|
||||
mock_repo.count_by_user.return_value = 10
|
||||
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert exc_info.value.dimension == "max_voiceovers"
|
||||
assert exc_info.value.limit == 10
|
||||
assert exc_info.value.used == 10
|
||||
mock_repo.create.assert_not_called()
|
||||
|
||||
def test_quota_free_plan_over_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""free 套餐(上限10),当前 15 个,拒绝创建"""
|
||||
mock_repo.count_by_user.return_value = 15
|
||||
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert exc_info.value.dimension == "max_voiceovers"
|
||||
assert exc_info.value.limit == 10
|
||||
assert exc_info.value.used == 15
|
||||
|
||||
def test_quota_free_plan_just_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""free 套餐(上限10),当前 9 个,允许创建(边界)"""
|
||||
mock_repo.count_by_user.return_value = 9
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_basic_plan_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""basic 套餐(上限100),当前 50 个,允许创建"""
|
||||
mock_repo.count_by_user.return_value = 50
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="basic")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_basic_plan_at_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""basic 套餐(上限100),当前 100 个,拒绝创建"""
|
||||
mock_repo.count_by_user.return_value = 100
|
||||
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
create_use_case.execute(sample_create_command, plan_name="basic")
|
||||
|
||||
assert exc_info.value.dimension == "max_voiceovers"
|
||||
assert exc_info.value.limit == 100
|
||||
assert exc_info.value.used == 100
|
||||
|
||||
def test_quota_basic_plan_just_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""basic 套餐(上限100),当前 99 个,允许创建(边界)"""
|
||||
mock_repo.count_by_user.return_value = 99
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="basic")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_premium_plan_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""premium 套餐(上限100),当前 50 个,允许创建"""
|
||||
mock_repo.count_by_user.return_value = 50
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="premium")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_premium_plan_at_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""premium 套餐(上限100),当前 100 个,拒绝创建"""
|
||||
mock_repo.count_by_user.return_value = 100
|
||||
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
create_use_case.execute(sample_create_command, plan_name="premium")
|
||||
|
||||
assert exc_info.value.dimension == "max_voiceovers"
|
||||
assert exc_info.value.limit == 100
|
||||
assert exc_info.value.used == 100
|
||||
|
||||
def test_quota_premium_plan_just_under_limit(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""premium 套餐(上限100),当前 99 个,允许创建(边界)"""
|
||||
mock_repo.count_by_user.return_value = 99
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name="premium")
|
||||
|
||||
assert result is not None
|
||||
mock_repo.create.assert_called_once()
|
||||
|
||||
def test_quota_zero_usage(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""新用户零使用量,所有套餐均可创建"""
|
||||
mock_repo.count_by_user.return_value = 0
|
||||
|
||||
for plan in ["free", "basic", "premium"]:
|
||||
mock_repo.create.reset_mock()
|
||||
mock_repo.count_by_user.reset_mock()
|
||||
mock_repo.count_by_user.return_value = 0
|
||||
|
||||
result = create_use_case.execute(sample_create_command, plan_name=plan)
|
||||
assert result is not None, f"{plan} 套餐零使用量应允许创建"
|
||||
|
||||
def test_quota_unknown_plan_defaults_to_zero(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""未知套餐名默认配额为 0,即使 0 使用量也无法创建"""
|
||||
mock_repo.count_by_user.return_value = 0
|
||||
|
||||
with pytest.raises(QuotaExceededError):
|
||||
create_use_case.execute(sample_create_command, plan_name="unknown_plan")
|
||||
|
||||
def test_quota_exceeded_error_attributes(self, create_use_case, mock_repo, sample_create_command):
|
||||
"""QuotaExceededError 异常属性完整性"""
|
||||
mock_repo.count_by_user.return_value = 10
|
||||
|
||||
with pytest.raises(QuotaExceededError) as exc_info:
|
||||
create_use_case.execute(sample_create_command, plan_name="free")
|
||||
|
||||
err = exc_info.value
|
||||
assert hasattr(err, "dimension")
|
||||
assert hasattr(err, "limit")
|
||||
assert hasattr(err, "used")
|
||||
assert "max_voiceovers" in str(err)
|
||||
assert "10" in str(err)
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 3. UpdateVoiceLibraryUseCase 测试
|
||||
# ===========================================================================
|
||||
|
||||
class TestUpdateVoiceLibraryUseCase:
|
||||
"""配音库更新 UseCase 测试"""
|
||||
|
||||
def test_update_success_all_fields(self, update_use_case, mock_repo, existing_voice_item):
|
||||
"""测试全字段更新成功"""
|
||||
mock_repo.get.return_value = existing_voice_item
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="existing-voice-001",
|
||||
user_id="user-001",
|
||||
name="更新后的名称",
|
||||
text="更新后的文本",
|
||||
voice_provider="new_provider",
|
||||
voice_id="new-voice-id-456",
|
||||
voice_name="新声音",
|
||||
audio_url="https://oss.example.com/new.wav",
|
||||
duration=5.0,
|
||||
file_size=80000,
|
||||
status="processing",
|
||||
tags=["新标签"],
|
||||
metadata_={"updated": True},
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
assert result.name == "更新后的名称"
|
||||
assert result.text == "更新后的文本"
|
||||
assert result.voice_provider == "new_provider"
|
||||
assert result.voice_name == "新声音"
|
||||
assert result.audio_url == "https://oss.example.com/new.wav"
|
||||
assert result.duration == 5.0
|
||||
assert result.file_size == 80000
|
||||
assert result.status == "processing"
|
||||
assert result.tags == ["新标签"]
|
||||
assert result.metadata_ == {"updated": True}
|
||||
|
||||
mock_repo.update.assert_called_once()
|
||||
|
||||
def test_update_voice_id_field_mapping(self, update_use_case, mock_repo, existing_voice_item):
|
||||
"""
|
||||
【P0 回归】验证 update 时 voice_id 字段映射正确
|
||||
|
||||
PR#74 修复了 API 路由层将 command.id 错误传给 voice_id 的 bug。
|
||||
此测试确保 UpdateVoiceLibraryCommand 中 voice_id 字段
|
||||
被正确写入 VoiceLibraryItem.voice_id,而非被 item.id 覆盖。
|
||||
"""
|
||||
mock_repo.get.return_value = existing_voice_item
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="existing-voice-001",
|
||||
user_id="user-001",
|
||||
voice_id="completely-different-voice-id",
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
# 核心断言:voice_id 应被更新为新值
|
||||
assert result.voice_id == "completely-different-voice-id", \
|
||||
"voice_id 应被更新为 command.voice_id 的值"
|
||||
# item 自身的 id 保持不变
|
||||
assert result.id == "existing-voice-001"
|
||||
|
||||
def test_update_partial_only_voice_id(self, update_use_case, mock_repo, existing_voice_item):
|
||||
"""测试仅更新 voice_id 一个字段"""
|
||||
mock_repo.get.return_value = existing_voice_item
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="existing-voice-001",
|
||||
user_id="user-001",
|
||||
voice_id="only-voice-id-changed",
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
assert result.voice_id == "only-voice-id-changed"
|
||||
# 其他字段保持不变
|
||||
assert result.name == "旧配音"
|
||||
assert result.text == "旧文本"
|
||||
assert result.voice_provider == "old_provider"
|
||||
assert result.voice_name == "旧声音"
|
||||
assert result.audio_url == "https://oss.example.com/old.wav"
|
||||
assert result.duration == 1.0
|
||||
assert result.file_size == 16000
|
||||
|
||||
def test_update_partial_only_name(self, update_use_case, mock_repo, existing_voice_item):
|
||||
"""测试仅更新 name"""
|
||||
mock_repo.get.return_value = existing_voice_item
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="existing-voice-001",
|
||||
user_id="user-001",
|
||||
name="仅改名",
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
assert result.name == "仅改名"
|
||||
assert result.voice_id == "old-voice-id" # voice_id 不变
|
||||
|
||||
def test_update_not_found(self, update_use_case, mock_repo):
|
||||
"""测试更新不存在的条目"""
|
||||
mock_repo.get.return_value = None
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="nonexistent-id",
|
||||
user_id="user-001",
|
||||
name="不存在",
|
||||
)
|
||||
|
||||
with pytest.raises(NotFoundError, match="nonexistent-id"):
|
||||
update_use_case.execute(command)
|
||||
|
||||
mock_repo.update.assert_not_called()
|
||||
|
||||
def test_update_wrong_user(self, update_use_case, mock_repo):
|
||||
"""测试用户隔离 - 不能更新其他用户的条目"""
|
||||
mock_repo.get.return_value = None # repo 返回 None 表示找不到(不同 user_id)
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="existing-voice-001",
|
||||
user_id="other-user-999",
|
||||
name="恶意修改",
|
||||
)
|
||||
|
||||
with pytest.raises(NotFoundError):
|
||||
update_use_case.execute(command)
|
||||
|
||||
def test_update_none_fields_not_changed(self, update_use_case, mock_repo, existing_voice_item):
|
||||
"""测试 None 字段不覆盖原有值"""
|
||||
mock_repo.get.return_value = existing_voice_item
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="existing-voice-001",
|
||||
user_id="user-001",
|
||||
# 所有可选字段保持 None
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
# 所有字段应保持不变
|
||||
assert result.name == "旧配音"
|
||||
assert result.text == "旧文本"
|
||||
assert result.voice_id == "old-voice-id"
|
||||
assert result.voice_provider == "old_provider"
|
||||
assert result.voice_name == "旧声音"
|
||||
assert result.audio_url == "https://oss.example.com/old.wav"
|
||||
assert result.duration == 1.0
|
||||
assert result.file_size == 16000
|
||||
assert result.status == "completed"
|
||||
|
||||
def test_update_voice_id_empty_string(self, update_use_case, mock_repo, existing_voice_item):
|
||||
"""测试 voice_id 更新为空字符串(合法场景:清除 voice_id)"""
|
||||
mock_repo.get.return_value = existing_voice_item
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="existing-voice-001",
|
||||
user_id="user-001",
|
||||
voice_id="",
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
assert result.voice_id == ""
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 4. DeleteVoiceLibraryUseCase 测试
|
||||
# ===========================================================================
|
||||
|
||||
class TestDeleteVoiceLibraryUseCase:
|
||||
"""配音库删除 UseCase 测试"""
|
||||
|
||||
def test_delete_success(self, mock_repo):
|
||||
"""测试删除成功"""
|
||||
mock_repo.delete.return_value = True
|
||||
use_case = DeleteVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("voice-001", "user-001")
|
||||
|
||||
assert result is True
|
||||
mock_repo.delete.assert_called_once_with("voice-001", "user-001")
|
||||
|
||||
def test_delete_not_found(self, mock_repo):
|
||||
"""测试删除不存在的条目"""
|
||||
mock_repo.delete.return_value = False
|
||||
use_case = DeleteVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("nonexistent", "user-001")
|
||||
|
||||
assert result is False
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 5. GetVoiceLibraryUseCase 测试
|
||||
# ===========================================================================
|
||||
|
||||
class TestGetVoiceLibraryUseCase:
|
||||
"""配音库查询 UseCase 测试"""
|
||||
|
||||
def test_get_existing(self, mock_repo):
|
||||
"""测试查询存在的条目"""
|
||||
expected = VoiceLibraryItem(
|
||||
id="v-001",
|
||||
user_id="user-001",
|
||||
name="测试",
|
||||
voice_id="voice-xyz",
|
||||
)
|
||||
mock_repo.get.return_value = expected
|
||||
use_case = GetVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("v-001", "user-001")
|
||||
|
||||
assert result is not None
|
||||
assert result.id == "v-001"
|
||||
assert result.voice_id == "voice-xyz"
|
||||
mock_repo.get.assert_called_once_with("v-001", "user-001")
|
||||
|
||||
def test_get_not_found(self, mock_repo):
|
||||
"""测试查询不存在的条目"""
|
||||
mock_repo.get.return_value = None
|
||||
use_case = GetVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("nonexistent", "user-001")
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 6. ListVoiceLibraryUseCase 测试
|
||||
# ===========================================================================
|
||||
|
||||
class TestListVoiceLibraryUseCase:
|
||||
"""配音库列表 UseCase 测试"""
|
||||
|
||||
def test_list_default(self, mock_repo):
|
||||
"""测试默认列表查询"""
|
||||
items = [
|
||||
VoiceLibraryItem(id="v1", user_id="user-001", name="A"),
|
||||
VoiceLibraryItem(id="v2", user_id="user-001", name="B"),
|
||||
]
|
||||
mock_repo.list_by_user.return_value = items
|
||||
use_case = ListVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("user-001")
|
||||
|
||||
assert len(result) == 2
|
||||
mock_repo.list_by_user.assert_called_once_with("user-001", status=None, skip=0, limit=50)
|
||||
|
||||
def test_list_with_status_filter(self, mock_repo):
|
||||
"""测试按状态筛选"""
|
||||
mock_repo.list_by_user.return_value = []
|
||||
use_case = ListVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
use_case.execute("user-001", status="completed", skip=10, limit=20)
|
||||
|
||||
mock_repo.list_by_user.assert_called_once_with(
|
||||
"user-001", status="completed", skip=10, limit=20
|
||||
)
|
||||
|
||||
def test_list_empty(self, mock_repo):
|
||||
"""测试空列表"""
|
||||
mock_repo.list_by_user.return_value = []
|
||||
use_case = ListVoiceLibraryUseCase(repository=mock_repo)
|
||||
|
||||
result = use_case.execute("user-001")
|
||||
|
||||
assert result == []
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# 7. voice_id 与 id 字段隔离专项回归测试
|
||||
# ===========================================================================
|
||||
|
||||
class TestVoiceIdFieldIsolation:
|
||||
"""
|
||||
PR#74 P0 Bug 回归:voice_id 与 item.id 字段隔离
|
||||
|
||||
原 bug:API 路由层误将 command.id(item 主键)用作 voice_id,
|
||||
导致 voice_id 字段值错误。本测试类从 UseCase 层验证
|
||||
这两个字段在整个 CRUD 生命周期中互不干扰。
|
||||
"""
|
||||
|
||||
def test_create_id_and_voice_id_are_independent(self, create_use_case, mock_repo):
|
||||
"""创建时 id 自动生成,voice_id 来自 command"""
|
||||
command = CreateVoiceLibraryCommand(
|
||||
user_id="user-001",
|
||||
name="隔离测试",
|
||||
voice_id="tts-voice-001",
|
||||
voice_provider="openai",
|
||||
)
|
||||
|
||||
result = create_use_case.execute(command, plan_name="free")
|
||||
|
||||
assert result.id != result.voice_id, "id 和 voice_id 应为不同值"
|
||||
assert result.voice_id == "tts-voice-001"
|
||||
assert len(result.id) == 32 # UUID hex
|
||||
|
||||
def test_update_voice_id_does_not_change_id(self, update_use_case, mock_repo):
|
||||
"""更新 voice_id 不影响 item 主键 id"""
|
||||
existing = VoiceLibraryItem(
|
||||
id="stable-id-001",
|
||||
user_id="user-001",
|
||||
name="测试",
|
||||
voice_id="old-voice",
|
||||
)
|
||||
mock_repo.get.return_value = existing
|
||||
|
||||
command = UpdateVoiceLibraryCommand(
|
||||
id="stable-id-001",
|
||||
user_id="user-001",
|
||||
voice_id="new-voice-999",
|
||||
)
|
||||
|
||||
result = update_use_case.execute(command)
|
||||
|
||||
assert result.id == "stable-id-001", "item 主键 id 不应改变"
|
||||
assert result.voice_id == "new-voice-999", "voice_id 应被更新"
|
||||
|
||||
def test_create_then_update_voice_id_preserves_id(self, create_use_case, update_use_case, mock_repo):
|
||||
"""创建后再更新 voice_id,id 始终不变"""
|
||||
# 创建
|
||||
create_cmd = CreateVoiceLibraryCommand(
|
||||
user_id="user-001",
|
||||
name="生命周期测试",
|
||||
voice_id="initial-voice",
|
||||
)
|
||||
created = create_use_case.execute(create_cmd, plan_name="free")
|
||||
original_id = created.id
|
||||
|
||||
# 更新
|
||||
mock_repo.get.return_value = created
|
||||
update_cmd = UpdateVoiceLibraryCommand(
|
||||
id=original_id,
|
||||
user_id="user-001",
|
||||
voice_id="updated-voice",
|
||||
)
|
||||
updated = update_use_case.execute(update_cmd)
|
||||
|
||||
assert updated.id == original_id, "经过创建和更新,id 应保持一致"
|
||||
assert updated.voice_id == "updated-voice"
|
||||
assert updated.voice_id != "initial-voice"
|
||||
Reference in New Issue
Block a user