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.enums.celery_queue import CeleryQueue from app.enums.generation_task import ( ALLOWED_GENERATION_MODES, ChatGenerationPipelineStage, ChatGenerationTaskEventType, ChatGenerationTaskStatus, 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_task_event, log_provider_call 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.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 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 _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: return is_final_poll_due(task, now=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), "poll_started_at": datetime_to_epoch(task.poll_started_at) if getattr(task, "poll_started_at", None) else None, "poll_interval_seconds": int(getattr(task, "poll_interval_seconds", 0) 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 from app.services.generation.ai.task_group_service import aggregate_parent_for_child await notify_chat_generation_task_finished(db, task) await aggregate_parent_for_child(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=ChatGenerationPipelineStage.TIMEOUT.value, ) task.next_poll_at = None await _notify_finished(db, task) await db.commit() await remove_poll_active(task.id) await log_task_event( task, event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value, to_status=ChatGenerationTaskStatus.FAILED.value, to_stage=ChatGenerationPipelineStage.TIMEOUT.value, ) 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=ChatGenerationPipelineStage.FAILED.value, ) task.next_poll_at = None await _notify_finished(db, task) await db.commit() await remove_poll_active(task.id) await log_task_event( task, event_type=ChatGenerationTaskEventType.POLL_FAILED.value, message=task.error_message, detail=detail, ) async def _skip_not_due(task: ChatGenerationTask) -> None: next_poll_at = ensure_aware_utc(task.next_poll_at) if next_poll_at is None: return await register_poll_active( task, check_at=next_poll_at, next_poll_at=next_poll_at, reason="poll_task_not_due", ) await log_task_event( task, event_type=ChatGenerationTaskEventType.POLL_SKIP_NOT_DUE.value, message="视频任务尚未到下一次轮询时间,本次 poll 跳过", detail={"next_poll_at": next_poll_at, "pipeline_stage": task.pipeline_stage}, ) async def _schedule_next_poll( task: ChatGenerationTask, *, reason: str, default_delay_seconds: int | None = None, ) -> None: current_time = _now() if is_video_generation_task(task): schedule = build_video_pending_poll_schedule(task, now=current_time) else: schedule = build_default_poll_schedule( task, now=current_time, delay_seconds=default_delay_seconds, reason=reason, ) task.next_poll_at = schedule.next_poll_at task.poll_interval_seconds = schedule.poll_interval_seconds await log_task_event( task, event_type=ChatGenerationTaskEventType.POLL_SCHEDULED.value, message=f"已登记下一次轮询。reason={schedule.reason}", detail={ "delay_seconds": schedule.delay_seconds, "next_poll_at": schedule.next_poll_at, "direct_countdown": schedule.direct_countdown, "poll_interval_seconds": schedule.poll_interval_seconds, "source_reason": reason, }, ) await register_poll_active( task, check_at=schedule.next_poll_at, next_poll_at=schedule.next_poll_at, reason=schedule.reason, ) if schedule.direct_countdown: poll_generation_task.apply_async( args=[task.id], queue=POLL_QUEUE, countdown=max(0, int(schedule.delay_seconds)), ) async def _run(task_id: str, *, force_due: bool = False): 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 != ChatGenerationTaskStatus.GENERATING.value or task.pipeline_stage not in ( ChatGenerationPipelineStage.WAITING_REMOTE.value, ChatGenerationPipelineStage.POLLING.value, ): await remove_poll_active(task.id) return current_time = _now() if is_video_generation_task(task): ensure_video_poll_fields(task, now=current_time) # dispatcher / recovery 已经在投递前确认到期时,会传 force_due=True。 # 这样可以避免投递侧为了防重复消费临时写入的 next_poll_at, # 又被当前 worker 当成“业务下一次轮询时间”而误判未到期。 if not force_due and is_poll_not_due(task, now=current_time): await db.commit() await _skip_not_due(task) return await db.commit() final_poll_before_timeout = _deadline_expired(task, current_time) 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=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT.value, message="任务已到 deadline,执行最后一次供应商查询后再判定超时", detail={"deadline_at": task.deadline_at, "stage": task.pipeline_stage}, ) # 标记本次正在轮询,并登记 poll lease。 # 如果 worker 在供应商接口调用过程中退出,启动恢复会在 lease 过期后重新投递。 task.pipeline_stage = ChatGenerationPipelineStage.POLLING.value task.poll_count = (task.poll_count or 0) + 1 task.last_poll_at = _now() task.next_poll_at = _poll_lease_until(task.last_poll_at) await db.commit() await register_poll_active( task, check_at=_poll_lease_until(task.last_poll_at), reason="polling_lease", next_poll_at=task.next_poll_at, ) 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 == GenerationType.IMAGE.value: 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 = ChatGenerationPipelineStage.RESULT_READY.value task.retry_count = 0 task.next_poll_at = None await db.commit() await remove_poll_active(task.id) await log_task_event(task, event_type=ChatGenerationTaskEventType.POLL_SUCCESS.value, to_stage=ChatGenerationPipelineStage.RESULT_READY.value) 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=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT_PENDING.value, message=f"最终查询后供应商仍未完成,按超时处理。status={status}", detail=poll_result, ) await _mark_timeout(db, task, message="任务轮询超时") return # 供应商仍在 pending / running 时,把阶段从 polling 改回 waiting_remote。 # 视频任务写入 next_poll_at,由 Beat dispatcher 到期投递;短间隔可保留 countdown 兼容。 task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value task.retry_count = 0 await _schedule_next_poll(task, reason="poll_pending_next") await db.commit() await log_task_event( task, event_type=ChatGenerationTaskEventType.POLL_PENDING.value, message=f"status={status}", detail={"next_poll_at": task.next_poll_at, "poll_interval_seconds": task.poll_interval_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=ChatGenerationTaskEventType.FINAL_POLL_BEFORE_TIMEOUT_ERROR.value, message=str(exc), ) await _mark_timeout(db, task, message="任务轮询超时") return task.retry_count = (task.retry_count or 0) + 1 # 视频轮询的临时异常不再 3 次内直接退款;继续降频到 24 小时最终 deadline。 if is_video_generation_task(task): task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value await _schedule_next_poll(task, reason="poll_exception_retry") await db.commit() return 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 = ChatGenerationPipelineStage.WAITING_REMOTE.value delay_seconds = int(settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS or 30) * int(task.retry_count or 1) default_schedule = build_default_poll_schedule( task, now=_now(), delay_seconds=delay_seconds, reason="poll_exception_retry", ) task.next_poll_at = default_schedule.next_poll_at await db.commit() await register_poll_active( task, check_at=_poll_check_at(delay_seconds=delay_seconds), next_poll_at=default_schedule.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, force_due: bool = False): try: return run_async(_run(task_id, force_due=bool(force_due))) 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()