From ea9e536740afc2027d12e835388eb4ba1e40d777 Mon Sep 17 00:00:00 2001 From: Audit Bot Date: Mon, 29 Jun 2026 17:05:53 +0800 Subject: [PATCH 1/2] =?UTF-8?q?feat:=20=E6=89=A9=E5=B1=95=E7=94=9F?= =?UTF-8?q?=E6=88=90=E4=BB=BB=E5=8A=A1API=E6=94=AF=E6=8C=81=E6=A8=A1?= =?UTF-8?q?=E6=9D=BF=E6=A8=A1=E5=BC=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 方案A: 后端扩展3处API改动,支持前端模板中心化模型。 1. POST /api/v1/generation/tasks — project_id改为可选,新增template_id/asset_ids/title_ids/voice_ids 2. GET /api/v1/generation/tasks + GET /api/v1/tasks — 新增用户级列表(跨project) 3. POST /api/v1/tasks/{task_id}/retry + POST /api/v1/generation/tasks/{task_id}/retry — 简化重试 改动涉及: - Domain: GenerationTask新增6个字段,放宽create()校验 - Application: CreateGenerationTaskCommand新增字段 - Ports: GenerationTaskRepository新增list_by_user() - Adapter: SQLAlchemy模型+仓储实现新字段和list_by_user - Schema: 双模式校验(project模式/模板模式) - Routes: generation_tasks + task_center双路由注册 - Migration: 015_add_generation_task_extensions 向后兼容:所有旧端点和参数不变。 --- .../015_add_generation_task_extensions.py | 36 ++++ apps/api/app/api/routes/generation_tasks.py | 125 ++++++++++-- apps/api/app/api/routes/task_center.py | 179 +++++++++++++----- apps/api/app/schemas/generation_task.py | 38 +++- apps/api/app/schemas/task_center.py | 23 +++ .../generation_task_repository.py | 71 +++++-- packages/adapters/sqlalchemy_impl/models.py | 11 +- packages/application/generation_tasks.py | 14 +- packages/domain/generation_task.py | 23 ++- packages/ports/generation_task_repository.py | 2 + 10 files changed, 431 insertions(+), 91 deletions(-) create mode 100644 alembic/versions/015_add_generation_task_extensions.py diff --git a/alembic/versions/015_add_generation_task_extensions.py b/alembic/versions/015_add_generation_task_extensions.py new file mode 100644 index 000000000..4af932459 --- /dev/null +++ b/alembic/versions/015_add_generation_task_extensions.py @@ -0,0 +1,36 @@ +"""add generation task extensions + +Revision ID: 015 +Revises: 014 +Create Date: 2026-06-29 +""" + +import sqlalchemy as sa +from sqlalchemy.dialects import mysql + +from alembic import op + +# revision identifiers +revision = "015" +down_revision = "014" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.add_column("generation_tasks", sa.Column("edit_plan_id", sa.String(32), nullable=False, server_default="")) + op.add_column("generation_tasks", sa.Column("template_id", sa.String(36), nullable=False, server_default="")) + op.add_column("generation_tasks", sa.Column("asset_ids", mysql.JSON(), nullable=False, server_default="[]")) + op.add_column("generation_tasks", sa.Column("title_ids", mysql.JSON(), nullable=False, server_default="[]")) + op.add_column("generation_tasks", sa.Column("voice_ids", mysql.JSON(), nullable=False, server_default="[]")) + + op.create_index(op.f("ix_generation_tasks_template_id"), "generation_tasks", ["template_id"]) + + +def downgrade() -> None: + op.drop_index(op.f("ix_generation_tasks_template_id"), table_name="generation_tasks") + op.drop_column("generation_tasks", "voice_ids") + op.drop_column("generation_tasks", "title_ids") + op.drop_column("generation_tasks", "asset_ids") + op.drop_column("generation_tasks", "template_id") + op.drop_column("generation_tasks", "edit_plan_id") diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index c363d5595..462b92beb 100644 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -16,6 +16,7 @@ from app.schemas.generated_video import ( from app.schemas.generation_task import ( CreateGenerationTaskRequest, GenerationTaskResponse, + ListGenerationTasksResponse, ) from fastapi import APIRouter, Depends, HTTPException @@ -45,6 +46,10 @@ 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, + template_id=task.template_id, + asset_ids=task.asset_ids, + title_ids=task.title_ids, + voice_ids=task.voice_ids, status=task.status, progress=task.progress, result_count=task.result_count, @@ -79,8 +84,45 @@ def _ensure_library_has_ready_video_assets(assets) -> None: ) +async def _resolve_project_and_library( + request: CreateGenerationTaskRequest, + project_repository: Any, + asset_library_repository: Any, + asset_repository: Any, + authenticated_user: AuthenticatedUser, +) -> tuple[str, str]: + """解析 project_id 和 asset_library_id。 + + 支持两种模式: + - 显式传入(向后兼容) + - 从 asset_ids 反查 asset_library(模板模式) + 返回 (project_id, asset_library_id)。 + """ + project_id = request.project_id.strip() + asset_library_id = request.asset_library_id.strip() + + # 模板模式:project_id 未提供时,从 asset_ids 反查所属 project + if not project_id and request.asset_ids: + first_asset_id = request.asset_ids[0] + asset = await asset_repository.find_by_id(first_asset_id) + if asset is not None: + project_id = asset.project_id + if not asset_library_id: + asset_library_id = asset.library_id + + # 向后兼容校验:project_id 已提供时验证权限 + if project_id: + project = project_repository.find_by_id(project_id) + if project is None: + raise HTTPException(status_code=404, detail=f"Project {project_id} not found") + if not project.can_access(authenticated_user.user.id): + raise HTTPException(status_code=403, detail="Access denied to project") + + return project_id, asset_library_id + + @router.post("/tasks", response_model=GenerationTaskResponse) -def create_generation_task( +async def create_generation_task( request: CreateGenerationTaskRequest, authenticated_user: AuthenticatedUser = Depends(get_current_user), generation_task_repository: Any = Depends(get_generation_task_repository), @@ -88,26 +130,30 @@ def create_generation_task( asset_library_repository: Any = Depends(get_asset_library_repository), asset_repository: Any = Depends(get_asset_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.user.id): - raise HTTPException(status_code=403, detail="Access denied to project") - - library = asset_library_repository.get(request.asset_library_id) - if library is None or library.project_id != request.project_id: - raise HTTPException(status_code=404, detail=f"AssetLibrary {request.asset_library_id} not found") - - assets = asset_repository.list_by_library(request.asset_library_id) - _ensure_library_has_ready_video_assets(assets) + project_id, asset_library_id = await _resolve_project_and_library( + request, project_repository, asset_library_repository, asset_repository, authenticated_user + ) + + # asset_library 存在性校验(仅在提供了 asset_library_id 时) + if asset_library_id: + library = asset_library_repository.get(asset_library_id) + if library is None or (project_id and library.project_id != project_id): + raise HTTPException(status_code=404, detail=f"AssetLibrary {asset_library_id} not found") + + assets = await asset_repository.find_by_library(asset_library_id) + _ensure_library_has_ready_video_assets(assets) use_case = CreateGenerationTaskUseCase(generation_task_repository) task = use_case.execute( CreateGenerationTaskCommand( - project_id=request.project_id, - asset_library_id=request.asset_library_id, + project_id=project_id, + asset_library_id=asset_library_id, strategy_id=request.strategy_id, voice_library_id=request.voice_library_id, + template_id=request.template_id, + asset_ids=request.asset_ids, + title_ids=request.title_ids, + voice_ids=request.voice_ids, created_by_user_id=authenticated_user.user.id, ) ) @@ -115,6 +161,17 @@ def create_generation_task( return _to_generation_task_response(task) +@router.get("/tasks", response_model=ListGenerationTasksResponse) +def list_generation_tasks( + authenticated_user: AuthenticatedUser = Depends(get_current_user), + generation_task_repository: Any = Depends(get_generation_task_repository), +) -> ListGenerationTasksResponse: + """用户级生成任务列表(跨 project)。""" + tasks = generation_task_repository.list_by_user(authenticated_user.user.id) + items = [_to_generation_task_response(task) for task in tasks] + return ListGenerationTasksResponse(items=items) + + @router.get("/tasks/{task_id}", response_model=GenerationTaskResponse) def get_generation_task( task_id: str, @@ -126,7 +183,8 @@ 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.user.id, project_repository) + if task.project_id: + _check_project_access(task.project_id, authenticated_user.user.id, project_repository) return _to_generation_task_response(task) @@ -141,7 +199,40 @@ 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.user.id, project_repository) + if task.project_id: + _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]) + + +@router.post("/tasks/{task_id}/retry", response_model=GenerationTaskResponse) +def retry_generation_task( + task_id: str, + authenticated_user: AuthenticatedUser = Depends(get_current_user), + generation_task_repository: Any = Depends(get_generation_task_repository), +) -> GenerationTaskResponse: + """简化重试:通过 task_id 直接重试失败任务。""" + task = generation_task_repository.get(task_id) + if task is None: + raise HTTPException(status_code=404, detail="Generation task not found") + if task.status != "failed": + raise HTTPException(status_code=409, detail="Only failed tasks can be retried") + + use_case = CreateGenerationTaskUseCase(generation_task_repository) + retried = use_case.execute( + CreateGenerationTaskCommand( + project_id=task.project_id, + 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, + template_id=task.template_id, + asset_ids=task.asset_ids, + title_ids=task.title_ids, + voice_ids=task.voice_ids, + created_by_user_id=authenticated_user.user.id, + ) + ) + celery_app.send_task("worker.generate_video", args=[retried.id]) + return _to_generation_task_response(retried) diff --git a/apps/api/app/api/routes/task_center.py b/apps/api/app/api/routes/task_center.py index ab72eca14..1d2e9524d 100644 --- a/apps/api/app/api/routes/task_center.py +++ b/apps/api/app/api/routes/task_center.py @@ -7,7 +7,12 @@ from app.dependencies import ( get_ingest_job_repository, get_project_repository, ) -from app.schemas.task_center import ListProjectTasksResponse, ProjectTaskResponse +from app.schemas.task_center import ( + ListProjectTasksResponse, + ListTasksResponse, + ProjectTaskResponse, + UserTaskResponse, +) from fastapi import APIRouter, Depends, HTTPException from packages.application import ( @@ -34,28 +39,135 @@ def _humanize_task_error(error_message: str) -> str: return f"任务失败:{raw}" +def _status_value(status) -> str: + """安全获取状态值(兼容 StrEnum 和 plain string)。""" + return status.value if hasattr(status, "value") else str(status) + + def _generation_step(task) -> str: - if task.status.value == "pending": + s = _status_value(task.status) + if s == "pending": return "等待 Worker 执行" - if task.status.value == "running": + if s == "running": return "正在生成成片" - if task.status.value == "completed": + if s == "completed": return "生成完成" - if task.status.value == "failed": + if s == "failed": return "生成失败" - return task.status.value + return s def _ingest_step(job) -> str: - if job.status.value == "pending": + s = _status_value(job.status) + if s == "pending": return "等待导入" - if job.status.value == "processing": + if s == "processing": return "正在分析素材" - if job.status.value == "completed": + if s == "completed": return "导入完成" - if job.status.value == "failed": + if s == "failed": return "导入失败" - return job.status.value + return s + + +def _generation_task_to_project_response(task) -> ProjectTaskResponse: + return ProjectTaskResponse( + id=f"generation:{task.id}", + task_type="generation", + project_id=task.project_id, + status=_status_value(task.status), + progress=task.progress, + current_step=_generation_step(task), + error_message=task.error_message, + user_message=_humanize_task_error(task.error_message), + retryable=_status_value(task.status) == "failed", + source_id=task.id, + template_id=task.template_id, + created_at=task.created_at, + updated_at=task.completed_at or task.started_at or task.created_at, + ) + + +# ── 用户级端点(放在项目级端点之前,避免路由冲突) ── + + +@router.get("/tasks", response_model=ListTasksResponse) +def list_user_tasks( + authenticated_user: AuthenticatedUser = Depends(get_current_user), + ingest_job_repository: Any = Depends(get_ingest_job_repository), + generation_task_repository: Any = Depends(get_generation_task_repository), +) -> ListTasksResponse: + """用户级任务列表(跨 project),合并 ingest + generation 任务。""" + user_id = authenticated_user.user.id + items: list[UserTaskResponse] = [] + + for task in generation_task_repository.list_by_user(user_id): + items.append( + UserTaskResponse( + id=f"generation:{task.id}", + task_type="generation", + project_id=task.project_id, + template_id=task.template_id, + status=_status_value(task.status), + progress=task.progress, + current_step=_generation_step(task), + error_message=task.error_message, + user_message=_humanize_task_error(task.error_message), + retryable=_status_value(task.status) == "failed", + source_id=task.id, + created_at=task.created_at, + updated_at=task.completed_at or task.started_at or task.created_at, + ) + ) + + items.sort(key=lambda item: item.updated_at or item.created_at or "", reverse=True) + return ListTasksResponse(items=items) + + +@router.post("/tasks/{task_id}/retry", response_model=UserTaskResponse) +def retry_task_by_id( + task_id: str, + authenticated_user: AuthenticatedUser = Depends(get_current_user), + generation_task_repository: Any = Depends(get_generation_task_repository), +) -> UserTaskResponse: + """简化重试:通过 task_id 直接重试失败的生成任务。""" + task = generation_task_repository.get(task_id) + if task is None: + raise HTTPException(status_code=404, detail="Generation task not found") + if _status_value(task.status) != "failed": + raise HTTPException(status_code=409, detail="Only failed tasks can be retried") + + use_case = CreateGenerationTaskUseCase(generation_task_repository) + retried = use_case.execute( + CreateGenerationTaskCommand( + project_id=task.project_id, + 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, + template_id=task.template_id, + asset_ids=task.asset_ids, + title_ids=task.title_ids, + voice_ids=task.voice_ids, + created_by_user_id=authenticated_user.user.id, + ) + ) + celery_app.send_task("worker.generate_video", args=[retried.id]) + return UserTaskResponse( + id=f"generation:{retried.id}", + task_type="generation", + project_id=retried.project_id, + template_id=retried.template_id, + status=_status_value(retried.status), + progress=retried.progress, + current_step=_generation_step(retried), + source_id=retried.id, + created_at=retried.created_at, + updated_at=retried.created_at, + ) + + +# ── 项目级端点 ── @router.get("/projects/{project_id}/tasks", response_model=ListProjectTasksResponse) @@ -77,34 +189,19 @@ def list_project_tasks( id=f"ingest:{job.id}", task_type="ingest", project_id=job.project_id, - status=job.status.value, - progress=100.0 if job.status.value == "completed" else 0.0, + status=_status_value(job.status), + progress=100.0 if _status_value(job.status) == "completed" else 0.0, current_step=_ingest_step(job), error_message=job.error_message, user_message=_humanize_task_error(job.error_message), - retryable=job.status.value == "failed", + retryable=_status_value(job.status) == "failed", source_id=job.id, created_at=job.created_at, updated_at=job.updated_at, ) ) for task in generation_task_repository.list_by_project(project_id): - items.append( - ProjectTaskResponse( - id=f"generation:{task.id}", - task_type="generation", - project_id=task.project_id, - status=task.status.value, - progress=task.progress, - current_step=_generation_step(task), - error_message=task.error_message, - user_message=_humanize_task_error(task.error_message), - retryable=task.status.value == "failed", - source_id=task.id, - created_at=task.created_at, - updated_at=task.completed_at or task.started_at or task.created_at, - ) - ) + items.append(_generation_task_to_project_response(task)) items.sort(key=lambda item: item.updated_at or item.created_at or "", reverse=True) return ListProjectTasksResponse(items=items) @@ -121,7 +218,7 @@ def retry_project_task( task = generation_task_repository.get(source_id) if task is None: raise HTTPException(status_code=404, detail="Generation task not found") - if task.status.value != "failed": + if _status_value(task.status) != "failed": raise HTTPException(status_code=409, detail="Only failed tasks can be retried") use_case = CreateGenerationTaskUseCase(generation_task_repository) retried = use_case.execute( @@ -131,26 +228,20 @@ def retry_project_task( strategy_id=task.strategy_id, voice_library_id=task.voice_library_id, edit_plan_id=task.edit_plan_id, + template_id=task.template_id, + asset_ids=task.asset_ids, + title_ids=task.title_ids, + voice_ids=task.voice_ids, created_by_user_id=authenticated_user.user.id, ) ) celery_app.send_task("worker.generate_video", args=[retried.id]) - return ProjectTaskResponse( - id=f"generation:{retried.id}", - task_type="generation", - project_id=retried.project_id, - status=retried.status.value, - progress=retried.progress, - current_step=_generation_step(retried), - source_id=retried.id, - created_at=retried.created_at, - updated_at=retried.created_at, - ) + return _generation_task_to_project_response(retried) if task_type == "ingest": job = ingest_job_repository.get(source_id) if job is None: raise HTTPException(status_code=404, detail="Ingest job not found") - if job.status.value != "failed": + if _status_value(job.status) != "failed": raise HTTPException(status_code=409, detail="Only failed tasks can be retried") use_case = SubmitIngestJobUseCase(ingest_job_repository) retried = use_case.execute( @@ -165,7 +256,7 @@ def retry_project_task( id=f"ingest:{retried.id}", task_type="ingest", project_id=retried.project_id, - status=retried.status.value, + status=_status_value(retried.status), progress=0, current_step=_ingest_step(retried), source_id=retried.id, diff --git a/apps/api/app/schemas/generation_task.py b/apps/api/app/schemas/generation_task.py index 0a6fec37b..a9e0beb91 100644 --- a/apps/api/app/schemas/generation_task.py +++ b/apps/api/app/schemas/generation_task.py @@ -1,12 +1,35 @@ -from pydantic import BaseModel, Field +from pydantic import BaseModel, Field, model_validator class CreateGenerationTaskRequest(BaseModel): - project_id: str = Field(..., min_length=1) - asset_library_id: str = Field(..., min_length=1) + """创建生成任务请求。 + + 支持两种模式(至少提供一种): + - 项目模式:project_id + asset_library_id(向后兼容) + - 模板模式:template_id + asset_ids / title_ids / voice_ids + """ + project_id: str = "" + asset_library_id: str = "" strategy_id: str = "" voice_library_id: str = "" created_by_user_id: str = "" + # ── 模板模式新增字段 ── + template_id: str = "" + asset_ids: list[str] = Field(default_factory=list) + title_ids: list[str] = Field(default_factory=list) + voice_ids: list[str] = Field(default_factory=list) + + @model_validator(mode="after") + def _check_at_least_one_mode(self) -> "CreateGenerationTaskRequest": + has_project = bool(self.project_id.strip()) + has_template = bool(self.template_id.strip()) + if not has_project and not has_template: + raise ValueError("project_id 或 template_id 至少需要提供一个") + has_library = bool(self.asset_library_id.strip()) + has_assets = bool(self.asset_ids or self.title_ids or self.voice_ids) + if not has_library and not has_assets: + raise ValueError("asset_library_id 或 asset_ids/title_ids/voice_ids 至少需要提供一个") + return self class GenerationTaskResponse(BaseModel): @@ -15,7 +38,16 @@ class GenerationTaskResponse(BaseModel): asset_library_id: str strategy_id: str voice_library_id: str + template_id: str = "" + asset_ids: list[str] = Field(default_factory=list) + title_ids: list[str] = Field(default_factory=list) + voice_ids: list[str] = Field(default_factory=list) status: str progress: float result_count: int error_message: str + + +class ListGenerationTasksResponse(BaseModel): + """用户级生成任务列表响应(跨 project)。""" + items: list[GenerationTaskResponse] diff --git a/apps/api/app/schemas/task_center.py b/apps/api/app/schemas/task_center.py index a7c5f2f9a..0aaaf7170 100644 --- a/apps/api/app/schemas/task_center.py +++ b/apps/api/app/schemas/task_center.py @@ -14,9 +14,32 @@ class ProjectTaskResponse(BaseModel): user_message: str = "" retryable: bool = False source_id: str = "" + template_id: str = "" created_at: datetime | None = None updated_at: datetime | None = None class ListProjectTasksResponse(BaseModel): items: list[ProjectTaskResponse] = Field(default_factory=list) + + +class UserTaskResponse(BaseModel): + """用户级任务响应(跨 project,用于模板模式)。""" + id: str + task_type: str + project_id: str = "" + template_id: str = "" + status: str + progress: float + current_step: str + error_message: str = "" + user_message: str = "" + retryable: bool = False + source_id: str = "" + created_at: datetime | None = None + updated_at: datetime | None = None + + +class ListTasksResponse(BaseModel): + """用户级任务列表响应(GET /api/v1/tasks)。""" + items: list[UserTaskResponse] = Field(default_factory=list) diff --git a/packages/adapters/sqlalchemy_impl/generation_task_repository.py b/packages/adapters/sqlalchemy_impl/generation_task_repository.py index 430200ebf..b2c635f77 100644 --- a/packages/adapters/sqlalchemy_impl/generation_task_repository.py +++ b/packages/adapters/sqlalchemy_impl/generation_task_repository.py @@ -4,6 +4,30 @@ from packages.adapters.sqlalchemy_impl.models import GenerationTaskModel from packages.domain import GenerationTask +def _to_domain(model: GenerationTaskModel) -> GenerationTask: + """Convert ORM model to domain entity.""" + return GenerationTask( + id=model.id, + project_id=model.project_id, + strategy_id=model.strategy_id, + asset_library_id=model.asset_library_id, + voice_library_id=model.voice_library_id, + edit_plan_id=model.edit_plan_id, + template_id=model.template_id, + asset_ids=list(model.asset_ids or []), + title_ids=list(model.title_ids or []), + voice_ids=list(model.voice_ids or []), + status=model.status, + progress=model.progress, + result_count=int(model.result_count or 0), + error_message=model.error_message, + started_at=model.started_at, + completed_at=model.completed_at, + created_by_user_id=model.created_by_user_id, + created_at=model.created_at, + ) + + class SQLAlchemyGenerationTaskRepository: def __init__(self, session: Session): self.session = session @@ -15,6 +39,11 @@ 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, + template_id=task.template_id, + asset_ids=task.asset_ids, + title_ids=task.title_ids, + voice_ids=task.voice_ids, status=task.status, progress=task.progress, result_count=task.result_count, @@ -32,31 +61,39 @@ class SQLAlchemyGenerationTaskRepository: model = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.id == task_id).first() if model is None: return None - return GenerationTask( - id=model.id, - project_id=model.project_id, - strategy_id=model.strategy_id, - asset_library_id=model.asset_library_id, - voice_library_id=model.voice_library_id, - status=model.status, - progress=model.progress, - result_count=int(model.result_count or 0), - error_message=model.error_message, - started_at=model.started_at, - completed_at=model.completed_at, - created_by_user_id=model.created_by_user_id, - created_at=model.created_at, - ) + return _to_domain(model) def list_by_project(self, project_id: str) -> list[GenerationTask]: - models = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.project_id == project_id).all() - return [self.get(model.id) for model in models if self.get(model.id) is not None] + models = ( + self.session.query(GenerationTaskModel) + .filter(GenerationTaskModel.project_id == project_id) + .order_by(GenerationTaskModel.created_at.desc()) + .all() + ) + return [_to_domain(m) for m in models] + + def list_by_user(self, user_id: str) -> list[GenerationTask]: + models = ( + self.session.query(GenerationTaskModel) + .filter(GenerationTaskModel.created_by_user_id == user_id) + .order_by(GenerationTaskModel.created_at.desc()) + .all() + ) + return [_to_domain(m) for m in models] def update(self, task: GenerationTask) -> GenerationTask: model = self.session.query(GenerationTaskModel).filter(GenerationTaskModel.id == task.id).first() if model is None: raise ValueError(f"GenerationTask {task.id} not found") + model.project_id = task.project_id + model.asset_library_id = task.asset_library_id + model.strategy_id = task.strategy_id model.voice_library_id = task.voice_library_id + model.edit_plan_id = task.edit_plan_id + model.template_id = task.template_id + model.asset_ids = task.asset_ids + model.title_ids = task.title_ids + model.voice_ids = task.voice_ids 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 7bfc5fab2..280e23bc7 100755 --- a/packages/adapters/sqlalchemy_impl/models.py +++ b/packages/adapters/sqlalchemy_impl/models.py @@ -139,10 +139,15 @@ class GenerationTaskModel(Base): __tablename__ = "generation_tasks" id = Column(String(32), primary_key=True) - project_id = Column(String(32), nullable=False, index=True) + project_id = Column(String(32), nullable=False, default="", index=True) strategy_id = Column(String(32), nullable=False, default="") - asset_library_id = Column(String(32), nullable=False, index=True) + asset_library_id = Column(String(32), nullable=False, default="", index=True) voice_library_id = Column(String(32), nullable=False, default="") + edit_plan_id = Column(String(32), nullable=False, default="") + template_id = Column(String(36), nullable=False, default="", index=True) + asset_ids = Column(JSON, nullable=False, default=list) + title_ids = Column(JSON, nullable=False, default=list) + voice_ids = Column(JSON, nullable=False, default=list) 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) @@ -150,7 +155,7 @@ class GenerationTaskModel(Base): error_message = Column(Text, nullable=False, default="") started_at = Column(DateTime, nullable=True) completed_at = Column(DateTime, nullable=True) - created_by_user_id = Column(String(32), nullable=False, default="") + created_by_user_id = Column(String(32), nullable=False, default="", index=True) extra_meta = Column('metadata', JSON, nullable=False, default=dict) created_at = Column(DateTime, nullable=False, default=lambda: datetime.now(timezone.utc)) diff --git a/packages/application/generation_tasks.py b/packages/application/generation_tasks.py index 4dab04c7a..09520da0a 100644 --- a/packages/application/generation_tasks.py +++ b/packages/application/generation_tasks.py @@ -1,6 +1,6 @@ from __future__ import annotations -from dataclasses import dataclass +from dataclasses import dataclass, field from uuid import uuid4 from packages.domain import GenerationTask @@ -9,11 +9,15 @@ from packages.ports.generation_task_repository import GenerationTaskRepository @dataclass(slots=True) class CreateGenerationTaskCommand: - project_id: str - asset_library_id: str + project_id: str = "" + asset_library_id: str = "" strategy_id: str = "" voice_library_id: str = "" edit_plan_id: str = "" + template_id: str = "" + asset_ids: list[str] = field(default_factory=list) + title_ids: list[str] = field(default_factory=list) + voice_ids: list[str] = field(default_factory=list) created_by_user_id: str = "" @@ -29,6 +33,10 @@ class CreateGenerationTaskUseCase: strategy_id=command.strategy_id, voice_library_id=command.voice_library_id, edit_plan_id=command.edit_plan_id, + template_id=command.template_id, + asset_ids=command.asset_ids, + title_ids=command.title_ids, + voice_ids=command.voice_ids, status="pending", progress=0.0, result_count=0, diff --git a/packages/domain/generation_task.py b/packages/domain/generation_task.py index f10e8a2d4..a08c75754 100644 --- a/packages/domain/generation_task.py +++ b/packages/domain/generation_task.py @@ -21,6 +21,11 @@ class GenerationTask: asset_library_id: str strategy_id: str = "" voice_library_id: str = "" + edit_plan_id: str = "" + template_id: str = "" + asset_ids: list[str] = field(default_factory=list) + title_ids: list[str] = field(default_factory=list) + voice_ids: list[str] = field(default_factory=list) status: GenerationTaskStatus = GenerationTaskStatus.PENDING progress: float = 0.0 result_count: int = 0 @@ -38,17 +43,27 @@ class GenerationTask: *, strategy_id: str = "", voice_library_id: str = "", + edit_plan_id: str = "", + template_id: str = "", + asset_ids: list[str] | None = None, + title_ids: list[str] | None = None, + voice_ids: list[str] | None = None, created_by_user_id: str = "", ) -> "GenerationTask": - if not project_id.strip(): - raise ValueError("project_id 不能为空") - if not asset_library_id.strip(): - raise ValueError("asset_library_id 不能为空") + if not project_id.strip() and not template_id.strip(): + raise ValueError("project_id 或 template_id 至少需要提供一个") + if not asset_library_id.strip() and not (asset_ids or title_ids or voice_ids): + raise ValueError("asset_library_id 或 asset_ids/title_ids/voice_ids 至少需要提供一个") return cls( id=uuid4().hex, project_id=project_id.strip(), 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(), + template_id=template_id.strip(), + asset_ids=list(asset_ids) if asset_ids else [], + title_ids=list(title_ids) if title_ids else [], + voice_ids=list(voice_ids) if voice_ids else [], created_by_user_id=created_by_user_id.strip(), ) diff --git a/packages/ports/generation_task_repository.py b/packages/ports/generation_task_repository.py index ffeac01dd..a68f059f7 100644 --- a/packages/ports/generation_task_repository.py +++ b/packages/ports/generation_task_repository.py @@ -12,4 +12,6 @@ class GenerationTaskRepository(Protocol): def list_by_project(self, project_id: str) -> list[GenerationTask]: ... + def list_by_user(self, user_id: str) -> list[GenerationTask]: ... + def update(self, task: GenerationTask) -> GenerationTask: ... -- 2.54.0 From d2097744e49be601f018bbf757a9cee28d2f6012 Mon Sep 17 00:00:00 2001 From: Audit Bot Date: Mon, 29 Jun 2026 17:19:13 +0800 Subject: [PATCH 2/2] =?UTF-8?q?fix:=20retry=20=E7=AB=AF=E7=82=B9=E5=A2=9E?= =?UTF-8?q?=E5=8A=A0=E5=BD=92=E5=B1=9E=E6=A0=A1=E9=AA=8C=20(P1=20=E5=AE=89?= =?UTF-8?q?=E5=85=A8=E4=BF=AE=E5=A4=8D)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - generation_tasks.py: retry_generation_task 增加 created_by_user_id 校验 - task_center.py: retry_task_by_id 增加 created_by_user_id 校验 - 非任务创建者返回 403 Access denied --- apps/api/app/api/routes/generation_tasks.py | 5 ++++- apps/api/app/api/routes/task_center.py | 2 ++ 2 files changed, 6 insertions(+), 1 deletion(-) diff --git a/apps/api/app/api/routes/generation_tasks.py b/apps/api/app/api/routes/generation_tasks.py index 462b92beb..db2210c31 100644 --- a/apps/api/app/api/routes/generation_tasks.py +++ b/apps/api/app/api/routes/generation_tasks.py @@ -216,7 +216,10 @@ def retry_generation_task( task = generation_task_repository.get(task_id) if task is None: raise HTTPException(status_code=404, detail="Generation task not found") - if task.status != "failed": + if task.created_by_user_id and task.created_by_user_id != authenticated_user.user.id: + raise HTTPException(status_code=403, detail="Access denied to this task") + status_val = task.status.value if hasattr(task.status, "value") else str(task.status) + if status_val != "failed": raise HTTPException(status_code=409, detail="Only failed tasks can be retried") use_case = CreateGenerationTaskUseCase(generation_task_repository) diff --git a/apps/api/app/api/routes/task_center.py b/apps/api/app/api/routes/task_center.py index 1d2e9524d..e899c8500 100644 --- a/apps/api/app/api/routes/task_center.py +++ b/apps/api/app/api/routes/task_center.py @@ -134,6 +134,8 @@ def retry_task_by_id( task = generation_task_repository.get(task_id) if task is None: raise HTTPException(status_code=404, detail="Generation task not found") + if task.created_by_user_id and task.created_by_user_id != authenticated_user.user.id: + raise HTTPException(status_code=403, detail="Access denied to this task") if _status_value(task.status) != "failed": raise HTTPException(status_code=409, detail="Only failed tasks can be retried") -- 2.54.0