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

1""" 

2Production-grade multi-backend cache with tiered architecture. 

3 

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 

13 

14Copyright 2026 AgentOS. All rights reserved. 

15""" 

16 

17from __future__ import annotations 

18 

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) 

32 

33import logging 

34 

35logger = logging.getLogger("agentos.cache") 

36 

37T = TypeVar("T") 

38 

39# --------------------------------------------------------------------------- 

40# Exceptions 

41# --------------------------------------------------------------------------- 

42 

43class CacheError(Exception): 

44 """Base cache error.""" 

45 

46 

47class CacheBackendUnavailable(CacheError): 

48 """Backend is down or unreachable.""" 

49 

50 

51class SerializationError(CacheError): 

52 """Failed to serialize/deserialize a cached value.""" 

53 

54 

55# --------------------------------------------------------------------------- 

56# Serializer 

57# --------------------------------------------------------------------------- 

58 

59class Serializer(ABC): 

60 """Serialization interface for cache values.""" 

61 

62 @abstractmethod 

63 def dumps(self, value: Any) -> bytes: ... 

64 

65 @abstractmethod 

66 def loads(self, data: bytes) -> Any: ... 

67 

68 

69class PickleSerializer(Serializer): 

70 def dumps(self, value: Any) -> bytes: 

71 return pickle.dumps(value, protocol=pickle.HIGHEST_PROTOCOL) 

72 

73 def loads(self, data: bytes) -> Any: 

74 return pickle.loads(data) 

75 

76 

77class JSONSerializer(Serializer): 

78 def dumps(self, value: Any) -> bytes: 

79 return json.dumps(value, default=str, separators=(",", ":")).encode() 

80 

81 def loads(self, data: bytes) -> Any: 

82 return json.loads(data.decode()) 

83 

84 

85# --------------------------------------------------------------------------- 

86# Stats 

87# --------------------------------------------------------------------------- 

88 

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 

97 

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 

102 

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 } 

112 

113 

114# --------------------------------------------------------------------------- 

115# Backends 

116# --------------------------------------------------------------------------- 

117 

118class CacheBackend(ABC): 

119 """Abstract cache backend.""" 

120 

121 @abstractmethod 

122 async def get(self, key: str) -> Optional[bytes]: ... 

123 

124 @abstractmethod 

125 async def set(self, key: str, value: bytes, ttl: Optional[float] = None) -> None: ... 

126 

127 @abstractmethod 

128 async def delete(self, key: str) -> bool: ... 

129 

130 @abstractmethod 

131 async def exists(self, key: str) -> bool: ... 

132 

133 @abstractmethod 

134 async def clear(self) -> None: ... 

135 

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)))} 

140 

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())) 

143 

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) 

147 

148 

149class MemoryCacheBackend(CacheBackend): 

150 """In-memory LRU cache with TTL support and stampede protection.""" 

151 

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) 

157 

158 @property 

159 def size(self) -> int: 

160 return len(self._store) 

161 

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] 

171 

172 def _evict_lru(self): 

173 while len(self._store) > self._max_size: 

174 self._store.popitem(last=False) 

175 

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 

189 

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() 

195 

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 

201 

202 async def exists(self, key: str) -> bool: 

203 return (await self.get(key)) is not None 

204 

205 async def clear(self) -> None: 

206 self._store.clear() 

207 

208 

209class RedisCacheBackend(CacheBackend): 

210 """Redis cache backend using async Redis client. 

211 

212 Requires: pip install redis[hiredis] 

213 """ 

214 

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 

222 

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) 

232 

233 def _key(self, raw: str) -> str: 

234 return f"{self._prefix}{raw}" 

235 

236 async def get(self, key: str) -> Optional[bytes]: 

237 await self._ensure_client() 

238 return await self._client.get(self._key(key)) 

239 

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) 

247 

248 async def delete(self, key: str) -> bool: 

249 await self._ensure_client() 

250 return bool(await self._client.delete(self._key(key))) 

251 

252 async def exists(self, key: str) -> bool: 

253 await self._ensure_client() 

254 return bool(await self._client.exists(self._key(key))) 

255 

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 

266 

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) 

270 

271 

272# --------------------------------------------------------------------------- 

273# Cache Manager 

274# --------------------------------------------------------------------------- 

275 

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 

286 

287 

288class Cache(Generic[T]): 

289 """High-level cache API with tiered backends and stampede protection. 

290 

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 """ 

297 

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() 

307 

308 @property 

309 def stats(self) -> CacheStats: 

310 return self._stats 

311 

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 

317 

318 # -- Core ops -- 

319 

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 

335 

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) 

344 

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 

355 

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 

361 

362 # -- Atomic get-or-set with stampede protection -- 

363 

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 

376 

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 

393 

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 

397 

398 # -- Bulk ops -- 

399 

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 

419 

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 

430 

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 

441 

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) 

448 

449 

450# --------------------------------------------------------------------------- 

451# Tiered Cache 

452# --------------------------------------------------------------------------- 

453 

454class TieredCache(Generic[T]): 

455 """Two-tier cache: L1 (fast, small) → L2 (slower, larger). 

456 

457 L1: typically MemoryCacheBackend 

458 L2: typically RedisCacheBackend 

459 """ 

460 

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 

466 

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 

477 

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 ) 

483 

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 

490 

491 async def clear(self) -> None: 

492 await asyncio.gather(self.l1.clear(), self.l2.clear()) 

493 

494 @property 

495 def stats(self) -> Dict[str, Any]: 

496 return { 

497 "l1": self.l1.stats.snapshot(), 

498 "l2": self.l2.stats.snapshot(), 

499 } 

500 

501 

502# --------------------------------------------------------------------------- 

503# Decorator 

504# --------------------------------------------------------------------------- 

505 

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. 

513 

514 Usage: 

515 user_cache = Cache[dict](MemoryCacheBackend()) 

516 

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 

537 

538 

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