from app.tasks.async_runner import run_async import json from datetime import datetime, timedelta, timezone from typing import Any from sqlalchemy import select from app.config import settings 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, log_provider_call from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once 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 from app.services.redis_registry_service import ( datetime_to_epoch, ensure_aware_utc, redis_remove_registry_item, redis_upsert_registry_item, utc_now, ) from app.tasks.celery_app import celery_app ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"} POLL_QUEUE = "gen_provider_poll" def _now() -> datetime: return datetime.now(timezone.utc) def _is_success(status: str) -> bool: return status in ("succeeded", "success", "completed", "done") def _is_failed(status: str) -> bool: return status in ("failed", "error", "canceled", "cancelled") def _engine_snapshot(task: ChatGenerationTask) -> dict: try: return json.loads(task.engine_snapshot_json or "{}") except Exception: return {} def _deadline_expired(task: ChatGenerationTask, now: datetime | None = None) -> bool: deadline_at = ensure_aware_utc(task.deadline_at) return bool(deadline_at and deadline_at <= (now or _now())) def _poll_check_at(*, delay_seconds: int | float | None = None, now: datetime | None = None) -> datetime: current_time = now or _now() delay = int(delay_seconds or settings.CHATAPI_ASYNC_POLL_INTERVAL_SECONDS or 30) grace = int(settings.POLL_TASK_QUEUE_TIMEOUT_SECONDS or 120) return current_time + timedelta(seconds=max(1, delay) + max(0, grace)) def _poll_lease_until(now: datetime | None = None) -> datetime: current_time = now or _now() return current_time + timedelta(seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300)) def _build_poll_active_payload( task: ChatGenerationTask, *, stage: str, reason: str, next_poll_at: datetime | None = None, check_at: datetime | None = None, ) -> dict[str, Any]: current_time = utc_now() checked_next_poll_at = ensure_aware_utc(next_poll_at) checked_check_at = ensure_aware_utc(check_at) return { "task_id": task.id, "provider_task_id": task.provider_task_id, "seedance_task_id": task.seedance_task_id, "generation_mode": task.generation_mode, "gen_type": task.gen_type, "stage": stage, "queue": POLL_QUEUE, "poll_count": int(task.poll_count or 0), "retry_count": int(task.retry_count or 0), "last_poll_at": datetime_to_epoch(task.last_poll_at) if task.last_poll_at else None, "next_poll_at": datetime_to_epoch(checked_next_poll_at) if checked_next_poll_at else None, "deadline_at": datetime_to_epoch(task.deadline_at) if task.deadline_at else None, "check_at": datetime_to_epoch(checked_check_at) if checked_check_at else None, "updated_at": datetime_to_epoch(current_time), "reason": reason, } async def register_poll_active( task: ChatGenerationTask, *, check_at: datetime, reason: str, next_poll_at: datetime | None = None, ) -> None: payload = _build_poll_active_payload( task, stage=task.pipeline_stage or "", reason=reason, next_poll_at=next_poll_at, check_at=check_at, ) await redis_upsert_registry_item( hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY, zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY, item_id=task.id, payload=payload, check_at=check_at, log_context="poll_active", ) async def remove_poll_active(task_id: str) -> None: await redis_remove_registry_item( hash_key=settings.POLL_ACTIVE_REDIS_HASH_KEY, zset_key=settings.POLL_ACTIVE_REDIS_ZSET_KEY, item_id=task_id, log_context="poll_active", ) async def _notify_finished(db, task: ChatGenerationTask) -> None: from app.services.generation_module_hook_service import notify_chat_generation_task_finished await notify_chat_generation_task_finished(db, task) async def _reload_task(db, task_id: str) -> ChatGenerationTask | None: """ rollback 后重新查询任务对象。 说明: - SQLAlchemy rollback 后,当前 ORM 对象可能进入过期状态。 - 后续继续访问旧 task.retry_count / task.id 等字段,有概率触发异步懒加载异常。 - 所以 poll/download 的异常分支统一 rollback 后重新 select。 """ result = await db.execute( select(ChatGenerationTask).where( ChatGenerationTask.id == task_id, ChatGenerationTask.deleted_at.is_(None), ).with_for_update().limit(1) ) return result.scalar_one_or_none() async def _mark_timeout(db, task: ChatGenerationTask, *, message: str = "任务轮询超时") -> None: await mark_chat_generation_task_failed_and_refund_once( db, task=task, error_message=message, pipeline_stage="timeout", ) await _notify_finished(db, task) await db.commit() await remove_poll_active(task.id) await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout") async def _mark_failed(db, task: ChatGenerationTask, *, message: str, detail: Any = None) -> None: await mark_chat_generation_task_failed_and_refund_once( db, task=task, error_message=message, pipeline_stage="failed", ) await _notify_finished(db, task) await db.commit() await remove_poll_active(task.id) await log_task_event(task, event_type="POLL_FAILED", message=task.error_message, detail=detail) async def _run(task_id: str): async with async_session() as db: result = await db.execute(select(ChatGenerationTask).where( ChatGenerationTask.id == task_id, ChatGenerationTask.deleted_at.is_(None), ).with_for_update().limit(1)) task = result.scalar_one_or_none() if not task: await remove_poll_active(task_id) return if task.generation_mode not in ALLOWED_GENERATION_MODES: await remove_poll_active(task.id) return # 只处理正在生成,且处于远程等待/轮询中的任务。 if task.status != "generating" or task.pipeline_stage not in ("waiting_remote", "polling"): await remove_poll_active(task.id) return final_poll_before_timeout = _deadline_expired(task) if final_poll_before_timeout and not (task.seedance_task_id or task.provider_task_id): await _mark_timeout(db, task, message="任务轮询超时") return if not (task.seedance_task_id or task.provider_task_id): await _mark_failed(db, task, message="缺少外部任务ID") return if final_poll_before_timeout: await log_task_event( task, event_type="FINAL_POLL_BEFORE_TIMEOUT", message="任务已到 deadline,执行最后一次供应商查询后再判定超时", detail={"deadline_at": task.deadline_at, "stage": task.pipeline_stage}, ) # 标记本次正在轮询,并登记 poll lease。 # 如果 worker 在供应商接口调用过程中退出,启动恢复会在 lease 过期后重新投递。 task.pipeline_stage = "polling" task.poll_count = (task.poll_count or 0) + 1 task.last_poll_at = _now() await db.commit() await register_poll_active( task, check_at=_poll_lease_until(task.last_poll_at), reason="polling_lease", ) try: poll_result = await poll_provider_task(db, task) 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} snapshot = _engine_snapshot(task) await log_provider_call( task, provider=snapshot.get("provider") or "ark", api_type=f"{task.gen_type}_poll", model=snapshot.get("model_name"), engine_id=task.engine_id, status="success", provider_task_id=task.seedance_task_id or task.provider_task_id, response_data=provider_response, ) if _is_success(status): if task.gen_type == "image": task.remote_result_url = poll_result.get("image_url") task.image_tokens_used = poll_result.get("image_tokens", 0) or 0 else: task.remote_result_url = poll_result.get("video_url") task.video_tokens_used = poll_result.get("video_tokens", 0) or 0 task.provider_response_json = response_data await sync_chat_generation_task_media_token_snapshot(db, task, provider_response=response_data) if not task.remote_result_url: await _mark_failed(db, task, message="供应商任务成功但未返回结果URL", detail=poll_result) return task.pipeline_stage = "result_ready" task.retry_count = 0 await db.commit() await remove_poll_active(task.id) await log_task_event(task, event_type="POLL_SUCCESS", to_stage="result_ready") from app.tasks.generation_download_tasks import enqueue_download_task await enqueue_download_task(db, task, reason="poll_success_result_ready") return if _is_failed(status): task.provider_response_json = response_data await _mark_failed( db, task, message=poll_result.get("error") or f"供应商任务失败: {status}", detail=poll_result, ) return if final_poll_before_timeout: await log_task_event( task, event_type="FINAL_POLL_BEFORE_TIMEOUT_PENDING", message=f"最终查询后供应商仍未完成,按超时处理。status={status}", detail=poll_result, ) await _mark_timeout(db, task, message="任务轮询超时") return # 供应商仍在 pending / running 时,把阶段从 polling 改回 waiting_remote。 # 同时登记下一次 poll active,Celery countdown 丢失时可由恢复任务拉起。 task.pipeline_stage = "waiting_remote" task.retry_count = 0 await db.commit() await log_task_event(task, event_type="POLL_PENDING", message=f"status={status}") delay_seconds = int(settings.CHATAPI_ASYNC_POLL_INTERVAL_SECONDS or 30) next_poll_at = _now() + timedelta(seconds=max(1, delay_seconds)) await register_poll_active( task, check_at=_poll_check_at(delay_seconds=delay_seconds), next_poll_at=next_poll_at, reason="poll_pending_next", ) poll_generation_task.apply_async( args=[task.id], queue=POLL_QUEUE, countdown=delay_seconds, ) except Exception as exc: # 异常后先 rollback,再重新查询 task,不继续使用 rollback 前的旧 ORM 对象。 try: await db.rollback() except Exception: pass task = await _reload_task(db, task_id) if not task: await remove_poll_active(task_id) return if final_poll_before_timeout: await log_task_event( task, event_type="FINAL_POLL_BEFORE_TIMEOUT_ERROR", message=str(exc), ) await _mark_timeout(db, task, message="任务轮询超时") return task.retry_count = (task.retry_count or 0) + 1 if task.retry_count > settings.CHATAPI_ASYNC_MAX_RETRIES: error_message = extract_error_message(exc, "轮询") if callable(extract_error_message) else str(exc) await _mark_failed(db, task, message=error_message) else: # 临时轮询异常时,不让任务停在 polling。 # 回到 waiting_remote,等待下一次重试轮询。 task.pipeline_stage = "waiting_remote" await db.commit() delay_seconds = int(settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30) * int(task.retry_count or 1) next_poll_at = _now() + timedelta(seconds=max(1, delay_seconds)) await register_poll_active( task, check_at=_poll_check_at(delay_seconds=delay_seconds), next_poll_at=next_poll_at, reason="poll_exception_retry", ) poll_generation_task.apply_async( args=[task.id], queue=POLL_QUEUE, countdown=delay_seconds, ) 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): try: return run_async(_run(task_id)) except Exception as exc: # 只重试基础设施异常;供应商失败/业务失败已在 _run 内处理。 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()