1
This commit is contained in:
@@ -0,0 +1,230 @@
|
||||
"""add client-selectable multi generation and image batch claim
|
||||
|
||||
Revision ID: abae3e1c70f7
|
||||
Revises: 2026070902
|
||||
Create Date: 2026-07-15 10:49:31.803342
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = "abae3e1c70f7"
|
||||
down_revision: Union[str, None] = "2026070902"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
FK_CHAT_TASK_PARENT = "fk_chat_generation_tasks_parent_task_id"
|
||||
CK_CHAT_TASK_GENERATION_COUNT = "ck_chat_generation_tasks_generation_count"
|
||||
CK_CHAT_TASK_GENERATION_INDEX = "ck_chat_generation_tasks_generation_index"
|
||||
CK_IMAGE_ENGINE_MAX_GENERATION_COUNT = "ck_image_engines_max_generation_count"
|
||||
CK_IMAGE_ENGINE_MULTI_IMAGE_MAX = "ck_image_engines_multi_image_max_images"
|
||||
CK_IMAGE_ENGINE_MAX_REFERENCE = "ck_image_engines_max_reference_image_count"
|
||||
CK_VIDEO_ENGINE_MAX_GENERATION_COUNT = "ck_video_engines_max_generation_count"
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
# ChatGenerationTask:任务级实际生成数量、主子关联和图片批次执行租约。
|
||||
op.add_column(
|
||||
"chat_generation_tasks",
|
||||
sa.Column("parent_task_id", sa.String(length=32), nullable=True),
|
||||
)
|
||||
op.add_column(
|
||||
"chat_generation_tasks",
|
||||
sa.Column("generation_count", sa.Integer(), server_default=sa.text("1"), nullable=False),
|
||||
)
|
||||
op.add_column(
|
||||
"chat_generation_tasks",
|
||||
sa.Column("generation_index", sa.Integer(), nullable=True),
|
||||
)
|
||||
op.add_column(
|
||||
"chat_generation_tasks",
|
||||
sa.Column("provider_create_claim_token", sa.String(length=64), nullable=True),
|
||||
)
|
||||
op.add_column(
|
||||
"chat_generation_tasks",
|
||||
sa.Column("provider_create_lease_until", sa.DateTime(timezone=True), nullable=True),
|
||||
)
|
||||
op.add_column(
|
||||
"chat_generation_tasks",
|
||||
sa.Column("provider_create_started_at", sa.DateTime(timezone=True), nullable=True),
|
||||
)
|
||||
|
||||
op.create_check_constraint(
|
||||
CK_CHAT_TASK_GENERATION_COUNT,
|
||||
"chat_generation_tasks",
|
||||
"generation_count BETWEEN 1 AND 5",
|
||||
)
|
||||
op.create_check_constraint(
|
||||
CK_CHAT_TASK_GENERATION_INDEX,
|
||||
"chat_generation_tasks",
|
||||
"generation_index IS NULL OR generation_index > 0",
|
||||
)
|
||||
op.create_foreign_key(
|
||||
FK_CHAT_TASK_PARENT,
|
||||
"chat_generation_tasks",
|
||||
"chat_generation_tasks",
|
||||
["parent_task_id"],
|
||||
["id"],
|
||||
ondelete="RESTRICT",
|
||||
)
|
||||
op.create_index(
|
||||
"idx_chat_generation_tasks_parent",
|
||||
"chat_generation_tasks",
|
||||
["parent_task_id"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(
|
||||
"idx_chat_generation_tasks_user_mode_created",
|
||||
"chat_generation_tasks",
|
||||
["user_id", "generation_mode", "created_at"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(
|
||||
"ix_chat_generation_tasks_provider_create_claim_token",
|
||||
"chat_generation_tasks",
|
||||
["provider_create_claim_token"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(
|
||||
"ix_chat_generation_tasks_provider_create_lease_until",
|
||||
"chat_generation_tasks",
|
||||
["provider_create_lease_until"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(
|
||||
"uq_chat_generation_tasks_parent_index",
|
||||
"chat_generation_tasks",
|
||||
["parent_task_id", "generation_index"],
|
||||
unique=True,
|
||||
postgresql_where=sa.text(
|
||||
"parent_task_id IS NOT NULL AND generation_index IS NOT NULL"
|
||||
),
|
||||
)
|
||||
op.create_index(
|
||||
"uq_chat_generation_tasks_user_chat_idempotency",
|
||||
"chat_generation_tasks",
|
||||
["user_id", "idempotency_key"],
|
||||
unique=True,
|
||||
postgresql_where=sa.text(
|
||||
"deleted_at IS NULL "
|
||||
"AND idempotency_key IS NOT NULL "
|
||||
"AND generation_mode IN ('chatapi_async', 'chatapi_main')"
|
||||
),
|
||||
)
|
||||
|
||||
# ImageEngine:管理后台只配置是否允许客户端多份生成和数量上限。
|
||||
op.add_column(
|
||||
"image_engines",
|
||||
sa.Column(
|
||||
"multi_generation_enabled",
|
||||
sa.Boolean(),
|
||||
server_default=sa.text("false"),
|
||||
nullable=False,
|
||||
),
|
||||
)
|
||||
op.add_column(
|
||||
"image_engines",
|
||||
sa.Column("max_generation_count", sa.Integer(), server_default=sa.text("1"), nullable=False),
|
||||
)
|
||||
op.add_column(
|
||||
"image_engines",
|
||||
sa.Column("multi_image_max_images", sa.Integer(), server_default=sa.text("15"), nullable=False),
|
||||
)
|
||||
op.add_column(
|
||||
"image_engines",
|
||||
sa.Column("max_reference_image_count", sa.Integer(), server_default=sa.text("14"), nullable=False),
|
||||
)
|
||||
op.add_column(
|
||||
"image_engines",
|
||||
sa.Column("output_format", sa.String(length=16), server_default=sa.text("''"), nullable=False),
|
||||
)
|
||||
op.create_check_constraint(
|
||||
CK_IMAGE_ENGINE_MAX_GENERATION_COUNT,
|
||||
"image_engines",
|
||||
"max_generation_count BETWEEN 1 AND 5",
|
||||
)
|
||||
op.create_check_constraint(
|
||||
CK_IMAGE_ENGINE_MULTI_IMAGE_MAX,
|
||||
"image_engines",
|
||||
"multi_image_max_images BETWEEN 1 AND 15",
|
||||
)
|
||||
op.create_check_constraint(
|
||||
CK_IMAGE_ENGINE_MAX_REFERENCE,
|
||||
"image_engines",
|
||||
"max_reference_image_count BETWEEN 0 AND 14",
|
||||
)
|
||||
|
||||
# VideoEngine:管理后台只配置是否允许客户端多份生成和数量上限。
|
||||
op.add_column(
|
||||
"video_engines",
|
||||
sa.Column(
|
||||
"multi_generation_enabled",
|
||||
sa.Boolean(),
|
||||
server_default=sa.text("false"),
|
||||
nullable=False,
|
||||
),
|
||||
)
|
||||
op.add_column(
|
||||
"video_engines",
|
||||
sa.Column("max_generation_count", sa.Integer(), server_default=sa.text("1"), nullable=False),
|
||||
)
|
||||
op.create_check_constraint(
|
||||
CK_VIDEO_ENGINE_MAX_GENERATION_COUNT,
|
||||
"video_engines",
|
||||
"max_generation_count BETWEEN 1 AND 5",
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_constraint(CK_VIDEO_ENGINE_MAX_GENERATION_COUNT, "video_engines", type_="check")
|
||||
op.drop_column("video_engines", "max_generation_count")
|
||||
op.drop_column("video_engines", "multi_generation_enabled")
|
||||
|
||||
op.drop_constraint(CK_IMAGE_ENGINE_MAX_REFERENCE, "image_engines", type_="check")
|
||||
op.drop_constraint(CK_IMAGE_ENGINE_MULTI_IMAGE_MAX, "image_engines", type_="check")
|
||||
op.drop_constraint(CK_IMAGE_ENGINE_MAX_GENERATION_COUNT, "image_engines", type_="check")
|
||||
op.drop_column("image_engines", "output_format")
|
||||
op.drop_column("image_engines", "max_reference_image_count")
|
||||
op.drop_column("image_engines", "multi_image_max_images")
|
||||
op.drop_column("image_engines", "max_generation_count")
|
||||
op.drop_column("image_engines", "multi_generation_enabled")
|
||||
|
||||
op.drop_index(
|
||||
"uq_chat_generation_tasks_user_chat_idempotency",
|
||||
table_name="chat_generation_tasks",
|
||||
postgresql_where=sa.text(
|
||||
"deleted_at IS NULL "
|
||||
"AND idempotency_key IS NOT NULL "
|
||||
"AND generation_mode IN ('chatapi_async', 'chatapi_main')"
|
||||
),
|
||||
)
|
||||
op.drop_index(
|
||||
"uq_chat_generation_tasks_parent_index",
|
||||
table_name="chat_generation_tasks",
|
||||
postgresql_where=sa.text(
|
||||
"parent_task_id IS NOT NULL AND generation_index IS NOT NULL"
|
||||
),
|
||||
)
|
||||
op.drop_index(
|
||||
"ix_chat_generation_tasks_provider_create_lease_until",
|
||||
table_name="chat_generation_tasks",
|
||||
)
|
||||
op.drop_index(
|
||||
"ix_chat_generation_tasks_provider_create_claim_token",
|
||||
table_name="chat_generation_tasks",
|
||||
)
|
||||
op.drop_index("idx_chat_generation_tasks_user_mode_created", table_name="chat_generation_tasks")
|
||||
op.drop_index("idx_chat_generation_tasks_parent", table_name="chat_generation_tasks")
|
||||
op.drop_constraint(FK_CHAT_TASK_PARENT, "chat_generation_tasks", type_="foreignkey")
|
||||
op.drop_constraint(CK_CHAT_TASK_GENERATION_INDEX, "chat_generation_tasks", type_="check")
|
||||
op.drop_constraint(CK_CHAT_TASK_GENERATION_COUNT, "chat_generation_tasks", type_="check")
|
||||
op.drop_column("chat_generation_tasks", "provider_create_started_at")
|
||||
op.drop_column("chat_generation_tasks", "provider_create_lease_until")
|
||||
op.drop_column("chat_generation_tasks", "provider_create_claim_token")
|
||||
op.drop_column("chat_generation_tasks", "generation_index")
|
||||
op.drop_column("chat_generation_tasks", "generation_count")
|
||||
op.drop_column("chat_generation_tasks", "parent_task_id")
|
||||
@@ -55,12 +55,12 @@ from app.services.payment import sync_pending_orders, process_refund
|
||||
from app.services.resource_capacity_service import batch_get_user_resource_capacity_usage, get_user_resource_capacity_usage
|
||||
from app.services.team_service import batch_get_team_name_map, set_frontend_user_team
|
||||
|
||||
from app.services.generation_billing_service import (
|
||||
from app.services.generation.billing_service import (
|
||||
OWNER_GENERATION_RECORD,
|
||||
charge_generation_media_by_params,
|
||||
get_next_credit_attempt_no,
|
||||
)
|
||||
from app.services.generation_refund_service import mark_generation_record_failed_and_refund_once
|
||||
from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once
|
||||
from app.utils.id_gen import generate_id
|
||||
from app.schemas.generation import GenerationType, ASPECT_RATIOS, RESOLUTIONS
|
||||
|
||||
|
||||
@@ -41,7 +41,7 @@ from app.services.resource_capacity_service import assert_user_resource_capacity
|
||||
from app.services.upload_resource import delete_unbound_upload_resource, upload_reference_file, cleanup_upload_resource_files_after_commit
|
||||
from app.services.upload_resource.log_service import log_upload_resource_exception, safe_rollback_with_log
|
||||
from app.enums.upload_resource import UploadResourceEventEnum, UploadResourceModuleEnum, UploadResourceTypeEnum
|
||||
from app.services.generation_billing_service import (
|
||||
from app.services.generation.billing_service import (
|
||||
CHARGE_TEXT_PROMPT,
|
||||
OWNER_GENERATION_RECORD,
|
||||
build_credit_biz_key,
|
||||
@@ -49,7 +49,7 @@ from app.services.generation_billing_service import (
|
||||
charge_generation_media_for_record,
|
||||
get_next_credit_attempt_no,
|
||||
)
|
||||
from app.services.generation_refund_service import mark_generation_record_failed_and_refund_once
|
||||
from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once
|
||||
from app.services.media_token_usage_snapshot_service import sync_generation_record_media_token_snapshot
|
||||
from app.services.credit_record_meta_service import build_generation_record_prompt_meta
|
||||
from app.services.video_cover_service import async_create_video_cover_for_local_video
|
||||
|
||||
@@ -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} 个"
|
||||
),
|
||||
)
|
||||
|
||||
@@ -45,5 +45,9 @@ async def list_active_engines(
|
||||
"supported_sizes": sizes,
|
||||
"default_size": e.default_size,
|
||||
"max_image_count": e.max_image_count,
|
||||
"multi_generation_enabled": bool(getattr(e, "multi_generation_enabled", False)),
|
||||
"max_generation_count": int(getattr(e, "max_generation_count", 1) or 1),
|
||||
"multi_image_max_images": int(getattr(e, "multi_image_max_images", 15) or 15),
|
||||
"max_reference_image_count": int(getattr(e, "max_reference_image_count", 14) or 0),
|
||||
})
|
||||
return {"items": items}
|
||||
@@ -52,6 +52,8 @@ async def list_active_engines(
|
||||
"max_image_count": e.max_image_count,
|
||||
"max_video_count": e.max_video_count,
|
||||
"max_audio_count": e.max_audio_count,
|
||||
"multi_generation_enabled": bool(getattr(e, "multi_generation_enabled", False)),
|
||||
"max_generation_count": int(getattr(e, "max_generation_count", 1) or 1),
|
||||
"supports_first_last_frame": e.supports_first_last_frame,
|
||||
"supports_universal_reference": e.supports_universal_reference,
|
||||
})
|
||||
|
||||
@@ -18,3 +18,4 @@ from app.enums.celery_queue import *
|
||||
from app.enums.audio_reference import *
|
||||
|
||||
from app.enums.private_portrait import *
|
||||
from app.enums.generation_provider import *
|
||||
|
||||
@@ -100,3 +100,7 @@ VIDEO_SCHEMA_MAX_SECTION_COUNT = 40
|
||||
VIDEO_SCHEMA_MAX_FIELD_COUNT_PER_SECTION = 80
|
||||
VIDEO_SCHEMA_MAX_TIME_RULE_COUNT = 30
|
||||
VIDEO_SCHEMA_MAX_SEGMENT_COUNT_PER_RULE = 12
|
||||
|
||||
|
||||
MIN_GENERATION_COUNT = 1
|
||||
MAX_GENERATION_COUNT = 5
|
||||
|
||||
@@ -45,17 +45,28 @@ GENERATION_HISTORY_MODULE_SOURCES: tuple[GenerationHistorySourceEnum, ...] = (
|
||||
"""需要回填 module_generation_projects/module_generation_steps 的模块来源集合。"""
|
||||
|
||||
|
||||
GENERATION_HISTORY_SOURCE_TO_TASK_MODE: dict[GenerationHistorySourceEnum, GenerationMode] = {
|
||||
GenerationHistorySourceEnum.CHAT_TASK: GenerationMode.CHATAPI_ASYNC,
|
||||
GenerationHistorySourceEnum.HOT_OPENING_REPLICATE: GenerationMode.HOT_OPENING_REPLICATE,
|
||||
GenerationHistorySourceEnum.SHOT_REPLICATE: GenerationMode.SHOT_REPLICATE,
|
||||
GENERATION_HISTORY_SOURCE_TO_TASK_MODES: dict[GenerationHistorySourceEnum, tuple[GenerationMode, ...]] = {
|
||||
GenerationHistorySourceEnum.CHAT_TASK: (
|
||||
GenerationMode.CHATAPI_ASYNC,
|
||||
GenerationMode.CHATAPI_CHILD,
|
||||
),
|
||||
GenerationHistorySourceEnum.HOT_OPENING_REPLICATE: (GenerationMode.HOT_OPENING_REPLICATE,),
|
||||
GenerationHistorySourceEnum.SHOT_REPLICATE: (GenerationMode.SHOT_REPLICATE,),
|
||||
}
|
||||
"""history_source 到 ChatGenerationTask.generation_mode 的映射。"""
|
||||
"""history_source 到 ChatGenerationTask.generation_mode 集合的映射。"""
|
||||
|
||||
|
||||
GENERATION_HISTORY_SOURCE_TO_TASK_MODE: dict[GenerationHistorySourceEnum, GenerationMode] = {
|
||||
history_source: task_modes[0]
|
||||
for history_source, task_modes in GENERATION_HISTORY_SOURCE_TO_TASK_MODES.items()
|
||||
}
|
||||
"""兼容旧调用的单一模式映射;新查询应使用 GENERATION_HISTORY_SOURCE_TO_TASK_MODES。"""
|
||||
|
||||
|
||||
GENERATION_HISTORY_TASK_MODE_VALUE_TO_SOURCE: dict[str, GenerationHistorySourceEnum] = {
|
||||
task_mode.value: history_source
|
||||
for history_source, task_mode in GENERATION_HISTORY_SOURCE_TO_TASK_MODE.items()
|
||||
for history_source, task_modes in GENERATION_HISTORY_SOURCE_TO_TASK_MODES.items()
|
||||
for task_mode in task_modes
|
||||
}
|
||||
"""ChatGenerationTask.generation_mode 字符串值到 history_source 的映射。"""
|
||||
|
||||
@@ -103,11 +114,17 @@ def get_generation_history_source_label(source: GenerationHistorySourceEnum | st
|
||||
|
||||
|
||||
def get_generation_history_task_mode(source: GenerationHistorySourceEnum) -> GenerationMode | None:
|
||||
"""获取 history_source 对应的 ChatGenerationTask.generation_mode。"""
|
||||
"""兼容旧调用:返回 history_source 对应的第一个任务模式。"""
|
||||
|
||||
return GENERATION_HISTORY_SOURCE_TO_TASK_MODE.get(source)
|
||||
|
||||
|
||||
def get_generation_history_task_modes(source: GenerationHistorySourceEnum) -> tuple[GenerationMode, ...]:
|
||||
"""获取 history_source 对应的全部 ChatGenerationTask.generation_mode。"""
|
||||
|
||||
return GENERATION_HISTORY_SOURCE_TO_TASK_MODES.get(source, ())
|
||||
|
||||
|
||||
def is_generation_history_chat_task_source(source: GenerationHistorySourceEnum) -> bool:
|
||||
"""判断当前来源是否走 chat_generation_tasks 表。"""
|
||||
|
||||
@@ -121,3 +138,7 @@ def is_generation_history_module_source(source: GenerationHistorySourceEnum) ->
|
||||
|
||||
|
||||
MAX_BATCH_DELETE_COUNT = 30
|
||||
|
||||
|
||||
HISTORY_DAY_PAGE_SIZE_MAX = 10
|
||||
HISTORY_GROUP_ITEM_LIMIT = 10
|
||||
|
||||
@@ -0,0 +1,40 @@
|
||||
from enum import StrEnum
|
||||
|
||||
|
||||
class GenerationProviderResultType(StrEnum):
|
||||
IMAGE = "image"
|
||||
VIDEO = "video"
|
||||
|
||||
|
||||
class GenerationProviderTaskPhase(StrEnum):
|
||||
SUBMITTED = "submitted"
|
||||
POLLING = "polling"
|
||||
RESULT_READY = "result_ready"
|
||||
DOWNLOAD_PENDING = "download_pending"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
|
||||
|
||||
class ImageProviderErrorType(StrEnum):
|
||||
TIMEOUT = "timeout"
|
||||
NETWORK = "network"
|
||||
RATE_LIMIT = "rate_limit"
|
||||
AUTH = "auth"
|
||||
INVALID_REQUEST = "invalid_request"
|
||||
CAPABILITY_MISMATCH = "capability_mismatch"
|
||||
CONTENT_REJECTED = "content_rejected"
|
||||
PROVIDER_INTERNAL = "provider_internal"
|
||||
INVALID_RESPONSE = "invalid_response"
|
||||
UNKNOWN = "unknown"
|
||||
|
||||
|
||||
IMAGE_MULTI_OUTPUT_MIN = 1
|
||||
IMAGE_MULTI_OUTPUT_MAX = 15
|
||||
IMAGE_MULTI_REFERENCE_MAX = 14
|
||||
IMAGE_PROVIDER_CLAIM_LEASE_SECONDS = 10 * 60
|
||||
|
||||
MULTI_IMAGE_PROMPT_TEMPLATE = (
|
||||
"请严格生成恰好{count}张内容相关但画面具有明显差异的图片。"
|
||||
"每张图片必须作为独立图片分别输出,不要把多个画面拼接到同一张图片中,"
|
||||
"不要生成九宫格、分镜图、组合图或包含多张子图的单张图片。"
|
||||
)
|
||||
@@ -3,6 +3,8 @@ from enum import Enum
|
||||
|
||||
class GenerationMode(str, Enum):
|
||||
CHATAPI_ASYNC = "chatapi_async"
|
||||
CHATAPI_MAIN = "chatapi_main"
|
||||
CHATAPI_CHILD = "chatapi_child"
|
||||
HOT_OPENING_REPLICATE = "hot_opening_replicate"
|
||||
SHOT_REPLICATE = "shot_replicate"
|
||||
|
||||
@@ -19,6 +21,15 @@ class ChatGenerationTaskStatus(str, Enum):
|
||||
FAILED = "failed"
|
||||
|
||||
|
||||
class ChatGenerationDisplayStatus(str, Enum):
|
||||
PENDING = "pending"
|
||||
GENERATING = "generating"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
DOWNLOAD_FAILED = "download_failed"
|
||||
DELETED = "deleted"
|
||||
|
||||
|
||||
class ChatGenerationPipelineStage(str, Enum):
|
||||
QUEUED = "queued"
|
||||
PREPARING = "preparing"
|
||||
@@ -36,6 +47,31 @@ class ChatGenerationPipelineStage(str, Enum):
|
||||
|
||||
|
||||
class ChatGenerationTaskEventType(str, Enum):
|
||||
TASK_CREATED = "TASK_CREATED"
|
||||
IDEMPOTENCY_HIT = "IDEMPOTENCY_HIT"
|
||||
BATCH_CREATE_START = "BATCH_CREATE_START"
|
||||
BATCH_MAIN_CREATED = "BATCH_MAIN_CREATED"
|
||||
BATCH_CHILDREN_CREATED = "BATCH_CHILDREN_CREATED"
|
||||
BATCH_BILLING_SUCCESS = "BATCH_BILLING_SUCCESS"
|
||||
BATCH_COMMIT_SUCCESS = "BATCH_COMMIT_SUCCESS"
|
||||
CHILD_ENQUEUE_START = "CHILD_ENQUEUE_START"
|
||||
CHILD_ENQUEUE_SUCCESS = "CHILD_ENQUEUE_SUCCESS"
|
||||
CHILD_ENQUEUE_FAILED = "CHILD_ENQUEUE_FAILED"
|
||||
IMAGE_MAIN_CLAIM_ACQUIRED = "IMAGE_MAIN_CLAIM_ACQUIRED"
|
||||
IMAGE_MAIN_CLAIM_REJECTED = "IMAGE_MAIN_CLAIM_REJECTED"
|
||||
IMAGE_MAIN_CLAIM_EXPIRED = "IMAGE_MAIN_CLAIM_EXPIRED"
|
||||
IMAGE_BATCH_PROVIDER_START = "IMAGE_BATCH_PROVIDER_START"
|
||||
IMAGE_BATCH_PROVIDER_SUCCESS = "IMAGE_BATCH_PROVIDER_SUCCESS"
|
||||
IMAGE_BATCH_PROVIDER_FAILED = "IMAGE_BATCH_PROVIDER_FAILED"
|
||||
IMAGE_BATCH_SPLIT_START = "IMAGE_BATCH_SPLIT_START"
|
||||
IMAGE_BATCH_SPLIT_SUCCESS = "IMAGE_BATCH_SPLIT_SUCCESS"
|
||||
IMAGE_BATCH_SPLIT_FAILED = "IMAGE_BATCH_SPLIT_FAILED"
|
||||
MAIN_STATUS_AGGREGATED = "MAIN_STATUS_AGGREGATED"
|
||||
CHILD_RESOURCE_DELETE_START = "CHILD_RESOURCE_DELETE_START"
|
||||
CHILD_RESOURCE_DELETE_SUCCESS = "CHILD_RESOURCE_DELETE_SUCCESS"
|
||||
BATCH_GROUP_DELETE_SUCCESS = "BATCH_GROUP_DELETE_SUCCESS"
|
||||
BATCH_RECOVERY_RECONCILED = "BATCH_RECOVERY_RECONCILED"
|
||||
|
||||
PROMPT_CONCAT_START = "PROMPT_CONCAT_START"
|
||||
PROMPT_CONCAT_SUCCESS = "PROMPT_CONCAT_SUCCESS"
|
||||
|
||||
@@ -87,8 +123,23 @@ class ChatGenerationTaskEventType(str, Enum):
|
||||
TASK_FAILED = "TASK_FAILED"
|
||||
|
||||
|
||||
ALLOWED_GENERATION_MODES = {
|
||||
CHAT_TOP_LEVEL_MODES = {
|
||||
GenerationMode.CHATAPI_ASYNC.value,
|
||||
GenerationMode.CHATAPI_MAIN.value,
|
||||
}
|
||||
|
||||
CHAT_RESOURCE_MODES = {
|
||||
GenerationMode.CHATAPI_ASYNC.value,
|
||||
GenerationMode.CHATAPI_CHILD.value,
|
||||
}
|
||||
|
||||
CHAT_EXECUTABLE_MODES = {
|
||||
GenerationMode.CHATAPI_ASYNC.value,
|
||||
GenerationMode.CHATAPI_CHILD.value,
|
||||
}
|
||||
|
||||
ALLOWED_GENERATION_MODES = {
|
||||
*CHAT_EXECUTABLE_MODES,
|
||||
GenerationMode.HOT_OPENING_REPLICATE.value,
|
||||
GenerationMode.SHOT_REPLICATE.value,
|
||||
}
|
||||
|
||||
@@ -38,17 +38,28 @@ RECENT_GENERATION_CHAT_TASK_MODULES: tuple[RecentGenerationModuleEnum, ...] = (
|
||||
"""来自 chat_generation_tasks 表的模块集合。"""
|
||||
|
||||
|
||||
RECENT_GENERATION_MODULE_TO_TASK_MODE: dict[RecentGenerationModuleEnum, GenerationMode] = {
|
||||
RecentGenerationModuleEnum.CHAT_AI: GenerationMode.CHATAPI_ASYNC,
|
||||
RecentGenerationModuleEnum.HOT_OPENING_REPLICATE: GenerationMode.HOT_OPENING_REPLICATE,
|
||||
RecentGenerationModuleEnum.SHOT_REPLICATE: GenerationMode.SHOT_REPLICATE,
|
||||
RECENT_GENERATION_MODULE_TO_TASK_MODES: dict[RecentGenerationModuleEnum, tuple[GenerationMode, ...]] = {
|
||||
RecentGenerationModuleEnum.CHAT_AI: (
|
||||
GenerationMode.CHATAPI_ASYNC,
|
||||
GenerationMode.CHATAPI_CHILD,
|
||||
),
|
||||
RecentGenerationModuleEnum.HOT_OPENING_REPLICATE: (GenerationMode.HOT_OPENING_REPLICATE,),
|
||||
RecentGenerationModuleEnum.SHOT_REPLICATE: (GenerationMode.SHOT_REPLICATE,),
|
||||
}
|
||||
"""最近生成记录模块枚举到 ChatGenerationTask.generation_mode 的映射。"""
|
||||
"""最近生成记录模块枚举到 ChatGenerationTask.generation_mode 集合的映射。"""
|
||||
|
||||
|
||||
RECENT_GENERATION_MODULE_TO_TASK_MODE: dict[RecentGenerationModuleEnum, GenerationMode] = {
|
||||
module: task_modes[0]
|
||||
for module, task_modes in RECENT_GENERATION_MODULE_TO_TASK_MODES.items()
|
||||
}
|
||||
"""兼容旧调用的单一任务模式映射。"""
|
||||
|
||||
|
||||
RECENT_GENERATION_TASK_MODE_VALUE_TO_MODULE: dict[str, RecentGenerationModuleEnum] = {
|
||||
task_mode.value: module
|
||||
for module, task_mode in RECENT_GENERATION_MODULE_TO_TASK_MODE.items()
|
||||
for module, task_modes in RECENT_GENERATION_MODULE_TO_TASK_MODES.items()
|
||||
for task_mode in task_modes
|
||||
}
|
||||
"""ChatGenerationTask.generation_mode 字符串值到最近生成记录模块枚举的映射。"""
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
from datetime import datetime
|
||||
|
||||
from sqlalchemy import DateTime, Float, ForeignKey, Index, Integer, String, Text, text
|
||||
from sqlalchemy import CheckConstraint, DateTime, Float, ForeignKey, Index, Integer, String, Text, text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.models.base import Base, TimestampMixin, SoftDeleteMixin
|
||||
@@ -26,6 +26,19 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin):
|
||||
unique=True,
|
||||
postgresql_where=text("deleted_at IS NULL AND idempotency_key IS NOT NULL"),
|
||||
),
|
||||
# AI 创作顶层任务在 chatapi_async/chatapi_main 之间切换时,
|
||||
# 同一个前端幂等键也只能创建一组任务。
|
||||
Index(
|
||||
"uq_chat_generation_tasks_user_chat_idempotency",
|
||||
"user_id",
|
||||
"idempotency_key",
|
||||
unique=True,
|
||||
postgresql_where=text(
|
||||
"deleted_at IS NULL "
|
||||
"AND idempotency_key IS NOT NULL "
|
||||
"AND generation_mode IN ('chatapi_async', 'chatapi_main')"
|
||||
),
|
||||
),
|
||||
# 视频 24 小时降频轮询调度使用。
|
||||
Index(
|
||||
"idx_chat_generation_tasks_next_poll_at",
|
||||
@@ -37,6 +50,17 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin):
|
||||
"AND next_poll_at IS NOT NULL"
|
||||
),
|
||||
),
|
||||
Index(
|
||||
"uq_chat_generation_tasks_parent_index",
|
||||
"parent_task_id",
|
||||
"generation_index",
|
||||
unique=True,
|
||||
postgresql_where=text("parent_task_id IS NOT NULL AND generation_index IS NOT NULL"),
|
||||
),
|
||||
Index("idx_chat_generation_tasks_parent", "parent_task_id"),
|
||||
Index("idx_chat_generation_tasks_user_mode_created", "user_id", "generation_mode", "created_at"),
|
||||
CheckConstraint("generation_count BETWEEN 1 AND 5", name="ck_chat_generation_tasks_generation_count"),
|
||||
CheckConstraint("generation_index IS NULL OR generation_index > 0", name="ck_chat_generation_tasks_generation_index"),
|
||||
)
|
||||
|
||||
|
||||
@@ -59,6 +83,17 @@ class ChatGenerationTask(Base, TimestampMixin, SoftDeleteMixin):
|
||||
status: Mapped[str] = mapped_column(String(32), default="generating", index=True)
|
||||
pipeline_stage: Mapped[str | None] = mapped_column(String(32), nullable=True, index=True)
|
||||
generation_mode: Mapped[str] = mapped_column(String(32), default="chatapi_async", index=True)
|
||||
parent_task_id: Mapped[str | None] = mapped_column(
|
||||
String(32), ForeignKey("chat_generation_tasks.id", ondelete="RESTRICT"), nullable=True
|
||||
)
|
||||
generation_count: Mapped[int] = mapped_column(Integer, default=1, server_default="1", nullable=False)
|
||||
generation_index: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
|
||||
# 图片主任务同步调用供应商时的分布式执行租约。
|
||||
# 防止重复 Celery 消息或恢复任务同时触发多次组图请求。
|
||||
provider_create_claim_token: Mapped[str | None] = mapped_column(String(64), nullable=True, index=True)
|
||||
provider_create_lease_until: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True, index=True)
|
||||
provider_create_started_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
media_references: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
provider_task_id: Mapped[str | None] = mapped_column(String(128), nullable=True, index=True)
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from sqlalchemy import Boolean, Integer, String, Text
|
||||
from sqlalchemy import Boolean, CheckConstraint, Integer, String, Text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.models.base import Base, TimestampMixin
|
||||
@@ -6,6 +6,11 @@ from app.models.base import Base, TimestampMixin
|
||||
|
||||
class ImageEngine(Base, TimestampMixin):
|
||||
__tablename__ = "image_engines"
|
||||
__table_args__ = (
|
||||
CheckConstraint("max_generation_count BETWEEN 1 AND 5", name="ck_image_engines_max_generation_count"),
|
||||
CheckConstraint("multi_image_max_images BETWEEN 1 AND 15", name="ck_image_engines_multi_image_max_images"),
|
||||
CheckConstraint("max_reference_image_count BETWEEN 0 AND 14", name="ck_image_engines_max_reference_image_count"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
name: Mapped[str] = mapped_column(String(64), nullable=False)
|
||||
@@ -18,6 +23,26 @@ class ImageEngine(Base, TimestampMixin):
|
||||
supported_sizes: Mapped[str] = mapped_column(Text, default='{}')
|
||||
default_size: Mapped[str] = mapped_column(String(32), default="2K")
|
||||
max_image_count: Mapped[int] = mapped_column(Integer, default=0)
|
||||
|
||||
# 管理后台只配置能力开关与数量上限;本次实际生成数量保存在 ChatGenerationTask.generation_count。
|
||||
multi_generation_enabled: Mapped[bool] = mapped_column(
|
||||
Boolean,
|
||||
default=False,
|
||||
server_default="false",
|
||||
nullable=False,
|
||||
)
|
||||
max_generation_count: Mapped[int] = mapped_column(
|
||||
Integer,
|
||||
default=1,
|
||||
server_default="1",
|
||||
nullable=False,
|
||||
)
|
||||
|
||||
# 火山组图接口能力约束。多份图片始终只调用一次 sequential_image_generation=auto 接口。
|
||||
multi_image_max_images: Mapped[int] = mapped_column(Integer, default=15, server_default="15", nullable=False)
|
||||
max_reference_image_count: Mapped[int] = mapped_column(Integer, default=14, server_default="14", nullable=False)
|
||||
# 留空表示不向供应商传 output_format;用于兼容不支持该参数的模型。
|
||||
output_format: Mapped[str] = mapped_column(String(16), default="", server_default="", nullable=False)
|
||||
generate_url: Mapped[str | None] = mapped_column(String(512), nullable=True, default="")
|
||||
is_active: Mapped[bool] = mapped_column(Boolean, default=True)
|
||||
priority: Mapped[int] = mapped_column(Integer, default=0)
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from sqlalchemy import Boolean, Integer, String
|
||||
from sqlalchemy import Boolean, CheckConstraint, Integer, String
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from app.models.base import Base, TimestampMixin
|
||||
@@ -6,6 +6,9 @@ from app.models.base import Base, TimestampMixin
|
||||
|
||||
class VideoEngine(Base, TimestampMixin):
|
||||
__tablename__ = "video_engines"
|
||||
__table_args__ = (
|
||||
CheckConstraint("max_generation_count BETWEEN 1 AND 5", name="ck_video_engines_max_generation_count"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(32), primary_key=True)
|
||||
name: Mapped[str] = mapped_column(String(64), nullable=False)
|
||||
@@ -20,6 +23,21 @@ class VideoEngine(Base, TimestampMixin):
|
||||
max_image_count: Mapped[int] = mapped_column(Integer, default=2)
|
||||
max_video_count: Mapped[int] = mapped_column(Integer, default=0)
|
||||
max_audio_count: Mapped[int] = mapped_column(Integer, default=0)
|
||||
|
||||
# 管理后台只配置能力开关与数量上限;本次实际生成数量保存在 ChatGenerationTask.generation_count。
|
||||
multi_generation_enabled: Mapped[bool] = mapped_column(
|
||||
Boolean,
|
||||
default=False,
|
||||
server_default="false",
|
||||
nullable=False,
|
||||
)
|
||||
max_generation_count: Mapped[int] = mapped_column(
|
||||
Integer,
|
||||
default=1,
|
||||
server_default="1",
|
||||
nullable=False,
|
||||
)
|
||||
|
||||
supports_first_last_frame: Mapped[bool] = mapped_column(Boolean, default=False)
|
||||
supports_universal_reference: Mapped[bool] = mapped_column(Boolean, default=True)
|
||||
generate_url: Mapped[str | None] = mapped_column(String(512), nullable=True, default="")
|
||||
|
||||
@@ -108,6 +108,7 @@ class GenerationAITaskCreate(BaseModel):
|
||||
}
|
||||
],
|
||||
"idempotency_key": "frontend-submit-uuid-001",
|
||||
"generation_count": 3,
|
||||
"image_size": "2K",
|
||||
"image_proportion": "1:1",
|
||||
"image_px": "2048x2048",
|
||||
@@ -122,6 +123,7 @@ class GenerationAITaskCreate(BaseModel):
|
||||
"engine_id": None,
|
||||
"media_references": None,
|
||||
"idempotency_key": "frontend-submit-uuid-002",
|
||||
"generation_count": 2,
|
||||
"image_size": None,
|
||||
"image_proportion": None,
|
||||
"image_px": None,
|
||||
@@ -169,11 +171,21 @@ class GenerationAITaskCreate(BaseModel):
|
||||
max_length=64,
|
||||
description=(
|
||||
"幂等键,用于防止前端重复提交、网络重试导致重复创建任务和重复扣费。"
|
||||
"同一用户、同一 idempotency_key、同一 generation_mode 下重复请求会返回已有任务。"
|
||||
"同一用户、同一 idempotency_key 的 AI 创作顶层请求会返回已有任务,即使客户端再次传入不同生成数量也不会重复创建。"
|
||||
"建议前端每次点击生成时生成 UUID;同一次请求失败重试时复用同一个 UUID。"
|
||||
),
|
||||
examples=["frontend-submit-uuid-001"],
|
||||
)
|
||||
generation_count: int = Field(
|
||||
1,
|
||||
ge=1,
|
||||
le=5,
|
||||
description=(
|
||||
"客户端本次实际选择的生成数量,默认 1。后端会按当前引擎的多份生成开关、"
|
||||
"最大生成数量以及图片参考图总量限制再次校验。"
|
||||
),
|
||||
examples=[3],
|
||||
)
|
||||
|
||||
# image params
|
||||
image_size: str | None = Field(
|
||||
@@ -227,6 +239,10 @@ class GenerationAIImageEngineOptionOut(BaseModel):
|
||||
default_size: str | None = Field(None, description="默认图片分辨率档位,例如 2K")
|
||||
priority: int = Field(0, description="引擎优先级,数值越大越优先")
|
||||
max_image_count: int = Field(0, description="最大图片数量")
|
||||
multi_generation_enabled: bool = Field(False, description="是否允许客户端选择生成多份图片")
|
||||
max_generation_count: int = Field(1, ge=1, le=5, description="客户端本次最多可选择的图片生成数量")
|
||||
multi_image_max_images: int = Field(15, ge=1, le=15, description="单次组图输入与输出总图片上限")
|
||||
max_reference_image_count: int = Field(14, ge=0, le=14, description="允许的最大参考图片数量")
|
||||
|
||||
|
||||
class GenerationAIVideoEngineOptionOut(BaseModel):
|
||||
@@ -244,6 +260,8 @@ class GenerationAIVideoEngineOptionOut(BaseModel):
|
||||
max_image_count: int | None = Field(None, description="最大图片数量")
|
||||
max_video_count: int | None = Field(None, description="最大视频数量")
|
||||
max_audio_count: int | None = Field(None, description="最大参考音频数量,0 表示不支持音频参考")
|
||||
multi_generation_enabled: bool = Field(False, description="是否允许客户端选择生成多个视频")
|
||||
max_generation_count: int = Field(1, ge=1, le=5, description="客户端本次最多可选择的视频生成数量")
|
||||
supports_first_last_frame: bool = Field(False, description="是否支持首帧和最后一帧")
|
||||
supports_universal_reference: bool = Field(False, description="是否支持通用参考")
|
||||
|
||||
@@ -416,8 +434,12 @@ class GenerationAITaskOut(BaseModel):
|
||||
gen_type: str = Field(..., description="生成类型:image=图片,video=视频")
|
||||
generation_mode: str | None = Field(
|
||||
None,
|
||||
description="生成模式。当前异步Chat生成任务一般为 chatapi_async",
|
||||
description="生成模式:chatapi_async=单份任务,chatapi_main=多份主任务,chatapi_child=多份子任务",
|
||||
)
|
||||
parent_task_id: str | None = Field(None, description="多份生成子任务关联的主任务ID")
|
||||
generation_count: int = Field(1, ge=1, le=5, description="本次实际生成数量快照")
|
||||
generation_index: int | None = Field(None, ge=1, le=5, description="子任务生成序号,从1开始")
|
||||
display_status: str | None = Field(None, description="前端展示状态,例如 download_failed、deleted")
|
||||
pipeline_stage: str | None = Field(
|
||||
None,
|
||||
description=(
|
||||
@@ -468,6 +490,7 @@ class GenerationAITaskOut(BaseModel):
|
||||
error_message: str | None = Field(None, description="错误信息。成功任务一般为 null")
|
||||
created_at: NaiveDatetimeOptional = Field(None, description="任务创建时间")
|
||||
generated_at: NaiveDatetimeOptional = Field(None, description="任务生成完成时间")
|
||||
child_items: list["GenerationAITaskOut"] = Field(default_factory=list, description="多份生成子任务列表,按 generation_index 升序")
|
||||
|
||||
|
||||
class GenerationAITaskListOut(BaseModel):
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import BaseModel, Field, model_validator
|
||||
|
||||
from app.enums.generation_provider import IMAGE_MULTI_OUTPUT_MAX, IMAGE_MULTI_REFERENCE_MAX
|
||||
from app.schemas.common import NaiveDatetime
|
||||
|
||||
|
||||
@@ -13,10 +14,48 @@ class ImageEngineCreate(BaseModel):
|
||||
supported_sizes: str = Field(default='{}')
|
||||
default_size: str = Field(default="2K", max_length=32)
|
||||
max_image_count: int = Field(default=0)
|
||||
|
||||
multi_generation_enabled: bool = Field(
|
||||
default=False,
|
||||
description="是否允许客户端选择生成多份图片;关闭时客户端只能选择 1 份",
|
||||
)
|
||||
max_generation_count: int = Field(
|
||||
default=1,
|
||||
ge=1,
|
||||
le=5,
|
||||
description="客户端单次最多可选择的生成数量,范围 1-5",
|
||||
)
|
||||
multi_image_max_images: int = Field(
|
||||
default=IMAGE_MULTI_OUTPUT_MAX,
|
||||
ge=1,
|
||||
le=IMAGE_MULTI_OUTPUT_MAX,
|
||||
description="火山组图接口输入参考图与输出图片总上限",
|
||||
)
|
||||
max_reference_image_count: int = Field(
|
||||
default=IMAGE_MULTI_REFERENCE_MAX,
|
||||
ge=0,
|
||||
le=IMAGE_MULTI_REFERENCE_MAX,
|
||||
description="图片引擎允许的最大参考图片数量",
|
||||
)
|
||||
output_format: str = Field(
|
||||
default="",
|
||||
max_length=16,
|
||||
description="供应商输出格式;留空表示不传该参数,用于兼容不支持 output_format 的模型",
|
||||
)
|
||||
generate_url: str = Field(default="", max_length=512)
|
||||
is_active: bool = True
|
||||
priority: int = 0
|
||||
|
||||
@model_validator(mode="after")
|
||||
def validate_multi_generation_capability(self):
|
||||
if self.max_generation_count > self.multi_image_max_images:
|
||||
raise ValueError("max_generation_count 不能大于 multi_image_max_images")
|
||||
normalized_output_format = (self.output_format or "").lower().strip()
|
||||
if normalized_output_format not in {"", "png", "jpeg"}:
|
||||
raise ValueError("output_format 仅支持留空、png 或 jpeg")
|
||||
self.output_format = normalized_output_format
|
||||
return self
|
||||
|
||||
|
||||
class ImageEngineOut(ImageEngineCreate):
|
||||
id: str
|
||||
@@ -33,6 +72,10 @@ class ImageEnginePublic(BaseModel):
|
||||
supported_sizes: dict[str, dict[str, str]] = {}
|
||||
default_size: str = "2K"
|
||||
max_image_count: int = 0
|
||||
multi_generation_enabled: bool = False
|
||||
max_generation_count: int = 1
|
||||
multi_image_max_images: int = IMAGE_MULTI_OUTPUT_MAX
|
||||
max_reference_image_count: int = IMAGE_MULTI_REFERENCE_MAX
|
||||
|
||||
|
||||
class ImageEngineListResponse(BaseModel):
|
||||
|
||||
@@ -16,6 +16,16 @@ class VideoEngineCreate(BaseModel):
|
||||
max_image_count: int = Field(default=2)
|
||||
max_video_count: int = Field(default=0)
|
||||
max_audio_count: int = Field(default=0, ge=0, le=3, description="最大参考音频数量,0 表示不支持音频参考")
|
||||
multi_generation_enabled: bool = Field(
|
||||
default=False,
|
||||
description="是否允许客户端选择生成多个视频;关闭时客户端只能选择 1 份",
|
||||
)
|
||||
max_generation_count: int = Field(
|
||||
default=1,
|
||||
ge=1,
|
||||
le=5,
|
||||
description="客户端单次最多可选择的生成数量,范围 1-5",
|
||||
)
|
||||
supports_first_last_frame: bool = Field(default=False, description="是否支持首尾帧模式")
|
||||
supports_universal_reference: bool = Field(default=True, description="是否支持全能参考模式")
|
||||
generate_url: str = Field(default="", max_length=512)
|
||||
@@ -41,6 +51,8 @@ class VideoEnginePublic(BaseModel):
|
||||
max_image_count: int = 2
|
||||
max_video_count: int = 0
|
||||
max_audio_count: int = 0
|
||||
multi_generation_enabled: bool = False
|
||||
max_generation_count: int = 1
|
||||
supports_first_last_frame: bool = False
|
||||
supports_universal_reference: bool = True
|
||||
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""生成任务领域服务。"""
|
||||
@@ -0,0 +1 @@
|
||||
"""AI 创作生成编排服务。"""
|
||||
@@ -0,0 +1,122 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.enums.common import MAX_GENERATION_COUNT, MIN_GENERATION_COUNT
|
||||
from app.enums.generation_provider import IMAGE_MULTI_OUTPUT_MAX, IMAGE_MULTI_REFERENCE_MAX
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.models.video_engine import VideoEngine
|
||||
|
||||
IMAGE_DEFAULT_SIZE = "2K"
|
||||
IMAGE_DEFAULT_PROPORTION = "1:1"
|
||||
IMAGE_DEFAULT_PX = "2048x2048"
|
||||
VIDEO_DEFAULT_DURATION = 4
|
||||
VIDEO_DEFAULT_RATIO = "16:9"
|
||||
VIDEO_DEFAULT_RESOLUTION = "480p"
|
||||
|
||||
|
||||
def normalize_px(value: str | None) -> str | None:
|
||||
if not value:
|
||||
return value
|
||||
return value.replace("×", "x").replace("X", "x").replace("×x", "x").replace("x×", "x")
|
||||
|
||||
|
||||
def parse_json_list(value: str | None, fallback: list):
|
||||
try:
|
||||
parsed = json.loads(value or "")
|
||||
return parsed if isinstance(parsed, list) else fallback
|
||||
except Exception:
|
||||
return fallback
|
||||
|
||||
|
||||
def image_supported_sizes(engine: ImageEngine) -> dict:
|
||||
try:
|
||||
data = json.loads(engine.supported_sizes or "{}")
|
||||
return data if isinstance(data, dict) else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def normalize_generation_count(value: int | None) -> int:
|
||||
try:
|
||||
count = int(value or MIN_GENERATION_COUNT)
|
||||
except (TypeError, ValueError):
|
||||
count = MIN_GENERATION_COUNT
|
||||
return min(MAX_GENERATION_COUNT, max(MIN_GENERATION_COUNT, count))
|
||||
|
||||
|
||||
async def get_image_engine(db: AsyncSession, engine_id: str | None) -> ImageEngine:
|
||||
query = select(ImageEngine).where(ImageEngine.is_active == True)
|
||||
if engine_id:
|
||||
query = query.where(ImageEngine.id == engine_id)
|
||||
else:
|
||||
query = query.order_by(ImageEngine.priority.desc()).limit(1)
|
||||
result = await db.execute(query)
|
||||
engine = result.scalar_one_or_none()
|
||||
if not engine:
|
||||
raise HTTPException(status_code=400, detail="没有可用的图片引擎")
|
||||
return engine
|
||||
|
||||
|
||||
async def get_video_engine(db: AsyncSession, engine_id: str | None) -> VideoEngine:
|
||||
query = select(VideoEngine).where(VideoEngine.is_active == True)
|
||||
if engine_id:
|
||||
query = query.where(VideoEngine.id == engine_id)
|
||||
else:
|
||||
query = query.order_by(VideoEngine.priority.desc())
|
||||
result = await db.execute(query.limit(1))
|
||||
engine = result.scalar_one_or_none()
|
||||
if not engine:
|
||||
raise HTTPException(status_code=400, detail="没有可用的视频引擎")
|
||||
return engine
|
||||
|
||||
|
||||
def build_image_snapshot(engine: ImageEngine, size: str, proportion: str, px: str) -> dict:
|
||||
return {
|
||||
"engine_type": "image",
|
||||
"id": engine.id,
|
||||
"name": engine.name,
|
||||
"provider": engine.provider,
|
||||
"api_base": engine.api_base,
|
||||
"api_key_masked": "****" if engine.api_key else "",
|
||||
"model_name": engine.model_name,
|
||||
"generate_url": engine.generate_url,
|
||||
"supported_models": parse_json_list(engine.supported_models, []),
|
||||
"default_size": engine.default_size,
|
||||
"multi_generation_enabled": bool(getattr(engine, "multi_generation_enabled", False)),
|
||||
"max_generation_count": normalize_generation_count(getattr(engine, "max_generation_count", 1)),
|
||||
"multi_image_max_images": int(getattr(engine, "multi_image_max_images", IMAGE_MULTI_OUTPUT_MAX) or IMAGE_MULTI_OUTPUT_MAX),
|
||||
"max_reference_image_count": int(getattr(engine, "max_reference_image_count", IMAGE_MULTI_REFERENCE_MAX) or 0),
|
||||
"output_format": (getattr(engine, "output_format", "") or "").lower().strip(),
|
||||
"selected_size": size,
|
||||
"selected_proportion": proportion,
|
||||
"selected_px": px,
|
||||
}
|
||||
|
||||
|
||||
def build_video_snapshot(engine: VideoEngine, ratio: str, resolution: str, duration: int) -> dict:
|
||||
return {
|
||||
"engine_type": "video",
|
||||
"id": engine.id,
|
||||
"name": engine.name,
|
||||
"provider": engine.provider,
|
||||
"api_base": engine.api_base,
|
||||
"api_key_masked": "****" if engine.api_key else "",
|
||||
"model_name": engine.model_name,
|
||||
"generate_url": engine.generate_url,
|
||||
"query_url": engine.query_url,
|
||||
"supported_ratios": parse_json_list(engine.supported_ratios, []),
|
||||
"supported_resolutions": parse_json_list(engine.supported_resolutions, []),
|
||||
"supported_durations": parse_json_list(engine.supported_durations, []),
|
||||
"max_duration": engine.max_duration,
|
||||
"max_audio_count": engine.max_audio_count,
|
||||
"multi_generation_enabled": bool(getattr(engine, "multi_generation_enabled", False)),
|
||||
"max_generation_count": normalize_generation_count(getattr(engine, "max_generation_count", 1)),
|
||||
"selected_ratio": ratio,
|
||||
"selected_resolution": resolution,
|
||||
"selected_duration": duration,
|
||||
}
|
||||
@@ -0,0 +1,532 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from types import SimpleNamespace
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.enums.generation_provider import IMAGE_PROVIDER_CLAIM_LEASE_SECONDS
|
||||
from app.enums.generation_task import (
|
||||
ChatGenerationPipelineStage,
|
||||
ChatGenerationTaskEventType,
|
||||
ChatGenerationTaskStatus,
|
||||
GenerationMode,
|
||||
GenerationType,
|
||||
)
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.services.generation.ai.task_group_service import aggregate_main_task_status, load_children_map
|
||||
from app.services.generation.log_service import log_task_event
|
||||
from app.services.generation.provider_service import (
|
||||
create_image_sync_batch_result_with_engine,
|
||||
get_runtime_engine,
|
||||
)
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.image_gen import ImageProviderError
|
||||
from app.services.operation_log_service import build_exception_detail, log_operation_event
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class ImageBatchClaim:
|
||||
acquired: bool
|
||||
main_task_id: str
|
||||
claim_token: str | None = None
|
||||
task_snapshot: SimpleNamespace | None = None
|
||||
runtime_engine: SimpleNamespace | None = None
|
||||
existing_child_ids: list[str] | None = None
|
||||
reason: str | None = None
|
||||
|
||||
|
||||
def _now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
def _json(value) -> str | None:
|
||||
if value is None:
|
||||
return None
|
||||
return json.dumps(value, ensure_ascii=False, default=str)
|
||||
|
||||
|
||||
def _aware(value: datetime | None) -> datetime | None:
|
||||
if value is None:
|
||||
return None
|
||||
if value.tzinfo is None:
|
||||
return value.replace(tzinfo=timezone.utc)
|
||||
return value.astimezone(timezone.utc)
|
||||
|
||||
|
||||
def _lease_alive(task: ChatGenerationTask, now: datetime | None = None) -> bool:
|
||||
lease_until = _aware(task.provider_create_lease_until)
|
||||
return bool(task.provider_create_claim_token and lease_until and lease_until > (now or _now()))
|
||||
|
||||
|
||||
def _task_snapshot(main: ChatGenerationTask) -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
id=str(main.id),
|
||||
user_id=str(main.user_id),
|
||||
generation_mode=str(main.generation_mode),
|
||||
generation_count=int(main.generation_count or 1),
|
||||
original_prompt=main.original_prompt,
|
||||
optimized_prompt=main.optimized_prompt,
|
||||
media_references=main.media_references,
|
||||
gen_type=main.gen_type,
|
||||
duration=main.duration,
|
||||
aspect_ratio=main.aspect_ratio,
|
||||
resolution=main.resolution,
|
||||
image_size=main.image_size,
|
||||
image_proportion=main.image_proportion,
|
||||
image_px=main.image_px,
|
||||
engine_id=main.engine_id,
|
||||
)
|
||||
|
||||
|
||||
async def _claim_image_main_batch(db: AsyncSession, main_task_id: str) -> ImageBatchClaim:
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == main_task_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
ChatGenerationTask.gen_type == GenerationType.IMAGE.value,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
main = result.scalar_one_or_none()
|
||||
if not main:
|
||||
await db.rollback()
|
||||
return ImageBatchClaim(False, main_task_id, reason="main_missing")
|
||||
|
||||
children_map = await load_children_map(db, [main.id], include_deleted=True)
|
||||
existing_children = children_map.get(main.id, [])
|
||||
if existing_children:
|
||||
child_ids = [str(child.id) for child in existing_children if child.deleted_at is None]
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await db.commit()
|
||||
return ImageBatchClaim(False, main_task_id, existing_child_ids=child_ids, reason="already_split")
|
||||
|
||||
if main.status != ChatGenerationTaskStatus.GENERATING.value:
|
||||
status = str(main.status)
|
||||
await db.rollback()
|
||||
return ImageBatchClaim(False, main_task_id, reason=f"status_{status}")
|
||||
|
||||
now = _now()
|
||||
if _lease_alive(main, now):
|
||||
user_id = str(main.user_id)
|
||||
group_id = str(main.id)
|
||||
lease_until = main.provider_create_lease_until
|
||||
await db.rollback()
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="IMAGE_MAIN_CLAIM_REJECTED",
|
||||
event_status="skipped",
|
||||
source="celery",
|
||||
user_id=user_id,
|
||||
group_id=group_id,
|
||||
task_id=group_id,
|
||||
detail={"reason": "lease_alive", "lease_until": lease_until},
|
||||
)
|
||||
return ImageBatchClaim(False, main_task_id, reason="lease_alive")
|
||||
|
||||
deadline = _aware(main.deadline_at)
|
||||
if deadline and deadline <= now:
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=main,
|
||||
error_message="图片批量生成任务超时",
|
||||
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
)
|
||||
await db.commit()
|
||||
return ImageBatchClaim(False, main_task_id, reason="deadline_expired")
|
||||
|
||||
claim_token = uuid4().hex
|
||||
main.provider_create_claim_token = claim_token
|
||||
main.provider_create_started_at = now
|
||||
main.provider_create_lease_until = now + timedelta(seconds=IMAGE_PROVIDER_CLAIM_LEASE_SECONDS)
|
||||
main.pipeline_stage = ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value
|
||||
runtime_engine = await get_runtime_engine(db, main)
|
||||
snapshot = _task_snapshot(main)
|
||||
user_id = str(main.user_id)
|
||||
generation_count = int(main.generation_count or 1)
|
||||
lease_until = main.provider_create_lease_until
|
||||
await db.commit()
|
||||
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="IMAGE_MAIN_CLAIM_ACQUIRED",
|
||||
event_status="success",
|
||||
source="celery",
|
||||
user_id=user_id,
|
||||
group_id=main_task_id,
|
||||
task_id=main_task_id,
|
||||
detail={
|
||||
"generation_count": generation_count,
|
||||
"claim_token_suffix": claim_token[-8:],
|
||||
"lease_until": lease_until,
|
||||
},
|
||||
)
|
||||
return ImageBatchClaim(
|
||||
True,
|
||||
main_task_id,
|
||||
claim_token=claim_token,
|
||||
task_snapshot=snapshot,
|
||||
runtime_engine=runtime_engine,
|
||||
)
|
||||
|
||||
|
||||
def _validate_provider_batch(provider_result: dict, generation_count: int) -> list[dict]:
|
||||
items = provider_result.get("items") or []
|
||||
if not isinstance(items, list):
|
||||
raise RuntimeError("图片供应商返回 items 结构异常")
|
||||
|
||||
success_items: list[dict] = []
|
||||
errors: list[str] = []
|
||||
for position, item in enumerate(items, start=1):
|
||||
if not isinstance(item, dict):
|
||||
errors.append(f"第{position}项返回结构无效")
|
||||
continue
|
||||
if item.get("error_message") or item.get("error_code"):
|
||||
errors.append(
|
||||
f"第{position}项: {item.get('error_message') or item.get('error_code') or '生成失败'}"
|
||||
)
|
||||
continue
|
||||
remote_url = str(item.get("remote_result_url") or "").strip()
|
||||
if not remote_url:
|
||||
errors.append(f"第{position}项: 供应商未返回图片地址")
|
||||
continue
|
||||
normalized = dict(item)
|
||||
normalized["generation_index"] = position
|
||||
success_items.append(normalized)
|
||||
|
||||
generated_images = int(provider_result.get("generated_images") or 0)
|
||||
if generated_images and generated_images != len(success_items):
|
||||
errors.append(
|
||||
f"usage.generated_images={generated_images} 与有效图片数 {len(success_items)} 不一致"
|
||||
)
|
||||
if len(items) != generation_count:
|
||||
errors.append(f"返回条目数应为 {generation_count},实际 {len(items)}")
|
||||
if len(success_items) != generation_count:
|
||||
errors.append(f"成功图片数应为 {generation_count},实际 {len(success_items)}")
|
||||
if errors:
|
||||
raise RuntimeError("图片组图未全部成功;" + ";".join(errors))
|
||||
return success_items
|
||||
|
||||
|
||||
async def _fail_claimed_main(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
main_task_id: str,
|
||||
claim_token: str,
|
||||
error_message: str,
|
||||
event_type: ChatGenerationTaskEventType,
|
||||
exception: Exception | None = None,
|
||||
) -> bool:
|
||||
try:
|
||||
await db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == main_task_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
main = result.scalar_one_or_none()
|
||||
if not main or main.provider_create_claim_token != claim_token:
|
||||
await db.rollback()
|
||||
return False
|
||||
|
||||
existing_map = await load_children_map(db, [main.id], include_deleted=True)
|
||||
if existing_map.get(main.id):
|
||||
# child 已经落库后不再允许图片生成退款。
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await db.commit()
|
||||
return False
|
||||
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=main,
|
||||
error_message=error_message,
|
||||
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
)
|
||||
task_id = str(main.id)
|
||||
user_id = str(main.user_id)
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
task_id=task_id,
|
||||
event_type=event_type.value,
|
||||
to_status=ChatGenerationTaskStatus.FAILED.value,
|
||||
to_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
message=error_message,
|
||||
)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type=event_type.value,
|
||||
event_status="failed",
|
||||
source="celery",
|
||||
user_id=user_id,
|
||||
group_id=task_id,
|
||||
task_id=task_id,
|
||||
message=error_message,
|
||||
detail=build_exception_detail(exception) if exception else {"message": error_message},
|
||||
error=error_message,
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
async def _split_children(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
main_task_id: str,
|
||||
claim_token: str,
|
||||
provider_result: dict,
|
||||
provider_items: list[dict],
|
||||
) -> list[str]:
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == main_task_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
ChatGenerationTask.gen_type == GenerationType.IMAGE.value,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
main = result.scalar_one_or_none()
|
||||
if not main:
|
||||
raise RuntimeError("图片主任务不存在或已删除")
|
||||
if main.provider_create_claim_token != claim_token:
|
||||
raise RuntimeError("图片主任务执行租约已失效,拒绝拆分子任务")
|
||||
if main.status != ChatGenerationTaskStatus.GENERATING.value:
|
||||
raise RuntimeError(f"图片主任务当前状态不允许拆分: {main.status}")
|
||||
|
||||
existing_map = await load_children_map(db, [main.id], include_deleted=True)
|
||||
existing = existing_map.get(main.id, [])
|
||||
if existing:
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await db.commit()
|
||||
return [str(child.id) for child in existing if child.deleted_at is None]
|
||||
|
||||
expected_count = max(1, int(main.generation_count or 1))
|
||||
if len(provider_items) != expected_count:
|
||||
raise RuntimeError(f"图片批量拆分数量不一致,期望 {expected_count},实际 {len(provider_items)}")
|
||||
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_SPLIT_START.value,
|
||||
event_status="started",
|
||||
source="celery",
|
||||
user_id=main.user_id,
|
||||
group_id=main.id,
|
||||
task_id=main.id,
|
||||
detail={"generation_count": expected_count},
|
||||
)
|
||||
|
||||
children: list[ChatGenerationTask] = []
|
||||
for item in provider_items:
|
||||
index = int(item.get("generation_index") or 0)
|
||||
if index < 1 or index > expected_count:
|
||||
raise RuntimeError(f"无效的图片生成序号: {index}")
|
||||
child = ChatGenerationTask(
|
||||
id=generate_id(),
|
||||
user_id=main.user_id,
|
||||
original_prompt=main.original_prompt,
|
||||
optimized_prompt=main.optimized_prompt,
|
||||
gen_type=main.gen_type,
|
||||
image_size=main.image_size,
|
||||
image_proportion=main.image_proportion,
|
||||
image_px=main.image_px,
|
||||
status=ChatGenerationTaskStatus.GENERATING.value,
|
||||
pipeline_stage=ChatGenerationPipelineStage.RESULT_READY.value,
|
||||
generation_mode=GenerationMode.CHATAPI_CHILD.value,
|
||||
parent_task_id=main.id,
|
||||
generation_count=expected_count,
|
||||
generation_index=index,
|
||||
media_references=main.media_references,
|
||||
remote_result_url=item.get("remote_result_url"),
|
||||
engine_id=main.engine_id,
|
||||
engine_snapshot_json=main.engine_snapshot_json,
|
||||
provider_response_json=_json(item.get("response_data") or {}),
|
||||
# 图片生成计费和 token 都归属于 main;child 只负责下载和资源展示。
|
||||
credits_cost=0,
|
||||
image_tokens_used=0,
|
||||
deadline_at=main.deadline_at,
|
||||
)
|
||||
children.append(child)
|
||||
|
||||
children.sort(key=lambda child: int(child.generation_index or 0))
|
||||
db.add_all(children)
|
||||
main.provider_response_json = _json(provider_result.get("response_data") or provider_result)
|
||||
main.image_tokens_used = int(provider_result.get("image_tokens") or 0)
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await db.flush()
|
||||
child_ids = [str(child.id) for child in children]
|
||||
main_id = str(main.id)
|
||||
main_user_id = str(main.user_id)
|
||||
await aggregate_main_task_status(db, parent_task_id=main_id)
|
||||
await db.commit()
|
||||
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_SPLIT_SUCCESS.value,
|
||||
event_status="success",
|
||||
source="celery",
|
||||
user_id=main_user_id,
|
||||
group_id=main_id,
|
||||
task_id=main_id,
|
||||
detail={"child_task_ids": child_ids},
|
||||
)
|
||||
return child_ids
|
||||
|
||||
|
||||
async def _enqueue_child_downloads(db: AsyncSession, child_ids: list[str]) -> dict[str, list[str]]:
|
||||
if not child_ids:
|
||||
return {"enqueued": [], "failed": []}
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id.in_(child_ids),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.order_by(ChatGenerationTask.generation_index.asc())
|
||||
)
|
||||
children = list(result.scalars().all())
|
||||
from app.tasks.generation_download_tasks import enqueue_download_task
|
||||
|
||||
enqueued: list[str] = []
|
||||
failed: list[str] = []
|
||||
for child in children:
|
||||
if child.status == ChatGenerationTaskStatus.COMPLETED.value:
|
||||
continue
|
||||
if child.pipeline_stage in {
|
||||
ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value,
|
||||
ChatGenerationPipelineStage.DOWNLOADING.value,
|
||||
ChatGenerationPipelineStage.RETRY_WAITING.value,
|
||||
}:
|
||||
continue
|
||||
celery_task_id = await enqueue_download_task(db, child, reason="image_batch_split")
|
||||
if celery_task_id:
|
||||
enqueued.append(str(child.id))
|
||||
else:
|
||||
failed.append(str(child.id))
|
||||
|
||||
if children and children[0].parent_task_id:
|
||||
await aggregate_main_task_status(db, parent_task_id=str(children[0].parent_task_id))
|
||||
await db.commit()
|
||||
return {"enqueued": enqueued, "failed": failed}
|
||||
|
||||
|
||||
async def run_image_main_batch(db: AsyncSession, main_task: ChatGenerationTask) -> list[str]:
|
||||
"""单次同步组图,全部成功后原子拆分 child。
|
||||
|
||||
绝不在组图 API 失败后退化为 N 次单图请求。
|
||||
"""
|
||||
main_task_id = str(main_task.id)
|
||||
claim = await _claim_image_main_batch(db, main_task_id)
|
||||
if claim.existing_child_ids is not None:
|
||||
await _enqueue_child_downloads(db, claim.existing_child_ids)
|
||||
return claim.existing_child_ids
|
||||
if not claim.acquired or not claim.claim_token or not claim.task_snapshot or not claim.runtime_engine:
|
||||
return []
|
||||
|
||||
generation_count = max(1, int(claim.task_snapshot.generation_count or 1))
|
||||
try:
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_PROVIDER_START.value,
|
||||
event_status="started",
|
||||
source="celery",
|
||||
user_id=claim.task_snapshot.user_id,
|
||||
group_id=main_task_id,
|
||||
task_id=main_task_id,
|
||||
detail={"generation_count": generation_count},
|
||||
)
|
||||
provider_result = await create_image_sync_batch_result_with_engine(
|
||||
claim.task_snapshot,
|
||||
claim.runtime_engine,
|
||||
generation_count=generation_count,
|
||||
)
|
||||
provider_items = _validate_provider_batch(provider_result, generation_count)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_PROVIDER_SUCCESS.value,
|
||||
event_status="success",
|
||||
source="celery",
|
||||
user_id=claim.task_snapshot.user_id,
|
||||
group_id=main_task_id,
|
||||
task_id=main_task_id,
|
||||
detail={
|
||||
"generation_count": generation_count,
|
||||
"result_count": len(provider_items),
|
||||
"image_tokens": int(provider_result.get("image_tokens") or 0),
|
||||
"single_provider_request": True,
|
||||
"fallback_to_single_requests": False,
|
||||
},
|
||||
)
|
||||
except Exception as exc:
|
||||
message = exc.safe_message if isinstance(exc, ImageProviderError) else str(exc)
|
||||
await _fail_claimed_main(
|
||||
db,
|
||||
main_task_id=main_task_id,
|
||||
claim_token=claim.claim_token,
|
||||
error_message=message or "图片批量生成失败",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_PROVIDER_FAILED,
|
||||
exception=exc,
|
||||
)
|
||||
return []
|
||||
|
||||
try:
|
||||
child_ids = await _split_children(
|
||||
db,
|
||||
main_task_id=main_task_id,
|
||||
claim_token=claim.claim_token,
|
||||
provider_result=provider_result,
|
||||
provider_items=provider_items,
|
||||
)
|
||||
except Exception as exc:
|
||||
await _fail_claimed_main(
|
||||
db,
|
||||
main_task_id=main_task_id,
|
||||
claim_token=claim.claim_token,
|
||||
error_message=f"图片批量结果拆分失败: {exc}",
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_BATCH_SPLIT_FAILED,
|
||||
exception=exc,
|
||||
)
|
||||
return []
|
||||
|
||||
# child 已提交后,下载投递失败不属于图片生成失败,不退款、不重新请求供应商。
|
||||
enqueue_result = await _enqueue_child_downloads(db, child_ids)
|
||||
if enqueue_result["failed"]:
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="DOWNLOAD_ENQUEUE_FAILED",
|
||||
event_status="failed",
|
||||
source="celery",
|
||||
user_id=claim.task_snapshot.user_id,
|
||||
group_id=main_task_id,
|
||||
task_id=main_task_id,
|
||||
detail={
|
||||
"failed_child_task_ids": enqueue_result["failed"],
|
||||
"enqueued_child_task_ids": enqueue_result["enqueued"],
|
||||
"provider_regenerated": False,
|
||||
"generation_refunded": False,
|
||||
},
|
||||
)
|
||||
return child_ids
|
||||
+140
-358
@@ -1,77 +1,54 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime, timedelta, timezone, date
|
||||
from datetime import datetime, date
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import and_, func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.config import settings
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.models.generation_record import GenerationRecord
|
||||
from app.models.project import Project
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.models.user import User
|
||||
from app.models.video_engine import VideoEngine
|
||||
from app.enums.audio_reference import (
|
||||
AUDIO_ALLOWED_EXTENSIONS,
|
||||
AUDIO_MAX_COUNT_LIMIT,
|
||||
AUDIO_MAX_DURATION_SECONDS,
|
||||
AUDIO_MAX_TOTAL_DURATION_SECONDS,
|
||||
AUDIO_MIN_DURATION_SECONDS,
|
||||
)
|
||||
from app.enums.generation_task import CHAT_TOP_LEVEL_MODES, GenerationMode
|
||||
from app.enums.generation_history import (
|
||||
GenerationHistorySourceEnum,
|
||||
get_generation_history_source_label,
|
||||
get_generation_history_task_mode,
|
||||
get_generation_history_task_modes,
|
||||
normalize_generation_history_source,
|
||||
HISTORY_DAY_PAGE_SIZE_MAX,
|
||||
HISTORY_GROUP_ITEM_LIMIT,
|
||||
)
|
||||
from app.schemas.generation_ai import (
|
||||
GenerationAIEngineGroupOut,
|
||||
GenerationAIEngineOptionsOut,
|
||||
GenerationAIImageEngineOptionOut,
|
||||
GenerationAIRecordHistoryItemOut,
|
||||
GenerationAITaskCreate,
|
||||
GenerationAITaskOut,
|
||||
GenerationAIVideoEngineOptionOut,
|
||||
)
|
||||
from app.services.generation_billing_service import (
|
||||
OWNER_CHAT_GENERATION_TASK,
|
||||
charge_generation_media_by_params,
|
||||
)
|
||||
from app.services.resource_accounting_service import (
|
||||
SOURCE_MODEL_CHAT_TASK,
|
||||
SOURCE_MODEL_GENERATION_RECORD,
|
||||
batch_get_generated_resource_info_map,
|
||||
soft_delete_chat_task_resources,
|
||||
)
|
||||
from app.services.resource_signed_url_service import build_resource_signed_url
|
||||
from app.services.generation_history_meta_service import (
|
||||
from app.services.generation.history_meta_service import (
|
||||
GenerationHistoryMeta,
|
||||
batch_load_generation_history_meta_map,
|
||||
build_empty_history_meta,
|
||||
)
|
||||
from app.services.resource_capacity_service import assert_user_resource_capacity_available
|
||||
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls, resolve_private_portrait_reference_display_urls, resolve_private_portrait_references
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
IMAGE_DEFAULT_SIZE = "2K"
|
||||
IMAGE_DEFAULT_PROPORTION = "1:1"
|
||||
IMAGE_DEFAULT_PX = "2048x2048"
|
||||
VIDEO_DEFAULT_DURATION = 4
|
||||
VIDEO_DEFAULT_RATIO = "16:9"
|
||||
VIDEO_DEFAULT_RESOLUTION = "480p"
|
||||
|
||||
HISTORY_DAY_PAGE_SIZE_MAX = 10
|
||||
HISTORY_GROUP_ITEM_LIMIT = 10
|
||||
|
||||
def normalize_px(value: str | None) -> str | None:
|
||||
if not value:
|
||||
return value
|
||||
return value.replace("×", "x").replace("X", "x").replace("×x", "x").replace("x×", "x")
|
||||
|
||||
from app.services.generation.ai.task_group_service import get_display_status, load_children_map
|
||||
from app.services.generation.ai.engine_service import (
|
||||
image_supported_sizes,
|
||||
normalize_generation_count,
|
||||
parse_json_list,
|
||||
)
|
||||
from app.services.private_portrait.reference_resolver import batch_resolve_private_portrait_reference_display_urls
|
||||
|
||||
def _json(data: Any) -> str | None:
|
||||
if data is None:
|
||||
@@ -88,7 +65,12 @@ def _parse_json(text: str | None):
|
||||
return None
|
||||
|
||||
|
||||
async def _resolve_task_reference_display_map(db: AsyncSession, tasks: list[ChatGenerationTask], *, user_id: str | None = None) -> dict[str, list[dict] | None]:
|
||||
async def _resolve_task_reference_display_map(
|
||||
db: AsyncSession,
|
||||
tasks: list[ChatGenerationTask],
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
) -> dict[str, list[dict] | None]:
|
||||
return await batch_resolve_private_portrait_reference_display_urls(
|
||||
db,
|
||||
{task.id: _parse_json(task.media_references) for task in tasks},
|
||||
@@ -96,7 +78,12 @@ async def _resolve_task_reference_display_map(db: AsyncSession, tasks: list[Chat
|
||||
)
|
||||
|
||||
|
||||
async def _resolve_generation_record_reference_display_map(db: AsyncSession, records: list[GenerationRecord], *, user_id: str | None = None) -> dict[str, list[dict] | None]:
|
||||
async def _resolve_generation_record_reference_display_map(
|
||||
db: AsyncSession,
|
||||
records: list[GenerationRecord],
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
) -> dict[str, list[dict] | None]:
|
||||
return await batch_resolve_private_portrait_reference_display_urls(
|
||||
db,
|
||||
{record.id: _parse_json(record.media_references) for record in records},
|
||||
@@ -104,90 +91,6 @@ async def _resolve_generation_record_reference_display_map(db: AsyncSession, rec
|
||||
)
|
||||
|
||||
|
||||
async def _get_image_engine(db: AsyncSession, engine_id: str | None) -> ImageEngine:
|
||||
query = select(ImageEngine).where(ImageEngine.is_active == True)
|
||||
if engine_id:
|
||||
query = query.where(ImageEngine.id == engine_id)
|
||||
else:
|
||||
query = query.order_by(ImageEngine.priority.desc()).limit(1)
|
||||
result = await db.execute(query)
|
||||
engine = result.scalar_one_or_none()
|
||||
if not engine:
|
||||
raise HTTPException(status_code=400, detail="没有可用的图片引擎")
|
||||
return engine
|
||||
|
||||
|
||||
async def _get_video_engine(db: AsyncSession, engine_id: str | None) -> VideoEngine:
|
||||
query = select(VideoEngine).where(VideoEngine.is_active == True)
|
||||
if engine_id:
|
||||
query = query.where(VideoEngine.id == engine_id)
|
||||
else:
|
||||
query = query.order_by(VideoEngine.priority.desc())
|
||||
|
||||
query = query.limit(1)
|
||||
result = await db.execute(query)
|
||||
engine = result.scalar_one_or_none()
|
||||
if not engine:
|
||||
raise HTTPException(status_code=400, detail="没有可用的视频引擎")
|
||||
return engine
|
||||
|
||||
|
||||
def _image_supported_sizes(engine: ImageEngine) -> dict:
|
||||
try:
|
||||
data = json.loads(engine.supported_sizes or "{}")
|
||||
return data if isinstance(data, dict) else {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def _parse_list(value: str | None, fallback: list):
|
||||
try:
|
||||
parsed = json.loads(value or "")
|
||||
return parsed if isinstance(parsed, list) else fallback
|
||||
except Exception:
|
||||
return fallback
|
||||
|
||||
|
||||
def _build_image_snapshot(engine: ImageEngine, size: str, proportion: str, px: str) -> dict:
|
||||
return {
|
||||
"engine_type": "image",
|
||||
"id": engine.id,
|
||||
"name": engine.name,
|
||||
"provider": engine.provider,
|
||||
"api_base": engine.api_base,
|
||||
"api_key_masked": "****" if engine.api_key else "",
|
||||
"model_name": engine.model_name,
|
||||
"generate_url": engine.generate_url,
|
||||
"supported_models": _parse_list(engine.supported_models, []),
|
||||
"default_size": engine.default_size,
|
||||
"selected_size": size,
|
||||
"selected_proportion": proportion,
|
||||
"selected_px": px,
|
||||
}
|
||||
|
||||
|
||||
def _build_video_snapshot(engine: VideoEngine, ratio: str, resolution: str, duration: int) -> dict:
|
||||
return {
|
||||
"engine_type": "video",
|
||||
"id": engine.id,
|
||||
"name": engine.name,
|
||||
"provider": engine.provider,
|
||||
"api_base": engine.api_base,
|
||||
"api_key_masked": "****" if engine.api_key else "",
|
||||
"model_name": engine.model_name,
|
||||
"generate_url": engine.generate_url,
|
||||
"query_url": engine.query_url,
|
||||
"supported_ratios": _parse_list(engine.supported_ratios, []),
|
||||
"supported_resolutions": _parse_list(engine.supported_resolutions, []),
|
||||
"supported_durations": _parse_list(engine.supported_durations, []),
|
||||
"max_duration": engine.max_duration,
|
||||
"max_audio_count": engine.max_audio_count,
|
||||
"selected_ratio": ratio,
|
||||
"selected_resolution": resolution,
|
||||
"selected_duration": duration,
|
||||
}
|
||||
|
||||
|
||||
async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEngineOptionsOut:
|
||||
"""获取当前启用的图片/视频生成引擎,供前端创建任务时选择 engine_id。"""
|
||||
image_result = await db.execute(
|
||||
@@ -207,11 +110,15 @@ async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEng
|
||||
name=engine.name,
|
||||
provider=engine.provider,
|
||||
model_name=engine.model_name,
|
||||
supported_models=_parse_list(engine.supported_models, []),
|
||||
supported_sizes=_image_supported_sizes(engine),
|
||||
supported_models=parse_json_list(engine.supported_models, []),
|
||||
supported_sizes=image_supported_sizes(engine),
|
||||
default_size=engine.default_size,
|
||||
priority=engine.priority or 0,
|
||||
max_image_count=engine.max_image_count,
|
||||
multi_generation_enabled=bool(getattr(engine, "multi_generation_enabled", False)),
|
||||
max_generation_count=normalize_generation_count(getattr(engine, "max_generation_count", 1)),
|
||||
multi_image_max_images=int(getattr(engine, "multi_image_max_images", 15) or 15),
|
||||
max_reference_image_count=int(getattr(engine, "max_reference_image_count", 14) or 0),
|
||||
)
|
||||
for engine in image_result.scalars().all()
|
||||
]
|
||||
@@ -221,9 +128,9 @@ async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEng
|
||||
name=engine.name,
|
||||
provider=engine.provider,
|
||||
model_name=engine.model_name,
|
||||
supported_ratios=_parse_list(engine.supported_ratios, []),
|
||||
supported_resolutions=_parse_list(engine.supported_resolutions, []),
|
||||
supported_durations=_parse_list(engine.supported_durations, []),
|
||||
supported_ratios=parse_json_list(engine.supported_ratios, []),
|
||||
supported_resolutions=parse_json_list(engine.supported_resolutions, []),
|
||||
supported_durations=parse_json_list(engine.supported_durations, []),
|
||||
max_duration=engine.max_duration,
|
||||
priority=engine.priority or 0,
|
||||
max_image_count=engine.max_image_count,
|
||||
@@ -231,6 +138,8 @@ async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEng
|
||||
max_audio_count=engine.max_audio_count,
|
||||
supports_first_last_frame=engine.supports_first_last_frame,
|
||||
supports_universal_reference=engine.supports_universal_reference,
|
||||
multi_generation_enabled=bool(getattr(engine, "multi_generation_enabled", False)),
|
||||
max_generation_count=normalize_generation_count(getattr(engine, "max_generation_count", 1)),
|
||||
)
|
||||
for engine in video_result.scalars().all()
|
||||
]
|
||||
@@ -240,193 +149,6 @@ async def list_generation_ai_engine_options(db: AsyncSession) -> GenerationAIEng
|
||||
)
|
||||
|
||||
|
||||
async def create_async_generation_task(db: AsyncSession, current_user: User, req: GenerationAITaskCreate) -> ChatGenerationTask:
|
||||
"""Create a project-independent chat generation task.
|
||||
|
||||
Important: this writes chat_generation_tasks, not generation_records, so chat
|
||||
image/video generation no longer needs or validates a project_id.
|
||||
"""
|
||||
gen_type = req.gen_type.lower().strip()
|
||||
if gen_type not in ("image", "video"):
|
||||
raise HTTPException(status_code=400, detail="gen_type 仅支持 image 或 video")
|
||||
|
||||
if req.idempotency_key:
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.user_id == current_user.id,
|
||||
ChatGenerationTask.idempotency_key == req.idempotency_key,
|
||||
ChatGenerationTask.generation_mode == "chatapi_async",
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
).order_by(ChatGenerationTask.created_at.desc()).limit(1)
|
||||
)
|
||||
existing = result.scalar_one_or_none()
|
||||
if existing:
|
||||
return existing
|
||||
|
||||
refs = [r.model_dump(exclude_none=True) for r in (req.media_references or [])]
|
||||
refs = await resolve_private_portrait_references(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
media_references=refs,
|
||||
gen_type=gen_type,
|
||||
)
|
||||
now = datetime.now(timezone.utc)
|
||||
task_id = generate_id()
|
||||
|
||||
await assert_user_resource_capacity_available(db, current_user.id)
|
||||
|
||||
if gen_type == "image":
|
||||
if any((r.get("type") or "").lower() == "audio" for r in refs):
|
||||
raise HTTPException(status_code=400, detail="图片生成不支持音频参考素材")
|
||||
engine = await _get_image_engine(db, req.engine_id)
|
||||
sizes = _image_supported_sizes(engine)
|
||||
size = req.image_size or engine.default_size or IMAGE_DEFAULT_SIZE
|
||||
proportion = req.image_proportion or IMAGE_DEFAULT_PROPORTION
|
||||
px = normalize_px(req.image_px)
|
||||
if sizes:
|
||||
if size not in sizes:
|
||||
raise HTTPException(status_code=400, detail=f"图片分辨率档位不支持: {size}")
|
||||
if proportion not in sizes.get(size, {}):
|
||||
raise HTTPException(status_code=400, detail=f"图片比例不支持: {proportion}")
|
||||
px = px or normalize_px((sizes.get(size) or {}).get(proportion))
|
||||
px = px or IMAGE_DEFAULT_PX
|
||||
media_billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
record_id=task_id,
|
||||
gen_type="image",
|
||||
image_size=size,
|
||||
engine_id=engine.id,
|
||||
project_name="AI生成任务",
|
||||
description_prefix="AI创作-",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
attempt_no=1,
|
||||
)
|
||||
snapshot = _build_image_snapshot(engine, size, proportion, px)
|
||||
task = ChatGenerationTask(
|
||||
id=task_id,
|
||||
user_id=current_user.id,
|
||||
original_prompt=req.original_prompt,
|
||||
gen_type="image",
|
||||
image_size=size,
|
||||
image_proportion=proportion,
|
||||
image_px=px,
|
||||
status="generating",
|
||||
generation_mode="chatapi_async",
|
||||
pipeline_stage="queued",
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=_json(snapshot),
|
||||
media_references=_json(refs) if refs else None,
|
||||
credits_cost=round(media_billing.total_charged, 2),
|
||||
idempotency_key=req.idempotency_key,
|
||||
deadline_at=now + timedelta(minutes=settings.CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES),
|
||||
)
|
||||
else:
|
||||
engine = await _get_video_engine(db, req.engine_id)
|
||||
ratio = req.aspect_ratio or VIDEO_DEFAULT_RATIO
|
||||
resolution = req.resolution or VIDEO_DEFAULT_RESOLUTION
|
||||
duration = req.duration or VIDEO_DEFAULT_DURATION
|
||||
ratios = _parse_list(engine.supported_ratios, [])
|
||||
resolutions = _parse_list(engine.supported_resolutions, [])
|
||||
durations = _parse_list(engine.supported_durations, [])
|
||||
if ratios and ratio not in ratios:
|
||||
raise HTTPException(status_code=400, detail=f"视频比例不支持: {ratio}")
|
||||
if resolutions and resolution not in resolutions:
|
||||
raise HTTPException(status_code=400, detail=f"视频分辨率不支持: {resolution}")
|
||||
if durations and duration not in durations:
|
||||
raise HTTPException(status_code=400, detail=f"视频时长不支持: {duration}")
|
||||
if engine.max_duration and duration > engine.max_duration:
|
||||
raise HTTPException(status_code=400, detail=f"视频时长不能超过 {engine.max_duration} 秒")
|
||||
|
||||
input_video_duration = 0.0
|
||||
if refs:
|
||||
video_refs = [r for r in refs if (r.get("type") or "").lower() == "video"]
|
||||
for ref in video_refs:
|
||||
ref_duration = float(ref.get("duration") or 0)
|
||||
if ref_duration < 2:
|
||||
raise HTTPException(status_code=400, detail=f"视频素材最短不能少于 2 秒")
|
||||
input_video_duration += ref_duration
|
||||
if input_video_duration > 15:
|
||||
raise HTTPException(status_code=400, detail=f"所有视频素材总时长不能超过 15 秒,当前 {input_video_duration:.1f} 秒")
|
||||
|
||||
audio_refs = [r for r in refs if (r.get("type") or "").lower() == "audio"]
|
||||
if audio_refs:
|
||||
max_audio_count = int(engine.max_audio_count or 0)
|
||||
if max_audio_count <= 0:
|
||||
raise HTTPException(status_code=400, detail="当前视频引擎不支持音频参考素材")
|
||||
if max_audio_count > AUDIO_MAX_COUNT_LIMIT:
|
||||
max_audio_count = AUDIO_MAX_COUNT_LIMIT
|
||||
if len(audio_refs) > max_audio_count:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"参考音频最多可传 {max_audio_count} 段,当前 {len(audio_refs)} 段",
|
||||
)
|
||||
|
||||
input_audio_duration = 0.0
|
||||
for ref in audio_refs:
|
||||
raw_duration = ref.get("duration")
|
||||
if raw_duration is None:
|
||||
raw_duration = 0.0
|
||||
try:
|
||||
ref_duration = float(raw_duration)
|
||||
except (TypeError, ValueError):
|
||||
ref_duration = 0.0
|
||||
|
||||
if ref_duration < AUDIO_MIN_DURATION_SECONDS or ref_duration > AUDIO_MAX_DURATION_SECONDS:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"单段参考音频时长必须在 {AUDIO_MIN_DURATION_SECONDS}-{AUDIO_MAX_DURATION_SECONDS} 秒之间",
|
||||
)
|
||||
input_audio_duration += ref_duration
|
||||
|
||||
if input_audio_duration > AUDIO_MAX_TOTAL_DURATION_SECONDS:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"所有参考音频总时长不能超过 {AUDIO_MAX_TOTAL_DURATION_SECONDS} 秒,当前 {input_audio_duration:.1f} 秒",
|
||||
)
|
||||
|
||||
media_billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
record_id=task_id,
|
||||
gen_type="video",
|
||||
duration=duration,
|
||||
resolution=resolution,
|
||||
engine_id=engine.id,
|
||||
input_video_duration=input_video_duration if input_video_duration > 0 else None,
|
||||
project_name="AI生成任务",
|
||||
description_prefix="AI创作-",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
attempt_no=1,
|
||||
)
|
||||
snapshot = _build_video_snapshot(engine, ratio, resolution, duration)
|
||||
task = ChatGenerationTask(
|
||||
id=task_id,
|
||||
user_id=current_user.id,
|
||||
original_prompt=req.original_prompt,
|
||||
gen_type="video",
|
||||
duration=duration,
|
||||
aspect_ratio=ratio,
|
||||
resolution=resolution,
|
||||
image_size=req.image_size or IMAGE_DEFAULT_SIZE,
|
||||
image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION,
|
||||
image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX,
|
||||
status="generating",
|
||||
generation_mode="chatapi_async",
|
||||
pipeline_stage="queued",
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=_json(snapshot),
|
||||
media_references=_json(refs) if refs else None,
|
||||
credits_cost=round(media_billing.total_charged, 2),
|
||||
idempotency_key=req.idempotency_key,
|
||||
deadline_at=now + timedelta(hours=settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS),
|
||||
)
|
||||
|
||||
db.add(task)
|
||||
await db.flush()
|
||||
return task
|
||||
|
||||
|
||||
def _resolve_error_message(error_message: str | None) -> str | None:
|
||||
"""匹配 ARK_ERRORS 字典,将原始错误码转换为友好提示。
|
||||
与 app/api/v1/generation.py 的 _record_to_out 保持一致。
|
||||
@@ -479,26 +201,36 @@ def record_to_out(
|
||||
file_name: str | None = None,
|
||||
history_meta: GenerationHistoryMeta | None = None,
|
||||
media_references: list[dict] | None = None,
|
||||
child_items: list[GenerationAITaskOut] | None = None,
|
||||
) -> GenerationAITaskOut:
|
||||
refs = media_references if media_references is not None else _parse_json(task.media_references)
|
||||
snapshot = engine_snapshot_out(_parse_json(task.engine_snapshot_json))
|
||||
|
||||
source = GenerationHistorySourceEnum.CHAT_TASK
|
||||
try:
|
||||
source = GenerationHistorySourceEnum(
|
||||
"chat_task" if task.generation_mode == "chatapi_async" else str(task.generation_mode or "chat_task")
|
||||
)
|
||||
if task.generation_mode in {
|
||||
GenerationMode.CHATAPI_ASYNC.value,
|
||||
GenerationMode.CHATAPI_MAIN.value,
|
||||
GenerationMode.CHATAPI_CHILD.value,
|
||||
}:
|
||||
source = GenerationHistorySourceEnum.CHAT_TASK
|
||||
else:
|
||||
source = GenerationHistorySourceEnum(str(task.generation_mode or "chat_task"))
|
||||
except ValueError:
|
||||
source = GenerationHistorySourceEnum.CHAT_TASK
|
||||
meta = history_meta or build_empty_history_meta(source)
|
||||
|
||||
is_deleted = task.deleted_at is not None
|
||||
is_main = task.generation_mode == GenerationMode.CHATAPI_MAIN.value
|
||||
hide_resource = is_deleted or is_main
|
||||
|
||||
return GenerationAITaskOut(
|
||||
id=task.id,
|
||||
user_id=task.user_id if is_admin else None,
|
||||
user_name=getattr(task, "username", None) if is_admin else None,
|
||||
project_id=None,
|
||||
generated_resource_id=generated_resource_id,
|
||||
file_name=file_name,
|
||||
generated_resource_id=None if hide_resource else generated_resource_id,
|
||||
file_name=None if hide_resource else file_name,
|
||||
history_source=meta.get("history_source"),
|
||||
history_source_label=meta.get("history_source_label"),
|
||||
module_project_id=meta.get("module_project_id"),
|
||||
@@ -515,10 +247,14 @@ def record_to_out(
|
||||
shot_segment_label=meta.get("shot_segment_label"),
|
||||
gen_type=task.gen_type,
|
||||
generation_mode=task.generation_mode,
|
||||
parent_task_id=task.parent_task_id,
|
||||
generation_count=max(1, min(5, int(task.generation_count or 1))),
|
||||
generation_index=task.generation_index,
|
||||
display_status=get_display_status(task),
|
||||
pipeline_stage=task.pipeline_stage,
|
||||
status=task.status,
|
||||
original_prompt=task.original_prompt,
|
||||
# optimized_prompt=task.optimized_prompt,
|
||||
optimized_prompt=task.optimized_prompt,
|
||||
duration=task.duration,
|
||||
aspect_ratio=task.aspect_ratio,
|
||||
resolution=task.resolution,
|
||||
@@ -528,10 +264,9 @@ def record_to_out(
|
||||
media_references=refs,
|
||||
provider_task_id=task.provider_task_id,
|
||||
seedance_task_id=task.seedance_task_id,
|
||||
# remote_result_url=task.remote_result_url,
|
||||
image_url=build_resource_signed_url(task.image_url) if task.image_url else "",
|
||||
video_url=build_resource_signed_url(task.video_url) if task.video_url else "",
|
||||
video_cover_url=build_resource_signed_url(task.video_cover_url) if task.video_cover_url else "",
|
||||
image_url="" if hide_resource else (build_resource_signed_url(task.image_url) if task.image_url else ""),
|
||||
video_url="" if hide_resource else (build_resource_signed_url(task.video_url) if task.video_url else ""),
|
||||
video_cover_url="" if hide_resource else (build_resource_signed_url(task.video_cover_url) if task.video_cover_url else ""),
|
||||
engine_id=task.engine_id,
|
||||
engine_snapshot=snapshot,
|
||||
credits_cost=task.credits_cost or 0.0,
|
||||
@@ -541,33 +276,90 @@ def record_to_out(
|
||||
video_tokens_used=task.video_tokens_used or 0,
|
||||
retry_count=task.retry_count or 0,
|
||||
poll_count=task.poll_count or 0,
|
||||
error_message=_resolve_error_message(task.error_message),
|
||||
error_message=task.error_message if is_main else _resolve_error_message(task.error_message),
|
||||
created_at=task.created_at,
|
||||
generated_at=task.generated_at,
|
||||
child_items=child_items or [],
|
||||
)
|
||||
|
||||
def engine_snapshot_out(snapshot: dict) -> dict:
|
||||
"""
|
||||
从完整的 engine_snapshot 中过滤出需要返回的字段
|
||||
"""
|
||||
"""从完整引擎快照中过滤前端允许展示的字段。"""
|
||||
if not snapshot:
|
||||
return {}
|
||||
keys = (
|
||||
"engine_type", "id", "name", "provider", "model_name",
|
||||
"supported_models", "default_size", "selected_size",
|
||||
"selected_proportion", "selected_px", "supported_ratios",
|
||||
"supported_resolutions", "supported_durations", "max_duration",
|
||||
"max_audio_count", "selected_ratio", "selected_resolution",
|
||||
"selected_duration", "generation_count", "multi_generation_enabled",
|
||||
"max_generation_count", "multi_image_max_images", "max_reference_image_count", "output_format",
|
||||
)
|
||||
result = {key: snapshot.get(key) for key in keys if key in snapshot}
|
||||
result.setdefault("generation_count", 1)
|
||||
return result
|
||||
|
||||
return {
|
||||
"engine_type": snapshot.get("engine_type"),
|
||||
"id": snapshot.get("id"),
|
||||
"name": snapshot.get("name"),
|
||||
"provider": snapshot.get("provider"),
|
||||
# "api_base": snapshot.get("api_base"),
|
||||
# "api_key_masked": snapshot.get("api_key_masked"),
|
||||
"model_name": snapshot.get("model_name"),
|
||||
# "generate_url": snapshot.get("generate_url"),
|
||||
"supported_models": snapshot.get("supported_models", []),
|
||||
"default_size": snapshot.get("default_size"),
|
||||
"selected_size": snapshot.get("selected_size"),
|
||||
"selected_proportion": snapshot.get("selected_proportion"),
|
||||
"selected_px": snapshot.get("selected_px")
|
||||
}
|
||||
|
||||
async def build_task_out_list(
|
||||
db: AsyncSession,
|
||||
tasks: list[ChatGenerationTask],
|
||||
*,
|
||||
is_admin: bool = False,
|
||||
viewer_user_id: str | None = None,
|
||||
) -> list[GenerationAITaskOut]:
|
||||
"""批量回填主任务子项、资源账本和参考素材,避免列表 N+1。"""
|
||||
if not tasks:
|
||||
return []
|
||||
parent_ids = [
|
||||
task.id for task in tasks
|
||||
if task.generation_mode == GenerationMode.CHATAPI_MAIN.value
|
||||
]
|
||||
children_map = await load_children_map(db, parent_ids, include_deleted=True)
|
||||
children = [child for items in children_map.values() for child in items]
|
||||
resource_task_ids = [
|
||||
task.id for task in [*tasks, *children]
|
||||
if task.generation_mode != GenerationMode.CHATAPI_MAIN.value and task.deleted_at is None
|
||||
]
|
||||
resource_info_map = await batch_get_generated_resource_info_map(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=resource_task_ids,
|
||||
)
|
||||
reference_display_map = await _resolve_task_reference_display_map(
|
||||
db,
|
||||
tasks,
|
||||
user_id=viewer_user_id,
|
||||
)
|
||||
|
||||
output: list[GenerationAITaskOut] = []
|
||||
for task in tasks:
|
||||
refs = reference_display_map.get(task.id)
|
||||
child_out: list[GenerationAITaskOut] = []
|
||||
for child in children_map.get(task.id, []):
|
||||
if is_admin:
|
||||
child.username = getattr(task, "username", None)
|
||||
resource = resource_info_map.get(child.id, {})
|
||||
child_out.append(
|
||||
record_to_out(
|
||||
child,
|
||||
is_admin=is_admin,
|
||||
generated_resource_id=resource.get("resource_id"),
|
||||
file_name=resource.get("file_name"),
|
||||
media_references=refs,
|
||||
)
|
||||
)
|
||||
resource = resource_info_map.get(task.id, {})
|
||||
output.append(
|
||||
record_to_out(
|
||||
task,
|
||||
is_admin=is_admin,
|
||||
generated_resource_id=resource.get("resource_id"),
|
||||
file_name=resource.get("file_name"),
|
||||
media_references=refs,
|
||||
child_items=child_out,
|
||||
)
|
||||
)
|
||||
return output
|
||||
|
||||
async def list_async_generation_tasks(
|
||||
db: AsyncSession,
|
||||
@@ -594,7 +386,7 @@ async def list_async_generation_tasks(
|
||||
query = select(ChatGenerationTask)
|
||||
|
||||
query = query.where(
|
||||
ChatGenerationTask.generation_mode == "chatapi_async",
|
||||
ChatGenerationTask.generation_mode.in_(list(CHAT_TOP_LEVEL_MODES)),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
|
||||
@@ -623,7 +415,7 @@ async def list_async_generation_tasks(
|
||||
total = (await db.execute(count_query)).scalar_one()
|
||||
|
||||
result = await db.execute(
|
||||
query.order_by(ChatGenerationTask.created_at.desc())
|
||||
query.order_by(ChatGenerationTask.created_at.desc(), ChatGenerationTask.id.desc())
|
||||
.offset((page - 1) * page_size)
|
||||
.limit(page_size)
|
||||
)
|
||||
@@ -677,12 +469,12 @@ def _parse_history_date(value: str) -> date:
|
||||
|
||||
|
||||
def _history_base_filters(user_id: str, gen_type: str, source: GenerationHistorySourceEnum):
|
||||
task_mode = get_generation_history_task_mode(source)
|
||||
if not task_mode:
|
||||
task_modes = get_generation_history_task_modes(source)
|
||||
if not task_modes:
|
||||
raise HTTPException(status_code=400, detail="history_source 不支持查询 ChatGenerationTask 历史")
|
||||
return [
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.generation_mode == task_mode.value,
|
||||
ChatGenerationTask.generation_mode.in_([mode.value for mode in task_modes]),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
ChatGenerationTask.status == "completed",
|
||||
ChatGenerationTask.gen_type == gen_type,
|
||||
@@ -1193,13 +985,3 @@ async def list_generation_history_day_items(
|
||||
],
|
||||
}
|
||||
|
||||
async def soft_delete_chat_generation_task(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
task: ChatGenerationTask,
|
||||
deleted_at: datetime | None = None,
|
||||
) -> int:
|
||||
"""软删 ChatGenerationTask 并联动软删资源账本,返回释放的 active 空间字节数。"""
|
||||
deleted_at = deleted_at or datetime.now(timezone.utc)
|
||||
task.deleted_at = deleted_at
|
||||
return await soft_delete_chat_task_resources(db, task.id, deleted_at=deleted_at)
|
||||
@@ -0,0 +1,587 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.audio_reference import (
|
||||
AUDIO_MAX_COUNT_LIMIT,
|
||||
AUDIO_MAX_DURATION_SECONDS,
|
||||
AUDIO_MAX_TOTAL_DURATION_SECONDS,
|
||||
AUDIO_MIN_DURATION_SECONDS,
|
||||
)
|
||||
from app.enums.generation_task import CHAT_TOP_LEVEL_MODES, GenerationMode, GenerationType
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.models.user import User
|
||||
from app.schemas.generation_ai import GenerationAITaskCreate
|
||||
from app.services.generation.ai.engine_service import (
|
||||
IMAGE_DEFAULT_PROPORTION,
|
||||
IMAGE_DEFAULT_PX,
|
||||
IMAGE_DEFAULT_SIZE,
|
||||
VIDEO_DEFAULT_DURATION,
|
||||
VIDEO_DEFAULT_RATIO,
|
||||
VIDEO_DEFAULT_RESOLUTION,
|
||||
build_image_snapshot,
|
||||
build_video_snapshot,
|
||||
get_image_engine,
|
||||
get_video_engine,
|
||||
image_supported_sizes,
|
||||
normalize_generation_count,
|
||||
normalize_px,
|
||||
parse_json_list,
|
||||
)
|
||||
from app.services.generation.billing_service import OWNER_CHAT_GENERATION_TASK, charge_generation_media_by_params
|
||||
from app.services.operation_log_service import log_operation_event
|
||||
from app.services.private_portrait.reference_resolver import resolve_private_portrait_references
|
||||
from app.services.resource_capacity_service import assert_user_resource_capacity_available
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
|
||||
@dataclass(slots=True)
|
||||
class GenerationTaskCreateResult:
|
||||
top_level_task_id: str
|
||||
enqueue_task_ids: list[str] = field(default_factory=list)
|
||||
child_task_ids: list[str] = field(default_factory=list)
|
||||
generation_count: int = 1
|
||||
gen_type: str = GenerationType.IMAGE.value
|
||||
created: bool = True
|
||||
|
||||
|
||||
def _json(data: Any) -> str | None:
|
||||
if data is None:
|
||||
return None
|
||||
return json.dumps(data, ensure_ascii=False, default=str)
|
||||
|
||||
|
||||
async def find_existing_top_level_task(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
user_id: str,
|
||||
idempotency_key: str | None,
|
||||
) -> ChatGenerationTask | None:
|
||||
if not idempotency_key:
|
||||
return None
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.idempotency_key == idempotency_key,
|
||||
ChatGenerationTask.generation_mode.in_(list(CHAT_TOP_LEVEL_MODES)),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.order_by(ChatGenerationTask.created_at.desc())
|
||||
.limit(1)
|
||||
)
|
||||
return result.scalar_one_or_none()
|
||||
|
||||
|
||||
def _validate_video_references(refs: list[dict], *, max_audio_count: int) -> float:
|
||||
input_video_duration = 0.0
|
||||
for ref in refs:
|
||||
if (ref.get("type") or "").lower() != GenerationType.VIDEO.value:
|
||||
continue
|
||||
try:
|
||||
ref_duration = float(ref.get("duration") or 0)
|
||||
except (TypeError, ValueError):
|
||||
ref_duration = 0.0
|
||||
if ref_duration < 2:
|
||||
raise HTTPException(status_code=400, detail="视频素材最短不能少于 2 秒")
|
||||
input_video_duration += ref_duration
|
||||
if input_video_duration > 15:
|
||||
raise HTTPException(status_code=400, detail=f"所有视频素材总时长不能超过 15 秒,当前 {input_video_duration:.1f} 秒")
|
||||
|
||||
audio_refs = [ref for ref in refs if (ref.get("type") or "").lower() == "audio"]
|
||||
if audio_refs:
|
||||
allowed_count = min(AUDIO_MAX_COUNT_LIMIT, max(0, int(max_audio_count or 0)))
|
||||
if allowed_count <= 0:
|
||||
raise HTTPException(status_code=400, detail="当前视频引擎不支持音频参考素材")
|
||||
if len(audio_refs) > allowed_count:
|
||||
raise HTTPException(status_code=400, detail=f"参考音频最多可传 {allowed_count} 段,当前 {len(audio_refs)} 段")
|
||||
|
||||
input_audio_duration = 0.0
|
||||
for ref in audio_refs:
|
||||
try:
|
||||
ref_duration = float(ref.get("duration") or 0)
|
||||
except (TypeError, ValueError):
|
||||
ref_duration = 0.0
|
||||
if ref_duration < AUDIO_MIN_DURATION_SECONDS or ref_duration > AUDIO_MAX_DURATION_SECONDS:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"单段参考音频时长必须在 {AUDIO_MIN_DURATION_SECONDS}-{AUDIO_MAX_DURATION_SECONDS} 秒之间",
|
||||
)
|
||||
input_audio_duration += ref_duration
|
||||
if input_audio_duration > AUDIO_MAX_TOTAL_DURATION_SECONDS:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"所有参考音频总时长不能超过 {AUDIO_MAX_TOTAL_DURATION_SECONDS} 秒,当前 {input_audio_duration:.1f} 秒",
|
||||
)
|
||||
return input_video_duration
|
||||
|
||||
|
||||
def _base_task_kwargs(
|
||||
*,
|
||||
task_id: str,
|
||||
user_id: str,
|
||||
req: GenerationAITaskCreate,
|
||||
gen_type: str,
|
||||
generation_mode: str,
|
||||
generation_count: int,
|
||||
engine_id: str,
|
||||
engine_snapshot_json: str,
|
||||
media_references_json: str | None,
|
||||
deadline_at: datetime,
|
||||
parent_task_id: str | None = None,
|
||||
generation_index: int | None = None,
|
||||
credits_cost: float = 0.0,
|
||||
idempotency_key: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"id": task_id,
|
||||
"user_id": user_id,
|
||||
"original_prompt": req.original_prompt,
|
||||
"gen_type": gen_type,
|
||||
"status": "generating",
|
||||
"generation_mode": generation_mode,
|
||||
"pipeline_stage": "queued",
|
||||
"parent_task_id": parent_task_id,
|
||||
"generation_count": generation_count,
|
||||
"generation_index": generation_index,
|
||||
"engine_id": engine_id,
|
||||
"engine_snapshot_json": engine_snapshot_json,
|
||||
"media_references": media_references_json,
|
||||
"credits_cost": round(float(credits_cost or 0), 2),
|
||||
"idempotency_key": idempotency_key,
|
||||
"deadline_at": deadline_at,
|
||||
}
|
||||
|
||||
|
||||
async def create_generation_task_group(
|
||||
db: AsyncSession,
|
||||
current_user: User,
|
||||
req: GenerationAITaskCreate,
|
||||
) -> GenerationTaskCreateResult:
|
||||
"""创建单份 chatapi_async 或多份 chatapi_main/chatapi_child 任务组。
|
||||
|
||||
本函数只 flush,不主动 commit。调用方提交成功后才能投递 Celery。
|
||||
"""
|
||||
gen_type = (req.gen_type or "").lower().strip()
|
||||
if gen_type not in (GenerationType.IMAGE.value, GenerationType.VIDEO.value):
|
||||
raise HTTPException(status_code=400, detail="gen_type 仅支持 image 或 video")
|
||||
|
||||
existing = await find_existing_top_level_task(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
idempotency_key=req.idempotency_key,
|
||||
)
|
||||
if existing:
|
||||
return GenerationTaskCreateResult(
|
||||
top_level_task_id=existing.id,
|
||||
generation_count=int(existing.generation_count or 1),
|
||||
gen_type=existing.gen_type,
|
||||
created=False,
|
||||
)
|
||||
|
||||
refs = [item.model_dump(exclude_none=True) for item in (req.media_references or [])]
|
||||
refs = await resolve_private_portrait_references(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
media_references=refs,
|
||||
gen_type=gen_type,
|
||||
)
|
||||
media_references_json = _json(refs) if refs else None
|
||||
await assert_user_resource_capacity_available(db, current_user.id)
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
main_id = generate_id()
|
||||
child_ids: list[str] = []
|
||||
enqueue_ids: list[str] = []
|
||||
total_billed_credits = 0.0
|
||||
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="BATCH_CREATE_START",
|
||||
event_status="started",
|
||||
source="service",
|
||||
user_id=current_user.id,
|
||||
group_id=main_id,
|
||||
detail={
|
||||
"gen_type": gen_type,
|
||||
"requested_generation_count": normalize_generation_count(req.generation_count),
|
||||
"idempotency_key_present": bool(req.idempotency_key),
|
||||
},
|
||||
)
|
||||
|
||||
if gen_type == GenerationType.IMAGE.value:
|
||||
if any((ref.get("type") or "").lower() == "audio" for ref in refs):
|
||||
raise HTTPException(status_code=400, detail="图片生成不支持音频参考素材")
|
||||
|
||||
engine = await get_image_engine(db, req.engine_id)
|
||||
generation_count = normalize_generation_count(req.generation_count)
|
||||
multi_generation_enabled = bool(getattr(engine, "multi_generation_enabled", False))
|
||||
max_generation_count = normalize_generation_count(getattr(engine, "max_generation_count", 1))
|
||||
if generation_count > 1 and not multi_generation_enabled:
|
||||
raise HTTPException(status_code=400, detail="当前图片引擎未开启多份生成,本次生成数量只能为 1")
|
||||
if generation_count > max_generation_count:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"当前图片引擎本次最多允许生成 {max_generation_count} 份",
|
||||
)
|
||||
|
||||
reference_image_count = sum(
|
||||
1 for ref in refs if (ref.get("type") or "").lower() == GenerationType.IMAGE.value
|
||||
)
|
||||
max_reference_count = max(0, int(getattr(engine, "max_reference_image_count", 14) or 0))
|
||||
multi_image_max_images = max(1, int(getattr(engine, "multi_image_max_images", 15) or 15))
|
||||
if reference_image_count > max_reference_count:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"当前图片引擎最多支持 {max_reference_count} 张参考图,当前 {reference_image_count} 张",
|
||||
)
|
||||
if generation_count > 1 and reference_image_count + generation_count > multi_image_max_images:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=(
|
||||
f"参考图数量与生成数量合计不能超过 {multi_image_max_images} 张,"
|
||||
f"当前参考图 {reference_image_count} 张、生成 {generation_count} 张"
|
||||
),
|
||||
)
|
||||
|
||||
sizes = image_supported_sizes(engine)
|
||||
size = req.image_size or engine.default_size or IMAGE_DEFAULT_SIZE
|
||||
proportion = req.image_proportion or IMAGE_DEFAULT_PROPORTION
|
||||
px = normalize_px(req.image_px)
|
||||
if sizes:
|
||||
if size not in sizes:
|
||||
raise HTTPException(status_code=400, detail=f"图片分辨率档位不支持: {size}")
|
||||
if proportion not in sizes.get(size, {}):
|
||||
raise HTTPException(status_code=400, detail=f"图片比例不支持: {proportion}")
|
||||
px = px or normalize_px((sizes.get(size) or {}).get(proportion))
|
||||
px = px or IMAGE_DEFAULT_PX
|
||||
|
||||
mode = GenerationMode.CHATAPI_ASYNC.value if generation_count == 1 else GenerationMode.CHATAPI_MAIN.value
|
||||
billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
record_id=main_id,
|
||||
gen_type=GenerationType.IMAGE.value,
|
||||
image_size=size,
|
||||
engine_id=engine.id,
|
||||
project_name="AI生成任务",
|
||||
description_prefix="AI创作-",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
attempt_no=1,
|
||||
quantity=generation_count,
|
||||
)
|
||||
image_snapshot = build_image_snapshot(engine, size, proportion, px)
|
||||
image_snapshot["generation_count"] = generation_count
|
||||
snapshot_json = _json(image_snapshot) or "{}"
|
||||
total_billed_credits = round(float(billing.total_charged or 0), 2)
|
||||
task = ChatGenerationTask(
|
||||
**_base_task_kwargs(
|
||||
task_id=main_id,
|
||||
user_id=current_user.id,
|
||||
req=req,
|
||||
gen_type=GenerationType.IMAGE.value,
|
||||
generation_mode=mode,
|
||||
generation_count=generation_count,
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=snapshot_json,
|
||||
media_references_json=media_references_json,
|
||||
deadline_at=now + timedelta(minutes=settings.CHATAPI_ASYNC_IMAGE_DEADLINE_MINUTES),
|
||||
credits_cost=billing.total_charged,
|
||||
idempotency_key=req.idempotency_key,
|
||||
),
|
||||
image_size=size,
|
||||
image_proportion=proportion,
|
||||
image_px=px,
|
||||
)
|
||||
db.add(task)
|
||||
enqueue_ids.append(task.id)
|
||||
else:
|
||||
engine = await get_video_engine(db, req.engine_id)
|
||||
generation_count = normalize_generation_count(req.generation_count)
|
||||
multi_generation_enabled = bool(getattr(engine, "multi_generation_enabled", False))
|
||||
max_generation_count = normalize_generation_count(getattr(engine, "max_generation_count", 1))
|
||||
if generation_count > 1 and not multi_generation_enabled:
|
||||
raise HTTPException(status_code=400, detail="当前视频引擎未开启多份生成,本次生成数量只能为 1")
|
||||
if generation_count > max_generation_count:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"当前视频引擎本次最多允许生成 {max_generation_count} 份",
|
||||
)
|
||||
ratio = req.aspect_ratio or VIDEO_DEFAULT_RATIO
|
||||
resolution = req.resolution or VIDEO_DEFAULT_RESOLUTION
|
||||
duration = req.duration or VIDEO_DEFAULT_DURATION
|
||||
ratios = parse_json_list(engine.supported_ratios, [])
|
||||
resolutions = parse_json_list(engine.supported_resolutions, [])
|
||||
durations = parse_json_list(engine.supported_durations, [])
|
||||
if ratios and ratio not in ratios:
|
||||
raise HTTPException(status_code=400, detail=f"视频比例不支持: {ratio}")
|
||||
if resolutions and resolution not in resolutions:
|
||||
raise HTTPException(status_code=400, detail=f"视频分辨率不支持: {resolution}")
|
||||
if durations and duration not in durations:
|
||||
raise HTTPException(status_code=400, detail=f"视频时长不支持: {duration}")
|
||||
if engine.max_duration and duration > engine.max_duration:
|
||||
raise HTTPException(status_code=400, detail=f"视频时长不能超过 {engine.max_duration} 秒")
|
||||
|
||||
input_video_duration = _validate_video_references(refs, max_audio_count=engine.max_audio_count)
|
||||
video_snapshot = build_video_snapshot(engine, ratio, resolution, duration)
|
||||
video_snapshot["generation_count"] = generation_count
|
||||
snapshot_json = _json(video_snapshot) or "{}"
|
||||
deadline_at = now + timedelta(hours=settings.CHATAPI_ASYNC_VIDEO_FINAL_DEADLINE_HOURS)
|
||||
|
||||
if generation_count == 1:
|
||||
billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
record_id=main_id,
|
||||
gen_type=GenerationType.VIDEO.value,
|
||||
duration=duration,
|
||||
resolution=resolution,
|
||||
engine_id=engine.id,
|
||||
input_video_duration=input_video_duration if input_video_duration > 0 else None,
|
||||
project_name="AI生成任务",
|
||||
description_prefix="AI创作-",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
attempt_no=1,
|
||||
)
|
||||
total_billed_credits = round(float(billing.total_charged or 0), 2)
|
||||
task = ChatGenerationTask(
|
||||
**_base_task_kwargs(
|
||||
task_id=main_id,
|
||||
user_id=current_user.id,
|
||||
req=req,
|
||||
gen_type=GenerationType.VIDEO.value,
|
||||
generation_mode=GenerationMode.CHATAPI_ASYNC.value,
|
||||
generation_count=1,
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=snapshot_json,
|
||||
media_references_json=media_references_json,
|
||||
deadline_at=deadline_at,
|
||||
credits_cost=billing.total_charged,
|
||||
idempotency_key=req.idempotency_key,
|
||||
),
|
||||
duration=duration,
|
||||
aspect_ratio=ratio,
|
||||
resolution=resolution,
|
||||
image_size=req.image_size or IMAGE_DEFAULT_SIZE,
|
||||
image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION,
|
||||
image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX,
|
||||
)
|
||||
db.add(task)
|
||||
enqueue_ids.append(task.id)
|
||||
else:
|
||||
main_task = ChatGenerationTask(
|
||||
**_base_task_kwargs(
|
||||
task_id=main_id,
|
||||
user_id=current_user.id,
|
||||
req=req,
|
||||
gen_type=GenerationType.VIDEO.value,
|
||||
generation_mode=GenerationMode.CHATAPI_MAIN.value,
|
||||
generation_count=generation_count,
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=snapshot_json,
|
||||
media_references_json=media_references_json,
|
||||
deadline_at=deadline_at,
|
||||
idempotency_key=req.idempotency_key,
|
||||
),
|
||||
duration=duration,
|
||||
aspect_ratio=ratio,
|
||||
resolution=resolution,
|
||||
image_size=req.image_size or IMAGE_DEFAULT_SIZE,
|
||||
image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION,
|
||||
image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX,
|
||||
)
|
||||
db.add(main_task)
|
||||
await db.flush()
|
||||
|
||||
total_credits = 0.0
|
||||
children: list[ChatGenerationTask] = []
|
||||
for generation_index in range(1, generation_count + 1):
|
||||
child_id = generate_id()
|
||||
billing = await charge_generation_media_by_params(
|
||||
db,
|
||||
user_id=current_user.id,
|
||||
record_id=child_id,
|
||||
gen_type=GenerationType.VIDEO.value,
|
||||
duration=duration,
|
||||
resolution=resolution,
|
||||
engine_id=engine.id,
|
||||
input_video_duration=input_video_duration if input_video_duration > 0 else None,
|
||||
project_name="AI生成任务",
|
||||
description_prefix=f"AI创作-第{generation_index}份-",
|
||||
owner_type=OWNER_CHAT_GENERATION_TASK,
|
||||
attempt_no=1,
|
||||
)
|
||||
child = ChatGenerationTask(
|
||||
**_base_task_kwargs(
|
||||
task_id=child_id,
|
||||
user_id=current_user.id,
|
||||
req=req,
|
||||
gen_type=GenerationType.VIDEO.value,
|
||||
generation_mode=GenerationMode.CHATAPI_CHILD.value,
|
||||
generation_count=generation_count,
|
||||
generation_index=generation_index,
|
||||
parent_task_id=main_id,
|
||||
engine_id=engine.id,
|
||||
engine_snapshot_json=snapshot_json,
|
||||
media_references_json=media_references_json,
|
||||
deadline_at=deadline_at,
|
||||
credits_cost=billing.total_charged,
|
||||
),
|
||||
duration=duration,
|
||||
aspect_ratio=ratio,
|
||||
resolution=resolution,
|
||||
image_size=req.image_size or IMAGE_DEFAULT_SIZE,
|
||||
image_proportion=req.image_proportion or IMAGE_DEFAULT_PROPORTION,
|
||||
image_px=normalize_px(req.image_px) or IMAGE_DEFAULT_PX,
|
||||
)
|
||||
children.append(child)
|
||||
child_ids.append(child_id)
|
||||
enqueue_ids.append(child_id)
|
||||
total_credits = round(total_credits + billing.total_charged, 2)
|
||||
db.add_all(children)
|
||||
main_task.credits_cost = total_credits
|
||||
total_billed_credits = total_credits
|
||||
|
||||
await db.flush()
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="BATCH_BILLING_SUCCESS",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=current_user.id,
|
||||
group_id=main_id,
|
||||
detail={
|
||||
"gen_type": gen_type,
|
||||
"generation_count": generation_count,
|
||||
"total_billed_credits": total_billed_credits,
|
||||
},
|
||||
)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="BATCH_CHILDREN_CREATED" if child_ids else "BATCH_MAIN_CREATED",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=current_user.id,
|
||||
group_id=main_id,
|
||||
detail={
|
||||
"gen_type": gen_type,
|
||||
"generation_count": generation_count,
|
||||
"child_task_ids": child_ids,
|
||||
"enqueue_task_ids": enqueue_ids,
|
||||
},
|
||||
)
|
||||
return GenerationTaskCreateResult(
|
||||
top_level_task_id=main_id,
|
||||
enqueue_task_ids=enqueue_ids,
|
||||
child_task_ids=child_ids,
|
||||
generation_count=generation_count,
|
||||
gen_type=gen_type,
|
||||
created=True,
|
||||
)
|
||||
|
||||
|
||||
async def enqueue_created_generation_tasks(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
task_ids: list[str],
|
||||
) -> list[str]:
|
||||
"""在业务事务提交后投递任务;返回投递失败的任务ID。
|
||||
|
||||
投递失败会在补偿事务中将对应任务置为失败并幂等退款,视频子任务
|
||||
同时触发主任务状态汇总。调用方不应在初始事务提交前调用本函数。
|
||||
"""
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
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.tasks.generation_create_tasks import chatapi_create_generation_task
|
||||
|
||||
normalized_ids = list(dict.fromkeys(str(item) for item in task_ids if item))
|
||||
meta_result = await db.execute(
|
||||
select(
|
||||
ChatGenerationTask.id,
|
||||
ChatGenerationTask.user_id,
|
||||
ChatGenerationTask.parent_task_id,
|
||||
ChatGenerationTask.generation_index,
|
||||
).where(ChatGenerationTask.id.in_(normalized_ids))
|
||||
) if normalized_ids else None
|
||||
task_meta = {
|
||||
str(row.id): {
|
||||
"user_id": str(row.user_id),
|
||||
"parent_task_id": str(row.parent_task_id) if row.parent_task_id else None,
|
||||
"generation_index": row.generation_index,
|
||||
}
|
||||
for row in (meta_result.all() if meta_result is not None else [])
|
||||
}
|
||||
|
||||
failed_ids: list[str] = []
|
||||
for task_id in normalized_ids:
|
||||
meta = task_meta.get(task_id, {})
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_ENQUEUE_START",
|
||||
event_status="started",
|
||||
source="api",
|
||||
user_id=meta.get("user_id"),
|
||||
group_id=meta.get("parent_task_id") or task_id,
|
||||
task_id=task_id,
|
||||
detail={"generation_index": meta.get("generation_index")},
|
||||
)
|
||||
try:
|
||||
chatapi_create_generation_task.delay(task_id)
|
||||
await log_task_event(
|
||||
task_id=task_id,
|
||||
event_type="CHILD_ENQUEUE_SUCCESS",
|
||||
to_status="generating",
|
||||
to_stage="queued",
|
||||
detail={"task_id": task_id},
|
||||
)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_ENQUEUE_SUCCESS",
|
||||
event_status="success",
|
||||
source="api",
|
||||
user_id=meta.get("user_id"),
|
||||
group_id=meta.get("parent_task_id") or task_id,
|
||||
task_id=task_id,
|
||||
detail={"generation_index": meta.get("generation_index")},
|
||||
)
|
||||
except Exception as exc:
|
||||
failed_ids.append(task_id)
|
||||
await db.rollback()
|
||||
failed_task = await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task_id=task_id,
|
||||
error_message=f"任务队列投递失败: {exc}",
|
||||
pipeline_stage="failed",
|
||||
)
|
||||
await aggregate_parent_for_child(db, failed_task)
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
task_id=task_id,
|
||||
event_type="CHILD_ENQUEUE_FAILED",
|
||||
to_status="failed",
|
||||
to_stage="failed",
|
||||
message=str(exc),
|
||||
detail={"task_id": task_id},
|
||||
)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_ENQUEUE_FAILED",
|
||||
event_status="failed",
|
||||
source="api",
|
||||
user_id=getattr(failed_task, "user_id", None),
|
||||
group_id=getattr(failed_task, "parent_task_id", None) or task_id,
|
||||
task_id=task_id,
|
||||
message=str(exc),
|
||||
detail={"physical_files_deleted": False},
|
||||
error=str(exc),
|
||||
)
|
||||
return failed_ids
|
||||
@@ -0,0 +1,426 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import Counter, defaultdict
|
||||
from datetime import datetime, timezone
|
||||
from typing import Iterable, Sequence
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.enums.generation_task import (
|
||||
ChatGenerationPipelineStage,
|
||||
ChatGenerationTaskStatus,
|
||||
GenerationMode,
|
||||
)
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.services.operation_log_service import log_operation_event
|
||||
from app.services.resource_accounting_service import (
|
||||
SOURCE_MODEL_CHAT_TASK,
|
||||
soft_delete_resources_by_source,
|
||||
)
|
||||
|
||||
|
||||
ACTIVE_STAGES = {
|
||||
ChatGenerationPipelineStage.QUEUED.value,
|
||||
ChatGenerationPipelineStage.PREPARING.value,
|
||||
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
|
||||
ChatGenerationPipelineStage.WAITING_REMOTE.value,
|
||||
ChatGenerationPipelineStage.POLLING.value,
|
||||
ChatGenerationPipelineStage.RESULT_READY.value,
|
||||
ChatGenerationPipelineStage.DOWNLOAD_QUEUED.value,
|
||||
ChatGenerationPipelineStage.DOWNLOADING.value,
|
||||
ChatGenerationPipelineStage.RETRY_WAITING.value,
|
||||
}
|
||||
|
||||
|
||||
def is_task_active(task: ChatGenerationTask) -> bool:
|
||||
return task.deleted_at is None and (
|
||||
task.status == ChatGenerationTaskStatus.GENERATING.value
|
||||
or (task.pipeline_stage or "") in ACTIVE_STAGES
|
||||
)
|
||||
|
||||
|
||||
def get_display_status(task: ChatGenerationTask) -> str:
|
||||
if task.deleted_at is not None:
|
||||
return "deleted"
|
||||
if task.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value:
|
||||
return "download_failed"
|
||||
return task.status or ChatGenerationTaskStatus.PENDING.value
|
||||
|
||||
|
||||
async def load_children_map(
|
||||
db: AsyncSession,
|
||||
parent_ids: Sequence[str] | Iterable[str],
|
||||
*,
|
||||
include_deleted: bool = True,
|
||||
) -> dict[str, list[ChatGenerationTask]]:
|
||||
ids = list(dict.fromkeys(str(item) for item in parent_ids if item))
|
||||
if not ids:
|
||||
return {}
|
||||
query = select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.parent_task_id.in_(ids),
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
|
||||
)
|
||||
if not include_deleted:
|
||||
query = query.where(ChatGenerationTask.deleted_at.is_(None))
|
||||
result = await db.execute(
|
||||
query.order_by(
|
||||
ChatGenerationTask.parent_task_id.asc(),
|
||||
ChatGenerationTask.generation_index.asc(),
|
||||
ChatGenerationTask.created_at.asc(),
|
||||
)
|
||||
)
|
||||
grouped: dict[str, list[ChatGenerationTask]] = defaultdict(list)
|
||||
for task in result.scalars().all():
|
||||
if task.parent_task_id:
|
||||
grouped[task.parent_task_id].append(task)
|
||||
return dict(grouped)
|
||||
|
||||
|
||||
async def load_task_and_children(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
task_id: str,
|
||||
user_id: str | None = None,
|
||||
include_deleted_children: bool = True,
|
||||
) -> tuple[ChatGenerationTask | None, list[ChatGenerationTask]]:
|
||||
query = select(ChatGenerationTask).where(ChatGenerationTask.id == task_id)
|
||||
if user_id:
|
||||
query = query.where(ChatGenerationTask.user_id == user_id)
|
||||
result = await db.execute(query.limit(1))
|
||||
task = result.scalar_one_or_none()
|
||||
if not task:
|
||||
return None, []
|
||||
if task.generation_mode == GenerationMode.CHATAPI_CHILD.value and task.parent_task_id:
|
||||
parent_result = await db.execute(
|
||||
select(ChatGenerationTask).where(ChatGenerationTask.id == task.parent_task_id).limit(1)
|
||||
)
|
||||
parent = parent_result.scalar_one_or_none()
|
||||
return parent or task, [task]
|
||||
if task.generation_mode != GenerationMode.CHATAPI_MAIN.value:
|
||||
return task, []
|
||||
children_map = await load_children_map(
|
||||
db,
|
||||
[task.id],
|
||||
include_deleted=include_deleted_children,
|
||||
)
|
||||
return task, children_map.get(task.id, [])
|
||||
|
||||
|
||||
def _generation_result_status(task: ChatGenerationTask) -> str:
|
||||
"""返回任务真实生成结果,不受资源软删除影响。"""
|
||||
if task.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value:
|
||||
return "download_failed"
|
||||
if task.status == ChatGenerationTaskStatus.FAILED.value or (task.pipeline_stage or "") in {
|
||||
ChatGenerationPipelineStage.FAILED.value,
|
||||
ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
}:
|
||||
return "failed"
|
||||
if is_task_active(task):
|
||||
return "generating"
|
||||
if task.status == ChatGenerationTaskStatus.COMPLETED.value:
|
||||
return "completed"
|
||||
return task.status or "pending"
|
||||
|
||||
|
||||
def _build_summary(children: list[ChatGenerationTask]) -> str | None:
|
||||
if not children:
|
||||
return None
|
||||
result_counters: Counter[str] = Counter(_generation_result_status(child) for child in children)
|
||||
labels = {
|
||||
"completed": "完成",
|
||||
"failed": "生成失败",
|
||||
"download_failed": "下载失败",
|
||||
"generating": "生成中",
|
||||
"pending": "待处理",
|
||||
}
|
||||
parts = [f"{count}项{labels.get(status, status)}" for status, count in result_counters.items() if count]
|
||||
deleted_count = sum(1 for child in children if child.deleted_at is not None)
|
||||
if deleted_count:
|
||||
parts.append(f"{deleted_count}项资源已删除")
|
||||
return f"{len(children)}项中" + ",".join(parts)
|
||||
|
||||
|
||||
async def aggregate_main_task_status(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
parent_task_id: str,
|
||||
) -> ChatGenerationTask | None:
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == parent_task_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
main = result.scalar_one_or_none()
|
||||
if not main or main.deleted_at is not None:
|
||||
return main
|
||||
|
||||
children_map = await load_children_map(db, [parent_task_id], include_deleted=True)
|
||||
children = children_map.get(parent_task_id, [])
|
||||
if not children:
|
||||
return main
|
||||
|
||||
previous_status = main.status
|
||||
previous_stage = main.pipeline_stage
|
||||
active_children = [child for child in children if is_task_active(child)]
|
||||
failed_children = [
|
||||
child
|
||||
for child in children
|
||||
if (
|
||||
child.status == ChatGenerationTaskStatus.FAILED.value
|
||||
or (child.pipeline_stage or "") in {
|
||||
ChatGenerationPipelineStage.FAILED.value,
|
||||
ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
ChatGenerationPipelineStage.DOWNLOAD_FAILED.value,
|
||||
}
|
||||
)
|
||||
]
|
||||
completed_children = [
|
||||
child for child in children if child.status == ChatGenerationTaskStatus.COMPLETED.value
|
||||
]
|
||||
|
||||
if active_children:
|
||||
main.status = ChatGenerationTaskStatus.GENERATING.value
|
||||
main.pipeline_stage = active_children[0].pipeline_stage or ChatGenerationPipelineStage.QUEUED.value
|
||||
main.generated_at = None
|
||||
main.error_message = _build_summary(children)
|
||||
elif failed_children:
|
||||
main.status = ChatGenerationTaskStatus.FAILED.value
|
||||
main.pipeline_stage = (
|
||||
ChatGenerationPipelineStage.DOWNLOAD_FAILED.value
|
||||
if any(child.pipeline_stage == ChatGenerationPipelineStage.DOWNLOAD_FAILED.value for child in failed_children)
|
||||
else ChatGenerationPipelineStage.FAILED.value
|
||||
)
|
||||
main.generated_at = max(
|
||||
(child.generated_at for child in completed_children if child.generated_at),
|
||||
default=datetime.now(timezone.utc),
|
||||
)
|
||||
main.error_message = _build_summary(children)
|
||||
else:
|
||||
# 所有子任务真实生成结果均成功;资源是否软删除不改变生成历史终态。
|
||||
main.status = ChatGenerationTaskStatus.COMPLETED.value
|
||||
main.pipeline_stage = ChatGenerationPipelineStage.DONE.value
|
||||
main.generated_at = max(
|
||||
(child.generated_at for child in children if child.generated_at),
|
||||
default=main.generated_at or datetime.now(timezone.utc),
|
||||
)
|
||||
main.error_message = None
|
||||
|
||||
if main.gen_type == "video":
|
||||
main.credits_cost = round(sum(float(child.credits_cost or 0) for child in children), 2)
|
||||
main.text_credits_cost = round(sum(float(child.text_credits_cost or 0) for child in children), 2)
|
||||
main.text_tokens_used = sum(int(child.text_tokens_used or 0) for child in children)
|
||||
main.image_tokens_used = sum(int(child.image_tokens_used or 0) for child in children)
|
||||
main.video_tokens_used = sum(int(child.video_tokens_used or 0) for child in children)
|
||||
main.retry_count = sum(int(child.retry_count or 0) for child in children)
|
||||
main.poll_count = sum(int(child.poll_count or 0) for child in children)
|
||||
|
||||
await db.flush()
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="MAIN_STATUS_AGGREGATED",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=main.user_id,
|
||||
group_id=main.id,
|
||||
task_id=main.id,
|
||||
detail={
|
||||
"before_status": previous_status,
|
||||
"before_stage": previous_stage,
|
||||
"after_status": main.status,
|
||||
"after_stage": main.pipeline_stage,
|
||||
"summary": _build_summary(children),
|
||||
},
|
||||
)
|
||||
return main
|
||||
|
||||
|
||||
async def aggregate_parent_for_child(db: AsyncSession, child: ChatGenerationTask | None) -> ChatGenerationTask | None:
|
||||
if not child or child.generation_mode != GenerationMode.CHATAPI_CHILD.value or not child.parent_task_id:
|
||||
return None
|
||||
# 项目关闭了 autoflush,先显式 flush 子任务的终态,确保聚合查询读取到本事务最新状态。
|
||||
await db.flush()
|
||||
return await aggregate_main_task_status(db, parent_task_id=str(child.parent_task_id))
|
||||
|
||||
|
||||
async def soft_delete_child_tasks_batch(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
child_task_ids: Sequence[str] | Iterable[str],
|
||||
user_id: str,
|
||||
deleted_at: datetime | None = None,
|
||||
require_completed: bool = False,
|
||||
) -> int:
|
||||
ids = list(dict.fromkeys(str(item) for item in child_task_ids if item))
|
||||
if not ids:
|
||||
return 0
|
||||
deleted_at = deleted_at or datetime.now(timezone.utc)
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id.in_(ids),
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
|
||||
)
|
||||
.with_for_update()
|
||||
)
|
||||
children = list(result.scalars().all())
|
||||
found_ids = {str(child.id) for child in children}
|
||||
missing_ids = [item for item in ids if item not in found_ids]
|
||||
if missing_ids:
|
||||
raise HTTPException(status_code=404, detail=f"子任务不存在: {','.join(missing_ids)}")
|
||||
|
||||
active_children = [child for child in children if child.deleted_at is None]
|
||||
running_ids = [child.id for child in active_children if is_task_active(child)]
|
||||
if running_ids:
|
||||
raise HTTPException(status_code=400, detail=f"仍有 {len(running_ids)} 个子任务生成中,暂不能删除")
|
||||
if require_completed:
|
||||
invalid_ids = [
|
||||
child.id for child in active_children
|
||||
if child.status != ChatGenerationTaskStatus.COMPLETED.value or child.generated_at is None
|
||||
]
|
||||
if invalid_ids:
|
||||
raise HTTPException(status_code=409, detail=f"只有生成完成的资源才能从素材云删除: {','.join(invalid_ids)}")
|
||||
|
||||
source_ids = [str(child.id) for child in active_children]
|
||||
freed_size = await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=source_ids,
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
parent_ids = list(dict.fromkeys(str(child.parent_task_id) for child in active_children if child.parent_task_id))
|
||||
for child in active_children:
|
||||
child.deleted_at = deleted_at
|
||||
await db.flush()
|
||||
for parent_id in parent_ids:
|
||||
await aggregate_main_task_status(db, parent_task_id=parent_id)
|
||||
return int(freed_size or 0)
|
||||
|
||||
|
||||
async def soft_delete_child_task(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
child_task_id: str,
|
||||
user_id: str,
|
||||
deleted_at: datetime | None = None,
|
||||
) -> int:
|
||||
deleted_at = deleted_at or datetime.now(timezone.utc)
|
||||
detail_result = await db.execute(
|
||||
select(
|
||||
ChatGenerationTask.parent_task_id,
|
||||
ChatGenerationTask.generation_index,
|
||||
).where(
|
||||
ChatGenerationTask.id == child_task_id,
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
|
||||
).limit(1)
|
||||
)
|
||||
detail = detail_result.one_or_none()
|
||||
if not detail:
|
||||
raise HTTPException(status_code=404, detail="子任务不存在")
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_RESOURCE_DELETE_START",
|
||||
event_status="started",
|
||||
source="service",
|
||||
user_id=user_id,
|
||||
group_id=detail.parent_task_id,
|
||||
task_id=child_task_id,
|
||||
detail={"generation_index": detail.generation_index},
|
||||
)
|
||||
freed_size = await soft_delete_child_tasks_batch(
|
||||
db,
|
||||
child_task_ids=[child_task_id],
|
||||
user_id=user_id,
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_RESOURCE_DELETE_SUCCESS",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=user_id,
|
||||
group_id=detail.parent_task_id,
|
||||
task_id=child_task_id,
|
||||
detail={"generation_index": detail.generation_index, "freed_size_bytes": freed_size},
|
||||
)
|
||||
return freed_size
|
||||
|
||||
|
||||
async def soft_delete_top_level_task_group(
|
||||
db: AsyncSession,
|
||||
*,
|
||||
task_id: str,
|
||||
user_id: str,
|
||||
deleted_at: datetime | None = None,
|
||||
) -> int:
|
||||
deleted_at = deleted_at or datetime.now(timezone.utc)
|
||||
result = await db.execute(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.user_id == user_id,
|
||||
ChatGenerationTask.generation_mode.in_(
|
||||
[GenerationMode.CHATAPI_ASYNC.value, GenerationMode.CHATAPI_MAIN.value]
|
||||
),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
.limit(1)
|
||||
)
|
||||
task = result.scalar_one_or_none()
|
||||
if not task:
|
||||
raise HTTPException(status_code=404, detail="任务不存在")
|
||||
|
||||
if task.generation_mode == GenerationMode.CHATAPI_ASYNC.value:
|
||||
if is_task_active(task):
|
||||
raise HTTPException(status_code=400, detail="当前任务正在生成中,暂不能删除")
|
||||
freed_size = await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=[task.id],
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
task.deleted_at = deleted_at
|
||||
await db.flush()
|
||||
return int(freed_size or 0)
|
||||
|
||||
children_map = await load_children_map(db, [task.id], include_deleted=True)
|
||||
children = children_map.get(task.id, [])
|
||||
active_ids = [child.id for child in children if is_task_active(child)]
|
||||
if active_ids:
|
||||
raise HTTPException(status_code=400, detail=f"任务组仍有 {len(active_ids)} 个子任务生成中,暂不能删除")
|
||||
|
||||
active_children = [child for child in children if child.deleted_at is None]
|
||||
child_ids = [child.id for child in active_children]
|
||||
freed_size = await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=child_ids,
|
||||
deleted_at=deleted_at,
|
||||
)
|
||||
for child in active_children:
|
||||
child.deleted_at = deleted_at
|
||||
task.deleted_at = deleted_at
|
||||
await db.flush()
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="BATCH_GROUP_DELETE_SUCCESS",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=user_id,
|
||||
group_id=task.id,
|
||||
task_id=task.id,
|
||||
detail={
|
||||
"child_task_ids": child_ids,
|
||||
"freed_size_bytes": int(freed_size or 0),
|
||||
"physical_files_deleted": False,
|
||||
},
|
||||
)
|
||||
return int(freed_size or 0)
|
||||
+8
-4
@@ -470,6 +470,7 @@ async def charge_generation_media_by_params(
|
||||
source_step_id: str | None = None,
|
||||
source_step_code: str | None = None,
|
||||
billing_scene: str | None = None,
|
||||
quantity: int = 1,
|
||||
) -> BillingSummary:
|
||||
"""图片/视频媒体生成扣费。
|
||||
|
||||
@@ -478,6 +479,7 @@ async def charge_generation_media_by_params(
|
||||
"""
|
||||
project_name = project_name or "AI生成任务"
|
||||
gen_type = (gen_type or "").lower().strip()
|
||||
quantity = max(1, int(quantity or 1))
|
||||
attempt_no = attempt_no or await get_next_credit_attempt_no(
|
||||
db,
|
||||
owner_type=owner_type,
|
||||
@@ -508,13 +510,14 @@ async def charge_generation_media_by_params(
|
||||
|
||||
if gen_type == "image":
|
||||
size = image_size or "2K"
|
||||
amount = await calc_image_credits(db, size, engine_id=engine_id)
|
||||
unit_amount = await calc_image_credits(db, size, engine_id=engine_id)
|
||||
amount = round(unit_amount * quantity, 2)
|
||||
items.append(
|
||||
await deduct_credits_locked_once(
|
||||
db,
|
||||
user_id=user_id,
|
||||
amount=amount,
|
||||
description=f"{description_prefix}图片生成",
|
||||
description=f"{description_prefix}图片生成" + (f"×{quantity}" if quantity > 1 else ""),
|
||||
related_id=record_id,
|
||||
charge_key=CHARGE_MEDIA,
|
||||
biz_key=biz_key,
|
||||
@@ -523,17 +526,18 @@ async def charge_generation_media_by_params(
|
||||
)
|
||||
)
|
||||
elif gen_type == "video":
|
||||
amount = await calc_video_credits(
|
||||
unit_amount = await calc_video_credits(
|
||||
db, duration or 5, resolution or "720p",
|
||||
engine_id=engine_id,
|
||||
input_video_duration=input_video_duration,
|
||||
)
|
||||
amount = round(unit_amount * quantity, 2)
|
||||
items.append(
|
||||
await deduct_credits_locked_once(
|
||||
db,
|
||||
user_id=user_id,
|
||||
amount=amount,
|
||||
description=f"{description_prefix}视频生成",
|
||||
description=f"{description_prefix}视频生成" + (f"×{quantity}" if quantity > 1 else ""),
|
||||
related_id=record_id,
|
||||
charge_key=CHARGE_MEDIA,
|
||||
biz_key=biz_key,
|
||||
+51
-11
@@ -1,6 +1,5 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from datetime import datetime, timezone
|
||||
from typing import Iterable, Sequence
|
||||
|
||||
@@ -24,12 +23,13 @@ from app.models.module_generation_step import ModuleGenerationStep
|
||||
from app.models.shot_replicate_segment import ShotReplicateSegment
|
||||
from app.models.user import User
|
||||
from app.schemas.generation_ai import GenerationAIHistoryBatchDeleteOut
|
||||
from app.services.generation.ai.task_group_service import soft_delete_child_tasks_batch
|
||||
from app.services.module_generation_flow_base_service import is_active_chat_generation_task
|
||||
from app.services.module_generation_log_service import log_module_event_file
|
||||
from app.services.operation_log_service import log_operation_event
|
||||
# from app.services.operation_log import log_operation
|
||||
from app.services.resource_accounting_service import (
|
||||
SOURCE_MODEL_CHAT_TASK,
|
||||
SOURCE_MODEL_GENERATION_RECORD,
|
||||
SOURCE_MODEL_SHOT_SEGMENT,
|
||||
soft_delete_generation_record_resources,
|
||||
soft_delete_resources_by_source,
|
||||
@@ -238,14 +238,19 @@ async def _delete_chat_tasks(
|
||||
.where(
|
||||
ChatGenerationTask.id.in_(ids),
|
||||
ChatGenerationTask.user_id == current_user.id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_ASYNC.value,
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
ChatGenerationTask.generation_mode.in_([
|
||||
GenerationMode.CHATAPI_ASYNC.value,
|
||||
GenerationMode.CHATAPI_CHILD.value,
|
||||
]),
|
||||
)
|
||||
.with_for_update()
|
||||
)
|
||||
tasks = list(result.scalars().all())
|
||||
_raise_missing_if_any(ids=ids, found_ids=[task.id for task in tasks], message="AI 创作记录不存在或已删除")
|
||||
|
||||
already_deleted_ids = [str(task.id) for task in tasks if task.deleted_at is not None]
|
||||
_raise_invalid_if_any(invalid_ids=already_deleted_ids, message="AI 创作记录不存在或已删除", status_code=404)
|
||||
|
||||
invalid_ids = [
|
||||
task.id
|
||||
for task in tasks
|
||||
@@ -253,14 +258,49 @@ async def _delete_chat_tasks(
|
||||
]
|
||||
_raise_invalid_if_any(invalid_ids=invalid_ids, message="AI 创作记录只有生成完成后才能删除")
|
||||
|
||||
freed_size = await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=[task.id for task in tasks],
|
||||
deleted_at=deleted_at,
|
||||
async_ids = [str(task.id) for task in tasks if task.generation_mode == GenerationMode.CHATAPI_ASYNC.value]
|
||||
child_ids = [str(task.id) for task in tasks if task.generation_mode == GenerationMode.CHATAPI_CHILD.value]
|
||||
|
||||
freed_size = 0
|
||||
if async_ids:
|
||||
freed_size += int(await soft_delete_resources_by_source(
|
||||
db,
|
||||
source_model=SOURCE_MODEL_CHAT_TASK,
|
||||
source_ids=async_ids,
|
||||
deleted_at=deleted_at,
|
||||
) or 0)
|
||||
async_id_set = set(async_ids)
|
||||
for task in tasks:
|
||||
if str(task.id) in async_id_set:
|
||||
task.deleted_at = deleted_at
|
||||
|
||||
if child_ids:
|
||||
freed_size += await soft_delete_child_tasks_batch(
|
||||
db,
|
||||
child_task_ids=child_ids,
|
||||
user_id=current_user.id,
|
||||
deleted_at=deleted_at,
|
||||
require_completed=True,
|
||||
)
|
||||
|
||||
parent_task_ids = list(dict.fromkeys(
|
||||
str(task.parent_task_id) for task in tasks if task.parent_task_id
|
||||
))
|
||||
log_operation_event(
|
||||
domain="generation_ai_batch",
|
||||
event_type="CHILD_RESOURCE_DELETE_SUCCESS",
|
||||
event_status="success",
|
||||
source="service",
|
||||
user_id=current_user.id,
|
||||
group_id=parent_task_ids[0] if len(parent_task_ids) == 1 else None,
|
||||
detail={
|
||||
"batch": True,
|
||||
"task_ids": [str(task.id) for task in tasks],
|
||||
"parent_task_ids": parent_task_ids,
|
||||
"freed_size_bytes": int(freed_size or 0),
|
||||
"physical_files_deleted": False,
|
||||
},
|
||||
)
|
||||
for task in tasks:
|
||||
task.deleted_at = deleted_at
|
||||
|
||||
return _build_out(
|
||||
source=source,
|
||||
+1
-1
@@ -14,7 +14,7 @@ from app.config import settings
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.models.model_config import ModelConfig
|
||||
from app.models.token_usage import TokenUsage
|
||||
from app.services.generation_log_service import log_provider_call
|
||||
from app.services.generation.log_service import log_provider_call
|
||||
from app.services.provider_limit import provider_limit
|
||||
from app.utils.id_gen import generate_id
|
||||
|
||||
+92
-38
@@ -13,10 +13,11 @@ from app.config import settings
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.models.video_engine import VideoEngine
|
||||
from app.services.generation_log_service import log_provider_call
|
||||
from app.services.image_gen import poll_image_task_status, submit_image_task
|
||||
from app.services.generation.log_service import log_provider_call
|
||||
from app.services.image_gen import ImageProviderError, poll_image_task_status, submit_image_task
|
||||
from app.services.provider_limit import provider_limit
|
||||
from app.services.video_gen import poll_task_status, submit_video_task
|
||||
from app.types.generation.provider import ImageProviderBatchResult
|
||||
|
||||
|
||||
def _loads(data: str | None) -> dict:
|
||||
@@ -29,8 +30,17 @@ def _loads(data: str | None) -> dict:
|
||||
return {}
|
||||
|
||||
|
||||
def _try_json(value: Any) -> Any:
|
||||
if not isinstance(value, str):
|
||||
return value
|
||||
try:
|
||||
return json.loads(value)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
async def get_runtime_engine(db: AsyncSession, task: ChatGenerationTask) -> Any:
|
||||
"""Use frozen snapshot for historical params, current DB row only for secret api_key."""
|
||||
"""使用任务快照冻结历史参数,只从当前引擎记录读取密钥。"""
|
||||
snapshot = _loads(task.engine_snapshot_json)
|
||||
if not task.engine_id:
|
||||
raise ValueError("缺少 engine_id")
|
||||
@@ -51,6 +61,31 @@ async def get_runtime_engine(db: AsyncSession, task: ChatGenerationTask) -> Any:
|
||||
generate_url=snapshot.get("generate_url") or getattr(engine, "generate_url", ""),
|
||||
query_url=snapshot.get("query_url") or getattr(engine, "query_url", ""),
|
||||
default_size=snapshot.get("default_size") or getattr(engine, "default_size", "2K"),
|
||||
multi_generation_enabled=bool(
|
||||
snapshot.get("multi_generation_enabled")
|
||||
if snapshot.get("multi_generation_enabled") is not None
|
||||
else getattr(engine, "multi_generation_enabled", False)
|
||||
),
|
||||
max_generation_count=int(
|
||||
snapshot.get("max_generation_count")
|
||||
or getattr(engine, "max_generation_count", 1)
|
||||
or 1
|
||||
),
|
||||
multi_image_max_images=int(
|
||||
snapshot.get("multi_image_max_images")
|
||||
or getattr(engine, "multi_image_max_images", 15)
|
||||
or 15
|
||||
),
|
||||
max_reference_image_count=int(
|
||||
snapshot.get("max_reference_image_count")
|
||||
if snapshot.get("max_reference_image_count") is not None
|
||||
else getattr(engine, "max_reference_image_count", 14)
|
||||
),
|
||||
output_format=(
|
||||
snapshot.get("output_format")
|
||||
if snapshot.get("output_format") is not None
|
||||
else getattr(engine, "output_format", "")
|
||||
) or "",
|
||||
)
|
||||
|
||||
|
||||
@@ -58,22 +93,16 @@ async def create_provider_task(db: AsyncSession, task: ChatGenerationTask) -> di
|
||||
if task.gen_type == "video":
|
||||
return await _create_video_task(db, task)
|
||||
if task.gen_type == "image":
|
||||
return await _create_image_sync_task(db, task)
|
||||
return await create_image_sync_result(db, task)
|
||||
raise ValueError(f"不支持的生成类型: {task.gen_type}")
|
||||
|
||||
|
||||
async def _create_video_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
|
||||
"""Create video provider task through the original Ark SDK async task API."""
|
||||
engine = await get_runtime_engine(db, task)
|
||||
started = time.perf_counter()
|
||||
async with provider_limit("ark_video_create", settings.ARK_VIDEO_CREATE_MAX_CONCURRENCY):
|
||||
try:
|
||||
provider_task_id = await submit_video_task(
|
||||
db,
|
||||
engine,
|
||||
task,
|
||||
include_media_references=True,
|
||||
)
|
||||
provider_task_id = await submit_video_task(None, engine, task, include_media_references=True)
|
||||
response = {"task_id": provider_task_id}
|
||||
await log_provider_call(
|
||||
task,
|
||||
@@ -101,65 +130,90 @@ async def _create_video_task(db: AsyncSession, task: ChatGenerationTask) -> dict
|
||||
raise
|
||||
|
||||
|
||||
async def _create_image_sync_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
|
||||
"""Run the original synchronous image generation SDK under Celery control.
|
||||
|
||||
The legacy image SDK returns a final remote image URL immediately. We do
|
||||
NOT use image_generation.tasks.create here, so image generation stays aligned
|
||||
with the old working flow while no longer blocking the FastAPI request.
|
||||
"""
|
||||
async def create_image_sync_batch_result(
|
||||
db: AsyncSession,
|
||||
task: ChatGenerationTask,
|
||||
*,
|
||||
generation_count: int,
|
||||
) -> ImageProviderBatchResult:
|
||||
engine = await get_runtime_engine(db, task)
|
||||
return await create_image_sync_batch_result_with_engine(
|
||||
task,
|
||||
engine,
|
||||
generation_count=generation_count,
|
||||
)
|
||||
|
||||
|
||||
async def create_image_sync_batch_result_with_engine(
|
||||
task: ChatGenerationTask,
|
||||
engine: Any,
|
||||
*,
|
||||
generation_count: int,
|
||||
) -> ImageProviderBatchResult:
|
||||
"""执行一次同步图片请求。
|
||||
|
||||
generation_count > 1 时是一次组图 API 调用;失败后绝不退化为多次单图调用。
|
||||
"""
|
||||
count = max(1, int(generation_count or 1))
|
||||
started = time.perf_counter()
|
||||
api_type = "image_sync_batch_create" if count > 1 else "image_sync_create"
|
||||
async with provider_limit("ark_image_sync_create", settings.ARK_IMAGE_CREATE_MAX_CONCURRENCY):
|
||||
try:
|
||||
result = await asyncio.to_thread(
|
||||
submit_image_task,
|
||||
db,
|
||||
None,
|
||||
engine,
|
||||
task,
|
||||
include_media_references=True,
|
||||
generation_count=count,
|
||||
)
|
||||
if result.get("error"):
|
||||
raise RuntimeError(result.get("error"))
|
||||
response_data = _try_json(result.get("response_data")) or result
|
||||
response_data = result.get("response_data") or result
|
||||
await log_provider_call(
|
||||
task,
|
||||
provider=engine.provider,
|
||||
api_type="image_sync_create",
|
||||
api_type=api_type,
|
||||
model=engine.model_name,
|
||||
engine_id=task.engine_id,
|
||||
status="success",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
provider_task_id=None,
|
||||
response_data=response_data,
|
||||
total_tokens=int(result.get("image_tokens", 0) or 0),
|
||||
)
|
||||
return {
|
||||
"task_id": None,
|
||||
"remote_result_url": result.get("image_url"),
|
||||
"image_tokens": result.get("image_tokens", 0) or 0,
|
||||
"response_data": response_data,
|
||||
}
|
||||
return result
|
||||
except Exception as exc:
|
||||
error_message = exc.safe_message if isinstance(exc, ImageProviderError) else str(exc)
|
||||
await log_provider_call(
|
||||
task,
|
||||
provider=engine.provider,
|
||||
api_type="image_sync_create",
|
||||
api_type=api_type,
|
||||
model=engine.model_name,
|
||||
engine_id=task.engine_id,
|
||||
status="failed",
|
||||
latency_ms=int((time.perf_counter() - started) * 1000),
|
||||
error_message=str(exc),
|
||||
error_message=error_message,
|
||||
response_data=exc.as_dict() if isinstance(exc, ImageProviderError) else None,
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
def _try_json(text: Any) -> Any:
|
||||
if not isinstance(text, str):
|
||||
return text
|
||||
try:
|
||||
return json.loads(text)
|
||||
except Exception:
|
||||
return None
|
||||
async def create_image_sync_result(db: AsyncSession, task: ChatGenerationTask) -> dict:
|
||||
result = await create_image_sync_batch_result(db, task, generation_count=1)
|
||||
items = result.get("items") or []
|
||||
if len(items) != 1:
|
||||
raise RuntimeError(f"图片供应商单图返回数量异常,期望 1,实际 {len(items)}")
|
||||
item = items[0]
|
||||
if item.get("error_message"):
|
||||
raise RuntimeError(item.get("error_message") or "图片生成失败")
|
||||
image_url = item.get("remote_result_url")
|
||||
if not image_url:
|
||||
raise RuntimeError("图片供应商未返回有效图片地址")
|
||||
return {
|
||||
"task_id": None,
|
||||
"remote_result_url": image_url,
|
||||
"image_tokens": int(result.get("image_tokens", 0) or 0),
|
||||
"response_data": result.get("response_data") or {},
|
||||
}
|
||||
|
||||
|
||||
async def poll_provider_task(db: AsyncSession, task: ChatGenerationTask) -> dict:
|
||||
+127
-4
@@ -15,6 +15,7 @@ from app.enums.generation_task import (
|
||||
ChatGenerationPipelineStage,
|
||||
ChatGenerationTaskEventType,
|
||||
ChatGenerationTaskStatus,
|
||||
GenerationMode,
|
||||
GenerationType,
|
||||
)
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
@@ -25,10 +26,10 @@ from app.services.celery_download_recovery_service import (
|
||||
postpone_download_active_check,
|
||||
remove_download_active,
|
||||
)
|
||||
from app.services.generation_log_service import log_task_event
|
||||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation_poll_schedule_service import ensure_video_poll_fields, is_poll_not_due, is_video_generation_task
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation.log_service import log_task_event
|
||||
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.poll_schedule_service import ensure_video_poll_fields, is_poll_not_due, is_video_generation_task
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.redis_registry_service import (
|
||||
redis_get_due_registry_ids,
|
||||
redis_get_registry_payloads,
|
||||
@@ -341,6 +342,8 @@ async def _mark_timeout(
|
||||
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
)
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
await aggregate_parent_for_child(db, task)
|
||||
await db.commit()
|
||||
await _remove_poll_active(task.id)
|
||||
await log_task_event(
|
||||
@@ -367,6 +370,8 @@ async def _mark_failed(
|
||||
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
)
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
await aggregate_parent_for_child(db, task)
|
||||
await db.commit()
|
||||
await _remove_poll_active(task.id)
|
||||
await log_task_event(task, event_type=event_type, message=task.error_message, detail=detail)
|
||||
@@ -590,6 +595,99 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
|
||||
checked_ids: set[str] = set()
|
||||
results: dict[str, int] = {}
|
||||
|
||||
# 图片多份主任务只补投递,不在恢复服务内直接调用供应商。
|
||||
# 有效 claim 未过期时必须跳过,防止与正在运行的 Worker 重复调用组图 API。
|
||||
from app.tasks.generation_create_tasks import chatapi_create_generation_task
|
||||
image_main_cursor: str | None = None
|
||||
image_main_batch_size = max(1, int(settings.GENERATION_RECOVERY_BATCH_SIZE or 100))
|
||||
while True:
|
||||
image_main_query = select(ChatGenerationTask).where(
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
ChatGenerationTask.gen_type == GenerationType.IMAGE.value,
|
||||
ChatGenerationTask.status == ChatGenerationTaskStatus.GENERATING.value,
|
||||
ChatGenerationTask.pipeline_stage.in_([
|
||||
ChatGenerationPipelineStage.QUEUED.value,
|
||||
ChatGenerationPipelineStage.PREPARING.value,
|
||||
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
|
||||
]),
|
||||
)
|
||||
if image_main_cursor:
|
||||
image_main_query = image_main_query.where(ChatGenerationTask.id > image_main_cursor)
|
||||
image_main_result = await db.execute(
|
||||
image_main_query.order_by(ChatGenerationTask.id.asc())
|
||||
.limit(image_main_batch_size)
|
||||
.with_for_update()
|
||||
)
|
||||
image_mains = list(image_main_result.scalars().all())
|
||||
if not image_mains:
|
||||
break
|
||||
|
||||
for main in image_mains:
|
||||
main_id = str(main.id)
|
||||
image_main_cursor = main_id
|
||||
checked_ids.add(main_id)
|
||||
|
||||
child_result = await db.execute(
|
||||
select(ChatGenerationTask.id)
|
||||
.where(
|
||||
ChatGenerationTask.parent_task_id == main_id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
|
||||
)
|
||||
.limit(1)
|
||||
)
|
||||
if child_result.scalar_one_or_none() is not None:
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await db.commit()
|
||||
results["image_main_already_split"] = results.get("image_main_already_split", 0) + 1
|
||||
continue
|
||||
|
||||
now = _now()
|
||||
lease_until = ensure_aware_utc(main.provider_create_lease_until)
|
||||
lease_alive = bool(main.provider_create_claim_token and lease_until and lease_until > now)
|
||||
if lease_alive:
|
||||
await db.commit()
|
||||
results["image_main_claim_alive"] = results.get("image_main_claim_alive", 0) + 1
|
||||
continue
|
||||
|
||||
if _is_expired(main.deadline_at, now):
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=main,
|
||||
error_message="图片批量生成任务超时",
|
||||
pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
)
|
||||
await db.commit()
|
||||
results["image_main_timeout"] = results.get("image_main_timeout", 0) + 1
|
||||
continue
|
||||
|
||||
if main.provider_create_claim_token or main.provider_create_lease_until:
|
||||
main.provider_create_claim_token = None
|
||||
main.provider_create_lease_until = None
|
||||
main.pipeline_stage = ChatGenerationPipelineStage.QUEUED.value
|
||||
await log_task_event(
|
||||
main,
|
||||
event_type=ChatGenerationTaskEventType.IMAGE_MAIN_CLAIM_EXPIRED.value,
|
||||
message="图片主任务供应商执行租约已过期,恢复重新投递",
|
||||
)
|
||||
await db.commit()
|
||||
try:
|
||||
chatapi_create_generation_task.apply_async(
|
||||
args=[main_id],
|
||||
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
|
||||
countdown=0,
|
||||
)
|
||||
results["recover_image_main_create"] = results.get("recover_image_main_create", 0) + 1
|
||||
except Exception as exc:
|
||||
logger.exception("恢复投递图片主任务失败 task_id=%s: %s", main_id, exc)
|
||||
results["recover_image_main_enqueue_failed"] = results.get("recover_image_main_enqueue_failed", 0) + 1
|
||||
|
||||
if len(image_mains) < image_main_batch_size:
|
||||
break
|
||||
|
||||
due_poll_ids = await redis_get_due_registry_ids(
|
||||
zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY,
|
||||
limit=int(settings.POLL_RECOVERY_BATCH_SIZE or settings.GENERATION_RECOVERY_BATCH_SIZE or 100),
|
||||
@@ -673,6 +771,31 @@ async def recover_generation_tasks_once(db: AsyncSession) -> dict[str, Any]:
|
||||
if len(tasks) < batch_size or progressed_this_round <= 0:
|
||||
break
|
||||
|
||||
# 子任务可能在 worker 中断前已进入终态但主任务尚未汇总,按稳定游标完整重算全部主任务。
|
||||
from app.services.generation.ai.task_group_service import aggregate_main_task_status
|
||||
reconciled = 0
|
||||
main_cursor: str | None = None
|
||||
while True:
|
||||
main_query = select(ChatGenerationTask.id).where(
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_MAIN.value,
|
||||
)
|
||||
if main_cursor:
|
||||
main_query = main_query.where(ChatGenerationTask.id > main_cursor)
|
||||
main_result = await db.execute(main_query.order_by(ChatGenerationTask.id.asc()).limit(batch_size))
|
||||
parent_ids = list(main_result.scalars().all())
|
||||
if not parent_ids:
|
||||
break
|
||||
for parent_task_id in parent_ids:
|
||||
main_cursor = str(parent_task_id)
|
||||
await aggregate_main_task_status(db, parent_task_id=str(parent_task_id))
|
||||
await db.commit()
|
||||
reconciled += 1
|
||||
if len(parent_ids) < batch_size:
|
||||
break
|
||||
if reconciled:
|
||||
results["reconcile_main"] = reconciled
|
||||
|
||||
return {
|
||||
"checked": len(checked_ids),
|
||||
"db_checked": total_db_checked,
|
||||
+2
-4
@@ -1,7 +1,5 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Iterable
|
||||
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
@@ -11,7 +9,7 @@ from app.models.credit_record import CreditRecord
|
||||
from app.models.generation_record import GenerationRecord
|
||||
from app.services.credits import refund_credits
|
||||
from app.services.credit_record_meta_service import build_refund_meta_from_charge
|
||||
from app.services.generation_billing_service import (
|
||||
from app.services.generation.billing_service import (
|
||||
CHARGE_MEDIA,
|
||||
OWNER_CHAT_GENERATION_TASK,
|
||||
OWNER_GENERATION_RECORD,
|
||||
@@ -182,7 +180,7 @@ async def mark_chat_generation_task_failed_and_refund_once(
|
||||
select(ChatGenerationTask)
|
||||
.where(
|
||||
ChatGenerationTask.id == task_id,
|
||||
ChatGenerationTask.generation_mode.in_(["chatapi_async", "hot_opening_replicate", "shot_replicate"]),
|
||||
ChatGenerationTask.generation_mode.in_(["chatapi_async", "chatapi_main", "chatapi_child", "hot_opening_replicate", "shot_replicate"]),
|
||||
ChatGenerationTask.deleted_at.is_(None),
|
||||
)
|
||||
.with_for_update()
|
||||
+10
-8
@@ -11,21 +11,21 @@ from app.config import settings
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.models.user import User
|
||||
from app.schemas.generation_ai import GenerationAIReference, GenerationAITaskCreate
|
||||
from app.services.generation_ai_service import (
|
||||
from app.services.generation.ai.engine_service import (
|
||||
IMAGE_DEFAULT_PROPORTION,
|
||||
IMAGE_DEFAULT_PX,
|
||||
IMAGE_DEFAULT_SIZE,
|
||||
VIDEO_DEFAULT_RATIO,
|
||||
VIDEO_DEFAULT_RESOLUTION,
|
||||
_build_image_snapshot,
|
||||
_build_video_snapshot,
|
||||
_get_image_engine,
|
||||
_get_video_engine,
|
||||
_image_supported_sizes,
|
||||
_parse_list,
|
||||
build_image_snapshot as _build_image_snapshot,
|
||||
build_video_snapshot as _build_video_snapshot,
|
||||
get_image_engine as _get_image_engine,
|
||||
get_video_engine as _get_video_engine,
|
||||
image_supported_sizes as _image_supported_sizes,
|
||||
normalize_px,
|
||||
parse_json_list as _parse_list,
|
||||
)
|
||||
from app.services.generation_billing_service import OWNER_CHAT_GENERATION_TASK, charge_generation_media_by_params
|
||||
from app.services.generation.billing_service import OWNER_CHAT_GENERATION_TASK, charge_generation_media_by_params
|
||||
from app.services.resource_capacity_service import assert_user_resource_capacity_available
|
||||
from app.services.private_portrait.reference_resolver import resolve_private_portrait_references
|
||||
from app.utils.id_gen import generate_id
|
||||
@@ -128,6 +128,7 @@ async def create_chat_generation_task_for_module(
|
||||
billing_scene=billing_scene,
|
||||
)
|
||||
snapshot = _build_image_snapshot(engine, size, proportion, px)
|
||||
snapshot["generation_count"] = 1
|
||||
task = ChatGenerationTask(
|
||||
id=task_id,
|
||||
user_id=current_user.id,
|
||||
@@ -182,6 +183,7 @@ async def create_chat_generation_task_for_module(
|
||||
billing_scene=billing_scene,
|
||||
)
|
||||
snapshot = _build_video_snapshot(engine, ratio, selected_resolution, selected_duration)
|
||||
snapshot["generation_count"] = 1
|
||||
task = ChatGenerationTask(
|
||||
id=task_id,
|
||||
user_id=current_user.id,
|
||||
@@ -1,14 +1,12 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from copy import deepcopy
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import String, cast, func, or_, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm.attributes import flag_modified
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.common import ModuleEventTypeEnum, ModuleProjectStatusEnum, ModulePromptTypeEnum, ModuleStepStatusEnum
|
||||
@@ -35,20 +33,19 @@ from app.schemas.hot_opening_replicate import (
|
||||
HotOpeningVideoGenerationOut,
|
||||
HotOpeningVideoPromptSchemaUpdateRequest,
|
||||
)
|
||||
from app.services.generation_ai_service import (
|
||||
from app.services.generation.ai.engine_service import (
|
||||
VIDEO_DEFAULT_DURATION,
|
||||
VIDEO_DEFAULT_RATIO,
|
||||
VIDEO_DEFAULT_RESOLUTION,
|
||||
_get_video_engine,
|
||||
_parse_list,
|
||||
get_video_engine,
|
||||
parse_json_list,
|
||||
)
|
||||
from app.services.generation_billing_service import charge_module_prompt_usage
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation_task_factory_service import create_chat_generation_task_for_module
|
||||
from app.services.generation.billing_service import charge_module_prompt_usage
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation.task_factory_service import create_chat_generation_task_for_module
|
||||
from app.services.hot_opening_video_prompt_service import build_final_video_prompt, optimize_hot_opening_video_prompt, patch_video_prompt_schema_from_client
|
||||
from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event
|
||||
from app.services.llm import optimize_prompt
|
||||
from app.services.resource_accounting_service import soft_delete_chat_task_resources
|
||||
from app.services.module_generation_flow_base_service import (
|
||||
chat_tasks_by_id as _base_chat_tasks_by_id,
|
||||
create_module_step as _base_create_step,
|
||||
@@ -1021,10 +1018,10 @@ async def generate_image_from_prompt(
|
||||
|
||||
|
||||
async def _resolve_video_prompt_config(db: AsyncSession, req: HotOpeningGenerateVideoPromptRequest) -> dict[str, Any]:
|
||||
engine = await _get_video_engine(db, req.engine_id)
|
||||
supported_ratios = _parse_list(engine.supported_ratios, [])
|
||||
supported_resolutions = _parse_list(engine.supported_resolutions, [])
|
||||
supported_durations = _parse_list(engine.supported_durations, [])
|
||||
engine = await get_video_engine(db, req.engine_id)
|
||||
supported_ratios = parse_json_list(engine.supported_ratios, [])
|
||||
supported_resolutions = parse_json_list(engine.supported_resolutions, [])
|
||||
supported_durations = parse_json_list(engine.supported_durations, [])
|
||||
|
||||
default_ratio = getattr(settings, "HOT_OPENING_DEFAULT_VIDEO_RATIO", None) or VIDEO_DEFAULT_RATIO
|
||||
default_resolution = getattr(settings, "HOT_OPENING_DEFAULT_VIDEO_RESOLUTION", None) or VIDEO_DEFAULT_RESOLUTION
|
||||
|
||||
@@ -1,9 +1,8 @@
|
||||
import base64
|
||||
import json
|
||||
import logging
|
||||
import mimetypes
|
||||
import os
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
from sqlalchemy import select
|
||||
@@ -11,10 +10,16 @@ from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from volcenginesdkarkruntime import AsyncArk
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.generation_provider import (
|
||||
MULTI_IMAGE_PROMPT_TEMPLATE,
|
||||
ImageProviderErrorType,
|
||||
)
|
||||
from app.enums.private_portrait import PRIVATE_PORTRAIT_ASSET_URI_PREFIX
|
||||
from app.models.image_engine import ImageEngine
|
||||
from app.services.log_config import is_enabled, LOG_DIR, LOG_DATE_FORMAT, encrypt_data
|
||||
from app.services.generation_provider_types import (
|
||||
from app.services.log_config import LOG_DATE_FORMAT, LOG_DIR, encrypt_data, is_enabled
|
||||
from app.types.generation.provider import (
|
||||
ImageProviderBatchResult,
|
||||
ImageProviderItem,
|
||||
ProviderGenerationRecordLike,
|
||||
ProviderImageEngineLike,
|
||||
)
|
||||
@@ -22,8 +27,39 @@ from app.services.generation_provider_types import (
|
||||
logger = logging.getLogger("videogen")
|
||||
|
||||
|
||||
class ImageProviderError(RuntimeError):
|
||||
"""可被生成任务状态机安全收敛的图片供应商异常。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
message: str,
|
||||
*,
|
||||
error_type: ImageProviderErrorType = ImageProviderErrorType.UNKNOWN,
|
||||
error_code: str | None = None,
|
||||
retryable: bool = False,
|
||||
http_status: int | None = None,
|
||||
provider_request_id: str | None = None,
|
||||
):
|
||||
super().__init__(message)
|
||||
self.safe_message = message
|
||||
self.error_type = error_type
|
||||
self.error_code = error_code
|
||||
self.retryable = retryable
|
||||
self.http_status = http_status
|
||||
self.provider_request_id = provider_request_id
|
||||
|
||||
def as_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"error_type": self.error_type.value,
|
||||
"error_code": self.error_code,
|
||||
"message": self.safe_message,
|
||||
"retryable": self.retryable,
|
||||
"http_status": self.http_status,
|
||||
"provider_request_id": self.provider_request_id,
|
||||
}
|
||||
|
||||
|
||||
def _log_image_request(engine: ProviderImageEngineLike, record_id: str, request_data: dict):
|
||||
"""Log image generation request to log/AiModel/YYYY-MM-DD.log"""
|
||||
if not is_enabled():
|
||||
return
|
||||
try:
|
||||
@@ -41,14 +77,13 @@ def _log_image_request(engine: ProviderImageEngineLike, record_id: str, request_
|
||||
"request": request_encrypted,
|
||||
"request_length": len(request_str),
|
||||
}
|
||||
with open(log_file, "a", encoding="utf-8") as f:
|
||||
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
with open(log_file, "a", encoding="utf-8") as file:
|
||||
file.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _log_image_response(record_id: str, response_data: dict, error: str | None = None):
|
||||
"""Log image generation response to log/AiModel/YYYY-MM-DD.log"""
|
||||
if not is_enabled():
|
||||
return
|
||||
try:
|
||||
@@ -63,17 +98,13 @@ def _log_image_response(record_id: str, response_data: dict, error: str | None =
|
||||
"response": response_encrypted,
|
||||
"error": error,
|
||||
}
|
||||
with open(log_file, "a", encoding="utf-8") as f:
|
||||
f.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
with open(log_file, "a", encoding="utf-8") as file:
|
||||
file.write(json.dumps(entry, ensure_ascii=False) + "\n")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
async def get_active_image_engine(db: AsyncSession) -> ImageEngine:
|
||||
"""Get the active image engine with highest priority."""
|
||||
result = await db.execute(
|
||||
select(ImageEngine)
|
||||
.where(ImageEngine.is_active == True)
|
||||
@@ -103,105 +134,291 @@ def _resolve_url(url: str) -> str:
|
||||
return f"{settings.BASE_URL.rstrip('/')}/{url.lstrip('/')}"
|
||||
|
||||
|
||||
def _value(obj: Any, name: str, default: Any = None) -> Any:
|
||||
if obj is None:
|
||||
return default
|
||||
if isinstance(obj, dict):
|
||||
return obj.get(name, default)
|
||||
return getattr(obj, name, default)
|
||||
|
||||
|
||||
def _jsonable(value: Any) -> Any:
|
||||
if value is None or isinstance(value, (str, int, float, bool)):
|
||||
return value
|
||||
if isinstance(value, dict):
|
||||
return {str(key): _jsonable(item) for key, item in value.items()}
|
||||
if isinstance(value, (list, tuple)):
|
||||
return [_jsonable(item) for item in value]
|
||||
if hasattr(value, "model_dump"):
|
||||
try:
|
||||
return _jsonable(value.model_dump())
|
||||
except Exception:
|
||||
pass
|
||||
if hasattr(value, "to_dict"):
|
||||
try:
|
||||
return _jsonable(value.to_dict())
|
||||
except Exception:
|
||||
pass
|
||||
result: dict[str, Any] = {}
|
||||
for key in ("url", "b64_json", "size", "output_format", "error", "code", "message"):
|
||||
item = getattr(value, key, None)
|
||||
if item is not None:
|
||||
result[key] = _jsonable(item)
|
||||
return result or str(value)
|
||||
|
||||
|
||||
def _safe_text(value: Any, *, limit: int = 1000) -> str:
|
||||
text = str(value or "").strip()
|
||||
return text[:limit]
|
||||
|
||||
|
||||
def _classify_provider_exception(exc: Exception) -> ImageProviderError:
|
||||
if isinstance(exc, ImageProviderError):
|
||||
return exc
|
||||
if isinstance(exc, (httpx.TimeoutException, TimeoutError)):
|
||||
return ImageProviderError(
|
||||
"图片生成请求超时,请稍后重试",
|
||||
error_type=ImageProviderErrorType.TIMEOUT,
|
||||
retryable=True,
|
||||
)
|
||||
|
||||
status_code = getattr(exc, "status_code", None)
|
||||
request_id = getattr(exc, "request_id", None) or getattr(exc, "x_request_id", None)
|
||||
code = getattr(exc, "code", None)
|
||||
raw_message = _safe_text(getattr(exc, "message", None) or exc)
|
||||
lowered = raw_message.lower()
|
||||
|
||||
if status_code == 429 or "rate limit" in lowered or "限流" in raw_message:
|
||||
error_type = ImageProviderErrorType.RATE_LIMIT
|
||||
retryable = True
|
||||
message = "图片生成请求过于频繁,请稍后重试"
|
||||
elif status_code in {401, 403} or "api key" in lowered or "unauthorized" in lowered:
|
||||
error_type = ImageProviderErrorType.AUTH
|
||||
retryable = False
|
||||
message = "图片引擎鉴权失败,请联系管理员检查配置"
|
||||
elif status_code and int(status_code) >= 500:
|
||||
error_type = ImageProviderErrorType.PROVIDER_INTERNAL
|
||||
retryable = True
|
||||
message = "图片供应商服务异常,请稍后重试"
|
||||
elif "sequential_image_generation" in lowered or "not support" in lowered or "unsupported" in lowered:
|
||||
error_type = ImageProviderErrorType.CAPABILITY_MISMATCH
|
||||
retryable = False
|
||||
message = "图片引擎组图能力配置与供应商实际能力不匹配,请联系管理员"
|
||||
elif "content" in lowered and ("risk" in lowered or "moderation" in lowered or "policy" in lowered):
|
||||
error_type = ImageProviderErrorType.CONTENT_REJECTED
|
||||
retryable = False
|
||||
message = "图片内容未通过供应商审核,请调整提示词后重试"
|
||||
elif status_code and 400 <= int(status_code) < 500:
|
||||
error_type = ImageProviderErrorType.INVALID_REQUEST
|
||||
retryable = False
|
||||
message = "图片生成参数不被供应商支持,请联系管理员检查引擎配置"
|
||||
elif isinstance(exc, httpx.HTTPError):
|
||||
error_type = ImageProviderErrorType.NETWORK
|
||||
retryable = True
|
||||
message = "图片供应商网络连接异常,请稍后重试"
|
||||
else:
|
||||
error_type = ImageProviderErrorType.UNKNOWN
|
||||
retryable = False
|
||||
message = raw_message or "图片生成失败"
|
||||
|
||||
return ImageProviderError(
|
||||
message,
|
||||
error_type=error_type,
|
||||
error_code=_safe_text(code, limit=128) or None,
|
||||
retryable=retryable,
|
||||
http_status=int(status_code) if status_code is not None else None,
|
||||
provider_request_id=_safe_text(request_id, limit=128) or None,
|
||||
)
|
||||
|
||||
|
||||
def build_multi_image_provider_prompt(prompt: str, generation_count: int) -> str:
|
||||
base_prompt = (prompt or "").strip()
|
||||
if generation_count <= 1:
|
||||
return base_prompt
|
||||
suffix = MULTI_IMAGE_PROMPT_TEMPLATE.format(count=generation_count)
|
||||
return f"{base_prompt}\n\n{suffix}" if base_prompt else suffix
|
||||
|
||||
|
||||
def submit_image_task(
|
||||
db,
|
||||
engine: ProviderImageEngineLike,
|
||||
record: ProviderGenerationRecordLike,
|
||||
*,
|
||||
include_media_references: bool,
|
||||
) -> dict:
|
||||
"""Submit an image generation task via Ark SDK. Returns image_url."""
|
||||
from volcenginesdkarkruntime import Ark
|
||||
|
||||
client = Ark(
|
||||
base_url=engine.api_base,
|
||||
api_key=engine.api_key,
|
||||
timeout=300,
|
||||
)
|
||||
generation_count: int = 1,
|
||||
) -> ImageProviderBatchResult:
|
||||
"""通过 Ark 同步图片接口生成单图或单次组图。
|
||||
|
||||
prompt = record.optimized_prompt or record.original_prompt
|
||||
image_urls = []
|
||||
generation_count > 1 时只执行一次 sequential_auto 请求;任何失败都直接抛出,
|
||||
绝不退化为多次单图请求。
|
||||
"""
|
||||
from volcenginesdkarkruntime import Ark
|
||||
|
||||
count = max(1, int(generation_count or 1))
|
||||
multi_generation_enabled = bool(getattr(engine, "multi_generation_enabled", False))
|
||||
max_generation_count = max(1, min(5, int(getattr(engine, "max_generation_count", 1) or 1)))
|
||||
if count > 1 and not multi_generation_enabled:
|
||||
raise ImageProviderError(
|
||||
"当前图片引擎未开启多份生成",
|
||||
error_type=ImageProviderErrorType.CAPABILITY_MISMATCH,
|
||||
)
|
||||
if count > max_generation_count:
|
||||
raise ImageProviderError(
|
||||
f"当前图片引擎最多允许生成 {max_generation_count} 份",
|
||||
error_type=ImageProviderErrorType.CAPABILITY_MISMATCH,
|
||||
)
|
||||
|
||||
client = Ark(base_url=engine.api_base, api_key=engine.api_key, timeout=300)
|
||||
original_prompt = record.optimized_prompt or record.original_prompt
|
||||
provider_prompt = build_multi_image_provider_prompt(original_prompt, count)
|
||||
image_urls: list[str] = []
|
||||
|
||||
if include_media_references and record.media_references:
|
||||
try:
|
||||
refs = json.loads(record.media_references)
|
||||
for ref in refs:
|
||||
ref_type = ref.get("type")
|
||||
ref_url = ref.get("url", "")
|
||||
if ref_type == "image" and ref_url:
|
||||
resolved = _resolve_url(ref_url)
|
||||
image_urls.append(resolved)
|
||||
for ref in refs if isinstance(refs, list) else []:
|
||||
if (ref.get("type") or "").lower() == "image" and ref.get("url"):
|
||||
image_urls.append(_resolve_url(ref["url"]))
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
pass
|
||||
image_urls = []
|
||||
|
||||
request_payload = {
|
||||
request_log_payload: dict[str, Any] = {
|
||||
"model": engine.model_name,
|
||||
"prompt": prompt,
|
||||
"prompt": provider_prompt,
|
||||
"size": record.image_size or engine.default_size,
|
||||
"sequential_image_generation": "disabled",
|
||||
"output_format": "png",
|
||||
"response_format": "url",
|
||||
"watermark": False,
|
||||
"include_media_references": include_media_references,
|
||||
}
|
||||
|
||||
request_sdk_payload: dict[str, Any] = dict(request_log_payload)
|
||||
if image_urls:
|
||||
request_payload["image"] = image_urls
|
||||
request_log_payload["image"] = image_urls
|
||||
request_sdk_payload["image"] = image_urls
|
||||
output_format = (getattr(engine, "output_format", "") or "").lower().strip()
|
||||
if output_format:
|
||||
request_log_payload["output_format"] = output_format
|
||||
request_sdk_payload["output_format"] = output_format
|
||||
if count > 1:
|
||||
try:
|
||||
from volcenginesdkarkruntime.types.images import SequentialImageGenerationOptions
|
||||
except Exception:
|
||||
try:
|
||||
from volcenginesdkarkruntime.types.images.image_generate_params import (
|
||||
SequentialImageGenerationOptions,
|
||||
)
|
||||
except Exception as import_exc:
|
||||
raise ImageProviderError(
|
||||
"当前图片引擎运行依赖缺少组图参数对象,请升级火山 Ark SDK 后重试",
|
||||
error_type=ImageProviderErrorType.CAPABILITY_MISMATCH,
|
||||
) from import_exc
|
||||
|
||||
_log_image_request(engine, record.id, request_payload)
|
||||
request_log_payload["sequential_image_generation"] = "auto"
|
||||
request_log_payload["sequential_image_generation_options"] = {"max_images": count}
|
||||
request_log_payload["stream"] = False
|
||||
|
||||
request_sdk_payload["sequential_image_generation"] = "auto"
|
||||
request_sdk_payload["sequential_image_generation_options"] = SequentialImageGenerationOptions(
|
||||
max_images=count,
|
||||
)
|
||||
request_sdk_payload["stream"] = False
|
||||
|
||||
_log_image_request(engine, record.id, request_log_payload)
|
||||
|
||||
try:
|
||||
result = client.images.generate(
|
||||
model=engine.model_name,
|
||||
prompt=prompt,
|
||||
size=record.image_size or engine.default_size,
|
||||
output_format="png",
|
||||
response_format="url",
|
||||
watermark=False,
|
||||
image=image_urls if image_urls else None,
|
||||
)
|
||||
image_url = result.data[0].url
|
||||
|
||||
response_data = {
|
||||
"model": result.model,
|
||||
"created": result.created,
|
||||
"data": [{"url": item.url, "size": item.size} for item in result.data] if result.data else [],
|
||||
"usage": {
|
||||
"generated_images": result.usage.generated_images if hasattr(result.usage, 'generated_images') else 0,
|
||||
"output_tokens": result.usage.output_tokens if hasattr(result.usage, 'output_tokens') else 0,
|
||||
"total_tokens": result.usage.total_tokens if hasattr(result.usage, 'total_tokens') else 0,
|
||||
result = client.images.generate(**request_sdk_payload)
|
||||
top_error = _value(result, "error")
|
||||
if top_error:
|
||||
error_code = _value(top_error, "code")
|
||||
error_message = _value(top_error, "message") or str(top_error)
|
||||
raise ImageProviderError(
|
||||
_safe_text(error_message) or "图片供应商返回失败",
|
||||
error_type=ImageProviderErrorType.INVALID_REQUEST,
|
||||
error_code=_safe_text(error_code, limit=128) or None,
|
||||
)
|
||||
|
||||
raw_data = _value(result, "data", []) or []
|
||||
if not isinstance(raw_data, (list, tuple)):
|
||||
raise ImageProviderError(
|
||||
"图片供应商返回 data 结构异常",
|
||||
error_type=ImageProviderErrorType.INVALID_RESPONSE,
|
||||
)
|
||||
|
||||
items: list[ImageProviderItem] = []
|
||||
response_items: list[dict[str, Any]] = []
|
||||
for index, raw_item in enumerate(raw_data, start=1):
|
||||
item_error = _value(raw_item, "error")
|
||||
if item_error:
|
||||
error_code = _safe_text(_value(item_error, "code"), limit=128)
|
||||
error_message = _safe_text(_value(item_error, "message") or item_error)
|
||||
items.append({
|
||||
"generation_index": index,
|
||||
"error_code": error_code,
|
||||
"error_message": error_message or "单张图片生成失败",
|
||||
"response_data": _jsonable(raw_item),
|
||||
})
|
||||
response_items.append(_jsonable(raw_item))
|
||||
continue
|
||||
|
||||
url = _safe_text(_value(raw_item, "url"), limit=4000)
|
||||
b64_json = _safe_text(_value(raw_item, "b64_json"), limit=100) if not url else ""
|
||||
item: ImageProviderItem = {
|
||||
"generation_index": index,
|
||||
"remote_result_url": url,
|
||||
"size": _safe_text(_value(raw_item, "size"), limit=64),
|
||||
"output_format": _safe_text(_value(raw_item, "output_format"), limit=32),
|
||||
"response_data": _jsonable(raw_item),
|
||||
}
|
||||
if b64_json:
|
||||
item["b64_json"] = b64_json
|
||||
items.append(item)
|
||||
response_items.append(_jsonable(raw_item))
|
||||
|
||||
usage = _value(result, "usage")
|
||||
generated_images = int(_value(usage, "generated_images", 0) or 0)
|
||||
total_tokens = int(_value(usage, "total_tokens", 0) or 0)
|
||||
response_data = {
|
||||
"model": _value(result, "model", engine.model_name),
|
||||
"created": _value(result, "created"),
|
||||
"data": response_items,
|
||||
"usage": {
|
||||
"generated_images": generated_images,
|
||||
"input_images": int(_value(usage, "input_images", 0) or 0),
|
||||
"output_tokens": int(_value(usage, "output_tokens", 0) or 0),
|
||||
"total_tokens": total_tokens,
|
||||
},
|
||||
}
|
||||
except httpx.TimeoutException:
|
||||
error_msg = "图片生成超时,请稍后重试"
|
||||
logger.error(f"Image generation timeout for record {record.id}")
|
||||
_log_image_response(record.id, {}, error_msg)
|
||||
raise TimeoutError(error_msg)
|
||||
except Exception as e:
|
||||
error_msg = str(e)
|
||||
logger.error(f"Image generation failed for record {record.id}: {error_msg}")
|
||||
_log_image_response(record.id, {}, error_msg)
|
||||
raise
|
||||
_log_image_response(record.id, response_data)
|
||||
return {
|
||||
"items": items,
|
||||
"model": str(response_data["model"] or ""),
|
||||
"created": int(response_data["created"] or 0),
|
||||
"generated_images": generated_images,
|
||||
"image_tokens": total_tokens,
|
||||
"response_data": response_data,
|
||||
}
|
||||
except Exception as exc:
|
||||
provider_error = _classify_provider_exception(exc)
|
||||
logger.error(
|
||||
"Image generation failed for record %s: type=%s code=%s message=%s",
|
||||
record.id,
|
||||
provider_error.error_type.value,
|
||||
provider_error.error_code,
|
||||
provider_error.safe_message,
|
||||
)
|
||||
_log_image_response(record.id, provider_error.as_dict(), provider_error.safe_message)
|
||||
raise provider_error from exc
|
||||
finally:
|
||||
client.close()
|
||||
|
||||
return {
|
||||
"image_url": image_url,
|
||||
"image_tokens": getattr(result.usage, "total_tokens", 0),
|
||||
"response_data": json.dumps(response_data, ensure_ascii=False, default=str),
|
||||
"error": str(result.error) if result.error else "",
|
||||
}
|
||||
try:
|
||||
client.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
async def poll_image_task_status(engine: ImageEngine, task_id: str) -> dict:
|
||||
"""Query image task status via Ark SDK. Returns {status, image_url, response_data}."""
|
||||
client = AsyncArk(
|
||||
base_url=engine.api_base,
|
||||
api_key=engine.api_key,
|
||||
)
|
||||
|
||||
result = await client.image_generation.tasks.get(task_id=task_id)
|
||||
await client.close()
|
||||
client = AsyncArk(base_url=engine.api_base, api_key=engine.api_key)
|
||||
try:
|
||||
result = await client.image_generation.tasks.get(task_id=task_id)
|
||||
finally:
|
||||
await client.close()
|
||||
|
||||
response_dict = {
|
||||
"id": result.id,
|
||||
@@ -239,13 +456,11 @@ async def poll_image_task_status(engine: ImageEngine, task_id: str) -> dict:
|
||||
|
||||
|
||||
async def download_image(image_url: str, dest_path: str) -> str:
|
||||
"""Download image to local storage."""
|
||||
os.makedirs(os.path.dirname(dest_path), exist_ok=True)
|
||||
|
||||
async with httpx.AsyncClient(timeout=300) as client:
|
||||
async with client.stream("GET", image_url) as response:
|
||||
response.raise_for_status()
|
||||
with open(dest_path, "wb") as f:
|
||||
with open(dest_path, "wb") as file:
|
||||
async for chunk in response.aiter_bytes(chunk_size=8192):
|
||||
f.write(chunk)
|
||||
return dest_path
|
||||
file.write(chunk)
|
||||
return dest_path
|
||||
|
||||
@@ -4,7 +4,7 @@ from collections.abc import Iterable
|
||||
from datetime import datetime
|
||||
from typing import Any, TypedDict
|
||||
|
||||
from sqlalchemy import and_, func, or_, select
|
||||
from sqlalchemy import and_, case, func, or_, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from app.enums.generation_task import GenerationType
|
||||
@@ -12,7 +12,7 @@ from app.enums.recent_generation import (
|
||||
RECENT_GENERATION_ALL_MODULES,
|
||||
RECENT_GENERATION_CHAT_TASK_MODULES,
|
||||
RECENT_GENERATION_COMPLETED_STATUS,
|
||||
RECENT_GENERATION_MODULE_TO_TASK_MODE,
|
||||
RECENT_GENERATION_MODULE_TO_TASK_MODES,
|
||||
RECENT_GENERATION_TASK_MODE_VALUE_TO_MODULE,
|
||||
RecentGenerationModuleEnum,
|
||||
RecentGenerationResourceTypeEnum,
|
||||
@@ -175,11 +175,12 @@ async def _list_chat_task_recent_rows(
|
||||
modules: list[RecentGenerationModuleEnum],
|
||||
limit: int,
|
||||
) -> list[dict[str, Any]]:
|
||||
task_mode_values = [
|
||||
RECENT_GENERATION_MODULE_TO_TASK_MODE[module].value
|
||||
task_mode_values = list(dict.fromkeys(
|
||||
task_mode.value
|
||||
for module in modules
|
||||
if module in RECENT_GENERATION_CHAT_TASK_MODULES
|
||||
]
|
||||
for task_mode in RECENT_GENERATION_MODULE_TO_TASK_MODES[module]
|
||||
))
|
||||
if not task_mode_values:
|
||||
return []
|
||||
|
||||
@@ -189,10 +190,19 @@ async def _list_chat_task_recent_rows(
|
||||
ChatGenerationTask.created_at,
|
||||
)
|
||||
|
||||
module_partition_expr = case(
|
||||
(
|
||||
ChatGenerationTask.generation_mode.in_(["chatapi_async", "chatapi_child"]),
|
||||
RecentGenerationModuleEnum.CHAT_AI.value,
|
||||
),
|
||||
else_=ChatGenerationTask.generation_mode,
|
||||
)
|
||||
|
||||
ranked_subquery = (
|
||||
select(
|
||||
ChatGenerationTask.id.label("generation_id"),
|
||||
ChatGenerationTask.generation_mode.label("generation_mode"),
|
||||
module_partition_expr.label("module_key"),
|
||||
ChatGenerationTask.gen_type.label("gen_type"),
|
||||
ChatGenerationTask.image_url.label("image_url"),
|
||||
ChatGenerationTask.video_url.label("video_url"),
|
||||
@@ -200,7 +210,7 @@ async def _list_chat_task_recent_rows(
|
||||
generated_time_expr.label("generated_time"),
|
||||
func.row_number()
|
||||
.over(
|
||||
partition_by=ChatGenerationTask.generation_mode,
|
||||
partition_by=module_partition_expr,
|
||||
order_by=(generated_time_expr.desc(), ChatGenerationTask.created_at.desc()),
|
||||
)
|
||||
.label("row_num"),
|
||||
@@ -218,7 +228,7 @@ async def _list_chat_task_recent_rows(
|
||||
stmt = (
|
||||
select(ranked_subquery)
|
||||
.where(ranked_subquery.c.row_num <= limit)
|
||||
.order_by(ranked_subquery.c.generation_mode.asc(), ranked_subquery.c.generated_time.desc())
|
||||
.order_by(ranked_subquery.c.module_key.asc(), ranked_subquery.c.generated_time.desc())
|
||||
)
|
||||
|
||||
return [dict(row) for row in (await db.execute(stmt)).mappings().all()]
|
||||
|
||||
@@ -1,14 +1,12 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from copy import deepcopy
|
||||
from datetime import datetime, timezone
|
||||
from datetime import datetime
|
||||
from typing import Any
|
||||
|
||||
from fastapi import HTTPException
|
||||
from sqlalchemy import func, select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
from sqlalchemy.orm.attributes import flag_modified
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.common import ModuleEventTypeEnum, ModuleProjectStatusEnum, ModulePromptTypeEnum, ModuleStepStatusEnum
|
||||
@@ -35,16 +33,16 @@ from app.schemas.shot_replicate import (
|
||||
ShotReplicateVideoGenerationOut,
|
||||
ShotReplicateVideoPromptSchemaUpdateRequest,
|
||||
)
|
||||
from app.services.generation_ai_service import (
|
||||
from app.services.generation.ai.engine_service import (
|
||||
VIDEO_DEFAULT_DURATION,
|
||||
VIDEO_DEFAULT_RATIO,
|
||||
VIDEO_DEFAULT_RESOLUTION,
|
||||
_get_video_engine,
|
||||
_parse_list,
|
||||
get_video_engine,
|
||||
parse_json_list,
|
||||
)
|
||||
from app.services.generation_billing_service import charge_module_prompt_usage
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation_task_factory_service import create_chat_generation_task_for_module
|
||||
from app.services.generation.billing_service import charge_module_prompt_usage
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation.task_factory_service import create_chat_generation_task_for_module
|
||||
from app.services.hot_opening_video_prompt_service import (
|
||||
build_final_video_prompt,
|
||||
optimize_hot_opening_video_prompt as optimize_shot_replicate_video_prompt,
|
||||
@@ -52,7 +50,6 @@ from app.services.hot_opening_video_prompt_service import (
|
||||
)
|
||||
from app.services.module_generation_log_service import log_module_error, log_module_event_file, log_module_prompt_event
|
||||
from app.services.llm import optimize_prompt
|
||||
from app.services.resource_accounting_service import soft_delete_chat_task_resources
|
||||
from app.services.module_generation_flow_base_service import (
|
||||
assert_project_has_no_active_chat_tasks as _base_assert_project_has_no_active_chat_tasks,
|
||||
chat_tasks_by_id as _base_chat_tasks_by_id,
|
||||
@@ -976,10 +973,10 @@ async def generate_image_from_prompt(
|
||||
|
||||
|
||||
async def _resolve_video_prompt_config(db: AsyncSession, req: ShotReplicateGenerateVideoPromptRequest) -> dict[str, Any]:
|
||||
engine = await _get_video_engine(db, req.engine_id)
|
||||
supported_ratios = _parse_list(engine.supported_ratios, [])
|
||||
supported_resolutions = _parse_list(engine.supported_resolutions, [])
|
||||
supported_durations = _parse_list(engine.supported_durations, [])
|
||||
engine = await get_video_engine(db, req.engine_id)
|
||||
supported_ratios = parse_json_list(engine.supported_ratios, [])
|
||||
supported_resolutions = parse_json_list(engine.supported_resolutions, [])
|
||||
supported_durations = parse_json_list(engine.supported_durations, [])
|
||||
|
||||
default_ratio = getattr(settings, "SHOT_REPLICATE_DEFAULT_VIDEO_RATIO", None) or VIDEO_DEFAULT_RATIO
|
||||
default_resolution = getattr(settings, "SHOT_REPLICATE_DEFAULT_VIDEO_RESOLUTION", None) or VIDEO_DEFAULT_RESOLUTION
|
||||
|
||||
@@ -14,7 +14,7 @@ from app.config import settings
|
||||
from app.enums.private_portrait import PRIVATE_PORTRAIT_ASSET_URI_PREFIX
|
||||
from app.models.video_engine import VideoEngine
|
||||
from app.services.log_config import is_enabled, LOG_DIR, LOG_DATE_FORMAT, encrypt_data
|
||||
from app.services.generation_provider_types import (
|
||||
from app.types.generation.provider import (
|
||||
ProviderGenerationRecordLike,
|
||||
ProviderVideoEngineLike,
|
||||
)
|
||||
|
||||
@@ -17,7 +17,7 @@ from app.services.resource_accounting_service import (
|
||||
)
|
||||
from app.services.video_cover_service import create_video_cover_for_local_video
|
||||
from app.config import settings
|
||||
from app.services.generation_refund_service import mark_generation_record_failed_and_refund_once
|
||||
from app.services.generation.refund_service import mark_generation_record_failed_and_refund_once
|
||||
|
||||
logger = logging.getLogger("videogen")
|
||||
|
||||
|
||||
@@ -18,10 +18,10 @@ from app.enums.generation_task import (
|
||||
from app.models.base import async_session
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.services.error_codes import extract_error_message
|
||||
from app.services.generation_log_service import log_task_event
|
||||
from app.services.generation_poll_schedule_service import ensure_video_poll_fields
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation_provider_service import create_provider_task
|
||||
from app.services.generation.log_service import log_task_event
|
||||
from app.services.generation.poll_schedule_service import ensure_video_poll_fields
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation.provider_service import create_provider_task
|
||||
from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
|
||||
from app.services.redis_registry_service import ensure_aware_utc
|
||||
from app.tasks.celery_app import celery_app
|
||||
@@ -138,14 +138,22 @@ async def _run(task_id: str):
|
||||
).with_for_update().limit(1))
|
||||
task = result.scalar_one_or_none()
|
||||
|
||||
if not task or task.generation_mode not in ALLOWED_GENERATION_MODES:
|
||||
is_image_main = bool(
|
||||
task
|
||||
and task.generation_mode == GenerationMode.CHATAPI_MAIN.value
|
||||
and task.gen_type == GenerationType.IMAGE.value
|
||||
and int(task.generation_count or 1) > 1
|
||||
)
|
||||
if not task or (task.generation_mode not in ALLOWED_GENERATION_MODES and not is_image_main):
|
||||
return
|
||||
|
||||
if task.status != ChatGenerationTaskStatus.GENERATING.value:
|
||||
return
|
||||
|
||||
deadline_at = ensure_aware_utc(task.deadline_at)
|
||||
if deadline_at and datetime.now(timezone.utc) > deadline_at:
|
||||
# 图片 main 的 deadline 与 provider claim 由 image_batch_service 原子处理,
|
||||
# 避免重复 Celery 消息在有效租约期间把正在执行的批次错误退款。
|
||||
if not is_image_main and deadline_at and datetime.now(timezone.utc) > deadline_at:
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=task,
|
||||
@@ -159,8 +167,10 @@ async def _run(task_id: str):
|
||||
to_status=ChatGenerationTaskStatus.FAILED.value,
|
||||
to_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
)
|
||||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
await aggregate_parent_for_child(db, task)
|
||||
await db.commit()
|
||||
return
|
||||
|
||||
@@ -202,6 +212,12 @@ async def _run(task_id: str):
|
||||
},
|
||||
)
|
||||
|
||||
if is_image_main:
|
||||
from app.services.generation.ai.image_batch_service import run_image_main_batch
|
||||
|
||||
await run_image_main_batch(db, task)
|
||||
return
|
||||
|
||||
if task.seedance_task_id or task.provider_task_id:
|
||||
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
|
||||
if task.gen_type == GenerationType.VIDEO.value:
|
||||
@@ -302,16 +318,41 @@ async def _run(task_id: str):
|
||||
|
||||
if task:
|
||||
error_message = extract_error_message(exc, "生成任务") if callable(extract_error_message) else str(exc)
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=task,
|
||||
error_message=error_message,
|
||||
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
)
|
||||
if is_image_main:
|
||||
# image_batch_service 负责供应商/拆分失败退款。若 child 已落库,
|
||||
# 顶层兜底绝不能再把 main 退款。
|
||||
child_result = await db.execute(
|
||||
select(ChatGenerationTask.id).where(
|
||||
ChatGenerationTask.parent_task_id == task.id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
|
||||
).limit(1)
|
||||
)
|
||||
has_children = child_result.scalar_one_or_none() is not None
|
||||
if not has_children:
|
||||
task.provider_create_claim_token = None
|
||||
task.provider_create_lease_until = None
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=task,
|
||||
error_message=error_message,
|
||||
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
)
|
||||
else:
|
||||
from app.services.generation.ai.task_group_service import aggregate_main_task_status
|
||||
await aggregate_main_task_status(db, parent_task_id=str(task.id))
|
||||
else:
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=task,
|
||||
error_message=error_message,
|
||||
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
)
|
||||
await db.commit()
|
||||
await log_task_event(task, event_type=ChatGenerationTaskEventType.TASK_FAILED.value, message=task.error_message)
|
||||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||||
await log_task_event(task, event_type=ChatGenerationTaskEventType.TASK_FAILED.value, message=error_message)
|
||||
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
await aggregate_parent_for_child(db, task)
|
||||
await db.commit()
|
||||
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@ from app.enums.generation_task import (
|
||||
ChatGenerationPipelineStage,
|
||||
ChatGenerationTaskEventType,
|
||||
ChatGenerationTaskStatus,
|
||||
GenerationMode,
|
||||
GenerationType,
|
||||
)
|
||||
from app.models.base import async_session
|
||||
@@ -28,9 +29,9 @@ from app.services.celery_download_recovery_service import (
|
||||
upsert_download_active,
|
||||
)
|
||||
from app.services.error_codes import extract_error_message
|
||||
from app.services.generation_download_service import download_generation_result
|
||||
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.generation.download_service import download_generation_result
|
||||
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.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
|
||||
from app.services.resource_accounting_service import record_chat_task_generated_resource
|
||||
from app.tasks.celery_app import celery_app
|
||||
@@ -275,9 +276,18 @@ async def enqueue_download_task(
|
||||
await db.commit()
|
||||
|
||||
check_at = _queue_timeout_at(now)
|
||||
await _register_active_from_task(task, check_at=check_at, priority=priority, reason=reason)
|
||||
try:
|
||||
await _register_active_from_task(task, check_at=check_at, priority=priority, reason=reason)
|
||||
except Exception as exc:
|
||||
# Redis active 注册表只用于恢复,不应阻止真实 Celery 投递。
|
||||
await _log_download_event(
|
||||
task,
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE_FAILED,
|
||||
message=f"下载恢复注册表写入失败: {exc}",
|
||||
detail={"reason": reason, "celery_task_id": celery_task_id},
|
||||
)
|
||||
|
||||
await _apply_download_async(
|
||||
applied = await _apply_download_async(
|
||||
task,
|
||||
priority=priority,
|
||||
countdown=countdown,
|
||||
@@ -285,6 +295,31 @@ async def enqueue_download_task(
|
||||
event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE if recover else ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE,
|
||||
failed_event_type=ChatGenerationTaskEventType.DOWNLOAD_RECOVERY_ENQUEUE_FAILED if recover else ChatGenerationTaskEventType.DOWNLOAD_ENQUEUE_FAILED,
|
||||
)
|
||||
if not applied:
|
||||
# apply_async 失败不能伪装成已投递。保留远程结果,进入下载恢复等待。
|
||||
try:
|
||||
await remove_download_active(task.id)
|
||||
except Exception:
|
||||
pass
|
||||
refreshed = await _reload_task(db, task.id)
|
||||
if refreshed and refreshed.status == ChatGenerationTaskStatus.GENERATING.value:
|
||||
retry_at = now + timedelta(seconds=int(settings.DOWNLOAD_TASK_RETRY_BACKOFF_SECONDS or 30))
|
||||
refreshed.pipeline_stage = DOWNLOAD_STAGE_RETRY_WAITING
|
||||
refreshed.download_next_retry_at = retry_at
|
||||
refreshed.download_last_error = "Celery 下载任务投递失败,等待恢复重试"
|
||||
refreshed.download_lease_until = None
|
||||
await db.commit()
|
||||
try:
|
||||
await _register_active_from_task(
|
||||
refreshed,
|
||||
check_at=retry_at,
|
||||
priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER,
|
||||
reason="enqueue_failed_wait_recovery",
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return None
|
||||
|
||||
if old_stage != DOWNLOAD_STAGE_QUEUED:
|
||||
# 独立记录阶段变化的上下文,便于和真正投递事件对照。
|
||||
await _log_download_event(
|
||||
@@ -466,19 +501,32 @@ async def _mark_download_failed(
|
||||
non_retryable: bool = False,
|
||||
) -> None:
|
||||
error_message = extract_error_message(exc, "下载") if callable(extract_error_message) else str(exc)
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=task,
|
||||
error_message=error_message,
|
||||
pipeline_stage=DOWNLOAD_STAGE_FAILED,
|
||||
is_image_child = (
|
||||
task.gen_type == GenerationType.IMAGE.value
|
||||
and task.generation_mode == GenerationMode.CHATAPI_CHILD.value
|
||||
)
|
||||
if is_image_child:
|
||||
# 图片生成费用属于 main;child 下载失败只记录下载终态,不退图片生成积分。
|
||||
task.status = ChatGenerationTaskStatus.FAILED.value
|
||||
task.pipeline_stage = DOWNLOAD_STAGE_FAILED
|
||||
task.error_message = error_message
|
||||
await db.flush()
|
||||
else:
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=task,
|
||||
error_message=error_message,
|
||||
pipeline_stage=DOWNLOAD_STAGE_FAILED,
|
||||
)
|
||||
task.download_last_error = error_message
|
||||
task.download_lease_until = None
|
||||
task.download_next_retry_at = None
|
||||
|
||||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
await aggregate_parent_for_child(db, task)
|
||||
await db.commit()
|
||||
await remove_download_active(task.id)
|
||||
|
||||
@@ -549,9 +597,11 @@ async def _run(task_id: str):
|
||||
)
|
||||
await sync_chat_generation_task_media_token_snapshot(db, task)
|
||||
|
||||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
await aggregate_parent_for_child(db, task)
|
||||
await db.commit()
|
||||
await remove_download_active(task.id)
|
||||
|
||||
|
||||
@@ -19,8 +19,8 @@ from app.enums.generation_task import (
|
||||
from app.models.base import async_session
|
||||
from app.models.chat_generation_task import ChatGenerationTask
|
||||
from app.services.error_codes import extract_error_message
|
||||
from app.services.generation_log_service import log_task_event, log_provider_call
|
||||
from app.services.generation_poll_schedule_service import (
|
||||
from app.services.generation.log_service import log_task_event, log_provider_call
|
||||
from app.services.generation.poll_schedule_service import (
|
||||
build_default_poll_schedule,
|
||||
build_video_pending_poll_schedule,
|
||||
ensure_video_poll_fields,
|
||||
@@ -28,8 +28,8 @@ from app.services.generation_poll_schedule_service import (
|
||||
is_poll_not_due,
|
||||
is_video_generation_task,
|
||||
)
|
||||
from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation_provider_service import poll_provider_task
|
||||
from app.services.generation.refund_service import mark_chat_generation_task_failed_and_refund_once
|
||||
from app.services.generation.provider_service import poll_provider_task
|
||||
from app.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot
|
||||
from app.services.redis_registry_service import (
|
||||
datetime_to_epoch,
|
||||
@@ -144,9 +144,11 @@ async def remove_poll_active(task_id: str) -> None:
|
||||
|
||||
|
||||
async def _notify_finished(db, task: ChatGenerationTask) -> None:
|
||||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
await aggregate_parent_for_child(db, task)
|
||||
|
||||
|
||||
async def _reload_task(db, task_id: str) -> ChatGenerationTask | None:
|
||||
|
||||
@@ -18,21 +18,21 @@ RecoveryRunner = Callable[[], Awaitable[Dict[str, Any]]]
|
||||
|
||||
|
||||
async def _run_download_once() -> Dict[str, Any]:
|
||||
from app.services.generation_recovery_service import recover_download_tasks_once
|
||||
from app.services.generation.recovery_service import recover_download_tasks_once
|
||||
|
||||
async with async_session() as db:
|
||||
return await recover_download_tasks_once(db)
|
||||
|
||||
|
||||
async def _run_generation_once() -> Dict[str, Any]:
|
||||
from app.services.generation_recovery_service import recover_generation_tasks_once
|
||||
from app.services.generation.recovery_service import recover_generation_tasks_once
|
||||
|
||||
async with async_session() as db:
|
||||
return await recover_generation_tasks_once(db)
|
||||
|
||||
|
||||
async def _run_due_poll_dispatch_once() -> Dict[str, Any]:
|
||||
from app.services.generation_recovery_service import dispatch_due_poll_tasks_once
|
||||
from app.services.generation.recovery_service import dispatch_due_poll_tasks_once
|
||||
|
||||
async with async_session() as db:
|
||||
return await dispatch_due_poll_tasks_once(db)
|
||||
|
||||
@@ -38,7 +38,7 @@ from app.services.module_async_recovery_service import (
|
||||
)
|
||||
from app.services.shot_replicate_taskset_service import refresh_task_set_split_summary
|
||||
from app.services.shot_video_analysis_service import analyze_video_for_shot_split
|
||||
from app.services.generation_billing_service import charge_shot_video_analysis_usage
|
||||
from app.services.generation.billing_service import charge_shot_video_analysis_usage
|
||||
from app.services.shot_video_split_service import split_video_segment_async
|
||||
from app.services.upload_video_asset_service import validate_split_range
|
||||
from app.services.upload_resource import record_shot_segment_upload_resource
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""应用级类型契约。"""
|
||||
@@ -0,0 +1 @@
|
||||
"""生成领域类型契约。"""
|
||||
+31
-6
@@ -1,14 +1,10 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Protocol
|
||||
from typing import Any, Protocol, TypedDict
|
||||
|
||||
|
||||
class ProviderGenerationRecordLike(Protocol):
|
||||
"""图片/视频供应商提交接口需要的任务字段协议。
|
||||
|
||||
GenerationRecord 与 ChatGenerationTask 都具备这些字段,但二者不是同一个 ORM 模型。
|
||||
使用 Protocol 可以避免把 submit_image_task / submit_video_task 错误限制为某一个具体模型。
|
||||
"""
|
||||
"""图片/视频供应商提交接口需要的任务字段协议。"""
|
||||
|
||||
id: str
|
||||
original_prompt: str
|
||||
@@ -21,16 +17,25 @@ class ProviderGenerationRecordLike(Protocol):
|
||||
image_size: str | None
|
||||
image_proportion: str | None
|
||||
image_px: str | None
|
||||
generation_count: int
|
||||
engine_id: str | None
|
||||
|
||||
|
||||
class ProviderImageEngineLike(Protocol):
|
||||
"""图片生成提交接口需要的引擎字段协议。"""
|
||||
|
||||
id: str
|
||||
name: str
|
||||
provider: str
|
||||
api_base: str
|
||||
api_key: str
|
||||
model_name: str
|
||||
default_size: str | None
|
||||
multi_generation_enabled: bool
|
||||
max_generation_count: int
|
||||
multi_image_max_images: int
|
||||
max_reference_image_count: int
|
||||
output_format: str
|
||||
|
||||
|
||||
class ProviderVideoEngineLike(Protocol):
|
||||
@@ -40,3 +45,23 @@ class ProviderVideoEngineLike(Protocol):
|
||||
api_base: str
|
||||
api_key: str
|
||||
model_name: str
|
||||
|
||||
|
||||
class ImageProviderItem(TypedDict, total=False):
|
||||
generation_index: int
|
||||
remote_result_url: str
|
||||
b64_json: str
|
||||
size: str
|
||||
output_format: str
|
||||
response_data: dict[str, Any]
|
||||
error_code: str
|
||||
error_message: str
|
||||
|
||||
|
||||
class ImageProviderBatchResult(TypedDict, total=False):
|
||||
items: list[ImageProviderItem]
|
||||
model: str
|
||||
created: int
|
||||
generated_images: int
|
||||
image_tokens: int
|
||||
response_data: dict[str, Any]
|
||||
Reference in New Issue
Block a user