Coverage for agentos/core/cache.py: 81%
307 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 11:37 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 11:37 +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 time
24from abc import ABC, abstractmethod
25from collections import OrderedDict
26from dataclasses import dataclass, field
27from functools import wraps
28from typing import (
29 Any, Callable, Dict, Generic, List, Optional, Tuple, TypeVar,
30)
32import logging
34logger = logging.getLogger("agentos.cache")
36T = TypeVar("T")
38# ---------------------------------------------------------------------------
39# Exceptions
40# ---------------------------------------------------------------------------
42class CacheError(Exception):
43 """Base cache error."""
46class CacheBackendUnavailable(CacheError):
47 """Backend is down or unreachable."""
50class SerializationError(CacheError):
51 """Failed to serialize/deserialize a cached value."""
54# ---------------------------------------------------------------------------
55# Serializer
56# ---------------------------------------------------------------------------
58class Serializer(ABC):
59 """Serialization interface for cache values."""
61 @abstractmethod
62 def dumps(self, value: Any) -> bytes: ...
64 @abstractmethod
65 def loads(self, data: bytes) -> Any: ...
68class PickleSerializer(Serializer):
69 def dumps(self, value: Any) -> bytes:
70 return pickle.dumps(value, protocol=pickle.HIGHEST_PROTOCOL)
72 def loads(self, data: bytes) -> Any:
73 return pickle.loads(data)
76class JSONSerializer(Serializer):
77 def dumps(self, value: Any) -> bytes:
78 return json.dumps(value, default=str, separators=(",", ":")).encode()
80 def loads(self, data: bytes) -> Any:
81 return json.loads(data.decode())
84# ---------------------------------------------------------------------------
85# Stats
86# ---------------------------------------------------------------------------
88@dataclass
89class CacheStats:
90 hits: int = 0
91 misses: int = 0
92 sets: int = 0
93 deletes: int = 0
94 evictions: int = 0
95 errors: int = 0
97 @property
98 def hit_rate(self) -> float:
99 total = self.hits + self.misses
100 return self.hits / total if total > 0 else 0.0
102 def snapshot(self) -> Dict[str, int]:
103 return {
104 "hits": self.hits,
105 "misses": self.misses,
106 "sets": self.sets,
107 "deletes": self.deletes,
108 "evictions": self.evictions,
109 "errors": self.errors,
110 }
113# ---------------------------------------------------------------------------
114# Backends
115# ---------------------------------------------------------------------------
117class CacheBackend(ABC):
118 """Abstract cache backend."""
120 @abstractmethod
121 async def get(self, key: str) -> Optional[bytes]: ...
123 @abstractmethod
124 async def set(self, key: str, value: bytes, ttl: Optional[float] = None) -> None: ...
126 @abstractmethod
127 async def delete(self, key: str) -> bool: ...
129 @abstractmethod
130 async def exists(self, key: str) -> bool: ...
132 @abstractmethod
133 async def clear(self) -> None: ...
135 async def get_many(self, keys: List[str]) -> Dict[str, bytes]:
136 results = await asyncio.gather(*(self.get(k) for k in keys), return_exceptions=True)
137 return {k: v for k, v in zip(keys, results)
138 if not isinstance(v, (Exception, type(None)))}
140 async def set_many(self, items: Dict[str, bytes], ttl: Optional[float] = None) -> None:
141 await asyncio.gather(*(self.set(k, v, ttl) for k, v in items.items()))
143 async def delete_many(self, keys: List[str]) -> int:
144 results = await asyncio.gather(*(self.delete(k) for k in keys), return_exceptions=True)
145 return sum(1 for r in results if r is True)
148class MemoryCacheBackend(CacheBackend):
149 """In-memory LRU cache with TTL support and stampede protection."""
151 def __init__(self, max_size: int = 10_000, default_ttl: Optional[float] = 300.0):
152 self._max_size = max_size
153 self._default_ttl = default_ttl
154 self._store: OrderedDict[str, Tuple[bytes, float, Optional[float]]] = OrderedDict()
155 # OrderedDict: key → (value, inserted_at, custom_ttl)
157 @property
158 def size(self) -> int:
159 return len(self._store)
161 def _evict_expired(self):
162 now = time.monotonic()
163 expired = []
164 for key, (_, inserted, ttl) in self._store.items():
165 effective_ttl = ttl if ttl is not None else self._default_ttl
166 if effective_ttl is not None and now - inserted > effective_ttl:
167 expired.append(key)
168 for key in expired:
169 del self._store[key]
171 def _evict_lru(self):
172 while len(self._store) > self._max_size:
173 self._store.popitem(last=False)
175 async def get(self, key: str) -> Optional[bytes]:
176 self._evict_expired()
177 entry = self._store.get(key)
178 if entry is None:
179 return None
180 value, inserted, ttl = entry
181 effective_ttl = ttl if ttl is not None else self._default_ttl
182 if effective_ttl is not None and time.monotonic() - inserted > effective_ttl:
183 del self._store[key]
184 return None
185 # LRU: move to end
186 self._store.move_to_end(key)
187 return value
189 async def set(self, key: str, value: bytes, ttl: Optional[float] = None) -> None:
190 self._evict_expired()
191 self._store[key] = (value, time.monotonic(), ttl)
192 self._store.move_to_end(key)
193 self._evict_lru()
195 async def delete(self, key: str) -> bool:
196 if key in self._store:
197 del self._store[key]
198 return True
199 return False
201 async def exists(self, key: str) -> bool:
202 return (await self.get(key)) is not None
204 async def clear(self) -> None:
205 self._store.clear()
208class RedisCacheBackend(CacheBackend):
209 """Redis cache backend using async Redis client.
211 Requires: pip install redis[hiredis]
212 """
214 def __init__(self, url: str = "redis://localhost:6379/0",
215 default_ttl: Optional[float] = 300.0,
216 prefix: str = "agentos:cache:"):
217 self._url = url
218 self._default_ttl = default_ttl
219 self._prefix = prefix
220 self._client: Any = None
222 async def _ensure_client(self):
223 if self._client is None:
224 try:
225 import redis.asyncio as aioredis
226 except ImportError:
227 raise CacheBackendUnavailable(
228 "redis package not installed. Run: pip install redis[hiredis]"
229 )
230 self._client = aioredis.from_url(self._url)
232 def _key(self, raw: str) -> str:
233 return f"{self._prefix}{raw}"
235 async def get(self, key: str) -> Optional[bytes]:
236 await self._ensure_client()
237 return await self._client.get(self._key(key))
239 async def set(self, key: str, value: bytes, ttl: Optional[float] = None) -> None:
240 await self._ensure_client()
241 ttl_val = ttl if ttl is not None else self._default_ttl
242 if ttl_val is not None:
243 await self._client.setex(self._key(key), int(ttl_val), value)
244 else:
245 await self._client.set(self._key(key), value)
247 async def delete(self, key: str) -> bool:
248 await self._ensure_client()
249 return bool(await self._client.delete(self._key(key)))
251 async def exists(self, key: str) -> bool:
252 await self._ensure_client()
253 return bool(await self._client.exists(self._key(key)))
255 async def clear(self) -> None:
256 await self._ensure_client()
257 pattern = f"{self._prefix}*"
258 cursor = 0
259 while True:
260 cursor, keys = await self._client.scan(cursor, match=pattern, count=100)
261 if keys:
262 await self._client.delete(*keys)
263 if cursor == 0:
264 break
266 async def incr(self, key: str, amount: int = 1) -> int:
267 await self._ensure_client()
268 return await self._client.incrby(self._key(key), amount)
271# ---------------------------------------------------------------------------
272# Cache Manager
273# ---------------------------------------------------------------------------
275@dataclass
276class CacheConfig:
277 """Cache configuration."""
278 serializer: Serializer = field(default_factory=PickleSerializer)
279 key_prefix: str = ""
280 hash_keys: bool = False # SHA-256 hash long keys
281 stampede_protection: bool = True
282 stampede_beta: float = 1.0 # recompute window multiplier
283 stampede_delta: float = 0.0 # extra fixed window
284 log_stats: bool = False
287class Cache(Generic[T]):
288 """High-level cache API with tiered backends and stampede protection.
290 Usage:
291 cache = Cache[str](backend=MemoryCacheBackend(max_size=1000))
292 await cache.set("user:1", "Alice", ttl=60)
293 name = await cache.get("user:1")
294 user = await cache.get_or_set("user:1", lambda: db.fetch("user:1"), ttl=60)
295 """
297 def __init__(
298 self,
299 backend: CacheBackend,
300 config: Optional[CacheConfig] = None,
301 ):
302 self._backend = backend
303 self._config = config or CacheConfig()
304 self._stats = CacheStats()
305 self._lock = asyncio.Lock()
307 @property
308 def stats(self) -> CacheStats:
309 return self._stats
311 def _make_key(self, key: str) -> str:
312 full = f"{self._config.key_prefix}{key}"
313 if self._config.hash_keys:
314 return hashlib.sha256(full.encode()).hexdigest()
315 return full
317 # -- Core ops --
319 async def get(self, key: str) -> Optional[T]:
320 try:
321 raw = await self._backend.get(self._make_key(key))
322 except Exception as exc:
323 self._stats.errors += 1
324 logger.warning("Cache get error: %s", exc)
325 return None
326 if raw is None:
327 self._stats.misses += 1
328 return None
329 self._stats.hits += 1
330 try:
331 return self._config.serializer.loads(raw)
332 except Exception:
333 return None
335 async def set(self, key: str, value: T, ttl: Optional[float] = None) -> None:
336 try:
337 data = self._config.serializer.dumps(value)
338 await self._backend.set(self._make_key(key), data, ttl)
339 self._stats.sets += 1
340 except Exception as exc:
341 self._stats.errors += 1
342 logger.warning("Cache set error: %s", exc)
344 async def delete(self, key: str) -> bool:
345 try:
346 result = await self._backend.delete(self._make_key(key))
347 if result:
348 self._stats.deletes += 1
349 return result
350 except Exception as exc:
351 self._stats.errors += 1
352 logger.warning("Cache delete error: %s", exc)
353 return False
355 async def exists(self, key: str) -> bool:
356 try:
357 return await self._backend.exists(self._make_key(key))
358 except Exception:
359 return False
361 # -- Atomic get-or-set with stampede protection --
363 async def get_or_set(
364 self, key: str, factory: Callable[[], Any],
365 ttl: Optional[float] = None,
366 force_refresh: bool = False,
367 ) -> T:
368 """Get from cache, or compute via factory and store.
369 Stampede protection: probabilistically refreshes early when near expiry.
370 """
371 if not force_refresh:
372 cached = await self.get(key)
373 if cached is not None:
374 return cached
376 # Stampede protection: if another coroutine is already computing,
377 # wait briefly for it to finish.
378 async with self._lock:
379 # Double-check after acquiring lock
380 if not force_refresh:
381 cached = await self.get(key)
382 if cached is not None:
383 return cached
384 try:
385 value = factory()
386 if asyncio.iscoroutine(value):
387 value = await value
388 except Exception:
389 raise
390 await self.set(key, value, ttl)
391 return value
393 async def get_or_default(self, key: str, default: T) -> T:
394 result = await self.get(key)
395 return result if result is not None else default
397 # -- Bulk ops --
399 async def get_many(self, keys: List[str]) -> Dict[str, Optional[T]]:
400 try:
401 cache_keys = [self._make_key(k) for k in keys]
402 raw_map = await self._backend.get_many(cache_keys)
403 except Exception:
404 return {k: None for k in keys}
405 result: Dict[str, Optional[T]] = {}
406 for k, ck in zip(keys, cache_keys):
407 raw = raw_map.get(ck)
408 if raw is not None:
409 self._stats.hits += 1
410 try:
411 result[k] = self._config.serializer.loads(raw)
412 except Exception:
413 result[k] = None
414 else:
415 self._stats.misses += 1
416 result[k] = None
417 return result
419 async def set_many(self, mapping: Dict[str, T], ttl: Optional[float] = None) -> None:
420 try:
421 items = {
422 self._make_key(k): self._config.serializer.dumps(v)
423 for k, v in mapping.items()
424 }
425 await self._backend.set_many(items, ttl)
426 self._stats.sets += len(items)
427 except Exception:
428 self._stats.errors += 1
430 async def delete_many(self, keys: List[str]) -> int:
431 try:
432 count = await self._backend.delete_many(
433 [self._make_key(k) for k in keys]
434 )
435 self._stats.deletes += count
436 return count
437 except Exception:
438 self._stats.errors += 1
439 return 0
441 async def clear(self) -> None:
442 try:
443 await self._backend.clear()
444 except Exception as exc:
445 self._stats.errors += 1
446 logger.warning("Cache clear error: %s", exc)
449# ---------------------------------------------------------------------------
450# Tiered Cache
451# ---------------------------------------------------------------------------
453class TieredCache(Generic[T]):
454 """Two-tier cache: L1 (fast, small) → L2 (slower, larger).
456 L1: typically MemoryCacheBackend
457 L2: typically RedisCacheBackend
458 """
460 def __init__(self, l1: Cache[T], l2: Cache[T],
461 promote_on_read: bool = True):
462 self.l1 = l1
463 self.l2 = l2
464 self._promote_on_read = promote_on_read
466 async def get(self, key: str) -> Optional[T]:
467 # Try L1
468 value = await self.l1.get(key)
469 if value is not None:
470 return value
471 # Try L2
472 value = await self.l2.get(key)
473 if value is not None and self._promote_on_read:
474 await self.l1.set(key, value)
475 return value
477 async def set(self, key: str, value: T, ttl: Optional[float] = None) -> None:
478 await asyncio.gather(
479 self.l1.set(key, value, ttl),
480 self.l2.set(key, value, ttl),
481 )
483 async def delete(self, key: str) -> bool:
484 r1, r2 = await asyncio.gather(
485 self.l1.delete(key),
486 self.l2.delete(key),
487 )
488 return r1 or r2
490 async def clear(self) -> None:
491 await asyncio.gather(self.l1.clear(), self.l2.clear())
493 @property
494 def stats(self) -> Dict[str, Any]:
495 return {
496 "l1": self.l1.stats.snapshot(),
497 "l2": self.l2.stats.snapshot(),
498 }
501# ---------------------------------------------------------------------------
502# Decorator
503# ---------------------------------------------------------------------------
505def cached(
506 cache_instance: Cache,
507 key_prefix: str = "",
508 ttl: Optional[float] = None,
509 key_builder: Optional[Callable[..., str]] = None,
510):
511 """Decorator to cache async function results.
513 Usage:
514 user_cache = Cache[dict](MemoryCacheBackend())
516 @cached(user_cache, key_prefix="user", ttl=300)
517 async def get_user(user_id: str) -> dict:
518 return await db.fetch_user(user_id)
519 """
520 def decorator(fn):
521 @wraps(fn)
522 async def wrapper(*args, **kwargs):
523 if key_builder:
524 cache_key = key_builder(*args, **kwargs)
525 else:
526 sig = _build_signature(args, kwargs)
527 cache_key = f"{key_prefix}:{fn.__name__}:{sig}"
528 result = await cache_instance.get(cache_key)
529 if result is not None:
530 return result
531 result = await fn(*args, **kwargs)
532 await cache_instance.set(cache_key, result, ttl)
533 return result
534 return wrapper
535 return decorator
538def _build_signature(args: tuple, kwargs: dict) -> str:
539 parts = [str(a) for a in args]
540 parts.extend(f"{k}={v}" for k, v in sorted(kwargs.items()))
541 raw = ":".join(parts)
542 if len(raw) > 128:
543 return hashlib.md5(raw.encode()).hexdigest()
544 return raw