from app.tasks.async_runner import run_async import json from datetime import datetime, timezone 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.tasks.celery_app import celery_app 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 {} 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 _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 or task.generation_mode != "chatapi_async": return # 只处理正在生成,且处于远程等待/轮询中的任务。 if task.status != "generating" or task.pipeline_stage not in ("waiting_remote", "polling"): return if task.deadline_at and datetime.now(timezone.utc) > task.deadline_at: await mark_chat_generation_task_failed_and_refund_once( db, task=task, error_message="任务轮询超时", pipeline_stage="timeout", ) await db.commit() await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout") return if not (task.seedance_task_id or task.provider_task_id): await mark_chat_generation_task_failed_and_refund_once( db, task=task, error_message="缺少外部任务ID", pipeline_stage="failed", ) await db.commit() await log_task_event(task, event_type="POLL_FAILED", message=task.error_message) return # 标记本次正在轮询。 # 注意:pending 后会再改回 waiting_remote,避免任务长期卡在 polling。 task.pipeline_stage = "polling" task.poll_count = (task.poll_count or 0) + 1 task.last_poll_at = datetime.now(timezone.utc) await db.commit() 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 if not task.remote_result_url: await mark_chat_generation_task_failed_and_refund_once( db, task=task, error_message="供应商任务成功但未返回结果URL", pipeline_stage="failed", ) await db.commit() await log_task_event(task, event_type="POLL_FAILED", message=task.error_message) return task.pipeline_stage = "result_ready" task.retry_count = 0 await db.commit() await log_task_event(task, event_type="POLL_SUCCESS", to_stage="result_ready") from app.tasks.generation_download_tasks import download_generation_result_task download_generation_result_task.delay(task.id) return if _is_failed(status): task.provider_response_json = response_data await mark_chat_generation_task_failed_and_refund_once( db, task=task, error_message=poll_result.get("error") or f"供应商任务失败: {status}", pipeline_stage="failed", ) await db.commit() await log_task_event(task, event_type="POLL_FAILED", message=task.error_message, detail=poll_result) return # 关键修改 1: # 供应商仍在 pending / running 时,把阶段从 polling 改回 waiting_remote。 # 这样数据库状态表示“等待下一次轮询”,不会长期停在 polling。 # 同时可以降低重复 Celery 消息形成多条轮询链的概率。 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}") poll_generation_task.apply_async( args=[task.id], countdown=settings.CHATAPI_ASYNC_POLL_INTERVAL_SECONDS, ) except Exception as exc: # 关键修改 2: # 异常后先 rollback,再重新查询 task,不继续使用 rollback 前的旧 ORM 对象。 try: await db.rollback() except Exception: pass task = await _reload_task(db, task_id) if not task: 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_chat_generation_task_failed_and_refund_once( db, task=task, error_message=error_message, pipeline_stage="failed", ) await db.commit() await log_task_event(task, event_type="POLL_FAILED", message=task.error_message) else: # 临时轮询异常时,不让任务停在 polling。 # 回到 waiting_remote,等待下一次重试轮询。 task.pipeline_stage = "waiting_remote" await db.commit() poll_generation_task.apply_async( args=[task.id], countdown=settings.CHATAPI_ASYNC_RETRY_BACKOFF_SECONDS * task.retry_count, ) 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): return run_async(_run(task_id)) 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()