from app.tasks.async_runner import run_async import json from datetime import datetime, timezone from sqlalchemy import select 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_billing_service import charge_chatapi_prompt_usage from app.services.generation_log_service import log_task_event from app.services.generation_prompt_service import build_prompt_with_chatapi from app.services.generation_provider_service import create_provider_task from app.tasks.celery_app import celery_app async def _run(task_id: str): async with async_session() as db: result = await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id == task_id)) task = result.scalar_one_or_none() if not task or task.generation_mode != "chatapi_async": return if task.status != "generating": return if task.deadline_at and datetime.now(timezone.utc) > task.deadline_at: task.status = "failed" task.pipeline_stage = "timeout" task.error_message = "任务超时" await db.commit() await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout") return if task.pipeline_stage not in ("queued", "preparing", "creating_provider_task"): return try: if not task.optimized_prompt: old_stage = task.pipeline_stage task.pipeline_stage = "preparing" await db.commit() await log_task_event(task, event_type="CHATAPI_START", from_stage=old_stage, to_stage="preparing") optimized, usage = await build_prompt_with_chatapi(db, task) await charge_chatapi_prompt_usage(db, record=task, usage=usage, project_name="AI生成任务") task.optimized_prompt = optimized task.text_tokens_used = usage["total_tokens"] await db.commit() await log_task_event(task, event_type="CHATAPI_SUCCESS", to_stage="preparing", detail=usage) if task.seedance_task_id or task.provider_task_id: task.pipeline_stage = "waiting_remote" await db.commit() elif task.remote_result_url: task.pipeline_stage = "result_ready" await db.commit() from app.tasks.generation_download_tasks import download_generation_result_task download_generation_result_task.delay(task.id) return else: old_stage = task.pipeline_stage task.pipeline_stage = "creating_provider_task" await db.commit() await log_task_event(task, event_type="PROVIDER_CREATE_START", from_stage=old_stage, to_stage="creating_provider_task") created = await create_provider_task(db, task) provider_task_id = created.get("task_id") if provider_task_id: task.provider_task_id = provider_task_id task.seedance_task_id = provider_task_id task.remote_result_url = created.get("remote_result_url") or task.remote_result_url if task.gen_type == "image": task.image_tokens_used = created.get("image_tokens", task.image_tokens_used or 0) or 0 task.provider_response_json = json.dumps(created.get("response_data") or {}, ensure_ascii=False, default=str) if task.remote_result_url and not task.seedance_task_id: # Sync image path: original SDK already returned final URL. task.pipeline_stage = "result_ready" else: # Video path: provider returns task id, poll later. task.pipeline_stage = "waiting_remote" task.status = "generating" await db.commit() await log_task_event(task, event_type="PROVIDER_CREATE_SUCCESS", to_stage=task.pipeline_stage, detail=created) if task.pipeline_stage == "result_ready": from app.tasks.generation_download_tasks import download_generation_result_task download_generation_result_task.delay(task.id) else: from app.tasks.generation_poll_tasks import poll_generation_task poll_generation_task.delay(task.id) except Exception as exc: try: await db.rollback() except Exception: pass result = await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id == task_id)) task = result.scalar_one_or_none() if task: task.status = "failed" task.error_message = extract_error_message(exc, "生成任务") if callable(extract_error_message) else str(exc) await db.commit() await log_task_event(task, event_type="TASK_FAILED", message=task.error_message) 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): 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") chatapi_create_generation_task = _DisabledTask()