314 lines
12 KiB
Python
314 lines
12 KiB
Python
"""Memory V2 - MemoryManager 编排层。
|
|
|
|
职责
|
|
====
|
|
编排 1 个内建 provider (``builtin``,始终存在) + 至多 1 个外部 provider
|
|
(P2 期的 ``hindsight``)。把系统提示词组装、召回、生命周期 hook 分发到各 provider。
|
|
任一 provider 出错都不会阻断其他 provider,也不会阻断对话。
|
|
|
|
典型用法 (由 MemoryMiddleware 每轮构建)
|
|
=======================================
|
|
::
|
|
|
|
manager = MemoryManager()
|
|
manager.add_provider(BuiltinFileProvider(...))
|
|
manager.add_provider(hindsight_provider) # 可选,至多 1 个
|
|
manager.initialize_all(user_id=..., agent_id=..., thread_id=..., base_dir=...)
|
|
|
|
block = manager.build_injection_block(query=user_msg, thread_id=tid)
|
|
# ... 把 block 注入系统提示词 ...
|
|
|
|
manager.sync_all(user_msg, assistant_msg, thread_id=tid)
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
from typing import Any
|
|
|
|
from deerflow.agents.memory.provider import MemoryProvider
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
#: 内建 provider 的固定名字。
|
|
BUILTIN_PROVIDER_NAME = "builtin"
|
|
KNOWLEDGE_PROVIDER_NAME = "knowledge"
|
|
_INTERNAL_PROVIDER_NAMES = {BUILTIN_PROVIDER_NAME, KNOWLEDGE_PROVIDER_NAME}
|
|
|
|
|
|
class MemoryManager:
|
|
"""编排内建 provider + 至多 1 个外部 provider。"""
|
|
|
|
def __init__(self) -> None:
|
|
self._providers: list[MemoryProvider] = []
|
|
self._has_external: bool = False
|
|
|
|
# ---- 注册 ------------------------------------------------------------
|
|
|
|
def add_provider(self, provider: MemoryProvider) -> bool:
|
|
"""注册一个 provider。
|
|
|
|
内建 provider (``name == "builtin"``) 始终接受。外部 provider 至多 1 个,
|
|
第二个会被拒绝并 log 警告。返回是否注册成功。
|
|
"""
|
|
is_internal = provider.name in _INTERNAL_PROVIDER_NAMES
|
|
if not is_internal:
|
|
if self._has_external:
|
|
existing = next(
|
|
(p.name for p in self._providers if p.name not in _INTERNAL_PROVIDER_NAMES),
|
|
"unknown",
|
|
)
|
|
logger.warning(
|
|
"拒绝注册记忆 provider '%s' —— 已存在外部 provider '%s'。"
|
|
"同一时间只允许 1 个外部 provider,请用 config.memory.provider 选择。",
|
|
provider.name,
|
|
existing,
|
|
)
|
|
return False
|
|
self._has_external = True
|
|
self._providers.append(provider)
|
|
logger.info("记忆 provider '%s' 已注册", provider.name)
|
|
return True
|
|
|
|
@property
|
|
def providers(self) -> list[MemoryProvider]:
|
|
"""按注册顺序返回所有 provider。"""
|
|
return list(self._providers)
|
|
|
|
def get_provider(self, name: str) -> MemoryProvider | None:
|
|
"""按名字取 provider,不存在返回 ``None``。"""
|
|
for p in self._providers:
|
|
if p.name == name:
|
|
return p
|
|
return None
|
|
|
|
# ---- 生命周期 ---------------------------------------------------------
|
|
|
|
def initialize_all(self, **kwargs: Any) -> None:
|
|
"""初始化所有 provider。单个 provider 失败不影响其他。"""
|
|
for provider in self._providers:
|
|
try:
|
|
provider.initialize(**kwargs)
|
|
except Exception as exc: # noqa: BLE001 - provider 失败必须隔离
|
|
logger.warning("记忆 provider '%s' initialize 失败:%s", provider.name, exc)
|
|
|
|
def shutdown_all(self) -> None:
|
|
"""关闭所有 provider (逆序,保证干净拆除)。"""
|
|
for provider in reversed(self._providers):
|
|
try:
|
|
provider.shutdown()
|
|
except Exception as exc: # noqa: BLE001
|
|
logger.warning("记忆 provider '%s' shutdown 失败:%s", provider.name, exc)
|
|
|
|
# ---- 系统提示词组装 / 召回 -------------------------------------------
|
|
|
|
def build_system_prompt(self) -> str:
|
|
"""收集所有 provider 的静态系统提示词块。"""
|
|
blocks: list[str] = []
|
|
for provider in self._providers:
|
|
try:
|
|
block = provider.system_prompt_block()
|
|
except Exception as exc: # noqa: BLE001
|
|
logger.warning("记忆 provider '%s' system_prompt_block 失败:%s", provider.name, exc)
|
|
continue
|
|
if block and block.strip():
|
|
blocks.append(block)
|
|
return "\n\n".join(blocks)
|
|
|
|
def prefetch_all(self, query: str, *, thread_id: str = "") -> str:
|
|
"""收集所有 provider 的召回上下文。单个失败不阻断其他。"""
|
|
parts: list[str] = []
|
|
for provider in self._providers:
|
|
try:
|
|
result = provider.prefetch(query, thread_id=thread_id)
|
|
except Exception as exc: # noqa: BLE001
|
|
logger.debug("记忆 provider '%s' prefetch 失败 (非致命):%s", provider.name, exc)
|
|
continue
|
|
if result and result.strip():
|
|
parts.append(result)
|
|
return "\n\n".join(parts)
|
|
|
|
def build_injection_block(self, query: str, *, thread_id: str = "") -> str:
|
|
"""组装注入系统提示词的完整记忆块。
|
|
|
|
= 所有 provider 的静态块 (frozen snapshot 等) + 动态召回块。
|
|
返回空串表示无内容可注入。
|
|
"""
|
|
static_part = self.build_system_prompt()
|
|
dynamic_part = self.prefetch_all(query, thread_id=thread_id)
|
|
parts = [p for p in (static_part, dynamic_part) if p and p.strip()]
|
|
return "\n\n".join(parts)
|
|
|
|
def queue_prefetch_all(self, query: str, *, thread_id: str = "") -> None:
|
|
"""通知所有 provider 后台预热下一轮召回。"""
|
|
for provider in self._providers:
|
|
try:
|
|
provider.queue_prefetch(query, thread_id=thread_id)
|
|
except Exception as exc: # noqa: BLE001
|
|
logger.debug("记忆 provider '%s' queue_prefetch 失败:%s", provider.name, exc)
|
|
|
|
# ---- 写入 ------------------------------------------------------------
|
|
|
|
def sync_all(
|
|
self,
|
|
user_content: str,
|
|
assistant_content: str,
|
|
*,
|
|
thread_id: str = "",
|
|
turn_index: int | None = None,
|
|
) -> None:
|
|
"""把一轮完成的对话同步到所有 provider。
|
|
|
|
``turn_index`` 是当前对话的轮次序号,透传给各 provider 用于轮次节流。
|
|
"""
|
|
for provider in self._providers:
|
|
try:
|
|
provider.sync_turn(
|
|
user_content,
|
|
assistant_content,
|
|
thread_id=thread_id,
|
|
turn_index=turn_index,
|
|
)
|
|
except Exception as exc: # noqa: BLE001
|
|
logger.warning("记忆 provider '%s' sync_turn 失败:%s", provider.name, exc)
|
|
|
|
def on_memory_write(
|
|
self,
|
|
action: str,
|
|
target: str,
|
|
content: str,
|
|
metadata: dict[str, Any] | None = None,
|
|
) -> None:
|
|
"""builtin 写入后,通知外部 provider 镜像。
|
|
|
|
跳过 builtin 自身 (它就是写入源,镜像回去没有意义)。
|
|
"""
|
|
for provider in self._providers:
|
|
if provider.name == BUILTIN_PROVIDER_NAME:
|
|
continue
|
|
try:
|
|
provider.on_memory_write(action, target, content, metadata)
|
|
except Exception as exc: # noqa: BLE001
|
|
logger.debug("记忆 provider '%s' on_memory_write 失败:%s", provider.name, exc)
|
|
|
|
# ---- 可选 hooks ------------------------------------------------------
|
|
|
|
def on_session_switch_all(
|
|
self,
|
|
new_thread_id: str,
|
|
*,
|
|
parent_thread_id: str = "",
|
|
reset: bool = False,
|
|
**kwargs: Any,
|
|
) -> None:
|
|
"""通知所有 provider thread_id 已切换。"""
|
|
if not new_thread_id:
|
|
return
|
|
for provider in self._providers:
|
|
try:
|
|
provider.on_session_switch(
|
|
new_thread_id,
|
|
parent_thread_id=parent_thread_id,
|
|
reset=reset,
|
|
**kwargs,
|
|
)
|
|
except Exception as exc: # noqa: BLE001
|
|
logger.debug("记忆 provider '%s' on_session_switch 失败:%s", provider.name, exc)
|
|
|
|
|
|
def build_memory_manager(
|
|
*,
|
|
user_id: str,
|
|
agent_id: str | None = None,
|
|
thread_id: str | None = None,
|
|
config: Any = None,
|
|
app_config: Any = None,
|
|
) -> MemoryManager:
|
|
"""按生效配置组装并初始化一个 MemoryManager。
|
|
|
|
这是构建 manager 的统一入口 —— 根据 ``config.version`` 选择 V1 或 V2 provider:
|
|
|
|
- ``version="v1"``: 仅注册 V1BuiltinProvider (memory.json + LLM 自动提取)
|
|
- ``version="v2"``: 注册 BuiltinFileProvider + 可选 HindsightProvider (USER.md/MEMORY.md)
|
|
|
|
所有 provider 已 ``initialize`` 完毕,可直接使用。
|
|
|
|
Parameters
|
|
----------
|
|
user_id / agent_id / thread_id:
|
|
请求上下文标识。
|
|
config:
|
|
生效的 :class:`MemoryConfig`;为 ``None`` 时按 ``user_id`` 解析 per-user 配置。
|
|
"""
|
|
from deerflow.config.memory_config import get_effective_memory_config
|
|
from deerflow.config.paths import get_paths
|
|
|
|
if config is None:
|
|
config = get_effective_memory_config(user_id)
|
|
|
|
manager = MemoryManager()
|
|
base_dir = str(get_paths().base_dir)
|
|
|
|
memory_enabled = bool(getattr(config, "enabled", True) and getattr(config.injection, "enabled", True))
|
|
|
|
if memory_enabled and getattr(config, "version", "v2") == "v1":
|
|
# V1 模式:原版 DeerFlow memory.json + LLM 自动提取
|
|
from deerflow.agents.memory.providers.v1_builtin import V1BuiltinProvider
|
|
|
|
v1_cfg = config.v1
|
|
manager.add_provider(
|
|
V1BuiltinProvider(
|
|
max_facts=v1_cfg.max_facts,
|
|
fact_confidence_threshold=v1_cfg.fact_confidence_threshold,
|
|
debounce_seconds=v1_cfg.debounce_seconds,
|
|
model_name=v1_cfg.model_name,
|
|
token_budget=v1_cfg.token_budget,
|
|
)
|
|
)
|
|
elif memory_enabled:
|
|
# V2 模式:USER.md / MEMORY.md + 可选 Hindsight
|
|
from deerflow.agents.memory.providers.builtin import BuiltinFileProvider
|
|
|
|
if config.builtin.enabled:
|
|
manager.add_provider(
|
|
BuiltinFileProvider(
|
|
memory_char_limit=config.builtin.memory_char_limit,
|
|
user_char_limit=config.builtin.user_char_limit,
|
|
deduplicate_on_load=config.builtin.deduplicate_on_load,
|
|
security=config.security,
|
|
)
|
|
)
|
|
|
|
if config.provider == "hindsight":
|
|
try:
|
|
from deerflow.agents.memory.providers.hindsight import HindsightProvider
|
|
|
|
hindsight = HindsightProvider(config.hindsight)
|
|
if hindsight.is_available():
|
|
manager.add_provider(hindsight)
|
|
else:
|
|
logger.info("Hindsight provider 已配置但不可用 (包未装或缺 api_url),退回纯 builtin")
|
|
except Exception as exc: # noqa: BLE001 - Hindsight 接入失败不能阻断记忆系统
|
|
logger.warning("加载 Hindsight provider 失败,退回纯 builtin:%s", exc)
|
|
|
|
try:
|
|
from deerflow.agents.memory.providers.knowledge import KnowledgeProvider
|
|
|
|
knowledge = KnowledgeProvider(app_config=app_config)
|
|
if knowledge.is_available():
|
|
manager.add_provider(knowledge)
|
|
except Exception as exc: # noqa: BLE001 - 知识库路由失败不能阻断对话
|
|
logger.warning("加载 Knowledge provider 失败:%s", exc)
|
|
|
|
manager.initialize_all(
|
|
user_id=user_id,
|
|
agent_id=agent_id,
|
|
thread_id=thread_id,
|
|
base_dir=base_dir,
|
|
app_config=app_config,
|
|
)
|
|
return manager
|
|
|
|
|
|
__all__ = ["MemoryManager", "BUILTIN_PROVIDER_NAME", "build_memory_manager"]
|