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), }