568 lines
20 KiB
Python
568 lines
20 KiB
Python
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
|