mirror of
https://github.com/bytedance/deer-flow.git
synced 2026-07-30 01:46:01 +00:00
feat(memory): add mem0 HTTP memory backend (#4528)
* feat(memory): add mem0 HTTP memory backend * fix(memory): address mem0 review feedback --------- Co-authored-by: Willem Jiang <willem.jiang@gmail.com>
This commit is contained in:
parent
43ed2b7d45
commit
352f247a81
@ -979,6 +979,12 @@ startup.
|
||||
|
||||
Across sessions, DeerFlow builds a persistent memory of your profile, preferences, and accumulated knowledge. The more you use it, the better it knows you — your writing style, your technical stack, your recurring workflows. Memory is stored locally and stays under your control.
|
||||
|
||||
DeerMem remains the default local backend. An opt-in `mem0` backend is also
|
||||
available for the hosted mem0 Platform API or API-compatible self-hosted
|
||||
servers. Its token-bearing `base_url` must use HTTPS by default; plaintext HTTP
|
||||
requires an explicit local-development opt-in. See the
|
||||
[mem0 backend guide](backend/packages/harness/deerflow/agents/memory/backends/mem0/README.md).
|
||||
|
||||
Memory updates now skip duplicate fact entries at apply time, so repeated preferences and context do not accumulate endlessly across sessions.
|
||||
|
||||
File-backed memory now separates global user context from agent facts. Each user has one `memory.json` containing only the project-independent `user` and `history` summaries; every fact is a canonical Markdown file below `agents/{agent_name}/facts/`. Existing lead-agent middleware, API, Settings, import/export, and embedded-client calls that omit `agent_name` resolve inside DeerMem to the reserved `__default__` bucket. That bucket is outside the valid custom-agent name grammar, so a real custom agent named `lead-agent` has a separate fact repository and deleting a custom agent cannot delete a memory-only directory without `config.yaml`. Public agent identifiers are case-insensitive and canonicalized to lowercase. Runtime/API readers still receive a compatibility `facts` array for the selected/default agent, so the frontend does not read agent facts from `memory.json`; structured Markdown `source` metadata is projected to the historical string field at the MemoryManager boundary. An unscoped Clear All first migrates facts from unread legacy per-agent JSON without adopting its soon-to-be-cleared summaries, then removes shared summaries and facts from every agent bucket while preserving agent configuration files, so a later read cannot resurrect skipped legacy facts; an explicitly agent-scoped clear removes only that agent's facts. On first normal read, old facts embedded in the user JSON are migrated automatically to `__default__`; facts written to the earlier implicit `lead-agent` bucket are also moved when that directory is not a real custom agent. Migration and normal writes notify the configured retrieval adapter only after durable storage locks are released. DeerMem uses a scope-aware SQLite FTS5/BM25 adapter by default, stores only rebuildable derived index data under `.retrieval/`, and rebuilds it in the background during Gateway startup or lazily on the first scoped search. A corrupt derived index is recreated automatically. Set `memory.backend_config.retrieval_adapter` to an empty string to disable it and use the local substring fallback. Chinese tokenization is optional; install the backend `memory-zh` extra (`uv sync --extra memory-zh`) for jieba-assisted sub-phrase search. Journaled writes, a shared user lock, and optimistic user-memory revisions prevent silent lost updates.
|
||||
|
||||
@ -248,6 +248,23 @@ Blocking-IO runtime gate (`tests/blocking_io/`):
|
||||
Boundary check (harness → app import firewall):
|
||||
- `tests/test_harness_boundary.py` — ensures `packages/harness/deerflow/` never imports from `app.*`
|
||||
|
||||
Memory backend async boundary:
|
||||
- `MemoryMiddleware.aafter_agent` calls `MemoryManager.aadd`; network-backed
|
||||
managers must override their `a*` methods to offload or use native async I/O.
|
||||
- The mem0 backend requires an HTTPS `base_url` by default because requests
|
||||
carry an API token. Plain HTTP requires the explicit
|
||||
`backend_config.allow_insecure_http: true` local-development opt-in.
|
||||
- Gateway memory routes offload the synchronous management contract with
|
||||
`asyncio.to_thread`, so backend file or HTTP I/O does not run on the ASGI
|
||||
event loop. Gateway startup and shutdown also resolve the manager off-loop,
|
||||
because a backend's `from_config` may perform a fail-fast connectivity check.
|
||||
- A backend may set `requires_passive_writes_in_tool_mode = True` when tool-mode
|
||||
search is supported but durable writes still depend on conversation-level
|
||||
extraction. Such backends receive memory tools and retain `MemoryMiddleware`.
|
||||
- Prompt recall rethrows `MemoryManagerError` only when backend config declares
|
||||
`failure_policy.read: fail_closed`; other recall errors preserve the existing
|
||||
log-and-empty-context behavior.
|
||||
|
||||
CI runs these regression tests for every pull request via [.github/workflows/backend-unit-tests.yml](../.github/workflows/backend-unit-tests.yml).
|
||||
|
||||
Agentic browser sessions are process-local. The Gateway startup safety gate rejects
|
||||
|
||||
@ -228,7 +228,7 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]:
|
||||
from deerflow.agents.memory import get_memory_manager
|
||||
|
||||
if startup_config.memory.enabled:
|
||||
manager = get_memory_manager()
|
||||
manager = await asyncio.to_thread(get_memory_manager)
|
||||
warm_retrieval = getattr(manager, "warm_retrieval", None)
|
||||
if callable(warm_retrieval):
|
||||
retrieval_warm_task = asyncio.create_task(
|
||||
@ -252,7 +252,7 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]:
|
||||
try:
|
||||
from deerflow.agents.memory import get_memory_manager
|
||||
|
||||
manager = get_memory_manager()
|
||||
manager = await asyncio.to_thread(get_memory_manager)
|
||||
warmed = await asyncio.wait_for(
|
||||
asyncio.to_thread(manager.warm),
|
||||
timeout=5,
|
||||
@ -411,7 +411,7 @@ async def lifespan(app: FastAPI) -> AsyncGenerator[None, None]:
|
||||
if app_cfg.memory.enabled:
|
||||
from deerflow.agents.memory import get_memory_manager
|
||||
|
||||
manager = get_memory_manager()
|
||||
manager = await asyncio.to_thread(get_memory_manager)
|
||||
flush_timeout = app_cfg.memory.shutdown_flush_timeout_seconds
|
||||
completed = await asyncio.to_thread(manager.shutdown_flush, flush_timeout)
|
||||
if completed:
|
||||
|
||||
@ -1,5 +1,6 @@
|
||||
"""Memory API router for retrieving and managing global memory data."""
|
||||
|
||||
import asyncio
|
||||
from typing import Any, Literal
|
||||
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
@ -146,7 +147,7 @@ def _unsupported_501(manager: object, label: str) -> HTTPException:
|
||||
)
|
||||
|
||||
|
||||
def _get_memory_or_501(manager: MemoryManager, user_id: str, label: str) -> dict[str, Any]:
|
||||
async def _get_memory_or_501(manager: MemoryManager, user_id: str, label: str) -> dict[str, Any]:
|
||||
"""Read the full memory doc; 501 if the backend doesn't expose one.
|
||||
|
||||
``get_memory`` is tier-2 (default ``raise NotImplementedError``); a minimal
|
||||
@ -157,7 +158,7 @@ def _get_memory_or_501(manager: MemoryManager, user_id: str, label: str) -> dict
|
||||
endpoint's verb, e.g. "get memory" / "export memory" / "reload memory").
|
||||
"""
|
||||
try:
|
||||
return manager.get_memory(user_id=user_id)
|
||||
return await asyncio.to_thread(manager.get_memory, user_id=user_id)
|
||||
except NotImplementedError:
|
||||
raise _unsupported_501(manager, label) from None
|
||||
except (MemoryConflictError, MemoryCorruptionError) as exc:
|
||||
@ -239,8 +240,8 @@ async def get_memory(http_request: Request) -> MemoryResponse:
|
||||
}
|
||||
```
|
||||
"""
|
||||
manager = get_memory_manager()
|
||||
memory_data = _get_memory_or_501(manager, _resolve_memory_user_id(http_request), "get memory")
|
||||
manager = await asyncio.to_thread(get_memory_manager)
|
||||
memory_data = await _get_memory_or_501(manager, _resolve_memory_user_id(http_request), "get memory")
|
||||
return MemoryResponse(**memory_data)
|
||||
|
||||
|
||||
@ -261,9 +262,9 @@ async def reload_memory(http_request: Request) -> MemoryResponse:
|
||||
The reloaded memory data.
|
||||
"""
|
||||
user_id = _resolve_memory_user_id(http_request)
|
||||
manager = get_memory_manager()
|
||||
manager = await asyncio.to_thread(get_memory_manager)
|
||||
try:
|
||||
memory_data = manager.reload_memory(user_id=user_id)
|
||||
memory_data = await asyncio.to_thread(manager.reload_memory, user_id=user_id)
|
||||
except NotImplementedError:
|
||||
# Non-DeerMem backends have no reload concept; fall back to get_memory
|
||||
# (read-only refresh, so degrading is safe and still useful -- vs fact
|
||||
@ -271,7 +272,7 @@ async def reload_memory(http_request: Request) -> MemoryResponse:
|
||||
# would hide data loss). If get_memory is also unsupported (a minimal
|
||||
# backend with no full doc), surface 501 rather than a raw 500: reads
|
||||
# degrade only when there is a doc to degrade to.
|
||||
memory_data = _get_memory_or_501(manager, user_id, "reload memory")
|
||||
memory_data = await _get_memory_or_501(manager, user_id, "reload memory")
|
||||
except (MemoryConflictError, MemoryCorruptionError) as exc:
|
||||
raise _map_memory_manager_error(exc) from exc
|
||||
return MemoryResponse(**memory_data)
|
||||
@ -286,9 +287,9 @@ async def reload_memory(http_request: Request) -> MemoryResponse:
|
||||
)
|
||||
async def clear_memory(http_request: Request) -> MemoryResponse:
|
||||
"""Clear all persisted memory data."""
|
||||
manager = get_memory_manager()
|
||||
manager = await asyncio.to_thread(get_memory_manager)
|
||||
try:
|
||||
memory_data = manager.clear_memory(user_id=_resolve_memory_user_id(http_request))
|
||||
memory_data = await asyncio.to_thread(manager.clear_memory, user_id=_resolve_memory_user_id(http_request))
|
||||
except NotImplementedError:
|
||||
raise _unsupported_501(manager, "clear memory") from None
|
||||
except (MemoryConflictError, MemoryCorruptionError) as exc:
|
||||
@ -308,9 +309,10 @@ async def clear_memory(http_request: Request) -> MemoryResponse:
|
||||
)
|
||||
async def create_memory_fact_endpoint(request: FactCreateRequest, http_request: Request) -> MemoryResponse:
|
||||
"""Create a single fact manually."""
|
||||
manager = get_memory_manager()
|
||||
manager = await asyncio.to_thread(get_memory_manager)
|
||||
try:
|
||||
memory_data, fact_id = manager.create_fact(
|
||||
memory_data, fact_id = await asyncio.to_thread(
|
||||
manager.create_fact,
|
||||
content=request.content,
|
||||
category=request.category,
|
||||
confidence=request.confidence,
|
||||
@ -340,9 +342,9 @@ async def create_memory_fact_endpoint(request: FactCreateRequest, http_request:
|
||||
)
|
||||
async def delete_memory_fact_endpoint(fact_id: str, http_request: Request) -> MemoryResponse:
|
||||
"""Delete a single fact from memory by fact id."""
|
||||
manager = get_memory_manager()
|
||||
manager = await asyncio.to_thread(get_memory_manager)
|
||||
try:
|
||||
memory_data = manager.delete_fact(fact_id, user_id=_resolve_memory_user_id(http_request))
|
||||
memory_data = await asyncio.to_thread(manager.delete_fact, fact_id, user_id=_resolve_memory_user_id(http_request))
|
||||
except NotImplementedError:
|
||||
raise _unsupported_501(manager, "delete fact") from None
|
||||
except KeyError as exc:
|
||||
@ -364,9 +366,10 @@ async def delete_memory_fact_endpoint(fact_id: str, http_request: Request) -> Me
|
||||
)
|
||||
async def update_memory_fact_endpoint(fact_id: str, request: FactPatchRequest, http_request: Request) -> MemoryResponse:
|
||||
"""Partially update a single fact manually."""
|
||||
manager = get_memory_manager()
|
||||
manager = await asyncio.to_thread(get_memory_manager)
|
||||
try:
|
||||
memory_data = manager.update_fact(
|
||||
memory_data = await asyncio.to_thread(
|
||||
manager.update_fact,
|
||||
fact_id=fact_id,
|
||||
content=request.content,
|
||||
category=request.category,
|
||||
@ -396,8 +399,8 @@ async def update_memory_fact_endpoint(fact_id: str, request: FactPatchRequest, h
|
||||
)
|
||||
async def export_memory(http_request: Request) -> MemoryResponse:
|
||||
"""Export the current memory data."""
|
||||
manager = get_memory_manager()
|
||||
memory_data = _get_memory_or_501(manager, _resolve_memory_user_id(http_request), "export memory")
|
||||
manager = await asyncio.to_thread(get_memory_manager)
|
||||
memory_data = await _get_memory_or_501(manager, _resolve_memory_user_id(http_request), "export memory")
|
||||
return MemoryResponse(**memory_data)
|
||||
|
||||
|
||||
@ -410,9 +413,13 @@ async def export_memory(http_request: Request) -> MemoryResponse:
|
||||
)
|
||||
async def import_memory(request: MemoryResponse, http_request: Request) -> MemoryResponse:
|
||||
"""Import and persist memory data."""
|
||||
manager = get_memory_manager()
|
||||
manager = await asyncio.to_thread(get_memory_manager)
|
||||
try:
|
||||
memory_data = manager.import_memory(request.model_dump(exclude_none=True), user_id=_resolve_memory_user_id(http_request))
|
||||
memory_data = await asyncio.to_thread(
|
||||
manager.import_memory,
|
||||
request.model_dump(exclude_none=True),
|
||||
user_id=_resolve_memory_user_id(http_request),
|
||||
)
|
||||
except NotImplementedError:
|
||||
raise _unsupported_501(manager, "import memory") from None
|
||||
except (MemoryConflictError, MemoryCorruptionError) as exc:
|
||||
@ -486,8 +493,8 @@ async def get_memory_status(http_request: Request) -> MemoryStatusResponse:
|
||||
Combined memory configuration and current data.
|
||||
"""
|
||||
config = get_memory_config()
|
||||
manager = get_memory_manager()
|
||||
memory_data = _get_memory_or_501(manager, _resolve_memory_user_id(http_request), "get memory status")
|
||||
manager = await asyncio.to_thread(get_memory_manager)
|
||||
memory_data = await _get_memory_or_501(manager, _resolve_memory_user_id(http_request), "get memory status")
|
||||
|
||||
return MemoryStatusResponse(
|
||||
config=MemoryConfigResponse(
|
||||
|
||||
@ -275,6 +275,7 @@ def _assemble_from_features(
|
||||
|
||||
memory_cfg: MemoryConfig = feat.memory_config or get_memory_config()
|
||||
if should_use_memory_tools(memory_cfg):
|
||||
from deerflow.agents.memory.manager import backend_requires_passive_writes_in_tool_mode
|
||||
from deerflow.agents.memory.tools import get_memory_tools
|
||||
|
||||
existing_names = {tool.name for tool in extra_tools}
|
||||
@ -284,8 +285,10 @@ def _assemble_from_features(
|
||||
continue
|
||||
extra_tools.append(memory_tool)
|
||||
existing_names.add(memory_tool.name)
|
||||
# MemoryMiddleware is intentionally NOT appended in tool mode.
|
||||
# The model drives memory via tools instead of passive middleware.
|
||||
if backend_requires_passive_writes_in_tool_mode(memory_cfg.manager_class):
|
||||
from deerflow.agents.middlewares.memory_middleware import MemoryMiddleware
|
||||
|
||||
chain.append(MemoryMiddleware(agent_name=name, memory_config=memory_cfg))
|
||||
else:
|
||||
if memory_cfg.mode == "tool" and not memory_cfg.enabled:
|
||||
logger.warning("memory.mode is 'tool' but memory.enabled is false; memory tools will not be registered.")
|
||||
|
||||
@ -384,9 +384,13 @@ def build_middlewares(
|
||||
# Add TitleMiddleware
|
||||
middlewares.append(TitleMiddleware(app_config=resolved_app_config))
|
||||
|
||||
# Add MemoryMiddleware (after TitleMiddleware) — skipped in enabled tool mode
|
||||
# Add MemoryMiddleware after TitleMiddleware. Tool mode normally skips it;
|
||||
# conversation-extraction backends may explicitly retain passive writes.
|
||||
if should_use_memory_tools(resolved_app_config.memory):
|
||||
pass
|
||||
from deerflow.agents.memory.manager import backend_requires_passive_writes_in_tool_mode
|
||||
|
||||
if backend_requires_passive_writes_in_tool_mode(resolved_app_config.memory.manager_class):
|
||||
middlewares.append(MemoryMiddleware(agent_name=agent_name, memory_config=resolved_app_config.memory))
|
||||
else:
|
||||
if resolved_app_config.memory.mode == "tool" and not resolved_app_config.memory.enabled:
|
||||
logger.warning("memory.mode is 'tool' but memory.enabled is false; memory tools will not be registered.")
|
||||
|
||||
@ -745,6 +745,7 @@ def _get_memory_context(
|
||||
Returns:
|
||||
Formatted memory context string wrapped in XML tags, or empty string if disabled.
|
||||
"""
|
||||
config = None
|
||||
try:
|
||||
from deerflow.agents.memory import get_memory_manager
|
||||
from deerflow.runtime.user_context import resolve_runtime_user_id
|
||||
@ -771,8 +772,13 @@ def _get_memory_context(
|
||||
{memory_content}
|
||||
</memory>
|
||||
"""
|
||||
except Exception:
|
||||
except Exception as exc:
|
||||
logger.exception("Failed to load memory context")
|
||||
from deerflow.agents.memory import MemoryManagerError
|
||||
|
||||
failure_policy = getattr(config, "backend_config", {}).get("failure_policy", {}) if config is not None else {}
|
||||
if isinstance(exc, MemoryManagerError) and failure_policy.get("read") == "fail_closed":
|
||||
raise
|
||||
return ""
|
||||
|
||||
|
||||
|
||||
@ -0,0 +1,73 @@
|
||||
# mem0 memory backend
|
||||
|
||||
Uses mem0 (Platform hosted API, or any API-compatible self-hosted server) as
|
||||
DeerFlow's memory store. Fully stateless in-process: dedup, fact extraction,
|
||||
and storage are server-side, so it is safe for multi-worker Gateway
|
||||
deployments.
|
||||
|
||||
## Configuration
|
||||
|
||||
```yaml
|
||||
memory:
|
||||
enabled: true
|
||||
injection_enabled: true
|
||||
manager_class: mem0
|
||||
mode: middleware # or "tool"
|
||||
backend_config:
|
||||
api_key_env: MEM0_API_KEY # key read from env, never in config.yaml
|
||||
base_url: https://api.mem0.ai # or your self-hosted mem0 server
|
||||
allow_insecure_http: false # true only for trusted local HTTP dev
|
||||
top_k: 8
|
||||
score_threshold: 0.1
|
||||
max_injection_chars: 12000
|
||||
timeout_seconds: 10
|
||||
startup_policy: fail_fast # fail_fast | tolerate
|
||||
failure_policy:
|
||||
read: fail_open # fail_open | fail_closed
|
||||
write: log_and_drop # log_and_drop | raise
|
||||
```
|
||||
|
||||
Set the key in the environment: `export MEM0_API_KEY=...`
|
||||
|
||||
`base_url` must use HTTPS because every request carries the API key. For a
|
||||
trusted local-development server that only exposes HTTP, opt in explicitly
|
||||
with `allow_insecure_http: true`; do not use that setting across an untrusted
|
||||
network.
|
||||
|
||||
## Identity mapping
|
||||
|
||||
| DeerFlow | mem0 |
|
||||
|---|---|
|
||||
| `user_id` | `user_id` |
|
||||
| `agent_name` | `agent_id` |
|
||||
| `thread_id` | `run_id` |
|
||||
|
||||
## Limitations
|
||||
|
||||
- `mode: middleware` recall is query-less (the `get_context` contract carries
|
||||
no query): the bucket's most recent `top_k` memories are injected. For
|
||||
query-aware semantic recall use `mode: tool`.
|
||||
- `mode: tool` retains the passive per-turn write middleware for this backend,
|
||||
because mem0 extracts and deduplicates facts from conversations through
|
||||
`add()`. The agent still gains query-aware `memory_search`, while new
|
||||
conversations continue accumulating memory even though fact CRUD is not
|
||||
available.
|
||||
- Fact CRUD, `import_memory`, and Settings-page memory editing are not
|
||||
implemented (gateway returns 501). DeerMem remains the default backend.
|
||||
- No migration of existing DeerMem data.
|
||||
- `log_and_drop` write policy is at-most-once: a failed write is dropped.
|
||||
- `memory_add`/`memory_update`/`memory_delete` are backed by fact CRUD, which
|
||||
this backend does not implement; they return a clear unsupported-operation
|
||||
error. Conversation writes still happen through the retained middleware.
|
||||
|
||||
## Async execution and failure behavior
|
||||
|
||||
The mem0 HTTP client is synchronous for compatibility with the
|
||||
`MemoryManager` contract. DeerFlow offloads it at every async boundary: the
|
||||
async middleware uses the manager's `a*` methods, and Gateway memory routes run
|
||||
sync management calls in worker threads. A slow mem0 request therefore does
|
||||
not block unrelated ASGI handlers or SSE heartbeats.
|
||||
|
||||
`failure_policy.read: fail_open` logs a recall failure and continues without
|
||||
new memory context. `fail_closed` propagates the backend error through prompt
|
||||
construction and aborts the run instead of silently degrading.
|
||||
@ -0,0 +1,9 @@
|
||||
"""mem0 memory backend -- HTTP client against the mem0 Platform API.
|
||||
|
||||
Drop-in contract: folder name == backend name == ``manager_class: mem0``.
|
||||
"""
|
||||
|
||||
from .mem0_manager import Mem0Manager
|
||||
|
||||
#: Discovered by the factory's ``_scan_backends`` under the folder name ``mem0``.
|
||||
MANAGER_CLASS = Mem0Manager
|
||||
@ -0,0 +1,128 @@
|
||||
"""Synchronous httpx client for the mem0 REST API (v3; delete is v1).
|
||||
|
||||
The MemoryManager contract is synchronous (DeerMem's LLM calls are sync too),
|
||||
so this client is a plain ``httpx.Client``. It is constructed with an
|
||||
optional ``transport`` so tests can inject ``httpx.MockTransport``.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
import httpx
|
||||
|
||||
|
||||
class Mem0APIError(RuntimeError):
|
||||
"""Any mem0 request failure (transport, 4xx/5xx)."""
|
||||
|
||||
|
||||
class Mem0AuthError(Mem0APIError):
|
||||
"""401 -- missing or invalid API key."""
|
||||
|
||||
|
||||
class Mem0Client:
|
||||
"""Thin wrapper over the mem0 endpoints DeerFlow uses."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
base_url: str,
|
||||
api_key: str,
|
||||
timeout_seconds: float = 10.0,
|
||||
transport: httpx.BaseTransport | None = None,
|
||||
) -> None:
|
||||
self._http = httpx.Client(
|
||||
base_url=base_url.rstrip("/"),
|
||||
headers={"Authorization": f"Token {api_key}", "Accept": "application/json"},
|
||||
timeout=timeout_seconds,
|
||||
transport=transport,
|
||||
)
|
||||
|
||||
def close(self) -> None:
|
||||
self._http.close()
|
||||
|
||||
def _request(self, method: str, path: str, **kwargs: Any) -> dict[str, Any]:
|
||||
try:
|
||||
resp = self._http.request(method, path, **kwargs)
|
||||
except httpx.HTTPError as e:
|
||||
raise Mem0APIError(f"mem0 request failed: {e}") from e
|
||||
if resp.status_code == 401:
|
||||
raise Mem0AuthError("mem0 authentication failed (check the API key)")
|
||||
if resp.status_code >= 400:
|
||||
raise Mem0APIError(f"mem0 {method} {path} -> {resp.status_code}: {resp.text[:200]}")
|
||||
if not resp.content:
|
||||
return {}
|
||||
try:
|
||||
return resp.json()
|
||||
except json.JSONDecodeError as e:
|
||||
raise Mem0APIError(f"mem0 {method} {path} returned malformed JSON: {e}") from e
|
||||
|
||||
def add_memories(
|
||||
self,
|
||||
*,
|
||||
messages: list[dict[str, str]],
|
||||
user_id: str | None = None,
|
||||
agent_id: str | None = None,
|
||||
run_id: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Queue extraction (async server-side; response carries an event_id)."""
|
||||
body: dict[str, Any] = {"messages": messages}
|
||||
if user_id:
|
||||
body["user_id"] = user_id
|
||||
if agent_id:
|
||||
body["agent_id"] = agent_id
|
||||
if run_id:
|
||||
body["run_id"] = run_id
|
||||
return self._request("POST", "/v3/memories/add/", json=body)
|
||||
|
||||
def search_memories(
|
||||
self,
|
||||
*,
|
||||
query: str,
|
||||
filters: dict[str, Any],
|
||||
top_k: int,
|
||||
threshold: float,
|
||||
) -> list[dict[str, Any]]:
|
||||
body = {"query": query, "filters": filters, "top_k": top_k, "threshold": threshold}
|
||||
return self._request("POST", "/v3/memories/search/", json=body).get("results", [])
|
||||
|
||||
def list_memories(
|
||||
self,
|
||||
*,
|
||||
filters: dict[str, Any],
|
||||
page_size: int = 200,
|
||||
max_items: int | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
"""List memories across pages until exhausted or ``max_items`` reached."""
|
||||
results: list[dict[str, Any]] = []
|
||||
page = 1
|
||||
while True:
|
||||
data = self._request(
|
||||
"POST",
|
||||
"/v3/memories/",
|
||||
params={"page": page, "page_size": page_size},
|
||||
json={"filters": filters},
|
||||
)
|
||||
results.extend(data.get("results", []))
|
||||
if not data.get("next") or (max_items is not None and len(results) >= max_items):
|
||||
return results[:max_items] if max_items is not None else results
|
||||
page += 1
|
||||
|
||||
def delete_all_memories(
|
||||
self,
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
agent_id: str | None = None,
|
||||
run_id: str | None = None,
|
||||
) -> None:
|
||||
params = {k: v for k, v in {"user_id": user_id, "agent_id": agent_id, "run_id": run_id}.items() if v}
|
||||
self._request("DELETE", "/v1/memories/", params=params)
|
||||
|
||||
def ping(self) -> None:
|
||||
"""Startup auth check: a 1-item list scoped to a sentinel user id.
|
||||
|
||||
Proves the API key works without touching real data (the sentinel
|
||||
bucket is always empty).
|
||||
"""
|
||||
self.list_memories(filters={"user_id": "__deerflow_startup_check__"}, page_size=1, max_items=1)
|
||||
@ -0,0 +1,123 @@
|
||||
"""mem0 backend config -- parses and validates ``backend_config``.
|
||||
|
||||
Follows the noop-template pattern: a plain dataclass + ``from_backend_config``.
|
||||
The host injects ``storage_path`` (and optionally ``should_keep_hidden_message``)
|
||||
into every backend's config dict; those keys are accepted and ignored. Any
|
||||
OTHER unknown key is rejected -- a typo in persistent-state config must fail
|
||||
fast, not silently fall back to defaults.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
from urllib.parse import urlsplit
|
||||
|
||||
#: Keys the host factory injects into backend_config; accepted and ignored.
|
||||
_HOST_INJECTED_KEYS = frozenset({"storage_path", "should_keep_hidden_message"})
|
||||
|
||||
_STARTUP_POLICIES = frozenset({"fail_fast", "tolerate"})
|
||||
_READ_POLICIES = frozenset({"fail_open", "fail_closed"})
|
||||
_WRITE_POLICIES = frozenset({"log_and_drop", "raise"})
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Mem0Config:
|
||||
"""Validated knobs for the mem0 HTTP backend."""
|
||||
|
||||
#: Name of the environment variable holding the mem0 API key. The key
|
||||
#: itself never appears in config.yaml.
|
||||
api_key_env: str = "MEM0_API_KEY"
|
||||
#: mem0 Platform API root; point at a self-hosted server for on-prem.
|
||||
base_url: str = "https://api.mem0.ai"
|
||||
#: Permit sending the API token over plaintext HTTP. Intended only for
|
||||
#: trusted local development networks.
|
||||
allow_insecure_http: bool = False
|
||||
#: Max memories injected by get_context / default search breadth (1-1000).
|
||||
top_k: int = 8
|
||||
#: Minimum relevance score for search() results (mem0 `threshold`, 0-1).
|
||||
score_threshold: float = 0.1
|
||||
#: Hard cap on the injection text returned by get_context.
|
||||
max_injection_chars: int = 12000
|
||||
#: Per-request HTTP timeout in seconds.
|
||||
timeout_seconds: float = 10.0
|
||||
#: "fail_fast" = auth-check in from_config; "tolerate" = defer to first use.
|
||||
startup_policy: str = "fail_fast"
|
||||
#: "fail_open" = recall errors inject nothing and continue;
|
||||
#: "fail_closed" = recall errors raise MemoryManagerError.
|
||||
read_policy: str = "fail_open"
|
||||
#: "log_and_drop" = write errors are logged and dropped (at-most-once);
|
||||
#: "raise" = write errors raise MemoryManagerError.
|
||||
write_policy: str = "log_and_drop"
|
||||
|
||||
@classmethod
|
||||
def from_backend_config(cls, backend_config: dict[str, Any] | None) -> Mem0Config:
|
||||
cfg = dict(backend_config or {})
|
||||
failure_policy = cfg.pop("failure_policy", {}) or {}
|
||||
unknown = (
|
||||
set(cfg)
|
||||
- {
|
||||
"api_key_env",
|
||||
"base_url",
|
||||
"allow_insecure_http",
|
||||
"top_k",
|
||||
"score_threshold",
|
||||
"max_injection_chars",
|
||||
"timeout_seconds",
|
||||
"startup_policy",
|
||||
}
|
||||
- _HOST_INJECTED_KEYS
|
||||
)
|
||||
if unknown:
|
||||
raise ValueError(f"mem0 backend_config has unknown keys: {sorted(unknown)}")
|
||||
if not isinstance(failure_policy, dict):
|
||||
raise ValueError("mem0 failure_policy must be a mapping {read, write}")
|
||||
unknown_fp = set(failure_policy) - {"read", "write"}
|
||||
if unknown_fp:
|
||||
raise ValueError(f"mem0 failure_policy has unknown keys: {sorted(unknown_fp)}")
|
||||
allow_insecure_http = cfg.get("allow_insecure_http", False)
|
||||
if not isinstance(allow_insecure_http, bool):
|
||||
raise ValueError("mem0 allow_insecure_http must be a boolean")
|
||||
|
||||
config = cls(
|
||||
api_key_env=str(cfg.get("api_key_env", "MEM0_API_KEY")),
|
||||
base_url=str(cfg.get("base_url", "https://api.mem0.ai")).rstrip("/"),
|
||||
allow_insecure_http=allow_insecure_http,
|
||||
top_k=int(cfg.get("top_k", 8)),
|
||||
score_threshold=float(cfg.get("score_threshold", 0.1)),
|
||||
max_injection_chars=int(cfg.get("max_injection_chars", 12000)),
|
||||
timeout_seconds=float(cfg.get("timeout_seconds", 10.0)),
|
||||
startup_policy=str(cfg.get("startup_policy", "fail_fast")),
|
||||
read_policy=str(failure_policy.get("read", "fail_open")),
|
||||
write_policy=str(failure_policy.get("write", "log_and_drop")),
|
||||
)
|
||||
if config.startup_policy not in _STARTUP_POLICIES:
|
||||
raise ValueError(f"mem0 startup_policy must be one of {sorted(_STARTUP_POLICIES)}")
|
||||
if config.read_policy not in _READ_POLICIES:
|
||||
raise ValueError(f"mem0 failure_policy.read must be one of {sorted(_READ_POLICIES)}")
|
||||
if config.write_policy not in _WRITE_POLICIES:
|
||||
raise ValueError(f"mem0 failure_policy.write must be one of {sorted(_WRITE_POLICIES)}")
|
||||
if not 1 <= config.top_k <= 1000:
|
||||
raise ValueError("mem0 top_k must be in [1, 1000]")
|
||||
if not 0.0 <= config.score_threshold <= 1.0:
|
||||
raise ValueError("mem0 score_threshold must be in [0, 1]")
|
||||
if config.max_injection_chars <= 0:
|
||||
raise ValueError("mem0 max_injection_chars must be positive")
|
||||
if config.timeout_seconds <= 0:
|
||||
raise ValueError("mem0 timeout_seconds must be positive")
|
||||
if not config.api_key_env.strip():
|
||||
raise ValueError("mem0 api_key_env must be a non-empty env var name")
|
||||
parsed_base_url = urlsplit(config.base_url)
|
||||
if parsed_base_url.scheme not in {"http", "https"} or not parsed_base_url.netloc:
|
||||
raise ValueError("mem0 base_url must be an absolute http:// or https:// URL")
|
||||
if parsed_base_url.scheme == "http" and not config.allow_insecure_http:
|
||||
raise ValueError("mem0 base_url must use https:// because it carries the API key; set allow_insecure_http: true only for trusted local development")
|
||||
return config
|
||||
|
||||
def resolve_api_key(self) -> str:
|
||||
"""Read the API key from the configured environment variable."""
|
||||
key = os.environ.get(self.api_key_env, "").strip()
|
||||
if not key:
|
||||
raise ValueError(f"mem0 API key missing: environment variable {self.api_key_env} is unset or empty")
|
||||
return key
|
||||
@ -0,0 +1,315 @@
|
||||
"""mem0 memory backend -- a stateless HTTP MemoryManager.
|
||||
|
||||
All state lives server-side in mem0 (dedup, extraction, storage): this backend
|
||||
keeps no queue, watermark, or cache, so it is safe for multi-worker Gateway
|
||||
deployments. Identity maps 1:1: (user_id, agent_name) -> mem0 (user_id,
|
||||
agent_id); thread_id -> mem0 run_id.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import Any, ClassVar, Literal
|
||||
|
||||
from pydantic import PrivateAttr
|
||||
|
||||
# ABC contract -- the ONE allowed `from deerflow` import in this backend folder.
|
||||
from deerflow.agents.memory.manager import MemoryManager, MemoryManagerError
|
||||
|
||||
from .client import Mem0APIError, Mem0Client
|
||||
from .config import Mem0Config
|
||||
from .message_filtering import extract_message_text, filter_messages_for_memory
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_ROLE_MAP = {"human": "user", "ai": "assistant"}
|
||||
|
||||
|
||||
def _build_filters(
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
agent_name: str | None = None,
|
||||
run_id: str | None = None,
|
||||
) -> dict[str, Any] | None:
|
||||
"""Build a mem0 ``filters`` object from the available identity parts.
|
||||
|
||||
Returns None when no entity id is available (mem0 requires at least one).
|
||||
"""
|
||||
parts: list[dict[str, Any]] = []
|
||||
if user_id:
|
||||
parts.append({"user_id": user_id})
|
||||
if agent_name:
|
||||
parts.append({"agent_id": agent_name})
|
||||
if run_id:
|
||||
parts.append({"run_id": run_id})
|
||||
if not parts:
|
||||
return None
|
||||
if len(parts) == 1:
|
||||
return parts[0]
|
||||
return {"AND": parts}
|
||||
|
||||
|
||||
def _to_fact(record: dict[str, Any]) -> dict[str, Any]:
|
||||
"""Map a mem0 record to the backend-neutral fact shape consumed by the
|
||||
host (agents/memory/tools.py): id/content/category/confidence/createdAt/
|
||||
source. mem0's relevance score doubles as confidence."""
|
||||
categories = record.get("categories") or []
|
||||
metadata = record.get("metadata") or {}
|
||||
return {
|
||||
"id": str(record.get("id", "")),
|
||||
"content": str(record.get("memory", "")),
|
||||
"category": str(categories[0]) if categories else "context",
|
||||
"confidence": float(record.get("score") or 0.0),
|
||||
"createdAt": str(record.get("created_at", "")),
|
||||
"source": str(metadata.get("source", "")),
|
||||
}
|
||||
|
||||
|
||||
class Mem0Manager(MemoryManager):
|
||||
"""MemoryManager backed by the mem0 Platform API (or compatible server)."""
|
||||
|
||||
# search() is overridden below -> flag must be True (contract invariant);
|
||||
# this also enables memory mode="tool".
|
||||
supports_search: ClassVar[bool] = True
|
||||
# mem0 extracts/deduplicates facts from full conversations through add();
|
||||
# its fact CRUD hooks are intentionally unsupported, so tool mode retains
|
||||
# passive writes while exposing query-aware search.
|
||||
requires_passive_writes_in_tool_mode: ClassVar[bool] = True
|
||||
|
||||
_config: Mem0Config = PrivateAttr()
|
||||
_client: Any = PrivateAttr(default=None) # Mem0Client; tests inject a fake
|
||||
|
||||
def model_post_init(self, __context: Any) -> None:
|
||||
self._config = Mem0Config.from_backend_config(self.backend_config)
|
||||
self._client = Mem0Client(
|
||||
base_url=self._config.base_url,
|
||||
api_key=self._config.resolve_api_key(),
|
||||
timeout_seconds=self._config.timeout_seconds,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
def from_config(
|
||||
cls,
|
||||
backend_config: dict[str, Any] | None = None,
|
||||
*,
|
||||
mode: Literal["middleware", "tool"] = "middleware",
|
||||
**host_hooks: Any,
|
||||
) -> Mem0Manager:
|
||||
"""Build the manager; ``fail_fast`` startup policy auth-checks via ping."""
|
||||
mgr = cls(backend_config=backend_config, mode=mode)
|
||||
if mgr._config.startup_policy == "fail_fast":
|
||||
mgr._client.ping()
|
||||
return mgr
|
||||
|
||||
def close(self) -> None:
|
||||
"""Release the underlying HTTP connection pool."""
|
||||
self._client.close()
|
||||
|
||||
# ── Error policies ───────────────────────────────────────────────────
|
||||
def _read_or_fallback(self, fallback: Any, fn: Any) -> Any:
|
||||
try:
|
||||
return fn()
|
||||
except Mem0APIError as e:
|
||||
if self._config.read_policy == "fail_open":
|
||||
logger.warning("mem0 read failed (%s); continuing without memory", e)
|
||||
return fallback
|
||||
raise MemoryManagerError(f"mem0 read failed: {e}") from e
|
||||
|
||||
def _write_or_drop(self, fn: Any) -> None:
|
||||
try:
|
||||
fn()
|
||||
except Mem0APIError as e:
|
||||
if self._config.write_policy == "log_and_drop":
|
||||
logger.warning("mem0 write failed (%s); dropping update", e)
|
||||
return
|
||||
raise MemoryManagerError(f"mem0 write failed: {e}") from e
|
||||
|
||||
# ── Tier 1: write ────────────────────────────────────────────────────
|
||||
def add(
|
||||
self,
|
||||
thread_id: str,
|
||||
messages: list[Any],
|
||||
*,
|
||||
agent_name: str | None = None,
|
||||
user_id: str | None = None,
|
||||
trace_id: str | None = None,
|
||||
) -> None:
|
||||
"""Submit the filtered conversation to mem0 for server-side extraction.
|
||||
|
||||
Fire-and-forget: mem0 processes asynchronously (response event_id is
|
||||
not polled). ``thread_id`` maps to mem0 ``run_id`` and always satisfies
|
||||
mem0's at-least-one-entity-id requirement.
|
||||
"""
|
||||
kept = filter_messages_for_memory(messages)
|
||||
payload = [{"role": _ROLE_MAP[getattr(m, "type", "")], "content": extract_message_text(m).strip()} for m in kept if getattr(m, "type", "") in _ROLE_MAP]
|
||||
payload = [p for p in payload if p["content"]]
|
||||
if not payload:
|
||||
return
|
||||
self._write_or_drop(
|
||||
lambda: self._client.add_memories(
|
||||
messages=payload,
|
||||
user_id=user_id,
|
||||
agent_id=agent_name,
|
||||
run_id=thread_id,
|
||||
)
|
||||
)
|
||||
|
||||
async def aadd(
|
||||
self,
|
||||
thread_id: str,
|
||||
messages: list[Any],
|
||||
*,
|
||||
agent_name: str | None = None,
|
||||
user_id: str | None = None,
|
||||
trace_id: str | None = None,
|
||||
) -> None:
|
||||
await asyncio.to_thread(
|
||||
self.add,
|
||||
thread_id,
|
||||
messages,
|
||||
agent_name=agent_name,
|
||||
user_id=user_id,
|
||||
trace_id=trace_id,
|
||||
)
|
||||
|
||||
# ── Tier 1: read-inject ──────────────────────────────────────────────
|
||||
def get_context(
|
||||
self,
|
||||
user_id: str | None,
|
||||
*,
|
||||
agent_name: str | None = None,
|
||||
thread_id: str | None = None,
|
||||
) -> str:
|
||||
"""Query-less recall: the contract passes no current query, so inject
|
||||
the bucket's most recent memories (top_k). Query-aware recall is
|
||||
available via search() in mode="tool"."""
|
||||
filters = _build_filters(user_id=user_id, agent_name=agent_name, run_id=thread_id)
|
||||
if filters is None:
|
||||
return ""
|
||||
top_k = self._config.top_k
|
||||
records = self._read_or_fallback(
|
||||
[],
|
||||
lambda: self._client.list_memories(
|
||||
filters=filters,
|
||||
page_size=min(top_k, 200),
|
||||
max_items=top_k,
|
||||
),
|
||||
)
|
||||
seen: set[str] = set()
|
||||
lines: list[str] = []
|
||||
for record in records:
|
||||
rid = record.get("id")
|
||||
if rid in seen:
|
||||
continue
|
||||
seen.add(rid)
|
||||
text = str(record.get("memory") or "").strip()
|
||||
if text:
|
||||
lines.append(f"- {text}")
|
||||
context = "\n".join(lines)
|
||||
if len(context) > self._config.max_injection_chars:
|
||||
context = context[: self._config.max_injection_chars]
|
||||
return context
|
||||
|
||||
async def aget_context(
|
||||
self,
|
||||
user_id: str | None,
|
||||
*,
|
||||
agent_name: str | None = None,
|
||||
thread_id: str | None = None,
|
||||
) -> str:
|
||||
return await asyncio.to_thread(
|
||||
self.get_context,
|
||||
user_id,
|
||||
agent_name=agent_name,
|
||||
thread_id=thread_id,
|
||||
)
|
||||
|
||||
# ── Tier 2: search ───────────────────────────────────────────────────
|
||||
def search(
|
||||
self,
|
||||
query: str,
|
||||
top_k: int = 5,
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
agent_name: str | None = None,
|
||||
category: str | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
filters = _build_filters(user_id=user_id, agent_name=agent_name)
|
||||
if filters is None:
|
||||
return []
|
||||
if category:
|
||||
parts = filters["AND"] if "AND" in filters else [filters]
|
||||
filters = {"AND": [*parts, {"categories": {"contains": category}}]}
|
||||
results = self._read_or_fallback(
|
||||
[],
|
||||
lambda: self._client.search_memories(
|
||||
query=query,
|
||||
filters=filters,
|
||||
top_k=top_k,
|
||||
threshold=self._config.score_threshold,
|
||||
),
|
||||
)
|
||||
return [_to_fact(r) for r in results]
|
||||
|
||||
async def asearch(
|
||||
self,
|
||||
query: str,
|
||||
top_k: int = 5,
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
agent_name: str | None = None,
|
||||
category: str | None = None,
|
||||
) -> list[dict[str, Any]]:
|
||||
return await asyncio.to_thread(
|
||||
self.search,
|
||||
query,
|
||||
top_k,
|
||||
user_id=user_id,
|
||||
agent_name=agent_name,
|
||||
category=category,
|
||||
)
|
||||
|
||||
# ── Tier 2: management ───────────────────────────────────────────────
|
||||
def get_memory(
|
||||
self,
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
agent_name: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
filters = _build_filters(user_id=user_id, agent_name=agent_name)
|
||||
if filters is None:
|
||||
return {"facts": []}
|
||||
records = self._read_or_fallback([], lambda: self._client.list_memories(filters=filters))
|
||||
return {"facts": [_to_fact(r) for r in records]}
|
||||
|
||||
def export_memory(
|
||||
self,
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
agent_name: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
return self.get_memory(user_id=user_id, agent_name=agent_name)
|
||||
|
||||
def clear_memory(
|
||||
self,
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
agent_name: str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""Clear the bucket. agent_name=None clears the user's whole memory;
|
||||
an explicit agent clears only that agent's bucket."""
|
||||
if not user_id and not agent_name:
|
||||
return {"facts": []}
|
||||
self._write_or_drop(lambda: self._client.delete_all_memories(user_id=user_id, agent_id=agent_name, run_id=None))
|
||||
return {"facts": []}
|
||||
|
||||
def delete_memory(
|
||||
self,
|
||||
*,
|
||||
user_id: str | None = None,
|
||||
agent_name: str | None = None,
|
||||
) -> None:
|
||||
if not user_id and not agent_name:
|
||||
return None
|
||||
self._write_or_drop(lambda: self._client.delete_all_memories(user_id=user_id, agent_id=agent_name, run_id=None))
|
||||
@ -0,0 +1,93 @@
|
||||
"""Message filtering for the mem0 write path -- self-contained mirror of
|
||||
DeerMem's ``filter_messages_for_memory`` rules (the portability rule forbids
|
||||
importing across backend folders, so the logic is duplicated, not shared).
|
||||
|
||||
Keeps: visible user inputs, well-formed human clarification answers, and final
|
||||
assistant responses. Drops: framework-internal ``hide_from_ui`` messages,
|
||||
tool-call AI messages, tool outputs, empty/upload-only turns.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from collections.abc import Mapping
|
||||
from copy import copy
|
||||
from typing import Any
|
||||
|
||||
_UPLOAD_BLOCK_RE = re.compile(r"<(?P<tag>uploaded_files|current_uploads)>[\s\S]*?</(?P=tag)>\n*", re.IGNORECASE)
|
||||
|
||||
|
||||
def extract_message_text(message: Any) -> str:
|
||||
"""Extract plain text from message content (str or content-block list)."""
|
||||
content = getattr(message, "content", "")
|
||||
if content is None:
|
||||
return ""
|
||||
if isinstance(content, list):
|
||||
parts: list[str] = []
|
||||
for part in content:
|
||||
if isinstance(part, str):
|
||||
parts.append(part)
|
||||
elif isinstance(part, dict):
|
||||
text = part.get("text")
|
||||
if isinstance(text, str):
|
||||
parts.append(text)
|
||||
return " ".join(parts)
|
||||
return str(content)
|
||||
|
||||
|
||||
def _non_empty_str(value: object) -> str | None:
|
||||
return value if isinstance(value, str) and value.strip() else None
|
||||
|
||||
|
||||
def _is_human_clarification_response(additional_kwargs: Any) -> bool:
|
||||
"""Structural check for a user-authored clarification answer carried in a
|
||||
hidden message (mirrors DeerMem's host-agnostic fallback)."""
|
||||
if not isinstance(additional_kwargs, Mapping):
|
||||
return False
|
||||
raw = additional_kwargs.get("human_input_response")
|
||||
if not isinstance(raw, Mapping):
|
||||
return False
|
||||
if raw.get("version") != 1 or raw.get("kind") != "human_input_response":
|
||||
return False
|
||||
if _non_empty_str(raw.get("source")) is None or _non_empty_str(raw.get("request_id")) is None or _non_empty_str(raw.get("value")) is None:
|
||||
return False
|
||||
response_kind = raw.get("response_kind")
|
||||
if response_kind == "text":
|
||||
return True
|
||||
if response_kind == "option":
|
||||
return _non_empty_str(raw.get("option_id")) is not None
|
||||
return False
|
||||
|
||||
|
||||
def filter_messages_for_memory(messages: list[Any]) -> list[Any]:
|
||||
"""Keep only user inputs and final assistant responses."""
|
||||
filtered: list[Any] = []
|
||||
skip_next_ai = False
|
||||
for msg in messages:
|
||||
msg_type = getattr(msg, "type", None)
|
||||
if msg_type == "human":
|
||||
additional_kwargs = getattr(msg, "additional_kwargs", {}) or {}
|
||||
if additional_kwargs.get("hide_from_ui") and not _is_human_clarification_response(additional_kwargs):
|
||||
continue
|
||||
text = extract_message_text(msg)
|
||||
if "<uploaded_files>" in text.lower() or "<current_uploads>" in text.lower():
|
||||
stripped = _UPLOAD_BLOCK_RE.sub("", text).strip()
|
||||
if not stripped:
|
||||
# Upload-only turn: the following AI ack carries no user content.
|
||||
skip_next_ai = True
|
||||
continue
|
||||
clean_msg = copy(msg)
|
||||
clean_msg.content = stripped
|
||||
filtered.append(clean_msg)
|
||||
skip_next_ai = False
|
||||
else:
|
||||
filtered.append(msg)
|
||||
skip_next_ai = False
|
||||
elif msg_type == "ai":
|
||||
if getattr(msg, "tool_calls", None):
|
||||
continue
|
||||
if skip_next_ai:
|
||||
skip_next_ai = False
|
||||
continue
|
||||
filtered.append(msg)
|
||||
return filtered
|
||||
@ -143,6 +143,10 @@ class MemoryManager(BaseModel):
|
||||
# that fails fast at instantiation rather than silently returning empty
|
||||
# results). Default False: a new backend must explicitly opt in to tool mode.
|
||||
supports_search: ClassVar[bool] = False
|
||||
# Backends that rely on conversation-level extraction instead of fact CRUD
|
||||
# can retain MemoryMiddleware writes while tool mode supplies query-aware
|
||||
# search. Most backends keep tool mode fully model-directed.
|
||||
requires_passive_writes_in_tool_mode: ClassVar[bool] = False
|
||||
|
||||
@model_validator(mode="after")
|
||||
def _check_invariants(self) -> MemoryManager:
|
||||
@ -585,6 +589,15 @@ def _resolve_manager_class(manager_class: str) -> type[MemoryManager]:
|
||||
)
|
||||
|
||||
|
||||
def backend_requires_passive_writes_in_tool_mode(manager_class: str) -> bool:
|
||||
"""Return whether a backend needs middleware writes in tool mode.
|
||||
|
||||
Resolve the class without constructing it so agent assembly does not run
|
||||
backend startup checks or perform network I/O.
|
||||
"""
|
||||
return _resolve_manager_class(manager_class).requires_passive_writes_in_tool_mode
|
||||
|
||||
|
||||
# ── Host-default hook providers (passed to from_config by the factory) ────
|
||||
#
|
||||
# These callables are the host's defaults for the slots a backend may consume
|
||||
|
||||
@ -3,10 +3,10 @@
|
||||
Exposes memory_search, memory_add, memory_update, memory_delete as
|
||||
LangChain @tool functions the model can call directly.
|
||||
|
||||
When memory.mode == "tool", these tools are registered on the agent
|
||||
instead of appending MemoryMiddleware. The model gains agency over
|
||||
its own persistent memory: it decides what to remember, when to
|
||||
search, and when to update or remove stale facts.
|
||||
When memory.mode == "tool", these tools are registered on the agent. Most
|
||||
backends omit MemoryMiddleware so the model drives persistence; a backend that
|
||||
sets ``requires_passive_writes_in_tool_mode`` retains conversation writes while
|
||||
the tools provide query-aware recall.
|
||||
|
||||
Backend-agnostic: every tool goes through the ``MemoryManager`` ABC
|
||||
(:func:`get_memory_manager`) -- ``search``/``get_memory`` are tier-2 methods;
|
||||
|
||||
@ -1,5 +1,6 @@
|
||||
"""Middleware for memory mechanism."""
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
from typing import TYPE_CHECKING, override
|
||||
|
||||
@ -49,17 +50,8 @@ class MemoryMiddleware(AgentMiddleware[MemoryMiddlewareState]):
|
||||
self._agent_name = agent_name
|
||||
self._memory_config = memory_config
|
||||
|
||||
@override
|
||||
def after_agent(self, state: MemoryMiddlewareState, runtime: Runtime) -> dict | None:
|
||||
"""Queue conversation for memory update after agent completes.
|
||||
|
||||
Args:
|
||||
state: The current agent state.
|
||||
runtime: The runtime context.
|
||||
|
||||
Returns:
|
||||
None (no state changes needed from this middleware).
|
||||
"""
|
||||
def _resolve_add_args(self, state: MemoryMiddlewareState, runtime: Runtime) -> tuple[str, list, str, str | None] | None:
|
||||
"""Resolve one write request without invoking the manager."""
|
||||
config = self._memory_config or get_memory_config()
|
||||
if not config.enabled:
|
||||
return None
|
||||
@ -95,6 +87,16 @@ class MemoryMiddleware(AgentMiddleware[MemoryMiddlewareState]):
|
||||
if trace_id is None:
|
||||
trace_id = get_current_trace_id()
|
||||
|
||||
return thread_id, messages, user_id, trace_id
|
||||
|
||||
@override
|
||||
def after_agent(self, state: MemoryMiddlewareState, runtime: Runtime) -> dict | None:
|
||||
"""Queue conversation for memory update after agent completes."""
|
||||
add_args = self._resolve_add_args(state, runtime)
|
||||
if add_args is None:
|
||||
return None
|
||||
thread_id, messages, user_id, trace_id = add_args
|
||||
|
||||
# Hand raw messages to the manager; the backend filters to user + final-AI
|
||||
# turns, validates, detects correction/reinforcement, and enqueues.
|
||||
get_memory_manager().add(
|
||||
@ -106,3 +108,20 @@ class MemoryMiddleware(AgentMiddleware[MemoryMiddlewareState]):
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
@override
|
||||
async def aafter_agent(self, state: MemoryMiddlewareState, runtime: Runtime) -> dict | None:
|
||||
"""Use the manager's async boundary on LangGraph's async execution path."""
|
||||
add_args = self._resolve_add_args(state, runtime)
|
||||
if add_args is None:
|
||||
return None
|
||||
thread_id, messages, user_id, trace_id = add_args
|
||||
manager = await asyncio.to_thread(get_memory_manager)
|
||||
await manager.aadd(
|
||||
thread_id,
|
||||
messages,
|
||||
agent_name=self._agent_name,
|
||||
user_id=user_id,
|
||||
trace_id=trace_id,
|
||||
)
|
||||
return None
|
||||
|
||||
@ -2,10 +2,12 @@
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import inspect
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
from unittest.mock import MagicMock
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from langchain.agents import create_agent
|
||||
@ -1184,6 +1186,23 @@ def test_memory_middleware_uses_explicit_memory_config_without_global_read(monke
|
||||
assert middleware.after_agent({"messages": []}, runtime=MagicMock(context={"thread_id": "thread-1"})) is None
|
||||
|
||||
|
||||
def test_memory_middleware_async_path_uses_async_manager_call(monkeypatch):
|
||||
from deerflow.agents.middlewares import memory_middleware as memory_middleware_module
|
||||
from deerflow.agents.middlewares.memory_middleware import MemoryMiddleware
|
||||
|
||||
manager = SimpleNamespace(aadd=AsyncMock(), add=MagicMock(side_effect=AssertionError("sync add must not run")))
|
||||
monkeypatch.setattr(memory_middleware_module, "get_memory_manager", lambda: manager)
|
||||
middleware = MemoryMiddleware(memory_config=MemoryConfig(enabled=True))
|
||||
runtime = MagicMock(context={"thread_id": "thread-1", "user_id": "user-1"})
|
||||
|
||||
result = asyncio.run(middleware.aafter_agent({"messages": [HumanMessage(content="hello")]}, runtime=runtime))
|
||||
|
||||
assert result is None
|
||||
manager.aadd.assert_awaited_once()
|
||||
assert manager.aadd.await_args.kwargs["user_id"] == "user-1"
|
||||
manager.add.assert_not_called()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Per-agent model settings (issue #4336)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
@ -4,6 +4,7 @@ from types import SimpleNamespace
|
||||
from typing import cast
|
||||
|
||||
import anyio
|
||||
import pytest
|
||||
|
||||
from deerflow.agents.lead_agent import prompt as prompt_module
|
||||
from deerflow.config.app_config import AppConfig
|
||||
@ -349,6 +350,37 @@ def test_get_memory_context_uses_explicit_app_config_without_global_config(monke
|
||||
}
|
||||
|
||||
|
||||
def test_get_memory_context_propagates_fail_closed_manager_error(monkeypatch):
|
||||
from deerflow.agents.memory import MemoryManagerError
|
||||
|
||||
explicit_config = SimpleNamespace(
|
||||
memory=SimpleNamespace(
|
||||
enabled=True,
|
||||
injection_enabled=True,
|
||||
backend_config={"failure_policy": {"read": "fail_closed"}},
|
||||
),
|
||||
)
|
||||
manager = SimpleNamespace(get_context=lambda *args, **kwargs: (_ for _ in ()).throw(MemoryManagerError("down")))
|
||||
monkeypatch.setattr("deerflow.agents.memory.get_memory_manager", lambda: manager)
|
||||
monkeypatch.setattr("deerflow.runtime.user_context.get_effective_user_id", lambda: "user-1")
|
||||
|
||||
with pytest.raises(MemoryManagerError, match="down"):
|
||||
prompt_module._get_memory_context("agent-a", app_config=explicit_config)
|
||||
|
||||
|
||||
def test_get_memory_context_swallows_manager_error_without_fail_closed(monkeypatch):
|
||||
from deerflow.agents.memory import MemoryManagerError
|
||||
|
||||
explicit_config = SimpleNamespace(
|
||||
memory=SimpleNamespace(enabled=True, injection_enabled=True, backend_config={}),
|
||||
)
|
||||
manager = SimpleNamespace(get_context=lambda *args, **kwargs: (_ for _ in ()).throw(MemoryManagerError("down")))
|
||||
monkeypatch.setattr("deerflow.agents.memory.get_memory_manager", lambda: manager)
|
||||
monkeypatch.setattr("deerflow.runtime.user_context.get_effective_user_id", lambda: "user-1")
|
||||
|
||||
assert prompt_module._get_memory_context("agent-a", app_config=explicit_config) == ""
|
||||
|
||||
|
||||
def test_get_memory_context_prefers_explicit_user_id(monkeypatch):
|
||||
explicit_config = SimpleNamespace(
|
||||
memory=SimpleNamespace(enabled=True, injection_enabled=True),
|
||||
|
||||
630
backend/tests/test_mem0_memory_backend.py
Normal file
630
backend/tests/test_mem0_memory_backend.py
Normal file
@ -0,0 +1,630 @@
|
||||
"""Unit tests for the mem0 HTTP memory backend (backends/mem0/)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import threading
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from deerflow.agents.memory.backends.mem0.client import Mem0APIError, Mem0AuthError, Mem0Client
|
||||
from deerflow.agents.memory.backends.mem0.config import Mem0Config
|
||||
|
||||
|
||||
class TestMem0Config:
|
||||
def test_defaults(self) -> None:
|
||||
cfg = Mem0Config.from_backend_config({})
|
||||
assert cfg.api_key_env == "MEM0_API_KEY"
|
||||
assert cfg.base_url == "https://api.mem0.ai"
|
||||
assert cfg.allow_insecure_http is False
|
||||
assert cfg.top_k == 8
|
||||
assert cfg.score_threshold == 0.1
|
||||
assert cfg.max_injection_chars == 12000
|
||||
assert cfg.timeout_seconds == 10.0
|
||||
assert cfg.startup_policy == "fail_fast"
|
||||
assert cfg.read_policy == "fail_open"
|
||||
assert cfg.write_policy == "log_and_drop"
|
||||
|
||||
def test_custom_values_and_nested_failure_policy(self) -> None:
|
||||
cfg = Mem0Config.from_backend_config(
|
||||
{
|
||||
"api_key_env": "MY_MEM0_KEY",
|
||||
"base_url": "http://mem0.local:8888/",
|
||||
"allow_insecure_http": True,
|
||||
"top_k": 5,
|
||||
"score_threshold": 0.3,
|
||||
"max_injection_chars": 4000,
|
||||
"timeout_seconds": 3,
|
||||
"startup_policy": "tolerate",
|
||||
"failure_policy": {"read": "fail_closed", "write": "raise"},
|
||||
}
|
||||
)
|
||||
assert cfg.api_key_env == "MY_MEM0_KEY"
|
||||
assert cfg.base_url == "http://mem0.local:8888" # trailing slash stripped
|
||||
assert cfg.allow_insecure_http is True
|
||||
assert cfg.top_k == 5
|
||||
assert cfg.score_threshold == 0.3
|
||||
assert cfg.max_injection_chars == 4000
|
||||
assert cfg.timeout_seconds == 3.0
|
||||
assert cfg.startup_policy == "tolerate"
|
||||
assert cfg.read_policy == "fail_closed"
|
||||
assert cfg.write_policy == "raise"
|
||||
|
||||
def test_unknown_keys_rejected_except_host_injected(self) -> None:
|
||||
with pytest.raises(ValueError, match="unknown"):
|
||||
Mem0Config.from_backend_config({"typo_knob": 1})
|
||||
# Host-injected keys must be tolerated (factory injects storage_path
|
||||
# into every backend's backend_config).
|
||||
cfg = Mem0Config.from_backend_config({"storage_path": "/tmp/x", "should_keep_hidden_message": None})
|
||||
assert cfg.base_url == "https://api.mem0.ai"
|
||||
|
||||
def test_insecure_http_requires_explicit_opt_in(self) -> None:
|
||||
with pytest.raises(ValueError, match="allow_insecure_http"):
|
||||
Mem0Config.from_backend_config({"base_url": "http://mem0.local:8888"})
|
||||
|
||||
@pytest.mark.parametrize("base_url", ["mem0.local:8888", "ftp://mem0.local", "https:///missing-host"])
|
||||
def test_invalid_base_url_rejected(self, base_url: str) -> None:
|
||||
with pytest.raises(ValueError, match="base_url"):
|
||||
Mem0Config.from_backend_config({"base_url": base_url, "allow_insecure_http": True})
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("key", "value"),
|
||||
[
|
||||
("startup_policy", "sometimes"),
|
||||
("top_k", 0),
|
||||
("top_k", 1001),
|
||||
("score_threshold", 1.5),
|
||||
("max_injection_chars", 0),
|
||||
("timeout_seconds", 0),
|
||||
],
|
||||
)
|
||||
def test_invalid_values_rejected(self, key: str, value: object) -> None:
|
||||
with pytest.raises(ValueError):
|
||||
Mem0Config.from_backend_config({key: value})
|
||||
|
||||
@pytest.mark.parametrize("policy", ["read", "write"])
|
||||
def test_invalid_failure_policy_rejected(self, policy: str) -> None:
|
||||
with pytest.raises(ValueError, match=policy):
|
||||
Mem0Config.from_backend_config({"failure_policy": {policy: "bogus"}})
|
||||
|
||||
def test_resolve_api_key(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
cfg = Mem0Config.from_backend_config({})
|
||||
monkeypatch.delenv("MEM0_API_KEY", raising=False)
|
||||
with pytest.raises(ValueError, match="MEM0_API_KEY"):
|
||||
cfg.resolve_api_key()
|
||||
monkeypatch.setenv("MEM0_API_KEY", " ")
|
||||
with pytest.raises(ValueError, match="MEM0_API_KEY"):
|
||||
cfg.resolve_api_key()
|
||||
monkeypatch.setenv("MEM0_API_KEY", "secret-key")
|
||||
assert cfg.resolve_api_key() == "secret-key"
|
||||
|
||||
|
||||
def _client(handler) -> Mem0Client:
|
||||
return Mem0Client(
|
||||
base_url="https://api.mem0.ai",
|
||||
api_key="test-key",
|
||||
transport=httpx.MockTransport(handler),
|
||||
)
|
||||
|
||||
|
||||
class TestMem0Client:
|
||||
def test_add_memories_payload(self) -> None:
|
||||
seen = {}
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
seen["path"] = request.url.path
|
||||
seen["auth"] = request.headers["authorization"]
|
||||
seen["body"] = httpx.QueryParams # placeholder, replaced below
|
||||
import json
|
||||
|
||||
seen["body"] = json.loads(request.content)
|
||||
return httpx.Response(200, json={"status": "PENDING", "event_id": "evt-1"})
|
||||
|
||||
client = _client(handler)
|
||||
result = client.add_memories(
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
user_id="u1",
|
||||
agent_id="lead_agent",
|
||||
run_id="t-1",
|
||||
)
|
||||
assert result["event_id"] == "evt-1"
|
||||
assert seen["path"] == "/v3/memories/add/"
|
||||
assert seen["auth"] == "Token test-key"
|
||||
assert seen["body"] == {
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"user_id": "u1",
|
||||
"agent_id": "lead_agent",
|
||||
"run_id": "t-1",
|
||||
}
|
||||
|
||||
def test_search_memories_returns_results(self) -> None:
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
import json
|
||||
|
||||
body = json.loads(request.content)
|
||||
assert request.url.path == "/v3/memories/search/"
|
||||
assert body == {
|
||||
"query": "hobbies",
|
||||
"filters": {"user_id": "u1"},
|
||||
"top_k": 5,
|
||||
"threshold": 0.2,
|
||||
}
|
||||
return httpx.Response(200, json={"results": [{"id": "m1", "memory": "likes cricket", "score": 0.9}]})
|
||||
|
||||
results = _client(handler).search_memories(query="hobbies", filters={"user_id": "u1"}, top_k=5, threshold=0.2)
|
||||
assert results == [{"id": "m1", "memory": "likes cricket", "score": 0.9}]
|
||||
|
||||
def test_list_memories_paginates_and_respects_max_items(self) -> None:
|
||||
pages = {
|
||||
1: {"results": [{"id": "a"}, {"id": "b"}], "next": "https://x/?page=2"},
|
||||
2: {"results": [{"id": "c"}], "next": None},
|
||||
}
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
page = int(request.url.params["page"])
|
||||
return httpx.Response(200, json=pages[page])
|
||||
|
||||
client = _client(handler)
|
||||
assert client.list_memories(filters={"user_id": "u1"}) == [{"id": "a"}, {"id": "b"}, {"id": "c"}]
|
||||
assert client.list_memories(filters={"user_id": "u1"}, max_items=2) == [{"id": "a"}, {"id": "b"}]
|
||||
|
||||
def test_delete_all_memories_uses_query_params(self) -> None:
|
||||
seen = {}
|
||||
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
seen["method"] = request.method
|
||||
seen["path"] = request.url.path
|
||||
seen["params"] = dict(request.url.params)
|
||||
return httpx.Response(200, json={"message": "deleted"})
|
||||
|
||||
_client(handler).delete_all_memories(user_id="u1", agent_id="lead_agent", run_id=None)
|
||||
assert seen == {
|
||||
"method": "DELETE",
|
||||
"path": "/v1/memories/",
|
||||
"params": {"user_id": "u1", "agent_id": "lead_agent"},
|
||||
}
|
||||
|
||||
def test_401_raises_auth_error(self) -> None:
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(401, json={"detail": "invalid key"})
|
||||
|
||||
with pytest.raises(Mem0AuthError):
|
||||
_client(handler).ping()
|
||||
|
||||
def test_other_4xx_raises_api_error(self) -> None:
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(400, json={"error": "bad request"})
|
||||
|
||||
with pytest.raises(Mem0APIError, match="400"):
|
||||
_client(handler).list_memories(filters={"user_id": "u1"})
|
||||
|
||||
def test_transport_error_raises_api_error(self) -> None:
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
raise httpx.ConnectError("boom")
|
||||
|
||||
with pytest.raises(Mem0APIError, match="boom"):
|
||||
_client(handler).ping()
|
||||
|
||||
def test_malformed_json_raises_api_error(self) -> None:
|
||||
def handler(request: httpx.Request) -> httpx.Response:
|
||||
return httpx.Response(200, content=b"not json{")
|
||||
|
||||
with pytest.raises(Mem0APIError, match="malformed JSON"):
|
||||
_client(handler).list_memories(filters={"user_id": "u1"})
|
||||
|
||||
|
||||
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage # noqa: E402
|
||||
|
||||
from deerflow.agents.memory.backends.mem0.message_filtering import ( # noqa: E402
|
||||
extract_message_text,
|
||||
filter_messages_for_memory,
|
||||
)
|
||||
|
||||
|
||||
def _clarification_kwargs() -> dict:
|
||||
return {
|
||||
"hide_from_ui": True,
|
||||
"human_input_response": {
|
||||
"version": 1,
|
||||
"kind": "human_input_response",
|
||||
"source": "clarification",
|
||||
"request_id": "req-1",
|
||||
"response_kind": "text",
|
||||
"value": "the user answered this",
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
class TestMessageFiltering:
|
||||
def test_keeps_user_and_final_assistant(self) -> None:
|
||||
msgs = [HumanMessage(content="hello"), AIMessage(content="hi there")]
|
||||
assert filter_messages_for_memory(msgs) == msgs
|
||||
|
||||
def test_drops_tool_messages_and_tool_call_ai(self) -> None:
|
||||
tool_ai = AIMessage(content="", tool_calls=[{"name": "t", "args": {}, "id": "1"}])
|
||||
msgs = [HumanMessage(content="q"), tool_ai, ToolMessage(content="out", tool_call_id="1"), AIMessage(content="a")]
|
||||
assert filter_messages_for_memory(msgs) == [msgs[0], msgs[3]]
|
||||
|
||||
def test_drops_hidden_framework_messages_keeps_clarification(self) -> None:
|
||||
hidden = HumanMessage(content="todo reminder", additional_kwargs={"hide_from_ui": True})
|
||||
clarification = HumanMessage(content="the user answered this", additional_kwargs=_clarification_kwargs())
|
||||
assert filter_messages_for_memory([hidden, clarification]) == [clarification]
|
||||
|
||||
def test_upload_only_human_drops_it_and_following_ai(self) -> None:
|
||||
upload_only = HumanMessage(content="<uploaded_files>\nfile.pdf\n</uploaded_files>")
|
||||
ack = AIMessage(content="I see your file")
|
||||
followup = HumanMessage(content="what is in it?")
|
||||
assert filter_messages_for_memory([upload_only, ack, followup]) == [followup]
|
||||
|
||||
def test_upload_block_stripped_from_mixed_message(self) -> None:
|
||||
mixed = HumanMessage(content="<current_uploads>\nf.txt\n</current_uploads>\nsummarize this")
|
||||
(kept,) = filter_messages_for_memory([mixed])
|
||||
assert extract_message_text(kept) == "summarize this"
|
||||
|
||||
def test_extract_message_text_handles_list_content(self) -> None:
|
||||
msg = AIMessage(content=[{"type": "text", "text": "part one"}, "part two"])
|
||||
assert extract_message_text(msg) == "part one part two"
|
||||
|
||||
def test_extract_message_text_treats_none_as_empty(self) -> None:
|
||||
msg = AIMessage(content="")
|
||||
msg.content = None
|
||||
assert extract_message_text(msg) == ""
|
||||
|
||||
|
||||
from typing import Any # noqa: E402
|
||||
|
||||
from deerflow.agents.memory.backends.mem0.mem0_manager import Mem0Manager # noqa: E402
|
||||
|
||||
|
||||
class FakeMem0Client:
|
||||
"""Test double injected as manager._client (records calls, returns fixtures)."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self.added: list[dict[str, Any]] = []
|
||||
self.deleted: list[dict[str, Any]] = []
|
||||
self.search_calls: list[dict[str, Any]] = []
|
||||
self.list_calls: list[dict[str, Any]] = []
|
||||
self.pings = 0
|
||||
self.search_results: list[dict[str, Any]] = []
|
||||
self.list_results: list[dict[str, Any]] = []
|
||||
self.error: Exception | None = None
|
||||
self.closed = False
|
||||
|
||||
def _maybe_raise(self) -> None:
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
|
||||
def add_memories(self, **kwargs: Any) -> dict[str, Any]:
|
||||
self._maybe_raise()
|
||||
self.added.append(kwargs)
|
||||
return {"status": "PENDING", "event_id": "evt-fake"}
|
||||
|
||||
def search_memories(self, **kwargs: Any) -> list[dict[str, Any]]:
|
||||
self._maybe_raise()
|
||||
self.search_calls.append(kwargs)
|
||||
return self.search_results
|
||||
|
||||
def list_memories(self, **kwargs: Any) -> list[dict[str, Any]]:
|
||||
self._maybe_raise()
|
||||
self.list_calls.append(kwargs)
|
||||
return self.list_results
|
||||
|
||||
def delete_all_memories(self, **kwargs: Any) -> None:
|
||||
self._maybe_raise()
|
||||
self.deleted.append(kwargs)
|
||||
|
||||
def ping(self) -> None:
|
||||
self._maybe_raise()
|
||||
self.pings += 1
|
||||
|
||||
def close(self) -> None:
|
||||
self.closed = True
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _mem0_api_key(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
"""Mem0Manager resolves the API key eagerly at construction; provide a dummy
|
||||
so the suite is hermetic (tests that need it missing delete it themselves)."""
|
||||
monkeypatch.setenv("MEM0_API_KEY", "test-key")
|
||||
|
||||
|
||||
def _manager(backend_config: dict | None = None, *, mode: str = "middleware") -> tuple[Mem0Manager, FakeMem0Client]:
|
||||
mgr = Mem0Manager(backend_config=backend_config or {}, mode=mode)
|
||||
fake = FakeMem0Client()
|
||||
mgr._client = fake
|
||||
return mgr, fake
|
||||
|
||||
|
||||
class TestMem0ManagerConstruction:
|
||||
def test_from_config_fail_fast_pings(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("MEM0_API_KEY", "k")
|
||||
fake = FakeMem0Client()
|
||||
monkeypatch.setattr(
|
||||
"deerflow.agents.memory.backends.mem0.mem0_manager.Mem0Client",
|
||||
lambda **kwargs: fake,
|
||||
)
|
||||
Mem0Manager.from_config({}, mode="middleware")
|
||||
assert fake.pings == 1
|
||||
|
||||
def test_from_config_tolerate_skips_ping(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setenv("MEM0_API_KEY", "k")
|
||||
fake = FakeMem0Client()
|
||||
monkeypatch.setattr(
|
||||
"deerflow.agents.memory.backends.mem0.mem0_manager.Mem0Client",
|
||||
lambda **kwargs: fake,
|
||||
)
|
||||
mgr = Mem0Manager.from_config({"startup_policy": "tolerate"}, mode="tool")
|
||||
assert fake.pings == 0
|
||||
assert mgr.mode == "tool"
|
||||
|
||||
def test_from_config_missing_key_raises(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.delenv("MEM0_API_KEY", raising=False)
|
||||
with pytest.raises(ValueError, match="MEM0_API_KEY"):
|
||||
Mem0Manager.from_config({})
|
||||
|
||||
def test_supports_search_enables_tool_mode(self) -> None:
|
||||
mgr, _fake = _manager(mode="tool")
|
||||
assert mgr.supports_search is True
|
||||
|
||||
def test_close_releases_http_client(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
mgr.close()
|
||||
assert fake.closed is True
|
||||
|
||||
|
||||
class TestMem0ManagerAdd:
|
||||
def test_add_maps_filtered_messages_and_identity(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
tool_ai = AIMessage(content="", tool_calls=[{"name": "t", "args": {}, "id": "1"}])
|
||||
mgr.add(
|
||||
"thread-1",
|
||||
[HumanMessage(content="I prefer dark mode"), tool_ai, AIMessage(content="Noted.")],
|
||||
agent_name="lead_agent",
|
||||
user_id="u1",
|
||||
)
|
||||
assert len(fake.added) == 1
|
||||
call = fake.added[0]
|
||||
assert call["user_id"] == "u1"
|
||||
assert call["agent_id"] == "lead_agent"
|
||||
assert call["run_id"] == "thread-1"
|
||||
assert call["messages"] == [
|
||||
{"role": "user", "content": "I prefer dark mode"},
|
||||
{"role": "assistant", "content": "Noted."},
|
||||
]
|
||||
|
||||
def test_add_without_optional_ids_uses_run_id_only(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
mgr.add("thread-9", [HumanMessage(content="hello")])
|
||||
call = fake.added[0]
|
||||
assert call["user_id"] is None
|
||||
assert call["agent_id"] is None
|
||||
assert call["run_id"] == "thread-9"
|
||||
|
||||
def test_add_empty_after_filter_is_noop(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
hidden = HumanMessage(content="internal", additional_kwargs={"hide_from_ui": True})
|
||||
mgr.add("thread-1", [hidden], user_id="u1")
|
||||
assert fake.added == []
|
||||
|
||||
def test_add_write_error_log_and_drop(self, caplog: pytest.LogCaptureFixture) -> None:
|
||||
mgr, fake = _manager()
|
||||
fake.error = Mem0APIError("server down")
|
||||
mgr.add("thread-1", [HumanMessage(content="hi")], user_id="u1") # must not raise
|
||||
assert any("mem0" in r.message for r in caplog.records)
|
||||
|
||||
def test_add_write_error_raise_policy(self) -> None:
|
||||
from deerflow.agents.memory.manager import MemoryManagerError
|
||||
|
||||
mgr, fake = _manager({"failure_policy": {"write": "raise"}})
|
||||
fake.error = Mem0APIError("server down")
|
||||
with pytest.raises(MemoryManagerError):
|
||||
mgr.add("thread-1", [HumanMessage(content="hi")], user_id="u1")
|
||||
|
||||
def test_async_add_offloads_sync_http_client(self) -> None:
|
||||
mgr, fake = _manager(mode="tool")
|
||||
event_loop_thread = threading.get_ident()
|
||||
called_from: list[int] = []
|
||||
original_add = fake.add_memories
|
||||
|
||||
def recording_add(**kwargs: Any) -> dict[str, Any]:
|
||||
called_from.append(threading.get_ident())
|
||||
return original_add(**kwargs)
|
||||
|
||||
fake.add_memories = recording_add
|
||||
asyncio.run(mgr.aadd("thread-1", [HumanMessage(content="hi")], user_id="u1"))
|
||||
|
||||
assert called_from and called_from[0] != event_loop_thread
|
||||
|
||||
|
||||
class TestMem0ManagerGetContext:
|
||||
def test_formats_dedupes_and_scopes(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
fake.list_results = [
|
||||
{"id": "m1", "memory": "likes cricket"},
|
||||
{"id": "m1", "memory": "likes cricket"}, # dup by id
|
||||
{"id": "m2", "memory": ""}, # empty dropped
|
||||
{"id": "m3", "memory": "lives in Austin"},
|
||||
]
|
||||
ctx = mgr.get_context("u1", agent_name="lead_agent", thread_id="t-1")
|
||||
assert ctx == "- likes cricket\n- lives in Austin"
|
||||
call = fake.list_calls[0]
|
||||
assert call["filters"] == {"AND": [{"user_id": "u1"}, {"agent_id": "lead_agent"}, {"run_id": "t-1"}]}
|
||||
assert call["max_items"] == 8 # default top_k
|
||||
|
||||
def test_no_identity_returns_empty(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
assert mgr.get_context(None) == ""
|
||||
assert fake.list_calls == []
|
||||
|
||||
def test_read_error_fail_open_returns_empty(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
fake.error = Mem0APIError("down")
|
||||
assert mgr.get_context("u1") == ""
|
||||
|
||||
def test_read_error_fail_closed_raises(self) -> None:
|
||||
from deerflow.agents.memory.manager import MemoryManagerError
|
||||
|
||||
mgr, fake = _manager({"failure_policy": {"read": "fail_closed"}})
|
||||
fake.error = Mem0APIError("down")
|
||||
with pytest.raises(MemoryManagerError):
|
||||
mgr.get_context("u1")
|
||||
|
||||
def test_truncates_to_max_injection_chars(self) -> None:
|
||||
mgr, fake = _manager({"max_injection_chars": 20})
|
||||
fake.list_results = [{"id": f"m{i}", "memory": "x" * 30} for i in range(3)]
|
||||
ctx = mgr.get_context("u1")
|
||||
assert len(ctx) <= 20
|
||||
|
||||
def test_async_get_context_offloads_sync_http_client(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
event_loop_thread = threading.get_ident()
|
||||
called_from: list[int] = []
|
||||
original_list = fake.list_memories
|
||||
|
||||
def recording_list(**kwargs: Any) -> list[dict[str, Any]]:
|
||||
called_from.append(threading.get_ident())
|
||||
return original_list(**kwargs)
|
||||
|
||||
fake.list_memories = recording_list
|
||||
asyncio.run(mgr.aget_context("u1"))
|
||||
|
||||
assert called_from and called_from[0] != event_loop_thread
|
||||
|
||||
|
||||
class TestMem0ManagerSearch:
|
||||
def test_maps_results_to_backend_neutral_shape(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
fake.search_results = [
|
||||
{
|
||||
"id": "m1",
|
||||
"memory": "likes cricket",
|
||||
"score": 0.9,
|
||||
"categories": ["hobbies"],
|
||||
"created_at": "2026-01-15T10:30:00Z",
|
||||
"metadata": {"source": "chat"},
|
||||
}
|
||||
]
|
||||
results = mgr.search("sports", top_k=5, user_id="u1")
|
||||
assert results == [
|
||||
{
|
||||
"id": "m1",
|
||||
"content": "likes cricket",
|
||||
"category": "hobbies",
|
||||
"confidence": 0.9,
|
||||
"createdAt": "2026-01-15T10:30:00Z",
|
||||
"source": "chat",
|
||||
}
|
||||
]
|
||||
call = fake.search_calls[0]
|
||||
assert call["filters"] == {"user_id": "u1"}
|
||||
assert call["threshold"] == 0.1 # default score_threshold
|
||||
|
||||
def test_category_filter_anded_in(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
mgr.search("q", user_id="u1", agent_name="lead_agent", category="preference")
|
||||
assert fake.search_calls[0]["filters"] == {"AND": [{"user_id": "u1"}, {"agent_id": "lead_agent"}, {"categories": {"contains": "preference"}}]}
|
||||
|
||||
def test_no_identity_returns_empty(self) -> None:
|
||||
mgr, _fake = _manager()
|
||||
assert mgr.search("q") == []
|
||||
|
||||
def test_async_search_offloads_sync_http_client(self) -> None:
|
||||
mgr, fake = _manager(mode="tool")
|
||||
event_loop_thread = threading.get_ident()
|
||||
called_from: list[int] = []
|
||||
original_search = fake.search_memories
|
||||
|
||||
def recording_search(**kwargs: Any) -> list[dict[str, Any]]:
|
||||
called_from.append(threading.get_ident())
|
||||
return original_search(**kwargs)
|
||||
|
||||
fake.search_memories = recording_search
|
||||
asyncio.run(mgr.asearch("q", user_id="u1"))
|
||||
|
||||
assert called_from and called_from[0] != event_loop_thread
|
||||
|
||||
|
||||
class TestMem0ManagerManage:
|
||||
def test_get_memory_maps_full_bucket(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
fake.list_results = [{"id": "m1", "memory": "likes cricket", "created_at": "2026-01-15T10:30:00Z"}]
|
||||
doc = mgr.get_memory(user_id="u1")
|
||||
assert doc["facts"][0]["id"] == "m1"
|
||||
assert doc["facts"][0]["content"] == "likes cricket"
|
||||
assert fake.list_calls[0].get("max_items") is None # full listing
|
||||
|
||||
def test_get_memory_no_identity_returns_empty_doc(self) -> None:
|
||||
mgr, _fake = _manager()
|
||||
assert mgr.get_memory() == {"facts": []}
|
||||
|
||||
def test_clear_memory_deletes_bucket_and_returns_empty(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
assert mgr.clear_memory(user_id="u1", agent_name="lead_agent") == {"facts": []}
|
||||
assert fake.deleted == [{"user_id": "u1", "agent_id": "lead_agent", "run_id": None}]
|
||||
|
||||
def test_clear_memory_user_wide_when_agent_none(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
mgr.clear_memory(user_id="u1")
|
||||
assert fake.deleted[0]["agent_id"] is None
|
||||
|
||||
def test_delete_memory_returns_none(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
assert mgr.delete_memory(user_id="u1") is None
|
||||
assert len(fake.deleted) == 1
|
||||
|
||||
def test_clear_memory_no_identity_is_noop(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
assert mgr.clear_memory() == {"facts": []}
|
||||
assert fake.deleted == []
|
||||
|
||||
def test_delete_memory_no_identity_is_noop(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
assert mgr.delete_memory() is None
|
||||
assert fake.deleted == []
|
||||
|
||||
def test_clear_memory_empty_string_identity_is_noop(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
assert mgr.clear_memory(user_id="", agent_name="") == {"facts": []}
|
||||
assert fake.deleted == []
|
||||
|
||||
def test_export_delegates_to_get_memory(self) -> None:
|
||||
mgr, fake = _manager()
|
||||
fake.list_results = [{"id": "m1", "memory": "x"}]
|
||||
assert mgr.export_memory(user_id="u1")["facts"][0]["id"] == "m1"
|
||||
|
||||
def test_tier3_defaults_raise_not_implemented(self) -> None:
|
||||
mgr, _fake = _manager()
|
||||
with pytest.raises(NotImplementedError):
|
||||
mgr.create_fact("x", user_id="u1")
|
||||
with pytest.raises(NotImplementedError):
|
||||
mgr.import_memory({"facts": []}, user_id="u1")
|
||||
|
||||
|
||||
class TestMem0Discovery:
|
||||
def test_scan_backends_registers_mem0(self) -> None:
|
||||
import deerflow.agents.memory.manager as manager_module
|
||||
|
||||
manager_module._backends_cache = None
|
||||
try:
|
||||
registry = manager_module._scan_backends()
|
||||
finally:
|
||||
manager_module._backends_cache = None
|
||||
assert registry["mem0"] is Mem0Manager
|
||||
|
||||
def test_factory_resolves_mem0(self, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
from deerflow.agents.memory.manager import get_memory_manager, reset_memory_manager
|
||||
from deerflow.config.memory_config import MemoryConfig, get_memory_config, set_memory_config
|
||||
|
||||
monkeypatch.setenv("MEM0_API_KEY", "k")
|
||||
monkeypatch.setattr(Mem0Client, "ping", lambda self: None)
|
||||
|
||||
original_config = get_memory_config()
|
||||
set_memory_config(MemoryConfig(manager_class="mem0"))
|
||||
reset_memory_manager()
|
||||
try:
|
||||
mgr = get_memory_manager()
|
||||
assert isinstance(mgr, Mem0Manager)
|
||||
finally:
|
||||
reset_memory_manager()
|
||||
set_memory_config(original_config)
|
||||
@ -1,5 +1,6 @@
|
||||
import asyncio
|
||||
import json
|
||||
import threading
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
@ -46,6 +47,27 @@ def test_export_memory_route_returns_current_memory() -> None:
|
||||
assert response.json()["facts"] == exported_memory["facts"]
|
||||
|
||||
|
||||
def test_get_memory_route_offloads_manager_call_from_event_loop() -> None:
|
||||
event_loop_thread = threading.get_ident()
|
||||
called_from: list[int] = []
|
||||
manager = MagicMock()
|
||||
|
||||
def get_memory(*, user_id: str) -> dict:
|
||||
called_from.append(threading.get_ident())
|
||||
return _sample_memory()
|
||||
|
||||
manager.get_memory.side_effect = get_memory
|
||||
request = SimpleNamespace()
|
||||
with (
|
||||
patch("app.gateway.routers.memory.get_memory_manager", return_value=manager),
|
||||
patch("app.gateway.routers.memory._resolve_memory_user_id", return_value="user-1"),
|
||||
):
|
||||
response = asyncio.run(memory.get_memory(request))
|
||||
|
||||
assert response.facts == []
|
||||
assert called_from and called_from[0] != event_loop_thread
|
||||
|
||||
|
||||
def test_export_memory_route_preserves_source_error() -> None:
|
||||
app = FastAPI()
|
||||
app.include_router(memory.router)
|
||||
|
||||
@ -416,6 +416,20 @@ class TestModeGating:
|
||||
assert MemoryMiddleware not in middleware_types
|
||||
assert "memory_add" in tool_names
|
||||
|
||||
def test_mem0_tool_mode_keeps_passive_write_middleware(self):
|
||||
"""mem0 has search tools but no fact CRUD, so tool mode must retain
|
||||
the per-turn middleware write path that feeds server-side extraction."""
|
||||
from deerflow.agents.factory import _assemble_from_features
|
||||
from deerflow.agents.features import RuntimeFeatures
|
||||
from deerflow.agents.middlewares.memory_middleware import MemoryMiddleware
|
||||
from deerflow.config.memory_config import MemoryConfig
|
||||
|
||||
config = MemoryConfig(enabled=True, mode="tool", manager_class="mem0")
|
||||
chain, extra_tools = _assemble_from_features(RuntimeFeatures(memory=True, memory_config=config), name="test-agent")
|
||||
|
||||
assert MemoryMiddleware in [type(m) for m in chain]
|
||||
assert "memory_search" in [tool.name for tool in extra_tools]
|
||||
|
||||
def test_middleware_mode_appends_middleware_not_tools(self, monkeypatch):
|
||||
"""When mode=middleware (default), MemoryMiddleware IS in the chain
|
||||
and memory tools are NOT in extra_tools."""
|
||||
|
||||
@ -1627,7 +1627,7 @@ summarization:
|
||||
# enabled - Master switch for the memory mechanism (call-site gate)
|
||||
# injection_enabled - Whether to inject memory into the system prompt (call-site gate)
|
||||
# shutdown_flush_timeout_seconds - Hard budget (s) to drain pending updates on Gateway graceful shutdown (default: 30)
|
||||
# manager_class - Backend selector: registered name (deermem/noop/openviking) or dotted path
|
||||
# manager_class - Backend selector: registered name (deermem/mem0/noop/openviking) or dotted path
|
||||
# backend_config - Backend-private config dict (passthrough; each backend self-interprets)
|
||||
#
|
||||
# DeerMem-private fields live under ``backend_config`` (NOT at the memory: top level):
|
||||
@ -1666,7 +1666,7 @@ memory:
|
||||
# gateway Helm deployment (see deploy/helm/deer-flow). Default 30s.
|
||||
shutdown_flush_timeout_seconds: 30.0
|
||||
# Memory backend selector. Either a registered backend name (matching a
|
||||
# backends/<name>/ folder that exposes MANAGER_CLASS, e.g. deermem / noop)
|
||||
# backends/<name>/ folder that exposes MANAGER_CLASS, e.g. deermem / mem0 / noop)
|
||||
# or a dotted import path to a MemoryManager subclass.
|
||||
manager_class: deermem
|
||||
# Memory operation mode:
|
||||
@ -1674,9 +1674,10 @@ memory:
|
||||
# tool - experimental opt-in; the model calls memory_search/memory_add/
|
||||
# memory_update/memory_delete directly. This gives the model agency over
|
||||
# memory writes, but effectiveness depends on model tool-use behavior.
|
||||
# Only one mode runs at a time. (tool mode calls the MemoryManager ABC --
|
||||
# memory_search/add/update/delete go through the active backend; backends
|
||||
# without fact-CRUD return a JSON error instead of crashing.)
|
||||
# Normally only one mode runs at a time. A backend that needs conversation-
|
||||
# level extraction may retain passive writes in tool mode while still
|
||||
# exposing query-aware search (mem0 does this). Tool calls go through the
|
||||
# active MemoryManager; unsupported fact CRUD returns a JSON error.
|
||||
mode: middleware
|
||||
# Backend-private config (a dict), passed verbatim to the backend __init__.
|
||||
# Each backend self-interprets it (DeerMem parses it into DeerMemConfig).
|
||||
|
||||
Loading…
x
Reference in New Issue
Block a user