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

324 lines
13 KiB
Python

"""Memory V2 - 记忆子系统配置。
配置结构 (详见 docs/MEMORY_V2_DESIGN_ZH.md §9)
==============================================
``MemoryConfig`` 采用嵌套结构::
memory:
enabled: true
provider: hindsight # 外部 provider 选哪个 (builtin 始终启用)
builtin: {...} # 本地 USER.md / MEMORY.md
hindsight: {...} # Hindsight HTTP 客户端
injection: {...} # 系统提示词注入
security: {...} # 内容安全扫描
两层配置覆盖
============
``get_effective_memory_config(user_id)`` 按优先级合并 (高 -> 低):
1. per-user 覆盖 数据库 ``user_memory_configs`` 表 (经进程内 cache 读取)
(有 per-user 记录时自动强制 enabled=True)
2. config.yaml ``memory:`` 段 (即 ``get_memory_config()``)
3. Pydantic 默认值
第 1 层只允许覆盖 :data:`USER_WRITABLE_FIELDS` 白名单内的字段。
cache 由 Gateway lifespan 启动时预热,写入时同步更新,保持同步读取路径。
"""
from __future__ import annotations
import copy
import logging
from typing import Any
from pydantic import BaseModel, Field
logger = logging.getLogger(__name__)
# ===========================================================================
# 嵌套子配置
# ===========================================================================
class BuiltinMemoryConfig(BaseModel):
"""本地文件 provider (USER.md / MEMORY.md) 配置。"""
enabled: bool = Field(default=True, description="是否启用 builtin 本地记忆 provider")
memory_char_limit: int = Field(
default=2200, ge=200, le=20000, description="MEMORY.md 字符配额"
)
user_char_limit: int = Field(
default=1375, ge=200, le=20000, description="USER.md 字符配额"
)
deduplicate_on_load: bool = Field(
default=True, description="加载时是否对条目去重"
)
class HindsightMemoryConfig(BaseModel):
"""Hindsight 外部 provider 配置。"""
mode: str = Field(
default="local_external",
description="连接模式:cloud / local_embedded / local_external",
)
api_url: str = Field(default="http://localhost:8765", description="Hindsight API 地址")
api_key: str = Field(default="", description="API key,建议用 $HINDSIGHT_API_KEY 环境变量")
bank_id_template: str = Field(
default="deerflow-user-{user_id}",
description="全局 user bank 命名模板,占位符 {user_id}/{agent_id}/{thread_id}",
)
work_bank_id_template: str = Field(
default="deerflow-work-{user_id}-{agent_id}",
description="agent 私有工作 bank 命名模板",
)
memory_mode: str = Field(
default="hybrid",
description="工具暴露策略:context(仅自动) / tools(仅工具) / hybrid(两者)",
)
bank_mission: str = Field(default="", description="reflect 推理的身份/框架描述")
bank_retain_mission: str = Field(default="", description="引导抽取内容的 mission")
recall_budget: str = Field(default="mid", description="召回深度:low / mid / high")
recall_max_tokens: int = Field(default=4096, ge=256, le=32768, description="召回结果 token 上限")
recall_max_input_chars: int = Field(default=800, ge=64, le=8000, description="召回 query 长度上限")
auto_recall: bool = Field(default=True, description="每轮对话前是否自动召回")
auto_retain: bool = Field(default=True, description="是否自动写入对话")
retain_async: bool = Field(default=True, description="retain 是否在 Hindsight 服务端异步处理")
retain_every_n_turns: int = Field(default=1, ge=1, le=50, description="每 N 轮 retain 一次")
retain_user_prefix: str = Field(default="User", description="自动 retain 时用户发言前缀")
retain_assistant_prefix: str = Field(default="Assistant", description="自动 retain 时助手发言前缀")
class InjectionConfig(BaseModel):
"""系统提示词记忆注入配置。"""
enabled: bool = Field(default=True, description="是否往系统提示词注入记忆")
max_tokens: int = Field(default=2000, ge=100, le=8000, description="注入 token 上限")
include_builtin: bool = Field(default=True, description="是否注入 builtin 本地记忆")
include_hindsight: bool = Field(default=True, description="是否注入 Hindsight 召回结果")
context_tag: str = Field(
default="memory-context", description="包裹外部召回内容的标签名"
)
class SecurityConfig(BaseModel):
"""记忆内容安全配置。"""
scan_content: bool = Field(default=True, description="是否扫描注入/外渗模式")
block_invisible_unicode: bool = Field(default=True, description="是否阻断不可见 unicode")
# ===========================================================================
# 顶层配置
# ===========================================================================
class V1MemoryConfig(BaseModel):
"""原版 DeerFlow V1 记忆 provider 配置 (memory.json + LLM 自动提取)。"""
max_facts: int = Field(default=100, ge=10, le=500, description="最多保留的事实条数")
fact_confidence_threshold: float = Field(default=0.7, ge=0.0, le=1.0, description="事实置信度阈值,低于此值不保存")
debounce_seconds: int = Field(default=30, ge=5, le=300, description="对话结束后延迟多少秒才触发 LLM 提取")
model_name: str = Field(default="", description="用于事实提取的 LLM 模型名;空则继承全局默认模型")
token_budget: int = Field(default=2000, ge=200, le=8000, description="注入系统提示词的 token 预算")
class MemoryConfig(BaseModel):
"""记忆子系统的全局配置 (config.yaml 的 ``memory:`` 段)。"""
enabled: bool = Field(default=True, description="记忆子系统总开关")
version: str = Field(
default="v2",
description="记忆系统版本:v1(原版 DeerFlow,memory.json + LLM 自动提取) 或 v2(USER.md/MEMORY.md)",
)
provider: str = Field(
default="builtin",
description="启用的外部 provider 名:builtin(无外部) 或 hindsight",
)
v1: V1MemoryConfig = Field(default_factory=V1MemoryConfig)
builtin: BuiltinMemoryConfig = Field(default_factory=BuiltinMemoryConfig)
hindsight: HindsightMemoryConfig = Field(default_factory=HindsightMemoryConfig)
injection: InjectionConfig = Field(default_factory=InjectionConfig)
security: SecurityConfig = Field(default_factory=SecurityConfig)
# ===========================================================================
# per-user 运行时覆盖 — 白名单
# ===========================================================================
# 用户唯一可覆盖的字段:召回深度。
# injection.* / builtin.* / enabled 始终由 config.yaml 控制,用户不可关闭。
USER_WRITABLE_FIELDS: frozenset[str] = frozenset(
{
"hindsight.recall_budget",
}
)
_VALID_RECALL_BUDGETS: frozenset[str] = frozenset({"low", "mid", "high"})
def _filter_to_whitelist(overrides: dict[str, Any]) -> dict[str, Any]:
"""把覆盖 dict 过滤为只含白名单字段的嵌套 dict。
输入支持嵌套形式 ``{"hindsight": {"recall_budget": "high"}}`` 或
点号扁平形式 ``{"hindsight.recall_budget": "high"}``。
非白名单字段被丢弃并 log。
"""
result: dict[str, Any] = {}
def _set_nested(path: str, value: Any) -> None:
if path not in USER_WRITABLE_FIELDS:
logger.warning("忽略非白名单的记忆配置覆盖字段:%s", path)
return
section, field = path.split(".", 1)
result.setdefault(section, {})[field] = value
for key, value in overrides.items():
if isinstance(value, dict):
for sub_key, sub_value in value.items():
_set_nested(f"{key}.{sub_key}", sub_value)
else:
_set_nested(key, value)
return result
def _deep_merge(base: dict[str, Any], override: dict[str, Any]) -> dict[str, Any]:
"""递归合并两个 dict,``override`` 优先。返回新 dict,不修改入参。"""
merged = copy.deepcopy(base)
for key, value in override.items():
if isinstance(value, dict) and isinstance(merged.get(key), dict):
merged[key] = _deep_merge(merged[key], value)
else:
merged[key] = value
return merged
# ===========================================================================
# 全局单例 + 加载入口
# ===========================================================================
_memory_config: MemoryConfig = MemoryConfig()
def get_memory_config() -> MemoryConfig:
"""返回 config.yaml 级别的记忆配置 (不含 per-user 覆盖)。"""
return _memory_config
def set_memory_config(config: MemoryConfig) -> None:
"""设置 config.yaml 级别的记忆配置 (测试与运行时重载用)。"""
global _memory_config
_memory_config = config
def load_memory_config_from_dict(config_dict: dict) -> None:
"""从 dict 加载记忆配置。"""
global _memory_config
_memory_config = MemoryConfig(**config_dict)
# ===========================================================================
# per-user recall_budget 进程内 cache
#
# Gateway lifespan 启动时调用 load_recall_budget_cache() 从 DB 批量加载。
# API 写入时同步更新 cache,使 get_effective_memory_config() 保持同步读取。
# ===========================================================================
_recall_budget_cache: dict[str, str] = {}
def load_recall_budget_cache(data: dict[str, str]) -> None:
"""用 DB 批量加载结果替换整个 cache (启动时调用一次)。"""
global _recall_budget_cache
_recall_budget_cache = dict(data)
logger.info("recall_budget cache loaded: %d user(s)", len(_recall_budget_cache))
def set_recall_budget_cache(user_id: str, value: str) -> None:
"""写入单个用户的 recall_budget cache。"""
_recall_budget_cache[user_id] = value
def delete_recall_budget_cache(user_id: str) -> None:
"""从 cache 中移除某用户的 recall_budget。"""
_recall_budget_cache.pop(user_id, None)
def get_recall_budget_from_cache(user_id: str) -> str | None:
"""返回 cache 中某用户的 recall_budget,未设置时返回 None。"""
return _recall_budget_cache.get(user_id)
def get_effective_memory_config(user_id: str) -> MemoryConfig:
"""返回某用户的生效记忆配置 (同步,从进程内 cache 读取 per-user 覆盖)。
优先级 (高 -> 低):per-user DB 覆盖 (经 cache) > config.yaml > 默认值。
有 per-user 记录时自动强制 enabled=True。
"""
base = _memory_config.model_dump()
recall_budget = _recall_budget_cache.get(user_id)
if recall_budget is not None:
base.setdefault("hindsight", {})["recall_budget"] = recall_budget
base["enabled"] = True
try:
return MemoryConfig(**base)
except Exception as exc: # noqa: BLE001 - 覆盖产生非法配置时降级
logger.warning("合并后的记忆配置非法,退回 config.yaml: %s", exc)
return _memory_config
# ===========================================================================
# Gateway 调用的异步 DB 对接函数
# ===========================================================================
async def save_user_memory_override_db(user_id: str, patch: dict[str, Any], store: Any) -> dict[str, Any]:
"""把白名单字段写入 DB store 并更新 cache,返回生效配置 dict。
``patch`` 支持嵌套或点号扁平形式;非白名单字段被忽略。
``store`` 为 ``UserMemoryConfigStore`` 实例 (避免循环 import,用 Any)。
"""
filtered = _filter_to_whitelist(patch)
recall_budget = (filtered.get("hindsight") or {}).get("recall_budget")
if recall_budget is not None:
if recall_budget not in _VALID_RECALL_BUDGETS:
raise ValueError(f"recall_budget 必须为 {sorted(_VALID_RECALL_BUDGETS)} 之一,收到:{recall_budget!r}")
await store.set_recall_budget(user_id, recall_budget)
set_recall_budget_cache(user_id, recall_budget)
return get_effective_memory_config(user_id).model_dump()
async def delete_user_memory_override_db(user_id: str, field: str | None, store: Any) -> dict[str, Any]:
"""撤销某用户的 per-user 覆盖并更新 cache,返回生效配置 dict。
``field=None`` 清空全部;``field="hindsight.recall_budget"`` 只清该字段。
"""
if field is None or field == "hindsight.recall_budget":
await store.delete(user_id)
delete_recall_budget_cache(user_id)
else:
logger.warning("忽略删除非白名单的记忆配置覆盖字段: %s", field)
return get_effective_memory_config(user_id).model_dump()
__all__ = [
"V1MemoryConfig",
"BuiltinMemoryConfig",
"HindsightMemoryConfig",
"InjectionConfig",
"SecurityConfig",
"MemoryConfig",
"USER_WRITABLE_FIELDS",
"get_memory_config",
"set_memory_config",
"load_memory_config_from_dict",
"get_effective_memory_config",
"load_recall_budget_cache",
"set_recall_budget_cache",
"delete_recall_budget_cache",
"get_recall_budget_from_cache",
"save_user_memory_override_db",
"delete_user_memory_override_db",
]