Files
video-gen/video-gen-api/app/services/celery_download_recovery_service.py
T

339 lines
9.2 KiB
Python

# app/services/celery_download_recovery_service.py
from __future__ import annotations
import inspect
import json
import logging
from datetime import datetime, timezone
from typing import Any, Dict, Iterable, List, Optional, Union
from app.config import settings
try:
from redis.exceptions import RedisError
except ImportError:
RedisError = RuntimeError # type: ignore[assignment]
logger = logging.getLogger("video_gen")
_redis_client: Optional[Any] = None
def utc_now() -> datetime:
return datetime.now(timezone.utc)
def ensure_aware_utc(value: Optional[datetime]) -> Optional[datetime]:
if value is None:
return None
if value.tzinfo is None:
return value.replace(tzinfo=timezone.utc)
return value.astimezone(timezone.utc)
def datetime_to_epoch(value: Optional[datetime]) -> int:
checked_value = ensure_aware_utc(value) or utc_now()
return int(checked_value.timestamp())
def _registry_redis_url() -> str:
return settings.CELERY_BROKER_URL or settings.REDIS_URL or ""
async def get_registry_redis() -> Optional[Any]:
global _redis_client
if _redis_client is not None:
return _redis_client
redis_url = _registry_redis_url()
if not redis_url:
return None
try:
from redis.asyncio import Redis
except ImportError as exc:
logger.warning(
"下载容灾 Redis 注册表不可用,redis 依赖未安装。error=%s",
exc,
)
return None
try:
redis_client = Redis.from_url(redis_url, decode_responses=True)
await redis_client.ping()
_redis_client = redis_client
return _redis_client
except (RedisError, OSError, RuntimeError) as exc:
logger.warning(
"下载容灾 Redis 注册表不可用,降级为仅 DB 容灾。error=%s",
exc,
)
_redis_client = None
return None
async def close_registry_redis() -> None:
global _redis_client
client = _redis_client
_redis_client = None
if client is None:
return
try:
close_method = getattr(client, "close", None)
if close_method is None:
return
close_result = close_method()
if inspect.isawaitable(close_result):
await close_result
except (RedisError, OSError, RuntimeError) as exc:
logger.debug(
"关闭下载容灾 Redis 注册表连接失败。error=%s",
exc,
)
def build_download_active_payload(
*,
record_id: str,
celery_task_id: Optional[str],
stage: str,
attempt: Optional[int] = None,
queue: str = "gen_result_download",
priority: Optional[int] = None,
enqueue_at: Optional[datetime] = None,
started_at: Optional[datetime] = None,
updated_at: Optional[datetime] = None,
lease_until: Optional[datetime] = None,
next_retry_at: Optional[datetime] = None,
check_at: Optional[datetime] = None,
reason: Optional[str] = None,
) -> Dict[str, Any]:
now = utc_now()
checked_updated_at = ensure_aware_utc(updated_at) or now
checked_enqueue_at = ensure_aware_utc(enqueue_at)
checked_started_at = ensure_aware_utc(started_at)
checked_lease_until = ensure_aware_utc(lease_until)
checked_next_retry_at = ensure_aware_utc(next_retry_at)
checked_check_at = ensure_aware_utc(check_at)
return {
"record_id": record_id,
"celery_task_id": celery_task_id,
"stage": stage,
"attempt": int(attempt or 0),
"queue": queue,
"priority": priority,
"enqueue_at": (
datetime_to_epoch(checked_enqueue_at)
if checked_enqueue_at
else None
),
"started_at": (
datetime_to_epoch(checked_started_at)
if checked_started_at
else None
),
"updated_at": datetime_to_epoch(checked_updated_at),
"lease_until": (
datetime_to_epoch(checked_lease_until)
if checked_lease_until
else None
),
"next_retry_at": (
datetime_to_epoch(checked_next_retry_at)
if checked_next_retry_at
else None
),
"check_at": (
datetime_to_epoch(checked_check_at)
if checked_check_at
else None
),
"reason": reason,
}
async def upsert_download_active(
*,
record_id: str,
payload: Dict[str, Any],
check_at: Optional[Union[datetime, int, float]],
) -> None:
redis = await get_registry_redis()
if redis is None:
return
if isinstance(check_at, datetime):
score = datetime_to_epoch(check_at)
elif check_at is None:
score = datetime_to_epoch(utc_now())
else:
score = int(float(check_at))
updated_payload = dict(payload)
updated_payload["check_at"] = score
try:
pipe: Any = redis.pipeline(transaction=True)
pipe.hset(
settings.DOWNLOAD_ACTIVE_REDIS_HASH_KEY,
record_id,
json.dumps(updated_payload, ensure_ascii=False, default=str),
)
pipe.zadd(
settings.DOWNLOAD_ACTIVE_REDIS_ZSET_KEY,
{record_id: score},
)
await pipe.execute()
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning(
"写入下载容灾 Redis 注册表失败。record_id=%s, error=%s",
record_id,
exc,
)
async def remove_download_active(record_id: str) -> None:
redis = await get_registry_redis()
if redis is None:
return
try:
pipe: Any = redis.pipeline(transaction=True)
pipe.hdel(settings.DOWNLOAD_ACTIVE_REDIS_HASH_KEY, record_id)
pipe.zrem(settings.DOWNLOAD_ACTIVE_REDIS_ZSET_KEY, record_id)
await pipe.execute()
except (RedisError, OSError, RuntimeError) as exc:
logger.warning(
"删除下载容灾 Redis 注册表失败。record_id=%s, error=%s",
record_id,
exc,
)
async def get_due_download_record_ids(
*,
limit: Optional[int] = None,
now: Optional[datetime] = None,
) -> List[str]:
redis = await get_registry_redis()
if redis is None:
return []
batch_limit = int(limit or settings.DOWNLOAD_RECOVERY_BATCH_SIZE or 100)
score = datetime_to_epoch(now or utc_now())
try:
result = await redis.zrangebyscore(
settings.DOWNLOAD_ACTIVE_REDIS_ZSET_KEY,
min="-inf",
max=score,
start=0,
num=batch_limit,
)
return [str(item) for item in result]
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning(
"扫描下载容灾 Redis ZSet 失败。error=%s",
exc,
)
return []
async def get_download_active_payloads(
record_ids: Iterable[str],
) -> Dict[str, Dict[str, Any]]:
cleaned_record_ids = [str(item) for item in record_ids if item]
if not cleaned_record_ids:
return {}
redis = await get_registry_redis()
if redis is None:
return {}
try:
raw_values = await redis.hmget(
settings.DOWNLOAD_ACTIVE_REDIS_HASH_KEY,
cleaned_record_ids,
)
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning(
"读取下载容灾 Redis Hash 失败。error=%s",
exc,
)
return {}
result: Dict[str, Dict[str, Any]] = {}
for record_id, raw in zip(cleaned_record_ids, raw_values):
if not raw:
continue
try:
value = json.loads(raw)
except (TypeError, ValueError, json.JSONDecodeError):
continue
if isinstance(value, dict):
result[record_id] = value
return result
async def postpone_download_active_check(
*,
record_id: str,
payload: Optional[Dict[str, Any]] = None,
check_at: Optional[Union[datetime, int, float]] = None,
) -> None:
redis = await get_registry_redis()
if redis is None:
return
if isinstance(check_at, datetime):
score = datetime_to_epoch(check_at)
elif check_at is None:
score = datetime_to_epoch(utc_now()) + int(
settings.DOWNLOAD_TASK_QUEUE_TIMEOUT_SECONDS or 300
)
else:
score = int(float(check_at))
try:
pipe: Any = redis.pipeline(transaction=True)
pipe.zadd(
settings.DOWNLOAD_ACTIVE_REDIS_ZSET_KEY,
{record_id: score},
)
if payload is not None:
updated_payload = dict(payload)
updated_payload["check_at"] = score
updated_payload["updated_at"] = datetime_to_epoch(utc_now())
pipe.hset(
settings.DOWNLOAD_ACTIVE_REDIS_HASH_KEY,
record_id,
json.dumps(
updated_payload,
ensure_ascii=False,
default=str,
),
)
await pipe.execute()
except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc:
logger.warning(
"刷新下载容灾 Redis 检查时间失败。record_id=%s, error=%s",
record_id,
exc,
)