Coverage for agentos/memory/consolidation.py: 37%

266 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-06 12:29 +0800

1""" 

2AgentOS v1.14.1 — 长期记忆巩固系统 (Memory Consolidation)。 

3 

4受 Letta/MemGPT 三层记忆体系启发,在虚拟内存分页器之上增加主动记忆巩固层。 

5核心机制: 

6- Reflection: 定期分析对话历史,提取关键事实、模式、用户偏好 

7- Consolidation: 将 Reflection 结果写入长期向量存储 

8- Retrieval: 智能检索历史记忆,注入后续对话上下文 

9 

10与 memory/pager.py 的关系: 

11- pager.py: 被动分页(上下文窗口溢出时 page_out / keyword search page_in) 

12- consolidation.py: 主动巩固(定期分析 → 提取 → 向量化存储) 

13""" 

14 

15from __future__ import annotations 

16 

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) 

27 

28# ── Memory Data Models ────────────────────── 

29 

30 

31class MemoryType(StrEnum): 

32 """记忆类型。""" 

33 

34 FACT = "fact" # 事实性信息 

35 PREFERENCE = "preference" # 用户偏好 

36 PATTERN = "pattern" # 行为模式 

37 DECISION = "decision" # 决策记录 

38 LESSON = "lesson" # 经验教训 

39 CONTEXT = "context" # 上下文摘要 

40 

41 

42class MemoryImportance(StrEnum): 

43 """记忆重要性。""" 

44 

45 LOW = "low" 

46 MEDIUM = "medium" 

47 HIGH = "high" 

48 CRITICAL = "critical" 

49 

50 

51@dataclass 

52class MemoryFragment: 

53 """记忆片段 — 从对话中提取的原子事实。""" 

54 

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) 

67 

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 } 

82 

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 ) 

98 

99 def touch(self) -> None: 

100 """更新访问时间。""" 

101 self.last_accessed = time.time() 

102 self.access_count += 1 

103 

104 

105@dataclass 

106class ReflectionResult: 

107 """一次 Reflection 的输出。""" 

108 

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) 

115 

116 @property 

117 def total_fragments(self) -> int: 

118 return len(self.fragments) 

119 

120 @property 

121 def has_insights(self) -> bool: 

122 return bool(self.fragments or self.summary or self.user_profile_update) 

123 

124 

125# ── Vector Store Interface ────────────────── 

126 

127 

128class VectorStoreBackend: 

129 """向量存储后端抽象。 

130 

131 支持多种后端: 内存、FAISS、Chroma、Pinecone 等。 

132 """ 

133 

134 async def add( 

135 self, 

136 fragments: list[MemoryFragment], 

137 embeddings: list[list[float]], 

138 ) -> list[str]: 

139 """批量添加记忆片段(含嵌入向量)。返回 memory_ids。""" 

140 raise NotImplementedError 

141 

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 

151 

152 async def delete(self, memory_ids: list[str]) -> int: 

153 """删除指定记忆。返回删除数量。""" 

154 raise NotImplementedError 

155 

156 async def count(self) -> int: 

157 """记忆总数。""" 

158 raise NotImplementedError 

159 

160 

161class InMemoryVectorStore(VectorStoreBackend): 

162 """内存向量存储(开发/测试用)。""" 

163 

164 def __init__(self): 

165 self._fragments: dict[str, MemoryFragment] = {} 

166 self._embeddings: dict[str, list[float]] = {} 

167 

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 

179 

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] 

195 

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)) 

207 

208 results.sort(key=lambda x: x[1], reverse=True) 

209 return results[:top_k] 

210 

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 

219 

220 async def count(self) -> int: 

221 return len(self._fragments) 

222 

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) 

231 

232 

233# ── Embedding Provider ────────────────────── 

234 

235 

236class EmbeddingProvider: 

237 """嵌入向量生成器抽象。""" 

238 

239 async def embed(self, texts: list[str]) -> list[list[float]]: 

240 """批量生成嵌入向量。""" 

241 raise NotImplementedError 

242 

243 async def embed_single(self, text: str) -> list[float]: 

244 """单条文本嵌入。""" 

245 results = await self.embed([text]) 

246 return results[0] 

247 

248 

249class SimpleHashEmbedding(EmbeddingProvider): 

250 """简单位次嵌入(开发/测试用,非语义向量)。 

251 

252 用字符 n-gram 哈希作为伪嵌入,提供基本的相似度。 

253 生产环境应替换为 OpenAI/Cohere 等真实嵌入模型。 

254 """ 

255 

256 def __init__(self, dim: int = 128): 

257 self.dim = dim 

258 

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 

274 

275 

276# ── Reflection Engine ─────────────────────── 

277 

278 

279class ReflectionConfig: 

280 """Reflection 触发配置。""" 

281 

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 

293 

294 

295class ReflectionEngine: 

296 """记忆反思引擎。 

297 

298 定期分析对话历史,提取: 

299 - Facts: 用户提到的具体信息 

300 - Preferences: 用户偏好与习惯 

301 - Patterns: 反复出现的行为模式 

302 - Lessons: 从错误中学到的经验 

303 

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 """ 

311 

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 

333 

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 

343 

344 async def reflect( 

345 self, 

346 messages: list[dict], 

347 existing_fragments: list[MemoryFragment] | None = None, 

348 ) -> ReflectionResult: 

349 """执行一次 Reflection。 

350 

351 Args: 

352 messages: 对话历史(dict 列表,含 role/content) 

353 existing_fragments: 已有的记忆片段(用于矛盾检测) 

354 

355 Returns: 

356 ReflectionResult 含新提取的记忆片段 

357 """ 

358 self._last_reflection_time = time.time() 

359 self._reflection_count += 1 

360 

361 # 1. 构建 Reflection prompt 

362 prompt = self._build_reflection_prompt(messages, existing_fragments) 

363 

364 # 2. 调用 LLM 提取记忆 

365 reflection_text = await self._llm_reflect(messages, prompt) 

366 

367 # 3. 解析 LLM 输出 

368 result = self._parse_reflection_output(reflection_text, len(messages)) 

369 

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 

376 

377 # 5. 存入向量库 

378 if result.fragments: 

379 await self._vector_store.add(result.fragments, embeddings) 

380 

381 # 6. 淘汰旧记忆 

382 if result.deprecated_ids: 

383 await self._vector_store.delete(result.deprecated_ids) 

384 

385 # Reset counter 

386 self._message_count_since_reflection = 0 

387 

388 return result 

389 

390 def record_message(self) -> None: 

391 """记录一条新消息(用于计数触发)。""" 

392 self._message_count_since_reflection += 1 

393 

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 """检索与查询相关的记忆。 

401 

402 Args: 

403 query: 查询文本 

404 top_k: 返回数量 

405 filter_types: 按类型过滤 

406 

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 

421 

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) 

434 

435 return f"""You are a memory consolidation system. Analyze the conversation and extract: 

436 

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 

441 

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 

446 

447{existing_str} 

448 

449Output JSON array only: 

450[{{"type": "...", "content": "...", "importance": "...", "confidence": 0.9, "tags": ["..."]}}] 

451 

452If nothing significant to extract, output empty array: []""" 

453 

454 def _parse_reflection_output( 

455 self, 

456 text: str, 

457 source_msg_count: int, 

458 ) -> ReflectionResult: 

459 """解析 LLM 输出的 JSON。""" 

460 result = ReflectionResult() 

461 

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 

487 

488 return result 

489 

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 } 

497 

498 # ── Persistence (v1.14.9) ──────────────── 

499 

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 } 

520 

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) 

526 

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 

534 

535 

536# ── Memory Context Injector ───────────────── 

537 

538 

539class MemoryContextInjector: 

540 """记忆上下文注入器。 

541 

542 在每次 Agent 对话开始时,自动检索相关历史记忆, 

543 注入到 system prompt 或上下文中。 

544 

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 """ 

550 

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 

560 

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 ) 

572 

573 if not fragments: 

574 return "" 

575 

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 ) 

582 

583 context = "\n".join(lines) 

584 if len(context) > self.max_context_length: 

585 context = context[: self.max_context_length] + "..." 

586 

587 return context 

588 

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 "" 

606 

607 lines = ["[Key Context]"] 

608 for frag in important[:3]: 

609 lines.append(f"- {frag.content}") 

610 

611 return "\n".join(lines) 

612 

613 

614# ── Memory Consolidation Pipeline ─────────── 

615 

616 

617class MemoryConsolidationPipeline: 

618 """记忆巩固流水线(一键集成)。 

619 

620 组合 ReflectionEngine + MemoryContextInjector, 

621 提供开箱即用的记忆系统。 

622 

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 """ 

631 

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) 

648 

649 def record_message(self) -> None: 

650 self._reflection_engine.record_message() 

651 

652 def should_reflect(self) -> bool: 

653 return self._reflection_engine.should_reflect( 

654 self._reflection_engine._message_count_since_reflection 

655 ) 

656 

657 async def reflect(self, messages: list[dict]) -> ReflectionResult: 

658 return await self._reflection_engine.reflect(messages) 

659 

660 async def get_context(self, query: str) -> str: 

661 return await self._injector.build_context(query) 

662 

663 async def get_condensed_context(self, query: str) -> str: 

664 return await self._injector.build_condensed_context(query) 

665 

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 } 

676 

677 # ── Persistence (v1.14.9) ──────────────── 

678 

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() 

682 

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)