Coverage for agentos/memory/retriever.py: 29%
178 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +0800
1"""
2Semantic Memory Retriever — Embedding-based memory retrieval with hybrid search.
4Supports semantic (embedding), keyword (BM25), and hybrid search across
5conversation memory, long-term memory, and working memory. Aligns with
6ConversationMemory window strategies and LongTermMemory persistence.
7"""
9from __future__ import annotations
11import math
12from collections import Counter
13from dataclasses import dataclass, field
14from enum import Enum
15from typing import Any, Callable, Optional
17import numpy as np
20class RetrievalStrategy(Enum):
22 """检索策略枚举。"""
24 SEMANTIC = "semantic"
25 KEYWORD = "keyword"
26 HYBRID = "hybrid"
27 RECENT = "recent"
30@dataclass
31class MemoryEntry:
32 """A single memory entry with content and metadata."""
34 id: str
35 content: str
36 metadata: dict[str, Any] = field(default_factory=dict)
37 embedding: Optional[list[float]] = None
38 timestamp: Optional[float] = None
39 importance: float = 0.5
40 source: str = "conversation" # conversation / long_term / working
43@dataclass
44class RetrievalResult:
45 """A single retrieval result with relevance score."""
47 entry: MemoryEntry
48 score: float
49 strategy: RetrievalStrategy
52@dataclass
53class RetrievalStats:
54 """Statistics for a retrieval operation."""
56 total_entries: int
57 retrieved: int
58 strategies_used: list[RetrievalStrategy] = field(default_factory=list)
59 latency_ms: float = 0.0
62class SemanticMemoryRetriever:
63 """
64 Semantic retrieval engine for AgentOS memory systems.
66 Supports three retrieval strategies:
67 - **semantic**: Cosine similarity over embeddings (requires embedder)
68 - **keyword**: BM25-style TF-IDF keyword matching (no embedder needed)
69 - **hybrid**: Weighted combination of semantic + keyword scores
71 Example::
73 retriever = SemanticMemoryRetriever(embedder=my_embedder)
74 results = retriever.retrieve(
75 "What did we discuss about deployment?",
76 top_k=5,
77 strategy=RetrievalStrategy.HYBRID,
78 )
79 for r in results:
80 print(f"[{r.score:.2f}] {r.entry.content[:80]}...")
81 """
83 def __init__(
84 self,
85 embedder: Optional[Callable[[str], list[float]]] = None,
86 hybrid_weight: float = 0.7,
87 min_keyword_score: float = 0.01,
88 default_top_k: int = 10,
89 ):
90 """
91 Args:
92 embedder: Callable that takes text and returns embedding vector.
93 hybrid_weight: Weight for semantic score in hybrid mode (0-1).
94 Remaining weight goes to keyword score.
95 min_keyword_score: Minimum BM25 score to include in results.
96 default_top_k: Default number of results to return.
97 """
98 self._embedder = embedder
99 self._hybrid_weight = hybrid_weight
100 self._min_keyword_score = min_keyword_score
101 self._default_top_k = default_top_k
102 self._entries: dict[str, MemoryEntry] = {}
103 self._idf_cache: dict[str, float] = {}
104 self._doc_freqs: Counter[str, int] = Counter()
105 self._total_docs: int = 0
107 def index(self, entries: list[MemoryEntry]) -> None:
108 """Add entries to the search index."""
109 for entry in entries:
110 self._entries[entry.id] = entry
111 if entry.embedding and self._embedder:
112 # Already has embedding, no need to re-embed
113 pass
114 elif self._embedder:
115 entry.embedding = self._embedder(entry.content)
117 # Update keyword index
118 tokens = self._tokenize(entry.content)
119 unique_tokens = set(tokens)
120 self._doc_freqs.update(unique_tokens)
121 self._total_docs += 1
123 def remove(self, entry_ids: list[str]) -> None:
124 """Remove entries from the index."""
125 for eid in entry_ids:
126 if eid in self._entries:
127 entry = self._entries.pop(eid)
128 unique_tokens = set(self._tokenize(entry.content))
129 for token in unique_tokens:
130 self._doc_freqs[token] = max(0, self._doc_freqs[token] - 1)
131 self._total_docs = max(0, self._total_docs - 1)
132 self._idf_cache.clear()
134 def retrieve(
135 self,
136 query: str,
137 top_k: Optional[int] = None,
138 strategy: RetrievalStrategy = RetrievalStrategy.HYBRID,
139 filter_source: Optional[str] = None,
140 min_importance: float = 0.0,
141 ) -> list[RetrievalResult]:
142 """
143 Retrieve the most relevant memories for a query.
145 Args:
146 query: Search query.
147 top_k: Number of results to return.
148 strategy: Retrieval strategy.
149 filter_source: Only return entries from this source.
150 min_importance: Minimum importance score filter.
152 Returns:
153 List of RetrievalResult sorted by relevance.
154 """
155 import time
156 start = time.perf_counter()
157 top_k = top_k or self._default_top_k
159 # Filter entries
160 candidates = [
161 e for e in self._entries.values()
162 if (filter_source is None or e.source == filter_source)
163 and e.importance >= min_importance
164 ]
166 if not candidates:
167 return []
169 if strategy == RetrievalStrategy.RECENT:
170 results = self._retrieve_recent(candidates, top_k)
171 elif strategy == RetrievalStrategy.KEYWORD:
172 results = self._retrieve_keyword(query, candidates, top_k)
173 elif strategy == RetrievalStrategy.SEMANTIC:
174 results = self._retrieve_semantic(query, candidates, top_k)
175 else: # HYBRID
176 results = self._retrieve_hybrid(query, candidates, top_k)
178 elapsed = (time.perf_counter() - start) * 1000
179 # Attach stats to results via a common approach
180 return results
182 def _retrieve_recent(
183 self, candidates: list[MemoryEntry], top_k: int,
184 ) -> list[RetrievalResult]:
185 """Return most recent entries sorted by timestamp."""
186 sorted_entries = sorted(
187 candidates,
188 key=lambda e: e.timestamp or 0,
189 reverse=True,
190 )
191 return [
192 RetrievalResult(
193 entry=e, score=1.0, strategy=RetrievalStrategy.RECENT,
194 )
195 for e in sorted_entries[:top_k]
196 ]
198 def _retrieve_keyword(
199 self, query: str, candidates: list[MemoryEntry], top_k: int,
200 ) -> list[RetrievalResult]:
201 """BM25-style keyword search."""
202 query_tokens = self._tokenize(query)
203 if not query_tokens:
204 return []
206 scores = []
207 for entry in candidates:
208 score = self._bm25_score(query_tokens, entry.content)
209 if score >= self._min_keyword_score:
210 scores.append((score, entry))
212 scores.sort(key=lambda x: x[0], reverse=True)
213 return [
214 RetrievalResult(
215 entry=e, score=s, strategy=RetrievalStrategy.KEYWORD,
216 )
217 for s, e in scores[:top_k]
218 ]
220 def _retrieve_semantic(
221 self, query: str, candidates: list[MemoryEntry], top_k: int,
222 ) -> list[RetrievalResult]:
223 """Cosine similarity semantic search."""
224 if not self._embedder:
225 return self._retrieve_keyword(query, candidates, top_k)
227 query_embedding = np.array(self._embedder(query))
228 scores = []
229 for entry in candidates:
230 if entry.embedding is None:
231 entry.embedding = self._embedder(entry.content)
232 entry_embedding = np.array(entry.embedding)
233 similarity = self._cosine_sim(query_embedding, entry_embedding)
234 scores.append((similarity, entry))
236 scores.sort(key=lambda x: x[0], reverse=True)
237 return [
238 RetrievalResult(
239 entry=e, score=float(s), strategy=RetrievalStrategy.SEMANTIC,
240 )
241 for s, e in scores[:top_k]
242 ]
244 def _retrieve_hybrid(
245 self, query: str, candidates: list[MemoryEntry], top_k: int,
246 ) -> list[RetrievalResult]:
247 """Weighted combination of semantic and keyword scores."""
248 query_tokens = self._tokenize(query)
249 has_embedder = self._embedder is not None
251 if has_embedder:
252 query_embedding = np.array(self._embedder(query))
254 scores = []
255 for entry in candidates:
256 kw_score = self._bm25_score(query_tokens, entry.content)
258 if has_embedder:
259 if entry.embedding is None:
260 entry.embedding = self._embedder(entry.content)
261 entry_embedding = np.array(entry.embedding)
262 sem_score = self._cosine_sim(query_embedding, entry_embedding)
263 combined = (
264 self._hybrid_weight * sem_score
265 + (1 - self._hybrid_weight) * kw_score
266 )
267 else:
268 combined = kw_score
270 if combined > 0:
271 scores.append((combined, entry))
273 scores.sort(key=lambda x: x[0], reverse=True)
274 return [
275 RetrievalResult(
276 entry=e, score=float(s), strategy=RetrievalStrategy.HYBRID,
277 )
278 for s, e in scores[:top_k]
279 ]
281 # --- BM25 implementation ---
283 @staticmethod
284 def _tokenize(text: str) -> list[str]:
285 """Simple word tokenizer."""
286 text = text.lower()
287 # Split on non-alphanumeric, keep sequences of 2+ chars
288 tokens = []
289 current = []
290 for ch in text:
291 if ch.isalnum():
292 current.append(ch)
293 else:
294 if len(current) >= 2:
295 tokens.append("".join(current))
296 current = []
297 if len(current) >= 2:
298 tokens.append("".join(current))
299 return tokens
301 def _idf(self, term: str) -> float:
302 """Inverse document frequency."""
303 if term not in self._idf_cache:
304 df = self._doc_freqs.get(term, 0)
305 if df == 0 or self._total_docs == 0:
306 self._idf_cache[term] = 0.0
307 else:
308 self._idf_cache[term] = math.log(
309 (self._total_docs - df + 0.5) / (df + 0.5) + 1.0
310 )
311 return self._idf_cache[term]
313 def _bm25_score(
314 self, query_tokens: list[str], document: str,
315 k1: float = 1.2, b: float = 0.75,
316 ) -> float:
317 """BM25 score for a document given query tokens."""
318 doc_tokens = self._tokenize(document)
319 doc_len = len(doc_tokens)
320 avg_doc_len = max(1, self._total_docs)
322 term_freqs = Counter(doc_tokens)
323 score = 0.0
325 for token in query_tokens:
326 tf = term_freqs.get(token, 0)
327 if tf == 0:
328 continue
329 idf = self._idf(token)
330 numerator = tf * (k1 + 1)
331 denominator = tf + k1 * (1 - b + b * doc_len / avg_doc_len)
332 score += idf * numerator / denominator
334 return round(score, 6)
336 # --- Utilities ---
338 @staticmethod
339 def _cosine_sim(a: np.ndarray, b: np.ndarray) -> float:
340 """Cosine similarity between two vectors."""
341 dot = np.dot(a, b)
342 norm_a = np.linalg.norm(a)
343 norm_b = np.linalg.norm(b)
344 if norm_a == 0 or norm_b == 0:
345 return 0.0
346 return float(dot / (norm_a * norm_b))
348 @property
349 def entry_count(self) -> int:
350 return len(self._entries)
352 def clear(self) -> None:
353 self._entries.clear()
354 self._idf_cache.clear()
355 self._doc_freqs.clear()
356 self._total_docs = 0
358 def get_stats(self) -> dict[str, Any]:
359 """Return index statistics."""
360 return {
361 "total_entries": len(self._entries),
362 "total_docs": self._total_docs,
363 "unique_terms": len(self._doc_freqs),
364 "has_embedder": self._embedder is not None,
365 "hybrid_weight": self._hybrid_weight,
366 }