248 lines
9.3 KiB
Python
248 lines
9.3 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
|
||
# 兼容 SimpleNamespace 快照(如 image_batch 的 task_snapshot)
|
||
if hasattr(owner, 'generation_mode'):
|
||
return GenerationOwnerType.CHAT_GENERATION_TASK.value
|
||
if hasattr(owner, 'include_media_references'):
|
||
return GenerationOwnerType.GENERATION_RECORD.value
|
||
raise TypeError(f"不支持的生成任务对象: {type(owner)!r}")
|
||
|
||
|
||
def owner_include_media_references(owner: GenerationOwner) -> bool:
|
||
"""返回本次供应商创建是否应携带附件。
|
||
|
||
ChatGenerationTask 延续原有行为;GenerationRecord 使用用户提交并持久化的开关。
|
||
"""
|
||
if isinstance(owner, ChatGenerationTask):
|
||
return True
|
||
if isinstance(owner, GenerationRecord):
|
||
return bool(owner.include_media_references)
|
||
# 兼容 SimpleNamespace 快照(如 image_batch 的 task_snapshot)
|
||
if hasattr(owner, 'generation_mode'):
|
||
return True
|
||
if hasattr(owner, 'include_media_references'):
|
||
return bool(owner.include_media_references)
|
||
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)
|
||
|
||
async def renew_generation_owner_claim_lease(
|
||
*,
|
||
owner_type: str | GenerationOwnerType | None,
|
||
owner_id: str,
|
||
attempt_no: int,
|
||
claim_field: str,
|
||
lease_field: str,
|
||
token: str,
|
||
lease_seconds: int,
|
||
) -> bool:
|
||
"""CAS 续期生成所有者租约,不加载 ORM 对象,避免 heartbeat 产生懒加载风险。"""
|
||
from datetime import datetime, timedelta, timezone
|
||
from sqlalchemy import update
|
||
from app.models.base import async_session
|
||
|
||
normalized = normalize_owner_type(owner_type)
|
||
model = ChatGenerationTask if normalized == GenerationOwnerType.CHAT_GENERATION_TASK.value else GenerationRecord
|
||
claim_column = getattr(model, claim_field)
|
||
now = datetime.now(timezone.utc)
|
||
async with async_session() as db:
|
||
result = await db.execute(
|
||
update(model)
|
||
.where(
|
||
model.id == owner_id,
|
||
model.deleted_at.is_(None),
|
||
model.generation_attempt_no == int(attempt_no),
|
||
claim_column == token,
|
||
)
|
||
.values({lease_field: now + timedelta(seconds=max(1, int(lease_seconds)))})
|
||
)
|
||
await db.commit()
|
||
return bool(result.rowcount == 1)
|