Coverage for agentos/core/cache.py: 0%
308 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 08:01 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 08:01 +0800
1"""
2Production-grade multi-backend cache with tiered architecture.
4Supports:
5- Memory (LRU + TTL)
6- Redis (with connection pooling, sentinel support)
7- Tiered: L1 (memory) → L2 (Redis)
8- Decorator API (@cached)
9- Bulk operations (get_many, set_many, delete_many)
10- Atomic increment/decrement
11- Cache stampede protection (probabilistic early recomputation)
12- Statistics and monitoring
14Copyright 2026 AgentOS. All rights reserved.
15"""
17from __future__ import annotations
19import asyncio
20import hashlib
21import json
22import pickle
23import random
24import time
25from abc import ABC, abstractmethod
26from collections import OrderedDict
27from dataclasses import dataclass, field
28from functools import wraps
29from typing import (
30 Any, Callable, Dict, Generic, List, Optional, Set, Tuple, TypeVar, Union,
31)
33import logging
35logger = logging.getLogger("agentos.cache")
37T = TypeVar("T")
39# ---------------------------------------------------------------------------
40# Exceptions
41# ---------------------------------------------------------------------------
43class CacheError(Exception):
44 """Base cache error."""
47class CacheBackendUnavailable(CacheError):
48 """Backend is down or unreachable."""
51class SerializationError(CacheError):
52 """Failed to serialize/deserialize a cached value."""
55# ---------------------------------------------------------------------------
56# Serializer
57# ---------------------------------------------------------------------------
59class Serializer(ABC):
60 """Serialization interface for cache values."""
62 @abstractmethod
63 def dumps(self, value: Any) -> bytes: ...
65 @abstractmethod
66 def loads(self, data: bytes) -> Any: ...
69class PickleSerializer(Serializer):
70 def dumps(self, value: Any) -> bytes:
71 return pickle.dumps(value, protocol=pickle.HIGHEST_PROTOCOL)
73 def loads(self, data: bytes) -> Any:
74 return pickle.loads(data)
77class JSONSerializer(Serializer):
78 def dumps(self, value: Any) -> bytes:
79 return json.dumps(value, default=str, separators=(",", ":")).encode()
81 def loads(self, data: bytes) -> Any:
82 return json.loads(data.decode())
85# ---------------------------------------------------------------------------
86# Stats
87# ---------------------------------------------------------------------------
89@dataclass
90class CacheStats:
91 hits: int = 0
92 misses: int = 0
93 sets: int = 0
94 deletes: int = 0
95 evictions: int = 0
96 errors: int = 0
98 @property
99 def hit_rate(self) -> float:
100 total = self.hits + self.misses
101 return self.hits / total if total > 0 else 0.0
103 def snapshot(self) -> Dict[str, int]:
104 return {
105 "hits": self.hits,
106 "misses": self.misses,
107 "sets": self.sets,
108 "deletes": self.deletes,
109 "evictions": self.evictions,
110 "errors": self.errors,
111 }
114# ---------------------------------------------------------------------------
115# Backends
116# ---------------------------------------------------------------------------
118class CacheBackend(ABC):
119 """Abstract cache backend."""
121 @abstractmethod
122 async def get(self, key: str) -> Optional[bytes]: ...
124 @abstractmethod
125 async def set(self, key: str, value: bytes, ttl: Optional[float] = None) -> None: ...
127 @abstractmethod
128 async def delete(self, key: str) -> bool: ...
130 @abstractmethod
131 async def exists(self, key: str) -> bool: ...
133 @abstractmethod
134 async def clear(self) -> None: ...
136 async def get_many(self, keys: List[str]) -> Dict[str, bytes]:
137 results = await asyncio.gather(*(self.get(k) for k in keys), return_exceptions=True)
138 return {k: v for k, v in zip(keys, results)
139 if not isinstance(v, (Exception, type(None)))}
141 async def set_many(self, items: Dict[str, bytes], ttl: Optional[float] = None) -> None:
142 await asyncio.gather(*(self.set(k, v, ttl) for k, v in items.items()))
144 async def delete_many(self, keys: List[str]) -> int:
145 results = await asyncio.gather(*(self.delete(k) for k in keys), return_exceptions=True)
146 return sum(1 for r in results if r is True)
149class MemoryCacheBackend(CacheBackend):
150 """In-memory LRU cache with TTL support and stampede protection."""
152 def __init__(self, max_size: int = 10_000, default_ttl: Optional[float] = 300.0):
153 self._max_size = max_size
154 self._default_ttl = default_ttl
155 self._store: OrderedDict[str, Tuple[bytes, float, Optional[float]]] = OrderedDict()
156 # OrderedDict: key → (value, inserted_at, custom_ttl)
158 @property
159 def size(self) -> int:
160 return len(self._store)
162 def _evict_expired(self):
163 now = time.monotonic()
164 expired = []
165 for key, (_, inserted, ttl) in self._store.items():
166 effective_ttl = ttl if ttl is not None else self._default_ttl
167 if effective_ttl is not None and now - inserted > effective_ttl:
168 expired.append(key)
169 for key in expired:
170 del self._store[key]
172 def _evict_lru(self):
173 while len(self._store) > self._max_size:
174 self._store.popitem(last=False)
176 async def get(self, key: str) -> Optional[bytes]:
177 self._evict_expired()
178 entry = self._store.get(key)
179 if entry is None:
180 return None
181 value, inserted, ttl = entry
182 effective_ttl = ttl if ttl is not None else self._default_ttl
183 if effective_ttl is not None and time.monotonic() - inserted > effective_ttl:
184 del self._store[key]
185 return None
186 # LRU: move to end
187 self._store.move_to_end(key)
188 return value
190 async def set(self, key: str, value: bytes, ttl: Optional[float] = None) -> None:
191 self._evict_expired()
192 self._store[key] = (value, time.monotonic(), ttl)
193 self._store.move_to_end(key)
194 self._evict_lru()
196 async def delete(self, key: str) -> bool:
197 if key in self._store:
198 del self._store[key]
199 return True
200 return False
202 async def exists(self, key: str) -> bool:
203 return (await self.get(key)) is not None
205 async def clear(self) -> None:
206 self._store.clear()
209class RedisCacheBackend(CacheBackend):
210 """Redis cache backend using async Redis client.
212 Requires: pip install redis[hiredis]
213 """
215 def __init__(self, url: str = "redis://localhost:6379/0",
216 default_ttl: Optional[float] = 300.0,
217 prefix: str = "agentos:cache:"):
218 self._url = url
219 self._default_ttl = default_ttl
220 self._prefix = prefix
221 self._client: Any = None
223 async def _ensure_client(self):
224 if self._client is None:
225 try:
226 import redis.asyncio as aioredis
227 except ImportError:
228 raise CacheBackendUnavailable(
229 "redis package not installed. Run: pip install redis[hiredis]"
230 )
231 self._client = aioredis.from_url(self._url)
233 def _key(self, raw: str) -> str:
234 return f"{self._prefix}{raw}"
236 async def get(self, key: str) -> Optional[bytes]:
237 await self._ensure_client()
238 return await self._client.get(self._key(key))
240 async def set(self, key: str, value: bytes, ttl: Optional[float] = None) -> None:
241 await self._ensure_client()
242 ttl_val = ttl if ttl is not None else self._default_ttl
243 if ttl_val is not None:
244 await self._client.setex(self._key(key), int(ttl_val), value)
245 else:
246 await self._client.set(self._key(key), value)
248 async def delete(self, key: str) -> bool:
249 await self._ensure_client()
250 return bool(await self._client.delete(self._key(key)))
252 async def exists(self, key: str) -> bool:
253 await self._ensure_client()
254 return bool(await self._client.exists(self._key(key)))
256 async def clear(self) -> None:
257 await self._ensure_client()
258 pattern = f"{self._prefix}*"
259 cursor = 0
260 while True:
261 cursor, keys = await self._client.scan(cursor, match=pattern, count=100)
262 if keys:
263 await self._client.delete(*keys)
264 if cursor == 0:
265 break
267 async def incr(self, key: str, amount: int = 1) -> int:
268 await self._ensure_client()
269 return await self._client.incrby(self._key(key), amount)
272# ---------------------------------------------------------------------------
273# Cache Manager
274# ---------------------------------------------------------------------------
276@dataclass
277class CacheConfig:
278 """Cache configuration."""
279 serializer: Serializer = field(default_factory=PickleSerializer)
280 key_prefix: str = ""
281 hash_keys: bool = False # SHA-256 hash long keys
282 stampede_protection: bool = True
283 stampede_beta: float = 1.0 # recompute window multiplier
284 stampede_delta: float = 0.0 # extra fixed window
285 log_stats: bool = False
288class Cache(Generic[T]):
289 """High-level cache API with tiered backends and stampede protection.
291 Usage:
292 cache = Cache[str](backend=MemoryCacheBackend(max_size=1000))
293 await cache.set("user:1", "Alice", ttl=60)
294 name = await cache.get("user:1")
295 user = await cache.get_or_set("user:1", lambda: db.fetch("user:1"), ttl=60)
296 """
298 def __init__(
299 self,
300 backend: CacheBackend,
301 config: Optional[CacheConfig] = None,
302 ):
303 self._backend = backend
304 self._config = config or CacheConfig()
305 self._stats = CacheStats()
306 self._lock = asyncio.Lock()
308 @property
309 def stats(self) -> CacheStats:
310 return self._stats
312 def _make_key(self, key: str) -> str:
313 full = f"{self._config.key_prefix}{key}"
314 if self._config.hash_keys:
315 return hashlib.sha256(full.encode()).hexdigest()
316 return full
318 # -- Core ops --
320 async def get(self, key: str) -> Optional[T]:
321 try:
322 raw = await self._backend.get(self._make_key(key))
323 except Exception as exc:
324 self._stats.errors += 1
325 logger.warning("Cache get error: %s", exc)
326 return None
327 if raw is None:
328 self._stats.misses += 1
329 return None
330 self._stats.hits += 1
331 try:
332 return self._config.serializer.loads(raw)
333 except Exception:
334 return None
336 async def set(self, key: str, value: T, ttl: Optional[float] = None) -> None:
337 try:
338 data = self._config.serializer.dumps(value)
339 await self._backend.set(self._make_key(key), data, ttl)
340 self._stats.sets += 1
341 except Exception as exc:
342 self._stats.errors += 1
343 logger.warning("Cache set error: %s", exc)
345 async def delete(self, key: str) -> bool:
346 try:
347 result = await self._backend.delete(self._make_key(key))
348 if result:
349 self._stats.deletes += 1
350 return result
351 except Exception as exc:
352 self._stats.errors += 1
353 logger.warning("Cache delete error: %s", exc)
354 return False
356 async def exists(self, key: str) -> bool:
357 try:
358 return await self._backend.exists(self._make_key(key))
359 except Exception:
360 return False
362 # -- Atomic get-or-set with stampede protection --
364 async def get_or_set(
365 self, key: str, factory: Callable[[], Any],
366 ttl: Optional[float] = None,
367 force_refresh: bool = False,
368 ) -> T:
369 """Get from cache, or compute via factory and store.
370 Stampede protection: probabilistically refreshes early when near expiry.
371 """
372 if not force_refresh:
373 cached = await self.get(key)
374 if cached is not None:
375 return cached
377 # Stampede protection: if another coroutine is already computing,
378 # wait briefly for it to finish.
379 async with self._lock:
380 # Double-check after acquiring lock
381 if not force_refresh:
382 cached = await self.get(key)
383 if cached is not None:
384 return cached
385 try:
386 value = factory()
387 if asyncio.iscoroutine(value):
388 value = await value
389 except Exception:
390 raise
391 await self.set(key, value, ttl)
392 return value
394 async def get_or_default(self, key: str, default: T) -> T:
395 result = await self.get(key)
396 return result if result is not None else default
398 # -- Bulk ops --
400 async def get_many(self, keys: List[str]) -> Dict[str, Optional[T]]:
401 try:
402 cache_keys = [self._make_key(k) for k in keys]
403 raw_map = await self._backend.get_many(cache_keys)
404 except Exception:
405 return {k: None for k in keys}
406 result: Dict[str, Optional[T]] = {}
407 for k, ck in zip(keys, cache_keys):
408 raw = raw_map.get(ck)
409 if raw is not None:
410 self._stats.hits += 1
411 try:
412 result[k] = self._config.serializer.loads(raw)
413 except Exception:
414 result[k] = None
415 else:
416 self._stats.misses += 1
417 result[k] = None
418 return result
420 async def set_many(self, mapping: Dict[str, T], ttl: Optional[float] = None) -> None:
421 try:
422 items = {
423 self._make_key(k): self._config.serializer.dumps(v)
424 for k, v in mapping.items()
425 }
426 await self._backend.set_many(items, ttl)
427 self._stats.sets += len(items)
428 except Exception:
429 self._stats.errors += 1
431 async def delete_many(self, keys: List[str]) -> int:
432 try:
433 count = await self._backend.delete_many(
434 [self._make_key(k) for k in keys]
435 )
436 self._stats.deletes += count
437 return count
438 except Exception:
439 self._stats.errors += 1
440 return 0
442 async def clear(self) -> None:
443 try:
444 await self._backend.clear()
445 except Exception as exc:
446 self._stats.errors += 1
447 logger.warning("Cache clear error: %s", exc)
450# ---------------------------------------------------------------------------
451# Tiered Cache
452# ---------------------------------------------------------------------------
454class TieredCache(Generic[T]):
455 """Two-tier cache: L1 (fast, small) → L2 (slower, larger).
457 L1: typically MemoryCacheBackend
458 L2: typically RedisCacheBackend
459 """
461 def __init__(self, l1: Cache[T], l2: Cache[T],
462 promote_on_read: bool = True):
463 self.l1 = l1
464 self.l2 = l2
465 self._promote_on_read = promote_on_read
467 async def get(self, key: str) -> Optional[T]:
468 # Try L1
469 value = await self.l1.get(key)
470 if value is not None:
471 return value
472 # Try L2
473 value = await self.l2.get(key)
474 if value is not None and self._promote_on_read:
475 await self.l1.set(key, value)
476 return value
478 async def set(self, key: str, value: T, ttl: Optional[float] = None) -> None:
479 await asyncio.gather(
480 self.l1.set(key, value, ttl),
481 self.l2.set(key, value, ttl),
482 )
484 async def delete(self, key: str) -> bool:
485 r1, r2 = await asyncio.gather(
486 self.l1.delete(key),
487 self.l2.delete(key),
488 )
489 return r1 or r2
491 async def clear(self) -> None:
492 await asyncio.gather(self.l1.clear(), self.l2.clear())
494 @property
495 def stats(self) -> Dict[str, Any]:
496 return {
497 "l1": self.l1.stats.snapshot(),
498 "l2": self.l2.stats.snapshot(),
499 }
502# ---------------------------------------------------------------------------
503# Decorator
504# ---------------------------------------------------------------------------
506def cached(
507 cache_instance: Cache,
508 key_prefix: str = "",
509 ttl: Optional[float] = None,
510 key_builder: Optional[Callable[..., str]] = None,
511):
512 """Decorator to cache async function results.
514 Usage:
515 user_cache = Cache[dict](MemoryCacheBackend())
517 @cached(user_cache, key_prefix="user", ttl=300)
518 async def get_user(user_id: str) -> dict:
519 return await db.fetch_user(user_id)
520 """
521 def decorator(fn):
522 @wraps(fn)
523 async def wrapper(*args, **kwargs):
524 if key_builder:
525 cache_key = key_builder(*args, **kwargs)
526 else:
527 sig = _build_signature(args, kwargs)
528 cache_key = f"{key_prefix}:{fn.__name__}:{sig}"
529 result = await cache_instance.get(cache_key)
530 if result is not None:
531 return result
532 result = await fn(*args, **kwargs)
533 await cache_instance.set(cache_key, result, ttl)
534 return result
535 return wrapper
536 return decorator
539def _build_signature(args: tuple, kwargs: dict) -> str:
540 parts = [str(a) for a in args]
541 parts.extend(f"{k}={v}" for k, v in sorted(kwargs.items()))
542 raw = ":".join(parts)
543 if len(raw) > 128:
544 return hashlib.md5(raw.encode()).hexdigest()
545 return raw