Coverage for agentos/memory/consolidation.py: 37%
266 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 21:26 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 21:26 +0800
1"""
2AgentOS v1.14.1 — 长期记忆巩固系统 (Memory Consolidation)。
4受 Letta/MemGPT 三层记忆体系启发,在虚拟内存分页器之上增加主动记忆巩固层。
5核心机制:
6- Reflection: 定期分析对话历史,提取关键事实、模式、用户偏好
7- Consolidation: 将 Reflection 结果写入长期向量存储
8- Retrieval: 智能检索历史记忆,注入后续对话上下文
10与 memory/pager.py 的关系:
11- pager.py: 被动分页(上下文窗口溢出时 page_out / keyword search page_in)
12- consolidation.py: 主动巩固(定期分析 → 提取 → 向量化存储)
13"""
15from __future__ import annotations
17import asyncio
18import json
19import time
20import uuid
21from collections.abc import Callable
22from dataclasses import dataclass, field
23from enum import StrEnum
24from typing import (
25 Any,
26)
28# ── Memory Data Models ──────────────────────
31class MemoryType(StrEnum):
32 """记忆类型。"""
34 FACT = "fact" # 事实性信息
35 PREFERENCE = "preference" # 用户偏好
36 PATTERN = "pattern" # 行为模式
37 DECISION = "decision" # 决策记录
38 LESSON = "lesson" # 经验教训
39 CONTEXT = "context" # 上下文摘要
42class MemoryImportance(StrEnum):
43 """记忆重要性。"""
45 LOW = "low"
46 MEDIUM = "medium"
47 HIGH = "high"
48 CRITICAL = "critical"
51@dataclass
52class MemoryFragment:
53 """记忆片段 — 从对话中提取的原子事实。"""
55 memory_id: str = field(default_factory=lambda: f"mem-{uuid.uuid4().hex[:12]}")
56 memory_type: MemoryType = MemoryType.FACT
57 content: str = ""
58 importance: MemoryImportance = MemoryImportance.MEDIUM
59 source_messages: list[int] = field(default_factory=list) # 源自哪些消息
60 tags: list[str] = field(default_factory=list)
61 confidence: float = 1.0 # 0.0~1.0 置信度
62 created_at: float = field(default_factory=time.time)
63 last_accessed: float = field(default_factory=time.time)
64 access_count: int = 0
65 embedding: list[float] | None = None # 向量嵌入(惰性计算)
66 metadata: dict[str, Any] = field(default_factory=dict)
68 def to_dict(self) -> dict:
69 return {
70 "memory_id": self.memory_id,
71 "memory_type": self.memory_type.value,
72 "content": self.content,
73 "importance": self.importance.value,
74 "source_messages": self.source_messages,
75 "tags": self.tags,
76 "confidence": self.confidence,
77 "created_at": self.created_at,
78 "last_accessed": self.last_accessed,
79 "access_count": self.access_count,
80 "metadata": self.metadata,
81 }
83 @classmethod
84 def from_dict(cls, d: dict) -> MemoryFragment:
85 return cls(
86 memory_id=d.get("memory_id", ""),
87 memory_type=MemoryType(d.get("memory_type", "fact")),
88 content=d.get("content", ""),
89 importance=MemoryImportance(d.get("importance", "medium")),
90 source_messages=d.get("source_messages", []),
91 tags=d.get("tags", []),
92 confidence=d.get("confidence", 1.0),
93 created_at=d.get("created_at", time.time()),
94 last_accessed=d.get("last_accessed", time.time()),
95 access_count=d.get("access_count", 0),
96 metadata=d.get("metadata", {}),
97 )
99 def touch(self) -> None:
100 """更新访问时间。"""
101 self.last_accessed = time.time()
102 self.access_count += 1
105@dataclass
106class ReflectionResult:
107 """一次 Reflection 的输出。"""
109 fragments: list[MemoryFragment] = field(default_factory=list)
110 summary: str = "" # 会话级摘要
111 contradictions: list[tuple[str, str]] = field(default_factory=list) # 新旧矛盾
112 deprecated_ids: list[str] = field(default_factory=list) # 需淘汰的旧记忆
113 user_profile_update: dict[str, Any] = field(default_factory=dict)
114 timestamp: float = field(default_factory=time.time)
116 @property
117 def total_fragments(self) -> int:
118 return len(self.fragments)
120 @property
121 def has_insights(self) -> bool:
122 return bool(self.fragments or self.summary or self.user_profile_update)
125# ── Vector Store Interface ──────────────────
128class VectorStoreBackend:
129 """向量存储后端抽象。
131 支持多种后端: 内存、FAISS、Chroma、Pinecone 等。
132 """
134 async def add(
135 self,
136 fragments: list[MemoryFragment],
137 embeddings: list[list[float]],
138 ) -> list[str]:
139 """批量添加记忆片段(含嵌入向量)。返回 memory_ids。"""
140 raise NotImplementedError
142 async def search(
143 self,
144 query_embedding: list[float],
145 top_k: int = 10,
146 filter_types: list[MemoryType] | None = None,
147 min_importance: MemoryImportance = MemoryImportance.LOW,
148 ) -> list[tuple[MemoryFragment, float]]:
149 """向量相似度搜索。返回 (fragment, score)。"""
150 raise NotImplementedError
152 async def delete(self, memory_ids: list[str]) -> int:
153 """删除指定记忆。返回删除数量。"""
154 raise NotImplementedError
156 async def count(self) -> int:
157 """记忆总数。"""
158 raise NotImplementedError
161class InMemoryVectorStore(VectorStoreBackend):
162 """内存向量存储(开发/测试用)。"""
164 def __init__(self):
165 self._fragments: dict[str, MemoryFragment] = {}
166 self._embeddings: dict[str, list[float]] = {}
168 async def add(
169 self,
170 fragments: list[MemoryFragment],
171 embeddings: list[list[float]],
172 ) -> list[str]:
173 ids = []
174 for frag, emb in zip(fragments, embeddings):
175 self._fragments[frag.memory_id] = frag
176 self._embeddings[frag.memory_id] = emb
177 ids.append(frag.memory_id)
178 return ids
180 async def search(
181 self,
182 query_embedding: list[float],
183 top_k: int = 10,
184 filter_types: list[MemoryType] | None = None,
185 min_importance: MemoryImportance = MemoryImportance.LOW,
186 ) -> list[tuple[MemoryFragment, float]]:
187 results = []
188 importance_rank = {
189 MemoryImportance.LOW: 0,
190 MemoryImportance.MEDIUM: 1,
191 MemoryImportance.HIGH: 2,
192 MemoryImportance.CRITICAL: 3,
193 }
194 min_rank = importance_rank[min_importance]
196 for mid, emb in self._embeddings.items():
197 frag = self._fragments[mid]
198 # Type filter
199 if filter_types and frag.memory_type not in filter_types:
200 continue
201 # Importance filter
202 if importance_rank[frag.importance] < min_rank:
203 continue
204 # Cosine similarity
205 score = self._cosine_similarity(query_embedding, emb)
206 results.append((frag, score))
208 results.sort(key=lambda x: x[1], reverse=True)
209 return results[:top_k]
211 async def delete(self, memory_ids: list[str]) -> int:
212 count = 0
213 for mid in memory_ids:
214 if mid in self._fragments:
215 del self._fragments[mid]
216 self._embeddings.pop(mid, None)
217 count += 1
218 return count
220 async def count(self) -> int:
221 return len(self._fragments)
223 @staticmethod
224 def _cosine_similarity(a: list[float], b: list[float]) -> float:
225 dot = sum(x * y for x, y in zip(a, b))
226 norm_a = sum(x * x for x in a) ** 0.5
227 norm_b = sum(x * x for x in b) ** 0.5
228 if norm_a == 0 or norm_b == 0:
229 return 0.0
230 return dot / (norm_a * norm_b)
233# ── Embedding Provider ──────────────────────
236class EmbeddingProvider:
237 """嵌入向量生成器抽象。"""
239 async def embed(self, texts: list[str]) -> list[list[float]]:
240 """批量生成嵌入向量。"""
241 raise NotImplementedError
243 async def embed_single(self, text: str) -> list[float]:
244 """单条文本嵌入。"""
245 results = await self.embed([text])
246 return results[0]
249class SimpleHashEmbedding(EmbeddingProvider):
250 """简单位次嵌入(开发/测试用,非语义向量)。
252 用字符 n-gram 哈希作为伪嵌入,提供基本的相似度。
253 生产环境应替换为 OpenAI/Cohere 等真实嵌入模型。
254 """
256 def __init__(self, dim: int = 128):
257 self.dim = dim
259 async def embed(self, texts: list[str]) -> list[list[float]]:
260 results = []
261 for text in texts:
262 vec = [0.0] * self.dim
263 # Character 3-gram hashing
264 for i in range(len(text) - 2):
265 gram = text[i : i + 3]
266 h = hash(gram) % self.dim
267 vec[h] += 1.0
268 # L2 normalize
269 norm = sum(v * v for v in vec) ** 0.5
270 if norm > 0:
271 vec = [v / norm for v in vec]
272 results.append(vec)
273 return results
276# ── Reflection Engine ───────────────────────
279class ReflectionConfig:
280 """Reflection 触发配置。"""
282 def __init__(
283 self,
284 min_messages_since_last: int = 10,
285 min_seconds_since_last: float = 300.0, # 5 分钟
286 max_conversation_turns: int = 50,
287 auto_reflect: bool = True,
288 ):
289 self.min_messages_since_last = min_messages_since_last
290 self.min_seconds_since_last = min_seconds_since_last
291 self.max_conversation_turns = max_conversation_turns
292 self.auto_reflect = auto_reflect
295class ReflectionEngine:
296 """记忆反思引擎。
298 定期分析对话历史,提取:
299 - Facts: 用户提到的具体信息
300 - Preferences: 用户偏好与习惯
301 - Patterns: 反复出现的行为模式
302 - Lessons: 从错误中学到的经验
304 Usage:
305 engine = ReflectionEngine(llm_reflect_fn, vector_store, embedding_provider)
306 # 在 agent loop 中定期调用
307 should_reflect = engine.should_reflect(message_count)
308 if should_reflect:
309 result = await engine.reflect(messages_history)
310 """
312 def __init__(
313 self,
314 llm_reflect_fn: Callable[[list[dict], str], Any] | None = None,
315 vector_store: Any | None = None,
316 embedding_provider: Any | None = None,
317 config: ReflectionConfig | None = None,
318 ):
319 """
320 Args:
321 llm_reflect_fn: LLM 调用函数,签名 (messages, prompt) -> reflection_text
322 vector_store: 向量存储后端
323 embedding_provider: 嵌入向量生成器
324 config: 触发配置
325 """
326 self._llm_reflect = llm_reflect_fn
327 self._vector_store = vector_store
328 self._embedding_provider = embedding_provider
329 self.config = config or ReflectionConfig()
330 self._last_reflection_time: float = 0.0
331 self._message_count_since_reflection: int = 0
332 self._reflection_count: int = 0
334 def should_reflect(self, current_message_count: int) -> bool:
335 """判断是否应该触发 Reflection。"""
336 if not self.config.auto_reflect:
337 return False
338 if self._reflection_count == 0 and current_message_count >= 5:
339 return True # 首次在 5 条消息后触发
340 msg_check = self._message_count_since_reflection >= self.config.min_messages_since_last
341 time_check = time.time() - self._last_reflection_time >= self.config.min_seconds_since_last
342 return msg_check or time_check
344 async def reflect(
345 self,
346 messages: list[dict],
347 existing_fragments: list[MemoryFragment] | None = None,
348 ) -> ReflectionResult:
349 """执行一次 Reflection。
351 Args:
352 messages: 对话历史(dict 列表,含 role/content)
353 existing_fragments: 已有的记忆片段(用于矛盾检测)
355 Returns:
356 ReflectionResult 含新提取的记忆片段
357 """
358 self._last_reflection_time = time.time()
359 self._reflection_count += 1
361 # 1. 构建 Reflection prompt
362 prompt = self._build_reflection_prompt(messages, existing_fragments)
364 # 2. 调用 LLM 提取记忆
365 reflection_text = await self._llm_reflect(messages, prompt)
367 # 3. 解析 LLM 输出
368 result = self._parse_reflection_output(reflection_text, len(messages))
370 # 4. 生成嵌入向量
371 if result.fragments:
372 texts = [f.content for f in result.fragments]
373 embeddings = await self._embedding_provider.embed(texts)
374 for frag, emb in zip(result.fragments, embeddings):
375 frag.embedding = emb
377 # 5. 存入向量库
378 if result.fragments:
379 await self._vector_store.add(result.fragments, embeddings)
381 # 6. 淘汰旧记忆
382 if result.deprecated_ids:
383 await self._vector_store.delete(result.deprecated_ids)
385 # Reset counter
386 self._message_count_since_reflection = 0
388 return result
390 def record_message(self) -> None:
391 """记录一条新消息(用于计数触发)。"""
392 self._message_count_since_reflection += 1
394 async def retrieve_relevant(
395 self,
396 query: str,
397 top_k: int = 5,
398 filter_types: list[MemoryType] | None = None,
399 ) -> list[MemoryFragment]:
400 """检索与查询相关的记忆。
402 Args:
403 query: 查询文本
404 top_k: 返回数量
405 filter_types: 按类型过滤
407 Returns:
408 相关记忆片段列表
409 """
410 query_embedding = await self._embedding_provider.embed_single(query)
411 results = await self._vector_store.search(
412 query_embedding,
413 top_k=top_k,
414 filter_types=filter_types,
415 )
416 fragments = []
417 for frag, score in results:
418 frag.touch()
419 fragments.append(frag)
420 return fragments
422 def _build_reflection_prompt(
423 self,
424 messages: list[dict],
425 existing_fragments: list[MemoryFragment] | None = None,
426 ) -> str:
427 """构建 Reflection prompt。"""
428 existing_str = ""
429 if existing_fragments:
430 existing_items = [
431 f"- [{f.memory_type.value}] {f.content}" for f in existing_fragments[:20]
432 ]
433 existing_str = "\n\nExisting memories:\n" + "\n".join(existing_items)
435 return f"""You are a memory consolidation system. Analyze the conversation and extract:
4371. FACTS: Specific information mentioned (names, dates, numbers, tools used, decisions made)
4382. PREFERENCES: User preferences, likes, dislikes, habits
4393. PATTERNS: Repeated behaviors, common workflows, recurring topics
4404. LESSONS: What went wrong, what worked, what to avoid next time
442For each extracted item, assign:
443- type: "fact" | "preference" | "pattern" | "decision" | "lesson"
444- importance: "low" | "medium" | "high" | "critical"
445- confidence: 0.0 to 1.0
447{existing_str}
449Output JSON array only:
450[{{"type": "...", "content": "...", "importance": "...", "confidence": 0.9, "tags": ["..."]}}]
452If nothing significant to extract, output empty array: []"""
454 def _parse_reflection_output(
455 self,
456 text: str,
457 source_msg_count: int,
458 ) -> ReflectionResult:
459 """解析 LLM 输出的 JSON。"""
460 result = ReflectionResult()
462 try:
463 # Extract JSON array
464 start = text.find("[")
465 end = text.rfind("]")
466 if start >= 0 and end > start:
467 json_str = text[start : end + 1]
468 items = json.loads(json_str)
469 for item in items:
470 frag = MemoryFragment(
471 memory_type=MemoryType(item.get("type", "fact")),
472 content=item.get("content", ""),
473 importance=MemoryImportance(item.get("importance", "medium")),
474 confidence=float(item.get("confidence", 1.0)),
475 tags=item.get("tags", []),
476 source_messages=list(
477 range(
478 max(0, source_msg_count - 20),
479 source_msg_count,
480 )
481 ),
482 )
483 if frag.content.strip():
484 result.fragments.append(frag)
485 except (json.JSONDecodeError, KeyError, ValueError):
486 pass
488 return result
490 @property
491 def stats(self) -> dict[str, Any]:
492 return {
493 "reflection_count": self._reflection_count,
494 "last_reflection_time": self._last_reflection_time,
495 "messages_since_last": self._message_count_since_reflection,
496 }
498 # ── Persistence (v1.14.9) ────────────────
500 def get_state(self) -> dict[str, Any]:
501 """Export ReflectionEngine state for persistence."""
502 return {
503 "reflection_count": self._reflection_count,
504 "last_reflection_time": self._last_reflection_time,
505 "message_count_since_reflection": self._message_count_since_reflection,
506 "vector_store_fragments": (
507 {mid: frag.to_dict() for mid, frag in self._vector_store._fragments.items()}
508 if hasattr(self._vector_store, "_fragments") and self._vector_store
509 else {}
510 ),
511 "vector_store_embeddings": {
512 mid: list(emb) if emb else []
513 for mid, emb in (
514 self._vector_store._embeddings.items()
515 if hasattr(self._vector_store, "_embeddings") and self._vector_store
516 else {}.items()
517 )
518 },
519 }
521 def restore_state(self, state: dict[str, Any]) -> None:
522 """Restore ReflectionEngine from a persisted snapshot."""
523 self._reflection_count = state.get("reflection_count", 0)
524 self._last_reflection_time = state.get("last_reflection_time", 0.0)
525 self._message_count_since_reflection = state.get("message_count_since_reflection", 0)
527 if self._vector_store and hasattr(self._vector_store, "_fragments"):
528 self._vector_store._fragments.clear()
529 self._vector_store._embeddings.clear()
530 for mid, frag_data in state.get("vector_store_fragments", {}).items():
531 self._vector_store._fragments[mid] = MemoryFragment.from_dict(frag_data)
532 for mid, emb in state.get("vector_store_embeddings", {}).items():
533 self._vector_store._embeddings[mid] = emb
536# ── Memory Context Injector ─────────────────
539class MemoryContextInjector:
540 """记忆上下文注入器。
542 在每次 Agent 对话开始时,自动检索相关历史记忆,
543 注入到 system prompt 或上下文中。
545 Usage:
546 injector = MemoryContextInjector(reflection_engine)
547 context = await injector.build_context("user query here")
548 messages.insert(0, {"role": "system", "content": context})
549 """
551 def __init__(
552 self,
553 reflection_engine: ReflectionEngine,
554 max_context_length: int = 2000,
555 max_fragments: int = 5,
556 ):
557 self._engine = reflection_engine
558 self.max_context_length = max_context_length
559 self.max_fragments = max_fragments
561 async def build_context(
562 self,
563 query: str,
564 include_types: list[MemoryType] | None = None,
565 ) -> str:
566 """构建上下文注入文本。"""
567 fragments = await self._engine.retrieve_relevant(
568 query,
569 top_k=self.max_fragments,
570 filter_types=include_types,
571 )
573 if not fragments:
574 return ""
576 lines = ["[Relevant Memories]"]
577 for frag in fragments:
578 lines.append(
579 f"- [{frag.memory_type.value}] {frag.content}"
580 f" (confidence: {frag.confidence:.0%})"
581 )
583 context = "\n".join(lines)
584 if len(context) > self.max_context_length:
585 context = context[: self.max_context_length] + "..."
587 return context
589 async def build_condensed_context(
590 self,
591 query: str,
592 ) -> str:
593 """构建紧凑上下文(仅高重要性记忆)。"""
594 fragments = await self._engine.retrieve_relevant(
595 query,
596 top_k=self.max_fragments,
597 )
598 # Filter: only HIGH/CRITICAL
599 important = [
600 f
601 for f in fragments
602 if f.importance in (MemoryImportance.HIGH, MemoryImportance.CRITICAL)
603 ]
604 if not important:
605 return ""
607 lines = ["[Key Context]"]
608 for frag in important[:3]:
609 lines.append(f"- {frag.content}")
611 return "\n".join(lines)
614# ── Memory Consolidation Pipeline ───────────
617class MemoryConsolidationPipeline:
618 """记忆巩固流水线(一键集成)。
620 组合 ReflectionEngine + MemoryContextInjector,
621 提供开箱即用的记忆系统。
623 Usage:
624 pipeline = MemoryConsolidationPipeline(llm_fn)
625 # 在 agent loop 中:
626 pipeline.record_message()
627 if pipeline.should_reflect():
628 await pipeline.reflect(messages)
629 context = await pipeline.get_context(user_query)
630 """
632 def __init__(
633 self,
634 llm_reflect_fn: Callable,
635 vector_store: VectorStoreBackend | None = None,
636 embedding_provider: EmbeddingProvider | None = None,
637 config: ReflectionConfig | None = None,
638 ):
639 self._vector_store = vector_store or InMemoryVectorStore()
640 self._embedding_provider = embedding_provider or SimpleHashEmbedding(128)
641 self._reflection_engine = ReflectionEngine(
642 llm_reflect_fn=llm_reflect_fn,
643 vector_store=self._vector_store,
644 embedding_provider=self._embedding_provider,
645 config=config,
646 )
647 self._injector = MemoryContextInjector(self._reflection_engine)
649 def record_message(self) -> None:
650 self._reflection_engine.record_message()
652 def should_reflect(self) -> bool:
653 return self._reflection_engine.should_reflect(
654 self._reflection_engine._message_count_since_reflection
655 )
657 async def reflect(self, messages: list[dict]) -> ReflectionResult:
658 return await self._reflection_engine.reflect(messages)
660 async def get_context(self, query: str) -> str:
661 return await self._injector.build_context(query)
663 async def get_condensed_context(self, query: str) -> str:
664 return await self._injector.build_condensed_context(query)
666 @property
667 def stats(self) -> dict[str, Any]:
668 return {
669 "reflection": self._reflection_engine.stats,
670 "total_memories": (
671 asyncio.get_event_loop().run_until_complete(self._vector_store.count())
672 if asyncio.get_event_loop().is_running()
673 else 0
674 ),
675 }
677 # ── Persistence (v1.14.9) ────────────────
679 def get_state(self) -> dict[str, Any]:
680 """Export consolidation pipeline state for persistence. Delegates to ReflectionEngine."""
681 return self._reflection_engine.get_state()
683 def restore_state(self, state: dict[str, Any]) -> None:
684 """Restore consolidation pipeline from a persisted snapshot."""
685 self._reflection_engine.restore_state(state)