deerflow-code/offline-backend-20260512/backend/app/gateway/routers/memory.py
2026-09-07 18:24:55 +08:00

495 lines
19 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""Memory V2 记忆 API 路由。
所有端点用 ``/api/memory/v2`` 前缀。Memory V2 重写后,旧的 facts/history
端点 (``/api/memory``、``/memory/facts`` 等) 已移除。
端点:
- ``GET /memory/v2`` 读取 builtin 本地记忆 (USER.md / MEMORY.md)
- ``POST /memory/v2/entries`` 新增一条记忆条目
- ``PUT /memory/v2/entries`` 替换一条记忆条目
- ``DELETE /memory/v2/entries`` 删除一条记忆条目
- ``GET /memory/v2/config`` 读取生效配置 + 用户可改字段白名单
- ``PATCH /memory/v2/config`` 写入 per-user 配置覆盖 (白名单字段)
- ``DELETE /memory/v2/config`` 撤销 per-user 配置覆盖
- ``POST /memory/v2/search`` 代理 Hindsight 检索 (recall / reflect)
可选 ``agent`` 查询参数:省略或传 ``global`` 表示全局桶;传 agent 名表示该
agent 的工作记忆桶。USER.md 始终是 per-user 全局的;MEMORY.md 按 agent 隔离。
"""
from __future__ import annotations
from typing import Any
from fastapi import APIRouter, HTTPException, Query, Request
from pydantic import BaseModel, Field
from app.gateway.deps import get_memory_config_store
from deerflow.config.agents_config import AGENT_NAME_PATTERN
from deerflow.config.memory_config import (
USER_WRITABLE_FIELDS,
delete_user_memory_override_db,
get_effective_memory_config,
save_user_memory_override_db,
)
from deerflow.config.paths import get_paths
from deerflow.runtime.user_context import get_effective_user_id
router = APIRouter(prefix="/api", tags=["memory"])
# 表示"全局桶 (无 agent)"的 agent 取值。
_GLOBAL_AGENT_VALUES: frozenset[str] = frozenset({"", "global", "__global__", "default"})
def _resolve_agent_name(agent: str | None) -> str | None:
"""校验 ``agent`` 查询参数,把全局哨兵规整为 ``None``。
返回 ``None`` 表示全局桶,否则返回小写的 agent 名。非法名抛 400。
"""
if agent is None:
return None
stripped = agent.strip()
if stripped.lower() in _GLOBAL_AGENT_VALUES:
return None
if not AGENT_NAME_PATTERN.match(stripped):
raise HTTPException(
status_code=400,
detail=f"Invalid agent name {stripped!r}: must match {AGENT_NAME_PATTERN.pattern}",
)
return stripped.lower()
_AgentQuery = Query(
None,
description="可选 agent 名。省略或传 'global' 表示全局桶;传 agent 名表示该 agent 的工作记忆桶。",
examples=["my-agent"],
pattern=r"^[A-Za-z0-9-]+$|^global$|^$",
)
# ===========================================================================
# 响应 / 请求模型
# ===========================================================================
class MemoryFact(BaseModel):
"""V1 记忆中的单条事实。"""
id: str = Field(default="", description="事实 UUID")
content: str = Field(default="", description="事实文本")
category: str = Field(default="knowledge", description="事实类别")
confidence: float = Field(default=0.0, description="置信度 0-1")
createdAt: str = Field(default="", description="创建时间 ISO 格式")
sourceError: str = Field(default="", description="纠错类事实:被纠正的原错误描述")
class MemoryV2DataResponse(BaseModel):
"""Memory 数据视图 —— V1/V2 共用,多余字段各版本置空。"""
# V2 字段
user_entries: list[str] = Field(default_factory=list, description="V2: USER.md 条目 (用户画像)")
memory_entries: list[str] = Field(default_factory=list, description="V2: MEMORY.md 条目 (agent 工作笔记)")
user_usage: str = Field(default="", description="字符/计数用量描述")
memory_usage: str = Field(default="", description="字符/计数用量描述")
# V1 字段
facts: list[MemoryFact] = Field(default_factory=list, description="V1: facts 列表 (含置信度)")
user_context: dict[str, str] = Field(default_factory=dict, description="V1: userContext 摘要")
history: dict[str, str] = Field(default_factory=dict, description="V1: history 历史摘要")
class MemoryV2ConfigResponse(BaseModel):
"""Memory 生效配置 (含三层覆盖合并结果) + 用户可改字段白名单 + 版本号。"""
config: dict[str, Any] = Field(..., description="按 per-user 覆盖合并后的生效配置")
writable_fields: list[str] = Field(..., description="用户可在运行时覆盖的字段 (点号路径)")
version: str = Field(default="v2", description="当前记忆版本:v1 或 v2")
class MemoryConfigPatchRequest(BaseModel):
"""Memory V2 配置覆盖请求。
overrides 支持嵌套形式 ``{"injection": {"enabled": false}}`` 或点号扁平形式
``{"injection.enabled": false}``;非白名单字段会被忽略。
"""
overrides: dict[str, Any] = Field(..., description="要写入 per-user 覆盖的字段")
class MemoryEntryCreateRequest(BaseModel):
"""新增一条 V2 记忆条目。"""
target: str = Field(..., description="memory(工作笔记) 或 user(用户画像)")
content: str = Field(..., min_length=1, description="条目内容")
class MemoryEntryReplaceRequest(BaseModel):
"""替换一条 V2 记忆条目。"""
target: str = Field(..., description="memory 或 user")
old_text: str = Field(..., min_length=1, description="短唯一子串,定位要替换的条目")
content: str = Field(..., min_length=1, description="新内容")
class MemorySearchRequest(BaseModel):
"""Hindsight 检索请求。"""
query: str = Field(..., min_length=1, description="检索查询")
mode: str = Field(default="recall", description="recall(返回原始片段) 或 reflect(LLM 综合)")
class MemorySearchResponse(BaseModel):
"""Hindsight 检索结果。"""
available: bool = Field(..., description="Hindsight 是否可用 (未配置/未装包时为 false)")
results: list[str] = Field(default_factory=list, description="recall 模式的检索结果")
text: str = Field(default="", description="reflect 模式的综合结论")
# ===========================================================================
# 辅助函数
# ===========================================================================
def _is_v1(user_id: str) -> bool:
"""判断当前用户是否处于 V1 记忆模式。"""
config = get_effective_memory_config(user_id)
return getattr(config, "version", "v2") == "v1"
def _build_v1_provider(user_id: str, agent_name: str | None = None):
"""构建并初始化 V1BuiltinProvider。"""
from deerflow.agents.memory.providers.v1_builtin import V1BuiltinProvider
config = get_effective_memory_config(user_id)
v1_cfg = config.v1
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,
)
provider.initialize(
user_id=user_id,
agent_id=agent_name,
base_dir=str(get_paths().base_dir),
)
return provider
def _build_builtin_provider(user_id: str, agent_name: str | None):
"""按生效配置构建并初始化一个 V2 builtin provider。"""
from deerflow.agents.memory.providers.builtin import BuiltinFileProvider
config = get_effective_memory_config(user_id)
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,
)
provider.initialize(
user_id=user_id,
agent_id=agent_name or "default",
base_dir=str(get_paths().base_dir),
)
return provider
def _validate_v2_target(target: str) -> str:
"""校验 V2 entries 的 target,非法时抛 400。"""
if target not in ("memory", "user"):
raise HTTPException(status_code=400, detail=f"非法 target '{target}',只能是 'memory' 或 'user'。")
return target
def _v2_data_response(provider) -> MemoryV2DataResponse:
"""从 V2 builtin provider 读出当前数据。"""
return MemoryV2DataResponse(
user_entries=provider.read("user"),
memory_entries=provider.read("memory"),
user_usage=provider.usage("user"),
memory_usage=provider.usage("memory"),
)
def _v1_data_response(provider) -> MemoryV2DataResponse:
"""从 V1 provider 读出当前数据,转换为统一响应格式。
支持新格式 {user: {workContext: {summary, updatedAt}}} 和老格式 {userContext: {workContext: str}}。
"""
data = provider.read_v1_data()
raw_facts = data.get("facts", [])
facts = [
MemoryFact(
id=str(f.get("id", "")),
content=str(f.get("content", "")),
category=str(f.get("category", "knowledge")),
confidence=float(f.get("confidence", 0)),
createdAt=str(f.get("createdAt", "")),
sourceError=str(f.get("sourceError", "")),
)
for f in raw_facts
]
def _extract_section(section: dict) -> dict[str, str]:
"""从 {key: {summary: str, updatedAt: str}} 或 {key: str} 中提取摘要字符串。"""
result = {}
for k, v in section.items():
if isinstance(v, dict):
summary = v.get("summary", "")
else:
summary = str(v)
if summary:
result[k] = summary
return result
# 新格式用 "user" key,老格式用 "userContext" key(已由 storage 迁移,此处兜底)
ctx_raw = data.get("user") or data.get("userContext") or {}
history_raw = data.get("history") or {}
return MemoryV2DataResponse(
facts=sorted(facts, key=lambda f: f.confidence, reverse=True),
user_context=_extract_section(ctx_raw),
history=_extract_section(history_raw),
memory_usage=provider.usage_text(),
)
def _raise_if_write_failed(result: dict) -> None:
"""builtin provider 写操作失败时抛 400。"""
if not result.get("success"):
raise HTTPException(status_code=400, detail=result.get("error", "记忆写入失败。"))
# ===========================================================================
# 数据端点
# ===========================================================================
@router.get(
"/memory/v2",
response_model=MemoryV2DataResponse,
summary="Get Memory Data",
description="读取记忆数据。V1 模式返回 facts 列表;V2 模式返回 USER.md / MEMORY.md 条目。",
)
async def get_memory_v2(agent: str | None = _AgentQuery) -> MemoryV2DataResponse:
"""读取当前用户的本地记忆。"""
agent_name = _resolve_agent_name(agent)
user_id = get_effective_user_id()
try:
if _is_v1(user_id):
provider = _build_v1_provider(user_id, agent_name)
return _v1_data_response(provider)
provider = _build_builtin_provider(user_id, agent_name)
return _v2_data_response(provider)
except Exception as exc: # noqa: BLE001
raise HTTPException(status_code=500, detail="Failed to read memory data.") from exc
def _assert_not_v1(user_id: str) -> None:
"""V1 模式下写入操作返回 400。"""
if _is_v1(user_id):
raise HTTPException(status_code=400, detail="V1 模式下记忆由系统自动提取,不支持手动修改。")
@router.post(
"/memory/v2/entries",
response_model=MemoryV2DataResponse,
summary="Add Memory V2 Entry",
description="手动往 USER.md / MEMORY.md 新增一条记忆条目 (仅 V2 模式)。",
)
async def add_memory_v2_entry(
request: MemoryEntryCreateRequest,
agent: str | None = _AgentQuery,
) -> MemoryV2DataResponse:
"""手动新增一条 V2 记忆条目。"""
user_id = get_effective_user_id()
_assert_not_v1(user_id)
_validate_v2_target(request.target)
try:
provider = _build_builtin_provider(user_id, _resolve_agent_name(agent))
result = provider.add(request.target, request.content)
except HTTPException:
raise
except Exception as exc: # noqa: BLE001
raise HTTPException(status_code=500, detail="Failed to add memory entry.") from exc
_raise_if_write_failed(result)
return _v2_data_response(provider)
@router.put(
"/memory/v2/entries",
response_model=MemoryV2DataResponse,
summary="Replace Memory V2 Entry",
description="用 old_text 定位一条记忆条目并替换为新内容 (仅 V2 模式)。",
)
async def replace_memory_v2_entry(
request: MemoryEntryReplaceRequest,
agent: str | None = _AgentQuery,
) -> MemoryV2DataResponse:
"""手动替换一条 V2 记忆条目。"""
user_id = get_effective_user_id()
_assert_not_v1(user_id)
_validate_v2_target(request.target)
try:
provider = _build_builtin_provider(user_id, _resolve_agent_name(agent))
result = provider.replace(request.target, request.old_text, request.content)
except HTTPException:
raise
except Exception as exc: # noqa: BLE001
raise HTTPException(status_code=500, detail="Failed to replace memory entry.") from exc
_raise_if_write_failed(result)
return _v2_data_response(provider)
@router.delete(
"/memory/v2/entries",
response_model=MemoryV2DataResponse,
summary="Delete Memory V2 Entry",
description="用 old_text 定位并删除一条记忆条目 (仅 V2 模式)。",
)
async def delete_memory_v2_entry(
target: str = Query(..., description="memory 或 user"),
old_text: str = Query(..., min_length=1, description="短唯一子串,定位要删除的条目"),
agent: str | None = _AgentQuery,
) -> MemoryV2DataResponse:
"""手动删除一条 V2 记忆条目。"""
user_id = get_effective_user_id()
_assert_not_v1(user_id)
_validate_v2_target(target)
try:
provider = _build_builtin_provider(user_id, _resolve_agent_name(agent))
result = provider.remove(target, old_text)
except HTTPException:
raise
except Exception as exc: # noqa: BLE001
raise HTTPException(status_code=500, detail="Failed to delete memory entry.") from exc
_raise_if_write_failed(result)
return _v2_data_response(provider)
# ===========================================================================
# 配置端点
# ===========================================================================
@router.get(
"/memory/v2/config",
response_model=MemoryV2ConfigResponse,
summary="Get Memory V2 Effective Config",
description="读取按 per-user 覆盖合并后的生效记忆配置,以及用户可改字段白名单。",
)
async def get_memory_v2_config() -> MemoryV2ConfigResponse:
"""读取当前用户的生效记忆配置。"""
user_id = get_effective_user_id()
config = get_effective_memory_config(user_id)
return MemoryV2ConfigResponse(
config=config.model_dump(),
writable_fields=sorted(USER_WRITABLE_FIELDS),
version=getattr(config, "version", "v2"),
)
@router.patch(
"/memory/v2/config",
response_model=MemoryV2ConfigResponse,
summary="Patch Memory V2 Config",
description="把 hindsight.recall_budget 写入当前用户的 per-user 配置覆盖 (DB 存储)。",
)
async def patch_memory_v2_config(request: MemoryConfigPatchRequest, req: Request) -> MemoryV2ConfigResponse:
"""更新当前用户的 per-user 记忆配置覆盖 (仅 hindsight.recall_budget)。"""
user_id = get_effective_user_id()
store = get_memory_config_store(req)
try:
await save_user_memory_override_db(user_id, request.overrides, store)
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
except Exception as exc: # noqa: BLE001
raise HTTPException(status_code=500, detail="Failed to save memory config override.") from exc
config = get_effective_memory_config(user_id)
return MemoryV2ConfigResponse(
config=config.model_dump(),
writable_fields=sorted(USER_WRITABLE_FIELDS),
version=getattr(config, "version", "v2"),
)
@router.delete(
"/memory/v2/config",
response_model=MemoryV2ConfigResponse,
summary="Delete Memory V2 Config Override",
description=(
"撤销当前用户的 per-user 配置覆盖。传 ?field=hindsight.recall_budget 只撤销该字段;"
"不传 field 则清空全部覆盖。"
),
)
async def delete_memory_v2_config(
req: Request,
field: str | None = Query(None, description="要撤销的点号路径字段;省略则清空全部覆盖"),
) -> MemoryV2ConfigResponse:
"""撤销当前用户的 per-user 记忆配置覆盖。"""
user_id = get_effective_user_id()
store = get_memory_config_store(req)
try:
await delete_user_memory_override_db(user_id, field, store)
except Exception as exc: # noqa: BLE001
raise HTTPException(status_code=500, detail="Failed to delete memory config override.") from exc
config = get_effective_memory_config(user_id)
return MemoryV2ConfigResponse(
config=config.model_dump(),
writable_fields=sorted(USER_WRITABLE_FIELDS),
version=getattr(config, "version", "v2"),
)
# ===========================================================================
# Hindsight 检索代理端点
# ===========================================================================
@router.post(
"/memory/v2/search",
response_model=MemorySearchResponse,
summary="Search Hindsight Memory",
description=(
"代理 Hindsight 检索。mode=recall 返回原始记忆片段;mode=reflect 返回 LLM 综合结论。"
"Hindsight 未配置或不可用时返回 available=false。"
),
)
async def search_memory_v2(
request: MemorySearchRequest,
agent: str | None = _AgentQuery,
) -> MemorySearchResponse:
"""代理 Hindsight 检索 (recall / reflect)。"""
user_id = get_effective_user_id()
config = get_effective_memory_config(user_id)
if not config.enabled or config.provider != "hindsight":
return MemorySearchResponse(available=False)
try:
from deerflow.agents.memory.providers.hindsight import HindsightProvider
provider = HindsightProvider(config.hindsight)
if not provider.is_available():
return MemorySearchResponse(available=False)
provider.initialize(
user_id=user_id,
agent_id=_resolve_agent_name(agent) or "default",
thread_id="",
base_dir="",
)
if request.mode == "reflect":
text = provider.reflect(request.query)
return MemorySearchResponse(available=True, text=text)
results = provider.recall(request.query)
return MemorySearchResponse(available=True, results=results)
except HTTPException:
raise
except Exception as exc: # noqa: BLE001
raise HTTPException(status_code=500, detail="Failed to search Hindsight memory.") from exc