161 lines
6.5 KiB
Python
161 lines
6.5 KiB
Python
from __future__ import annotations
|
|
|
|
from collections.abc import Awaitable, Callable
|
|
from typing import Any
|
|
|
|
from app.config import settings
|
|
from app.models.base import async_session
|
|
from app.services.redis_registry_service import RedisExecutionLockLease
|
|
from app.services.video_upscale.task_service import (
|
|
recover_video_upscale_tasks_once,
|
|
run_finalize_upscale,
|
|
run_local_upscale,
|
|
run_remote_poll,
|
|
run_remote_result_download,
|
|
run_remote_submit,
|
|
)
|
|
from app.tasks.async_runner import run_async
|
|
from app.tasks.celery_app import celery_app
|
|
|
|
|
|
async def _run_with_execution_lock(
|
|
upscale_task_id: str,
|
|
callback: Callable[[str, str, Callable[[], Awaitable[None]]], Awaitable[None]],
|
|
) -> None:
|
|
lock_key = f"{settings.VIDEO_UPSCALE_EXECUTION_LOCK_KEY_PREFIX}:{upscale_task_id}"
|
|
lease = await RedisExecutionLockLease.acquire(
|
|
lock_key=lock_key,
|
|
ttl_seconds=max(30, int(settings.VIDEO_UPSCALE_EXECUTION_LOCK_TTL_SECONDS or 900)),
|
|
log_context="video_upscale_execution",
|
|
renew_interval_seconds=max(1, int(settings.REDIS_EXECUTION_LOCK_RENEW_INTERVAL_SECONDS or 20)),
|
|
)
|
|
if lease is None:
|
|
# 重复消息已有其他 Worker 推进,不属于业务失败。
|
|
return
|
|
try:
|
|
await callback(upscale_task_id, lease.token, lease.ensure_owned)
|
|
await lease.ensure_owned()
|
|
finally:
|
|
await lease.close()
|
|
|
|
|
|
async def _run_local(upscale_task_id: str) -> None:
|
|
async def _execute(task_id: str, token: str, guard: Callable[[], Awaitable[None]]) -> None:
|
|
async with async_session() as db:
|
|
await run_local_upscale(db, task_id, execution_token=token, execution_guard=guard)
|
|
await _run_with_execution_lock(upscale_task_id, _execute)
|
|
|
|
|
|
async def _run_submit(upscale_task_id: str, *, count_attempt: bool = True) -> None:
|
|
async def _execute(task_id: str, token: str, guard: Callable[[], Awaitable[None]]) -> None:
|
|
async with async_session() as db:
|
|
await run_remote_submit(
|
|
db, task_id, count_attempt=count_attempt, execution_token=token, execution_guard=guard
|
|
)
|
|
await _run_with_execution_lock(upscale_task_id, _execute)
|
|
|
|
|
|
async def _run_poll(upscale_task_id: str) -> None:
|
|
async def _execute(task_id: str, token: str, guard: Callable[[], Awaitable[None]]) -> None:
|
|
async with async_session() as db:
|
|
await run_remote_poll(db, task_id, execution_token=token, execution_guard=guard)
|
|
await _run_with_execution_lock(upscale_task_id, _execute)
|
|
|
|
|
|
async def _run_download(upscale_task_id: str) -> None:
|
|
async def _execute(task_id: str, token: str, guard: Callable[[], Awaitable[None]]) -> None:
|
|
async with async_session() as db:
|
|
await run_remote_result_download(db, task_id, execution_token=token, execution_guard=guard)
|
|
await _run_with_execution_lock(upscale_task_id, _execute)
|
|
|
|
|
|
async def _run_finalize(upscale_task_id: str) -> None:
|
|
async def _execute(task_id: str, token: str, guard: Callable[[], Awaitable[None]]) -> None:
|
|
async with async_session() as db:
|
|
await run_finalize_upscale(db, task_id, execution_token=token, execution_guard=guard)
|
|
await _run_with_execution_lock(upscale_task_id, _execute)
|
|
|
|
|
|
async def _run_recovery() -> dict[str, Any]:
|
|
lease = await RedisExecutionLockLease.acquire(
|
|
lock_key=settings.VIDEO_UPSCALE_RECOVERY_LOCK_KEY,
|
|
ttl_seconds=max(30, int(settings.VIDEO_UPSCALE_EXECUTION_LOCK_TTL_SECONDS or 900)),
|
|
log_context="video_upscale_recovery",
|
|
renew_interval_seconds=max(1, int(settings.REDIS_EXECUTION_LOCK_RENEW_INTERVAL_SECONDS or 20)),
|
|
)
|
|
if lease is None:
|
|
return {"checked": 0, "results": {"lock_busy": 1}}
|
|
try:
|
|
async with async_session() as db:
|
|
result = await recover_video_upscale_tasks_once(db)
|
|
await lease.ensure_owned()
|
|
return result
|
|
finally:
|
|
await lease.close()
|
|
|
|
|
|
if celery_app:
|
|
|
|
@celery_app.task(name="video_upscale.execute_local", bind=True, max_retries=2)
|
|
def execute_local(self, upscale_task_id: str) -> None:
|
|
try:
|
|
return run_async(_run_local(upscale_task_id))
|
|
except Exception as exc:
|
|
raise self.retry(exc=exc, countdown=max(5, int(settings.VIDEO_UPSCALE_RETRY_BACKOFF_SECONDS or 60)))
|
|
|
|
|
|
@celery_app.task(name="video_upscale.submit_remote", bind=True, max_retries=2)
|
|
def submit_remote(self, upscale_task_id: str, count_attempt: bool = True) -> None:
|
|
try:
|
|
return run_async(_run_submit(upscale_task_id, count_attempt=count_attempt))
|
|
except Exception as exc:
|
|
raise self.retry(exc=exc, countdown=max(5, int(settings.VIDEO_UPSCALE_RETRY_BACKOFF_SECONDS or 60)))
|
|
|
|
|
|
@celery_app.task(name="video_upscale.poll_remote", bind=True, max_retries=2)
|
|
def poll_remote(self, upscale_task_id: str) -> None:
|
|
try:
|
|
return run_async(_run_poll(upscale_task_id))
|
|
except Exception as exc:
|
|
raise self.retry(exc=exc, countdown=max(5, int(settings.VIDEO_UPSCALE_RETRY_BACKOFF_SECONDS or 60)))
|
|
|
|
|
|
@celery_app.task(name="video_upscale.download_remote_result", bind=True, max_retries=2)
|
|
def download_remote_result(self, upscale_task_id: str) -> None:
|
|
try:
|
|
return run_async(_run_download(upscale_task_id))
|
|
except Exception as exc:
|
|
raise self.retry(exc=exc, countdown=max(5, int(settings.VIDEO_UPSCALE_RETRY_BACKOFF_SECONDS or 60)))
|
|
|
|
|
|
@celery_app.task(name="video_upscale.finalize", bind=True, max_retries=2)
|
|
def finalize(self, upscale_task_id: str) -> None:
|
|
try:
|
|
return run_async(_run_finalize(upscale_task_id))
|
|
except Exception as exc:
|
|
raise self.retry(exc=exc, countdown=max(5, int(settings.VIDEO_UPSCALE_RETRY_BACKOFF_SECONDS or 60)))
|
|
|
|
|
|
@celery_app.task(name="video_upscale.recover_once", bind=True, max_retries=2)
|
|
def recover_once(self) -> dict[str, Any]:
|
|
try:
|
|
return run_async(_run_recovery())
|
|
except Exception as exc:
|
|
raise self.retry(exc=exc, countdown=max(5, int(settings.VIDEO_UPSCALE_RETRY_BACKOFF_SECONDS or 60)))
|
|
|
|
else:
|
|
|
|
class _DisabledTask:
|
|
def delay(self, *args: Any, **kwargs: Any) -> None:
|
|
raise RuntimeError("Celery is disabled")
|
|
|
|
def apply_async(self, *args: Any, **kwargs: Any) -> None:
|
|
raise RuntimeError("Celery is disabled")
|
|
|
|
execute_local = _DisabledTask()
|
|
submit_remote = _DisabledTask()
|
|
poll_remote = _DisabledTask()
|
|
download_remote_result = _DisabledTask()
|
|
finalize = _DisabledTask()
|
|
recover_once = _DisabledTask()
|