import time import asyncio import math from typing import Dict, Optional, Tuple, List import numpy as np from collections import OrderedDict def cosine_sim(a: List[float], b: List[float]) -> float: if not a or not b: return 0.0 a_arr = np.array(a, dtype=np.float32) b_arr = np.array(b, dtype=np.float32) norm_a = np.linalg.norm(a_arr) norm_b = np.linalg.norm(b_arr) if norm_a == 0 or norm_b == 0: return 0.0 return float(np.dot(a_arr, b_arr) / (norm_a * norm_b)) class NexusCache: def __init__(self, max_size: int = 500, ttl_seconds: int = 3600, semantic_threshold: float = 0.93): self.max_size = max_size self.ttl = ttl_seconds self.semantic_threshold = semantic_threshold self.exact_cache: Dict[str, Dict] = {} self.semantic_entries: OrderedDict[str, Dict] = OrderedDict() self.lock = asyncio.Lock() async def get_exact(self, key: str) -> Optional[Dict]: async with self.lock: entry = self.exact_cache.get(key) if not entry: return None if time.time() - entry["timestamp"] > self.ttl: del self.exact_cache[key] # also remove from semantic if same key if key in self.semantic_entries: del self.semantic_entries[key] return None # update LRU if key in self.semantic_entries: self.semantic_entries.move_to_end(key) return entry async def get_semantic(self, embedding: List[float]) -> Tuple[Optional[Dict], float]: """ Search for semantically similar cached query Returns (entry, similarity) or (None, 0) """ async with self.lock: if not self.semantic_entries: return None, 0.0 best_score = 0.0 best_entry = None best_key = None # iterate from most recent to oldest for speed for k, entry in reversed(self.semantic_entries.items()): if time.time() - entry["timestamp"] > self.ttl: continue emb = entry.get("embedding") if not emb: continue sim = cosine_sim(embedding, emb) if sim > best_score: best_score = sim best_entry = entry best_key = k if sim > 0.98: # early exit if almost identical break if best_score >= self.semantic_threshold and best_entry: # move to end as LRU if best_key: self.semantic_entries.move_to_end(best_key) return best_entry, best_score return None, 0.0 async def set(self, key: str, embedding: List[float], response_data: Dict, sources: List): async with self.lock: now = time.time() entry = { "timestamp": now, "response": response_data, "embedding": embedding, "sources": sources, } self.exact_cache[key] = entry self.semantic_entries[key] = entry self.semantic_entries.move_to_end(key) # eviction if over size while len(self.semantic_entries) > self.max_size: oldest_key, _ = self.semantic_entries.popitem(last=False) if oldest_key in self.exact_cache: del self.exact_cache[oldest_key] async def cleanup(self): async with self.lock: now = time.time() expired_keys = [] for k, v in self.exact_cache.items(): if now - v["timestamp"] > self.ttl: expired_keys.append(k) for k in expired_keys: if k in self.exact_cache: del self.exact_cache[k] if k in self.semantic_entries: del self.semantic_entries[k] # also clean semantic entries that are not in exact but expired expired_sem = [] for k, v in self.semantic_entries.items(): if now - v["timestamp"] > self.ttl: expired_sem.append(k) for k in expired_sem: if k in self.semantic_entries: del self.semantic_entries[k] def stats(self) -> Dict: return { "exact_count": len(self.exact_cache), "semantic_count": len(self.semantic_entries), "max_size": self.max_size, "threshold": self.semantic_threshold, } def clear(self): self.exact_cache.clear() self.semantic_entries.clear()