mirror of
https://github.com/NousResearch/hermes-agent.git
synced 2026-07-23 16:36:23 +00:00
Local hardening: improve Mem0 explicit-memory classification and mirroring
This commit is contained in:
parent
458a94e425
commit
aae5638fd3
2 changed files with 426 additions and 22 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue