from __future__ import annotations import asyncio import json import uuid from datetime import datetime, timedelta, timezone from typing import Any, Optional from app.config import settings from app.enums.celery_queue import CeleryQueue, CeleryTaskName from app.enums.celery_runtime import CeleryRuntimeDomain from app.enums.generation_status import GenerationRecordPipelineStage from app.enums.generation_task import ( ALLOWED_GENERATION_MODES, ChatGenerationPipelineStage, ChatGenerationTaskEventType, GenerationMode, GenerationOwnerType, GenerationType, ) 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.pipeline.db_lock_service import DatabaseRowLockBusy from app.services.generation.pipeline.lifecycle_service import ( mark_owner_failed_and_refund_once, notify_owner_finished, ) from app.services.generation.pipeline.owner_service import ( GenerationOwner, is_attempt_current, load_generation_owner, normalize_owner_type, owner_is_generating, owner_mode, owner_provider_task_id, set_owner_provider_task_id, renew_generation_owner_claim_lease, ) from app.services.generation.poll_schedule_service import ensure_video_poll_fields 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, sync_generation_record_media_token_snapshot, ) from app.services.redis_registry_service import RedisExecutionLockError from app.services.celery_runtime.runtime_service import CeleryRuntimeLease, RuntimeIdentity from app.tasks.async_runner import run_async from app.tasks.celery_app import celery_app def _now() -> datetime: return datetime.now(timezone.utc) def _get_first_value(obj: Any, *field_names: str) -> Optional[Any]: for field_name in field_names: value = getattr(obj, field_name, None) if value is not None and str(value).strip() != "": return value return None def _clean(value: Any) -> str | None: text = str(value).strip() if value is not None else "" return text or None def _build_optimized_prompt_by_params(owner: GenerationOwner) -> str: base_prompt = (_clean(owner.original_prompt) or "").rstrip(",,。;; \n\t") gen_type = (_clean(owner.gen_type) or "").lower() generation_mode = owner_mode(owner) if ( generation_mode in { GenerationMode.HOT_OPENING_REPLICATE.value, GenerationMode.SHOT_REPLICATE.value, } and gen_type == GenerationType.VIDEO.value ): stripped = base_prompt.strip() if stripped.startswith(("{", "[")): return base_prompt parts: list[str] = [] if gen_type == GenerationType.VIDEO.value: parts.extend( [ f"时长:{_get_first_value(owner, 'duration') or 4}秒", f"画面比例:{_get_first_value(owner, 'aspect_ratio') or '16:9'}", f"分辨率:{_get_first_value(owner, 'provider_generation_resolution', 'resolution') or '480p'}", ] ) elif gen_type == GenerationType.IMAGE.value: parts.extend( [ f"分辨率:{_get_first_value(owner, 'image_size') or '2K'}", f"画布比例:{_get_first_value(owner, 'image_proportion') or '1:1'}", f"像素尺寸:{_get_first_value(owner, 'image_px') or '2048x2048'}", ] ) suffix = ",".join(parts) return f"{base_prompt},{suffix}" if base_prompt and suffix else base_prompt or suffix def _stage(owner: GenerationOwner, chat_stage: ChatGenerationPipelineStage) -> str: if isinstance(owner, ChatGenerationTask): return chat_stage.value try: return GenerationRecordPipelineStage(chat_stage.value).value except ValueError: return chat_stage.value def _lock_key(owner_type: str, owner_id: str, attempt_no: int) -> str: return ( f"{settings.GENERATION_CREATE_LOCK_KEY_PREFIX}:" f"{owner_type}:{owner_id}:attempt:{attempt_no}" ) async def _sync_media_snapshot( db, owner: GenerationOwner, provider_response: Any = None ) -> None: if isinstance(owner, ChatGenerationTask): await sync_chat_generation_task_media_token_snapshot( db, owner, provider_response=provider_response ) else: await sync_generation_record_media_token_snapshot( db, owner, provider_response=provider_response ) async def _stale(owner: GenerationOwner, attempt_no: int | None) -> bool: if is_attempt_current(owner, attempt_no): return False await log_task_event( owner, event_type=ChatGenerationTaskEventType.STALE_ATTEMPT_MESSAGE_SKIPPED.value, message="创建任务消息属于旧生成轮次,已跳过", detail={ "message_attempt": attempt_no, "current_attempt": owner.generation_attempt_no, }, ) return True async def _reload_owner_after_external_call( db, *, owner_type: str, owner_id: str, ) -> GenerationOwner | None: """Keep the provider result in the current Worker while briefly retrying a busy row lock.""" last_error: DatabaseRowLockBusy | None = None for retry_index in range(3): try: return await load_generation_owner( db, owner_type=owner_type, owner_id=owner_id, for_update=True, ) except DatabaseRowLockBusy as exc: last_error = exc await db.rollback() if retry_index < 2: await asyncio.sleep(1 + retry_index) raise last_error or DatabaseRowLockBusy() async def _resolve_attempt( task_id: str, *, owner_type: str, message_attempt: int | None ) -> int | None: async with async_session() as db: owner = await load_generation_owner( db, owner_type=owner_type, owner_id=task_id, for_update=False ) if not owner: return None if await _stale(owner, message_attempt): return None return int(owner.generation_attempt_no or 1) async def _dispatch_next_stage( db, owner: GenerationOwner, *, normalized_owner_type: str, ) -> None: if owner.pipeline_stage == _stage( owner, ChatGenerationPipelineStage.RESULT_READY ): from app.tasks.generation_download_tasks import enqueue_download_task await enqueue_download_task(db, owner, reason="create_result_ready") return from app.tasks.generation_poll_tasks import ( poll_generation_task, register_poll_active, ) try: check_at = _now() + timedelta( seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300) ) await register_poll_active( owner, reason="create_provider_success", check_at=check_at, next_poll_at=owner.next_poll_at, ) poll_generation_task.apply_async( args=[owner.id], kwargs={ "owner_type": normalized_owner_type, "generation_attempt_no": int(owner.generation_attempt_no or 1), "force_due": False, }, queue=CeleryQueue.GEN_PROVIDER_POLL.value, ) except Exception as exc: # Provider creation is already committed. A broker/registry failure is # an infrastructure enqueue failure, not a generation failure. Keep the # remote task ID and let due-poll/startup recovery enqueue it again. owner.next_poll_at = _now() await db.commit() await log_task_event( owner, event_type=ChatGenerationTaskEventType.POLL_SCHEDULED.value, message="供应商任务已创建,但轮询任务投递失败,等待恢复扫描", detail={ "error": str(exc), "provider_task_id": owner_provider_task_id(owner), "generation_refunded": False, }, to_stage=owner.pipeline_stage, ) async def _run( task_id: str, *, owner_type: str = GenerationOwnerType.CHAT_GENERATION_TASK.value, generation_attempt_no: int | None = None, ): normalized_owner_type = normalize_owner_type(owner_type) effective_attempt = await _resolve_attempt( task_id, owner_type=normalized_owner_type, message_attempt=generation_attempt_no, ) if effective_attempt is None: return lease_token = uuid.uuid4().hex async def _renew_db_claim(token: str) -> bool: return await renew_generation_owner_claim_lease( owner_type=normalized_owner_type, owner_id=task_id, attempt_no=effective_attempt, claim_field="provider_create_claim_token", lease_field="provider_create_lease_until", token=token, lease_seconds=int(settings.GENERATION_CREATE_LOCK_TTL_SECONDS or 600), ) lease = await CeleryRuntimeLease.acquire( identity=RuntimeIdentity( domain=CeleryRuntimeDomain.GENERATION_CREATE.value, owner_type=normalized_owner_type, owner_id=task_id, attempt_no=effective_attempt, task_name=CeleryTaskName.CHATAPI_CREATE.value, queue=CeleryQueue.GEN_CHATAPI_CREATE.value, ), lock_key=_lock_key(normalized_owner_type, task_id, effective_attempt), hash_key=settings.GENERATION_CREATE_ACTIVE_REDIS_HASH_KEY, zset_key=settings.GENERATION_CREATE_ACTIVE_REDIS_ZSET_KEY, token=lease_token, ttl_seconds=int(settings.GENERATION_CREATE_LOCK_TTL_SECONDS or 600), heartbeat_interval_seconds=int(settings.REDIS_EXECUTION_LOCK_RENEW_INTERVAL_SECONDS or 30), pipeline_stage=ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value, db_heartbeat=_renew_db_claim, ) if lease is None: return try: async with async_session() as db: owner = await load_generation_owner( db, owner_type=normalized_owner_type, owner_id=task_id, for_update=True, ) if not owner: return if not is_attempt_current(owner, effective_attempt): await db.rollback() return is_image_main = bool( isinstance(owner, ChatGenerationTask) and owner.generation_mode == GenerationMode.CHATAPI_MAIN.value and owner.gen_type == GenerationType.IMAGE.value and int(owner.generation_count or 1) > 1 ) if ( isinstance(owner, ChatGenerationTask) and owner.generation_mode not in ALLOWED_GENERATION_MODES and not is_image_main ): return if not owner_is_generating(owner): return if owner.deadline_at and _now() > owner.deadline_at and not is_image_main: await mark_owner_failed_and_refund_once( db, owner, error_message="任务超时", pipeline_stage=_stage( owner, ChatGenerationPipelineStage.TIMEOUT ), ) await db.commit() await notify_owner_finished(db, owner) await db.commit() await log_task_event( owner, event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value, to_stage=owner.pipeline_stage, ) return allowed_stages = { _stage(owner, ChatGenerationPipelineStage.QUEUED), _stage(owner, ChatGenerationPipelineStage.PREPARING), _stage(owner, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK), } if is_image_main: allowed_stages.add( _stage(owner, ChatGenerationPipelineStage.PROVIDER_RESULT_STAGED) ) if owner.pipeline_stage not in allowed_stages: return try: if not owner.optimized_prompt: owner.pipeline_stage = _stage( owner, ChatGenerationPipelineStage.PREPARING ) owner.optimized_prompt = _build_optimized_prompt_by_params(owner) owner.text_tokens_used = int(owner.text_tokens_used or 0) await db.commit() await log_task_event( owner, event_type=ChatGenerationTaskEventType.PROMPT_CONCAT_SUCCESS.value, to_stage=owner.pipeline_stage, ) if is_image_main: from app.services.generation.ai.image_batch_service import ( run_image_main_batch, ) await run_image_main_batch( db, owner, execution_token=lease.token, execution_guard=lease.ensure_owned, ) return if owner_provider_task_id(owner): owner.pipeline_stage = _stage( owner, ChatGenerationPipelineStage.WAITING_REMOTE ) if owner.gen_type == GenerationType.VIDEO.value: ensure_video_poll_fields(owner, now=_now()) owner.next_poll_at = _now() await db.commit() elif owner.remote_result_url: owner.pipeline_stage = _stage( owner, ChatGenerationPipelineStage.RESULT_READY ) await db.commit() else: current_time = _now() owner.pipeline_stage = _stage( owner, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK ) owner.provider_create_claim_token = lease.token owner.provider_create_started_at = current_time owner.provider_create_lease_until = current_time + timedelta( seconds=int( settings.GENERATION_CREATE_LOCK_TTL_SECONDS or 600 ) ) await db.commit() await log_task_event( owner, event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_START.value, to_stage=owner.pipeline_stage, ) created = await create_provider_task(db, owner) await lease.ensure_owned() owner = await _reload_owner_after_external_call( db, owner_type=normalized_owner_type, owner_id=task_id, ) if not owner: return if not is_attempt_current(owner, effective_attempt): await db.rollback() return if owner.provider_create_claim_token != lease.token: return provider_task_id = created.get("task_id") if provider_task_id: set_owner_provider_task_id(owner, str(provider_task_id)) owner.remote_result_url = ( created.get("remote_result_url") or owner.remote_result_url ) owner.provider_response_json = json.dumps( created.get("response_data") or {}, ensure_ascii=False, default=str, ) owner.provider_create_claim_token = None owner.provider_create_lease_until = None if owner.gen_type == GenerationType.IMAGE.value: owner.image_tokens_used = int( created.get( "image_tokens", owner.image_tokens_used or 0 ) or 0 ) await _sync_media_snapshot( db, owner, owner.provider_response_json ) if owner.remote_result_url and not owner_provider_task_id(owner): owner.pipeline_stage = _stage( owner, ChatGenerationPipelineStage.RESULT_READY ) else: owner.pipeline_stage = _stage( owner, ChatGenerationPipelineStage.WAITING_REMOTE ) if owner.gen_type == GenerationType.VIDEO.value: ensure_video_poll_fields(owner, now=_now()) owner.next_poll_at = _now() await db.commit() await log_task_event( owner, event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_SUCCESS.value, to_stage=owner.pipeline_stage, detail=created, ) except (RedisExecutionLockError, DatabaseRowLockBusy): # Redis ownership is mandatory. Do not convert infrastructure # lock loss into a business failure/refund. try: await db.rollback() except Exception: pass raise except Exception as exc: try: await db.rollback() except Exception: pass # 只有仍持有 Redis 执行权时,才能把供应商异常收敛为业务失败。 await lease.ensure_owned() owner = await load_generation_owner( db, owner_type=normalized_owner_type, owner_id=task_id, for_update=True, ) if not owner or not is_attempt_current(owner, effective_attempt): return error_message = ( extract_error_message(exc, "生成任务") if callable(extract_error_message) else str(exc) ) if is_image_main and isinstance(owner, ChatGenerationTask): from sqlalchemy import select child_result = await db.execute( select(ChatGenerationTask.id) .where(ChatGenerationTask.parent_task_id == owner.id) .limit(1) ) if child_result.scalar_one_or_none() is not None: from app.services.generation.ai.task_group_service import ( aggregate_main_task_status, ) await aggregate_main_task_status( db, parent_task_id=str(owner.id) ) else: await mark_owner_failed_and_refund_once( db, owner, error_message=error_message, pipeline_stage=_stage( owner, ChatGenerationPipelineStage.FAILED ), ) else: await mark_owner_failed_and_refund_once( db, owner, error_message=error_message, pipeline_stage=_stage( owner, ChatGenerationPipelineStage.FAILED ), ) await db.commit() await notify_owner_finished(db, owner) await db.commit() await log_task_event( owner, event_type=ChatGenerationTaskEventType.TASK_FAILED.value, message=error_message, ) return # The provider state is committed. Enqueue failures below are # recoverable infrastructure failures and must not trigger refunds. await _dispatch_next_stage( db, owner, normalized_owner_type=normalized_owner_type ) finally: await lease.close() if celery_app: @celery_app.task( name="generation.chatapi_create_generation_task", bind=True, max_retries=3, default_retry_delay=30, ) def chatapi_create_generation_task( self, task_id: str, owner_type: str = GenerationOwnerType.CHAT_GENERATION_TASK.value, generation_attempt_no: int | None = None, ): try: return run_async( _run( task_id, owner_type=owner_type, generation_attempt_no=generation_attempt_no, ) ) except Exception as exc: retries = int(getattr(self.request, "retries", 0) or 0) + 1 countdown = int( settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30 ) * max(1, retries) raise self.retry(exc=exc, countdown=countdown) else: class _DisabledTask: def delay(self, *args, **kwargs): raise RuntimeError("Celery is disabled") def apply_async(self, *args, **kwargs): raise RuntimeError("Celery is disabled") chatapi_create_generation_task = _DisabledTask()