# 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 dataclasses import dataclass, field 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] = {} class RedisExecutionLockError(RuntimeError): """Redis execution-lock infrastructure error. Execution locks are fail-closed: callers must stop the current Celery task and retry later instead of falling back to an unlocked database path. """ class RedisExecutionLockUnavailable(RedisExecutionLockError): """Redis is unavailable, so execution ownership cannot be established.""" class RedisExecutionLockLost(RedisExecutionLockError): """The current worker no longer owns the execution lock.""" _RELEASE_LOCK_SCRIPT = """ if redis.call('get', KEYS[1]) == ARGV[1] then return redis.call('del', KEYS[1]) else return 0 end """ _RENEW_LOCK_SCRIPT = """ if redis.call('get', KEYS[1]) == ARGV[1] then return redis.call('pexpire', KEYS[1], ARGV[2]) else return 0 end """ 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 缓存客户端,确保同一个客户端 只在创建它的事件循环里使用。普通 active 注册表在 Redis 不可用时返回 None;执行锁封装会把 None 转为 RedisExecutionLockUnavailable,严格中止 当前任务,不允许无锁执行外部副作用。 """ 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_acquire_execution_lock( *, lock_key: str, ttl_seconds: int, token: Optional[str] = None, log_context: str = "execution_lock", ) -> Optional[str]: """Fail-closed execution lock. Returns a token when acquired and ``None`` when another worker owns the lock. Redis connection/command failures raise :class:`RedisExecutionLockUnavailable`; callers must retry the Celery task and must not execute the external side effect without a lock. """ redis = await get_registry_redis() if redis is None: raise RedisExecutionLockUnavailable( f"Redis execution lock unavailable: context={log_context}, key={lock_key}" ) lock_token = token or uuid.uuid4().hex ttl_ms = max(1000, int(ttl_seconds or 60) * 1000) try: acquired = await redis.set(lock_key, lock_token, nx=True, px=ttl_ms) return lock_token if acquired else None except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc: logger.error( "获取 Redis 执行锁失败。context=%s, lock_key=%s, error=%s", log_context, lock_key, exc, ) raise RedisExecutionLockUnavailable( f"Redis execution lock acquire failed: {lock_key}: {exc}" ) from exc async def redis_check_lock_owner( *, lock_key: str, token: str, log_context: str = "execution_lock", ) -> bool: redis = await get_registry_redis() if redis is None: raise RedisExecutionLockUnavailable( f"Redis execution lock unavailable: context={log_context}, key={lock_key}" ) try: value = await redis.get(lock_key) except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc: logger.error( "检查 Redis 执行锁失败。context=%s, lock_key=%s, error=%s", log_context, lock_key, exc, ) raise RedisExecutionLockUnavailable( f"Redis execution lock check failed: {lock_key}: {exc}" ) from exc return bool(value and str(value) == str(token)) async def redis_renew_lock( *, lock_key: str, token: str, ttl_seconds: int, log_context: str = "execution_lock", ) -> bool: redis = await get_registry_redis() if redis is None: raise RedisExecutionLockUnavailable( f"Redis execution lock unavailable: context={log_context}, key={lock_key}" ) ttl_ms = max(1000, int(ttl_seconds or 60) * 1000) try: renewed = await redis.eval(_RENEW_LOCK_SCRIPT, 1, lock_key, token, ttl_ms) return bool(renewed) except (RedisError, OSError, RuntimeError, TypeError, ValueError) as exc: logger.error( "续期 Redis 执行锁失败。context=%s, lock_key=%s, error=%s", log_context, lock_key, exc, ) raise RedisExecutionLockUnavailable( f"Redis execution lock renew failed: {lock_key}: {exc}" ) from exc @dataclass(slots=True) class RedisExecutionLockLease: """Owned Redis execution lock with compare-and-expire heartbeat.""" lock_key: str token: str ttl_seconds: int log_context: str = "execution_lock" renew_interval_seconds: int | None = None _stop_event: asyncio.Event = field(default_factory=asyncio.Event, init=False, repr=False) _heartbeat_task: asyncio.Task[Any] | None = field(default=None, init=False, repr=False) _lost_error: RedisExecutionLockError | None = field(default=None, init=False, repr=False) @classmethod async def acquire( cls, *, lock_key: str, ttl_seconds: int, token: str | None = None, log_context: str = "execution_lock", renew_interval_seconds: int | None = None, ) -> "RedisExecutionLockLease | None": acquired_token = await redis_acquire_execution_lock( lock_key=lock_key, ttl_seconds=ttl_seconds, token=token, log_context=log_context, ) if not acquired_token: return None lease = cls( lock_key=lock_key, token=acquired_token, ttl_seconds=max(1, int(ttl_seconds or 60)), log_context=log_context, renew_interval_seconds=renew_interval_seconds, ) lease.start_heartbeat() return lease def start_heartbeat(self) -> None: if self._heartbeat_task is not None: return self._heartbeat_task = asyncio.create_task(self._heartbeat()) async def _heartbeat(self) -> None: interval = int( self.renew_interval_seconds or max(1, min(60, self.ttl_seconds // 3)) ) while not self._stop_event.is_set(): try: await asyncio.wait_for(self._stop_event.wait(), timeout=interval) return except asyncio.TimeoutError: pass try: renewed = await redis_renew_lock( lock_key=self.lock_key, token=self.token, ttl_seconds=self.ttl_seconds, log_context=self.log_context, ) if not renewed: self._lost_error = RedisExecutionLockLost( f"Redis execution lock ownership lost: {self.lock_key}" ) return except RedisExecutionLockError as exc: self._lost_error = exc return async def ensure_owned(self) -> None: if self._lost_error is not None: raise self._lost_error owned = await redis_check_lock_owner( lock_key=self.lock_key, token=self.token, log_context=self.log_context, ) if not owned: self._lost_error = RedisExecutionLockLost( f"Redis execution lock ownership lost: {self.lock_key}" ) raise self._lost_error async def close(self) -> None: self._stop_event.set() heartbeat = self._heartbeat_task if heartbeat is not None: try: await heartbeat except Exception: logger.debug( "Redis execution lock heartbeat close failed. context=%s key=%s", self.log_context, self.lock_key, exc_info=True, ) try: await redis_release_lock( lock_key=self.lock_key, token=self.token, log_context=self.log_context, ) except Exception: logger.debug( "Redis execution lock release failed. context=%s key=%s", self.log_context, self.lock_key, exc_info=True, ) 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 try: released = await redis.eval(_RELEASE_LOCK_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