495 lines
19 KiB
Python
495 lines
19 KiB
Python
"""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
|