项目/AI生成链路合并
This commit is contained in:
@@ -1,29 +1,51 @@
|
||||
from app.tasks.async_runner import run_async
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Any, Optional
|
||||
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.config import settings
|
||||
from app.enums.celery_queue import CeleryQueue
|
||||
from app.enums.generation_status import GenerationRecordPipelineStage
|
||||
from app.enums.generation_task import (
|
||||
ALLOWED_GENERATION_MODES,
|
||||
ChatGenerationPipelineStage,
|
||||
ChatGenerationTaskEventType,
|
||||
ChatGenerationTaskStatus,
|
||||
GenerationMode,
|
||||
GenerationOwnerType,
|
||||
GenerationType,
|
||||
)
|
||||
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.pipeline.db_lock_service import DatabaseRowLockBusy
|
||||
from app.services.generation.pipeline.lifecycle_service import (
|
||||
mark_owner_failed_and_refund_once,
|
||||
notify_owner_finished,
|
||||
)
|
||||
from app.services.generation.pipeline.owner_service import (
|
||||
GenerationOwner,
|
||||
is_attempt_current,
|
||||
load_generation_owner,
|
||||
normalize_owner_type,
|
||||
owner_is_generating,
|
||||
owner_mode,
|
||||
owner_provider_task_id,
|
||||
set_owner_provider_task_id,
|
||||
)
|
||||
from app.services.generation.poll_schedule_service import ensure_video_poll_fields
|
||||
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.services.media_token_usage_snapshot_service import (
|
||||
sync_chat_generation_task_media_token_snapshot,
|
||||
sync_generation_record_media_token_snapshot,
|
||||
)
|
||||
from app.services.redis_registry_service import (
|
||||
RedisExecutionLockError,
|
||||
RedisExecutionLockLease,
|
||||
)
|
||||
from app.tasks.async_runner import run_async
|
||||
from app.tasks.celery_app import celery_app
|
||||
|
||||
|
||||
@@ -32,341 +54,502 @@ def _now() -> datetime:
|
||||
|
||||
|
||||
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
|
||||
value = getattr(obj, field_name, None)
|
||||
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 _clean(value: Any) -> str | None:
|
||||
text = str(value).strip() if value is not None else ""
|
||||
return text or 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 {GenerationMode.HOT_OPENING_REPLICATE.value, GenerationMode.SHOT_REPLICATE.value} and gen_type == GenerationType.VIDEO.value:
|
||||
def _build_optimized_prompt_by_params(owner: GenerationOwner) -> str:
|
||||
base_prompt = (_clean(owner.original_prompt) or "").rstrip(",,。;; \n\t")
|
||||
gen_type = (_clean(owner.gen_type) or "").lower()
|
||||
generation_mode = owner_mode(owner)
|
||||
if (
|
||||
generation_mode
|
||||
in {
|
||||
GenerationMode.HOT_OPENING_REPLICATE.value,
|
||||
GenerationMode.SHOT_REPLICATE.value,
|
||||
}
|
||||
and gen_type == GenerationType.VIDEO.value
|
||||
):
|
||||
stripped = base_prompt.strip()
|
||||
if stripped.startswith("{") or stripped.startswith("["):
|
||||
if stripped.startswith(("{", "[")):
|
||||
return base_prompt
|
||||
|
||||
duration = _get_first_value(task, "duration")
|
||||
aspect_ratio = _get_first_value(task, "aspect_ratio")
|
||||
resolution = _get_first_value(task, "provider_generation_resolution", "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 = []
|
||||
|
||||
parts: list[str] = []
|
||||
if gen_type == GenerationType.VIDEO.value:
|
||||
# 时长:4秒,画面比例:16:9,分辨率:480p
|
||||
if duration:
|
||||
parts.append(f"时长:{duration}秒")
|
||||
parts.append(f"画面比例:{aspect_ratio}")
|
||||
parts.append(f"分辨率:{resolution}")
|
||||
else:
|
||||
parts.append("时长:4秒")
|
||||
parts.append("画面比例:16:9")
|
||||
parts.append("分辨率:480p")
|
||||
elif gen_type == GenerationType.IMAGE.value:
|
||||
if image_size:
|
||||
parts.append(f"分辨率:{image_size}")
|
||||
parts.append(f"画布比例:{image_proportion}")
|
||||
parts.append(f"像素尺寸:{image_px}")
|
||||
else:
|
||||
parts.append("分辨率:2K")
|
||||
parts.append("画布比例:1:1")
|
||||
parts.append("像素尺寸: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()
|
||||
|
||||
is_image_main = bool(
|
||||
task
|
||||
and task.generation_mode == GenerationMode.CHATAPI_MAIN.value
|
||||
and task.gen_type == GenerationType.IMAGE.value
|
||||
and int(task.generation_count or 1) > 1
|
||||
parts.extend(
|
||||
[
|
||||
f"时长:{_get_first_value(owner, 'duration') or 4}秒",
|
||||
f"画面比例:{_get_first_value(owner, 'aspect_ratio') or '16:9'}",
|
||||
f"分辨率:{_get_first_value(owner, 'provider_generation_resolution', 'resolution') or '480p'}",
|
||||
]
|
||||
)
|
||||
if not task or (task.generation_mode not in ALLOWED_GENERATION_MODES and not is_image_main):
|
||||
return
|
||||
elif gen_type == GenerationType.IMAGE.value:
|
||||
parts.extend(
|
||||
[
|
||||
f"分辨率:{_get_first_value(owner, 'image_size') or '2K'}",
|
||||
f"画布比例:{_get_first_value(owner, 'image_proportion') or '1:1'}",
|
||||
f"像素尺寸:{_get_first_value(owner, 'image_px') or '2048x2048'}",
|
||||
]
|
||||
)
|
||||
suffix = ",".join(parts)
|
||||
return f"{base_prompt},{suffix}" if base_prompt and suffix else base_prompt or suffix
|
||||
|
||||
if task.status != ChatGenerationTaskStatus.GENERATING.value:
|
||||
return
|
||||
|
||||
deadline_at = ensure_aware_utc(task.deadline_at)
|
||||
# 图片 main 的 deadline 与 provider claim 由 image_batch_service 原子处理,
|
||||
# 避免重复 Celery 消息在有效租约期间把正在执行的批次错误退款。
|
||||
if not is_image_main and 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=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
)
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
|
||||
to_status=ChatGenerationTaskStatus.FAILED.value,
|
||||
to_stage=ChatGenerationPipelineStage.TIMEOUT.value,
|
||||
)
|
||||
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
await aggregate_parent_for_child(db, task)
|
||||
await db.commit()
|
||||
return
|
||||
def _stage(owner: GenerationOwner, chat_stage: ChatGenerationPipelineStage) -> str:
|
||||
if isinstance(owner, ChatGenerationTask):
|
||||
return chat_stage.value
|
||||
try:
|
||||
return GenerationRecordPipelineStage(chat_stage.value).value
|
||||
except ValueError:
|
||||
return chat_stage.value
|
||||
|
||||
if task.pipeline_stage not in (
|
||||
ChatGenerationPipelineStage.QUEUED.value,
|
||||
ChatGenerationPipelineStage.PREPARING.value,
|
||||
ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
|
||||
):
|
||||
return
|
||||
|
||||
def _lock_key(owner_type: str, owner_id: str, attempt_no: int) -> str:
|
||||
return (
|
||||
f"{settings.GENERATION_CREATE_LOCK_KEY_PREFIX}:"
|
||||
f"{owner_type}:{owner_id}:attempt:{attempt_no}"
|
||||
)
|
||||
|
||||
|
||||
async def _sync_media_snapshot(
|
||||
db, owner: GenerationOwner, provider_response: Any = None
|
||||
) -> None:
|
||||
if isinstance(owner, ChatGenerationTask):
|
||||
await sync_chat_generation_task_media_token_snapshot(
|
||||
db, owner, provider_response=provider_response
|
||||
)
|
||||
else:
|
||||
await sync_generation_record_media_token_snapshot(
|
||||
db, owner, provider_response=provider_response
|
||||
)
|
||||
|
||||
|
||||
async def _stale(owner: GenerationOwner, attempt_no: int | None) -> bool:
|
||||
if is_attempt_current(owner, attempt_no):
|
||||
return False
|
||||
await log_task_event(
|
||||
owner,
|
||||
event_type=ChatGenerationTaskEventType.STALE_ATTEMPT_MESSAGE_SKIPPED.value,
|
||||
message="创建任务消息属于旧生成轮次,已跳过",
|
||||
detail={
|
||||
"message_attempt": attempt_no,
|
||||
"current_attempt": owner.generation_attempt_no,
|
||||
},
|
||||
)
|
||||
return True
|
||||
|
||||
|
||||
async def _reload_owner_after_external_call(
|
||||
db,
|
||||
*,
|
||||
owner_type: str,
|
||||
owner_id: str,
|
||||
) -> GenerationOwner | None:
|
||||
"""Keep the provider result in the current Worker while briefly retrying a busy row lock."""
|
||||
last_error: DatabaseRowLockBusy | None = None
|
||||
for retry_index in range(3):
|
||||
try:
|
||||
if not task.optimized_prompt:
|
||||
old_stage = task.pipeline_stage
|
||||
task.pipeline_stage = ChatGenerationPipelineStage.PREPARING.value
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type=ChatGenerationTaskEventType.PROMPT_CONCAT_START.value,
|
||||
from_stage=old_stage,
|
||||
to_stage=ChatGenerationPipelineStage.PREPARING.value,
|
||||
message="开始本地拼接提示词,不调用提词优化API",
|
||||
)
|
||||
return await load_generation_owner(
|
||||
db,
|
||||
owner_type=owner_type,
|
||||
owner_id=owner_id,
|
||||
for_update=True,
|
||||
)
|
||||
except DatabaseRowLockBusy as exc:
|
||||
last_error = exc
|
||||
await db.rollback()
|
||||
if retry_index < 2:
|
||||
await asyncio.sleep(1 + retry_index)
|
||||
raise last_error or DatabaseRowLockBusy()
|
||||
|
||||
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
|
||||
async def _resolve_attempt(
|
||||
task_id: str, *, owner_type: str, message_attempt: int | None
|
||||
) -> int | None:
|
||||
async with async_session() as db:
|
||||
owner = await load_generation_owner(
|
||||
db, owner_type=owner_type, owner_id=task_id, for_update=False
|
||||
)
|
||||
if not owner:
|
||||
return None
|
||||
if await _stale(owner, message_attempt):
|
||||
return None
|
||||
return int(owner.generation_attempt_no or 1)
|
||||
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type=ChatGenerationTaskEventType.PROMPT_CONCAT_SUCCESS.value,
|
||||
to_stage=ChatGenerationPipelineStage.PREPARING.value,
|
||||
detail={
|
||||
"optimized_prompt": optimized_prompt,
|
||||
"gen_type": task.gen_type,
|
||||
"message": "已完成本地提示词拼接,未调用提词优化API",
|
||||
},
|
||||
)
|
||||
|
||||
if is_image_main:
|
||||
from app.services.generation.ai.image_batch_service import run_image_main_batch
|
||||
async def _dispatch_next_stage(
|
||||
db,
|
||||
owner: GenerationOwner,
|
||||
*,
|
||||
normalized_owner_type: str,
|
||||
) -> None:
|
||||
if owner.pipeline_stage == _stage(
|
||||
owner, ChatGenerationPipelineStage.RESULT_READY
|
||||
):
|
||||
from app.tasks.generation_download_tasks import enqueue_download_task
|
||||
|
||||
await run_image_main_batch(db, task)
|
||||
await enqueue_download_task(db, owner, reason="create_result_ready")
|
||||
return
|
||||
|
||||
from app.tasks.generation_poll_tasks import (
|
||||
poll_generation_task,
|
||||
register_poll_active,
|
||||
)
|
||||
|
||||
try:
|
||||
check_at = _now() + timedelta(
|
||||
seconds=int(settings.POLL_TASK_LEASE_SECONDS or 300)
|
||||
)
|
||||
await register_poll_active(
|
||||
owner,
|
||||
reason="create_provider_success",
|
||||
check_at=check_at,
|
||||
next_poll_at=owner.next_poll_at,
|
||||
)
|
||||
poll_generation_task.apply_async(
|
||||
args=[owner.id],
|
||||
kwargs={
|
||||
"owner_type": normalized_owner_type,
|
||||
"generation_attempt_no": int(owner.generation_attempt_no or 1),
|
||||
"force_due": False,
|
||||
},
|
||||
queue=CeleryQueue.GEN_PROVIDER_POLL.value,
|
||||
)
|
||||
except Exception as exc:
|
||||
# Provider creation is already committed. A broker/registry failure is
|
||||
# an infrastructure enqueue failure, not a generation failure. Keep the
|
||||
# remote task ID and let due-poll/startup recovery enqueue it again.
|
||||
owner.next_poll_at = _now()
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
owner,
|
||||
event_type=ChatGenerationTaskEventType.POLL_SCHEDULED.value,
|
||||
message="供应商任务已创建,但轮询任务投递失败,等待恢复扫描",
|
||||
detail={
|
||||
"error": str(exc),
|
||||
"provider_task_id": owner_provider_task_id(owner),
|
||||
"generation_refunded": False,
|
||||
},
|
||||
to_stage=owner.pipeline_stage,
|
||||
)
|
||||
|
||||
|
||||
async def _run(
|
||||
task_id: str,
|
||||
*,
|
||||
owner_type: str = GenerationOwnerType.CHAT_GENERATION_TASK.value,
|
||||
generation_attempt_no: int | None = None,
|
||||
):
|
||||
normalized_owner_type = normalize_owner_type(owner_type)
|
||||
effective_attempt = await _resolve_attempt(
|
||||
task_id,
|
||||
owner_type=normalized_owner_type,
|
||||
message_attempt=generation_attempt_no,
|
||||
)
|
||||
if effective_attempt is None:
|
||||
return
|
||||
|
||||
lease = await RedisExecutionLockLease.acquire(
|
||||
lock_key=_lock_key(normalized_owner_type, task_id, effective_attempt),
|
||||
ttl_seconds=int(settings.GENERATION_CREATE_LOCK_TTL_SECONDS or 600),
|
||||
log_context="generation_create",
|
||||
renew_interval_seconds=int(
|
||||
settings.REDIS_EXECUTION_LOCK_RENEW_INTERVAL_SECONDS or 30
|
||||
),
|
||||
)
|
||||
if lease is None:
|
||||
return
|
||||
|
||||
try:
|
||||
async with async_session() as db:
|
||||
owner = await load_generation_owner(
|
||||
db,
|
||||
owner_type=normalized_owner_type,
|
||||
owner_id=task_id,
|
||||
for_update=True,
|
||||
)
|
||||
if not owner:
|
||||
return
|
||||
|
||||
if task.seedance_task_id or task.provider_task_id:
|
||||
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
|
||||
if task.gen_type == GenerationType.VIDEO.value:
|
||||
ensure_video_poll_fields(task, now=_now())
|
||||
task.next_poll_at = _now()
|
||||
await db.commit()
|
||||
|
||||
elif task.remote_result_url:
|
||||
task.pipeline_stage = ChatGenerationPipelineStage.RESULT_READY.value
|
||||
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 = ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_START.value,
|
||||
from_stage=old_stage,
|
||||
to_stage=ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
|
||||
)
|
||||
|
||||
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 == GenerationType.IMAGE.value:
|
||||
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 = ChatGenerationPipelineStage.RESULT_READY.value
|
||||
else:
|
||||
# 视频路径:provider 返回 task id,后续轮询。
|
||||
task.pipeline_stage = ChatGenerationPipelineStage.WAITING_REMOTE.value
|
||||
if task.gen_type == GenerationType.VIDEO.value:
|
||||
current_time = _now()
|
||||
ensure_video_poll_fields(task, now=current_time)
|
||||
task.next_poll_at = current_time
|
||||
task.poll_interval_seconds = int(task.poll_interval_seconds or 0)
|
||||
|
||||
task.status = ChatGenerationTaskStatus.GENERATING.value
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
task,
|
||||
event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_SUCCESS.value,
|
||||
to_stage=task.pipeline_stage,
|
||||
detail=created,
|
||||
)
|
||||
|
||||
if task.pipeline_stage == ChatGenerationPipelineStage.RESULT_READY.value:
|
||||
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)),
|
||||
next_poll_at=getattr(task, "next_poll_at", None),
|
||||
)
|
||||
poll_generation_task.apply_async(
|
||||
args=[task.id],
|
||||
queue=CeleryQueue.GEN_PROVIDER_POLL.value,
|
||||
countdown=0,
|
||||
)
|
||||
|
||||
except Exception as exc:
|
||||
try:
|
||||
if not is_attempt_current(owner, effective_attempt):
|
||||
await db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
return
|
||||
|
||||
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()
|
||||
is_image_main = bool(
|
||||
isinstance(owner, ChatGenerationTask)
|
||||
and owner.generation_mode == GenerationMode.CHATAPI_MAIN.value
|
||||
and owner.gen_type == GenerationType.IMAGE.value
|
||||
and int(owner.generation_count or 1) > 1
|
||||
)
|
||||
if (
|
||||
isinstance(owner, ChatGenerationTask)
|
||||
and owner.generation_mode not in ALLOWED_GENERATION_MODES
|
||||
and not is_image_main
|
||||
):
|
||||
return
|
||||
if not owner_is_generating(owner):
|
||||
return
|
||||
if owner.deadline_at and _now() > owner.deadline_at and not is_image_main:
|
||||
await mark_owner_failed_and_refund_once(
|
||||
db,
|
||||
owner,
|
||||
error_message="任务超时",
|
||||
pipeline_stage=_stage(
|
||||
owner, ChatGenerationPipelineStage.TIMEOUT
|
||||
),
|
||||
)
|
||||
await db.commit()
|
||||
await notify_owner_finished(db, owner)
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
owner,
|
||||
event_type=ChatGenerationTaskEventType.TASK_TIMEOUT.value,
|
||||
to_stage=owner.pipeline_stage,
|
||||
)
|
||||
return
|
||||
|
||||
if task:
|
||||
error_message = extract_error_message(exc, "生成任务") if callable(extract_error_message) else str(exc)
|
||||
if is_image_main:
|
||||
# image_batch_service 负责供应商/拆分失败退款。若 child 已落库,
|
||||
# 顶层兜底绝不能再把 main 退款。
|
||||
child_result = await db.execute(
|
||||
select(ChatGenerationTask.id).where(
|
||||
ChatGenerationTask.parent_task_id == task.id,
|
||||
ChatGenerationTask.generation_mode == GenerationMode.CHATAPI_CHILD.value,
|
||||
).limit(1)
|
||||
allowed_stages = {
|
||||
_stage(owner, ChatGenerationPipelineStage.QUEUED),
|
||||
_stage(owner, ChatGenerationPipelineStage.PREPARING),
|
||||
_stage(owner, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK),
|
||||
}
|
||||
if owner.pipeline_stage not in allowed_stages:
|
||||
return
|
||||
|
||||
try:
|
||||
if not owner.optimized_prompt:
|
||||
owner.pipeline_stage = _stage(
|
||||
owner, ChatGenerationPipelineStage.PREPARING
|
||||
)
|
||||
has_children = child_result.scalar_one_or_none() is not None
|
||||
if not has_children:
|
||||
task.provider_create_claim_token = None
|
||||
task.provider_create_lease_until = None
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
db,
|
||||
task=task,
|
||||
error_message=error_message,
|
||||
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
owner.optimized_prompt = _build_optimized_prompt_by_params(owner)
|
||||
owner.text_tokens_used = int(owner.text_tokens_used or 0)
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
owner,
|
||||
event_type=ChatGenerationTaskEventType.PROMPT_CONCAT_SUCCESS.value,
|
||||
to_stage=owner.pipeline_stage,
|
||||
)
|
||||
|
||||
if is_image_main:
|
||||
from app.services.generation.ai.image_batch_service import (
|
||||
run_image_main_batch,
|
||||
)
|
||||
|
||||
await run_image_main_batch(
|
||||
db,
|
||||
owner,
|
||||
execution_token=lease.token,
|
||||
execution_guard=lease.ensure_owned,
|
||||
)
|
||||
return
|
||||
|
||||
if owner_provider_task_id(owner):
|
||||
owner.pipeline_stage = _stage(
|
||||
owner, ChatGenerationPipelineStage.WAITING_REMOTE
|
||||
)
|
||||
if owner.gen_type == GenerationType.VIDEO.value:
|
||||
ensure_video_poll_fields(owner, now=_now())
|
||||
owner.next_poll_at = _now()
|
||||
await db.commit()
|
||||
elif owner.remote_result_url:
|
||||
owner.pipeline_stage = _stage(
|
||||
owner, ChatGenerationPipelineStage.RESULT_READY
|
||||
)
|
||||
await db.commit()
|
||||
else:
|
||||
current_time = _now()
|
||||
owner.pipeline_stage = _stage(
|
||||
owner, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK
|
||||
)
|
||||
owner.provider_create_claim_token = lease.token
|
||||
owner.provider_create_started_at = current_time
|
||||
owner.provider_create_lease_until = current_time + timedelta(
|
||||
seconds=int(
|
||||
settings.GENERATION_CREATE_LOCK_TTL_SECONDS or 600
|
||||
)
|
||||
)
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
owner,
|
||||
event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_START.value,
|
||||
to_stage=owner.pipeline_stage,
|
||||
)
|
||||
|
||||
created = await create_provider_task(db, owner)
|
||||
await lease.ensure_owned()
|
||||
owner = await _reload_owner_after_external_call(
|
||||
db,
|
||||
owner_type=normalized_owner_type,
|
||||
owner_id=task_id,
|
||||
)
|
||||
if not owner:
|
||||
return
|
||||
if not is_attempt_current(owner, effective_attempt):
|
||||
await db.rollback()
|
||||
return
|
||||
if owner.provider_create_claim_token != lease.token:
|
||||
return
|
||||
|
||||
provider_task_id = created.get("task_id")
|
||||
if provider_task_id:
|
||||
set_owner_provider_task_id(owner, str(provider_task_id))
|
||||
owner.remote_result_url = (
|
||||
created.get("remote_result_url") or owner.remote_result_url
|
||||
)
|
||||
owner.provider_response_json = json.dumps(
|
||||
created.get("response_data") or {},
|
||||
ensure_ascii=False,
|
||||
default=str,
|
||||
)
|
||||
owner.provider_create_claim_token = None
|
||||
owner.provider_create_lease_until = None
|
||||
if owner.gen_type == GenerationType.IMAGE.value:
|
||||
owner.image_tokens_used = int(
|
||||
created.get(
|
||||
"image_tokens", owner.image_tokens_used or 0
|
||||
)
|
||||
or 0
|
||||
)
|
||||
await _sync_media_snapshot(
|
||||
db, owner, owner.provider_response_json
|
||||
)
|
||||
|
||||
if owner.remote_result_url and not owner_provider_task_id(owner):
|
||||
owner.pipeline_stage = _stage(
|
||||
owner, ChatGenerationPipelineStage.RESULT_READY
|
||||
)
|
||||
else:
|
||||
from app.services.generation.ai.task_group_service import aggregate_main_task_status
|
||||
await aggregate_main_task_status(db, parent_task_id=str(task.id))
|
||||
owner.pipeline_stage = _stage(
|
||||
owner, ChatGenerationPipelineStage.WAITING_REMOTE
|
||||
)
|
||||
if owner.gen_type == GenerationType.VIDEO.value:
|
||||
ensure_video_poll_fields(owner, now=_now())
|
||||
owner.next_poll_at = _now()
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
owner,
|
||||
event_type=ChatGenerationTaskEventType.PROVIDER_CREATE_SUCCESS.value,
|
||||
to_stage=owner.pipeline_stage,
|
||||
detail=created,
|
||||
)
|
||||
except (RedisExecutionLockError, DatabaseRowLockBusy):
|
||||
# Redis ownership is mandatory. Do not convert infrastructure
|
||||
# lock loss into a business failure/refund.
|
||||
try:
|
||||
await db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
raise
|
||||
except Exception as exc:
|
||||
try:
|
||||
await db.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
# 只有仍持有 Redis 执行权时,才能把供应商异常收敛为业务失败。
|
||||
await lease.ensure_owned()
|
||||
owner = await load_generation_owner(
|
||||
db,
|
||||
owner_type=normalized_owner_type,
|
||||
owner_id=task_id,
|
||||
for_update=True,
|
||||
)
|
||||
if not owner or not is_attempt_current(owner, effective_attempt):
|
||||
return
|
||||
error_message = (
|
||||
extract_error_message(exc, "生成任务")
|
||||
if callable(extract_error_message)
|
||||
else str(exc)
|
||||
)
|
||||
if is_image_main and isinstance(owner, ChatGenerationTask):
|
||||
from sqlalchemy import select
|
||||
|
||||
child_result = await db.execute(
|
||||
select(ChatGenerationTask.id)
|
||||
.where(ChatGenerationTask.parent_task_id == owner.id)
|
||||
.limit(1)
|
||||
)
|
||||
if child_result.scalar_one_or_none() is not None:
|
||||
from app.services.generation.ai.task_group_service import (
|
||||
aggregate_main_task_status,
|
||||
)
|
||||
|
||||
await aggregate_main_task_status(
|
||||
db, parent_task_id=str(owner.id)
|
||||
)
|
||||
else:
|
||||
await mark_owner_failed_and_refund_once(
|
||||
db,
|
||||
owner,
|
||||
error_message=error_message,
|
||||
pipeline_stage=_stage(
|
||||
owner, ChatGenerationPipelineStage.FAILED
|
||||
),
|
||||
)
|
||||
else:
|
||||
await mark_chat_generation_task_failed_and_refund_once(
|
||||
await mark_owner_failed_and_refund_once(
|
||||
db,
|
||||
task=task,
|
||||
owner,
|
||||
error_message=error_message,
|
||||
pipeline_stage=ChatGenerationPipelineStage.FAILED.value,
|
||||
pipeline_stage=_stage(
|
||||
owner, ChatGenerationPipelineStage.FAILED
|
||||
),
|
||||
)
|
||||
await db.commit()
|
||||
await log_task_event(task, event_type=ChatGenerationTaskEventType.TASK_FAILED.value, message=error_message)
|
||||
from app.services.generation.module_hook_service import notify_chat_generation_task_finished
|
||||
from app.services.generation.ai.task_group_service import aggregate_parent_for_child
|
||||
await notify_chat_generation_task_finished(db, task)
|
||||
await aggregate_parent_for_child(db, task)
|
||||
await notify_owner_finished(db, owner)
|
||||
await db.commit()
|
||||
await log_task_event(
|
||||
owner,
|
||||
event_type=ChatGenerationTaskEventType.TASK_FAILED.value,
|
||||
message=error_message,
|
||||
)
|
||||
return
|
||||
|
||||
# The provider state is committed. Enqueue failures below are
|
||||
# recoverable infrastructure failures and must not trigger refunds.
|
||||
await _dispatch_next_stage(
|
||||
db, owner, normalized_owner_type=normalized_owner_type
|
||||
)
|
||||
finally:
|
||||
await lease.close()
|
||||
|
||||
|
||||
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):
|
||||
|
||||
@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,
|
||||
owner_type: str = GenerationOwnerType.CHAT_GENERATION_TASK.value,
|
||||
generation_attempt_no: int | None = None,
|
||||
):
|
||||
try:
|
||||
return run_async(_run(task_id))
|
||||
return run_async(
|
||||
_run(
|
||||
task_id,
|
||||
owner_type=owner_type,
|
||||
generation_attempt_no=generation_attempt_no,
|
||||
)
|
||||
)
|
||||
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)
|
||||
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")
|
||||
|
||||
Reference in New Issue
Block a user