Files
video-gen/video-gen-api/app/services/generation/ai/image_batch_service.py
T
2026-07-20 13:48:17 +08:00

568 lines
20 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
from __future__ import annotations
import json
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from typing import Awaitable, Callable
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.pipeline.db_lock_service import execute_with_lock_timeout
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.services.redis_registry_service import RedisExecutionLockError
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,
*,
execution_token: str,
) -> ImageBatchClaim:
result = await execute_with_lock_timeout(
db,
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 = execution_token
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 execute_with_lock_timeout(
db,
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 execute_with_lock_timeout(
db,
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_created_at = datetime.now(timezone.utc)
child = ChatGenerationTask(
id=generate_id(),
created_at=child_created_at,
resource_generation_started_at=child_created_at,
generation_attempt_no=1,
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,
*,
execution_token: str,
execution_guard: Callable[[], Awaitable[None]],
) -> list[str]:
"""单次同步组图,全部成功后原子拆分 child。
绝不在组图 API 失败后退化为 N 次单图请求。
"""
main_task_id = str(main_task.id)
claim = await _claim_image_main_batch(
db,
main_task_id,
execution_token=execution_token,
)
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,
)
await execution_guard()
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 RedisExecutionLockError:
raise
except Exception as exc:
await execution_guard()
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:
await execution_guard()
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 RedisExecutionLockError:
raise
except Exception as exc:
await execution_guard()
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