Files
video-gen/video-gen-api/app/services/generation/pipeline/owner_service.py
T
2026-07-22 14:48:29 +08:00

248 lines
9.3 KiB
Python
Raw 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
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)