CHAT生成任务不做提词优化|积分变动描述修订
This commit is contained in:
@@ -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,
|
||||
)
|
||||
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user