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_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 / duration_seconds / video_duration # aspect_ratio / ratio # resolution # image_size / size / pixel_size duration = _get_first_value(task, "duration", "duration_seconds", "video_duration") aspect_ratio = _get_first_value(task, "aspect_ratio", "ratio") resolution = _get_first_value(task, "resolution") image_size = _get_first_value(task, "image_size", "size", "pixel_size") parts = [] if gen_type == "video": duration_text = _format_duration(duration) aspect_ratio_text = _to_clean_str(aspect_ratio) resolution_text = _to_clean_str(resolution) if duration_text: parts.append(f"时长:{duration_text}") if aspect_ratio_text: parts.append(f"画面比例:{aspect_ratio_text}") if resolution_text: parts.append(f"分辨率:{resolution_text}") elif gen_type == "image": resolution_text = _to_clean_str(resolution) aspect_ratio_text = _to_clean_str(aspect_ratio) image_size_text = _to_clean_str(image_size) if resolution_text: if resolution_text.startswith("分辨率"): parts.append(resolution_text) else: parts.append(f"分辨率{resolution_text}") if aspect_ratio_text: if aspect_ratio_text.startswith("画布比例"): parts.append(aspect_ratio_text) else: parts.append(f"画布比例{aspect_ratio_text}") if image_size_text: if image_size_text.startswith("像素尺寸"): parts.append(image_size_text) else: parts.append(f"像素尺寸{image_size_text}") else: # 未知类型时尽量保守拼接已有参数,避免直接丢失生成参数。 duration_text = _format_duration(duration) aspect_ratio_text = _to_clean_str(aspect_ratio) resolution_text = _to_clean_str(resolution) image_size_text = _to_clean_str(image_size) if duration_text: parts.append(f"时长:{duration_text}") if aspect_ratio_text: parts.append(f"画面比例:{aspect_ratio_text}") if resolution_text: parts.append(f"分辨率:{resolution_text}") if image_size_text: parts.append(f"像素尺寸{image_size_text}") 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)) 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="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 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: # 同步图片路径:原 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 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()