308 lines
12 KiB
Python
308 lines
12 KiB
Python
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.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_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
|
||
|
||
ALLOWED_GENERATION_MODES = {"chatapi_async", "hot_opening_replicate", "shot_replicate"}
|
||
|
||
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 {"hot_opening_replicate", "shot_replicate"} and gen_type == "video":
|
||
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 == "video":
|
||
# 时长:4秒,画面比例:16:9,分辨率:480p
|
||
if duration:
|
||
parts.append(f"时长:{duration}秒")
|
||
parts.append(f"画面比例:{aspect_ratio}")
|
||
parts.append(f"分辨率:{resolution}")
|
||
else:
|
||
parts.append(f"时长:4秒")
|
||
parts.append(f"画面比例:16:9")
|
||
parts.append(f"分辨率:480p")
|
||
elif gen_type == "image":
|
||
if image_size :
|
||
parts.append(f"分辨率:{image_size}")
|
||
parts.append(f"画布比例:{image_proportion}")
|
||
parts.append(f"像素尺寸:{image_px}")
|
||
else:
|
||
parts.append(f"分辨率:2K")
|
||
parts.append(f"画布比例:1:1")
|
||
parts.append(f"像素尺寸: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()
|
||
|
||
if not task or task.generation_mode not in ALLOWED_GENERATION_MODES:
|
||
return
|
||
|
||
if task.status != "generating":
|
||
return
|
||
|
||
deadline_at = ensure_aware_utc(task.deadline_at)
|
||
if 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="timeout",
|
||
)
|
||
await db.commit()
|
||
await log_task_event(task, event_type="TASK_TIMEOUT", to_status="failed", to_stage="timeout")
|
||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||
await notify_chat_generation_task_finished(db, task)
|
||
await db.commit()
|
||
return
|
||
|
||
if task.pipeline_stage not in ("queued", "preparing", "creating_provider_task"):
|
||
return
|
||
|
||
try:
|
||
if not task.optimized_prompt:
|
||
old_stage = task.pipeline_stage
|
||
task.pipeline_stage = "preparing"
|
||
await db.commit()
|
||
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
|
||
|
||
await db.commit()
|
||
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 enqueue_download_task
|
||
|
||
await enqueue_download_task(db, task, reason="create_remote_result_ready")
|
||
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",
|
||
)
|
||
|
||
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,
|
||
)
|
||
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 = "result_ready"
|
||
else:
|
||
# 视频路径: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,
|
||
)
|
||
|
||
if task.pipeline_stage == "result_ready":
|
||
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)),
|
||
)
|
||
poll_generation_task.apply_async(
|
||
args=[task.id],
|
||
queue="gen_provider_poll",
|
||
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)
|
||
await mark_chat_generation_task_failed_and_refund_once(
|
||
db,
|
||
task=task,
|
||
error_message=error_message,
|
||
pipeline_stage="failed",
|
||
)
|
||
await db.commit()
|
||
await log_task_event(task, event_type="TASK_FAILED", message=task.error_message)
|
||
from app.services.generation_module_hook_service import notify_chat_generation_task_finished
|
||
await notify_chat_generation_task_finished(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() |