项目/AI生成链路合并

This commit is contained in:
2026-07-20 13:48:17 +08:00
parent 34ca98f9eb
commit fe5a59d725
73 changed files with 5819 additions and 2573 deletions
+483 -300
View File
@@ -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")