diff --git a/video-gen-api/app/services/generation_billing_service.py b/video-gen-api/app/services/generation_billing_service.py index daf77743..d8fa5ad0 100644 --- a/video-gen-api/app/services/generation_billing_service.py +++ b/video-gen-api/app/services/generation_billing_service.py @@ -170,7 +170,7 @@ async def deduct_credits_locked_once( type="consume", amount=-amount, balance_after=user.credits, - description=f"[{charge_key}] {description}", + description=description, related_id=related_id, ) ) @@ -203,7 +203,7 @@ async def charge_chatapi_prompt_usage( db, user_id=record.user_id, amount=text_credits, - description=f"ChatAPI提示词整理 - {project_name}", + description=f"ChatAPI提示词整理", related_id=record.id, charge_key=CHARGE_TEXT_PROMPT, ) @@ -220,7 +220,7 @@ async def charge_chatapi_prompt_usage( db, user_id=record.user_id, amount=file_parse_credits, - description=f"文件解析Token - {project_name}", + description=f"文件解析Token", related_id=record.id, charge_key=CHARGE_FILE_PARSE, ) @@ -237,7 +237,7 @@ async def charge_chatapi_prompt_usage( db, user_id=record.user_id, amount=vision_input_credits, - description=f"图片理解Token - {project_name}", + description=f"图片理解Token", related_id=record.id, charge_key=CHARGE_VISION_INPUT, ) @@ -278,7 +278,7 @@ async def charge_generation_media_by_params( db, user_id=user_id, amount=amount, - description=f"{description_prefix}图片生成 - {project_name}", + description=f"{description_prefix}图片生成", related_id=record_id, charge_key=CHARGE_MEDIA_IMAGE, ) @@ -290,7 +290,7 @@ async def charge_generation_media_by_params( db, user_id=user_id, amount=amount, - description=f"{description_prefix}视频生成 - {project_name}", + description=f"{description_prefix}视频生成", related_id=record_id, charge_key=CHARGE_MEDIA_VIDEO, ) diff --git a/video-gen-api/app/tasks/generation_create_tasks.py b/video-gen-api/app/tasks/generation_create_tasks.py index 51c7108f..40a45370 100644 --- a/video-gen-api/app/tasks/generation_create_tasks.py +++ b/video-gen-api/app/tasks/generation_create_tasks.py @@ -1,27 +1,151 @@ 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_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 +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" @@ -29,6 +153,7 @@ async def _run(task_id: str): 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 @@ -37,62 +162,108 @@ async def _run(task_id: str): 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") + 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 - 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) + 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") + 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) + + 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. + # 同步图片路径:原 SDK 已经返回最终 URL。 task.pipeline_stage = "result_ready" else: - # Video path: provider returns task id, poll later. + # 视频路径: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) + 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) @@ -108,6 +279,8 @@ 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() + + chatapi_create_generation_task = _DisabledTask() \ No newline at end of file