from __future__ import annotations import logging import json import uuid from datetime import datetime, timedelta, timezone from typing import Any 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, GenerationOwnerType, GenerationType, PROVIDER_FAILED_STATUSES, PROVIDER_SUCCESS_STATUSES, ) 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_provider_call, 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, load_generation_owner_for_update_retry, normalize_owner_type, owner_is_generating, owner_mode, owner_provider_task_id, owner_type_of, redis_owner_item_id, renew_generation_owner_claim_lease, ) from app.services.generation.poll_schedule_service import ( build_default_poll_schedule, build_video_pending_poll_schedule, ensure_video_poll_fields, is_final_poll_due, is_poll_not_due, is_video_generation_task, ) 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, sync_generation_record_media_token_snapshot, ) from app.services.redis_registry_service import ( RedisExecutionLockError, datetime_to_epoch, ensure_aware_utc, redis_get_registry_payloads, redis_remove_registry_item, redis_upsert_registry_item, utc_now, ) 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 logger = logging.getLogger("video_gen") POLL_QUEUE = CeleryQueue.GEN_PROVIDER_POLL.value def _now() -> datetime: return datetime.now(timezone.utc) def _is_success(status: str | None) -> bool: return str(status or "").lower() in PROVIDER_SUCCESS_STATUSES def _is_failed(status: str | None) -> bool: return str(status or "").lower() in PROVIDER_FAILED_STATUSES 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 _engine_snapshot(owner: GenerationOwner) -> dict: try: data = json.loads(owner.engine_snapshot_json or "{}") return data if isinstance(data, dict) else {} except Exception: return {} async def _log_poll_provider_call_after_commit( owner: GenerationOwner, *, provider_response: Any, ) -> None: """Provider logs use an independent session, so the owner row must be committed first.""" snapshot = _engine_snapshot(owner) await log_provider_call( owner, provider=snapshot.get("provider") or "ark", api_type=f"{owner.gen_type}_poll", model=snapshot.get("model_name"), engine_id=owner.engine_id, status="success", provider_task_id=owner_provider_task_id(owner), response_data=provider_response, ) def _registry_id(owner: GenerationOwner) -> str: return redis_owner_item_id( owner_type_of(owner), owner.id, int(owner.generation_attempt_no or 1) ) def _poll_lease_until(now: datetime | None = None) -> datetime: return (now or _now()) + timedelta( seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300) ) def _lock_key(owner_type: str, owner_id: str, attempt_no: int) -> str: return ( f"{settings.GENERATION_POLL_LOCK_KEY_PREFIX}:" f"{owner_type}:{owner_id}:attempt:{attempt_no}" ) def _build_poll_active_payload( owner: GenerationOwner, *, reason: str, next_poll_at: datetime | None, check_at: datetime | None, ) -> dict[str, Any]: return { "owner_type": owner_type_of(owner), "owner_id": owner.id, "task_id": owner.id, "generation_attempt_no": int(owner.generation_attempt_no or 1), "provider_task_id": owner_provider_task_id(owner), "generation_mode": owner_mode(owner), "gen_type": owner.gen_type, "stage": owner.pipeline_stage or "", "queue": POLL_QUEUE, "poll_count": int(owner.poll_count or 0), "poll_error_count": int(getattr(owner, "poll_error_count", 0) or 0), "manual_retry_count": int(getattr(owner, "manual_retry_count", 0) or 0), "poll_started_at": datetime_to_epoch(owner.poll_started_at) if owner.poll_started_at else None, "poll_interval_seconds": int(owner.poll_interval_seconds or 0), "last_poll_at": datetime_to_epoch(owner.last_poll_at) if owner.last_poll_at else None, "next_poll_at": datetime_to_epoch(ensure_aware_utc(next_poll_at)) if next_poll_at else None, "poll_lease_until": datetime_to_epoch(owner.poll_lease_until) if owner.poll_lease_until else None, "deadline_at": datetime_to_epoch(owner.deadline_at) if owner.deadline_at else None, "check_at": datetime_to_epoch(ensure_aware_utc(check_at)) if check_at else None, "updated_at": datetime_to_epoch(utc_now()), "reason": reason, } async def register_poll_active( owner: GenerationOwner, *, check_at: datetime, reason: str, next_poll_at: datetime | None = None, ) -> None: item_id = _registry_id(owner) payload = _build_poll_active_payload( owner, reason=reason, next_poll_at=next_poll_at, check_at=check_at, ) existing = await redis_get_registry_payloads( hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY, item_ids=[item_id], log_context="poll_active", ) if item_id in existing: merged = dict(existing[item_id]) merged.update(payload) payload = merged await redis_upsert_registry_item( hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY, zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY, item_id=item_id, payload=payload, check_at=check_at, log_context="poll_active", ) async def remove_poll_active( owner: GenerationOwner | None = None, *, owner_type: str | None = None, owner_id: str | None = None, attempt_no: int | None = None, ) -> None: if owner is not None: item_id = _registry_id(owner) else: item_id = redis_owner_item_id(owner_type, str(owner_id or ""), attempt_no) await redis_remove_registry_item( hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY, zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY, item_id=item_id, log_context="poll_active", ) async def _sync_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 _mark_failed( db, owner: GenerationOwner, *, message: str, stage: ChatGenerationPipelineStage, event_type: str, detail: Any = None, ) -> None: await mark_owner_failed_and_refund_once( db, owner, error_message=message, pipeline_stage=_stage(owner, stage), ) owner.next_poll_at = None owner.poll_claim_token = None owner.poll_lease_until = None await db.commit() await notify_owner_finished(db, owner) await db.commit() await remove_poll_active(owner) await log_task_event( owner, event_type=event_type, message=message, detail=detail, to_stage=owner.pipeline_stage, ) async def _schedule_next_poll( db, owner: GenerationOwner, *, reason: str, default_delay_seconds: int | None = None, ) -> None: current = _now() schedule = ( build_video_pending_poll_schedule(owner, now=current) if is_video_generation_task(owner) else build_default_poll_schedule( owner, now=current, delay_seconds=default_delay_seconds, reason=reason, ) ) owner.pipeline_stage = _stage(owner, ChatGenerationPipelineStage.WAITING_REMOTE) owner.next_poll_at = schedule.next_poll_at owner.poll_interval_seconds = schedule.poll_interval_seconds owner.poll_claim_token = None owner.poll_lease_until = None await db.commit() await register_poll_active( owner, check_at=schedule.next_poll_at, next_poll_at=schedule.next_poll_at, reason=schedule.reason, ) await log_task_event( owner, event_type=ChatGenerationTaskEventType.POLL_SCHEDULED.value, message=f"已登记下一次轮询。reason={schedule.reason}", detail={ "next_poll_at": schedule.next_poll_at, "delay_seconds": schedule.delay_seconds, }, ) if schedule.direct_countdown: poll_generation_task.apply_async( args=[owner.id], kwargs={ "owner_type": owner_type_of(owner), "generation_attempt_no": int(owner.generation_attempt_no or 1), "force_due": False, }, queue=POLL_QUEUE, countdown=max(0, int(schedule.delay_seconds)), ) 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 not is_attempt_current(owner, message_attempt): await remove_poll_active( owner_type=owner_type, owner_id=task_id, attempt_no=message_attempt, ) await log_task_event( owner, event_type=ChatGenerationTaskEventType.STALE_ATTEMPT_MESSAGE_SKIPPED.value, message="轮询消息属于旧生成轮次,已跳过", detail={ "message_attempt": message_attempt, "current_attempt": owner.generation_attempt_no, }, ) return None return int(owner.generation_attempt_no or 1) async def _restore_after_lock_error( *, owner_type: str, owner_id: str, attempt_no: int, claim_token: str, ) -> None: async with async_session() as db: owner = await load_generation_owner( db, owner_type=owner_type, owner_id=owner_id, for_update=True, ) if not owner or not is_attempt_current(owner, attempt_no): return if owner.poll_claim_token != claim_token: return owner.pipeline_stage = _stage(owner, ChatGenerationPipelineStage.WAITING_REMOTE) owner.poll_claim_token = None owner.poll_lease_until = None owner.next_poll_at = _now() await db.commit() await register_poll_active( owner, check_at=owner.next_poll_at, next_poll_at=owner.next_poll_at, reason="poll_execution_lock_error", ) async def _run( task_id: str, *, owner_type: str = GenerationOwnerType.CHAT_GENERATION_TASK.value, generation_attempt_no: int | None = None, force_due: bool = False, ): 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 token = uuid.uuid4().hex lease = await CeleryRuntimeLease.acquire( identity=RuntimeIdentity( domain=CeleryRuntimeDomain.GENERATION_POLL.value, owner_type=normalized_owner_type, owner_id=task_id, attempt_no=effective_attempt, task_name=CeleryTaskName.POLL_GENERATION.value, queue=POLL_QUEUE, registry_item_id=redis_owner_item_id( normalized_owner_type, task_id, effective_attempt ), ), lock_key=_lock_key(normalized_owner_type, task_id, effective_attempt), hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY, zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY, token=token, ttl_seconds=int(settings.GENERATION_POLL_LOCK_TTL_SECONDS or 300), heartbeat_interval_seconds=int(settings.REDIS_EXECUTION_LOCK_RENEW_INTERVAL_SECONDS or 30), pipeline_stage=ChatGenerationPipelineStage.POLLING.value, db_heartbeat=lambda owned_token: renew_generation_owner_claim_lease( owner_type=normalized_owner_type, owner_id=task_id, attempt_no=effective_attempt, claim_field="poll_claim_token", lease_field="poll_lease_until", token=owned_token, lease_seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300), ), ) if lease is None: return claim_started = False 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: await remove_poll_active( owner_type=normalized_owner_type, owner_id=task_id, attempt_no=effective_attempt, ) return if not is_attempt_current(owner, effective_attempt): await remove_poll_active(owner) return if ( isinstance(owner, ChatGenerationTask) and owner.generation_mode not in ALLOWED_GENERATION_MODES ): await remove_poll_active(owner) return if not owner_is_generating(owner) or owner.pipeline_stage not in { _stage(owner, ChatGenerationPipelineStage.WAITING_REMOTE), _stage(owner, ChatGenerationPipelineStage.POLLING), }: await remove_poll_active(owner) return current = _now() if is_video_generation_task(owner): ensure_video_poll_fields(owner, now=current) # Redis execution lock is the authoritative single-worker guard. # A stale database lease left by a crashed worker must not block the # worker that successfully acquired the current Redis lock. if not force_due and is_poll_not_due(owner, now=current): await db.commit() await register_poll_active( owner, check_at=owner.next_poll_at, next_poll_at=owner.next_poll_at, reason="poll_task_not_due", ) return final_poll = is_final_poll_due(owner, now=current) if not owner_provider_task_id(owner): await _mark_failed( db, owner, message="任务轮询超时" if final_poll else "缺少外部任务ID", stage=ChatGenerationPipelineStage.TIMEOUT if final_poll else ChatGenerationPipelineStage.FAILED, event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value if final_poll else ChatGenerationTaskEventType.POLL_FAILED.value, ) return owner.pipeline_stage = _stage(owner, ChatGenerationPipelineStage.POLLING) owner.poll_claim_token = lease.token owner.poll_lease_until = _poll_lease_until(current) owner.poll_count = int(owner.poll_count or 0) + 1 owner.last_poll_at = current await db.commit() claim_started = True await register_poll_active( owner, check_at=owner.poll_lease_until, next_poll_at=owner.next_poll_at, reason="polling_lease", ) try: poll_result = await poll_provider_task(db, owner) await lease.ensure_owned() owner = await load_generation_owner_for_update_retry( db, owner_type=normalized_owner_type, owner_id=task_id, ) if not owner or not is_attempt_current(owner, effective_attempt): await remove_poll_active( owner_type=normalized_owner_type, owner_id=task_id, attempt_no=effective_attempt, ) return if owner.poll_claim_token != lease.token: return status = poll_result.get("status") response_data = poll_result.get("response_data") try: provider_response = json.loads(response_data or "{}") except Exception: provider_response = {"raw": response_data} if _is_success(status): if owner.gen_type == GenerationType.IMAGE.value: owner.remote_result_url = poll_result.get("image_url") owner.image_tokens_used = int( poll_result.get("image_tokens", 0) or 0 ) else: owner.remote_result_url = poll_result.get("video_url") owner.video_tokens_used = int( poll_result.get("video_tokens", 0) or 0 ) owner.provider_response_json = response_data await _sync_snapshot(db, owner, response_data) if not owner.remote_result_url: await _mark_failed( db, owner, message="供应商任务成功但未返回结果URL", stage=ChatGenerationPipelineStage.FAILED, event_type=ChatGenerationTaskEventType.POLL_FAILED.value, detail=poll_result, ) await _log_poll_provider_call_after_commit( owner, provider_response=provider_response ) return owner.pipeline_stage = _stage( owner, ChatGenerationPipelineStage.RESULT_READY ) owner.poll_error_count = 0 owner.poll_claim_token = None owner.poll_lease_until = None owner.next_poll_at = None await db.commit() await _log_poll_provider_call_after_commit( owner, provider_response=provider_response ) await remove_poll_active(owner) await log_task_event( owner, event_type=ChatGenerationTaskEventType.POLL_SUCCESS.value, to_stage=owner.pipeline_stage, ) from app.tasks.generation_download_tasks import ( enqueue_download_task, ) await enqueue_download_task( db, owner, reason="poll_success_result_ready" ) return if _is_failed(status): owner.provider_response_json = response_data await _mark_failed( db, owner, message=poll_result.get("error") or f"供应商任务失败: {status}", stage=ChatGenerationPipelineStage.FAILED, event_type=ChatGenerationTaskEventType.POLL_FAILED.value, detail=poll_result, ) await _log_poll_provider_call_after_commit( owner, provider_response=provider_response ) return if final_poll: await _mark_failed( db, owner, message="任务轮询超时", stage=ChatGenerationPipelineStage.TIMEOUT, event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value, detail=poll_result, ) await _log_poll_provider_call_after_commit( owner, provider_response=provider_response ) return owner.poll_error_count = 0 await _schedule_next_poll( db, owner, reason="poll_pending_next" ) await _log_poll_provider_call_after_commit( owner, provider_response=provider_response ) await log_task_event( owner, event_type=ChatGenerationTaskEventType.POLL_PENDING.value, message=f"status={status}", ) except (RedisExecutionLockError, DatabaseRowLockBusy): try: await db.rollback() except Exception: pass raise except Exception as exc: try: await db.rollback() except Exception: pass await lease.ensure_owned() owner = await load_generation_owner_for_update_retry( db, owner_type=normalized_owner_type, owner_id=task_id, ) if not owner or not is_attempt_current(owner, effective_attempt): await remove_poll_active( owner_type=normalized_owner_type, owner_id=task_id, attempt_no=effective_attempt, ) return if owner.poll_claim_token != lease.token: return if is_final_poll_due(owner, now=_now()): await _mark_failed( db, owner, message="任务轮询超时", stage=ChatGenerationPipelineStage.TIMEOUT, event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value, ) return owner.poll_error_count = int(owner.poll_error_count or 0) + 1 if is_video_generation_task(owner): await _schedule_next_poll( db, owner, reason="poll_exception_retry" ) return if owner.poll_error_count > int( settings.CHATAPI_ASYNC_MAX_RETRIES or 3 ): error_message = ( extract_error_message(exc, "轮询") if callable(extract_error_message) else str(exc) ) await _mark_failed( db, owner, message=error_message, stage=ChatGenerationPipelineStage.FAILED, event_type=ChatGenerationTaskEventType.POLL_FAILED.value, ) else: delay = int( settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30 ) * owner.poll_error_count await _schedule_next_poll( db, owner, reason="poll_exception_retry", default_delay_seconds=delay, ) except (RedisExecutionLockError, DatabaseRowLockBusy): if claim_started: try: await _restore_after_lock_error( owner_type=normalized_owner_type, owner_id=task_id, attempt_no=effective_attempt, claim_token=lease.token, ) except Exception: logger.exception( "轮询执行锁异常后的状态恢复失败 owner_type=%s owner_id=%s attempt=%s", normalized_owner_type, task_id, effective_attempt, ) raise finally: await lease.close() if celery_app: @celery_app.task( name="generation.poll_generation_task", bind=True, max_retries=3, default_retry_delay=30, ) def poll_generation_task( self, task_id: str, force_due: bool = False, 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, force_due=bool(force_due), ) ) 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") poll_generation_task = _DisabledTask()