text2mem.py 13 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329
  1. from __future__ import annotations
  2. import re
  3. from typing import Any
  4. from ..config import settings
  5. from ..db import utcnow
  6. from ..embeddings import embedding_client, embedding_fingerprint
  7. from ..llm import LLMUnavailable
  8. from ..schemas import ChatResponse, MemoryItem, SystemDescriptor
  9. from .base import MemoryAgent
  10. class Text2MemAgent(MemoryAgent):
  11. """内置记忆实现,负责显式写入、混合检索和审计。"""
  12. id = "text2mem"
  13. _EXPLICIT_WRITE_RE = re.compile(r"记住|记下|请记录|请保存")
  14. _DECLARATIVE_RE = re.compile(
  15. r"我(?:偏好|喜欢|习惯|要求|正在|在开发|决定|选择|使用)|"
  16. r"项目(?:目前|现在|已经|使用|采用|数据库|前端|后端)|"
  17. r"(?:前端|后端|数据库|对话模型|向量模型)(?:使用|采用|选择|是|改成)|"
  18. r"更新一下|以后|上次|之前|曾经|长期规则|代码要求"
  19. )
  20. _QUESTION_RE = re.compile(
  21. r"[??]$|(?:什么|哪些|哪个|怎么|如何|是否|有没有|多少|为什么).*(?:[??]|$)"
  22. )
  23. _STOPWORDS = frozenset(
  24. {
  25. "什么",
  26. "哪些",
  27. "哪个",
  28. "怎么",
  29. "如何",
  30. "是否",
  31. "有没有",
  32. "请问",
  33. "告诉我",
  34. "我的",
  35. "当前",
  36. "相关",
  37. "一下",
  38. "遇到",
  39. "遇到了",
  40. "的",
  41. "吗",
  42. "呢",
  43. }
  44. )
  45. def __init__(self) -> None:
  46. """初始化内置实现及其系统说明。"""
  47. super().__init__()
  48. self.descriptor = SystemDescriptor(
  49. id="text2mem",
  50. name="Text2Mem",
  51. paradigm="IR 操作契约",
  52. description=(
  53. "当前实现 Encode / Retrieve、持久化、删除与审计;"
  54. "Update / Lock / Expire 等完整 IR 操作仍是后续项。"
  55. ),
  56. available=True,
  57. mode="native",
  58. status="ready",
  59. package="本项目内置教学实现",
  60. )
  61. @classmethod
  62. def _should_encode(cls, message: str) -> bool:
  63. """区分需要写入的陈述和只需要检索的问题。"""
  64. normalized = " ".join(message.strip().split())
  65. if not normalized:
  66. return False
  67. if cls._EXPLICIT_WRITE_RE.search(normalized):
  68. return True
  69. if cls._QUESTION_RE.search(normalized):
  70. return False
  71. return bool(cls._DECLARATIVE_RE.search(normalized))
  72. @staticmethod
  73. def _memory_type(content: str) -> str:
  74. """用规则把内容归入四类教学记忆。"""
  75. if re.search(r"上次|之前|曾经|遇到|发生|故障|失败", content):
  76. return "episodic"
  77. if re.search(r"流程|步骤|先.+再|以后.+按", content):
  78. return "procedural"
  79. if re.search(r"更新一下|进度|已完成|完成了|下一步|当前在做", content):
  80. return "task-state"
  81. return "semantic"
  82. @classmethod
  83. def _query_type(cls, message: str) -> str | None:
  84. """根据问题内容判断优先查找哪类记忆。"""
  85. if re.search(r"流程|步骤|怎么修复|如何修复|操作方法", message):
  86. return "procedural"
  87. if re.search(r"上次|之前|历史|曾经|遇到|发生|故障", message):
  88. return "episodic"
  89. if re.search(r"进度|完成|下一步|现在做到", message):
  90. return "task-state"
  91. if re.search(r"偏好|喜欢|技术栈|数据库|前端|后端|项目", message):
  92. return "semantic"
  93. return None
  94. @classmethod
  95. def _tokens(cls, text: str) -> set[str]:
  96. """提取英文单词和中文短词,供关键词匹配使用。"""
  97. tokens = {item.casefold() for item in re.findall(r"[A-Za-z][A-Za-z0-9_.+-]{1,}", text)}
  98. for chunk in re.findall(r"[\u4e00-\u9fff]{2,}", text):
  99. if chunk not in cls._STOPWORDS:
  100. tokens.add(chunk)
  101. if len(chunk) > 2:
  102. tokens.update(
  103. chunk[index : index + 2]
  104. for index in range(len(chunk) - 1)
  105. if chunk[index : index + 2] not in cls._STOPWORDS
  106. )
  107. return tokens
  108. async def _encode(self, content: str) -> tuple[MemoryItem, bool]:
  109. """去重后写入记忆,返回记忆及是否新建。"""
  110. normalized = " ".join(content.strip().split()).casefold()
  111. # 相同内容只保留一份,重复输入只记审计记录。
  112. for existing in await self.repo.list_memories(self.id):
  113. if " ".join(existing.content.strip().split()).casefold() == normalized:
  114. await self._audit(
  115. "ENC/Encode",
  116. existing.id,
  117. status="deduplicated",
  118. details={"source": "user_direct"},
  119. )
  120. return existing, False
  121. now = utcnow()
  122. # 新记忆在这里补齐分类、来源和时间信息。
  123. item = MemoryItem(
  124. id=self.repo.memory_id("t2m"),
  125. system=self.id,
  126. workspace_id=settings.workspace_id,
  127. content=content.strip(),
  128. memory_type=self._memory_type(content),
  129. source="user_direct",
  130. confidence=1.0,
  131. valid_from=now,
  132. created_at=now,
  133. updated_at=now,
  134. metadata={
  135. "ir": {"stage": "ENC", "op": "Encode"},
  136. "policy": {"confirmation": True, "locked": False},
  137. },
  138. )
  139. await self.repo.add_memory(item)
  140. await self._audit(
  141. "ENC/Encode",
  142. item.id,
  143. details={"source": "user_direct", "memory_type": item.memory_type},
  144. )
  145. return item, True
  146. async def _sync_embeddings(self, items: list[MemoryItem]) -> dict[str, Any]:
  147. """只更新内容指纹发生变化的向量。"""
  148. if not embedding_client.configured:
  149. return {"enabled": False, "embedded": 0, "removed": 0}
  150. current_ids = {item.id for item in items}
  151. previous_hashes = await self.repo.embedding_hashes(self.id)
  152. # 记忆已经删除时,旧向量不能继续留在搜索结果里。
  153. stale_ids = set(previous_hashes) - current_ids
  154. for memory_id in stale_ids:
  155. await self.repo.delete_embedding(self.id, memory_id)
  156. # 指纹没变的内容无需重复请求向量服务。
  157. pending = [
  158. item
  159. for item in items
  160. if previous_hashes.get(item.id) != embedding_fingerprint(item.content)
  161. ]
  162. if pending:
  163. vectors = await embedding_client.embed([item.content for item in pending])
  164. for item, vector in zip(pending, vectors):
  165. await self.repo.upsert_embedding(
  166. system=self.id,
  167. memory_id=item.id,
  168. content=item.content,
  169. source=item.source,
  170. content_hash=embedding_fingerprint(item.content),
  171. embedding=vector,
  172. )
  173. return {"enabled": True, "embedded": len(pending), "removed": len(stale_ids)}
  174. async def _retrieve(self, message: str, limit: int = 5) -> tuple[list[MemoryItem], dict[str, Any]]:
  175. """融合词面、记忆类型和向量分数进行召回。"""
  176. items = await self.repo.list_memories(self.id)
  177. if not items:
  178. return [], {"strategy": "empty", "embedding": {"enabled": False}}
  179. query_tokens = self._tokens(message)
  180. expected_type = self._query_type(message)
  181. token_overlap: dict[str, float] = {}
  182. lexical_scores: dict[str, float] = {}
  183. # 先看关键词是否对得上,同一类记忆再加一点分。
  184. for item in items:
  185. item_tokens = self._tokens(item.content)
  186. overlap = len(query_tokens & item_tokens) / max(1, len(query_tokens))
  187. token_overlap[item.id] = overlap
  188. type_bonus = 0.35 if expected_type and item.memory_type == expected_type else 0.0
  189. lexical_scores[item.id] = overlap + type_bonus
  190. semantic_scores: dict[str, float] = {}
  191. try:
  192. # 配置了向量服务时,再补一轮相近内容搜索。
  193. embedding_status = await self._sync_embeddings(items)
  194. if embedding_status["enabled"]:
  195. query_vector = (await embedding_client.embed([message]))[0]
  196. semantic_matches = await self.repo.search_embeddings(
  197. self.id,
  198. query_vector,
  199. limit=min(max(limit * 2, 5), 20),
  200. )
  201. semantic_scores = {
  202. str(match["memory_id"]): float(match["score"])
  203. for match in semantic_matches
  204. }
  205. except Exception as exc:
  206. # 向量搜索失败后继续使用关键词结果。
  207. embedding_status = {"enabled": True, "error": str(exc), "embedded": 0, "removed": 0}
  208. ranked: list[tuple[float, MemoryItem]] = []
  209. for item in items:
  210. lexical = lexical_scores[item.id]
  211. semantic = semantic_scores.get(item.id, 0.0)
  212. # 类型、关键词和内容都对不上时,不把这条记忆交给模型。
  213. if (
  214. expected_type
  215. and item.memory_type != expected_type
  216. and token_overlap[item.id] < 0.25
  217. and semantic < 0.65
  218. ):
  219. continue
  220. if lexical <= 0 and semantic < settings.embedding_min_score:
  221. continue
  222. ranked.append((lexical + max(semantic, 0.0), item))
  223. if not ranked and expected_type:
  224. # 没有直接命中时,至少返回问题所对应的同类记忆。
  225. ranked = [
  226. (0.1, item)
  227. for item in items
  228. if item.memory_type == expected_type
  229. ]
  230. ranked.sort(key=lambda pair: (pair[0], pair[1].updated_at), reverse=True)
  231. selected = [item for _, item in ranked[:limit]]
  232. return selected, {
  233. "strategy": "lexical+type+embedding" if semantic_scores else "lexical+type",
  234. "expected_type": expected_type,
  235. "embedding": embedding_status,
  236. "matches": [
  237. {
  238. "memory_id": item.id,
  239. "lexical": round(lexical_scores[item.id], 4),
  240. "semantic": round(semantic_scores.get(item.id, 0.0), 4),
  241. }
  242. for item in selected
  243. ],
  244. }
  245. async def chat(self, message: str) -> ChatResponse:
  246. """执行写入判断、记忆检索和回答生成。"""
  247. events: list[dict[str, Any]] = []
  248. # 先处理可能的写入,再查询本轮需要的上下文。
  249. if self._should_encode(message):
  250. item, created = await self._encode(message)
  251. events.append(
  252. {
  253. "event": "ENC/Encode",
  254. "memory_id": item.id,
  255. "content": item.content,
  256. "memory_type": item.memory_type,
  257. "status": "created" if created else "deduplicated",
  258. }
  259. )
  260. context, retrieval = await self._retrieve(message)
  261. system_prompt = (
  262. "你是一个 AI 编程助手。Text2Mem 只允许通过显式 IR 记忆操作读写记忆。"
  263. "回答时区分当前用户陈述、历史记忆和推断,不要把推断写成事实。"
  264. "只能依据给定的记忆证据回答历史问题;证据不足时要明确说明。"
  265. )
  266. context_text = "\n".join(
  267. f"- [{item.memory_type} | {item.source} | {item.updated_at.isoformat()}] {item.content}"
  268. for item in context
  269. ) or "(暂无命中记忆)"
  270. try:
  271. answer = await self.llm.chat(system_prompt, f"记忆上下文:\n{context_text}\n\n用户:{message}")
  272. mode = "llm"
  273. except LLMUnavailable:
  274. # 没配模型也能查看本轮写入和检索到的原始内容。
  275. mode = "local-fallback"
  276. answer = (
  277. "Text2Mem 本地模式:我已按显式 IR 规则处理本轮输入。\n\n"
  278. f"当前命中记忆:{context_text}\n\n"
  279. "配置 LLM_API_KEY 后,可让模型基于这些记忆生成自然语言回答。"
  280. )
  281. except Exception as exc:
  282. # 模型临时出错不影响已经完成的记忆写入和查询。
  283. mode = "local-fallback"
  284. retrieval["llm_error"] = str(exc)[:500]
  285. answer = (
  286. "模型调用失败,Text2Mem 已退回本地证据模式;记忆写入和检索结果不受影响。\n\n"
  287. f"当前命中记忆:{context_text}"
  288. )
  289. # 保存实际用到的记忆,方便之后检查回答来源。
  290. await self._audit(
  291. "RET/Retrieve",
  292. details={
  293. "query": message,
  294. "count": len(context),
  295. "memory_ids": [item.id for item in context],
  296. **retrieval,
  297. },
  298. )
  299. return ChatResponse(
  300. system="text2mem",
  301. answer=answer,
  302. mode=mode,
  303. memory_context=context,
  304. memory_events=events,
  305. audit_events=await self.audit(),
  306. )