94aead4342
CI/CD Pipeline / Validate Code Quality And Tests (push) Has been cancelled
CI/CD Pipeline / Unit Tests (push) Has been cancelled
CI/CD Pipeline / Integration Tests (push) Has been cancelled
CI/CD Pipeline / Frontend Lint (push) Has been cancelled
CI/CD Pipeline / Build & Push Staging (Watchtower auto-deploy) (push) Has been cancelled
CI/CD Pipeline / Staging E2E Tests (push) Has been cancelled
CI/CD Pipeline / Staging API Integration Tests (push) Has been cancelled
CI/CD Pipeline / Build Production Runtime Images (push) Has been cancelled
CI/CD Pipeline / Deploy Production (push) Has been cancelled
CI/CD Pipeline / Production Browser E2E (push) Has been cancelled
392 lines
15 KiB
Python
Executable File
392 lines
15 KiB
Python
Executable File
import logging
|
||
from typing import Any
|
||
|
||
from app.auth import AuthenticatedUser, get_current_user
|
||
from app.core.celery_app import celery_app
|
||
from app.core.task_enqueue import (
|
||
GLOBAL_PENDING_LIMIT,
|
||
USER_PENDING_LIMIT,
|
||
GlobalQueueFull,
|
||
UserPendingLimitExceeded,
|
||
safe_enqueue_generation_task,
|
||
)
|
||
from app.dependencies import (
|
||
get_generation_task_repository,
|
||
get_ingest_job_repository,
|
||
get_project_repository,
|
||
)
|
||
from app.schemas.task_center import (
|
||
ListProjectTasksResponse,
|
||
ListTasksResponse,
|
||
ProjectTaskResponse,
|
||
UserTaskResponse,
|
||
)
|
||
from fastapi import APIRouter, Depends, HTTPException, Query
|
||
|
||
from packages.application import (
|
||
CreateGenerationTaskCommand,
|
||
CreateGenerationTaskUseCase,
|
||
RetryGenerationTaskUseCase,
|
||
SubmitIngestJobCommand,
|
||
SubmitIngestJobUseCase,
|
||
)
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
router = APIRouter()
|
||
|
||
DEFAULT_PAGE_SIZE = 50
|
||
MAX_PAGE_SIZE = 200
|
||
|
||
|
||
def _humanize_task_error(error_message: str) -> str:
|
||
raw = (error_message or "").strip()
|
||
if not raw:
|
||
return ""
|
||
lower = raw.lower()
|
||
if "ffmpeg" in lower or "ffprobe" in lower or "invalid data" in lower or "moov atom" in lower:
|
||
return "视频素材格式无法识别,请重新导出为常见 MP4/H.264 后再试。"
|
||
if "oss" in lower or "bucket" in lower or "storage" in lower:
|
||
return "素材存储服务读取或写入失败,请稍后重试或联系小虾检查 OSS。"
|
||
if "not found" in lower or "no such file" in lower:
|
||
return "任务依赖的素材或文件不存在,请确认素材仍在项目中。"
|
||
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:
|
||
s = _status_value(task.status)
|
||
if s == "pending":
|
||
return "等待 Worker 执行"
|
||
if s == "running":
|
||
return "正在生成成片"
|
||
if s == "completed":
|
||
return "生成完成"
|
||
if s == "failed":
|
||
return "生成失败"
|
||
if s == "cancelled":
|
||
return "已取消"
|
||
return s
|
||
|
||
|
||
def _ingest_step(job) -> str:
|
||
s = _status_value(job.status)
|
||
if s == "pending":
|
||
return "等待导入"
|
||
if s == "processing":
|
||
return "正在分析素材"
|
||
if s == "completed":
|
||
return "导入完成"
|
||
if s == "failed":
|
||
return "导入失败"
|
||
return s
|
||
|
||
|
||
def _generation_task_to_user_response(task) -> UserTaskResponse:
|
||
return 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,
|
||
error_info=task.error_info or {},
|
||
user_message=_humanize_task_error(task.error_message),
|
||
retryable=_status_value(task.status) == "failed",
|
||
retry_count=task.retry_count or 0,
|
||
source_id=task.id,
|
||
created_at=task.created_at,
|
||
updated_at=task.completed_at or task.started_at or task.created_at,
|
||
)
|
||
|
||
|
||
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,
|
||
error_info=task.error_info or {},
|
||
user_message=_humanize_task_error(task.error_message),
|
||
retryable=_status_value(task.status) == "failed",
|
||
retry_count=task.retry_count or 0,
|
||
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,
|
||
)
|
||
|
||
|
||
def _validate_status(status: str | None) -> str | None:
|
||
"""校验状态值合法性。"""
|
||
if status is None:
|
||
return None
|
||
valid = {"pending", "running", "completed", "failed", "cancelled"}
|
||
if status not in valid:
|
||
raise HTTPException(
|
||
status_code=400,
|
||
detail=f"无效的状态筛选值: {status},允许值: {', '.join(sorted(valid))}",
|
||
)
|
||
return status
|
||
|
||
|
||
def _clamp_page_size(page_size: int) -> int:
|
||
if page_size <= 0:
|
||
return DEFAULT_PAGE_SIZE
|
||
if page_size > MAX_PAGE_SIZE:
|
||
return MAX_PAGE_SIZE
|
||
return page_size
|
||
|
||
|
||
# ── 用户级端点(放在项目级端点之前,避免路由冲突) ──
|
||
|
||
|
||
@router.get("/tasks", response_model=ListTasksResponse)
|
||
def list_user_tasks(
|
||
status: str | None = Query(None, description="按状态筛选:pending/running/completed/failed/cancelled"),
|
||
task_type: str | None = Query(None, description="按任务类型筛选:generation/ingest"),
|
||
page: int = Query(1, ge=1, description="页码,从1开始"),
|
||
page_size: int = Query(DEFAULT_PAGE_SIZE, ge=1, le=MAX_PAGE_SIZE, description="每页数量"),
|
||
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),支持状态/类型筛选和分页。"""
|
||
status = _validate_status(status)
|
||
page_size = _clamp_page_size(page_size)
|
||
user_id = authenticated_user.user.id
|
||
offset = (page - 1) * page_size
|
||
|
||
items: list[UserTaskResponse] = []
|
||
|
||
# 生成任务
|
||
if task_type is None or task_type == "generation":
|
||
gen_result = generation_task_repository.list_by_user_filtered(
|
||
user_id,
|
||
status=status,
|
||
limit=page_size + 1, # 多取一条判断是否还有下一页(简单起见这里用offset)
|
||
offset=offset,
|
||
)
|
||
for task in gen_result:
|
||
items.append(_generation_task_to_user_response(task))
|
||
|
||
# 按时间倒序
|
||
items.sort(key=lambda item: item.updated_at or item.created_at or "", reverse=True)
|
||
|
||
# 总数(仅generation,ingest暂不计入总数以保持简单)
|
||
total = generation_task_repository.count_by_user_filtered(user_id, status=status)
|
||
|
||
return ListTasksResponse(items=items[:page_size], total=total)
|
||
|
||
|
||
@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,retry_count+1)。"""
|
||
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")
|
||
|
||
user_id = authenticated_user.user.id
|
||
|
||
# 预检查
|
||
user_pending = generation_task_repository.count_pending_by_user(user_id)
|
||
global_pending = generation_task_repository.count_pending_total()
|
||
if user_pending >= USER_PENDING_LIMIT:
|
||
raise HTTPException(
|
||
status_code=429,
|
||
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
|
||
)
|
||
if global_pending >= GLOBAL_PENDING_LIMIT:
|
||
raise HTTPException(
|
||
status_code=503,
|
||
detail="系统繁忙,请稍后再试",
|
||
)
|
||
|
||
# 原地重试
|
||
use_case = RetryGenerationTaskUseCase(generation_task_repository)
|
||
retried = use_case.execute(task_id)
|
||
|
||
# 重新入队
|
||
try:
|
||
if not safe_enqueue_generation_task(
|
||
retried, generation_task_repository, user_id=user_id, log_prefix="[任务中心]"
|
||
):
|
||
logger.warning("[任务中心] 用户级重试入队失败: task_id=%s", retried.id)
|
||
except UserPendingLimitExceeded:
|
||
raise HTTPException(
|
||
status_code=429,
|
||
detail="您的待处理任务过多,请等待完成后再提交",
|
||
) from None
|
||
except GlobalQueueFull:
|
||
raise HTTPException(
|
||
status_code=503,
|
||
detail="系统繁忙,请稍后再试",
|
||
) from None
|
||
|
||
return _generation_task_to_user_response(retried)
|
||
|
||
|
||
# ── 项目级端点 ──
|
||
|
||
|
||
@router.get("/projects/{project_id}/tasks", response_model=ListProjectTasksResponse)
|
||
def list_project_tasks(
|
||
project_id: str,
|
||
status: str | None = Query(None, description="按状态筛选:pending/running/completed/failed/cancelled"),
|
||
task_type: str | None = Query(None, description="按任务类型筛选:generation/ingest"),
|
||
page: int = Query(1, ge=1, description="页码,从1开始"),
|
||
page_size: int = Query(DEFAULT_PAGE_SIZE, ge=1, le=MAX_PAGE_SIZE, description="每页数量"),
|
||
authenticated_user: AuthenticatedUser = Depends(get_current_user),
|
||
project_repository: Any = Depends(get_project_repository),
|
||
ingest_job_repository: Any = Depends(get_ingest_job_repository),
|
||
generation_task_repository: Any = Depends(get_generation_task_repository),
|
||
) -> ListProjectTasksResponse:
|
||
"""项目级任务列表,支持状态/类型筛选和分页。"""
|
||
project = project_repository.find_by_id(project_id)
|
||
if project is None:
|
||
raise HTTPException(status_code=404, detail="Project not found")
|
||
|
||
status = _validate_status(status)
|
||
page_size = _clamp_page_size(page_size)
|
||
offset = (page - 1) * page_size
|
||
|
||
items: list[ProjectTaskResponse] = []
|
||
|
||
# 导入任务
|
||
if task_type is None or task_type == "ingest":
|
||
for job in ingest_job_repository.list_by_project(project_id):
|
||
if status and _status_value(job.status) != status:
|
||
continue
|
||
items.append(
|
||
ProjectTaskResponse(
|
||
id=f"ingest:{job.id}",
|
||
task_type="ingest",
|
||
project_id=job.project_id,
|
||
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=_status_value(job.status) == "failed",
|
||
source_id=job.id,
|
||
created_at=job.created_at,
|
||
updated_at=job.updated_at,
|
||
)
|
||
)
|
||
|
||
# 生成任务
|
||
if task_type is None or task_type == "generation":
|
||
gen_items = generation_task_repository.list_by_project_filtered(
|
||
project_id,
|
||
status=status,
|
||
limit=page_size + 1,
|
||
offset=offset,
|
||
)
|
||
for task in gen_items:
|
||
items.append(_generation_task_to_project_response(task))
|
||
|
||
items.sort(key=lambda item: item.updated_at or item.created_at or "", reverse=True)
|
||
|
||
total = generation_task_repository.count_by_project_filtered(project_id, status=status)
|
||
|
||
return ListProjectTasksResponse(items=items[:page_size], total=total)
|
||
|
||
|
||
@router.post("/tasks/{task_type}/{source_id}/retry", response_model=ProjectTaskResponse)
|
||
def retry_project_task(
|
||
task_type: str,
|
||
source_id: str,
|
||
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),
|
||
) -> ProjectTaskResponse:
|
||
"""项目级任务重试。"""
|
||
if task_type == "generation":
|
||
task = generation_task_repository.get(source_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")
|
||
|
||
user_id = authenticated_user.user.id
|
||
|
||
# 预检查
|
||
user_pending = generation_task_repository.count_pending_by_user(user_id)
|
||
global_pending = generation_task_repository.count_pending_total()
|
||
if user_pending >= USER_PENDING_LIMIT:
|
||
raise HTTPException(
|
||
status_code=429,
|
||
detail=f"您的待处理任务过多(当前 {user_pending}/{USER_PENDING_LIMIT}),请等待完成后再提交",
|
||
)
|
||
if global_pending >= GLOBAL_PENDING_LIMIT:
|
||
raise HTTPException(
|
||
status_code=503,
|
||
detail="系统繁忙,请稍后再试",
|
||
)
|
||
|
||
# 原地重试
|
||
use_case = RetryGenerationTaskUseCase(generation_task_repository)
|
||
retried = use_case.execute(source_id)
|
||
|
||
try:
|
||
if not safe_enqueue_generation_task(
|
||
retried, generation_task_repository, user_id=user_id, log_prefix="[任务中心]"
|
||
):
|
||
logger.warning("[任务中心] 项目级重试入队失败: task_id=%s", retried.id)
|
||
except UserPendingLimitExceeded:
|
||
raise HTTPException(
|
||
status_code=429,
|
||
detail="您的待处理任务过多,请等待完成后再提交",
|
||
) from None
|
||
except GlobalQueueFull:
|
||
raise HTTPException(
|
||
status_code=503,
|
||
detail="系统繁忙,请稍后再试",
|
||
) from None
|
||
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 _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(
|
||
SubmitIngestJobCommand(
|
||
project_id=job.project_id,
|
||
library_id=job.library_id,
|
||
storage_key=job.storage_key,
|
||
)
|
||
)
|
||
celery_app.send_task("worker.ingest_asset", args=[retried.id])
|
||
return ProjectTaskResponse(
|
||
id=f"ingest:{retried.id}",
|
||
task_type="ingest",
|
||
project_id=retried.project_id,
|
||
status=_status_value(retried.status),
|
||
progress=0,
|
||
current_step=_ingest_step(retried),
|
||
source_id=retried.id,
|
||
created_at=retried.created_at,
|
||
updated_at=retried.updated_at,
|
||
)
|
||
raise HTTPException(status_code=400, detail="Unsupported task type")
|