129 lines
4.7 KiB
Python
129 lines
4.7 KiB
Python
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.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()
|
|
|
|
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}
|