# app/services/redis_registry_service.py from __future__ import annotations import asyncio import inspect import json import logging import os import threading import uuid 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: # pragma: no cover - redis 未安装时降级 RedisError = RuntimeError # type: ignore[assignment] logger = logging.getLogger("video_gen") _redis_clients: Dict[tuple[int, int, int], Any] = {} 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 normalize_registry_score(value: Optional[Union[datetime, int, float]]) -> int: if isinstance(value, datetime): return datetime_to_epoch(value) if value is None: return datetime_to_epoch(utc_now()) return int(float(value)) def registry_redis_url() -> str: """Celery 容灾注册表统一使用 Celery broker Redis。 不能改成只读 settings.REDIS_URL,否则线上 CELERY_BROKER_URL 使用独立 Redis DB 时,旧下载 active 注册表会被写到另一个库,导致恢复扫描失效。 """ return settings.CELERY_BROKER_URL or settings.REDIS_URL or "" def _is_supported_redis_url(redis_url: str) -> bool: if not redis_url: return False lowered = redis_url.lower() return lowered.startswith(("redis://", "rediss://", "unix://")) async def get_registry_redis() -> Optional[Any]: """获取 Celery 容灾 Redis 连接。 重点:redis.asyncio 的连接/连接池绑定 event loop,不能跨 loop 复用。 Celery -P threads 或 worker_ready + task 线程混用时,如果使用单个全局 Redis 客户端,会触发 got Future attached to a different loop。 因此这里按 pid + thread_id + event_loop_id 缓存客户端,确保同一个客户端 只在创建它的事件循环里使用。Redis 不可用时返回 None,调用方降级为 DB fallback,不能影响生成主链路。 """ redis_url = registry_redis_url() if not _is_supported_redis_url(redis_url): if redis_url: logger.warning( "Celery 容灾 Redis 注册表仅支持 redis/rediss/unix URL,当前 broker 不是 Redis,降级为 DB 容灾。url=%s", redis_url, ) return None try: from redis.asyncio import Redis except ImportError as exc: logger.warning( "Celery 容灾 Redis 注册表不可用,redis 依赖未安装。error=%s", exc, ) return None try: loop = asyncio.get_running_loop() except RuntimeError: return None client_key = (os.getpid(), threading.get_ident(), id(loop)) cached = _redis_clients.get(client_key) if cached is not None: return cached try: redis_client = Redis.from_url(redis_url, decode_responses=True) await redis_client.ping() _redis_clients[client_key] = redis_client return redis_client except (RedisError, OSError, RuntimeError) as exc: logger.warning( "Celery 容灾 Redis 注册表不可用,降级为仅 DB 容灾。error=%s", exc, ) _redis_clients.pop(client_key, None) return None async def close_registry_redis() -> None: """关闭当前进程内已缓存的 Redis 注册表连接。 关闭动作尽量只关闭当前 event loop 对应的客户端;如果调用方处于进程 退出阶段,则逐个尝试关闭,失败忽略,避免影响 worker 退出。 """ if not _redis_clients: return try: loop = asyncio.get_running_loop() current_key = (os.getpid(), threading.get_ident(), id(loop)) items = [(current_key, _redis_clients.pop(current_key, None))] except RuntimeError: items = list(_redis_clients.items()) _redis_clients.clear() for _, client in items: if client is None: continue try: close_method = getattr(client, "close", None) or getattr(client, "aclose", None) if close_method is None: continue close_result = close_method() if inspect.isawaitable(close_result): await close_result except (RedisError, OSError, RuntimeError) as exc: logger.debug("关闭 Celery 容灾 Redis 注册表连接失败。error=%s", exc) async def redis_upsert_registry_item( *, hash_key: str, zset_key: str, item_id: str, payload: Dict[str, Any], check_at: Optional[Union[datetime, int, float]], log_context: str = "registry", ) -> None: redis = await get_registry_redis() if redis is None: return score = normalize_registry_score(check_at) updated_payload = dict(payload) updated_payload["check_at"] = score updated_payload["updated_at"] = updated_payload.get("updated_at") or datetime_to_epoch(utc_now()) try: pipe: Any = redis.pipeline(transaction=True) pipe.hset(hash_key, item_id, json.dumps(updated_payload, ensure_ascii=False, default=str)) pipe.zadd(zset_key, {item_id: score}) await pipe.execute() except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc: logger.warning( "写入 Redis 注册表失败。context=%s, item_id=%s, error=%s", log_context, item_id, exc, ) async def redis_remove_registry_item( *, hash_key: str, zset_key: str, item_id: str, log_context: str = "registry", ) -> None: redis = await get_registry_redis() if redis is None: return try: pipe: Any = redis.pipeline(transaction=True) pipe.hdel(hash_key, item_id) pipe.zrem(zset_key, item_id) await pipe.execute() except (RedisError, OSError, RuntimeError) as exc: logger.warning( "删除 Redis 注册表失败。context=%s, item_id=%s, error=%s", log_context, item_id, exc, ) async def redis_get_due_registry_ids( *, zset_key: str, limit: Optional[int] = None, now: Optional[datetime] = None, log_context: str = "registry", ) -> List[str]: redis = await get_registry_redis() if redis is None: return [] batch_limit = int(limit or 100) score = datetime_to_epoch(now or utc_now()) try: result = await redis.zrangebyscore( 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 失败。context=%s, error=%s", log_context, exc) return [] async def redis_get_registry_payloads( *, hash_key: str, item_ids: Iterable[str], log_context: str = "registry", ) -> Dict[str, Dict[str, Any]]: cleaned_item_ids = [str(item) for item in item_ids if item] if not cleaned_item_ids: return {} redis = await get_registry_redis() if redis is None: return {} try: raw_values = await redis.hmget(hash_key, cleaned_item_ids) except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc: logger.warning("读取 Redis Hash 失败。context=%s, error=%s", log_context, exc) return {} result: Dict[str, Dict[str, Any]] = {} for item_id, raw in zip(cleaned_item_ids, raw_values): if not raw: continue try: value = json.loads(raw) except (TypeError, ValueError, json.JSONDecodeError): continue if isinstance(value, dict): result[item_id] = value return result async def redis_postpone_registry_item( *, hash_key: str, zset_key: str, item_id: str, payload: Optional[Dict[str, Any]] = None, check_at: Optional[Union[datetime, int, float]] = None, log_context: str = "registry", ) -> None: redis = await get_registry_redis() if redis is None: return score = normalize_registry_score(check_at) try: pipe: Any = redis.pipeline(transaction=True) pipe.zadd(zset_key, {item_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(hash_key, item_id, json.dumps(updated_payload, ensure_ascii=False, default=str)) await pipe.execute() except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc: logger.warning( "刷新 Redis 注册表检查时间失败。context=%s, item_id=%s, error=%s", log_context, item_id, exc, ) async def redis_acquire_lock( *, lock_key: str, ttl_seconds: int, token: Optional[str] = None, log_context: str = "lock", ) -> Optional[str]: """尝试获取 Redis 分布式锁。 返回 token 表示抢锁成功;返回 None 表示 Redis 不可用或锁已被其他 worker 持有。 """ redis = await get_registry_redis() if redis is None: return None lock_token = token or uuid.uuid4().hex ttl = max(1, int(ttl_seconds or 60)) try: acquired = await redis.set(lock_key, lock_token, nx=True, ex=ttl) return lock_token if acquired else None except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc: logger.warning("获取 Redis 锁失败。context=%s, lock_key=%s, error=%s", log_context, lock_key, exc) return None async def redis_release_lock( *, lock_key: str, token: str, log_context: str = "lock", ) -> bool: """只释放 token 匹配的锁,避免误删其他 worker 新抢到的锁。""" redis = await get_registry_redis() if redis is None: return False script = """ if redis.call('get', KEYS[1]) == ARGV[1] then return redis.call('del', KEYS[1]) else return 0 end """ try: released = await redis.eval(script, 1, lock_key, token) return bool(released) except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc: logger.warning("释放 Redis 锁失败。context=%s, lock_key=%s, error=%s", log_context, lock_key, exc) return False