ForcePilot/backend/package/yuxi/channel/message/idempotency.py

340 lines
10 KiB
Python
Raw Normal View History

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