feat: 扩展生成任务API支持模板模式(方案A) #109
@@ -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")
|
||||
@@ -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,43 @@ 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.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)
|
||||
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)
|
||||
|
||||
@@ -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,137 @@ 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 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")
|
||||
|
||||
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 +191,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 +220,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 +230,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 +258,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,
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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))
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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(),
|
||||
)
|
||||
|
||||
@@ -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: ...
|
||||
|
||||
Reference in New Issue
Block a user