from app.tasks.async_runner import run_async import json from datetime import datetime, timezone from typing import Any, Optional 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_log_service import log_task_event from app.services.generation_refund_service import mark_chat_generation_task_failed_and_refund_once from app.services.generation_provider_service import create_provider_task from app.tasks.celery_app import celery_app def _get_first_value(obj: Any, *field_names: str) -> Optional[Any]: """ 兼容不同版本字段名,避免字段调整后 Celery 任务直接报错。 """ for field_name in field_names: if hasattr(obj, field_name): value = getattr(obj, field_name) if value is not None and str(value).strip() != "": return value return None def _to_clean_str(value: Any) -> Optional[str]: if value is None: return None text = str(value).strip() return text if text else None def _format_duration(value: Any) -> Optional[str]: """ duration=4 -> 4秒 duration="4秒" -> 4秒 """ text = _to_clean_str(value) if not text: return None lower_text = text.lower() if text.endswith("秒") or lower_text.endswith("s") or lower_text.endswith("sec") or lower_text.endswith("seconds"): return text return f"{text}秒" def _build_optimized_prompt_by_params(task: ChatGenerationTask) -> str: """ 不调用提词优化 API,直接将 original_prompt 拼接上对应类型的生成参数。 视频示例: original_prompt,时长:4秒,画面比例:16:9,分辨率:480p 图片示例: original_prompt,分辨率2K,画布比例1:1,像素尺寸2048×2048 """ original_prompt = _to_clean_str(getattr(task, "original_prompt", None)) or "" base_prompt = original_prompt.rstrip(",,。;; \n\t") gen_type = (_to_clean_str(getattr(task, "gen_type", None)) or "").lower() duration = _get_first_value(task, "duration") aspect_ratio = _get_first_value(task, "aspect_ratio") resolution = _get_first_value(task, "resolution") image_size = _get_first_value(task, "image_size") image_px = _get_first_value(task, "image_px") image_proportion = _get_first_value(task, "image_proportion") parts = [] if gen_type == "video": # 时长:4秒,画面比例:16:9,分辨率:480p if duration: parts.append(f"时长:{duration}秒") parts.append(f"画面比例:{aspect_ratio}") parts.append(f"分辨率:{resolution}") else: parts.append(f"时长:4秒") parts.append(f"画面比例:16:9") parts.append(f"分辨率:480p") elif gen_type == "image": if image_size : parts.append(f"分辨率:{image_size}") parts.append(f"画布比例:{image_proportion}") parts.append(f"像素尺寸:{image_px}") else: parts.append(f"分辨率:2K") parts.append(f"画布比例:1:1") parts.append(f"像素尺寸:2048x2048") else: # 未知类型时返回原始字符 return base_prompt suffix = ",".join(parts) if base_prompt and suffix: return f"{base_prompt},{suffix}" if base_prompt: return base_prompt return suffix 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": 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 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="PROMPT_CONCAT_START", from_stage=old_stage, to_stage="preparing", message="开始本地拼接提示词,不调用提词优化API", ) optimized_prompt = _build_optimized_prompt_by_params(task) task.optimized_prompt = optimized_prompt # 不调用提词优化 API,因此不产生模型 token 消耗。 task.text_tokens_used = task.text_tokens_used or 0 await db.commit() await log_task_event( task, event_type="PROMPT_CONCAT_SUCCESS", to_stage="preparing", detail={ "optimized_prompt": optimized_prompt, "gen_type": task.gen_type, "message": "已完成本地提示词拼接,未调用提词优化API", }, ) 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 enqueue_download_task await enqueue_download_task(db, task, reason="create_remote_result_ready") 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: # 同步图片路径:原 SDK 已经返回最终 URL。 task.pipeline_stage = "result_ready" else: # 视频路径:provider 返回 task id,后续轮询。 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 enqueue_download_task await enqueue_download_task(db, task, reason="create_result_ready") 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, ChatGenerationTask.deleted_at.is_(None), ).with_for_update().limit(1)) task = result.scalar_one_or_none() if task: 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="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()