Files
video-gen/video-gen-api/app/api/v1/generation_ai.py
T

425 lines
15 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from fastapi import APIRouter, Body, Depends, HTTPException, Path, Query
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.dependencies import get_current_user, get_db
from app.models.chat_generation_task import ChatGenerationTask
from app.models.user import User
from app.schemas.generation_ai import (
GenerationAIHistoryDayItemsOut,
GenerationAIHistoryGroupedOut,
GenerationAIRetryOut,
GenerationAITaskCreate,
GenerationAITaskListOut,
GenerationAITaskOut,
)
from app.services.generation_ai_service import (
create_async_generation_task,
list_async_generation_tasks,
list_generation_history_day_items,
list_generation_history_grouped_days,
record_to_out,
)
from app.services.generation_log_service import log_task_event
from app.tasks.celery_app import celery_app
router = APIRouter(
prefix="/generation-ai",
tags=["generation-ai"],
)
@router.post(
"/tasks",
response_model=GenerationAITaskOut,
summary="创建AI图片/视频生成任务",
description=(
"创建一个项目无关的AI生成任务。"
"该接口用于 Chat 风格的图片/视频生成,不再绑定 project_id。"
"创建成功后会写入 chat_generation_tasks 表,并投递 Celery 异步任务。"
"支持 image 图片生成和 video 视频生成。"
"建议前端传入 idempotency_key,用于防止按钮连点、网络重试导致重复创建任务和重复扣费。"
),
responses={
200: {
"description": "任务创建成功,返回任务详情",
},
400: {
"description": "请求参数错误,例如 gen_type 不支持、图片尺寸不支持、视频参数不支持等",
},
401: {
"description": "未登录或 Token 无效",
},
503: {
"description": "Celery 未启用或消息队列未配置",
},
},
)
async def create_task(
req: GenerationAITaskCreate = Body(
...,
description="AI生成任务创建参数。gen_type=image 时使用图片参数;gen_type=video 时使用视频参数",
),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
if celery_app is None:
raise HTTPException(status_code=503, detail="Celery未启用:请配置 REDIS_URL 或 CELERY_BROKER_URL 后启动 worker")
task = await create_async_generation_task(db, current_user, req)
await db.commit()
await log_task_event(
task,
event_type="TASK_CREATED",
to_status="generating",
to_stage="queued",
detail={"gen_type": task.gen_type},
)
from app.tasks.generation_create_tasks import chatapi_create_generation_task
chatapi_create_generation_task.delay(task.id)
return record_to_out(task)
@router.get(
"/tasks",
response_model=GenerationAITaskListOut,
summary="获取AI生成任务列表",
description=(
"分页获取当前登录用户的AI生成任务列表。"
"可按生成类型 gen_type 和任务状态 status 过滤。"
"该接口返回的是普通任务列表,不按日期分组。"
"如果前端需要按生成日期分组展示历史记录,请使用 /generation-ai/history 接口。"
),
responses={
200: {
"description": "查询成功,返回任务总数和当前分页任务列表",
},
401: {
"description": "未登录或 Token 无效",
},
},
)
async def list_tasks(
gen_type: str | None = Query(
None,
description="生成类型筛选:image=图片任务,video=视频任务;为空表示不过滤生成类型",
examples=["image"],
),
status: str | None = Query(
None,
description=(
"任务状态筛选。常见值:generating=生成中,completed=已完成,failed=失败;"
"为空表示不过滤状态"
),
examples=["completed"],
),
page: int = Query(
1,
ge=1,
description="分页页码,从1开始",
examples=[1],
),
page_size: int = Query(
20,
ge=1,
le=100,
description="每页返回数量,范围 1~100",
examples=[20],
),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
total, items = await list_async_generation_tasks(
db,
current_user.id,
gen_type,
status,
page,
page_size,
)
# ====================== 在这里加排序(最新在前)======================
# 按 created_at 降序(没有则用 id 降序)
items_sorted = sorted(
items,
key=lambda x: x.created_at if x.created_at is not None else x.id,
reverse=False # 降序
)
return GenerationAITaskListOut(total=total, items=[record_to_out(i) for i in items_sorted])
@router.get(
"/history",
response_model=GenerationAIHistoryGroupedOut,
summary="获取AI生成历史日期分组",
description=(
"按生成完成日期倒序返回当前用户的AI生成历史记录。"
"默认查询 chat_generation_tasks / ChatGenerationTask 新任务历史。"
"当 history_source=generation_record 时,查询 generation_records / GenerationRecord 旧历史。"
"该接口只返回生成成功的任务,即 status=completed 且 generated_at 不为空的数据。"
"必须通过 gen_type 区分图片和视频。"
"分页对象是生成日期,不是单条记录。"
"每页最多返回10个生成日期分组,每个日期分组内最多返回该日期下倒序前10条生成记录。"
"如果某一天 total 大于10,前端可调用 /generation-ai/history/{generated_date} 加载该日期下的后续分页数据。"
),
responses={
200: {
"description": "查询成功,返回按生成日期分组的历史记录",
},
400: {
"description": "参数错误,例如 gen_type 不是 image 或 video,或 history_source 不支持",
},
401: {
"description": "未登录或 Token 无效",
},
},
)
async def list_history_grouped_days(
gen_type: str = Query(
...,
description="生成类型:image=图片历史,video=视频历史",
examples=["image"],
),
page: int = Query(
1,
ge=1,
description="日期分组分页页码,从1开始。注意:这里分页的是生成日期,不是单条生成记录",
examples=[1],
),
page_size: int = Query(
10,
ge=1,
le=10,
description="每页返回的生成日期数量,范围 1~10,最大只能获取10天",
examples=[10],
),
history_source: str | None = Query(
None,
description=(
"历史数据来源。默认不传或传 chat_task 查询 chat_generation_tasks / ChatGenerationTask"
"传 generation_record 查询 generation_records / GenerationRecord 旧历史数据"
),
examples=["generation_record"],
),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
return await list_generation_history_grouped_days(
db=db,
user_id=current_user.id,
gen_type=gen_type,
page=page,
page_size=page_size,
history_source=history_source,
)
@router.get(
"/history/{generated_date}",
response_model=GenerationAIHistoryDayItemsOut,
summary="获取指定日期下的AI生成历史分页",
description=(
"获取某一个生成日期下的生成成功记录分页。"
"默认查询 chat_generation_tasks / ChatGenerationTask 新任务历史。"
"当 history_source=generation_record 时,查询 generation_records / GenerationRecord 旧历史。"
"该接口用于前端在历史分组列表中继续加载某一天的后续记录。"
"例如 /history 接口中某一天 total=18,但 items 只返回前10条,"
"则前端可以调用本接口 page=2&page_size=10 获取该日期下剩余记录。"
"该接口同样必须通过 gen_type 区分图片和视频。"
),
responses={
200: {
"description": "查询成功,返回指定日期下的历史记录分页",
},
400: {
"description": "参数错误,例如 generated_date 格式不是 YYYY-MM-DDgen_type 不合法,或 history_source 不支持",
},
401: {
"description": "未登录或 Token 无效",
},
},
)
async def list_history_day_items(
generated_date: str = Path(
...,
description="生成日期,格式:YYYY-MM-DD,例如:2026-05-27",
examples=["2026-05-27"],
),
gen_type: str = Query(
...,
description="生成类型:image=图片历史,video=视频历史",
examples=["image"],
),
page: int = Query(
1,
ge=1,
description="当前日期下的记录分页页码,从1开始",
examples=[1],
),
page_size: int = Query(
10,
ge=1,
le=100,
description="当前日期下每页返回的生成记录数量,范围 1~100",
examples=[10],
),
history_source: str | None = Query(
None,
description=(
"历史数据来源。默认不传或传 chat_task 查询 chat_generation_tasks / ChatGenerationTask"
"传 generation_record 查询 generation_records / GenerationRecord 旧历史数据"
),
examples=["generation_record"],
),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
return await list_generation_history_day_items(
db=db,
user_id=current_user.id,
gen_type=gen_type,
generated_date=generated_date,
page=page,
page_size=page_size,
history_source=history_source,
)
@router.get(
"/tasks/{task_id}",
response_model=GenerationAITaskOut,
summary="获取AI生成任务详情",
description=(
"根据任务ID获取当前登录用户的AI生成任务详情。"
"只能查询当前用户自己的任务,且只查询 generation_mode=chatapi_async 的任务。"
"如果任务不存在或不属于当前用户,返回404。"
),
responses={
200: {
"description": "查询成功,返回任务详情",
},
401: {
"description": "未登录或 Token 无效",
},
404: {
"description": "任务不存在,或任务不属于当前用户",
},
},
)
async def get_task(
task_id: str = Path(
...,
description="AI生成任务ID",
examples=["0019e0a44895b6d837d"],
),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
result = await db.execute(
select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.user_id == current_user.id,
ChatGenerationTask.generation_mode == "chatapi_async",
)
)
task = result.scalar_one_or_none()
if not task:
raise HTTPException(status_code=404, detail="任务不存在")
return record_to_out(task)
@router.post(
"/tasks/{task_id}/retry",
response_model=GenerationAIRetryOut,
summary="重试失败的AI生成任务",
description=(
"重新投递一个失败的AI生成任务。"
"只有 status=failed 的任务允许重试。"
"重试次数超过3次后不允许继续重试。"
"后端会根据任务当前数据自动判断应从哪个流水线阶段恢复,"
"例如重新优化提示词、重新创建第三方任务、继续轮询远程结果或重新下载结果。"
),
responses={
200: {
"description": "任务重新投递成功",
},
400: {
"description": "任务状态不允许重试,或已超过最大重试次数",
},
401: {
"description": "未登录或 Token 无效",
},
404: {
"description": "任务不存在,或任务不属于当前用户",
},
503: {
"description": "Celery 未启用或消息队列未配置",
},
},
)
async def retry_task(
task_id: str = Path(
...,
description="需要重试的AI生成任务ID",
examples=["0019e0a44895b6d837d"],
),
current_user: User = Depends(get_current_user),
db: AsyncSession = Depends(get_db),
):
if celery_app is None:
raise HTTPException(status_code=503, detail="Celery未启用:请配置 REDIS_URL 或 CELERY_BROKER_URL 后启动 worker")
result = await db.execute(
select(ChatGenerationTask).where(
ChatGenerationTask.id == task_id,
ChatGenerationTask.user_id == current_user.id,
ChatGenerationTask.generation_mode == "chatapi_async",
)
)
task = result.scalar_one_or_none()
if not task:
raise HTTPException(status_code=404, detail="任务不存在")
if task.status != "failed":
raise HTTPException(status_code=400, detail="只有失败任务可以重试")
task.status = "generating"
task.error_message = None
task.retry_count = (task.retry_count or 0) + 1
if task.retry_count > 3:
raise HTTPException(status_code=400, detail="任务已超过最大重试次数")
# Resume from the earliest missing stage.
if not task.optimized_prompt:
task.pipeline_stage = "queued"
from app.tasks.generation_create_tasks import chatapi_create_generation_task
chatapi_create_generation_task.delay(task.id)
elif not task.seedance_task_id and not task.remote_result_url:
task.pipeline_stage = "creating_provider_task"
from app.tasks.generation_create_tasks import chatapi_create_generation_task
chatapi_create_generation_task.delay(task.id)
elif task.seedance_task_id and not task.remote_result_url:
task.pipeline_stage = "waiting_remote"
from app.tasks.generation_poll_tasks import poll_generation_task
poll_generation_task.delay(task.id)
elif task.remote_result_url and not (task.image_url or task.video_url):
task.pipeline_stage = "result_ready"
from app.tasks.generation_download_tasks import download_generation_result_task
download_generation_result_task.delay(task.id)
else:
task.status = "completed"
task.pipeline_stage = "done"
await db.commit()
return GenerationAIRetryOut(
id=task.id,
status=task.status,
pipeline_stage=task.pipeline_stage,
message="任务已重新投递",
)