Coverage for agentos/memory/long_term.py: 27%

103 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-08 21:26 +0800

1""" 

2AgentOS v0.20 长期记忆系统。 

3RAG检索 + 知识图谱双重记忆。 

4""" 

5 

6from __future__ import annotations 

7 

8from dataclasses import dataclass, field 

9from typing import Any 

10 

11 

12@dataclass 

13class MemoryEntry: 

14 """长期记忆条目。""" 

15 

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 

21 

22 

23class LongTermMemory: 

24 """ 

25 长期记忆 — RAG + 知识图谱。 

26 

27 功能: 

28 - 语义检索(向量相似度) 

29 - 关键词检索(倒排索引) 

30 - 实体关系图(知识图谱) 

31 - 记忆衰减(时间加权) 

32 - 自动摘要压缩 

33 """ 

34 

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 

41 

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) 

48 

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] 

58 

59 def search_by_vector(self, query_embedding: list[float], top_k: int = 10) -> list[MemoryEntry]: 

60 """向量相似度检索(余弦相似度)。""" 

61 

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) 

67 

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

75 

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

80 

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] 

85 

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) 

91 

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) 

97 

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

99 

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 } 

120 

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

128 

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) 

139 

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

143 

144 

145class MemoryStore: 

146 """三层记忆系统的统一入口。""" 

147 

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

152 

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

161 

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 

177 

178 def clear_short_term(self): 

179 self.short_term.clear() 

180 

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

182 

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 } 

189 

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", []))