Coverage for agentos/memory/long_term.py: 27%
103 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-09 07:12 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-09 07:12 +0800
1"""
2AgentOS v0.20 长期记忆系统。
3RAG检索 + 知识图谱双重记忆。
4"""
6from __future__ import annotations
8from dataclasses import dataclass, field
9from typing import Any
12@dataclass
13class MemoryEntry:
14 """长期记忆条目。"""
16 id: str
17 content: str
18 embedding: list[float] | None = None
19 metadata: dict[str, Any] = field(default_factory=dict)
20 created_at: float = 0.0
23class LongTermMemory:
24 """
25 长期记忆 — RAG + 知识图谱。
27 功能:
28 - 语义检索(向量相似度)
29 - 关键词检索(倒排索引)
30 - 实体关系图(知识图谱)
31 - 记忆衰减(时间加权)
32 - 自动摘要压缩
33 """
35 def __init__(self, embedding_dim: int = 1536, max_entries: int = 100000):
36 self._entries: dict[str, MemoryEntry] = {}
37 self._keyword_index: dict[str, set[str]] = {}
38 self._entity_graph: dict[str, set[tuple[str, str]]] = {}
39 self._embedding_dim = embedding_dim
40 self._max_entries = max_entries
42 def add(self, entry: MemoryEntry):
43 """添加记忆条目。"""
44 if len(self._entries) >= self._max_entries:
45 self._evict_oldest()
46 self._entries[entry.id] = entry
47 self._index_keywords(entry)
49 def search_by_keyword(self, query: str, top_k: int = 10) -> list[MemoryEntry]:
50 """关键词检索。"""
51 keywords = query.lower().split()
52 scores: dict[str, int] = {}
53 for kw in keywords:
54 for entry_id in self._keyword_index.get(kw, set()):
55 scores[entry_id] = scores.get(entry_id, 0) + 1
56 ranked = sorted(scores.items(), key=lambda x: x[1], reverse=True)[:top_k]
57 return [self._entries[eid] for eid, _ in ranked if eid in self._entries]
59 def search_by_vector(self, query_embedding: list[float], top_k: int = 10) -> list[MemoryEntry]:
60 """向量相似度检索(余弦相似度)。"""
62 def cosine(a, b):
63 dot = sum(x * y for x, y in zip(a, b))
64 norm_a = sum(x * x for x in a) ** 0.5
65 norm_b = sum(x * x for x in b) ** 0.5
66 return dot / (norm_a * norm_b + 1e-8)
68 scored = []
69 for entry in self._entries.values():
70 if entry.embedding:
71 sim = cosine(query_embedding, entry.embedding)
72 scored.append((sim, entry))
73 scored.sort(key=lambda x: x[0], reverse=True)
74 return [entry for _, entry in scored[:top_k]]
76 def add_relation(self, entity_a: str, relation: str, entity_b: str):
77 """添加知识图谱三元组。"""
78 self._entity_graph.setdefault(entity_a, set()).add((relation, entity_b))
79 self._entity_graph.setdefault(entity_b, set()).add((relation + "_reverse", entity_a))
81 def query_relations(self, entity: str, depth: int = 1) -> list[tuple[str, str]]:
82 """查询实体的关系。"""
83 results = list(self._entity_graph.get(entity, set()))
84 return results[:50]
86 def _index_keywords(self, entry: MemoryEntry):
87 for word in entry.content.lower().split():
88 clean = "".join(c for c in word if c.isalnum())
89 if clean and len(clean) > 1:
90 self._keyword_index.setdefault(clean, set()).add(entry.id)
92 def _evict_oldest(self):
93 oldest = min(self._entries.values(), key=lambda e: e.created_at)
94 del self._entries[oldest.id]
95 for kw_set in self._keyword_index.values():
96 kw_set.discard(oldest.id)
98 # ── Persistence (v1.14.9) ────────────────
100 def get_state(self) -> dict[str, Any]:
101 """Export LongTermMemory state for persistence."""
102 return {
103 "embedding_dim": self._embedding_dim,
104 "max_entries": self._max_entries,
105 "entries": {
106 eid: {
107 "id": entry.id,
108 "content": entry.content,
109 "embedding": entry.embedding,
110 "metadata": entry.metadata,
111 "created_at": entry.created_at,
112 }
113 for eid, entry in self._entries.items()
114 },
115 "entity_graph": {
116 entity: [(r, e) for r, e in relations]
117 for entity, relations in self._entity_graph.items()
118 },
119 }
121 def restore_state(self, state: dict[str, Any]) -> None:
122 """Restore LongTermMemory from a persisted snapshot."""
123 self._embedding_dim = state.get("embedding_dim", self._embedding_dim)
124 self._max_entries = state.get("max_entries", self._max_entries)
125 self._entries.clear()
126 self._keyword_index.clear()
127 self._entity_graph.clear()
129 for eid, entry_data in state.get("entries", {}).items():
130 entry = MemoryEntry(
131 id=entry_data.get("id", eid),
132 content=entry_data.get("content", ""),
133 embedding=entry_data.get("embedding"),
134 metadata=entry_data.get("metadata", {}),
135 created_at=entry_data.get("created_at", 0.0),
136 )
137 self._entries[eid] = entry
138 self._index_keywords(entry)
140 for entity, relations in state.get("entity_graph", {}).items():
141 for rel, target in relations:
142 self._entity_graph.setdefault(entity, set()).add((rel, target))
145class MemoryStore:
146 """三层记忆系统的统一入口。"""
148 def __init__(self, long_term: LongTermMemory | None = None):
149 self.working: dict[str, Any] = {}
150 self.short_term: list[dict] = []
151 self.long_term = long_term or LongTermMemory()
153 def remember(self, key: str, value: Any, long_term: bool = False):
154 """存储记忆。"""
155 if long_term:
156 entry = MemoryEntry(id=key, content=str(value), created_at=__import__("time").time())
157 self.long_term.add(entry)
158 else:
159 self.working[key] = value
160 self.short_term.append({"key": key, "value": value})
162 def recall(self, query: str, use_long_term: bool = True) -> list[Any]:
163 """检索记忆。"""
164 results = []
165 # 工作记忆优先
166 if query in self.working:
167 results.append(self.working[query])
168 # 短期记忆
169 for item in self.short_term:
170 if query.lower() in item["key"].lower():
171 results.append(item["value"])
172 # 长期记忆
173 if use_long_term and not results:
174 long_results = self.long_term.search_by_keyword(query)
175 results.extend([e.content for e in long_results])
176 return results if results else None
178 def clear_short_term(self):
179 self.short_term.clear()
181 # ── Persistence (v1.14.9) ────────────────
183 def get_state(self) -> dict[str, Any]:
184 """Export MemoryStore state for persistence."""
185 return {
186 "working": self.working,
187 "short_term": self.short_term,
188 }
190 def restore_state(self, state: dict[str, Any]) -> None:
191 """Restore MemoryStore from a persisted snapshot."""
192 self.working = dict(state.get("working", {}))
193 self.short_term = list(state.get("short_term", []))