AI创作批量生成任务 main V1 init
This commit is contained in:
@@ -1,7 +1,8 @@
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime
|
||||
|
||||
from fastapi import APIRouter, Body, Depends, HTTPException, Path, Query
|
||||
from sqlalchemy import and_, select
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.dependencies import get_current_user, get_db
|
||||
@@ -19,25 +20,35 @@ from app.schemas.generation_ai import (
|
||||
GenerationAITaskListOut,
|
||||
GenerationAITaskOut,
|
||||
)
|
||||
from app.services.generation_ai_service import (
|
||||
create_async_generation_task,
|
||||
from app.services.generation.ai.service import (
|
||||
build_task_out_list,
|
||||
list_generation_ai_engine_options,
|
||||
list_async_generation_tasks,
|
||||
list_generation_history_day_items,
|
||||
list_generation_history_grouped_days,
|
||||
record_to_out,
|
||||
soft_delete_chat_generation_task,
|
||||
)
|
||||
from app.services.generation_billing_service import (
|
||||
from app.enums.generation_task import ChatGenerationPipelineStage, ChatGenerationTaskStatus, GenerationMode
|
||||
from app.services.generation.ai.task_create_service import (
|
||||
GenerationTaskCreateResult,
|
||||
create_generation_task_group,
|
||||
enqueue_created_generation_tasks,
|
||||
find_existing_top_level_task,
|
||||
)
|
||||
from app.services.generation.ai.task_group_service import (
|
||||
aggregate_main_task_status,
|
||||
load_children_map,
|
||||
soft_delete_child_task,
|
||||
soft_delete_top_level_task_group,
|
||||
)
|
||||
from app.services.generation.billing_service import (
|
||||
OWNER_CHAT_GENERATION_TASK,
|
||||
charge_generation_media_by_params,
|
||||
get_next_credit_attempt_no,
|
||||
)
|
||||
from app.services.generation_history_delete_service import batch_delete_generation_history_items
|
||||
from app.services.generation_log_service import log_task_event
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls, resolve_private_portrait_reference_display_urls
|
||||
from app.services.generation.history_delete_service import batch_delete_generation_history_items
|
||||
from app.services.generation.log_service import log_task_event
|
||||
from app.services.resource_capacity_service import assert_user_resource_capacity_available
|
||||
from app.services.operation_log_service import log_operation_event
|
||||
from app.tasks.celery_app import celery_app
|
||||
|
||||
router = APIRouter(
|
||||
@@ -146,6 +157,7 @@ async def create_task(
|
||||
...,
|
||||
description=(
|
||||
"AI生成任务创建参数。gen_type=image 时使用图片参数;gen_type=video 时使用视频参数。"
|
||||
"generation_count 为客户端本次选择的生成数量,默认1,后端会按引擎开关和数量上限校验。"
|
||||
"枚举:gen_type=image/video;media_references[].type=image/video/audio;"
|
||||
"media_references[].source=upload_resource/private_portrait_asset/空;"
|
||||
"media_references[].role=first_frame/last_frame/reference_image/reference_video/reference_audio。"
|
||||
@@ -157,34 +169,88 @@ async def create_task(
|
||||
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()
|
||||
try:
|
||||
create_result = await create_generation_task_group(db, current_user, req)
|
||||
top_level_task_id = str(create_result.top_level_task_id)
|
||||
enqueue_task_ids = list(create_result.enqueue_task_ids)
|
||||
await db.commit()
|
||||
except IntegrityError:
|
||||
# 并发重复请求可能同时通过预查询;唯一索引负责兜底。
|
||||
# 回滚本次任务和计费后,按幂等键返回已经成功提交的顶层任务。
|
||||
await db.rollback()
|
||||
existing = await find_existing_top_level_task(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
idempotency_key=req.idempotency_key,
|
||||
)
|
||||
if not existing:
|
||||
raise
|
||||
create_result = GenerationTaskCreateResult(
|
||||
top_level_task_id=str(existing.id),
|
||||
generation_count=int(existing.generation_count or 1),
|
||||
gen_type=str(existing.gen_type),
|
||||
created=False,
|
||||
)
|
||||
top_level_task_id = str(existing.id)
|
||||
enqueue_task_ids = []
|
||||
|
||||
if create_result.created:
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="BATCH_COMMIT_SUCCESS",
|
||||
event_status="success",
|
||||
source="api",
|
||||
user_id=current_user.id,
|
||||
group_id=top_level_task_id,
|
||||
task_id=top_level_task_id,
|
||||
detail={
|
||||
"gen_type": create_result.gen_type,
|
||||
"generation_count": create_result.generation_count,
|
||||
"child_task_ids": create_result.child_task_ids,
|
||||
"physical_files_deleted": False,
|
||||
},
|
||||
)
|
||||
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type="TASK_CREATED",
|
||||
to_status="generating",
|
||||
to_stage="queued",
|
||||
detail={"gen_type": task.gen_type},
|
||||
task_id=top_level_task_id,
|
||||
event_type=(
|
||||
"TASK_CREATED" if create_result.created else "IDEMPOTENCY_HIT"
|
||||
),
|
||||
to_status="generating" if create_result.created else None,
|
||||
to_stage="queued" if create_result.created else None,
|
||||
detail={
|
||||
"gen_type": create_result.gen_type,
|
||||
"generation_count": create_result.generation_count,
|
||||
"child_task_ids": create_result.child_task_ids,
|
||||
"created": create_result.created,
|
||||
},
|
||||
)
|
||||
|
||||
from app.tasks.generation_create_tasks import chatapi_create_generation_task
|
||||
|
||||
try:
|
||||
chatapi_create_generation_task.delay(task.id)
|
||||
except Exception as exc:
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
failed_enqueue_ids: list[str] = []
|
||||
if create_result.created and enqueue_task_ids:
|
||||
failed_enqueue_ids = await enqueue_created_generation_tasks(
|
||||
db,
|
||||
task_id=task.id,
|
||||
error_message=f"任务队列投递失败: {exc}",
|
||||
pipeline_stage="failed",
|
||||
task_ids=enqueue_task_ids,
|
||||
)
|
||||
await db.commit()
|
||||
raise HTTPException(status_code=503, detail="任务队列投递失败,请稍后重试")
|
||||
|
||||
refs = await resolve_private_portrait_reference_display_urls(db, record_to_out(task).media_references, user_id=current_user.id)
|
||||
return record_to_out(task, media_references=refs)
|
||||
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.id == top_level_task_id,
|
||||
ChatGenerationTask.user_id == current_user.id,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
).limit(1)
|
||||
)
|
||||
task = result.scalar_one_or_none()
|
||||
if not task:
|
||||
raise HTTPException(status_code=404, detail="任务创建后未找到")
|
||||
output = await build_task_out_list(
|
||||
db,
|
||||
[task],
|
||||
viewer_user_id=current_user.id,
|
||||
)
|
||||
if failed_enqueue_ids and len(failed_enqueue_ids) == len(enqueue_task_ids):
|
||||
raise HTTPException(status_code=503, detail="任务已创建,但任务队列投递失败,请稍后重试")
|
||||
return output[0]
|
||||
|
||||
@router.get(
|
||||
"/tasks",
|
||||
@@ -261,10 +327,8 @@ async def list_tasks(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
is_admin = False
|
||||
if current_user.user_type == 'admin':
|
||||
is_admin = True
|
||||
else:
|
||||
is_admin = current_user.user_type == "admin"
|
||||
if not is_admin:
|
||||
user_id = current_user.id
|
||||
|
||||
total, items = await list_async_generation_tasks(
|
||||
@@ -280,24 +344,17 @@ async def list_tasks(
|
||||
created_start=created_start,
|
||||
created_end=created_end,
|
||||
)
|
||||
|
||||
# ====================== 在这里加排序(最新在前)======================
|
||||
if not is_admin:
|
||||
# 按 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 # 升序
|
||||
)
|
||||
else:
|
||||
items_sorted = items
|
||||
refs_map = await batch_resolve_private_portrait_reference_display_urls(
|
||||
# 同一个 API 同时服务管理后台和客户端:
|
||||
# - 管理员保持数据库倒序,最新记录在列表上方;
|
||||
# - 普通用户先查询最新一页,再仅反转当前页,聊天消息从旧到新排列。
|
||||
items_for_output = items if is_admin else list(reversed(items))
|
||||
out_items = await build_task_out_list(
|
||||
db,
|
||||
{item.id: record_to_out(task=item, is_admin=is_admin).media_references for item in items_sorted},
|
||||
user_id=None if is_admin else current_user.id,
|
||||
items_for_output,
|
||||
is_admin=is_admin,
|
||||
viewer_user_id=None if is_admin else current_user.id,
|
||||
)
|
||||
return GenerationAITaskListOut(total=total, items=[record_to_out(task=i, is_admin=is_admin, media_references=refs_map.get(i.id)) for i in items_sorted])
|
||||
|
||||
return GenerationAITaskListOut(total=total, items=out_items)
|
||||
|
||||
@router.get(
|
||||
"/history",
|
||||
@@ -513,8 +570,9 @@ async def list_history_day_items(
|
||||
summary="获取AI生成任务详情",
|
||||
description=(
|
||||
"根据任务ID获取当前登录用户的AI生成任务详情。"
|
||||
"只能查询当前用户自己的任务,且只查询 generation_mode=chatapi_async 的任务。"
|
||||
"如果任务不存在或不属于当前用户,返回404。"
|
||||
"支持 chatapi_async、chatapi_main 和未删除的 chatapi_child。"
|
||||
"查询 chatapi_main 时返回按 generation_index 升序排列的 child_items。"
|
||||
"已软删除 child 只在父任务 child_items 中保留槽位,不能通过 child ID 单独查询。"
|
||||
),
|
||||
responses={
|
||||
200: {
|
||||
@@ -541,17 +599,19 @@ async def get_task(
|
||||
select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.user_id == current_user.id,
|
||||
ChatGenerationTask.generation_mode == "chatapi_async",
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.limit(1)
|
||||
).limit(1)
|
||||
)
|
||||
task = result.scalar_one_or_none()
|
||||
if not task:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
refs = await resolve_private_portrait_reference_display_urls(db, record_to_out(task).media_references, user_id=current_user.id)
|
||||
return record_to_out(task, media_references=refs)
|
||||
|
||||
if task.deleted_at is not None:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
output = await build_task_out_list(
|
||||
db,
|
||||
[task],
|
||||
viewer_user_id=current_user.id,
|
||||
)
|
||||
return output[0]
|
||||
|
||||
@router.delete(
|
||||
"/tasks/{task_id}",
|
||||
@@ -587,38 +647,33 @@ async def delete_task(
|
||||
current_user: User = Depends(get_current_user),
|
||||
db: AsyncSession = Depends(get_db),
|
||||
):
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask).where(
|
||||
mode_result = await db.execute(
|
||||
select(ChatGenerationTask.generation_mode).where(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.user_id == current_user.id,
|
||||
ChatGenerationTask.generation_mode == "chatapi_async",
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
).limit(1)
|
||||
)
|
||||
generation_mode = mode_result.scalar_one_or_none()
|
||||
if generation_mode == GenerationMode.CHATAPI_CHILD.value:
|
||||
freed_size_bytes = await soft_delete_child_task(
|
||||
db,
|
||||
child_task_id=task_id,
|
||||
user_id=current_user.id,
|
||||
)
|
||||
.limit(1)
|
||||
)
|
||||
task = result.scalar_one_or_none()
|
||||
if not task:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
|
||||
if task.status == "generating":
|
||||
raise HTTPException(status_code=400, detail="当前任务正在生成中,暂不能删除")
|
||||
|
||||
deleted_at = datetime.now(timezone.utc)
|
||||
freed_size_bytes = await soft_delete_chat_generation_task(
|
||||
db,
|
||||
task=task,
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
await db.flush()
|
||||
|
||||
else:
|
||||
freed_size_bytes = await soft_delete_top_level_task_group(
|
||||
db,
|
||||
task_id=task_id,
|
||||
user_id=current_user.id,
|
||||
)
|
||||
await db.commit()
|
||||
return GenerationAITaskDeleteOut(
|
||||
message="任务已删除",
|
||||
task_id=task.id,
|
||||
task_id=task_id,
|
||||
deleted=True,
|
||||
freed_size_bytes=freed_size_bytes,
|
||||
)
|
||||
|
||||
|
||||
@router.post(
|
||||
"/tasks/{task_id}/retry",
|
||||
response_model=GenerationAIRetryOut,
|
||||
@@ -664,74 +719,146 @@ async def retry_task(
|
||||
select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.user_id == current_user.id,
|
||||
ChatGenerationTask.generation_mode == "chatapi_async",
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
).with_for_update().limit(1)
|
||||
)
|
||||
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="只有失败任务可以重试")
|
||||
|
||||
retry_targets: list[ChatGenerationTask]
|
||||
retrying_group_children = False
|
||||
if task.generation_mode == GenerationMode.CHATAPI_MAIN.value:
|
||||
children_map = await load_children_map(db, [task.id], include_deleted=False)
|
||||
children = children_map.get(task.id, [])
|
||||
if task.gen_type == "video":
|
||||
retry_targets = [
|
||||
child for child in children
|
||||
if child.status == ChatGenerationTaskStatus.FAILED.value
|
||||
]
|
||||
retrying_group_children = True
|
||||
if not retry_targets:
|
||||
raise HTTPException(status_code=400, detail="当前视频任务组没有可重试的失败子任务")
|
||||
elif children:
|
||||
# 图片供应商全部成功后才会拆子任务;已有子任务时只允许重试下载,
|
||||
# 不能再次扣费并覆盖原有生成序号。
|
||||
retry_targets = [
|
||||
child for child in children
|
||||
if child.status == ChatGenerationTaskStatus.FAILED.value
|
||||
and child.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value
|
||||
and bool(child.remote_result_url)
|
||||
]
|
||||
retrying_group_children = True
|
||||
if not retry_targets:
|
||||
raise HTTPException(status_code=400, detail="当前图片任务组没有可重试的下载失败子任务")
|
||||
else:
|
||||
# 图片批次在供应商阶段整批失败时尚未创建子任务,可整批重新生成并重新计费。
|
||||
if task.status != ChatGenerationTaskStatus.FAILED.value:
|
||||
raise HTTPException(status_code=400, detail="只有失败任务可以重试")
|
||||
retry_targets = [task]
|
||||
else:
|
||||
if task.status != ChatGenerationTaskStatus.FAILED.value:
|
||||
raise HTTPException(status_code=400, detail="只有失败任务可以重试")
|
||||
retry_targets = [task]
|
||||
|
||||
await assert_user_resource_capacity_available(db, current_user.id)
|
||||
enqueue_ids: list[str] = []
|
||||
download_retry_ids: list[str] = []
|
||||
for target in retry_targets:
|
||||
if int(target.retry_count or 0) >= 3:
|
||||
raise HTTPException(status_code=400, detail=f"任务 {target.id} 已超过最大重试次数")
|
||||
|
||||
attempt_no = await get_next_credit_attempt_no(
|
||||
db,
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
owner_id=task.id,
|
||||
)
|
||||
media_billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=task.user_id,
|
||||
record_id=task.id,
|
||||
gen_type=task.gen_type,
|
||||
image_size=task.image_size,
|
||||
duration=task.duration,
|
||||
resolution=task.resolution,
|
||||
engine_id=task.engine_id,
|
||||
project_name="AI生成任务",
|
||||
description_prefix="Chat任务重试",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
attempt_no=attempt_no,
|
||||
)
|
||||
is_download_retry = bool(
|
||||
target.remote_result_url
|
||||
and target.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value
|
||||
)
|
||||
if not is_download_retry:
|
||||
attempt_no = await get_next_credit_attempt_no(
|
||||
db,
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
owner_id=target.id,
|
||||
)
|
||||
quantity = int(target.generation_count or 1) if (
|
||||
target.generation_mode == GenerationMode.CHATAPI_MAIN.value and target.gen_type == "image"
|
||||
) else 1
|
||||
media_billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=target.user_id,
|
||||
record_id=target.id,
|
||||
gen_type=target.gen_type,
|
||||
image_size=target.image_size,
|
||||
duration=target.duration,
|
||||
resolution=target.resolution,
|
||||
engine_id=target.engine_id,
|
||||
project_name="AI生成任务",
|
||||
description_prefix="Chat任务重试",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
attempt_no=attempt_no,
|
||||
quantity=quantity,
|
||||
)
|
||||
target.credits_cost = round(float(target.credits_cost or 0) + media_billing.total_charged, 2)
|
||||
target.provider_task_id = None
|
||||
target.seedance_task_id = None
|
||||
target.remote_result_url = None
|
||||
target.provider_response_json = None
|
||||
target.provider_create_claim_token = None
|
||||
target.provider_create_lease_until = None
|
||||
target.provider_create_started_at = None
|
||||
target.image_url = None
|
||||
target.video_url = None
|
||||
target.video_cover_url = None
|
||||
target.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
|
||||
enqueue_ids.append(str(target.id))
|
||||
else:
|
||||
target.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value
|
||||
download_retry_ids.append(str(target.id))
|
||||
|
||||
task.status = "generating"
|
||||
task.pipeline_stage = "queued"
|
||||
task.error_message = None
|
||||
task.poll_count = 0
|
||||
task.last_poll_at = None
|
||||
task.provider_task_id = None
|
||||
task.seedance_task_id = None
|
||||
task.remote_result_url = None
|
||||
task.provider_response_json = None
|
||||
task.image_url = None
|
||||
task.video_url = None
|
||||
task.video_cover_url = None
|
||||
task.generated_at = None
|
||||
task.credits_cost = round(float(task.credits_cost or 0) + media_billing.total_charged, 2)
|
||||
target.status = ChatGenerationTaskStatus.GENERATING.value
|
||||
target.error_message = None
|
||||
target.poll_count = 0
|
||||
target.last_poll_at = None
|
||||
target.generated_at = None
|
||||
target.retry_count = int(target.retry_count or 0) + 1
|
||||
|
||||
if retrying_group_children:
|
||||
await db.flush()
|
||||
await aggregate_main_task_status(db, parent_task_id=str(task.id))
|
||||
|
||||
refreshed_task_id = str(task.id)
|
||||
await db.commit()
|
||||
|
||||
from app.tasks.generation_create_tasks import chatapi_create_generation_task
|
||||
failed_enqueue_ids = await enqueue_created_generation_tasks(db, task_ids=enqueue_ids) if enqueue_ids else []
|
||||
failed_download_enqueue_ids: list[str] = []
|
||||
if download_retry_ids:
|
||||
from app.tasks.generation_download_tasks import enqueue_download_task
|
||||
for target_id in download_retry_ids:
|
||||
target_result = await db.execute(
|
||||
select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.id == target_id,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
).limit(1)
|
||||
)
|
||||
target = target_result.scalar_one_or_none()
|
||||
if not target or not await enqueue_download_task(db, target, recover=True, reason="manual_retry"):
|
||||
failed_download_enqueue_ids.append(target_id)
|
||||
|
||||
try:
|
||||
chatapi_create_generation_task.delay(task.id)
|
||||
except Exception as exc:
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task_id=task.id,
|
||||
error_message=f"任务队列投递失败: {exc}",
|
||||
pipeline_stage="failed",
|
||||
)
|
||||
await db.commit()
|
||||
raise HTTPException(status_code=503, detail="任务队列投递失败,请稍后重试")
|
||||
requested_enqueue_count = len(enqueue_ids) + len(download_retry_ids)
|
||||
failed_total_count = len(failed_enqueue_ids) + len(failed_download_enqueue_ids)
|
||||
if requested_enqueue_count and failed_total_count == requested_enqueue_count:
|
||||
raise HTTPException(status_code=503, detail="任务状态已重置,但任务队列投递全部失败,将由恢复任务继续处理")
|
||||
|
||||
refreshed = await db.execute(
|
||||
select(ChatGenerationTask).where(ChatGenerationTask.id == refreshed_task_id).limit(1)
|
||||
)
|
||||
refreshed_task = refreshed.scalar_one_or_none()
|
||||
if not refreshed_task:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
return GenerationAIRetryOut(
|
||||
id=task.id,
|
||||
status=task.status,
|
||||
pipeline_stage=task.pipeline_stage,
|
||||
message="任务已重新扣费并重新投递",
|
||||
id=refreshed_task.id,
|
||||
status=refreshed_task.status,
|
||||
pipeline_stage=refreshed_task.pipeline_stage,
|
||||
message=(
|
||||
f"请求重试 {len(retry_targets)} 个任务,成功投递 {max(0, requested_enqueue_count - failed_total_count)} 个,"
|
||||
f"投递失败 {failed_total_count} 个"
|
||||
),
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user