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

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

31 

32import logging 

33 

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

35 

36T = TypeVar("T") 

37 

38# --------------------------------------------------------------------------- 

39# Exceptions 

40# --------------------------------------------------------------------------- 

41 

42class CacheError(Exception): 

43 """Base cache error.""" 

44 

45 

46class CacheBackendUnavailable(CacheError): 

47 """Backend is down or unreachable.""" 

48 

49 

50class SerializationError(CacheError): 

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

52 

53 

54# --------------------------------------------------------------------------- 

55# Serializer 

56# --------------------------------------------------------------------------- 

57 

58class Serializer(ABC): 

59 """Serialization interface for cache values.""" 

60 

61 @abstractmethod 

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

63 

64 @abstractmethod 

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

66 

67 

68class PickleSerializer(Serializer): 

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

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

71 

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

73 return pickle.loads(data) 

74 

75 

76class JSONSerializer(Serializer): 

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

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

79 

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

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

82 

83 

84# --------------------------------------------------------------------------- 

85# Stats 

86# --------------------------------------------------------------------------- 

87 

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 

96 

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 

101 

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 } 

111 

112 

113# --------------------------------------------------------------------------- 

114# Backends 

115# --------------------------------------------------------------------------- 

116 

117class CacheBackend(ABC): 

118 """Abstract cache backend.""" 

119 

120 @abstractmethod 

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

122 

123 @abstractmethod 

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

125 

126 @abstractmethod 

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

128 

129 @abstractmethod 

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

131 

132 @abstractmethod 

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

134 

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

139 

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

142 

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) 

146 

147 

148class MemoryCacheBackend(CacheBackend): 

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

150 

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) 

156 

157 @property 

158 def size(self) -> int: 

159 return len(self._store) 

160 

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] 

170 

171 def _evict_lru(self): 

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

173 self._store.popitem(last=False) 

174 

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 

188 

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

194 

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 

200 

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

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

203 

204 async def clear(self) -> None: 

205 self._store.clear() 

206 

207 

208class RedisCacheBackend(CacheBackend): 

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

210 

211 Requires: pip install redis[hiredis] 

212 """ 

213 

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 

221 

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) 

231 

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

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

234 

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

236 await self._ensure_client() 

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

238 

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) 

246 

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

248 await self._ensure_client() 

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

250 

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

252 await self._ensure_client() 

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

254 

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 

265 

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) 

269 

270 

271# --------------------------------------------------------------------------- 

272# Cache Manager 

273# --------------------------------------------------------------------------- 

274 

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 

285 

286 

287class Cache(Generic[T]): 

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

289 

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

296 

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

306 

307 @property 

308 def stats(self) -> CacheStats: 

309 return self._stats 

310 

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 

316 

317 # -- Core ops -- 

318 

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 

334 

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) 

343 

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 

354 

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 

360 

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

362 

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 

375 

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 

392 

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 

396 

397 # -- Bulk ops -- 

398 

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 

418 

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 

429 

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 

440 

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) 

447 

448 

449# --------------------------------------------------------------------------- 

450# Tiered Cache 

451# --------------------------------------------------------------------------- 

452 

453class TieredCache(Generic[T]): 

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

455 

456 L1: typically MemoryCacheBackend 

457 L2: typically RedisCacheBackend 

458 """ 

459 

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 

465 

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 

476 

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 ) 

482 

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 

489 

490 async def clear(self) -> None: 

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

492 

493 @property 

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

495 return { 

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

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

498 } 

499 

500 

501# --------------------------------------------------------------------------- 

502# Decorator 

503# --------------------------------------------------------------------------- 

504 

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. 

512 

513 Usage: 

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

515 

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 

536 

537 

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