Files
video-gen/video-gen-api/app/tasks/generation_create_tasks.py
2026-07-22 14:48:29 +08:00

589 lines
22 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 __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()