324 lines
13 KiB
Python
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",
|
|
]
|