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