Files
video-gen/video-gen-api/app/tasks/generation_create_tasks.py

300 lines
11 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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.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,
)
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):
return run_async(_run(task_id))
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()