File size: 4,807 Bytes
51501f2 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 | 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()
|