193 lines
7.1 KiB
Python
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)
|