from __future__ import annotations from datetime import datetime, timedelta, timezone from typing import Any from sqlalchemy import select from sqlalchemy.ext.asyncio import AsyncSession from app.config import settings from app.enums.shot_replicate import ShotSplitStatusEnum from app.models.shot_replicate_segment import ShotReplicateSegment from app.models.shot_replicate_task_set import ShotReplicateTaskSet from app.services.shot_replicate_taskset_service import refresh_task_set_split_summary from app.services.module_async_recovery_service import register_shot_split_task from app.tasks.celery_app import celery_app def _now() -> datetime: return datetime.now(timezone.utc) def _ensure_aware(value: datetime | None) -> datetime | None: if value is None: return None if value.tzinfo is None: return value.replace(tzinfo=timezone.utc) return value.astimezone(timezone.utc) def _expired(value: datetime | None, now: datetime | None = None) -> bool: checked = _ensure_aware(value) if checked is None: return True return checked <= (now or _now()) def _queue_timeout(segment: ShotReplicateSegment, now: datetime | None = None) -> bool: enqueued_at = _ensure_aware(segment.split_enqueued_at) if enqueued_at is None: return True return enqueued_at + timedelta(seconds=int(settings.SHOT_SPLIT_PENDING_TIMEOUT_SECONDS or 300)) <= (now or _now()) async def recover_one_split_segment(db: AsyncSession, segment: ShotReplicateSegment, *, source: str = "startup_db") -> str: from app.tasks.shot_replicate_tasks import split_one_segment if not segment: return "skip_missing_segment" if segment.deleted_at is not None: return "skip_deleted" if segment.split_status == ShotSplitStatusEnum.COMPLETED.value: return "skip_completed" if segment.split_status == ShotSplitStatusEnum.FAILED.value: return "skip_failed" current_time = _now() should_recover = False if segment.split_status == ShotSplitStatusEnum.PENDING.value: should_recover = _queue_timeout(segment, current_time) elif segment.split_status == ShotSplitStatusEnum.PROCESSING.value: should_recover = _expired(segment.split_lease_until, current_time) elif segment.split_status == ShotSplitStatusEnum.RETRY_WAITING.value: should_recover = _expired(segment.split_next_retry_at, current_time) if not should_recover: return f"skip_{segment.split_status}_not_due" if int(segment.split_retry_count or 0) >= int(settings.SHOT_SPLIT_MAX_RETRY_COUNT or 3): segment.split_status = ShotSplitStatusEnum.FAILED.value segment.split_last_error = segment.split_last_error or f"{source} 恢复时超过最大重试次数" segment.split_lease_until = None segment.split_next_retry_at = None await refresh_task_set_split_summary(db, segment.task_set_id) await db.commit() return "mark_failed_max_retry" segment.split_status = ShotSplitStatusEnum.PENDING.value segment.split_enqueued_at = current_time segment.split_lease_until = None segment.split_next_retry_at = None await refresh_task_set_split_summary(db, segment.task_set_id) await db.commit() await register_shot_split_task(segment.id, task_set_id=segment.task_set_id) if celery_app: split_one_segment.apply_async( args=[segment.id], queue="gen_result_download", priority=settings.DOWNLOAD_TASK_PRIORITY_RECOVER, countdown=0, ) return f"recover_{source}" async def recover_shot_split_tasks_once(db: AsyncSession) -> dict[str, Any]: """拆镜 ffmpeg 任务容灾恢复。独立扫描 shot_replicate_segments,不复用 Chat 下载 active registry。""" batch_size = int(settings.SHOT_SPLIT_RECOVERY_BATCH_SIZE or 50) result = await db.execute( select(ShotReplicateSegment) .where( ShotReplicateSegment.deleted_at.is_(None), ShotReplicateSegment.split_status.in_( [ ShotSplitStatusEnum.PENDING.value, ShotSplitStatusEnum.PROCESSING.value, ShotSplitStatusEnum.RETRY_WAITING.value, ] ), ) .order_by(ShotReplicateSegment.updated_at.asc()) .limit(batch_size) .with_for_update(skip_locked=True) ) segments = list(result.scalars().all()) checked = 0 results: dict[str, int] = {} touched_task_set_ids: set[str] = set() for segment in segments: action = await recover_one_split_segment(db, segment, source="startup_db") checked += 1 touched_task_set_ids.add(segment.task_set_id) results[action] = results.get(action, 0) + 1 for task_set_id in touched_task_set_ids: await refresh_task_set_split_summary(db, task_set_id) await db.commit() return {"checked": checked, "results": results}