From a43ddb4b63730b523b370e389ce934c1eda80c60 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?API=E6=96=87=E6=A1=A3=E7=BB=B4=E6=8A=A4Agent?= Date: Sun, 28 Jun 2026 15:02:10 +0800 Subject: [PATCH 1/4] =?UTF-8?q?feat:=20Phase=201=20-=20=E6=A0=B8=E5=BF=83?= =?UTF-8?q?=E9=87=8D=E6=9E=84=EF=BC=88=E5=8E=BBProject=E5=B1=82/=E6=A0=87?= =?UTF-8?q?=E9=A2=98=E5=BA=93API/=E9=85=8D=E9=9F=B3=E5=BA=93API/=E6=B8=85?= =?UTF-8?q?=E7=90=86=E5=BA=9F=E5=BC=83=E4=BB=A3=E7=A0=81=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ## 主要变更 ### 1. 清理废弃代码(18个文件删除) - 删除 TaskModel/MilestoneModel/TaskIssueModel 及相关文件 - 删除 ProjectTitleModel/EditPlanModel/EditPlanClipModel 及相关文件 - 清理 domain/ports/adapters/application/api 各层引用 - 从 GenerationTaskModel 移除 edit_plan_id 字段 ### 2. 新建标题库 API(/api/v1/titles) - Domain: TitleLibraryItem 数据类 - Ports: TitleLibraryRepository 接口 - Adapters: SQLAlchemy 实现(软删除) - Application: CRUD Use Cases + 配额检查(max_titles: free=50, basic=500, premium=500) - API: GET/POST/PUT/DELETE 端点 - Schema: Pydantic 请求/响应模型 ### 3. 新建配音库 API(/api/v1/voices) - Domain: VoiceLibraryItem 数据类 - Ports: VoiceLibraryRepository 接口 - Adapters: SQLAlchemy 实现(软删除) - Application: CRUD Use Cases + 配额检查(max_voiceovers: free=10, basic=100, premium=100) - API: GET/POST/PUT/DELETE 端点 - Schema: Pydantic 请求/响应模型 ### 4. 去掉 Project 层依赖 - 修复 authenticated_user.id → authenticated_user.user.id bug - asset_libraries.py: project_id 改为可选查询参数 - generated_videos.py: project_id 改为可选查询参数 - 无 project_id 时通过 find_accessible_projects 获取用户可访问的所有项目 ### 5. 数据库迁移 - 创建 011_phase1_core_refactor.py - 删除 6 个废弃表:tasks, milestones, task_issues, project_titles, edit_plans, edit_plan_clips - 从 generation_tasks 表删除 edit_plan_id 列 ### 6. 其他改进 - 迁移 EditingMode 到独立模块 packages/domain/editing_mode.py - 注册 titles_router 和 voices_router - 添加 get_title_library_repository 和 get_voice_library_repository 依赖 - 更新 domain/ports __init__.py 导出新实体和仓储接口 ## 技术细节 - 遵循六边形架构模式 - 配额检查通过 QuotaRegistry 实现 - 软删除:标题库用 is_active=False,配音库用 status='deleted' - 配音库支持可选的 project_id 关联 ## 破坏性变更 - 删除 6 个废弃表(需先备份数据) - 删除 /api/v1/edit-plans, /api/v1/project-titles, /api/v1/project-management 端点 - generation_tasks API 不再包含 edit_plan_id 字段 --- alembic/versions/011_phase1_core_refactor.py | 144 ++++++ apps/api/app/api/router.py | 23 +- apps/api/app/api/routes/asset_libraries.py | 30 +- apps/api/app/api/routes/assets.py | 8 +- apps/api/app/api/routes/edit_plans.py | 287 ----------- apps/api/app/api/routes/generated_videos.py | 24 +- apps/api/app/api/routes/generation_tasks.py | 29 +- apps/api/app/api/routes/project_management.py | 466 ------------------ apps/api/app/api/routes/project_titles.py | 91 ---- apps/api/app/api/routes/titles.py | 155 ++++++ apps/api/app/api/routes/voices.py | 170 +++++++ apps/api/app/dependencies.py | 31 +- apps/api/app/schemas/edit_plan.py | 60 --- apps/api/app/schemas/generation_task.py | 2 - apps/api/app/schemas/project_title.py | 35 -- apps/api/app/schemas/title_library.py | 42 ++ apps/api/app/schemas/voice_library.py | 57 +++ apps/worker/video_processing/editing_modes.py | 2 +- apps/worker/worker_app/core/title_usage.py | 6 +- .../worker_app/tasks/edit_plan_generator.py | 358 -------------- apps/worker/worker_app/tasks/generation.py | 6 +- .../project_management_repositories.py | 92 ---- .../generation_task_repository.py | 3 - packages/adapters/sqlalchemy_impl/models.py | 87 ---- .../project_management_repositories.py | 241 --------- .../project_title_repository.py | 52 -- .../title_library_repository.py | 114 +++++ .../voice_library_repository.py | 125 +++++ packages/adapters/sqlite_tracker/__init__.py | 5 - .../project_management_repositories.py | 201 -------- .../application/get_task_detail_use_case.py | 17 - .../project_management_use_cases.py | 146 ------ .../application/title_library/__init__.py | 20 + .../application/title_library/commands.py | 29 ++ .../application/title_library/use_cases.py | 111 +++++ packages/application/update_task_use_case.py | 35 -- .../application/voice_library/__init__.py | 20 + .../application/voice_library/commands.py | 39 ++ .../application/voice_library/use_cases.py | 125 +++++ packages/domain/__init__.py | 18 +- packages/domain/edit_plan.py | 379 -------------- packages/domain/editing_mode.py | 11 + packages/domain/entities.py | 70 --- packages/domain/generation_task.py | 3 - packages/domain/project_management.py | 232 --------- packages/domain/title_library.py | 23 + packages/domain/voice_library.py | 27 + packages/ports/__init__.py | 12 +- .../ports/project_management_repositories.py | 102 ---- packages/ports/project_title_repository.py | 9 - packages/ports/title_library_repository.py | 36 ++ packages/ports/voice_library_repository.py | 35 ++ 52 files changed, 1378 insertions(+), 3067 deletions(-) create mode 100644 alembic/versions/011_phase1_core_refactor.py delete mode 100644 apps/api/app/api/routes/edit_plans.py delete mode 100644 apps/api/app/api/routes/project_management.py delete mode 100644 apps/api/app/api/routes/project_titles.py create mode 100644 apps/api/app/api/routes/titles.py create mode 100644 apps/api/app/api/routes/voices.py delete mode 100644 apps/api/app/schemas/edit_plan.py delete mode 100644 apps/api/app/schemas/project_title.py create mode 100644 apps/api/app/schemas/title_library.py create mode 100644 apps/api/app/schemas/voice_library.py delete mode 100644 apps/worker/worker_app/tasks/edit_plan_generator.py delete mode 100644 packages/adapters/in_memory/project_management_repositories.py delete mode 100644 packages/adapters/sqlalchemy_impl/project_management_repositories.py delete mode 100644 packages/adapters/sqlalchemy_impl/project_title_repository.py create mode 100644 packages/adapters/sqlalchemy_impl/title_library_repository.py create mode 100644 packages/adapters/sqlalchemy_impl/voice_library_repository.py delete mode 100644 packages/adapters/sqlite_tracker/project_management_repositories.py delete mode 100644 packages/application/get_task_detail_use_case.py delete mode 100644 packages/application/project_management_use_cases.py create mode 100644 packages/application/title_library/__init__.py create mode 100644 packages/application/title_library/commands.py create mode 100644 packages/application/title_library/use_cases.py delete mode 100644 packages/application/update_task_use_case.py create mode 100644 packages/application/voice_library/__init__.py create mode 100644 packages/application/voice_library/commands.py create mode 100644 packages/application/voice_library/use_cases.py delete mode 100644 packages/domain/edit_plan.py create mode 100644 packages/domain/editing_mode.py delete mode 100644 packages/domain/project_management.py create mode 100644 packages/domain/title_library.py create mode 100644 packages/domain/voice_library.py delete mode 100644 packages/ports/project_management_repositories.py delete mode 100644 packages/ports/project_title_repository.py create mode 100644 packages/ports/title_library_repository.py create mode 100644 packages/ports/voice_library_repository.py diff --git a/alembic/versions/011_phase1_core_refactor.py b/alembic/versions/011_phase1_core_refactor.py new file mode 100644 index 000000000..52e64c875 --- /dev/null +++ b/alembic/versions/011_phase1_core_refactor.py @@ -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() + ) + """)) diff --git a/apps/api/app/api/router.py b/apps/api/app/api/router.py index 0b7decf45..049378ae5 100644 --- a/apps/api/app/api/router.py +++ b/apps/api/app/api/router.py @@ -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"], ) diff --git a/apps/api/app/api/routes/asset_libraries.py b/apps/api/app/api/routes/asset_libraries.py index 1bd0acaa5..c8d25228a 100644 --- a/apps/api/app/api/routes/asset_libraries.py +++ b/apps/api/app/api/routes/asset_libraries.py @@ -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( diff --git a/apps/api/app/api/routes/assets.py b/apps/api/app/api/routes/assets.py index 7fa7ef839..aa37bfe99 100644 --- a/apps/api/app/api/routes/assets.py +++ b/apps/api/app/api/routes/assets.py @@ -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) diff --git a/apps/api/app/api/routes/edit_plans.py b/apps/api/app/api/routes/edit_plans.py deleted file mode 100644 index ce446cb28..000000000 --- a/apps/api/app/api/routes/edit_plans.py +++ /dev/null @@ -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}) diff --git a/apps/api/app/api/routes/generated_videos.py b/apps/api/app/api/routes/generated_videos.py index 092b07153..30a1b6aa8 100644 --- a/apps/api/app/api/routes/generated_videos.py +++ b/apps/api/app/api/routes/generated_videos.py @@ -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: diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index ceb3f204c..ab1721442 100644 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -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]) diff --git a/apps/api/app/api/routes/project_management.py b/apps/api/app/api/routes/project_management.py deleted file mode 100644 index b346b3932..000000000 --- a/apps/api/app/api/routes/project_management.py +++ /dev/null @@ -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, - ) diff --git a/apps/api/app/api/routes/project_titles.py b/apps/api/app/api/routes/project_titles.py deleted file mode 100644 index 90acea423..000000000 --- a/apps/api/app/api/routes/project_titles.py +++ /dev/null @@ -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)) diff --git a/apps/api/app/api/routes/titles.py b/apps/api/app/api/routes/titles.py new file mode 100644 index 000000000..5c4bdd3ca --- /dev/null +++ b/apps/api/app/api/routes/titles.py @@ -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") diff --git a/apps/api/app/api/routes/voices.py b/apps/api/app/api/routes/voices.py new file mode 100644 index 000000000..fba24a7b5 --- /dev/null +++ b/apps/api/app/api/routes/voices.py @@ -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") diff --git a/apps/api/app/dependencies.py b/apps/api/app/dependencies.py index 000dab852..fe6a993e2 100644 --- a/apps/api/app/dependencies.py +++ b/apps/api/app/dependencies.py @@ -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) diff --git a/apps/api/app/schemas/edit_plan.py b/apps/api/app/schemas/edit_plan.py deleted file mode 100644 index b1059e7fa..000000000 --- a/apps/api/app/schemas/edit_plan.py +++ /dev/null @@ -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 = "" diff --git a/apps/api/app/schemas/generation_task.py b/apps/api/app/schemas/generation_task.py index 7ab783be4..0a6fec37b 100644 --- a/apps/api/app/schemas/generation_task.py +++ b/apps/api/app/schemas/generation_task.py @@ -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 diff --git a/apps/api/app/schemas/project_title.py b/apps/api/app/schemas/project_title.py deleted file mode 100644 index 59862cd98..000000000 --- a/apps/api/app/schemas/project_title.py +++ /dev/null @@ -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 diff --git a/apps/api/app/schemas/title_library.py b/apps/api/app/schemas/title_library.py new file mode 100644 index 000000000..8af89af77 --- /dev/null +++ b/apps/api/app/schemas/title_library.py @@ -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 diff --git a/apps/api/app/schemas/voice_library.py b/apps/api/app/schemas/voice_library.py new file mode 100644 index 000000000..f906e0554 --- /dev/null +++ b/apps/api/app/schemas/voice_library.py @@ -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 diff --git a/apps/worker/video_processing/editing_modes.py b/apps/worker/video_processing/editing_modes.py index 8d1b03033..3b3968a14 100644 --- a/apps/worker/video_processing/editing_modes.py +++ b/apps/worker/video_processing/editing_modes.py @@ -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): diff --git a/apps/worker/worker_app/core/title_usage.py b/apps/worker/worker_app/core/title_usage.py index 4c075a0f8..0d0ffc981 100644 --- a/apps/worker/worker_app/core/title_usage.py +++ b/apps/worker/worker_app/core/title_usage.py @@ -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) diff --git a/apps/worker/worker_app/tasks/edit_plan_generator.py b/apps/worker/worker_app/tasks/edit_plan_generator.py deleted file mode 100644 index 54788492e..000000000 --- a/apps/worker/worker_app/tasks/edit_plan_generator.py +++ /dev/null @@ -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}") diff --git a/apps/worker/worker_app/tasks/generation.py b/apps/worker/worker_app/tasks/generation.py index 5bf7e4fb0..8c8b3c313 100755 --- a/apps/worker/worker_app/tasks/generation.py +++ b/apps/worker/worker_app/tasks/generation.py @@ -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}" diff --git a/packages/adapters/in_memory/project_management_repositories.py b/packages/adapters/in_memory/project_management_repositories.py deleted file mode 100644 index 2f1132af0..000000000 --- a/packages/adapters/in_memory/project_management_repositories.py +++ /dev/null @@ -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) diff --git a/packages/adapters/sqlalchemy_impl/generation_task_repository.py b/packages/adapters/sqlalchemy_impl/generation_task_repository.py index 91f75449e..430200ebf 100644 --- a/packages/adapters/sqlalchemy_impl/generation_task_repository.py +++ b/packages/adapters/sqlalchemy_impl/generation_task_repository.py @@ -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 diff --git a/packages/adapters/sqlalchemy_impl/models.py b/packages/adapters/sqlalchemy_impl/models.py index d4aef271c..390db0dd4 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -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" diff --git a/packages/adapters/sqlalchemy_impl/project_management_repositories.py b/packages/adapters/sqlalchemy_impl/project_management_repositories.py deleted file mode 100644 index 4dbe3cebc..000000000 --- a/packages/adapters/sqlalchemy_impl/project_management_repositories.py +++ /dev/null @@ -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, - ) diff --git a/packages/adapters/sqlalchemy_impl/project_title_repository.py b/packages/adapters/sqlalchemy_impl/project_title_repository.py deleted file mode 100644 index 2e677072c..000000000 --- a/packages/adapters/sqlalchemy_impl/project_title_repository.py +++ /dev/null @@ -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 diff --git a/packages/adapters/sqlalchemy_impl/title_library_repository.py b/packages/adapters/sqlalchemy_impl/title_library_repository.py new file mode 100644 index 000000000..3b5a47d73 --- /dev/null +++ b/packages/adapters/sqlalchemy_impl/title_library_repository.py @@ -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, + ) diff --git a/packages/adapters/sqlalchemy_impl/voice_library_repository.py b/packages/adapters/sqlalchemy_impl/voice_library_repository.py new file mode 100644 index 000000000..c790e480e --- /dev/null +++ b/packages/adapters/sqlalchemy_impl/voice_library_repository.py @@ -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, + ) diff --git a/packages/adapters/sqlite_tracker/__init__.py b/packages/adapters/sqlite_tracker/__init__.py index 9f17f75ac..8ef8ef9d3 100644 --- a/packages/adapters/sqlite_tracker/__init__.py +++ b/packages/adapters/sqlite_tracker/__init__.py @@ -1,10 +1,5 @@ """SQLite Tracker Adapter""" -from .project_management_repositories import ( - SQLiteMilestoneRepository, - SQLiteTaskIssueRepository, - SQLiteTaskRepository, -) __all__ = [ "SQLiteTaskRepository", diff --git a/packages/adapters/sqlite_tracker/project_management_repositories.py b/packages/adapters/sqlite_tracker/project_management_repositories.py deleted file mode 100644 index dc920320f..000000000 --- a/packages/adapters/sqlite_tracker/project_management_repositories.py +++ /dev/null @@ -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 diff --git a/packages/application/get_task_detail_use_case.py b/packages/application/get_task_detail_use_case.py deleted file mode 100644 index f7ecf370f..000000000 --- a/packages/application/get_task_detail_use_case.py +++ /dev/null @@ -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 diff --git a/packages/application/project_management_use_cases.py b/packages/application/project_management_use_cases.py deleted file mode 100644 index b97a7dea4..000000000 --- a/packages/application/project_management_use_cases.py +++ /dev/null @@ -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) diff --git a/packages/application/title_library/__init__.py b/packages/application/title_library/__init__.py new file mode 100644 index 000000000..384a9411e --- /dev/null +++ b/packages/application/title_library/__init__.py @@ -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", +] diff --git a/packages/application/title_library/commands.py b/packages/application/title_library/commands.py new file mode 100644 index 000000000..d65acbf2e --- /dev/null +++ b/packages/application/title_library/commands.py @@ -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 diff --git a/packages/application/title_library/use_cases.py b/packages/application/title_library/use_cases.py new file mode 100644 index 000000000..f6fe8d4e4 --- /dev/null +++ b/packages/application/title_library/use_cases.py @@ -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 diff --git a/packages/application/update_task_use_case.py b/packages/application/update_task_use_case.py deleted file mode 100644 index 7c742284c..000000000 --- a/packages/application/update_task_use_case.py +++ /dev/null @@ -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 diff --git a/packages/application/voice_library/__init__.py b/packages/application/voice_library/__init__.py new file mode 100644 index 000000000..ef25f3343 --- /dev/null +++ b/packages/application/voice_library/__init__.py @@ -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", +] diff --git a/packages/application/voice_library/commands.py b/packages/application/voice_library/commands.py new file mode 100644 index 000000000..210a99f54 --- /dev/null +++ b/packages/application/voice_library/commands.py @@ -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 diff --git a/packages/application/voice_library/use_cases.py b/packages/application/voice_library/use_cases.py new file mode 100644 index 000000000..f03021d79 --- /dev/null +++ b/packages/application/voice_library/use_cases.py @@ -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.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.id is not None: + existing.voice_id = command.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 diff --git a/packages/domain/__init__.py b/packages/domain/__init__.py index c068bb87f..9ac6df2e0 100644 --- a/packages/domain/__init__.py +++ b/packages/domain/__init__.py @@ -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", ] diff --git a/packages/domain/edit_plan.py b/packages/domain/edit_plan.py deleted file mode 100644 index 7faf5d8db..000000000 --- a/packages/domain/edit_plan.py +++ /dev/null @@ -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}") diff --git a/packages/domain/editing_mode.py b/packages/domain/editing_mode.py new file mode 100644 index 000000000..9acc85a81 --- /dev/null +++ b/packages/domain/editing_mode.py @@ -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" # 口播+画中画组合模式 diff --git a/packages/domain/entities.py b/packages/domain/entities.py index 4de800777..db4dff14b 100644 --- a/packages/domain/entities.py +++ b/packages/domain/entities.py @@ -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)) diff --git a/packages/domain/generation_task.py b/packages/domain/generation_task.py index ffa0e0552..f10e8a2d4 100644 --- a/packages/domain/generation_task.py +++ b/packages/domain/generation_task.py @@ -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(), ) diff --git a/packages/domain/project_management.py b/packages/domain/project_management.py deleted file mode 100644 index 401cdd9f2..000000000 --- a/packages/domain/project_management.py +++ /dev/null @@ -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) diff --git a/packages/domain/title_library.py b/packages/domain/title_library.py new file mode 100644 index 000000000..65d75c5f6 --- /dev/null +++ b/packages/domain/title_library.py @@ -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)) diff --git a/packages/domain/voice_library.py b/packages/domain/voice_library.py new file mode 100644 index 000000000..683c38558 --- /dev/null +++ b/packages/domain/voice_library.py @@ -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)) diff --git a/packages/ports/__init__.py b/packages/ports/__init__.py index ece885bb6..5d915d17b 100644 --- a/packages/ports/__init__.py +++ b/packages/ports/__init__.py @@ -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", ] diff --git a/packages/ports/project_management_repositories.py b/packages/ports/project_management_repositories.py deleted file mode 100644 index 66dfa784b..000000000 --- a/packages/ports/project_management_repositories.py +++ /dev/null @@ -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 diff --git a/packages/ports/project_title_repository.py b/packages/ports/project_title_repository.py deleted file mode 100644 index 389585460..000000000 --- a/packages/ports/project_title_repository.py +++ /dev/null @@ -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: ... diff --git a/packages/ports/title_library_repository.py b/packages/ports/title_library_repository.py new file mode 100644 index 000000000..e74a50efc --- /dev/null +++ b/packages/ports/title_library_repository.py @@ -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: + ... diff --git a/packages/ports/voice_library_repository.py b/packages/ports/voice_library_repository.py new file mode 100644 index 000000000..f4922cb1a --- /dev/null +++ b/packages/ports/voice_library_repository.py @@ -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: + ... -- 2.54.0 From abe2bfab632aedeb9979232dcb0047e1715c7257 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?API=E6=96=87=E6=A1=A3=E7=BB=B4=E6=8A=A4Agent?= Date: Sun, 28 Jun 2026 15:33:00 +0800 Subject: [PATCH 2/4] fix: resolve P0 bugs in voice library use cases and clean up obsolete tests - Fix voice_id field mapping in CreateVoiceLibraryUseCase (line 60) - Fix voice_id field mapping in UpdateVoiceLibraryUseCase (lines 88-89) - Delete obsolete test files for removed ProjectTitle modules - Remove edit_plans.py reference from architecture boundaries test --- .../application/voice_library/use_cases.py | 6 +- tests/unit/test_architecture_boundaries.py | 1 - tests/unit/test_project_title_generation.py | 71 ------------------- tests/unit/test_project_title_repository.py | 50 ------------- 4 files changed, 3 insertions(+), 125 deletions(-) delete mode 100644 tests/unit/test_project_title_generation.py delete mode 100644 tests/unit/test_project_title_repository.py diff --git a/packages/application/voice_library/use_cases.py b/packages/application/voice_library/use_cases.py index f03021d79..2c3ae5b58 100644 --- a/packages/application/voice_library/use_cases.py +++ b/packages/application/voice_library/use_cases.py @@ -57,7 +57,7 @@ class CreateVoiceLibraryUseCase: name=command.name, text=command.text, voice_provider=command.voice_provider, - voice_id=command.id, + voice_id=command.voice_id, voice_name=command.voice_name, audio_url=command.audio_url, duration=command.duration, @@ -85,8 +85,8 @@ class UpdateVoiceLibraryUseCase: existing.text = command.text if command.voice_provider is not None: existing.voice_provider = command.voice_provider - if command.id is not None: - existing.voice_id = command.id + 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: diff --git a/tests/unit/test_architecture_boundaries.py b/tests/unit/test_architecture_boundaries.py index 016a8f8e5..4828e5f42 100644 --- a/tests/unit/test_architecture_boundaries.py +++ b/tests/unit/test_architecture_boundaries.py @@ -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"), } diff --git a/tests/unit/test_project_title_generation.py b/tests/unit/test_project_title_generation.py deleted file mode 100644 index b9b3877ea..000000000 --- a/tests/unit/test_project_title_generation.py +++ /dev/null @@ -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 diff --git a/tests/unit/test_project_title_repository.py b/tests/unit/test_project_title_repository.py deleted file mode 100644 index bd07fbbf3..000000000 --- a/tests/unit/test_project_title_repository.py +++ /dev/null @@ -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) == [] -- 2.54.0 From 6b819678d6acfea30d7d1658856272a32c8603a1 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sun, 28 Jun 2026 15:43:49 +0800 Subject: [PATCH 3/4] =?UTF-8?q?test:=20=E6=B7=BB=E5=8A=A0=E9=85=8D?= =?UTF-8?q?=E9=9F=B3=E5=BA=93=20UseCase=20=E5=9B=9E=E5=BD=92=E6=B5=8B?= =?UTF-8?q?=E8=AF=95=EF=BC=88PR#74=20P0=20bug=20=E4=BF=AE=E5=A4=8D?= =?UTF-8?q?=E9=AA=8C=E8=AF=81=20+=20=E9=85=8D=E9=A2=9D=E9=80=BB=E8=BE=91?= =?UTF-8?q?=E8=A6=86=E7=9B=96=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tests/unit/test_voice_library_use_cases.py | 685 +++++++++++++++++++++ 1 file changed, 685 insertions(+) create mode 100644 tests/unit/test_voice_library_use_cases.py diff --git a/tests/unit/test_voice_library_use_cases.py b/tests/unit/test_voice_library_use_cases.py new file mode 100644 index 000000000..6e1ccfd7d --- /dev/null +++ b/tests/unit/test_voice_library_use_cases.py @@ -0,0 +1,685 @@ +""" +配音库(Voice Library)Use Case 回归测试 + +测试目标: +1. CreateVoiceLibraryUseCase - 创建配音库条目,验证 voice_id 字段映射正确(PR#74 P0 bug 修复) +2. UpdateVoiceLibraryUseCase - 更新配音库条目,验证 voice_id 字段映射正确 +3. 配额逻辑覆盖 - free=10, basic=100, premium=100 +4. 边界条件与异常场景 +""" + +from unittest.mock import Mock, MagicMock, call + +import pytest + +from packages.application.voice_library.commands import ( + CreateVoiceLibraryCommand, + UpdateVoiceLibraryCommand, +) +from packages.application.voice_library.use_cases import ( + CreateVoiceLibraryUseCase, + UpdateVoiceLibraryUseCase, + DeleteVoiceLibraryUseCase, + GetVoiceLibraryUseCase, + ListVoiceLibraryUseCase, + QuotaExceededError, + NotFoundError, +) +from packages.domain.voice_library import VoiceLibraryItem + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + +@pytest.fixture +def mock_repo(): + """创建 Mock 仓储""" + repo = Mock() + repo.count_by_user = Mock(return_value=0) + repo.create = Mock(side_effect=lambda item: item) + repo.update = Mock(side_effect=lambda item: item) + repo.get = Mock(return_value=None) + repo.delete = Mock(return_value=True) + repo.list_by_user = Mock(return_value=[]) + return repo + + +@pytest.fixture +def create_use_case(mock_repo): + return CreateVoiceLibraryUseCase(repository=mock_repo) + + +@pytest.fixture +def update_use_case(mock_repo): + return UpdateVoiceLibraryUseCase(repository=mock_repo) + + +@pytest.fixture +def sample_create_command(): + """标准创建命令""" + return CreateVoiceLibraryCommand( + user_id="user-001", + name="测试配音", + text="你好世界", + voice_provider="aliyun", + voice_id="voice-abc-123", + voice_name="小云", + audio_url="https://oss.example.com/audio/abc.wav", + duration=3.5, + file_size=56000, + status="completed", + project_id="proj-001", + tags=["测试", "中文"], + metadata_={"source": "unit_test"}, + ) + + +@pytest.fixture +def existing_voice_item(): + """模拟已存在的配音条目""" + return VoiceLibraryItem( + id="existing-voice-001", + user_id="user-001", + name="旧配音", + text="旧文本", + voice_provider="old_provider", + voice_id="old-voice-id", + voice_name="旧声音", + audio_url="https://oss.example.com/old.wav", + duration=1.0, + file_size=16000, + status="completed", + project_id="proj-001", + tags=["旧"], + metadata_={}, + ) + + +# =========================================================================== +# 1. CreateVoiceLibraryUseCase 测试 +# =========================================================================== + +class TestCreateVoiceLibraryUseCase: + """配音库创建 UseCase 测试""" + + def test_create_success_all_fields(self, create_use_case, mock_repo, sample_create_command): + """测试创建成功 - 所有字段完整传入""" + result = create_use_case.execute(sample_create_command, plan_name="free") + + assert result is not None + assert result.user_id == "user-001" + assert result.name == "测试配音" + assert result.text == "你好世界" + assert result.voice_provider == "aliyun" + assert result.voice_name == "小云" + assert result.audio_url == "https://oss.example.com/audio/abc.wav" + assert result.duration == 3.5 + assert result.file_size == 56000 + assert result.status == "completed" + assert result.project_id == "proj-001" + assert result.tags == ["测试", "中文"] + assert result.metadata_ == {"source": "unit_test"} + + mock_repo.count_by_user.assert_called_once_with("user-001") + mock_repo.create.assert_called_once() + + def test_create_voice_id_field_mapping(self, create_use_case, mock_repo): + """ + 【P0 回归】验证 voice_id 字段映射正确 + + PR#74 修复了 command.id 被错误使用的问题。 + 此测试确保 CreateVoiceLibraryCommand 中的 voice_id 字段 + 被正确传递到 VoiceLibraryItem 的 voice_id 属性上, + 而非被其他字段(如 item 自身的 id)覆盖。 + """ + command = CreateVoiceLibraryCommand( + user_id="user-001", + name="voice_id 回归测试", + voice_id="specific-voice-id-xyz", + voice_provider="azure", + voice_name="Azure Xiaoxiao", + ) + + result = create_use_case.execute(command, plan_name="free") + + # 核心断言:voice_id 必须来自 command.voice_id + assert result.voice_id == "specific-voice-id-xyz", \ + "voice_id 应来自 command.voice_id,而非其他字段" + # 同时确保 item 自身生成的 id 与 voice_id 不同 + assert result.id != "specific-voice-id-xyz", \ + "item.id(UUID)不应与 voice_id 混淆" + + def test_create_voice_id_empty_string(self, create_use_case, mock_repo): + """测试 voice_id 为空字符串的合法场景""" + command = CreateVoiceLibraryCommand( + user_id="user-001", + name="无 voice_id 配音", + voice_id="", + voice_provider="custom", + ) + + result = create_use_case.execute(command, plan_name="free") + + assert result.voice_id == "" + + def test_create_default_values(self, create_use_case, mock_repo): + """测试默认值填充""" + command = CreateVoiceLibraryCommand( + user_id="user-001", + name="最小化创建", + ) + + result = create_use_case.execute(command, plan_name="free") + + assert result.text == "" + assert result.voice_provider == "" + assert result.voice_id == "" + assert result.voice_name == "" + assert result.audio_url == "" + assert result.duration == 0 + assert result.file_size == 0 + assert result.status == "completed" + assert result.project_id is None + assert result.tags == [] + assert result.metadata_ == {} + + def test_create_generates_uuid(self, create_use_case, mock_repo): + """测试创建时自动生成 UUID 作为 id""" + command = CreateVoiceLibraryCommand( + user_id="user-001", + name="UUID 测试", + ) + + result = create_use_case.execute(command, plan_name="free") + + assert result.id is not None + assert len(result.id) == 32 # uuid4().hex 长度为 32 + assert result.id.isalnum() + + +# =========================================================================== +# 2. 配额逻辑测试(Create 时的配额检查) +# =========================================================================== + +class TestCreateVoiceLibraryQuota: + """配音库创建配额检查测试""" + + def test_quota_free_plan_under_limit(self, create_use_case, mock_repo, sample_create_command): + """free 套餐(上限10),当前 5 个,允许创建""" + mock_repo.count_by_user.return_value = 5 + + result = create_use_case.execute(sample_create_command, plan_name="free") + + assert result is not None + mock_repo.create.assert_called_once() + + def test_quota_free_plan_at_limit(self, create_use_case, mock_repo, sample_create_command): + """free 套餐(上限10),当前 10 个,拒绝创建""" + mock_repo.count_by_user.return_value = 10 + + with pytest.raises(QuotaExceededError) as exc_info: + create_use_case.execute(sample_create_command, plan_name="free") + + assert exc_info.value.dimension == "max_voiceovers" + assert exc_info.value.limit == 10 + assert exc_info.value.used == 10 + mock_repo.create.assert_not_called() + + def test_quota_free_plan_over_limit(self, create_use_case, mock_repo, sample_create_command): + """free 套餐(上限10),当前 15 个,拒绝创建""" + mock_repo.count_by_user.return_value = 15 + + with pytest.raises(QuotaExceededError) as exc_info: + create_use_case.execute(sample_create_command, plan_name="free") + + assert exc_info.value.dimension == "max_voiceovers" + assert exc_info.value.limit == 10 + assert exc_info.value.used == 15 + + def test_quota_free_plan_just_under_limit(self, create_use_case, mock_repo, sample_create_command): + """free 套餐(上限10),当前 9 个,允许创建(边界)""" + mock_repo.count_by_user.return_value = 9 + + result = create_use_case.execute(sample_create_command, plan_name="free") + + assert result is not None + mock_repo.create.assert_called_once() + + def test_quota_basic_plan_under_limit(self, create_use_case, mock_repo, sample_create_command): + """basic 套餐(上限100),当前 50 个,允许创建""" + mock_repo.count_by_user.return_value = 50 + + result = create_use_case.execute(sample_create_command, plan_name="basic") + + assert result is not None + mock_repo.create.assert_called_once() + + def test_quota_basic_plan_at_limit(self, create_use_case, mock_repo, sample_create_command): + """basic 套餐(上限100),当前 100 个,拒绝创建""" + mock_repo.count_by_user.return_value = 100 + + with pytest.raises(QuotaExceededError) as exc_info: + create_use_case.execute(sample_create_command, plan_name="basic") + + assert exc_info.value.dimension == "max_voiceovers" + assert exc_info.value.limit == 100 + assert exc_info.value.used == 100 + + def test_quota_basic_plan_just_under_limit(self, create_use_case, mock_repo, sample_create_command): + """basic 套餐(上限100),当前 99 个,允许创建(边界)""" + mock_repo.count_by_user.return_value = 99 + + result = create_use_case.execute(sample_create_command, plan_name="basic") + + assert result is not None + mock_repo.create.assert_called_once() + + def test_quota_premium_plan_under_limit(self, create_use_case, mock_repo, sample_create_command): + """premium 套餐(上限100),当前 50 个,允许创建""" + mock_repo.count_by_user.return_value = 50 + + result = create_use_case.execute(sample_create_command, plan_name="premium") + + assert result is not None + mock_repo.create.assert_called_once() + + def test_quota_premium_plan_at_limit(self, create_use_case, mock_repo, sample_create_command): + """premium 套餐(上限100),当前 100 个,拒绝创建""" + mock_repo.count_by_user.return_value = 100 + + with pytest.raises(QuotaExceededError) as exc_info: + create_use_case.execute(sample_create_command, plan_name="premium") + + assert exc_info.value.dimension == "max_voiceovers" + assert exc_info.value.limit == 100 + assert exc_info.value.used == 100 + + def test_quota_premium_plan_just_under_limit(self, create_use_case, mock_repo, sample_create_command): + """premium 套餐(上限100),当前 99 个,允许创建(边界)""" + mock_repo.count_by_user.return_value = 99 + + result = create_use_case.execute(sample_create_command, plan_name="premium") + + assert result is not None + mock_repo.create.assert_called_once() + + def test_quota_zero_usage(self, create_use_case, mock_repo, sample_create_command): + """新用户零使用量,所有套餐均可创建""" + mock_repo.count_by_user.return_value = 0 + + for plan in ["free", "basic", "premium"]: + mock_repo.create.reset_mock() + mock_repo.count_by_user.reset_mock() + mock_repo.count_by_user.return_value = 0 + + result = create_use_case.execute(sample_create_command, plan_name=plan) + assert result is not None, f"{plan} 套餐零使用量应允许创建" + + def test_quota_unknown_plan_defaults_to_zero(self, create_use_case, mock_repo, sample_create_command): + """未知套餐名默认配额为 0,即使 0 使用量也无法创建""" + mock_repo.count_by_user.return_value = 0 + + with pytest.raises(QuotaExceededError): + create_use_case.execute(sample_create_command, plan_name="unknown_plan") + + def test_quota_exceeded_error_attributes(self, create_use_case, mock_repo, sample_create_command): + """QuotaExceededError 异常属性完整性""" + mock_repo.count_by_user.return_value = 10 + + with pytest.raises(QuotaExceededError) as exc_info: + create_use_case.execute(sample_create_command, plan_name="free") + + err = exc_info.value + assert hasattr(err, "dimension") + assert hasattr(err, "limit") + assert hasattr(err, "used") + assert "max_voiceovers" in str(err) + assert "10" in str(err) + + +# =========================================================================== +# 3. UpdateVoiceLibraryUseCase 测试 +# =========================================================================== + +class TestUpdateVoiceLibraryUseCase: + """配音库更新 UseCase 测试""" + + def test_update_success_all_fields(self, update_use_case, mock_repo, existing_voice_item): + """测试全字段更新成功""" + mock_repo.get.return_value = existing_voice_item + + command = UpdateVoiceLibraryCommand( + id="existing-voice-001", + user_id="user-001", + name="更新后的名称", + text="更新后的文本", + voice_provider="new_provider", + voice_id="new-voice-id-456", + voice_name="新声音", + audio_url="https://oss.example.com/new.wav", + duration=5.0, + file_size=80000, + status="processing", + tags=["新标签"], + metadata_={"updated": True}, + ) + + result = update_use_case.execute(command) + + assert result.name == "更新后的名称" + assert result.text == "更新后的文本" + assert result.voice_provider == "new_provider" + assert result.voice_name == "新声音" + assert result.audio_url == "https://oss.example.com/new.wav" + assert result.duration == 5.0 + assert result.file_size == 80000 + assert result.status == "processing" + assert result.tags == ["新标签"] + assert result.metadata_ == {"updated": True} + + mock_repo.update.assert_called_once() + + def test_update_voice_id_field_mapping(self, update_use_case, mock_repo, existing_voice_item): + """ + 【P0 回归】验证 update 时 voice_id 字段映射正确 + + PR#74 修复了 API 路由层将 command.id 错误传给 voice_id 的 bug。 + 此测试确保 UpdateVoiceLibraryCommand 中 voice_id 字段 + 被正确写入 VoiceLibraryItem.voice_id,而非被 item.id 覆盖。 + """ + mock_repo.get.return_value = existing_voice_item + + command = UpdateVoiceLibraryCommand( + id="existing-voice-001", + user_id="user-001", + voice_id="completely-different-voice-id", + ) + + result = update_use_case.execute(command) + + # 核心断言:voice_id 应被更新为新值 + assert result.voice_id == "completely-different-voice-id", \ + "voice_id 应被更新为 command.voice_id 的值" + # item 自身的 id 保持不变 + assert result.id == "existing-voice-001" + + def test_update_partial_only_voice_id(self, update_use_case, mock_repo, existing_voice_item): + """测试仅更新 voice_id 一个字段""" + mock_repo.get.return_value = existing_voice_item + + command = UpdateVoiceLibraryCommand( + id="existing-voice-001", + user_id="user-001", + voice_id="only-voice-id-changed", + ) + + result = update_use_case.execute(command) + + assert result.voice_id == "only-voice-id-changed" + # 其他字段保持不变 + assert result.name == "旧配音" + assert result.text == "旧文本" + assert result.voice_provider == "old_provider" + assert result.voice_name == "旧声音" + assert result.audio_url == "https://oss.example.com/old.wav" + assert result.duration == 1.0 + assert result.file_size == 16000 + + def test_update_partial_only_name(self, update_use_case, mock_repo, existing_voice_item): + """测试仅更新 name""" + mock_repo.get.return_value = existing_voice_item + + command = UpdateVoiceLibraryCommand( + id="existing-voice-001", + user_id="user-001", + name="仅改名", + ) + + result = update_use_case.execute(command) + + assert result.name == "仅改名" + assert result.voice_id == "old-voice-id" # voice_id 不变 + + def test_update_not_found(self, update_use_case, mock_repo): + """测试更新不存在的条目""" + mock_repo.get.return_value = None + + command = UpdateVoiceLibraryCommand( + id="nonexistent-id", + user_id="user-001", + name="不存在", + ) + + with pytest.raises(NotFoundError, match="nonexistent-id"): + update_use_case.execute(command) + + mock_repo.update.assert_not_called() + + def test_update_wrong_user(self, update_use_case, mock_repo): + """测试用户隔离 - 不能更新其他用户的条目""" + mock_repo.get.return_value = None # repo 返回 None 表示找不到(不同 user_id) + + command = UpdateVoiceLibraryCommand( + id="existing-voice-001", + user_id="other-user-999", + name="恶意修改", + ) + + with pytest.raises(NotFoundError): + update_use_case.execute(command) + + def test_update_none_fields_not_changed(self, update_use_case, mock_repo, existing_voice_item): + """测试 None 字段不覆盖原有值""" + mock_repo.get.return_value = existing_voice_item + + command = UpdateVoiceLibraryCommand( + id="existing-voice-001", + user_id="user-001", + # 所有可选字段保持 None + ) + + result = update_use_case.execute(command) + + # 所有字段应保持不变 + assert result.name == "旧配音" + assert result.text == "旧文本" + assert result.voice_id == "old-voice-id" + assert result.voice_provider == "old_provider" + assert result.voice_name == "旧声音" + assert result.audio_url == "https://oss.example.com/old.wav" + assert result.duration == 1.0 + assert result.file_size == 16000 + assert result.status == "completed" + + def test_update_voice_id_empty_string(self, update_use_case, mock_repo, existing_voice_item): + """测试 voice_id 更新为空字符串(合法场景:清除 voice_id)""" + mock_repo.get.return_value = existing_voice_item + + command = UpdateVoiceLibraryCommand( + id="existing-voice-001", + user_id="user-001", + voice_id="", + ) + + result = update_use_case.execute(command) + + assert result.voice_id == "" + + +# =========================================================================== +# 4. DeleteVoiceLibraryUseCase 测试 +# =========================================================================== + +class TestDeleteVoiceLibraryUseCase: + """配音库删除 UseCase 测试""" + + def test_delete_success(self, mock_repo): + """测试删除成功""" + mock_repo.delete.return_value = True + use_case = DeleteVoiceLibraryUseCase(repository=mock_repo) + + result = use_case.execute("voice-001", "user-001") + + assert result is True + mock_repo.delete.assert_called_once_with("voice-001", "user-001") + + def test_delete_not_found(self, mock_repo): + """测试删除不存在的条目""" + mock_repo.delete.return_value = False + use_case = DeleteVoiceLibraryUseCase(repository=mock_repo) + + result = use_case.execute("nonexistent", "user-001") + + assert result is False + + +# =========================================================================== +# 5. GetVoiceLibraryUseCase 测试 +# =========================================================================== + +class TestGetVoiceLibraryUseCase: + """配音库查询 UseCase 测试""" + + def test_get_existing(self, mock_repo): + """测试查询存在的条目""" + expected = VoiceLibraryItem( + id="v-001", + user_id="user-001", + name="测试", + voice_id="voice-xyz", + ) + mock_repo.get.return_value = expected + use_case = GetVoiceLibraryUseCase(repository=mock_repo) + + result = use_case.execute("v-001", "user-001") + + assert result is not None + assert result.id == "v-001" + assert result.voice_id == "voice-xyz" + mock_repo.get.assert_called_once_with("v-001", "user-001") + + def test_get_not_found(self, mock_repo): + """测试查询不存在的条目""" + mock_repo.get.return_value = None + use_case = GetVoiceLibraryUseCase(repository=mock_repo) + + result = use_case.execute("nonexistent", "user-001") + + assert result is None + + +# =========================================================================== +# 6. ListVoiceLibraryUseCase 测试 +# =========================================================================== + +class TestListVoiceLibraryUseCase: + """配音库列表 UseCase 测试""" + + def test_list_default(self, mock_repo): + """测试默认列表查询""" + items = [ + VoiceLibraryItem(id="v1", user_id="user-001", name="A"), + VoiceLibraryItem(id="v2", user_id="user-001", name="B"), + ] + mock_repo.list_by_user.return_value = items + use_case = ListVoiceLibraryUseCase(repository=mock_repo) + + result = use_case.execute("user-001") + + assert len(result) == 2 + mock_repo.list_by_user.assert_called_once_with("user-001", status=None, skip=0, limit=50) + + def test_list_with_status_filter(self, mock_repo): + """测试按状态筛选""" + mock_repo.list_by_user.return_value = [] + use_case = ListVoiceLibraryUseCase(repository=mock_repo) + + use_case.execute("user-001", status="completed", skip=10, limit=20) + + mock_repo.list_by_user.assert_called_once_with( + "user-001", status="completed", skip=10, limit=20 + ) + + def test_list_empty(self, mock_repo): + """测试空列表""" + mock_repo.list_by_user.return_value = [] + use_case = ListVoiceLibraryUseCase(repository=mock_repo) + + result = use_case.execute("user-001") + + assert result == [] + + +# =========================================================================== +# 7. voice_id 与 id 字段隔离专项回归测试 +# =========================================================================== + +class TestVoiceIdFieldIsolation: + """ + PR#74 P0 Bug 回归:voice_id 与 item.id 字段隔离 + + 原 bug:API 路由层误将 command.id(item 主键)用作 voice_id, + 导致 voice_id 字段值错误。本测试类从 UseCase 层验证 + 这两个字段在整个 CRUD 生命周期中互不干扰。 + """ + + def test_create_id_and_voice_id_are_independent(self, create_use_case, mock_repo): + """创建时 id 自动生成,voice_id 来自 command""" + command = CreateVoiceLibraryCommand( + user_id="user-001", + name="隔离测试", + voice_id="tts-voice-001", + voice_provider="openai", + ) + + result = create_use_case.execute(command, plan_name="free") + + assert result.id != result.voice_id, "id 和 voice_id 应为不同值" + assert result.voice_id == "tts-voice-001" + assert len(result.id) == 32 # UUID hex + + def test_update_voice_id_does_not_change_id(self, update_use_case, mock_repo): + """更新 voice_id 不影响 item 主键 id""" + existing = VoiceLibraryItem( + id="stable-id-001", + user_id="user-001", + name="测试", + voice_id="old-voice", + ) + mock_repo.get.return_value = existing + + command = UpdateVoiceLibraryCommand( + id="stable-id-001", + user_id="user-001", + voice_id="new-voice-999", + ) + + result = update_use_case.execute(command) + + assert result.id == "stable-id-001", "item 主键 id 不应改变" + assert result.voice_id == "new-voice-999", "voice_id 应被更新" + + def test_create_then_update_voice_id_preserves_id(self, create_use_case, update_use_case, mock_repo): + """创建后再更新 voice_id,id 始终不变""" + # 创建 + create_cmd = CreateVoiceLibraryCommand( + user_id="user-001", + name="生命周期测试", + voice_id="initial-voice", + ) + created = create_use_case.execute(create_cmd, plan_name="free") + original_id = created.id + + # 更新 + mock_repo.get.return_value = created + update_cmd = UpdateVoiceLibraryCommand( + id=original_id, + user_id="user-001", + voice_id="updated-voice", + ) + updated = update_use_case.execute(update_cmd) + + assert updated.id == original_id, "经过创建和更新,id 应保持一致" + assert updated.voice_id == "updated-voice" + assert updated.voice_id != "initial-voice" -- 2.54.0 From 2f7e50bff9a95b20bdaa45e24bff1594a8677b06 Mon Sep 17 00:00:00 2001 From: xiaoxia Date: Sun, 28 Jun 2026 15:43:57 +0800 Subject: [PATCH 4/4] =?UTF-8?q?test:=20=E6=B7=BB=E5=8A=A0=E6=A0=87?= =?UTF-8?q?=E9=A2=98=E5=BA=93=20UseCase=20=E5=9B=9E=E5=BD=92=E6=B5=8B?= =?UTF-8?q?=E8=AF=95=EF=BC=88=E9=85=8D=E9=A2=9D=E9=80=BB=E8=BE=91=20free?= =?UTF-8?q?=3D50/basic=3D500/premium=3D500=20=E8=A6=86=E7=9B=96=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- tests/unit/test_title_library_use_cases.py | 488 +++++++++++++++++++++ 1 file changed, 488 insertions(+) create mode 100644 tests/unit/test_title_library_use_cases.py diff --git a/tests/unit/test_title_library_use_cases.py b/tests/unit/test_title_library_use_cases.py new file mode 100644 index 000000000..77a181faa --- /dev/null +++ b/tests/unit/test_title_library_use_cases.py @@ -0,0 +1,488 @@ +""" +标题库(Title Library)Use Case 回归测试 + +测试目标: +1. CreateTitleLibraryUseCase - 创建标题库条目 +2. UpdateTitleLibraryUseCase - 更新标题库条目 +3. 配额逻辑覆盖 - titles: free=50, basic=500, premium=500 +4. 边界条件与异常场景 +""" + +from unittest.mock import Mock + +import pytest + +from packages.application.title_library.commands import ( + CreateTitleLibraryCommand, + UpdateTitleLibraryCommand, +) +from packages.application.title_library.use_cases import ( + CreateTitleLibraryUseCase, + UpdateTitleLibraryUseCase, + DeleteTitleLibraryUseCase, + GetTitleLibraryUseCase, + ListTitleLibraryUseCase, + QuotaExceededError, + NotFoundError, +) +from packages.domain.title_library import TitleLibraryItem + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + +@pytest.fixture +def mock_repo(): + """创建 Mock 仓储""" + repo = Mock() + repo.count_by_user = Mock(return_value=0) + repo.create = Mock(side_effect=lambda item: item) + repo.update = Mock(side_effect=lambda item: item) + repo.get = Mock(return_value=None) + repo.delete = Mock(return_value=True) + repo.list_by_user = Mock(return_value=[]) + return repo + + +@pytest.fixture +def create_use_case(mock_repo): + return CreateTitleLibraryUseCase(repository=mock_repo) + + +@pytest.fixture +def update_use_case(mock_repo): + return UpdateTitleLibraryUseCase(repository=mock_repo) + + +@pytest.fixture +def sample_create_command(): + """标准创建命令""" + return CreateTitleLibraryCommand( + user_id="user-001", + name="测试标题", + text="这是一个测试标题文本", + category="新闻", + description="用于测试的标题", + tags=["测试", "新闻"], + metadata_={"source": "unit_test"}, + ) + + +@pytest.fixture +def existing_title_item(): + """模拟已存在的标题条目""" + return TitleLibraryItem( + id="existing-title-001", + user_id="user-001", + name="旧标题", + text="旧文本", + category="旧分类", + description="旧描述", + tags=["旧"], + is_active=True, + metadata_={}, + ) + + +# =========================================================================== +# 1. CreateTitleLibraryUseCase 测试 +# =========================================================================== + +class TestCreateTitleLibraryUseCase: + """标题库创建 UseCase 测试""" + + def test_create_success_all_fields(self, create_use_case, mock_repo, sample_create_command): + """测试创建成功 - 所有字段完整传入""" + result = create_use_case.execute(sample_create_command, plan_name="free") + + assert result is not None + assert result.user_id == "user-001" + assert result.name == "测试标题" + assert result.text == "这是一个测试标题文本" + assert result.category == "新闻" + assert result.description == "用于测试的标题" + assert result.tags == ["测试", "新闻"] + assert result.metadata_ == {"source": "unit_test"} + + mock_repo.count_by_user.assert_called_once_with("user-001") + mock_repo.create.assert_called_once() + + def test_create_generates_uuid(self, create_use_case, mock_repo, sample_create_command): + """测试创建时自动生成 UUID 作为 id""" + result = create_use_case.execute(sample_create_command, plan_name="free") + + assert result.id is not None + assert len(result.id) == 32 # uuid4().hex 长度为 32 + assert result.id.isalnum() + + def test_create_default_values(self, create_use_case, mock_repo): + """测试默认值填充""" + command = CreateTitleLibraryCommand( + user_id="user-001", + name="最小化创建", + text="文本", + ) + + result = create_use_case.execute(command, plan_name="free") + + assert result.category == "default" + assert result.description == "" + assert result.tags == [] + assert result.metadata_ == {} + + +# =========================================================================== +# 2. 配额逻辑测试(titles: free=50, basic=500, premium=500) +# =========================================================================== + +class TestCreateTitleLibraryQuota: + """标题库创建配额检查测试""" + + def test_quota_free_plan_under_limit(self, create_use_case, mock_repo, sample_create_command): + """free 套餐(上限50),当前 25 个,允许创建""" + mock_repo.count_by_user.return_value = 25 + + result = create_use_case.execute(sample_create_command, plan_name="free") + + assert result is not None + mock_repo.create.assert_called_once() + + def test_quota_free_plan_at_limit(self, create_use_case, mock_repo, sample_create_command): + """free 套餐(上限50),当前 50 个,拒绝创建""" + mock_repo.count_by_user.return_value = 50 + + with pytest.raises(QuotaExceededError) as exc_info: + create_use_case.execute(sample_create_command, plan_name="free") + + assert exc_info.value.dimension == "max_titles" + assert exc_info.value.limit == 50 + assert exc_info.value.used == 50 + mock_repo.create.assert_not_called() + + def test_quota_free_plan_just_under_limit(self, create_use_case, mock_repo, sample_create_command): + """free 套餐(上限50),当前 49 个,允许创建(边界)""" + mock_repo.count_by_user.return_value = 49 + + result = create_use_case.execute(sample_create_command, plan_name="free") + + assert result is not None + mock_repo.create.assert_called_once() + + def test_quota_free_plan_over_limit(self, create_use_case, mock_repo, sample_create_command): + """free 套餐(上限50),当前 60 个,拒绝创建""" + mock_repo.count_by_user.return_value = 60 + + with pytest.raises(QuotaExceededError) as exc_info: + create_use_case.execute(sample_create_command, plan_name="free") + + assert exc_info.value.dimension == "max_titles" + assert exc_info.value.limit == 50 + assert exc_info.value.used == 60 + + def test_quota_basic_plan_under_limit(self, create_use_case, mock_repo, sample_create_command): + """basic 套餐(上限500),当前 200 个,允许创建""" + mock_repo.count_by_user.return_value = 200 + + result = create_use_case.execute(sample_create_command, plan_name="basic") + + assert result is not None + mock_repo.create.assert_called_once() + + def test_quota_basic_plan_at_limit(self, create_use_case, mock_repo, sample_create_command): + """basic 套餐(上限500),当前 500 个,拒绝创建""" + mock_repo.count_by_user.return_value = 500 + + with pytest.raises(QuotaExceededError) as exc_info: + create_use_case.execute(sample_create_command, plan_name="basic") + + assert exc_info.value.dimension == "max_titles" + assert exc_info.value.limit == 500 + assert exc_info.value.used == 500 + + def test_quota_basic_plan_just_under_limit(self, create_use_case, mock_repo, sample_create_command): + """basic 套餐(上限500),当前 499 个,允许创建(边界)""" + mock_repo.count_by_user.return_value = 499 + + result = create_use_case.execute(sample_create_command, plan_name="basic") + + assert result is not None + mock_repo.create.assert_called_once() + + def test_quota_premium_plan_under_limit(self, create_use_case, mock_repo, sample_create_command): + """premium 套餐(上限500),当前 250 个,允许创建""" + mock_repo.count_by_user.return_value = 250 + + result = create_use_case.execute(sample_create_command, plan_name="premium") + + assert result is not None + mock_repo.create.assert_called_once() + + def test_quota_premium_plan_at_limit(self, create_use_case, mock_repo, sample_create_command): + """premium 套餐(上限500),当前 500 个,拒绝创建""" + mock_repo.count_by_user.return_value = 500 + + with pytest.raises(QuotaExceededError) as exc_info: + create_use_case.execute(sample_create_command, plan_name="premium") + + assert exc_info.value.dimension == "max_titles" + assert exc_info.value.limit == 500 + assert exc_info.value.used == 500 + + def test_quota_premium_plan_just_under_limit(self, create_use_case, mock_repo, sample_create_command): + """premium 套餐(上限500),当前 499 个,允许创建(边界)""" + mock_repo.count_by_user.return_value = 499 + + result = create_use_case.execute(sample_create_command, plan_name="premium") + + assert result is not None + mock_repo.create.assert_called_once() + + def test_quota_zero_usage_all_plans(self, create_use_case, mock_repo, sample_create_command): + """新用户零使用量,所有套餐均可创建""" + mock_repo.count_by_user.return_value = 0 + + for plan in ["free", "basic", "premium"]: + mock_repo.create.reset_mock() + mock_repo.count_by_user.reset_mock() + mock_repo.count_by_user.return_value = 0 + + result = create_use_case.execute(sample_create_command, plan_name=plan) + assert result is not None, f"{plan} 套餐零使用量应允许创建" + + def test_quota_unknown_plan_defaults_to_zero(self, create_use_case, mock_repo, sample_create_command): + """未知套餐名默认配额为 0,无法创建""" + mock_repo.count_by_user.return_value = 0 + + with pytest.raises(QuotaExceededError): + create_use_case.execute(sample_create_command, plan_name="unknown_plan") + + def test_quota_exceeded_error_attributes(self, create_use_case, mock_repo, sample_create_command): + """QuotaExceededError 异常属性完整性""" + mock_repo.count_by_user.return_value = 50 + + with pytest.raises(QuotaExceededError) as exc_info: + create_use_case.execute(sample_create_command, plan_name="free") + + err = exc_info.value + assert hasattr(err, "dimension") + assert hasattr(err, "limit") + assert hasattr(err, "used") + assert "max_titles" in str(err) + assert "50" in str(err) + + +# =========================================================================== +# 3. UpdateTitleLibraryUseCase 测试 +# =========================================================================== + +class TestUpdateTitleLibraryUseCase: + """标题库更新 UseCase 测试""" + + def test_update_success_all_fields(self, update_use_case, mock_repo, existing_title_item): + """测试全字段更新成功""" + mock_repo.get.return_value = existing_title_item + + command = UpdateTitleLibraryCommand( + title_id="existing-title-001", + user_id="user-001", + name="更新后标题", + text="更新后文本", + category="新分类", + description="新描述", + tags=["新标签"], + is_active=False, + metadata_={"updated": True}, + ) + + result = update_use_case.execute(command) + + assert result.name == "更新后标题" + assert result.text == "更新后文本" + assert result.category == "新分类" + assert result.description == "新描述" + assert result.tags == ["新标签"] + assert result.is_active is False + assert result.metadata_ == {"updated": True} + + mock_repo.update.assert_called_once() + + def test_update_partial_only_name(self, update_use_case, mock_repo, existing_title_item): + """测试仅更新 name""" + mock_repo.get.return_value = existing_title_item + + command = UpdateTitleLibraryCommand( + title_id="existing-title-001", + user_id="user-001", + name="仅改名", + ) + + result = update_use_case.execute(command) + + assert result.name == "仅改名" + # 其他字段保持不变 + assert result.text == "旧文本" + assert result.category == "旧分类" + assert result.description == "旧描述" + + def test_update_partial_only_is_active(self, update_use_case, mock_repo, existing_title_item): + """测试仅更新 is_active(软删除/恢复)""" + mock_repo.get.return_value = existing_title_item + + command = UpdateTitleLibraryCommand( + title_id="existing-title-001", + user_id="user-001", + is_active=False, + ) + + result = update_use_case.execute(command) + + assert result.is_active is False + assert result.name == "旧标题" # 其他字段不变 + + def test_update_not_found(self, update_use_case, mock_repo): + """测试更新不存在的条目""" + mock_repo.get.return_value = None + + command = UpdateTitleLibraryCommand( + title_id="nonexistent-id", + user_id="user-001", + name="不存在", + ) + + with pytest.raises(NotFoundError, match="nonexistent-id"): + update_use_case.execute(command) + + mock_repo.update.assert_not_called() + + def test_update_wrong_user(self, update_use_case, mock_repo): + """测试用户隔离""" + mock_repo.get.return_value = None + + command = UpdateTitleLibraryCommand( + title_id="existing-title-001", + user_id="other-user-999", + name="恶意修改", + ) + + with pytest.raises(NotFoundError): + update_use_case.execute(command) + + def test_update_none_fields_not_changed(self, update_use_case, mock_repo, existing_title_item): + """测试 None 字段不覆盖原有值""" + mock_repo.get.return_value = existing_title_item + + command = UpdateTitleLibraryCommand( + title_id="existing-title-001", + user_id="user-001", + ) + + result = update_use_case.execute(command) + + assert result.name == "旧标题" + assert result.text == "旧文本" + assert result.category == "旧分类" + assert result.is_active is True + + +# =========================================================================== +# 4. DeleteTitleLibraryUseCase 测试 +# =========================================================================== + +class TestDeleteTitleLibraryUseCase: + """标题库删除 UseCase 测试""" + + def test_delete_success(self, mock_repo): + """测试删除成功""" + mock_repo.delete.return_value = True + use_case = DeleteTitleLibraryUseCase(repository=mock_repo) + + result = use_case.execute("title-001", "user-001") + + assert result is True + mock_repo.delete.assert_called_once_with("title-001", "user-001") + + def test_delete_not_found(self, mock_repo): + """测试删除不存在的条目""" + mock_repo.delete.return_value = False + use_case = DeleteTitleLibraryUseCase(repository=mock_repo) + + result = use_case.execute("nonexistent", "user-001") + + assert result is False + + +# =========================================================================== +# 5. GetTitleLibraryUseCase 测试 +# =========================================================================== + +class TestGetTitleLibraryUseCase: + """标题库查询 UseCase 测试""" + + def test_get_existing(self, mock_repo): + """测试查询存在的条目""" + expected = TitleLibraryItem( + id="t-001", + user_id="user-001", + name="测试", + text="文本", + ) + mock_repo.get.return_value = expected + use_case = GetTitleLibraryUseCase(repository=mock_repo) + + result = use_case.execute("t-001", "user-001") + + assert result is not None + assert result.id == "t-001" + mock_repo.get.assert_called_once_with("t-001", "user-001") + + def test_get_not_found(self, mock_repo): + """测试查询不存在的条目""" + mock_repo.get.return_value = None + use_case = GetTitleLibraryUseCase(repository=mock_repo) + + result = use_case.execute("nonexistent", "user-001") + + assert result is None + + +# =========================================================================== +# 6. ListTitleLibraryUseCase 测试 +# =========================================================================== + +class TestListTitleLibraryUseCase: + """标题库列表 UseCase 测试""" + + def test_list_default(self, mock_repo): + """测试默认列表查询""" + items = [ + TitleLibraryItem(id="t1", user_id="user-001", name="A", text="a"), + TitleLibraryItem(id="t2", user_id="user-001", name="B", text="b"), + ] + mock_repo.list_by_user.return_value = items + use_case = ListTitleLibraryUseCase(repository=mock_repo) + + result = use_case.execute("user-001") + + assert len(result) == 2 + mock_repo.list_by_user.assert_called_once_with("user-001", category=None, skip=0, limit=50) + + def test_list_with_category_filter(self, mock_repo): + """测试按分类筛选""" + mock_repo.list_by_user.return_value = [] + use_case = ListTitleLibraryUseCase(repository=mock_repo) + + use_case.execute("user-001", category="新闻", skip=5, limit=10) + + mock_repo.list_by_user.assert_called_once_with( + "user-001", category="新闻", skip=5, limit=10 + ) + + def test_list_empty(self, mock_repo): + """测试空列表""" + mock_repo.list_by_user.return_value = [] + use_case = ListTitleLibraryUseCase(repository=mock_repo) + + result = use_case.execute("user-001") + + assert result == [] -- 2.54.0