From 352f247a81fc2a4980e37345a41a75ed1cac6890 Mon Sep 17 00:00:00 2001 From: Vanzeren <53075619+Vanzeren@users.noreply.github.com> Date: Wed, 29 Jul 2026 07:11:20 +0800 Subject: [PATCH] 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 --- README.md | 6 + backend/AGENTS.md | 17 + backend/app/gateway/app.py | 6 +- backend/app/gateway/routers/memory.py | 49 +- .../harness/deerflow/agents/factory.py | 7 +- .../deerflow/agents/lead_agent/agent.py | 8 +- .../deerflow/agents/lead_agent/prompt.py | 8 +- .../agents/memory/backends/mem0/README.md | 73 ++ .../agents/memory/backends/mem0/__init__.py | 9 + .../agents/memory/backends/mem0/client.py | 128 ++++ .../agents/memory/backends/mem0/config.py | 123 ++++ .../memory/backends/mem0/mem0_manager.py | 315 +++++++++ .../memory/backends/mem0/message_filtering.py | 93 +++ .../harness/deerflow/agents/memory/manager.py | 13 + .../harness/deerflow/agents/memory/tools.py | 8 +- .../agents/middlewares/memory_middleware.py | 41 +- .../tests/test_lead_agent_model_resolution.py | 21 +- backend/tests/test_lead_agent_prompt.py | 32 + backend/tests/test_mem0_memory_backend.py | 630 ++++++++++++++++++ backend/tests/test_memory_router.py | 22 + backend/tests/test_memory_tools.py | 14 + config.example.yaml | 11 +- 22 files changed, 1584 insertions(+), 50 deletions(-) create mode 100644 backend/packages/harness/deerflow/agents/memory/backends/mem0/README.md create mode 100644 backend/packages/harness/deerflow/agents/memory/backends/mem0/__init__.py create mode 100644 backend/packages/harness/deerflow/agents/memory/backends/mem0/client.py create mode 100644 backend/packages/harness/deerflow/agents/memory/backends/mem0/config.py create mode 100644 backend/packages/harness/deerflow/agents/memory/backends/mem0/mem0_manager.py create mode 100644 backend/packages/harness/deerflow/agents/memory/backends/mem0/message_filtering.py create mode 100644 backend/tests/test_mem0_memory_backend.py diff --git a/README.md b/README.md index 37cc3d762..a04992226 100644 --- a/README.md +++ b/README.md @@ -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. diff --git a/backend/AGENTS.md b/backend/AGENTS.md index 68d292d82..83720d2c3 100644 --- a/backend/AGENTS.md +++ b/backend/AGENTS.md @@ -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 diff --git a/backend/app/gateway/app.py b/backend/app/gateway/app.py index 1c97050a7..705f9be7d 100644 --- a/backend/app/gateway/app.py +++ b/backend/app/gateway/app.py @@ -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: diff --git a/backend/app/gateway/routers/memory.py b/backend/app/gateway/routers/memory.py index d4f38703e..291983df2 100644 --- a/backend/app/gateway/routers/memory.py +++ b/backend/app/gateway/routers/memory.py @@ -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( diff --git a/backend/packages/harness/deerflow/agents/factory.py b/backend/packages/harness/deerflow/agents/factory.py index 64dcb45b6..9270d3428 100644 --- a/backend/packages/harness/deerflow/agents/factory.py +++ b/backend/packages/harness/deerflow/agents/factory.py @@ -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.") diff --git a/backend/packages/harness/deerflow/agents/lead_agent/agent.py b/backend/packages/harness/deerflow/agents/lead_agent/agent.py index 2ca9768f9..0555f9b9c 100644 --- a/backend/packages/harness/deerflow/agents/lead_agent/agent.py +++ b/backend/packages/harness/deerflow/agents/lead_agent/agent.py @@ -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.") diff --git a/backend/packages/harness/deerflow/agents/lead_agent/prompt.py b/backend/packages/harness/deerflow/agents/lead_agent/prompt.py index 7707e0896..3781fee38 100644 --- a/backend/packages/harness/deerflow/agents/lead_agent/prompt.py +++ b/backend/packages/harness/deerflow/agents/lead_agent/prompt.py @@ -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} """ - 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 "" diff --git a/backend/packages/harness/deerflow/agents/memory/backends/mem0/README.md b/backend/packages/harness/deerflow/agents/memory/backends/mem0/README.md new file mode 100644 index 000000000..09e44ee92 --- /dev/null +++ b/backend/packages/harness/deerflow/agents/memory/backends/mem0/README.md @@ -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. diff --git a/backend/packages/harness/deerflow/agents/memory/backends/mem0/__init__.py b/backend/packages/harness/deerflow/agents/memory/backends/mem0/__init__.py new file mode 100644 index 000000000..93a75e6e2 --- /dev/null +++ b/backend/packages/harness/deerflow/agents/memory/backends/mem0/__init__.py @@ -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 diff --git a/backend/packages/harness/deerflow/agents/memory/backends/mem0/client.py b/backend/packages/harness/deerflow/agents/memory/backends/mem0/client.py new file mode 100644 index 000000000..02513ec41 --- /dev/null +++ b/backend/packages/harness/deerflow/agents/memory/backends/mem0/client.py @@ -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) diff --git a/backend/packages/harness/deerflow/agents/memory/backends/mem0/config.py b/backend/packages/harness/deerflow/agents/memory/backends/mem0/config.py new file mode 100644 index 000000000..ce48c8cf5 --- /dev/null +++ b/backend/packages/harness/deerflow/agents/memory/backends/mem0/config.py @@ -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 diff --git a/backend/packages/harness/deerflow/agents/memory/backends/mem0/mem0_manager.py b/backend/packages/harness/deerflow/agents/memory/backends/mem0/mem0_manager.py new file mode 100644 index 000000000..6207c7f33 --- /dev/null +++ b/backend/packages/harness/deerflow/agents/memory/backends/mem0/mem0_manager.py @@ -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)) diff --git a/backend/packages/harness/deerflow/agents/memory/backends/mem0/message_filtering.py b/backend/packages/harness/deerflow/agents/memory/backends/mem0/message_filtering.py new file mode 100644 index 000000000..974a253cc --- /dev/null +++ b/backend/packages/harness/deerflow/agents/memory/backends/mem0/message_filtering.py @@ -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"<(?Puploaded_files|current_uploads)>[\s\S]*?\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 "" in text.lower() or "" 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 diff --git a/backend/packages/harness/deerflow/agents/memory/manager.py b/backend/packages/harness/deerflow/agents/memory/manager.py index dc002fb5e..89f16a049 100644 --- a/backend/packages/harness/deerflow/agents/memory/manager.py +++ b/backend/packages/harness/deerflow/agents/memory/manager.py @@ -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 diff --git a/backend/packages/harness/deerflow/agents/memory/tools.py b/backend/packages/harness/deerflow/agents/memory/tools.py index 20dcda958..5be20b046 100644 --- a/backend/packages/harness/deerflow/agents/memory/tools.py +++ b/backend/packages/harness/deerflow/agents/memory/tools.py @@ -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; diff --git a/backend/packages/harness/deerflow/agents/middlewares/memory_middleware.py b/backend/packages/harness/deerflow/agents/middlewares/memory_middleware.py index 4d7faccf1..e0e946726 100644 --- a/backend/packages/harness/deerflow/agents/middlewares/memory_middleware.py +++ b/backend/packages/harness/deerflow/agents/middlewares/memory_middleware.py @@ -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 diff --git a/backend/tests/test_lead_agent_model_resolution.py b/backend/tests/test_lead_agent_model_resolution.py index 2a8c974de..804dca9ad 100644 --- a/backend/tests/test_lead_agent_model_resolution.py +++ b/backend/tests/test_lead_agent_model_resolution.py @@ -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) # --------------------------------------------------------------------------- diff --git a/backend/tests/test_lead_agent_prompt.py b/backend/tests/test_lead_agent_prompt.py index 5d8a7ffb2..20150325d 100644 --- a/backend/tests/test_lead_agent_prompt.py +++ b/backend/tests/test_lead_agent_prompt.py @@ -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), diff --git a/backend/tests/test_mem0_memory_backend.py b/backend/tests/test_mem0_memory_backend.py new file mode 100644 index 000000000..070db8392 --- /dev/null +++ b/backend/tests/test_mem0_memory_backend.py @@ -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="\nfile.pdf\n") + 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="\nf.txt\n\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) diff --git a/backend/tests/test_memory_router.py b/backend/tests/test_memory_router.py index 4031e3c9d..7d2e3bd25 100644 --- a/backend/tests/test_memory_router.py +++ b/backend/tests/test_memory_router.py @@ -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) diff --git a/backend/tests/test_memory_tools.py b/backend/tests/test_memory_tools.py index 9b163c1b4..c1ec8195e 100644 --- a/backend/tests/test_memory_tools.py +++ b/backend/tests/test_memory_tools.py @@ -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.""" diff --git a/config.example.yaml b/config.example.yaml index 74a431074..81f4bbf66 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -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// folder that exposes MANAGER_CLASS, e.g. deermem / noop) + # backends// 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).