deerflow-code/offline-backend-20260512/backend/packages/harness/deerflow/agents/memory/manager.py
2026-09-07 18:24:55 +08:00

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"]