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