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 staged_provider_result: dict | 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 _parse_staged_provider_result(raw: str | None) -> dict | None: if not raw: return None try: value = json.loads(raw) except (TypeError, ValueError, json.JSONDecodeError): return None if not isinstance(value, dict): return None items = value.get("items") if not isinstance(items, list) or not items: return None return value 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() staged_provider_result = None if main.pipeline_stage == ChatGenerationPipelineStage.PROVIDER_RESULT_STAGED.value: staged_provider_result = _parse_staged_provider_result(main.provider_response_json) if staged_provider_result is None: main.pipeline_stage = ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value main.provider_response_json = None 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 and staged_provider_result is None: 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.PROVIDER_RESULT_STAGED.value if staged_provider_result is not None else 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, staged_provider_result=staged_provider_result, reason="provider_result_staged" if staged_provider_result is not None else None, ) 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 _stage_provider_result( db: AsyncSession, *, main_task_id: str, claim_token: str, provider_result: dict, ) -> None: 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.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}") main.provider_response_json = _json(provider_result) main.image_tokens_used = int(provider_result.get("image_tokens") or 0) main.pipeline_stage = ChatGenerationPipelineStage.PROVIDER_RESULT_STAGED.value user_id_snapshot = str(main.user_id) generation_count_snapshot = int(main.generation_count or 1) image_tokens_snapshot = int(main.image_tokens_used or 0) await db.commit() log_operation_event( domain="generation_ai_batch", event_type="IMAGE_BATCH_PROVIDER_RESULT_STAGED", event_status="success", source="celery", user_id=user_id_snapshot, group_id=main_task_id, task_id=main_task_id, detail={ "generation_count": generation_count_snapshot, "image_tokens": image_tokens_snapshot, }, ) 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)) provider_result = claim.staged_provider_result if provider_result is None: 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) await _stage_provider_result( db, main_task_id=main_task_id, claim_token=claim.claim_token, provider_result=provider_result, ) 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, "provider_result_staged": True, }, ) 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 [] else: provider_items = _validate_provider_batch(provider_result, generation_count) log_operation_event( domain="generation_ai_batch", event_type="IMAGE_BATCH_STAGED_RESULT_RECOVERED", event_status="success", source="recovery", user_id=claim.task_snapshot.user_id, group_id=main_task_id, task_id=main_task_id, detail={"generation_count": generation_count, "provider_regenerated": False}, ) 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: # 供应商结果已经落库;拆分失败只记录并等待恢复,绝不退款或重新调用供应商。 try: await db.rollback() except Exception: pass log_operation_event( domain="generation_ai_batch", event_type=ChatGenerationTaskEventType.IMAGE_BATCH_SPLIT_FAILED.value, event_status="failed", source="celery", user_id=claim.task_snapshot.user_id, group_id=main_task_id, task_id=main_task_id, message=f"图片批量结果拆分失败: {exc}", detail={"provider_result_staged": True, "provider_regenerated": False}, error=str(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