Coverage for merco/memory/strategy.py: 89%

94 statements  

« prev     ^ index     » next       coverage.py v7.15.0, created at 2026-07-07 14:04 +0800

1"""Memory 保存触发策略 — 监听 Hook 事件,构造 SaveItem 喂给 Pipeline""" 

2from __future__ import annotations 

3 

4import hashlib 

5import json 

6import logging 

7import re 

8from abc import ABC, abstractmethod 

9 

10from .save_pipeline import SaveItem 

11 

12logger = logging.getLogger("merco.memory.strategy") 

13 

14 

15class MemorySaveStrategy(ABC): 

16 """监听事件,构造 SaveItem 喂给 Pipeline""" 

17 

18 name: str = "" 

19 

20 def __init__(self, pipeline): 

21 self.pipeline = pipeline 

22 

23 @abstractmethod 

24 async def on_event(self, event: str, **kwargs) -> None: ... 

25 

26 

27class ExplicitRememberStrategy(MemorySaveStrategy): 

28 """/remember <text> 显式存一条记忆""" 

29 name = "explicit_remember" 

30 

31 def subscribe(self, hooks) -> None: 

32 """注册到 HookRegistry""" 

33 hooks.on("command.remember", self._on_remember) 

34 

35 async def on_event(self, event: str, **kwargs) -> None: 

36 """兼容直接调用(测试用)""" 

37 await self._on_remember(**kwargs) 

38 

39 async def _on_remember(self, text: str, key: str = "", **kwargs) -> None: 

40 if not key: 

41 key = self._derive_key(text) 

42 item = SaveItem(key=key, value=text, source="user") 

43 await self.pipeline.save(item) 

44 

45 @staticmethod 

46 def _derive_key(text: str) -> str: 

47 """从文本生成稳定 key: user_<前20字净化>_<hash8>""" 

48 h = hashlib.md5(text.encode()).hexdigest()[:8] 

49 prefix = re.sub(r"\W+", "_", text[:20].strip()).strip("_") 

50 return f"user_{prefix}_{h}" if prefix else f"user_{h}" 

51 

52 

53class SessionEndExtractStrategy(MemorySaveStrategy): 

54 """session.destroy 时用 LLM 抽取 1-3 条 insight 记忆""" 

55 name = "session_end_extract" 

56 

57 EXTRACT_PROMPT = """从以下对话中抽取 1-3 条值得长期记住的关键信息(用户偏好、事实、决策)。 

58仅返回 JSON 数组,每条形如 {{"key": "snake_case_key", "value": "原文", "tags": ["tag1"]}}。 

59没有值得记的就返回 []。 

60 

61对话: 

62{messages} 

63""" 

64 

65 def __init__(self, pipeline, llm, *, 

66 session_store=None, max_per_session: int = 3, 

67 min_messages: int = 5): 

68 super().__init__(pipeline) 

69 self.llm = llm 

70 self._session_store = session_store 

71 self.max = max_per_session 

72 self.min_msgs = min_messages 

73 

74 def subscribe(self, hooks) -> None: 

75 hooks.on("session.destroy", self._on_destroy) 

76 

77 async def on_event(self, event: str, **kwargs) -> None: 

78 """兼容直接调用(测试用)""" 

79 await self._on_destroy(**kwargs) 

80 

81 async def _on_destroy(self, session_id: str = "", **kwargs) -> None: 

82 if not self._session_store or not session_id: 

83 return 

84 try: 

85 messages = self._session_store.load_messages(session_id) 

86 except Exception as e: 

87 logger.warning("加载 session 消息失败: %s", e) 

88 return 

89 if not messages or len(messages) < self.min_msgs: 

90 return 

91 

92 prompt = self.EXTRACT_PROMPT.format( 

93 messages=self._format_messages(messages) 

94 ) 

95 try: 

96 response = await self.llm.chat( 

97 [{"role": "user", "content": prompt}], 

98 tools=None, tool_choice="none", 

99 ) 

100 except Exception as e: 

101 logger.warning("LLM 抽取失败,跳过: %s", e) 

102 return 

103 

104 items = self._parse_llm_output(response.get("content", ""), session_id) 

105 for item in items: 

106 await self.pipeline.save(item) 

107 

108 @staticmethod 

109 def _format_messages(messages: list) -> str: 

110 """压缩消息为 LLM 提示(仅 role + content)""" 

111 lines = [] 

112 for m in messages: 

113 role = m.get("role", "?") 

114 content = (m.get("content") or "").strip() 

115 if content: 

116 lines.append(f"[{role}] {content[:200]}") 

117 return "\n".join(lines) 

118 

119 def _parse_llm_output(self, content: str, session_id: str) -> list[SaveItem]: 

120 """解析 LLM JSON 输出为 SaveItem 列表""" 

121 content = (content or "").strip() 

122 # 尝试提取 ```json ... ``` 包裹 

123 if content.startswith("```"): 

124 lines = content.split("\n") 

125 content = "\n".join(lines[1:-1]) if lines[-1].startswith("```") else "\n".join(lines[1:]) 

126 try: 

127 data = json.loads(content) 

128 except (ValueError, TypeError) as e: 

129 logger.warning("LLM 输出解析失败: %s", e) 

130 return [] 

131 if not isinstance(data, list): 

132 return [] 

133 items = [] 

134 for entry in data[:3]: # 兜底再 cap 

135 if not isinstance(entry, dict): 

136 continue 

137 key = entry.get("key", "").strip() 

138 value = entry.get("value", "").strip() 

139 if not key or not value: 

140 continue 

141 tags = entry.get("tags", []) or [] 

142 if not isinstance(tags, list): 

143 tags = [] 

144 items.append(SaveItem( 

145 key=key, value=value, source="extracted", 

146 tags=[str(t) for t in tags], session_id=session_id, 

147 )) 

148 return items[:self.max]