589 lines
22 KiB
Python
589 lines
22 KiB
Python
from __future__ import annotations
|
||
|
||
import asyncio
|
||
import json
|
||
import uuid
|
||
from datetime import datetime, timedelta, timezone
|
||
from typing import Any, Optional
|
||
|
||
from app.config import settings
|
||
from app.enums.celery_queue import CeleryQueue, CeleryTaskName
|
||
from app.enums.celery_runtime import CeleryRuntimeDomain
|
||
from app.enums.generation_status import GenerationRecordPipelineStage
|
||
from app.enums.generation_task import (
|
||
ALLOWED_GENERATION_MODES,
|
||
ChatGenerationPipelineStage,
|
||
ChatGenerationTaskEventType,
|
||
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,
|
||
renew_generation_owner_claim_lease,
|
||
)
|
||
from app.services.generation.poll_schedule_service import ensure_video_poll_fields
|
||
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,
|
||
sync_generation_record_media_token_snapshot,
|
||
)
|
||
from app.services.redis_registry_service import RedisExecutionLockError
|
||
from app.services.celery_runtime.runtime_service import CeleryRuntimeLease, RuntimeIdentity
|
||
from app.tasks.async_runner import run_async
|
||
from app.tasks.celery_app import celery_app
|
||
|
||
|
||
def _now() -> datetime:
|
||
return datetime.now(timezone.utc)
|
||
|
||
|
||
def _get_first_value(obj: Any, *field_names: str) -> Optional[Any]:
|
||
for field_name in field_names:
|
||
value = getattr(obj, field_name, None)
|
||
if value is not None and str(value).strip() != "":
|
||
return value
|
||
return None
|
||
|
||
|
||
def _clean(value: Any) -> str | None:
|
||
text = str(value).strip() if value is not None else ""
|
||
return text or None
|
||
|
||
|
||
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(("{", "[")):
|
||
return base_prompt
|
||
|
||
parts: list[str] = []
|
||
if gen_type == GenerationType.VIDEO.value:
|
||
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'}",
|
||
]
|
||
)
|
||
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
|
||
|
||
|
||
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
|
||
|
||
|
||
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:
|
||
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()
|
||
|
||
|
||
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)
|
||
|
||
|
||
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 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_token = uuid.uuid4().hex
|
||
|
||
async def _renew_db_claim(token: str) -> bool:
|
||
return await renew_generation_owner_claim_lease(
|
||
owner_type=normalized_owner_type,
|
||
owner_id=task_id,
|
||
attempt_no=effective_attempt,
|
||
claim_field="provider_create_claim_token",
|
||
lease_field="provider_create_lease_until",
|
||
token=token,
|
||
lease_seconds=int(settings.GENERATION_CREATE_LOCK_TTL_SECONDS or 600),
|
||
)
|
||
|
||
lease = await CeleryRuntimeLease.acquire(
|
||
identity=RuntimeIdentity(
|
||
domain=CeleryRuntimeDomain.GENERATION_CREATE.value,
|
||
owner_type=normalized_owner_type,
|
||
owner_id=task_id,
|
||
attempt_no=effective_attempt,
|
||
task_name=CeleryTaskName.CHATAPI_CREATE.value,
|
||
queue=CeleryQueue.GEN_CHATAPI_CREATE.value,
|
||
),
|
||
lock_key=_lock_key(normalized_owner_type, task_id, effective_attempt),
|
||
hash_key=settings.GENERATION_CREATE_ACTIVE_REDIS_HASH_KEY,
|
||
zset_key=settings.GENERATION_CREATE_ACTIVE_REDIS_ZSET_KEY,
|
||
token=lease_token,
|
||
ttl_seconds=int(settings.GENERATION_CREATE_LOCK_TTL_SECONDS or 600),
|
||
heartbeat_interval_seconds=int(settings.REDIS_EXECUTION_LOCK_RENEW_INTERVAL_SECONDS or 30),
|
||
pipeline_stage=ChatGenerationPipelineStage.CREATING_PROVIDER_TASK.value,
|
||
db_heartbeat=_renew_db_claim,
|
||
)
|
||
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 not is_attempt_current(owner, effective_attempt):
|
||
await db.rollback()
|
||
return
|
||
|
||
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
|
||
|
||
allowed_stages = {
|
||
_stage(owner, ChatGenerationPipelineStage.QUEUED),
|
||
_stage(owner, ChatGenerationPipelineStage.PREPARING),
|
||
_stage(owner, ChatGenerationPipelineStage.CREATING_PROVIDER_TASK),
|
||
}
|
||
if is_image_main:
|
||
allowed_stages.add(
|
||
_stage(owner, ChatGenerationPipelineStage.PROVIDER_RESULT_STAGED)
|
||
)
|
||
if owner.pipeline_stage not in allowed_stages:
|
||
return
|
||
|
||
try:
|
||
if not owner.optimized_prompt:
|
||
owner.pipeline_stage = _stage(
|
||
owner, ChatGenerationPipelineStage.PREPARING
|
||
)
|
||
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:
|
||
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_owner_failed_and_refund_once(
|
||
db,
|
||
owner,
|
||
error_message=error_message,
|
||
pipeline_stage=_stage(
|
||
owner, ChatGenerationPipelineStage.FAILED
|
||
),
|
||
)
|
||
await db.commit()
|
||
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,
|
||
owner_type: str = GenerationOwnerType.CHAT_GENERATION_TASK.value,
|
||
generation_attempt_no: int | None = None,
|
||
):
|
||
try:
|
||
return run_async(
|
||
_run(
|
||
task_id,
|
||
owner_type=owner_type,
|
||
generation_attempt_no=generation_attempt_no,
|
||
)
|
||
)
|
||
except Exception as exc:
|
||
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()
|