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
« 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
4import hashlib
5import json
6import logging
7import re
8from abc import ABC, abstractmethod
10from .save_pipeline import SaveItem
12logger = logging.getLogger("merco.memory.strategy")
15class MemorySaveStrategy(ABC):
16 """监听事件,构造 SaveItem 喂给 Pipeline"""
18 name: str = ""
20 def __init__(self, pipeline):
21 self.pipeline = pipeline
23 @abstractmethod
24 async def on_event(self, event: str, **kwargs) -> None: ...
27class ExplicitRememberStrategy(MemorySaveStrategy):
28 """/remember <text> 显式存一条记忆"""
29 name = "explicit_remember"
31 def subscribe(self, hooks) -> None:
32 """注册到 HookRegistry"""
33 hooks.on("command.remember", self._on_remember)
35 async def on_event(self, event: str, **kwargs) -> None:
36 """兼容直接调用(测试用)"""
37 await self._on_remember(**kwargs)
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)
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}"
53class SessionEndExtractStrategy(MemorySaveStrategy):
54 """session.destroy 时用 LLM 抽取 1-3 条 insight 记忆"""
55 name = "session_end_extract"
57 EXTRACT_PROMPT = """从以下对话中抽取 1-3 条值得长期记住的关键信息(用户偏好、事实、决策)。
58仅返回 JSON 数组,每条形如 {{"key": "snake_case_key", "value": "原文", "tags": ["tag1"]}}。
59没有值得记的就返回 []。
61对话:
62{messages}
63"""
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
74 def subscribe(self, hooks) -> None:
75 hooks.on("session.destroy", self._on_destroy)
77 async def on_event(self, event: str, **kwargs) -> None:
78 """兼容直接调用(测试用)"""
79 await self._on_destroy(**kwargs)
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
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
104 items = self._parse_llm_output(response.get("content", ""), session_id)
105 for item in items:
106 await self.pipeline.save(item)
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)
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]