feat: Phase 1 - 核心重构(去Project层/标题库API/配音库API/清理废弃代码) #74

Merged
xiaoxia merged 4 commits from feat/phase1-core-refactor into develop 2026-06-28 15:50:11 +08:00
57 changed files with 2551 additions and 3189 deletions
@@ -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
View File
@@ -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"],
)
+21 -9
View File
@@ -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(
+4 -4
View File
@@ -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)
-287
View File
@@ -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})
+18 -6
View File
@@ -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:
+5 -24
View File
@@ -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,
)
-91
View File
@@ -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))
+155
View File
@@ -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")
+170
View File
@@ -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")
+21 -10
View File
@@ -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)
-60
View File
@@ -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 = ""
-2
View File
@@ -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
-35
View File
@@ -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
+42
View File
@@ -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
+57
View File
@@ -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):
+3 -3
View File
@@ -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}")
+3 -3
View File
@@ -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 -13
View File
@@ -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",
]
-379
View File
@@ -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}")
+11
View File
@@ -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" # 口播+画中画组合模式
-70
View File
@@ -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))
-3
View File
@@ -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(),
)
-232
View File
@@ -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)
+23
View File
@@ -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))
+27
View File
@@ -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))
+4 -8
View File
@@ -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) == []
+488
View File
@@ -0,0 +1,488 @@
"""
标题库(Title LibraryUse 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 == []
+685
View File
@@ -0,0 +1,685 @@
"""
配音库(Voice LibraryUse 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.idUUID)不应与 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 字段隔离
原 bugAPI 路由层误将 command.iditem 主键)用作 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"