import asyncio import logging import os import time from abc import ABC, abstractmethod from enum import StrEnum from cachetools import LRUCache logger = logging.getLogger(__name__) _DEFAULT_TTL_SECONDS = 300 _IDEMPOTENCY_KEY_SEPARATOR = ":" _IDEMPOTENCY_NONCE_MARKER = ":nonce:" class ClaimStatus(StrEnum): INVALID = "invalid" DUPLICATE = "duplicate" INFLIGHT = "inflight" CLAIMED = "claimed" class IdempotencyBackend(ABC): """幂等性后端抽象基类""" @abstractmethod async def claim(self, msg_id: str, ttl: int = _DEFAULT_TTL_SECONDS) -> bool: ... @abstractmethod async def claim_with_status(self, msg_id: str) -> ClaimStatus: ... @abstractmethod async def check_and_set(self, msg_id: str, ttl: int = _DEFAULT_TTL_SECONDS) -> bool: ... @abstractmethod async def commit(self, msg_id: str) -> None: ... @abstractmethod async def release(self, msg_id: str) -> None: ... @abstractmethod async def release_and_forget(self, msg_id: str) -> None: ... @abstractmethod async def clear_inflight(self, msg_id: str | None = None) -> None: ... @abstractmethod async def reset(self) -> None: ... class InMemoryBackend(IdempotencyBackend): def __init__(self) -> None: self._cache: LRUCache = LRUCache(maxsize=10_000) self._inflight: set[str] = set() self._lock = asyncio.Lock() def _prune_expired(self, now: float) -> None: expired = [mid for mid, (ts, _ttl) in self._cache.items() if now - ts >= _ttl] for mid in expired: del self._cache[mid] self._inflight.discard(mid) async def check_and_set(self, msg_id: str, ttl: int = _DEFAULT_TTL_SECONDS) -> bool: if not msg_id: return False async with self._lock: now = time.monotonic() if msg_id in self._cache: timestamp, entry_ttl = self._cache[msg_id] if now - timestamp < entry_ttl: return False self._prune_expired(now) self._cache[msg_id] = (now, ttl) return True async def claim(self, msg_id: str, ttl: int = _DEFAULT_TTL_SECONDS) -> bool: if not msg_id: return False async with self._lock: self._prune_expired(time.monotonic()) if msg_id in self._inflight: return False now = time.monotonic() if msg_id in self._cache: timestamp, entry_ttl = self._cache[msg_id] if now - timestamp < entry_ttl: return False self._cache[msg_id] = (now, ttl) self._inflight.add(msg_id) return True async def claim_with_status(self, msg_id: str) -> ClaimStatus: if not msg_id: return ClaimStatus.INVALID async with self._lock: self._prune_expired(time.monotonic()) if msg_id in self._inflight: return ClaimStatus.INFLIGHT now = time.monotonic() if msg_id in self._cache: timestamp, entry_ttl = self._cache[msg_id] if now - timestamp < entry_ttl: return ClaimStatus.DUPLICATE self._cache[msg_id] = (now, _DEFAULT_TTL_SECONDS) self._inflight.add(msg_id) return ClaimStatus.CLAIMED async def commit(self, msg_id: str) -> None: async with self._lock: self._inflight.discard(msg_id) async def release(self, msg_id: str) -> None: async with self._lock: self._inflight.discard(msg_id) async def release_and_forget(self, msg_id: str) -> None: async with self._lock: self._inflight.discard(msg_id) self._cache.pop(msg_id, None) async def clear_inflight(self, msg_id: str | None = None) -> None: async with self._lock: if msg_id is not None: self._inflight.discard(msg_id) else: self._inflight.clear() async def reset(self) -> None: async with self._lock: self._cache.clear() self._inflight.clear() class RedisBackend(IdempotencyBackend): """基于 Redis 的分布式幂等性后端 使用 SET NX EX 实现原子性检查和设置,结合 SADD/SREM 管理 inflight 状态。 """ _INFLIGHT_SUFFIX = ":inflight" def __init__(self, redis_client, key_prefix: str = "yuxi:idempotency") -> None: self._redis = redis_client self._key_prefix = key_prefix def _cache_key(self, msg_id: str) -> str: return f"{self._key_prefix}:msg:{msg_id}" def _inflight_key(self, msg_id: str) -> str: return f"{self._key_prefix}{self._INFLIGHT_SUFFIX}" async def check_and_set(self, msg_id: str, ttl: int = _DEFAULT_TTL_SECONDS) -> bool: if not msg_id: return False key = self._cache_key(msg_id) return await self._redis.set(key, "1", nx=True, ex=ttl) or False async def claim(self, msg_id: str, ttl: int = _DEFAULT_TTL_SECONDS) -> bool: if not msg_id: return False inflight_key = self._inflight_key(msg_id) added = await self._redis.sadd(inflight_key, msg_id) if added == 0: return False key = self._cache_key(msg_id) set_ok = await self._redis.set(key, "1", nx=True, ex=ttl) if not set_ok: await self._redis.srem(inflight_key, msg_id) return False return True async def claim_with_status(self, msg_id: str) -> ClaimStatus: if not msg_id: return ClaimStatus.INVALID inflight_key = self._inflight_key(msg_id) if await self._redis.sismember(inflight_key, msg_id): return ClaimStatus.INFLIGHT key = self._cache_key(msg_id) if await self._redis.exists(key): return ClaimStatus.DUPLICATE added = await self._redis.sadd(inflight_key, msg_id) if added == 0: return ClaimStatus.INFLIGHT await self._redis.set(key, "1", ex=_DEFAULT_TTL_SECONDS) return ClaimStatus.CLAIMED async def commit(self, msg_id: str) -> None: inflight_key = self._inflight_key(msg_id) await self._redis.srem(inflight_key, msg_id) async def release(self, msg_id: str) -> None: inflight_key = self._inflight_key(msg_id) await self._redis.srem(inflight_key, msg_id) async def release_and_forget(self, msg_id: str) -> None: inflight_key = self._inflight_key(msg_id) cache_key = self._cache_key(msg_id) await self._redis.srem(inflight_key, msg_id) await self._redis.delete(cache_key) async def clear_inflight(self, msg_id: str | None = None) -> None: if msg_id is not None: inflight_key = self._inflight_key(msg_id) await self._redis.srem(inflight_key, msg_id) else: keys = await self._redis.keys(f"{self._key_prefix}{self._INFLIGHT_SUFFIX}:*") for key in keys: await self._redis.delete(key) async def reset(self) -> None: keys = await self._redis.keys(f"{self._key_prefix}:*") for key in keys: await self._redis.delete(key) _backend: IdempotencyBackend | None = None _backend_lock = asyncio.Lock() def _get_redis_client(): redis_url = os.getenv("REDIS_URL", os.getenv("YUXI_REDIS_URL", "")) if not redis_url: return None try: from redis.asyncio import Redis return Redis.from_url(redis_url, decode_responses=True) except Exception: logger.warning("Failed to create Redis client for idempotency, falling back to in-memory") return None async def get_backend() -> IdempotencyBackend: global _backend if _backend is not None: return _backend async with _backend_lock: if _backend is not None: return _backend redis_client = _get_redis_client() if redis_client is not None: _backend = RedisBackend(redis_client) logger.info("Idempotency backend: Redis") else: _backend = InMemoryBackend() logger.info("Idempotency backend: InMemory (single-process)") return _backend async def _ensure_backend() -> IdempotencyBackend: if _backend is not None: return _backend return await get_backend() async def check_and_set(msg_id: str, ttl: int = _DEFAULT_TTL_SECONDS) -> bool: backend = await _ensure_backend() return await backend.check_and_set(msg_id, ttl) async def claim(msg_id: str, ttl: int = _DEFAULT_TTL_SECONDS) -> bool: backend = await _ensure_backend() return await backend.claim(msg_id, ttl) async def claim_with_status(msg_id: str) -> ClaimStatus: backend = await _ensure_backend() return await backend.claim_with_status(msg_id) async def commit(msg_id: str) -> None: backend = await _ensure_backend() await backend.commit(msg_id) async def release(msg_id: str) -> None: backend = await _ensure_backend() await backend.release(msg_id) async def release_and_forget(msg_id: str) -> None: backend = await _ensure_backend() await backend.release_and_forget(msg_id) async def consume(msg_id: str) -> bool: if not await claim(msg_id): return False await release(msg_id) return True def build_idempotency_key(prefix: str, key: str, nonce: str | None = None) -> str: base = f"{prefix}{_IDEMPOTENCY_KEY_SEPARATOR}{key}" if nonce: return f"{base}{_IDEMPOTENCY_NONCE_MARKER}{nonce}" return base def parse_idempotency_key(idempotency_key: str) -> tuple[str, str] | None: if not idempotency_key: return None sep_idx = idempotency_key.find(_IDEMPOTENCY_KEY_SEPARATOR) if sep_idx < 0: return None prefix = idempotency_key[:sep_idx] body = idempotency_key[sep_idx + 1 :] nonce_marker = body.rfind(_IDEMPOTENCY_NONCE_MARKER) if nonce_marker >= 0: return prefix, body[:nonce_marker] return prefix, body async def clear_inflight(msg_id: str | None = None) -> None: backend = await _ensure_backend() await backend.clear_inflight(msg_id) async def reset() -> None: backend = await _ensure_backend() await backend.reset() def reset_sync() -> None: global _backend if _backend is not None: if isinstance(_backend, InMemoryBackend): _backend._cache.clear() _backend._inflight.clear() else: import asyncio asyncio.get_event_loop().run_until_complete(reset()) else: _backend = InMemoryBackend() reset_for_tests = reset_sync