119 lines
4.2 KiB
Python
119 lines
4.2 KiB
Python
from __future__ import annotations
|
|
|
|
from typing import Any, TypeAlias
|
|
|
|
from sqlalchemy import select
|
|
from sqlalchemy.ext.asyncio import AsyncSession
|
|
|
|
from app.enums.generation_status import GenerationRecordPipelineStage, GenerationStatus
|
|
from app.enums.generation_task import ChatGenerationPipelineStage, ChatGenerationTaskStatus
|
|
from app.models.chat_generation_task import ChatGenerationTask
|
|
from app.models.generation_record import GenerationRecord
|
|
from app.models.video_upscale_task import VideoUpscaleTask
|
|
from app.services.generation.pipeline.db_lock_service import apply_short_lock_timeout
|
|
|
|
VideoUpscaleOwner: TypeAlias = ChatGenerationTask | GenerationRecord
|
|
|
|
|
|
def owner_type(owner: VideoUpscaleOwner | None) -> str | None:
|
|
if isinstance(owner, ChatGenerationTask):
|
|
return "chat_generation_task"
|
|
if isinstance(owner, GenerationRecord):
|
|
return "generation_record"
|
|
return None
|
|
|
|
|
|
def owner_id(owner: VideoUpscaleOwner | None) -> str | None:
|
|
return str(getattr(owner, "id", "") or "") or None
|
|
|
|
|
|
def owner_is_generating(owner: VideoUpscaleOwner) -> bool:
|
|
if isinstance(owner, ChatGenerationTask):
|
|
return owner.status == ChatGenerationTaskStatus.GENERATING.value
|
|
return owner.status == GenerationStatus.generating.value
|
|
|
|
|
|
def owner_is_completed(owner: VideoUpscaleOwner) -> bool:
|
|
if isinstance(owner, ChatGenerationTask):
|
|
return owner.status == ChatGenerationTaskStatus.COMPLETED.value
|
|
return owner.status == GenerationStatus.completed.value
|
|
|
|
|
|
def set_owner_stage(owner: VideoUpscaleOwner, stage: str) -> None:
|
|
owner.pipeline_stage = stage
|
|
|
|
|
|
def upscale_stage_value(owner: VideoUpscaleOwner, chat_stage: ChatGenerationPipelineStage | str) -> str:
|
|
value = chat_stage.value if hasattr(chat_stage, "value") else str(chat_stage)
|
|
if isinstance(owner, ChatGenerationTask):
|
|
return value
|
|
try:
|
|
return GenerationRecordPipelineStage(value).value
|
|
except ValueError:
|
|
return value
|
|
|
|
|
|
async def load_upscale_owner(
|
|
db: AsyncSession,
|
|
upscale: VideoUpscaleTask,
|
|
*,
|
|
for_update: bool,
|
|
) -> VideoUpscaleOwner | None:
|
|
if upscale.chat_generation_task_id:
|
|
query = select(ChatGenerationTask).where(
|
|
ChatGenerationTask.id == upscale.chat_generation_task_id,
|
|
ChatGenerationTask.deleted_at.is_(None),
|
|
)
|
|
elif upscale.generation_record_id:
|
|
query = select(GenerationRecord).where(
|
|
GenerationRecord.id == upscale.generation_record_id,
|
|
GenerationRecord.deleted_at.is_(None),
|
|
)
|
|
else:
|
|
return None
|
|
if for_update:
|
|
await apply_short_lock_timeout(db)
|
|
query = query.with_for_update()
|
|
result = await db.execute(query.limit(1))
|
|
return result.scalar_one_or_none()
|
|
|
|
|
|
async def mark_owner_upscale_failed(
|
|
db: AsyncSession,
|
|
owner: VideoUpscaleOwner,
|
|
*,
|
|
error_message: str,
|
|
) -> None:
|
|
if isinstance(owner, ChatGenerationTask):
|
|
owner.status = ChatGenerationTaskStatus.FAILED.value
|
|
owner.pipeline_stage = ChatGenerationPipelineStage.UPSCALE_FAILED.value
|
|
owner.error_message = error_message
|
|
return
|
|
|
|
owner.status = GenerationStatus.failed.value
|
|
owner.pipeline_stage = GenerationRecordPipelineStage.UPSCALE_FAILED.value
|
|
owner.error_message = error_message
|
|
|
|
|
|
def restore_owner_for_upscale_retry(owner: VideoUpscaleOwner) -> None:
|
|
if isinstance(owner, ChatGenerationTask):
|
|
owner.status = ChatGenerationTaskStatus.GENERATING.value
|
|
owner.pipeline_stage = ChatGenerationPipelineStage.UPSCALE_QUEUED.value
|
|
else:
|
|
owner.status = GenerationStatus.generating.value
|
|
owner.pipeline_stage = GenerationRecordPipelineStage.UPSCALE_QUEUED.value
|
|
owner.error_message = None
|
|
|
|
|
|
def owner_context(owner: VideoUpscaleOwner | None) -> dict[str, Any]:
|
|
if owner is None:
|
|
return {}
|
|
return {
|
|
"owner_type": owner_type(owner),
|
|
"owner_id": owner_id(owner),
|
|
"chat_generation_task_id": owner.id if isinstance(owner, ChatGenerationTask) else None,
|
|
"generation_record_id": owner.id if isinstance(owner, GenerationRecord) else None,
|
|
"project_id": getattr(owner, "project_id", None),
|
|
"generation_mode": getattr(owner, "generation_mode", None),
|
|
}
|