from app.tasks.async_runner import run_async import json from datetime import datetime, timedelta, timezone from typing import Any, Optional 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, GenerationMode, GenerationType, ) 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.poll_schedule_service import ensure_video_poll_fields 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.services.media_token_usage_snapshot_service import sync_chat_generation_task_media_token_snapshot from app.services.redis_registry_service import ensure_aware_utc from app.tasks.celery_app import celery_app def _now() -> datetime: return datetime.now(timezone.utc) 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() generation_mode = _to_clean_str(getattr(task, "generation_mode", None)) or "" # 爆款开头复刻第5步的视频生成,original_prompt 已经是视频提词 JSON schema。 # 不能再追加“时长/比例/分辨率”中文参数,否则会污染 schema。 if generation_mode in {GenerationMode.HOT_OPENING_REPLICATE.value, GenerationMode.SHOT_REPLICATE.value} and gen_type == GenerationType.VIDEO.value: stripped = base_prompt.strip() if stripped.startswith("{") or stripped.startswith("["): return base_prompt 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 == GenerationType.VIDEO.value: # 时长:4秒,画面比例:16:9,分辨率:480p if duration: parts.append(f"时长:{duration}秒") parts.append(f"画面比例:{aspect_ratio}") parts.append(f"分辨率:{resolution}") else: parts.append("时长:4秒") parts.append("画面比例:16:9") parts.append("分辨率:480p") elif gen_type == GenerationType.IMAGE.value: if image_size: parts.append(f"分辨率:{image_size}") parts.append(f"画布比例:{image_proportion}") parts.append(f"像素尺寸:{image_px}") else: parts.append("分辨率:2K") parts.append("画布比例:1:1") parts.append("像素尺寸: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() is_image_main = bool( task and task.generation_mode == GenerationMode.CHATAPI_MAIN.value and task.gen_type == GenerationType.IMAGE.value and int(task.generation_count or 1) > 1 ) if not task or (task.generation_mode not in ALLOWED_GENERATION_MODES and not is_image_main): return if task.status != ChatGenerationTaskStatus.GENERATING.value: return deadline_at = ensure_aware_utc(task.deadline_at) # 图片 main 的 deadline 与 provider claim 由 image_batch_service 原子处理, # 避免重复 Celery 消息在有效租约期间把正在执行的批次错误退款。 if not is_image_main and deadline_at and datetime.now(timezone.utc) > deadline_at: await mark_chat_generation_task_failed_and_refund_once( db, task=task, error_message="任务超时", pipeline_stage=ChatGenerationPipelineStage.TIMEOUT.value, ) await db.commit() await log_task_event( task, event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value, to_status=ChatGenerationTaskStatus.FAILED.value, to_stage=ChatGenerationPipelineStage.TIMEOUT.value, ) 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) await db.commit() return if task.pipeline_stage not in ( ChatGenerationPipelineStage.QUEUED.value, ChatGenerationPipelineStage.PREPARING.value, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value, ): return try: if not task.optimized_prompt: old_stage = task.pipeline_stage task.pipeline_stage = ChatGenerationPipelineStage.PREPARING.value await db.commit() await log_task_event( task, event_type=ChatGenerationTaskEventType.PROMPT_CONCAT_START.value, from_stage=old_stage, to_stage=ChatGenerationPipelineStage.PREPARING.value, 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=ChatGenerationTaskEventType.PROMPT_CONCAT_SUCCESS.value, to_stage=ChatGenerationPipelineStage.PREPARING.value, detail={ "optimized_prompt": optimized_prompt, "gen_type": task.gen_type, "message": "已完成本地提示词拼接,未调用提词优化API", }, ) if is_image_main: from app.services.generation.ai.image_batch_service import run_image_main_batch await run_image_main_batch(db, task) return if task.seedance_task_id or task.provider_task_id: task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value if task.gen_type == GenerationType.VIDEO.value: ensure_video_poll_fields(task, now=_now()) task.next_poll_at = _now() await db.commit() elif task.remote_result_url: task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value 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 = ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value await db.commit() await log_task_event( task, event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_START.value, from_stage=old_stage, to_stage=ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value, ) 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 == GenerationType.IMAGE.value: 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, ) await sync_chat_generation_task_media_token_snapshot(db, task, provider_response=task.provider_response_json) if task.remote_result_url and not task.seedance_task_id: # 同步图片路径:原 SDK 已经返回最终 URL。 task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value else: # 视频路径:provider 返回 task id,后续轮询。 task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value if task.gen_type == GenerationType.VIDEO.value: current_time = _now() ensure_video_poll_fields(task, now=current_time) task.next_poll_at = current_time task.poll_interval_seconds = int(task.poll_interval_seconds or 0) task.status = ChatGenerationTaskStatus.GENERATING.value await db.commit() await log_task_event( task, event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_SUCCESS.value, to_stage=task.pipeline_stage, detail=created, ) if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value: 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, register_poll_active await register_poll_active( task, reason="create_provider_success", check_at=_now() + timedelta(seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300)), next_poll_at=getattr(task, "next_poll_at", None), ) poll_generation_task.apply_async( args=[task.id], queue=CeleryQueue.GEN_PROVIDER_POLL.value, countdown=0, ) 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) if is_image_main: # image_batch_service 负责供应商/拆分失败退款。若 child 已落库, # 顶层兜底绝不能再把 main 退款。 child_result = await db.execute( select(ChatGenerationTask.id).where( ChatGenerationTask.parent_task_id == task.id, ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value, ).limit(1) ) has_children = child_result.scalar_one_or_none() is not None if not has_children: task.provider_create_claim_token = None task.provider_create_lease_until = None await mark_chat_generation_task_failed_and_refund_once( db, task=task, error_message=error_message, pipeline_stage=ChatGenerationPipelineStage.FAILED.value, ) else: from app.services.generation.ai.task_group_service import aggregate_main_task_status await aggregate_main_task_status(db, parent_task_id=str(task.id)) else: await mark_chat_generation_task_failed_and_refund_once( db, task=task, error_message=error_message, pipeline_stage=ChatGenerationPipelineStage.FAILED.value, ) await db.commit() await log_task_event(task, event_type=ChatGenerationTaskEventType.TASK_FAILED.value, message=error_message) 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) await db.commit() 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): try: return run_async(_run(task_id)) except Exception as exc: # 只处理 run_async/连接池/worker 中断等基础设施异常;业务异常已在 _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") chatapi_create_generation_task = _DisabledTask()