refactor(integrations): rename dsh-mem0 to deepseek-plugin, strands-mem0 to mem0-strands (#7098)
This commit is contained in:
@@ -0,0 +1,214 @@
|
||||
"""Tests for Mem0ServiceClient: backend routing, call shapes, and response shaping.
|
||||
|
||||
The fakes mirror the *real* mem0ai signatures: a keyword-only ``search`` that
|
||||
rejects top-level entity params, and a fixed-signature OSS ``add`` with no
|
||||
``**kwargs``. So a call shape the real SDK would reject fails here too, which is
|
||||
what the earlier permissive fakes did not do.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from mem0_strands.client import Mem0ServiceClient, _extract_results, _is_platform_client
|
||||
|
||||
|
||||
class FakeMemoryClient:
|
||||
"""Stand-in for mem0.MemoryClient (platform: add/search take **kwargs)."""
|
||||
|
||||
def __init__(self):
|
||||
self.add_calls = []
|
||||
self.search_calls = []
|
||||
|
||||
def add(self, messages, **kwargs):
|
||||
self.add_calls.append((messages, kwargs))
|
||||
return {"results": [{"id": "m1"}]}
|
||||
|
||||
def search(self, query, **kwargs):
|
||||
self.search_calls.append((query, kwargs))
|
||||
return {"results": [{"id": "m1", "memory": "hi"}]}
|
||||
|
||||
|
||||
class FakeMemory:
|
||||
"""Stand-in for mem0.Memory (OSS) with the real, strict signatures.
|
||||
|
||||
``search`` is keyword-only and rejects top-level entity params; ``add`` has a
|
||||
fixed signature with no ``**kwargs`` (so ``source`` or ``app_id`` is a
|
||||
``TypeError``), exactly like the shipped SDK.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.add_calls = []
|
||||
self.search_calls = []
|
||||
|
||||
def add(
|
||||
self,
|
||||
messages,
|
||||
*,
|
||||
user_id=None,
|
||||
agent_id=None,
|
||||
run_id=None,
|
||||
metadata=None,
|
||||
infer=True,
|
||||
timestamp=None,
|
||||
expiration_date=None,
|
||||
memory_type=None,
|
||||
prompt=None,
|
||||
):
|
||||
self.add_calls.append(
|
||||
(
|
||||
messages,
|
||||
{"user_id": user_id, "agent_id": agent_id, "run_id": run_id, "metadata": metadata, "infer": infer},
|
||||
)
|
||||
)
|
||||
return {"results": []}
|
||||
|
||||
def search(self, query, *, top_k=20, filters=None, threshold=0.1, **kwargs):
|
||||
rejected = kwargs.keys() & {"user_id", "agent_id", "run_id", "app_id"}
|
||||
if rejected:
|
||||
raise ValueError(f"Top-level entity parameters {set(rejected)} are not supported in search().")
|
||||
self.search_calls.append((query, {"top_k": top_k, "filters": filters}))
|
||||
return {"results": []}
|
||||
|
||||
|
||||
class FakeAsyncMemoryClient:
|
||||
"""Stand-in for mem0.AsyncMemoryClient: coroutine add/search."""
|
||||
|
||||
async def add(self, messages, **kwargs): # pragma: no cover - never called
|
||||
return {}
|
||||
|
||||
async def search(self, query, **kwargs): # pragma: no cover - never called
|
||||
return {}
|
||||
|
||||
|
||||
def platform_client():
|
||||
"""A Mem0ServiceClient wrapping a fake platform client."""
|
||||
fake = FakeMemoryClient()
|
||||
fake.__class__.__name__ = "MemoryClient"
|
||||
return Mem0ServiceClient(client=fake), fake
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# backend detection
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_detects_platform_by_class_name():
|
||||
assert _is_platform_client(FakeMemory()) is False
|
||||
fake = FakeMemoryClient()
|
||||
fake.__class__.__name__ = "MemoryClient"
|
||||
assert _is_platform_client(fake) is True
|
||||
|
||||
|
||||
def test_injected_client_sets_platform_flag():
|
||||
fake = FakeMemoryClient()
|
||||
fake.__class__.__name__ = "MemoryClient"
|
||||
assert Mem0ServiceClient(client=fake).is_platform is True
|
||||
assert Mem0ServiceClient(client=FakeMemory()).is_platform is False
|
||||
|
||||
|
||||
def test_async_client_is_rejected():
|
||||
"""An async Mem0 client cannot be driven from a worker thread; reject it loudly."""
|
||||
with pytest.raises(ValueError, match="Async Mem0 clients are not supported"):
|
||||
Mem0ServiceClient(client=FakeAsyncMemoryClient())
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# write routing
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_store_memory_is_verbatim_and_tagged():
|
||||
"""Platform store_memory writes infer=False, scope top-level, and a source tag."""
|
||||
client, fake = platform_client()
|
||||
|
||||
client.store_memory("a fact", {"user_id": "alex"}, {"k": "v"})
|
||||
|
||||
messages, kwargs = fake.add_calls[0]
|
||||
assert messages == "a fact"
|
||||
assert kwargs["infer"] is False
|
||||
assert kwargs["user_id"] == "alex"
|
||||
assert kwargs["metadata"] == {"k": "v"}
|
||||
assert kwargs["source"] == "STRANDS"
|
||||
|
||||
|
||||
def test_store_messages_infers_and_tags():
|
||||
"""Platform store_messages hands turns to Mem0 with infer=True and a source tag."""
|
||||
client, fake = platform_client()
|
||||
|
||||
turns = [{"role": "user", "content": "hi"}]
|
||||
client.store_messages(turns, {"user_id": "alex"})
|
||||
|
||||
messages, kwargs = fake.add_calls[0]
|
||||
assert messages == turns
|
||||
assert kwargs["infer"] is True
|
||||
assert kwargs["source"] == "STRANDS"
|
||||
|
||||
|
||||
def test_oss_writes_omit_source():
|
||||
"""OSS Memory.add has no source parameter, so the tag must be platform-only.
|
||||
|
||||
(If the code passed source here, FakeMemory.add would raise TypeError.)
|
||||
"""
|
||||
fake = FakeMemory()
|
||||
client = Mem0ServiceClient(client=fake)
|
||||
|
||||
client.store_memory("a fact", {"user_id": "alex"}, None)
|
||||
client.store_messages([{"role": "user", "content": "hi"}], {"user_id": "alex"})
|
||||
|
||||
assert len(fake.add_calls) == 2
|
||||
for _, kwargs in fake.add_calls:
|
||||
assert "source" not in kwargs
|
||||
|
||||
|
||||
def test_oss_add_app_id_is_rejected():
|
||||
"""app_id is platform-only; the OSS path fails loudly rather than TypeError-ing."""
|
||||
client = Mem0ServiceClient(client=FakeMemory())
|
||||
with pytest.raises(ValueError, match="platform-only"):
|
||||
client.store_memory("f", {"user_id": "alex", "app_id": "app1"}, None)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# search routing (filters + top_k on both backends)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_platform_search_uses_filters():
|
||||
"""Platform search passes scope inside filters with top_k, never top-level."""
|
||||
client, fake = platform_client()
|
||||
|
||||
client.search_memories("q", {"user_id": "alex"}, 5)
|
||||
|
||||
_, kwargs = fake.search_calls[0]
|
||||
assert kwargs["filters"] == {"user_id": "alex"}
|
||||
assert kwargs["top_k"] == 5
|
||||
assert "user_id" not in kwargs
|
||||
|
||||
|
||||
def test_oss_search_uses_filters():
|
||||
"""OSS search also takes filters + top_k. The strict fake would raise on the
|
||||
old top-level/limit call shape, so this is the regression test for blocker 1."""
|
||||
fake = FakeMemory()
|
||||
client = Mem0ServiceClient(client=fake)
|
||||
|
||||
client.search_memories("q", {"user_id": "alex"}, 5)
|
||||
|
||||
_, recorded = fake.search_calls[0]
|
||||
assert recorded["filters"] == {"user_id": "alex"}
|
||||
assert recorded["top_k"] == 5
|
||||
|
||||
|
||||
def test_oss_search_app_id_is_rejected():
|
||||
client = Mem0ServiceClient(client=FakeMemory())
|
||||
with pytest.raises(ValueError, match="platform-only"):
|
||||
client.search_memories("q", {"user_id": "alex", "app_id": "app1"}, 5)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# response normalization
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_extract_results_shapes():
|
||||
assert _extract_results({"results": [{"id": 1}]}) == [{"id": 1}]
|
||||
assert _extract_results([{"id": 1}]) == [{"id": 1}]
|
||||
assert _extract_results({"nope": 1}) == []
|
||||
assert _extract_results(None) == []
|
||||
@@ -0,0 +1,255 @@
|
||||
"""Tests for the Mem0MemoryStore (Strands MemoryStore integration).
|
||||
|
||||
The store is exercised with a mocked Mem0ServiceClient, so no live Mem0 server
|
||||
(or the ``mem0ai`` SDK) is required.
|
||||
"""
|
||||
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
import pytest
|
||||
from strands.memory import MemoryEntry, MemoryStore
|
||||
from strands.memory.types import _has_method, _has_write_sink
|
||||
|
||||
from mem0_strands import Mem0MemoryStore
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_client():
|
||||
"""A mocked Mem0ServiceClient."""
|
||||
return MagicMock()
|
||||
|
||||
|
||||
def make_store(mock_client, **kwargs):
|
||||
"""Build a store wired to the mocked client (default scope: user_id=alex)."""
|
||||
kwargs.setdefault("user_id", "alex")
|
||||
return Mem0MemoryStore(client=mock_client, **kwargs)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Construction / protocol conformance
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_requires_a_scope():
|
||||
"""At least one of user_id / agent_id / run_id / app_id is mandatory."""
|
||||
with pytest.raises(ValueError, match="at least one of"):
|
||||
Mem0MemoryStore()
|
||||
|
||||
|
||||
def test_app_id_with_oss_config_rejected_at_construction():
|
||||
"""app_id is platform-only; pairing it with an OSS config fails at construction."""
|
||||
with pytest.raises(ValueError, match="platform-only"):
|
||||
Mem0MemoryStore(app_id="app1", config={"vector_store": {"provider": "qdrant"}})
|
||||
|
||||
|
||||
def test_scope_collects_only_set_fields(mock_client):
|
||||
"""Only the provided entity fields end up in the scope."""
|
||||
store = Mem0MemoryStore(client=mock_client, user_id="alex", agent_id="assistant")
|
||||
assert store.scope == {"user_id": "alex", "agent_id": "assistant"}
|
||||
|
||||
|
||||
def test_is_a_memory_store(mock_client):
|
||||
"""The store is a genuine MemoryStore subclass (MemoryStore is a
|
||||
non-runtime-checkable Protocol, so check the MRO rather than isinstance)."""
|
||||
store = make_store(mock_client)
|
||||
assert MemoryStore in type(store).__mro__
|
||||
|
||||
|
||||
def test_protocol_attributes_default(mock_client):
|
||||
"""Protocol attributes take sensible, writable-by-default values."""
|
||||
store = make_store(mock_client)
|
||||
assert store.name == "mem0"
|
||||
assert store.description is not None
|
||||
assert store.max_search_results is None
|
||||
assert store.writable is True
|
||||
assert store.extraction is None
|
||||
assert store.scope == {"user_id": "alex"}
|
||||
|
||||
|
||||
def test_protocol_attributes_override(mock_client):
|
||||
"""Config fields are honored."""
|
||||
store = make_store(
|
||||
mock_client,
|
||||
name="notes",
|
||||
description="d",
|
||||
max_search_results=3,
|
||||
writable=False,
|
||||
extraction=True,
|
||||
metadata={"team": "growth"},
|
||||
)
|
||||
assert store.name == "notes"
|
||||
assert store.max_search_results == 3
|
||||
assert store.writable is False
|
||||
assert store.extraction is True
|
||||
assert store.metadata == {"team": "growth"}
|
||||
|
||||
|
||||
def test_write_sink_detection(mock_client):
|
||||
"""Both `add` and `add_messages` are real sinks -- extraction defaults to
|
||||
Mem0's server-side path (add_messages), not a client-side ModelExtractor."""
|
||||
store = make_store(mock_client)
|
||||
assert _has_method(store, "search") is True
|
||||
assert _has_method(store, "add") is True
|
||||
assert _has_method(store, "add_messages") is True
|
||||
assert _has_method(store, "initialize") is False
|
||||
assert _has_method(store, "get_tools") is False
|
||||
assert _has_write_sink(store) is True
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# search
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_search_maps_to_memory_entries(mock_client):
|
||||
"""Mem0 hits are mapped to MemoryEntry with metadata preserved."""
|
||||
mock_client.search_memories.return_value = [
|
||||
{
|
||||
"id": "mem-1",
|
||||
"memory": "Alex prefers dark roast",
|
||||
"score": 0.91,
|
||||
"categories": ["preferences"],
|
||||
"created_at": "2026-07-02T00:00:00Z",
|
||||
"user_id": "alex",
|
||||
"metadata": {"category": "prefs"},
|
||||
}
|
||||
]
|
||||
store = make_store(mock_client)
|
||||
|
||||
results = await store.search("coffee")
|
||||
|
||||
assert len(results) == 1
|
||||
entry = results[0]
|
||||
assert isinstance(entry, MemoryEntry)
|
||||
assert entry.content == "Alex prefers dark roast"
|
||||
assert entry.metadata["id"] == "mem-1"
|
||||
assert entry.metadata["score"] == 0.91
|
||||
assert entry.metadata["categories"] == ["preferences"]
|
||||
assert entry.metadata["category"] == "prefs"
|
||||
|
||||
|
||||
async def test_search_default_top_k(mock_client):
|
||||
"""With no options and no configured max, the default top_k is used."""
|
||||
mock_client.search_memories.return_value = []
|
||||
store = make_store(mock_client)
|
||||
|
||||
await store.search("q")
|
||||
|
||||
mock_client.search_memories.assert_called_once_with("q", {"user_id": "alex"}, 5)
|
||||
|
||||
|
||||
async def test_search_options_override_top_k(mock_client):
|
||||
"""SearchOptions.max_search_results wins over the configured default."""
|
||||
mock_client.search_memories.return_value = []
|
||||
store = make_store(mock_client, max_search_results=3)
|
||||
|
||||
await store.search("q", {"max_search_results": 10})
|
||||
|
||||
mock_client.search_memories.assert_called_once_with("q", {"user_id": "alex"}, 10)
|
||||
|
||||
|
||||
async def test_search_config_top_k(mock_client):
|
||||
"""The configured max is used when options omit it."""
|
||||
mock_client.search_memories.return_value = []
|
||||
store = make_store(mock_client, max_search_results=7)
|
||||
|
||||
await store.search("q")
|
||||
|
||||
mock_client.search_memories.assert_called_once_with("q", {"user_id": "alex"}, 7)
|
||||
|
||||
|
||||
async def test_search_handles_missing_content(mock_client):
|
||||
"""A hit without memory text maps to an empty string, not None."""
|
||||
mock_client.search_memories.return_value = [{"id": "mem-2"}]
|
||||
store = make_store(mock_client)
|
||||
|
||||
results = await store.search("q")
|
||||
|
||||
assert results[0].content == ""
|
||||
assert results[0].metadata == {"id": "mem-2"}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# add / add_messages
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def test_add_writes_a_verbatim_fact(mock_client):
|
||||
"""add() forwards content, scope and merged metadata to store_memory."""
|
||||
stored = {"id": "mem-9"}
|
||||
mock_client.store_memory.return_value = stored
|
||||
store = make_store(mock_client, metadata={"team": "growth"})
|
||||
|
||||
result = await store.add("new fact", {"source": "chat"})
|
||||
|
||||
assert result == stored
|
||||
mock_client.store_memory.assert_called_once_with(
|
||||
"new fact", {"user_id": "alex"}, {"team": "growth", "source": "chat"}
|
||||
)
|
||||
|
||||
|
||||
async def test_add_without_metadata_uses_store_default(mock_client):
|
||||
"""With no per-call metadata, the store's default metadata is used."""
|
||||
store = make_store(mock_client, metadata={"team": "growth"})
|
||||
|
||||
await store.add("fact")
|
||||
|
||||
mock_client.store_memory.assert_called_once_with("fact", {"user_id": "alex"}, {"team": "growth"})
|
||||
|
||||
|
||||
async def test_add_messages_renders_content_blocks(mock_client):
|
||||
"""add_messages renders Strands content blocks to text before sending.
|
||||
|
||||
Strands hands content as list[ContentBlock] (a text block is ``{"text": ...}``);
|
||||
mem0 keeps only text parts, so the store must flatten each turn to a string.
|
||||
"""
|
||||
messages = [
|
||||
{"role": "user", "content": [{"text": "I love hiking"}]},
|
||||
{"role": "assistant", "content": [{"text": "Noted!"}]},
|
||||
]
|
||||
store = make_store(mock_client)
|
||||
|
||||
await store.add_messages(messages)
|
||||
|
||||
mock_client.store_messages.assert_called_once_with(
|
||||
[{"role": "user", "content": "I love hiking"}, {"role": "assistant", "content": "Noted!"}],
|
||||
{"user_id": "alex"},
|
||||
)
|
||||
|
||||
|
||||
async def test_add_messages_skips_empty_turns(mock_client):
|
||||
"""A turn with no text (a pure tool-use turn) renders to nothing and is not sent."""
|
||||
store = make_store(mock_client)
|
||||
|
||||
result = await store.add_messages([{"role": "assistant", "content": [{"toolUse": {"name": "x"}}]}])
|
||||
|
||||
assert result is None
|
||||
mock_client.store_messages.assert_not_called()
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# lazy client construction
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_client_constructed_lazily(monkeypatch):
|
||||
"""No Mem0ServiceClient is built until the client property is accessed."""
|
||||
calls = {"n": 0}
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, api_key=None, host=None, config=None, client=None):
|
||||
calls["n"] += 1
|
||||
self.api_key = api_key
|
||||
|
||||
monkeypatch.setattr("mem0_strands.store.Mem0ServiceClient", FakeClient)
|
||||
|
||||
store = Mem0MemoryStore(user_id="alex", api_key="m0-x")
|
||||
assert calls["n"] == 0 # not built yet
|
||||
|
||||
client = store.client
|
||||
assert calls["n"] == 1
|
||||
assert client.api_key == "m0-x"
|
||||
|
||||
# Second access reuses the same instance.
|
||||
assert store.client is client
|
||||
assert calls["n"] == 1
|
||||
Reference in New Issue
Block a user