Coverage for agentos/memory/retriever.py: 29%
179 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 12:20 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 12:20 +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 collections.abc import Callable
14from dataclasses import dataclass, field
15from enum import Enum
16from typing import Any
18import numpy as np
21class 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: list[float] | None = None
38 timestamp: float | None = 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: Callable[[str], list[float]] | None = 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: int | None = None,
138 strategy: RetrievalStrategy = RetrievalStrategy.HYBRID,
139 filter_source: str | None = 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
157 start = time.perf_counter()
158 top_k = top_k or self._default_top_k
160 # Filter entries
161 candidates = [
162 e
163 for e in self._entries.values()
164 if (filter_source is None or e.source == filter_source)
165 and e.importance >= min_importance
166 ]
168 if not candidates:
169 return []
171 if strategy == RetrievalStrategy.RECENT:
172 results = self._retrieve_recent(candidates, top_k)
173 elif strategy == RetrievalStrategy.KEYWORD:
174 results = self._retrieve_keyword(query, candidates, top_k)
175 elif strategy == RetrievalStrategy.SEMANTIC:
176 results = self._retrieve_semantic(query, candidates, top_k)
177 else: # HYBRID
178 results = self._retrieve_hybrid(query, candidates, top_k)
180 (time.perf_counter() - start) * 1000
181 # Attach stats to results via a common approach
182 return results
184 def _retrieve_recent(
185 self,
186 candidates: list[MemoryEntry],
187 top_k: int,
188 ) -> list[RetrievalResult]:
189 """Return most recent entries sorted by timestamp."""
190 sorted_entries = sorted(
191 candidates,
192 key=lambda e: e.timestamp or 0,
193 reverse=True,
194 )
195 return [
196 RetrievalResult(
197 entry=e,
198 score=1.0,
199 strategy=RetrievalStrategy.RECENT,
200 )
201 for e in sorted_entries[:top_k]
202 ]
204 def _retrieve_keyword(
205 self,
206 query: str,
207 candidates: list[MemoryEntry],
208 top_k: int,
209 ) -> list[RetrievalResult]:
210 """BM25-style keyword search."""
211 query_tokens = self._tokenize(query)
212 if not query_tokens:
213 return []
215 scores = []
216 for entry in candidates:
217 score = self._bm25_score(query_tokens, entry.content)
218 if score >= self._min_keyword_score:
219 scores.append((score, entry))
221 scores.sort(key=lambda x: x[0], reverse=True)
222 return [
223 RetrievalResult(
224 entry=e,
225 score=s,
226 strategy=RetrievalStrategy.KEYWORD,
227 )
228 for s, e in scores[:top_k]
229 ]
231 def _retrieve_semantic(
232 self,
233 query: str,
234 candidates: list[MemoryEntry],
235 top_k: int,
236 ) -> list[RetrievalResult]:
237 """Cosine similarity semantic search."""
238 if not self._embedder:
239 return self._retrieve_keyword(query, candidates, top_k)
241 query_embedding = np.array(self._embedder(query))
242 scores = []
243 for entry in candidates:
244 if entry.embedding is None:
245 entry.embedding = self._embedder(entry.content)
246 entry_embedding = np.array(entry.embedding)
247 similarity = self._cosine_sim(query_embedding, entry_embedding)
248 scores.append((similarity, entry))
250 scores.sort(key=lambda x: x[0], reverse=True)
251 return [
252 RetrievalResult(
253 entry=e,
254 score=float(s),
255 strategy=RetrievalStrategy.SEMANTIC,
256 )
257 for s, e in scores[:top_k]
258 ]
260 def _retrieve_hybrid(
261 self,
262 query: str,
263 candidates: list[MemoryEntry],
264 top_k: int,
265 ) -> list[RetrievalResult]:
266 """Weighted combination of semantic and keyword scores."""
267 query_tokens = self._tokenize(query)
268 has_embedder = self._embedder is not None
270 if has_embedder:
271 query_embedding = np.array(self._embedder(query))
273 scores = []
274 for entry in candidates:
275 kw_score = self._bm25_score(query_tokens, entry.content)
277 if has_embedder:
278 if entry.embedding is None:
279 entry.embedding = self._embedder(entry.content)
280 entry_embedding = np.array(entry.embedding)
281 sem_score = self._cosine_sim(query_embedding, entry_embedding)
282 combined = self._hybrid_weight * sem_score + (1 - self._hybrid_weight) * kw_score
283 else:
284 combined = kw_score
286 if combined > 0:
287 scores.append((combined, entry))
289 scores.sort(key=lambda x: x[0], reverse=True)
290 return [
291 RetrievalResult(
292 entry=e,
293 score=float(s),
294 strategy=RetrievalStrategy.HYBRID,
295 )
296 for s, e in scores[:top_k]
297 ]
299 # --- BM25 implementation ---
301 @staticmethod
302 def _tokenize(text: str) -> list[str]:
303 """Simple word tokenizer."""
304 text = text.lower()
305 # Split on non-alphanumeric, keep sequences of 2+ chars
306 tokens = []
307 current = []
308 for ch in text:
309 if ch.isalnum():
310 current.append(ch)
311 else:
312 if len(current) >= 2:
313 tokens.append("".join(current))
314 current = []
315 if len(current) >= 2:
316 tokens.append("".join(current))
317 return tokens
319 def _idf(self, term: str) -> float:
320 """Inverse document frequency."""
321 if term not in self._idf_cache:
322 df = self._doc_freqs.get(term, 0)
323 if df == 0 or self._total_docs == 0:
324 self._idf_cache[term] = 0.0
325 else:
326 self._idf_cache[term] = math.log((self._total_docs - df + 0.5) / (df + 0.5) + 1.0)
327 return self._idf_cache[term]
329 def _bm25_score(
330 self,
331 query_tokens: list[str],
332 document: str,
333 k1: float = 1.2,
334 b: float = 0.75,
335 ) -> float:
336 """BM25 score for a document given query tokens."""
337 doc_tokens = self._tokenize(document)
338 doc_len = len(doc_tokens)
339 avg_doc_len = max(1, self._total_docs)
341 term_freqs = Counter(doc_tokens)
342 score = 0.0
344 for token in query_tokens:
345 tf = term_freqs.get(token, 0)
346 if tf == 0:
347 continue
348 idf = self._idf(token)
349 numerator = tf * (k1 + 1)
350 denominator = tf + k1 * (1 - b + b * doc_len / avg_doc_len)
351 score += idf * numerator / denominator
353 return round(score, 6)
355 # --- Utilities ---
357 @staticmethod
358 def _cosine_sim(a: np.ndarray, b: np.ndarray) -> float:
359 """Cosine similarity between two vectors."""
360 dot = np.dot(a, b)
361 norm_a = np.linalg.norm(a)
362 norm_b = np.linalg.norm(b)
363 if norm_a == 0 or norm_b == 0:
364 return 0.0
365 return float(dot / (norm_a * norm_b))
367 @property
368 def entry_count(self) -> int:
369 return len(self._entries)
371 def clear(self) -> None:
372 self._entries.clear()
373 self._idf_cache.clear()
374 self._doc_freqs.clear()
375 self._total_docs = 0
377 def get_stats(self) -> dict[str, Any]:
378 """Return index statistics."""
379 return {
380 "total_entries": len(self._entries),
381 "total_docs": self._total_docs,
382 "unique_terms": len(self._doc_freqs),
383 "has_embedder": self._embedder is not None,
384 "hybrid_weight": self._hybrid_weight,
385 }