CHAT生成任务不做提词优化|积分变动描述修订

This commit is contained in:
2026-05-27 17:59:12 +08:00
parent 56b9bdaef0
commit 5021cfd35b
2 changed files with 193 additions and 20 deletions
@@ -170,7 +170,7 @@ async def deduct_credits_locked_once(
type="consume", type="consume",
amount=-amount, amount=-amount,
balance_after=user.credits, balance_after=user.credits,
description=f"[{charge_key}] {description}", description=description,
related_id=related_id, related_id=related_id,
) )
) )
@@ -203,7 +203,7 @@ async def charge_chatapi_prompt_usage(
db, db,
user_id=record.user_id, user_id=record.user_id,
amount=text_credits, amount=text_credits,
description=f"ChatAPI提示词整理 - {project_name}", description=f"ChatAPI提示词整理",
related_id=record.id, related_id=record.id,
charge_key=CHARGE_TEXT_PROMPT, charge_key=CHARGE_TEXT_PROMPT,
) )
@@ -220,7 +220,7 @@ async def charge_chatapi_prompt_usage(
db, db,
user_id=record.user_id, user_id=record.user_id,
amount=file_parse_credits, amount=file_parse_credits,
description=f"文件解析Token - {project_name}", description=f"文件解析Token",
related_id=record.id, related_id=record.id,
charge_key=CHARGE_FILE_PARSE, charge_key=CHARGE_FILE_PARSE,
) )
@@ -237,7 +237,7 @@ async def charge_chatapi_prompt_usage(
db, db,
user_id=record.user_id, user_id=record.user_id,
amount=vision_input_credits, amount=vision_input_credits,
description=f"图片理解Token - {project_name}", description=f"图片理解Token",
related_id=record.id, related_id=record.id,
charge_key=CHARGE_VISION_INPUT, charge_key=CHARGE_VISION_INPUT,
) )
@@ -278,7 +278,7 @@ async def charge_generation_media_by_params(
db, db,
user_id=user_id, user_id=user_id,
amount=amount, amount=amount,
description=f"{description_prefix}图片生成 - {project_name}", description=f"{description_prefix}图片生成",
related_id=record_id, related_id=record_id,
charge_key=CHARGE_MEDIA_IMAGE, charge_key=CHARGE_MEDIA_IMAGE,
) )
@@ -290,7 +290,7 @@ async def charge_generation_media_by_params(
db, db,
user_id=user_id, user_id=user_id,
amount=amount, amount=amount,
description=f"{description_prefix}视频生成 - {project_name}", description=f"{description_prefix}视频生成",
related_id=record_id, related_id=record_id,
charge_key=CHARGE_MEDIA_VIDEO, charge_key=CHARGE_MEDIA_VIDEO,
) )
@@ -1,27 +1,151 @@
from app.tasks.async_runner import run_async from app.tasks.async_runner import run_async
import json import json
from datetime import datetime, timezone from datetime import datetime, timezone
from typing import Any, Optional
from sqlalchemy import select from sqlalchemy import select
from app.models.base import async_session from app.models.base import async_session
from app.models.chat_generation_task import ChatGenerationTask from app.models.chat_generation_task import ChatGenerationTask
from app.services.error_codes import extract_error_message 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_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.services.generation_provider_service import create_provider_task
from app.tasks.celery_app import celery_app 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 def _run(task_id: str):
async with async_session() as db: async with async_session() as db:
result = await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id == task_id)) result = await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id == task_id))
task = result.scalar_one_or_none() task = result.scalar_one_or_none()
if not task or task.generation_mode != "chatapi_async": if not task or task.generation_mode != "chatapi_async":
return return
if task.status != "generating": if task.status != "generating":
return return
if task.deadline_at and datetime.now(timezone.utc) > task.deadline_at: if task.deadline_at and datetime.now(timezone.utc) > task.deadline_at:
task.status = "failed" task.status = "failed"
task.pipeline_stage = "timeout" task.pipeline_stage = "timeout"
@@ -29,6 +153,7 @@ async def _run(task_id: str):
await db.commit() await db.commit()
await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout") await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout")
return return
if task.pipeline_stage not in ("queued", "preparing", "creating_provider_task"): if task.pipeline_stage not in ("queued", "preparing", "creating_provider_task"):
return return
@@ -37,62 +162,108 @@ async def _run(task_id: str):
old_stage = task.pipeline_stage old_stage = task.pipeline_stage
task.pipeline_stage = "preparing" task.pipeline_stage = "preparing"
await db.commit() 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 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: if task.seedance_task_id or task.provider_task_id:
task.pipeline_stage = "waiting_remote" task.pipeline_stage = "waiting_remote"
await db.commit() await db.commit()
elif task.remote_result_url: elif task.remote_result_url:
task.pipeline_stage = "result_ready" task.pipeline_stage = "result_ready"
await db.commit() await db.commit()
from app.tasks.generation_download_tasks import download_generation_result_task from app.tasks.generation_download_tasks import download_generation_result_task
download_generation_result_task.delay(task.id) download_generation_result_task.delay(task.id)
return return
else: else:
old_stage = task.pipeline_stage old_stage = task.pipeline_stage
task.pipeline_stage = "creating_provider_task" task.pipeline_stage = "creating_provider_task"
await db.commit() 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) created = await create_provider_task(db, task)
provider_task_id = created.get("task_id") provider_task_id = created.get("task_id")
if provider_task_id: if provider_task_id:
task.provider_task_id = provider_task_id task.provider_task_id = provider_task_id
task.seedance_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 task.remote_result_url = created.get("remote_result_url") or task.remote_result_url
if task.gen_type == "image": if task.gen_type == "image":
task.image_tokens_used = created.get("image_tokens", task.image_tokens_used or 0) or 0 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: 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" task.pipeline_stage = "result_ready"
else: else:
# Video path: provider returns task id, poll later. # 视频路径:provider 返回 task id,后续轮询。
task.pipeline_stage = "waiting_remote" task.pipeline_stage = "waiting_remote"
task.status = "generating" task.status = "generating"
await db.commit() 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": if task.pipeline_stage == "result_ready":
from app.tasks.generation_download_tasks import download_generation_result_task from app.tasks.generation_download_tasks import download_generation_result_task
download_generation_result_task.delay(task.id) download_generation_result_task.delay(task.id)
else: else:
from app.tasks.generation_poll_tasks import poll_generation_task from app.tasks.generation_poll_tasks import poll_generation_task
poll_generation_task.delay(task.id) poll_generation_task.delay(task.id)
except Exception as exc: except Exception as exc:
try: try:
await db.rollback() await db.rollback()
except Exception: except Exception:
pass pass
result = await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id == task_id)) result = await db.execute(select(ChatGenerationTask).where(ChatGenerationTask.id == task_id))
task = result.scalar_one_or_none() task = result.scalar_one_or_none()
if task: if task:
task.status = "failed" task.status = "failed"
task.error_message = extract_error_message(exc, "生成任务") if callable(extract_error_message) else str(exc) task.error_message = extract_error_message(exc, "生成任务") if callable(extract_error_message) else str(exc)
@@ -108,6 +279,8 @@ else:
class _DisabledTask: class _DisabledTask:
def delay(self, *args, **kwargs): def delay(self, *args, **kwargs):
raise RuntimeError("Celery is disabled") raise RuntimeError("Celery is disabled")
def apply_async(self, *args, **kwargs): def apply_async(self, *args, **kwargs):
raise RuntimeError("Celery is disabled") raise RuntimeError("Celery is disabled")
chatapi_create_generation_task = _DisabledTask()
chatapi_create_generation_task = _DisabledTask()