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