Files
video-gen/video-gen-api/app/services/generation/pipeline/owner_service.py
T
2026-07-20 13:48:17 +08:00

193 lines
7.1 KiB
Python

from __future__ import annotations
import asyncio
from dataclasses import dataclass
from typing import TypeAlias
from sqlalchemy import select
from sqlalchemy.ext.asyncio import AsyncSession
from app.enums.generation_status import GenerationStatus
from app.enums.generation_task import (
ChatGenerationTaskStatus,
GenerationMode,
GenerationOwnerType,
)
from app.models.chat_generation_task import ChatGenerationTask
from app.models.generation_record import GenerationRecord
from app.services.generation.pipeline.db_lock_service import (
DatabaseRowLockBusy,
apply_short_lock_timeout,
raise_if_database_lock_busy,
)
GenerationOwner: TypeAlias = ChatGenerationTask | GenerationRecord
@dataclass(frozen=True, slots=True)
class GenerationOwnerRef:
owner_type: str
owner_id: str
generation_attempt_no: int | None = None
@classmethod
def from_owner(cls, owner: GenerationOwner) -> "GenerationOwnerRef":
return cls(
owner_type=owner_type_of(owner),
owner_id=str(owner.id),
generation_attempt_no=int(getattr(owner, "generation_attempt_no", 1) or 1),
)
def normalize_owner_type(value: str | GenerationOwnerType | None) -> str:
if isinstance(value, GenerationOwnerType):
return value.value
text = str(value or "").strip().lower()
if not text:
return GenerationOwnerType.CHAT_GENERATION_TASK.value
if text not in {item.value for item in GenerationOwnerType}:
raise ValueError(f"不支持的生成任务所有者类型: {value}")
return text
def owner_type_of(owner: GenerationOwner) -> str:
if isinstance(owner, ChatGenerationTask):
return GenerationOwnerType.CHAT_GENERATION_TASK.value
if isinstance(owner, GenerationRecord):
return GenerationOwnerType.GENERATION_RECORD.value
raise TypeError(f"不支持的生成任务对象: {type(owner)!r}")
def owner_mode(owner: GenerationOwner) -> str:
if isinstance(owner, ChatGenerationTask):
return str(owner.generation_mode or GenerationMode.CHATAPI_ASYNC.value)
return GenerationMode.GENERATION_RECORD.value
def owner_provider_task_id(owner: GenerationOwner) -> str | None:
return str(getattr(owner, "provider_task_id", None) or getattr(owner, "seedance_task_id", None) or "") or None
def set_owner_provider_task_id(owner: GenerationOwner, value: str | None) -> None:
if hasattr(owner, "provider_task_id"):
owner.provider_task_id = value
owner.seedance_task_id = value
def owner_is_generating(owner: GenerationOwner) -> bool:
if isinstance(owner, ChatGenerationTask):
return owner.status == ChatGenerationTaskStatus.GENERATING.value
return owner.status == GenerationStatus.generating.value
def owner_is_completed(owner: GenerationOwner) -> bool:
if isinstance(owner, ChatGenerationTask):
return owner.status == ChatGenerationTaskStatus.COMPLETED.value
return owner.status == GenerationStatus.completed.value
def owner_is_failed(owner: GenerationOwner) -> bool:
if isinstance(owner, ChatGenerationTask):
return owner.status == ChatGenerationTaskStatus.FAILED.value
return owner.status == GenerationStatus.failed.value
def set_owner_generating(owner: GenerationOwner) -> None:
owner.status = ChatGenerationTaskStatus.GENERATING.value if isinstance(owner, ChatGenerationTask) else GenerationStatus.generating.value
def set_owner_completed(owner: GenerationOwner) -> None:
owner.status = ChatGenerationTaskStatus.COMPLETED.value if isinstance(owner, ChatGenerationTask) else GenerationStatus.completed.value
def set_owner_failed(owner: GenerationOwner) -> None:
owner.status = ChatGenerationTaskStatus.FAILED.value if isinstance(owner, ChatGenerationTask) else GenerationStatus.failed.value
def is_attempt_current(owner: GenerationOwner, attempt_no: int | None) -> bool:
if attempt_no is None:
return True
return int(getattr(owner, "generation_attempt_no", 1) or 1) == int(attempt_no)
async def load_generation_owner(
db: AsyncSession,
*,
owner_type: str | GenerationOwnerType | None,
owner_id: str,
for_update: bool = False,
include_deleted: bool = False,
) -> GenerationOwner | None:
normalized = normalize_owner_type(owner_type)
if normalized == GenerationOwnerType.CHAT_GENERATION_TASK.value:
query = select(ChatGenerationTask).where(ChatGenerationTask.id == owner_id)
if not include_deleted:
query = query.where(ChatGenerationTask.deleted_at.is_(None))
else:
query = select(GenerationRecord).where(GenerationRecord.id == owner_id)
if not include_deleted:
query = query.where(GenerationRecord.deleted_at.is_(None))
if for_update:
await apply_short_lock_timeout(db)
query = query.with_for_update().execution_options(populate_existing=True)
try:
result = await db.execute(query.limit(1))
except Exception as exc:
raise_if_database_lock_busy(exc)
raise
return result.scalar_one_or_none()
async def load_generation_owner_for_update_retry(
db: AsyncSession,
*,
owner_type: str | GenerationOwnerType | None,
owner_id: str,
attempts: int = 3,
) -> GenerationOwner | None:
"""Retry a short owner row lock inside the same Worker.
This is intended after an external call/download has already completed so a
transient row lock does not force the Worker to repeat that external side effect.
"""
last_error: DatabaseRowLockBusy | None = None
for retry_index in range(max(1, int(attempts or 1))):
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 + 1 < max(1, int(attempts or 1)):
await asyncio.sleep(1 + retry_index)
raise last_error or DatabaseRowLockBusy()
def redis_owner_item_id(owner_type: str | GenerationOwnerType | None, owner_id: str, attempt_no: int | None = None) -> str:
normalized = normalize_owner_type(owner_type)
if attempt_no is None:
return f"{normalized}:{owner_id}"
return f"{normalized}:{owner_id}:attempt:{int(attempt_no)}"
def parse_redis_owner_item_id(value: str) -> GenerationOwnerRef:
text = str(value or "").strip()
for owner_type in (GenerationOwnerType.CHAT_GENERATION_TASK.value, GenerationOwnerType.GENERATION_RECORD.value):
prefix = f"{owner_type}:"
if text.startswith(prefix):
rest = text[len(prefix):]
marker = ":attempt:"
if marker in rest:
owner_id, attempt = rest.rsplit(marker, 1)
try:
return GenerationOwnerRef(owner_type, owner_id, int(attempt))
except ValueError:
return GenerationOwnerRef(owner_type, owner_id, None)
return GenerationOwnerRef(owner_type, rest, None)
# Historical Redis/Celery identifiers always belonged to ChatGenerationTask.
return GenerationOwnerRef(GenerationOwnerType.CHAT_GENERATION_TASK.value, text, None)