diff --git a/plugins/memory/mem0/__init__.py b/plugins/memory/mem0/__init__.py index 32d1f6ff7002..31b879b79f03 100644 --- a/plugins/memory/mem0/__init__.py +++ b/plugins/memory/mem0/__init__.py @@ -18,9 +18,10 @@ from __future__ import annotations import json import logging import os +import re import threading import time -from typing import Any, Dict, List +from typing import Any, Dict, List, Optional from agent.memory_provider import MemoryProvider from tools.registry import tool_error @@ -100,17 +101,40 @@ CONCLUDE_SCHEMA = { "name": "mem0_conclude", "description": ( "Store a durable fact about the user. Stored verbatim (no LLM extraction). " - "Use for explicit preferences, corrections, or decisions." + "Use for explicit preferences, corrections, decisions, or operational facts. " + "Optional metadata lets callers separate preferences, doctrine, facts, and history." ), "parameters": { "type": "object", "properties": { "conclusion": {"type": "string", "description": "The fact to store."}, + "memory_class": { + "type": "string", + "enum": ["preference", "doctrine", "factual", "temporal"], + "description": "Optional class override for the stored memory.", + }, + "time_scope": { + "type": "string", + "enum": ["current", "past", "future", "timeless"], + "description": "Optional time-scope override for the stored memory.", + }, + "metadata": { + "type": "object", + "description": "Optional extra metadata stored alongside the memory.", + }, }, "required": ["conclusion"], }, } +_EXPLICIT_PREFIX_RE = re.compile(r"^\[(?:PREF|RULE|FACT|HIST)(?:/(?:CURRENT|PAST|FUTURE|TIMELESS))?\]\s*") +_MEMORY_CLASS_PREFIX = { + "preference": "PREF", + "doctrine": "RULE", + "factual": "FACT", + "temporal": "HIST", +} + # --------------------------------------------------------------------------- # MemoryProvider implementation @@ -127,10 +151,12 @@ class Mem0MemoryProvider(MemoryProvider): self._user_id = "hermes-user" self._agent_id = "hermes" self._rerank = True + self._sync_turn_mode = "full" self._prefetch_result = "" self._prefetch_lock = threading.Lock() self._prefetch_thread = None self._sync_thread = None + self._write_thread = None # Circuit breaker state self._consecutive_failures = 0 self._breaker_open_until = 0.0 @@ -163,6 +189,7 @@ class Mem0MemoryProvider(MemoryProvider): {"key": "user_id", "description": "User identifier", "default": "hermes-user"}, {"key": "agent_id", "description": "Agent identifier", "default": "hermes"}, {"key": "rerank", "description": "Enable reranking for recall", "default": "true", "choices": ["true", "false"]}, + {"key": "sync_turn_mode", "description": "Automatic turn-to-memory extraction mode", "default": "full", "choices": ["full", "off"]}, ] def _get_client(self): @@ -208,11 +235,26 @@ class Mem0MemoryProvider(MemoryProvider): self._user_id = kwargs.get("user_id") or self._config.get("user_id", "hermes-user") self._agent_id = self._config.get("agent_id", "hermes") self._rerank = self._config.get("rerank", True) + if "sync_turn_mode" in self._config: + self._sync_turn_mode = str(self._config.get("sync_turn_mode", "full") or "full").strip().lower() + elif self._config.get("sync_turns") is False: + self._sync_turn_mode = "off" + else: + self._sync_turn_mode = "full" + if self._sync_turn_mode not in {"full", "off"}: + self._sync_turn_mode = "full" def _read_filters(self) -> Dict[str, Any]: """Filters for search/get_all — scoped to user only for cross-session recall.""" return {"user_id": self._user_id} + def _scoped_read_filters(self) -> Dict[str, Any]: + """Fallback filters for stores that only return agent-scoped memories.""" + filters = self._read_filters().copy() + if self._agent_id: + filters["agent_id"] = self._agent_id + return filters + def _write_filters(self) -> Dict[str, Any]: """Filters for add — scoped to user + agent for attribution.""" return {"user_id": self._user_id, "agent_id": self._agent_id} @@ -226,6 +268,107 @@ class Mem0MemoryProvider(MemoryProvider): return response return [] + @staticmethod + def _strip_explicit_prefix(text: str) -> str: + return _EXPLICIT_PREFIX_RE.sub("", (text or "").strip()) + + @classmethod + def _normalize_memory_class(cls, value: Optional[str], *, target: str = "memory", content: str = "") -> str: + raw = str(value or "").strip().lower() + aliases = { + "pref": "preference", + "preference": "preference", + "user": "preference", + "rule": "doctrine", + "doctrine": "doctrine", + "policy": "doctrine", + "fact": "factual", + "factual": "factual", + "current_fact": "factual", + "history": "temporal", + "historical": "temporal", + "temporal": "temporal", + "past": "temporal", + } + if raw in aliases: + return aliases[raw] + if target == "user": + return "preference" + base = cls._strip_explicit_prefix(content).lower() + if re.search(r"\b(do not|don't|never|must|should|unless|only if|expected to|keep .* on)\b", base): + return "doctrine" + if re.search(r"\b(previously|earlier|before|used to|was changed|retired|on 20\d{2}-\d{2}-\d{2}|as of)\b", base): + return "temporal" + return "factual" + + @staticmethod + def _normalize_time_scope(value: Optional[str], content: str = "") -> str: + raw = str(value or "").strip().lower() + aliases = { + "current": "current", + "now": "current", + "present": "current", + "past": "past", + "previous": "past", + "historical": "past", + "future": "future", + "planned": "future", + "timeless": "timeless", + "rule": "timeless", + } + if raw in aliases: + return aliases[raw] + base = Mem0MemoryProvider._strip_explicit_prefix(content).lower() + if re.search(r"\b(plan|planned|later|upcoming|will|future)\b", base): + return "future" + if re.search(r"\b(always|never|must|should|unless|only if|expected to)\b", base): + return "timeless" + if re.search(r"\b(previously|earlier|before|used to|was|were|retired|on 20\d{2}-\d{2}-\d{2}|as of)\b", base): + return "past" + return "current" + + @classmethod + def _decorate_explicit_memory( + cls, + content: str, + *, + target: str = "memory", + memory_class: Optional[str] = None, + time_scope: Optional[str] = None, + ) -> str: + base = cls._strip_explicit_prefix(content) + memory_class_norm = cls._normalize_memory_class(memory_class, target=target, content=base) + time_scope_norm = cls._normalize_time_scope(time_scope, base) + prefix = _MEMORY_CLASS_PREFIX.get(memory_class_norm, "FACT") + if memory_class_norm == "preference": + return f"[{prefix}] {base}" + return f"[{prefix}/{time_scope_norm.upper()}] {base}" + + @classmethod + def _build_explicit_metadata( + cls, + *, + target: str, + content: str, + memory_class: Optional[str] = None, + time_scope: Optional[str] = None, + metadata: Optional[Dict[str, Any]] = None, + source: str = "mem0_conclude", + action: str = "add", + ) -> Dict[str, Any]: + merged = dict(metadata or {}) + merged.setdefault("memory_target", target) + merged.setdefault("source", source) + merged.setdefault("write_action", action) + merged["memory_class"] = cls._normalize_memory_class(memory_class, target=target, content=content) + merged["time_scope"] = cls._normalize_time_scope(time_scope, content) + merged["explicit_memory"] = True + return merged + + @classmethod + def _format_memory_for_display(cls, text: str) -> str: + return cls._strip_explicit_prefix(text) + def system_prompt_block(self) -> str: return ( "# Mem0 Memory\n" @@ -234,6 +377,30 @@ class Mem0MemoryProvider(MemoryProvider): "mem0_profile for a full overview." ) + def _search_with_fallback(self, client: Any, *, query: str, rerank: bool, top_k: int) -> list: + """Search user-wide first, then fall back to agent-scoped reads if empty.""" + results = self._unwrap_results(client.search( + query=query, + filters=self._read_filters(), + rerank=rerank, + top_k=top_k, + )) + if results or not self._agent_id: + return results + return self._unwrap_results(client.search( + query=query, + filters=self._scoped_read_filters(), + rerank=rerank, + top_k=top_k, + )) + + def _get_all_with_fallback(self, client: Any) -> list: + """Load user-wide memories first, then fall back to agent-scoped reads if empty.""" + memories = self._unwrap_results(client.get_all(filters=self._read_filters())) + if memories or not self._agent_id: + return memories + return self._unwrap_results(client.get_all(filters=self._scoped_read_filters())) + def prefetch(self, query: str, *, session_id: str = "") -> str: if self._prefetch_thread and self._prefetch_thread.is_alive(): self._prefetch_thread.join(timeout=3.0) @@ -251,16 +418,16 @@ class Mem0MemoryProvider(MemoryProvider): def _run(): try: client = self._get_client() - results = self._unwrap_results(client.search( + results = self._search_with_fallback( + client, query=query, - filters=self._read_filters(), rerank=self._rerank, top_k=5, - )) + ) if results: - lines = [r.get("memory", "") for r in results if r.get("memory")] + lines = [self._format_memory_for_display(r.get("memory", "")) for r in results if r.get("memory")] with self._prefetch_lock: - self._prefetch_result = "\n".join(f"- {l}" for l in lines) + self._prefetch_result = "\n".join(f"- {l}" for l in lines if l) self._record_success() except Exception as e: self._record_failure() @@ -273,6 +440,8 @@ class Mem0MemoryProvider(MemoryProvider): """Send the turn to Mem0 for server-side fact extraction (non-blocking).""" if self._is_breaker_open(): return + if self._sync_turn_mode == "off": + return def _sync(): try: @@ -310,12 +479,12 @@ class Mem0MemoryProvider(MemoryProvider): if tool_name == "mem0_profile": try: - memories = self._unwrap_results(client.get_all(filters=self._read_filters())) + memories = self._get_all_with_fallback(client) self._record_success() if not memories: return json.dumps({"result": "No memories stored yet."}) - lines = [m.get("memory", "") for m in memories if m.get("memory")] - return json.dumps({"result": "\n".join(lines), "count": len(lines)}) + lines = [self._format_memory_for_display(m.get("memory", "")) for m in memories if m.get("memory")] + return json.dumps({"result": "\n".join(l for l in lines if l), "count": len(lines)}) except Exception as e: self._record_failure() return tool_error(f"Failed to fetch profile: {e}") @@ -327,16 +496,23 @@ class Mem0MemoryProvider(MemoryProvider): rerank = args.get("rerank", False) top_k = min(int(args.get("top_k", 10)), 50) try: - results = self._unwrap_results(client.search( + results = self._search_with_fallback( + client, query=query, - filters=self._read_filters(), rerank=rerank, top_k=top_k, - )) + ) self._record_success() if not results: return json.dumps({"result": "No relevant memories found."}) - items = [{"memory": r.get("memory", ""), "score": r.get("score", 0)} for r in results] + items = [ + { + "memory": self._format_memory_for_display(r.get("memory", "")), + "score": r.get("score", 0), + "metadata": r.get("metadata", {}), + } + for r in results + ] return json.dumps({"results": items, "count": len(items)}) except Exception as e: self._record_failure() @@ -347,21 +523,89 @@ class Mem0MemoryProvider(MemoryProvider): if not conclusion: return tool_error("Missing required parameter: conclusion") try: + explicit_text = self._decorate_explicit_memory( + conclusion, + target="memory", + memory_class=args.get("memory_class"), + time_scope=args.get("time_scope"), + ) + explicit_metadata = self._build_explicit_metadata( + target="memory", + content=conclusion, + memory_class=args.get("memory_class"), + time_scope=args.get("time_scope"), + metadata=args.get("metadata") if isinstance(args.get("metadata"), dict) else None, + source="mem0_conclude", + action="add", + ) client.add( - [{"role": "user", "content": conclusion}], + [{"role": "user", "content": explicit_text}], **self._write_filters(), infer=False, + metadata=explicit_metadata, ) self._record_success() - return json.dumps({"result": "Fact stored."}) + return json.dumps({ + "result": "Fact stored.", + "stored": self._format_memory_for_display(explicit_text), + "memory_class": explicit_metadata.get("memory_class"), + "time_scope": explicit_metadata.get("time_scope"), + }) except Exception as e: self._record_failure() return tool_error(f"Failed to store: {e}") return tool_error(f"Unknown tool: {tool_name}") + def on_memory_write( + self, + action: str, + target: str, + content: str, + metadata: Optional[Dict[str, Any]] = None, + ) -> None: + """Mirror built-in memory writes to Mem0 using explicit classification metadata.""" + if self._is_breaker_open(): + return + if action not in {"add", "replace"} or not (content or "").strip(): + return + + def _write(): + try: + client = self._get_client() + explicit_text = self._decorate_explicit_memory( + content, + target=target, + memory_class=(metadata or {}).get("memory_class") if isinstance(metadata, dict) else None, + time_scope=(metadata or {}).get("time_scope") if isinstance(metadata, dict) else None, + ) + explicit_metadata = self._build_explicit_metadata( + target=target, + content=content, + memory_class=(metadata or {}).get("memory_class") if isinstance(metadata, dict) else None, + time_scope=(metadata or {}).get("time_scope") if isinstance(metadata, dict) else None, + metadata=metadata, + source="builtin_memory_tool", + action=action, + ) + client.add( + [{"role": "user", "content": explicit_text}], + **self._write_filters(), + infer=False, + metadata=explicit_metadata, + ) + self._record_success() + except Exception as e: + self._record_failure() + logger.debug("Mem0 on_memory_write failed: %s", e) + + if self._write_thread and self._write_thread.is_alive(): + self._write_thread.join(timeout=2.0) + self._write_thread = threading.Thread(target=_write, daemon=True, name="mem0-memwrite") + self._write_thread.start() + def shutdown(self) -> None: - for t in (self._prefetch_thread, self._sync_thread): + for t in (self._prefetch_thread, self._sync_thread, self._write_thread): if t and t.is_alive(): t.join(timeout=5.0) with self._client_lock: diff --git a/tests/plugins/memory/test_mem0_v2.py b/tests/plugins/memory/test_mem0_v2.py index 6f60771f5c48..ca5d23e5e19f 100644 --- a/tests/plugins/memory/test_mem0_v2.py +++ b/tests/plugins/memory/test_mem0_v2.py @@ -12,19 +12,29 @@ from plugins.memory.mem0 import Mem0MemoryProvider class FakeClientV2: """Fake Mem0 client that returns v2-style dict responses and captures call kwargs.""" - def __init__(self, search_results=None, all_results=None): + def __init__(self, search_results=None, all_results=None, search_results_sequence=None, all_results_sequence=None): self._search_results = search_results or {"results": []} self._all_results = all_results or {"results": []} + self._search_results_sequence = list(search_results_sequence or []) + self._all_results_sequence = list(all_results_sequence or []) self.captured_search = {} self.captured_get_all = {} + self.search_calls = [] + self.get_all_calls = [] self.captured_add = [] def search(self, **kwargs): self.captured_search = kwargs + self.search_calls.append(kwargs) + if self._search_results_sequence: + return self._search_results_sequence.pop(0) return self._search_results def get_all(self, **kwargs): self.captured_get_all = kwargs + self.get_all_calls.append(kwargs) + if self._all_results_sequence: + return self._all_results_sequence.pop(0) return self._all_results def add(self, messages, **kwargs): @@ -48,7 +58,7 @@ class TestMem0FiltersV2: return provider def test_search_uses_filters(self, monkeypatch): - client = FakeClientV2() + client = FakeClientV2(search_results={"results": [{"memory": "hello", "score": 0.9}]}) provider = self._make_provider(monkeypatch, client) provider.handle_tool_call("mem0_search", {"query": "hello", "top_k": 3, "rerank": False}) @@ -61,7 +71,7 @@ class TestMem0FiltersV2: assert "user_id" not in {k for k in client.captured_search if k != "filters"} def test_profile_uses_filters(self, monkeypatch): - client = FakeClientV2() + client = FakeClientV2(all_results={"results": [{"memory": "alpha"}]}) provider = self._make_provider(monkeypatch, client) provider.handle_tool_call("mem0_profile", {}) @@ -70,7 +80,7 @@ class TestMem0FiltersV2: assert "user_id" not in {k for k in client.captured_get_all if k != "filters"} def test_prefetch_uses_filters(self, monkeypatch): - client = FakeClientV2() + client = FakeClientV2(search_results={"results": [{"memory": "hello"}]}) provider = self._make_provider(monkeypatch, client) provider.queue_prefetch("hello") @@ -96,13 +106,17 @@ class TestMem0FiltersV2: client = FakeClientV2() provider = self._make_provider(monkeypatch, client) - provider.handle_tool_call("mem0_conclude", {"conclusion": "user likes dark mode"}) + result = json.loads(provider.handle_tool_call("mem0_conclude", {"conclusion": "user likes dark mode"})) assert len(client.captured_add) == 1 call = client.captured_add[0] assert call["user_id"] == "u123" assert call["agent_id"] == "hermes" assert call["infer"] is False + assert call["messages"][0]["content"] == "[FACT/CURRENT] user likes dark mode" + assert call["metadata"]["memory_class"] == "factual" + assert call["metadata"]["time_scope"] == "current" + assert result["stored"] == "user likes dark mode" def test_read_filters_no_agent_id(self): """Read filters should use user_id only — cross-session recall across agents.""" @@ -111,6 +125,54 @@ class TestMem0FiltersV2: provider._agent_id = "hermes" assert provider._read_filters() == {"user_id": "u123"} + def test_search_falls_back_to_agent_scope_when_user_wide_empty(self, monkeypatch): + client = FakeClientV2( + search_results_sequence=[ + {"results": []}, + {"results": [{"memory": "agent memory", "score": 0.8}]}, + ] + ) + provider = self._make_provider(monkeypatch, client) + + result = json.loads(provider.handle_tool_call("mem0_search", {"query": "hello", "top_k": 3})) + + assert result["count"] == 1 + assert len(client.search_calls) == 2 + assert client.search_calls[0]["filters"] == {"user_id": "u123"} + assert client.search_calls[1]["filters"] == {"user_id": "u123", "agent_id": "hermes"} + + def test_profile_falls_back_to_agent_scope_when_user_wide_empty(self, monkeypatch): + client = FakeClientV2( + all_results_sequence=[ + {"results": []}, + {"results": [{"memory": "agent memory"}]}, + ] + ) + provider = self._make_provider(monkeypatch, client) + + result = json.loads(provider.handle_tool_call("mem0_profile", {})) + + assert result["count"] == 1 + assert len(client.get_all_calls) == 2 + assert client.get_all_calls[0]["filters"] == {"user_id": "u123"} + assert client.get_all_calls[1]["filters"] == {"user_id": "u123", "agent_id": "hermes"} + + def test_prefetch_falls_back_to_agent_scope_when_user_wide_empty(self, monkeypatch): + client = FakeClientV2( + search_results_sequence=[ + {"results": []}, + {"results": [{"memory": "agent memory"}]}, + ] + ) + provider = self._make_provider(monkeypatch, client) + + provider.queue_prefetch("hello") + provider._prefetch_thread.join(timeout=2) + + assert len(client.search_calls) == 2 + assert client.search_calls[0]["filters"] == {"user_id": "u123"} + assert client.search_calls[1]["filters"] == {"user_id": "u123", "agent_id": "hermes"} + def test_write_filters_include_agent_id(self): """Write filters should include agent_id for attribution.""" provider = Mem0MemoryProvider() @@ -225,3 +287,101 @@ class TestMem0Defaults: provider.initialize("test") assert provider._agent_id == "hermes" + + def test_default_sync_turn_mode_full(self, monkeypatch, tmp_path): + monkeypatch.setenv("MEM0_API_KEY", "test-key") + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + + provider = Mem0MemoryProvider() + provider.initialize("test") + + assert provider._sync_turn_mode == "full" + + def test_legacy_sync_turns_false_maps_to_off(self, monkeypatch, tmp_path): + monkeypatch.setenv("MEM0_API_KEY", "test-key") + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + (tmp_path / "mem0.json").write_text(json.dumps({"sync_turns": False})) + + provider = Mem0MemoryProvider() + provider.initialize("test") + + assert provider._sync_turn_mode == "off" + + +class TestMem0ExplicitMemoryFormatting: + def test_decorate_preference_memory(self): + assert Mem0MemoryProvider._decorate_explicit_memory( + "User prefers concise responses.", + target="user", + ) == "[PREF] User prefers concise responses." + + def test_decorate_doctrine_memory(self): + assert Mem0MemoryProvider._decorate_explicit_memory( + "OpenClaw is retired and should be ignored unless live Hermes residue still matters.", + target="memory", + ).startswith("[RULE/TIMELESS]") + + def test_format_memory_for_display_strips_prefix(self): + assert Mem0MemoryProvider._format_memory_for_display( + "[FACT/PAST] Previously Dockerized apps were hosted locally." + ) == "Previously Dockerized apps were hosted locally." + + def test_on_memory_write_mirrors_explicit_metadata(self, monkeypatch): + client = FakeClientV2() + provider = Mem0MemoryProvider() + provider.initialize("test-session") + provider._user_id = "u123" + provider._agent_id = "hermes" + monkeypatch.setattr(provider, "_get_client", lambda: client) + + provider.on_memory_write( + "add", + "user", + "User prefers concise responses.", + metadata={"tool_name": "memory", "platform": "telegram"}, + ) + assert provider._write_thread is not None + provider._write_thread.join(timeout=2) + + assert len(client.captured_add) == 1 + call = client.captured_add[0] + assert call["messages"][0]["content"] == "[PREF] User prefers concise responses." + assert call["metadata"]["memory_class"] == "preference" + assert call["metadata"]["time_scope"] == "current" + assert call["metadata"]["source"] == "builtin_memory_tool" + + +class TestMem0SyncTurnMode: + def _make_provider(self, monkeypatch, client, tmp_path, sync_turn_mode="full"): + monkeypatch.setenv("MEM0_API_KEY", "test-key") + monkeypatch.setenv("HERMES_HOME", str(tmp_path)) + mem0_cfg = tmp_path / "mem0.json" + mem0_cfg.write_text(json.dumps({"sync_turn_mode": sync_turn_mode})) + provider = Mem0MemoryProvider() + provider.initialize("test-session") + provider._user_id = "u123" + provider._agent_id = "hermes" + monkeypatch.setattr(provider, "_get_client", lambda: client) + return provider + + def test_sync_turn_mode_off_skips_turn_write(self, monkeypatch, tmp_path): + client = FakeClientV2() + provider = self._make_provider(monkeypatch, client, tmp_path, sync_turn_mode="off") + + provider.sync_turn("user said this", "assistant replied", session_id="s1") + + assert provider._sync_thread is None + assert client.captured_add == [] + + def test_sync_turn_mode_full_preserves_existing_behavior(self, monkeypatch, tmp_path): + client = FakeClientV2() + provider = self._make_provider(monkeypatch, client, tmp_path, sync_turn_mode="full") + + provider.sync_turn("user said this", "assistant replied", session_id="s1") + assert provider._sync_thread is not None + provider._sync_thread.join(timeout=2) + + assert len(client.captured_add) == 1 + call = client.captured_add[0] + assert call["user_id"] == "u123" + assert call["agent_id"] == "hermes"