156 lines
5.9 KiB
Python
156 lines
5.9 KiB
Python
"""Memory V2 - Hindsight 的 3 个 LangChain 工具。
|
|
|
|
仅当 ``config.memory.hindsight.memory_mode`` 为 ``tools`` 或 ``hybrid`` 时,
|
|
这些工具才会被 ``get_available_tools`` 加入工具集 (``context`` 模式不暴露)。
|
|
|
|
- hindsight_recall 多策略检索 (语义 + 实体图谱),查双 bank
|
|
- hindsight_reflect 跨记忆综合 (LLM 推理)
|
|
- hindsight_retain 主动存储一条记忆到全局 user bank
|
|
|
|
与 builtin 的 ``memory`` 工具一样,这些工具从 :class:`ToolRuntime` 与请求上下文
|
|
解析 user_id / agent_id,按生效配置懒建一个 :class:`HindsightProvider` 后调用。
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import logging
|
|
|
|
from langchain.tools import ToolRuntime, tool
|
|
from langgraph.config import get_config
|
|
from langgraph.typing import ContextT
|
|
|
|
from deerflow.agents.thread_state import ThreadState
|
|
from deerflow.config.memory_config import get_effective_memory_config
|
|
from deerflow.runtime.user_context import get_effective_user_id
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
|
|
def _resolve_agent_id(runtime: ToolRuntime[ContextT, ThreadState]) -> str:
|
|
"""从 runtime context / RunnableConfig 解析 agent_id,缺省 ``"default"``。"""
|
|
agent_id = runtime.context.get("agent_id") if runtime.context else None
|
|
if agent_id:
|
|
return str(agent_id)
|
|
runtime_config = getattr(runtime, "config", None) or {}
|
|
agent_id = runtime_config.get("configurable", {}).get("agent_id")
|
|
if agent_id:
|
|
return str(agent_id)
|
|
try:
|
|
agent_id = get_config().get("configurable", {}).get("agent_id")
|
|
except RuntimeError:
|
|
agent_id = None
|
|
return str(agent_id) if agent_id else "default"
|
|
|
|
|
|
def _build_provider(runtime: ToolRuntime[ContextT, ThreadState]):
|
|
"""按生效配置懒建并初始化一个 HindsightProvider;不可用时返回 ``None``。"""
|
|
from deerflow.agents.memory.providers.hindsight import HindsightProvider
|
|
|
|
user_id = get_effective_user_id()
|
|
config = get_effective_memory_config(user_id)
|
|
if not config.enabled or config.provider != "hindsight":
|
|
return None
|
|
provider = HindsightProvider(config.hindsight)
|
|
if not provider.is_available():
|
|
return None
|
|
thread_id = runtime.context.get("thread_id") if runtime.context else None
|
|
provider.initialize(
|
|
user_id=user_id,
|
|
agent_id=_resolve_agent_id(runtime),
|
|
thread_id=thread_id or "",
|
|
base_dir="",
|
|
)
|
|
return provider
|
|
|
|
|
|
@tool("hindsight_recall", parse_docstring=True)
|
|
def hindsight_recall_tool(
|
|
runtime: ToolRuntime[ContextT, ThreadState],
|
|
query: str,
|
|
) -> str:
|
|
"""检索 Hindsight 长期记忆,查找与 query 相关的历史信息。
|
|
|
|
用于回答关于用户、过往会话、之前讨论过的内容的问题。会同时检索全局用户
|
|
记忆和当前 agent 的工作记忆。
|
|
|
|
Args:
|
|
query: 检索查询,描述你想找回什么信息。
|
|
"""
|
|
provider = _build_provider(runtime)
|
|
if provider is None:
|
|
return json.dumps({"success": False, "error": "Hindsight 记忆不可用。"}, ensure_ascii=False)
|
|
try:
|
|
results = provider.recall(query)
|
|
except Exception as exc: # noqa: BLE001
|
|
logger.warning("hindsight_recall 失败:%s", exc)
|
|
return json.dumps({"success": False, "error": f"检索失败:{exc}"}, ensure_ascii=False)
|
|
if not results:
|
|
return json.dumps({"success": True, "result": "没有找到相关记忆。"}, ensure_ascii=False)
|
|
lines = "\n".join(f"{i}. {text}" for i, text in enumerate(results, 1))
|
|
return json.dumps({"success": True, "result": lines}, ensure_ascii=False)
|
|
|
|
|
|
@tool("hindsight_reflect", parse_docstring=True)
|
|
def hindsight_reflect_tool(
|
|
runtime: ToolRuntime[ContextT, ThreadState],
|
|
query: str,
|
|
) -> str:
|
|
"""对 Hindsight 记忆做跨记忆综合推理,得到一段综合性回答。
|
|
|
|
与 hindsight_recall 返回原始记忆片段不同,reflect 会让 LLM 综合多条记忆
|
|
给出推理结论。适合需要归纳、总结、跨记忆关联的问题。
|
|
|
|
Args:
|
|
query: 你希望综合分析的问题。
|
|
"""
|
|
provider = _build_provider(runtime)
|
|
if provider is None:
|
|
return json.dumps({"success": False, "error": "Hindsight 记忆不可用。"}, ensure_ascii=False)
|
|
try:
|
|
text = provider.reflect(query)
|
|
except Exception as exc: # noqa: BLE001
|
|
logger.warning("hindsight_reflect 失败:%s", exc)
|
|
return json.dumps({"success": False, "error": f"综合失败:{exc}"}, ensure_ascii=False)
|
|
return json.dumps(
|
|
{"success": True, "result": text or "没有找到相关记忆。"}, ensure_ascii=False
|
|
)
|
|
|
|
|
|
@tool("hindsight_retain", parse_docstring=True)
|
|
def hindsight_retain_tool(
|
|
runtime: ToolRuntime[ContextT, ThreadState],
|
|
content: str,
|
|
) -> str:
|
|
"""主动把一条信息存入 Hindsight 的全局用户记忆库。
|
|
|
|
用于记录跨会话、跨 agent 仍有价值的用户事实。Hindsight 会自动做实体抽取
|
|
和知识图谱关联。
|
|
|
|
Args:
|
|
content: 要存储的记忆内容。
|
|
"""
|
|
if not content or not content.strip():
|
|
return json.dumps({"success": False, "error": "content 不能为空。"}, ensure_ascii=False)
|
|
provider = _build_provider(runtime)
|
|
if provider is None:
|
|
return json.dumps({"success": False, "error": "Hindsight 记忆不可用。"}, ensure_ascii=False)
|
|
try:
|
|
provider.retain_to_user_bank(content.strip())
|
|
except Exception as exc: # noqa: BLE001
|
|
logger.warning("hindsight_retain 失败:%s", exc)
|
|
return json.dumps({"success": False, "error": f"存储失败:{exc}"}, ensure_ascii=False)
|
|
return json.dumps({"success": True, "result": "记忆已存储。"}, ensure_ascii=False)
|
|
|
|
|
|
#: 三个 Hindsight 工具,供 get_available_tools 在 tools/hybrid 模式下加入。
|
|
HINDSIGHT_TOOLS = [hindsight_recall_tool, hindsight_reflect_tool, hindsight_retain_tool]
|
|
|
|
|
|
__all__ = [
|
|
"hindsight_recall_tool",
|
|
"hindsight_reflect_tool",
|
|
"hindsight_retain_tool",
|
|
"HINDSIGHT_TOOLS",
|
|
]
|