Coverage for agentos/core/cache.py: 0%

308 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-09 09:19 +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 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) 

35 

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

37 

38T = TypeVar("T") 

39 

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

41# Exceptions 

42# --------------------------------------------------------------------------- 

43 

44 

45class CacheError(Exception): 

46 """Base cache error.""" 

47 

48 

49class CacheBackendUnavailable(CacheError): # noqa: N818 

50 """Backend is down or unreachable.""" 

51 

52 

53class SerializationError(CacheError): 

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

55 

56 

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

58# Serializer 

59# --------------------------------------------------------------------------- 

60 

61 

62class Serializer(ABC): 

63 """Serialization interface for cache values.""" 

64 

65 @abstractmethod 

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

67 

68 @abstractmethod 

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

70 

71 

72class PickleSerializer(Serializer): 

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

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

75 

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

77 return pickle.loads(data) 

78 

79 

80class JSONSerializer(Serializer): 

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

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

83 

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

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

86 

87 

88# --------------------------------------------------------------------------- 

89# Stats 

90# --------------------------------------------------------------------------- 

91 

92 

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 

101 

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 

106 

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 } 

116 

117 

118# --------------------------------------------------------------------------- 

119# Backends 

120# --------------------------------------------------------------------------- 

121 

122 

123class CacheBackend(ABC): 

124 """Abstract cache backend.""" 

125 

126 @abstractmethod 

127 async def get(self, key: str) -> bytes | None: ... 

128 

129 @abstractmethod 

130 async def set(self, key: str, value: bytes, ttl: float | None = None) -> None: ... 

131 

132 @abstractmethod 

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

134 

135 @abstractmethod 

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

137 

138 @abstractmethod 

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

140 

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

144 

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

147 

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) 

151 

152 

153class MemoryCacheBackend(CacheBackend): 

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

155 

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) 

161 

162 @property 

163 def size(self) -> int: 

164 return len(self._store) 

165 

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] 

175 

176 def _evict_lru(self): 

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

178 self._store.popitem(last=False) 

179 

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 

193 

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

199 

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 

205 

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

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

208 

209 async def clear(self) -> None: 

210 self._store.clear() 

211 

212 

213class RedisCacheBackend(CacheBackend): 

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

215 

216 Requires: pip install redis[hiredis] 

217 """ 

218 

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 

229 

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) 

239 

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

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

242 

243 async def get(self, key: str) -> bytes | None: 

244 await self._ensure_client() 

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

246 

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) 

254 

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

256 await self._ensure_client() 

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

258 

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

260 await self._ensure_client() 

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

262 

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 

273 

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) 

277 

278 

279# --------------------------------------------------------------------------- 

280# Cache Manager 

281# --------------------------------------------------------------------------- 

282 

283 

284@dataclass 

285class CacheConfig: 

286 """Cache configuration.""" 

287 

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 

295 

296 

297class Cache(Generic[T]): 

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

299 

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

306 

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

316 

317 @property 

318 def stats(self) -> CacheStats: 

319 return self._stats 

320 

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 

326 

327 # -- Core ops -- 

328 

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 

344 

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) 

353 

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 

364 

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 

370 

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

372 

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 

387 

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 

404 

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 

408 

409 # -- Bulk ops -- 

410 

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 

430 

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 

440 

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 

449 

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) 

456 

457 

458# --------------------------------------------------------------------------- 

459# Tiered Cache 

460# --------------------------------------------------------------------------- 

461 

462 

463class TieredCache(Generic[T]): 

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

465 

466 L1: typically MemoryCacheBackend 

467 L2: typically RedisCacheBackend 

468 """ 

469 

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 

474 

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 

485 

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 ) 

491 

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 

498 

499 async def clear(self) -> None: 

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

501 

502 @property 

503 def stats(self) -> dict[str, Any]: 

504 return { 

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

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

507 } 

508 

509 

510# --------------------------------------------------------------------------- 

511# Decorator 

512# --------------------------------------------------------------------------- 

513 

514 

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. 

522 

523 Usage: 

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

525 

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

530 

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 

545 

546 return wrapper 

547 

548 return decorator 

549 

550 

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