feat(integrations): strands-mem0 | Mem0 as a native Strands MemoryStore (#7021)

This commit is contained in:
Himanshu
2026-08-22 19:04:24 +05:30
committed by GitHub
parent 9b565da8e3
commit 8d5b7865bd
21 changed files with 1561 additions and 2 deletions
@@ -0,0 +1,37 @@
# strands-mem0 (Python)
Persistent long-term memory for [Strands Agents](https://github.com/strands-agents/sdk-python),
backed by [Mem0](https://mem0.ai).
See the [repository README](../README.md) for full usage. Quick start:
```bash
pip install strands-mem0
```
As a `MemoryStore` that plugs into the agent loop (Strands >= 1.45):
```python
from strands import Agent
from strands.memory import MemoryManager
from strands_mem0 import Mem0MemoryStore
store = Mem0MemoryStore(user_id="alex", writable=True, extraction=True)
agent = Agent(memory_manager=MemoryManager(stores=[store]))
```
Set `MEM0_API_KEY` for the hosted platform, or pass `config=...` for a self-hosted Mem0 OSS backend.
## Local development
```bash
pip install hatch
hatch run test # pytest (mocked client, no live server)
hatch run prepare # format + lint + typecheck + test
```
## Release
Publish a GitHub release tagged `strands-mem0-v*` (e.g. `strands-mem0-v0.1.0`). The
release router (`.github/workflows/release.yml`) dispatches `strands-mem0-cd.yml`,
which builds the wheel and publishes it to PyPI via trusted publishing (OIDC).
@@ -0,0 +1,88 @@
[build-system]
requires = ["hatchling"]
build-backend = "hatchling.build"
[project]
name = "strands-mem0"
version = "0.1.0"
description = "Persistent long-term memory for Strands agents, backed by Mem0."
readme = "README.md"
requires-python = ">=3.10"
license = "Apache-2.0"
authors = [
{name = "Mem0", email = "founders@mem0.ai"}
]
keywords = ["strands", "strands-agents", "agents", "ai", "memory", "mem0", "vector-search", "personalization"]
classifiers = [
"Development Status :: 4 - Beta",
"Intended Audience :: Developers",
"License :: OSI Approved :: Apache Software License",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12",
"Topic :: Scientific/Engineering :: Artificial Intelligence",
]
# strands-agents>=1.45.0: first release shipping the `strands.memory` module
# (MemoryStore / MemoryManager) that Mem0MemoryStore implements.
dependencies = [
"strands-agents>=1.45.0",
"mem0ai>=2.0.11",
]
[project.urls]
Homepage = "https://mem0.ai"
Documentation = "https://github.com/mem0ai/mem0/tree/main/integrations/strands-mem0#readme"
Repository = "https://github.com/mem0ai/mem0"
Issues = "https://github.com/mem0ai/mem0/issues"
[project.optional-dependencies]
dev = [
"pytest>=8.0.0,<9.0.0",
"pytest-asyncio>=0.25.0,<1.0.0",
"ruff>=0.11.0,<1.0.0",
"mypy>=1.15.0,<2.0.0",
"hatch",
]
[tool.hatch.build.targets.wheel]
packages = ["src/strands_mem0"]
[tool.hatch.envs.default]
dependencies = [
"pytest>=8.0.0,<9.0.0",
"pytest-asyncio>=0.25.0,<1.0.0",
"ruff>=0.11.0,<1.0.0",
"mypy>=1.15.0,<2.0.0",
]
[tool.hatch.envs.default.scripts]
test = "pytest {args}"
lint = "ruff check src tests"
format = "ruff format src tests"
typecheck = "mypy src"
prepare = ["format", "lint", "typecheck", "test"]
[tool.ruff]
line-length = 120
include = ["src/**/*.py", "tests/**/*.py"]
[tool.ruff.lint]
select = [
"E", # pycodestyle
"F", # pyflakes
"I", # isort
"B", # flake8-bugbear
]
[tool.mypy]
python_version = "3.10"
warn_return_any = true
warn_unused_configs = true
ignore_missing_imports = true
[tool.pytest.ini_options]
testpaths = ["tests"]
pythonpath = ["src"]
asyncio_mode = "auto"
@@ -0,0 +1,26 @@
"""Strands Mem0 -- persistent long-term memory for Strands agents, backed by Mem0.
:class:`Mem0MemoryStore` is a Strands ``MemoryStore`` that plugs into the agent
loop via a :class:`~strands.memory.MemoryManager`, with automatic memory injection
and extraction. It implements both write sinks, so ``extraction`` uses Mem0's
server-side extraction (no extra model call).
For the explicit, model-called tool (``store`` / ``retrieve`` / ``get`` / ``delete``),
use the ``mem0_memory`` tool from ``strands-agents-tools``; a store and the tool can
share one Mem0 backend and namespace.
"""
from importlib.metadata import PackageNotFoundError, version
from strands_mem0.client import Mem0ServiceClient
from strands_mem0.store import Mem0MemoryStore
__all__ = [
"Mem0MemoryStore",
"Mem0ServiceClient",
]
try:
__version__ = version("strands-mem0")
except PackageNotFoundError: # pragma: no cover - only when running from a source tree
__version__ = "0.0.0+unknown"
@@ -0,0 +1,176 @@
"""A thin wrapper around the Mem0 SDK used by :class:`~strands_mem0.store.Mem0MemoryStore`.
Both Mem0 backends -- the hosted platform (:class:`mem0.MemoryClient`) and
self-hosted OSS (:class:`mem0.Memory`) -- expose the same call shape to the store:
- **search** takes the entity scope inside a ``filters`` dict plus ``top_k``.
- **add** takes the entity scope as top-level keyword arguments.
The wrapper hides the two remaining differences:
- ``app_id`` is a platform-only scope; OSS ``Memory.add`` has no ``app_id``
parameter, so it is rejected up front for the OSS backend rather than surfacing
as a ``TypeError`` mid-call.
- the telemetry ``source`` tag is attached to platform writes only (OSS
``Memory.add`` has a fixed signature and would reject an unknown kwarg).
"""
from __future__ import annotations
import inspect
import os
from typing import Any
# Only the synchronous platform client is supported. ``AsyncMemoryClient``'s
# ``add`` / ``search`` are coroutine functions, so ``asyncio.to_thread`` would hand
# back an un-awaited coroutine and every write would silently no-op; it is rejected
# in ``__init__`` rather than listed here.
_PLATFORM_CLIENTS = {"MemoryClient"}
# Tags platform writes so Mem0's backend attributes the memory to this integration
# in telemetry (recognized values live in the backend's KNOWN_EVENT_SOURCES
# allowlist; unknown ones bucket into "OTHERS"). Platform only.
_SOURCE = "STRANDS"
def _is_platform_client(client: Any) -> bool:
"""Whether ``client`` is a hosted Mem0 platform client (vs an OSS ``Memory``)."""
return type(client).__name__ in _PLATFORM_CLIENTS
def _is_async_client(client: Any) -> bool:
"""Whether ``client``'s ``add`` / ``search`` are coroutine functions."""
return inspect.iscoroutinefunction(getattr(client, "add", None)) or inspect.iscoroutinefunction(
getattr(client, "search", None)
)
class Mem0ServiceClient:
"""Thin wrapper around the Mem0 SDK for the memory store.
Exactly one backend is selected at construction time:
- ``client`` given: use it as-is (a :class:`mem0.MemoryClient` or
:class:`mem0.Memory`); mainly for testing and advanced/OSS setups.
- ``config`` given: build a self-hosted :class:`mem0.Memory` from it.
- otherwise: build a hosted :class:`mem0.MemoryClient` from ``api_key`` /
``$MEM0_API_KEY`` (and optional ``host``).
"""
def __init__(
self,
api_key: str | None = None,
host: str | None = None,
config: dict[str, Any] | None = None,
client: Any | None = None,
) -> None:
"""Initialize the Mem0 client.
Args:
api_key: Mem0 platform API key. Falls back to ``$MEM0_API_KEY``.
host: Mem0 platform base URL. Defaults to the SDK default
(``https://api.mem0.ai``).
config: A Mem0 OSS config dict; when given, a self-hosted
:class:`mem0.Memory` is built instead of the platform client.
client: A pre-built Mem0 client to use directly (platform or OSS).
Raises:
ValueError: If ``client`` is an async Mem0 client (its coroutines
would never be awaited off the worker thread).
"""
if client is not None:
if _is_async_client(client):
raise ValueError(
"Async Mem0 clients are not supported. Pass a synchronous "
"mem0.MemoryClient (or a mem0.Memory / config): the store runs the "
"SDK in a worker thread, so an async client's coroutines would "
"never be awaited and every write would silently no-op."
)
self.mem0 = client
self.is_platform = _is_platform_client(client)
return
if config is not None:
try:
from mem0 import Memory
except ImportError as err: # pragma: no cover - exercised via install docs
raise ImportError(
"The mem0ai package is required. Install it with: pip install 'strands-mem0'"
) from err
self.mem0 = Memory.from_config(config)
self.is_platform = False
return
try:
from mem0 import MemoryClient
except ImportError as err: # pragma: no cover - exercised via install docs
raise ImportError("The mem0ai package is required. Install it with: pip install 'strands-mem0'") from err
api_key = api_key or os.environ.get("MEM0_API_KEY")
# MemoryClient(host=None) would override the SDK default with None, so only
# pass host when the caller actually set one.
self.mem0 = MemoryClient(api_key=api_key, host=host) if host else MemoryClient(api_key=api_key)
self.is_platform = True
def _check_scope(self, scope: dict[str, str]) -> None:
"""Reject scope the selected backend cannot honor.
``app_id`` exists only on the platform; the OSS ``Memory`` API has no
``app_id`` parameter, so we fail loudly here rather than let it surface as
a ``TypeError`` on ``add`` or silently miss on ``search``.
"""
if not self.is_platform and "app_id" in scope:
raise ValueError(
"app_id is a Mem0 platform-only scope. The OSS backend supports "
"user_id, agent_id, and run_id; drop app_id or use the platform client."
)
def _write_extras(self) -> dict[str, str]:
"""Extra kwargs attached to platform writes: the telemetry ``source`` tag."""
return {"source": _SOURCE} if self.is_platform else {}
def store_memory(
self,
content: str,
scope: dict[str, str],
metadata: dict[str, Any] | None = None,
) -> Any:
"""Store one discrete fact verbatim (``infer=False``).
Used by the store's ``add`` sink -- the content is already a distilled fact
(from the ``add_memory`` tool or a client-side extractor), so Mem0's own
extraction is skipped to preserve it exactly.
"""
self._check_scope(scope)
return self.mem0.add(content, metadata=metadata, infer=False, **self._write_extras(), **scope)
def store_messages(self, messages: list[dict[str, Any]], scope: dict[str, str]) -> Any:
"""Hand rendered conversation turns to Mem0 for server-side extraction (``infer=True``).
Used by the store's ``add_messages`` sink. Mem0 extracts and de-duplicates
facts on the server, so no client-side model call is needed.
"""
self._check_scope(scope)
return self.mem0.add(messages, infer=True, **self._write_extras(), **scope)
def search_memories(self, query: str, scope: dict[str, str], top_k: int) -> list[dict[str, Any]]:
"""Semantic recall scoped to the store's entity.
Both backends take the scope inside ``filters`` and honor ``top_k``; the
response is normalized to a plain list of memory dicts.
"""
self._check_scope(scope)
response = self.mem0.search(query, filters=dict(scope), top_k=top_k)
return _extract_results(response)
def _extract_results(response: Any) -> list[dict[str, Any]]:
"""Normalize a Mem0 search response to a list of memory dicts.
Mem0 returns ``{"results": [...]}`` (v1.1) or, on older paths, a bare list.
"""
if isinstance(response, dict):
results = response.get("results", [])
return list(results) if isinstance(results, list) else []
if isinstance(response, list):
return response
return []
@@ -0,0 +1,222 @@
"""A Strands ``MemoryStore`` backed by Mem0.
A memory store gives a Strands agent cross-session recall: a
:class:`~strands.memory.MemoryManager` searches it to recall facts and, when
writable, writes new ones -- either directly or via automatic extraction from the
conversation. Unlike the ``mem0_memory`` tool (which the model calls explicitly),
a store plugs into the agent loop out of the box, with memory injection and
extraction triggers handled by the manager.
``Mem0MemoryStore`` implements both write sinks, which is what sets it apart from a
vector-DB-style store:
- :meth:`add` writes a single distilled fact verbatim (``infer=False``). This is
the sink for the ``add_memory`` tool and for a client-side extractor.
- :meth:`add_messages` renders raw conversation turns to text and hands them to
Mem0 for **server-side extraction** (``infer=True``). Because this sink exists,
enabling ``extraction`` routes messages straight to Mem0's own extraction
pipeline -- no extra client-side model call, and Mem0's de-duplication applies.
Example:
```python
from strands import Agent
from strands.memory import MemoryManager
from strands_mem0 import Mem0MemoryStore
# Recall + write, with Mem0 extracting facts from the conversation server-side.
store = Mem0MemoryStore(user_id="alex", writable=True, extraction=True)
agent = Agent(memory_manager=MemoryManager(stores=[store]))
```
Configure the hosted platform via the ``api_key`` argument or the ``MEM0_API_KEY``
environment variable, or pass a Mem0 OSS ``config`` dict for a self-hosted backend.
``app_id`` scope is platform-only.
"""
from __future__ import annotations
import asyncio
from typing import Any
from strands.memory import AddMessagesContext, MemoryEntry, MemoryStore, SearchOptions
from strands.types.content import Message
from strands_mem0.client import Mem0ServiceClient
DEFAULT_MAX_SEARCH_RESULTS = 5
# Entity fields that scope a memory in Mem0. At least one must be set.
_SCOPE_FIELDS = ("user_id", "agent_id", "run_id", "app_id")
class Mem0MemoryStore(MemoryStore):
"""A Strands :class:`~strands.memory.MemoryStore` backed by Mem0.
Implements :meth:`search` (semantic recall), :meth:`add` (a verbatim
single-fact write sink) and :meth:`add_messages` (raw-message ingestion with
Mem0 server-side extraction). Because ``add_messages`` is implemented, enabling
``extraction`` uses Mem0's server-side extraction rather than a client-side
model call.
"""
def __init__(
self,
*,
user_id: str | None = None,
agent_id: str | None = None,
run_id: str | None = None,
app_id: str | None = None,
name: str = "mem0",
description: str | None = "Persistent long-term memory backed by Mem0.",
max_search_results: int | None = None,
writable: bool = True,
extraction: Any = None,
metadata: dict[str, Any] | None = None,
api_key: str | None = None,
host: str | None = None,
config: dict[str, Any] | None = None,
client: Mem0ServiceClient | None = None,
) -> None:
"""Initialize the store.
Args:
user_id: Mem0 user namespace that owns the memories.
agent_id: Mem0 agent namespace.
run_id: Mem0 run/session namespace.
app_id: Mem0 app namespace (platform only).
name: Unique store identifier, used to target it in tools.
description: Human-readable description, included in tool descriptions.
max_search_results: Default maximum results per search.
writable: Whether the store accepts writes.
extraction: Automatic-extraction config (``bool | ExtractionConfig``).
metadata: Default metadata merged into every write.
api_key: Mem0 platform API key (defaults to ``$MEM0_API_KEY``).
host: Mem0 platform base URL.
config: Mem0 OSS config dict for a self-hosted backend.
client: A pre-built :class:`~strands_mem0.client.Mem0ServiceClient`
(for testing, or to wrap your own raw Mem0 client via
``Mem0ServiceClient(client=...)``); when omitted, one is
constructed lazily on first use from ``api_key`` / ``config``.
Raises:
ValueError: If no entity scope (``user_id`` / ``agent_id`` / ``run_id``
/ ``app_id``) is provided.
"""
scope = {
"user_id": user_id,
"agent_id": agent_id,
"run_id": run_id,
"app_id": app_id,
}
self.scope = {key: value for key, value in scope.items() if value}
if not self.scope:
raise ValueError("Mem0MemoryStore requires at least one of user_id, agent_id, run_id, or app_id")
# app_id is platform-only. When a self-hosted OSS backend is requested via
# `config`, fail at construction rather than as a TypeError on the first
# write (OSS Memory.add has no app_id). The injected-client OSS case is
# caught in Mem0ServiceClient, which is the only place that knows the backend.
if "app_id" in self.scope and config is not None:
raise ValueError(
"app_id is a Mem0 platform-only scope and cannot be used with a self-hosted "
"config (OSS Memory has no app_id). Drop app_id or use the platform backend."
)
# MemoryStore Protocol attributes.
self.name = name
self.description = description
self.max_search_results = max_search_results
self.writable = writable
self.extraction = extraction
# Mem0-specific configuration.
self.metadata = metadata
self._api_key = api_key
self._host = host
self._config = config
self._client = client
@property
def client(self) -> Mem0ServiceClient:
"""The Mem0 service client, constructed lazily on first use.
Note: constructing the underlying SDK client can block (the platform client
validates the API key over HTTP; the OSS client builds embedders / vector
stores), so first use is deferred and always happens inside a worker thread
via :func:`asyncio.to_thread`, never on the event loop.
"""
if self._client is None:
self._client = Mem0ServiceClient(api_key=self._api_key, host=self._host, config=self._config)
return self._client
async def search(self, query: str, options: SearchOptions | None = None) -> list[MemoryEntry]:
"""Search Mem0 for entries matching ``query``, ordered by relevance."""
top_k = options.get("max_search_results") if options is not None else None
if top_k is None:
top_k = self.max_search_results
if top_k is None:
top_k = DEFAULT_MAX_SEARCH_RESULTS
# ``self.client`` is resolved inside the thread so lazy construction (a
# blocking call) does not run on the event loop.
memories = await asyncio.to_thread(lambda: self.client.search_memories(query, self.scope, top_k))
return [self._to_entry(memory) for memory in memories]
async def add(self, content: str, metadata: dict[str, Any] | None = None) -> Any:
"""Write a single distilled fact to Mem0 verbatim (``infer=False``).
Extraction writes are at-least-once, so this tolerates duplicate content;
Mem0 de-duplicates on the server.
"""
merged = self._merge_metadata(metadata)
return await asyncio.to_thread(lambda: self.client.store_memory(content, self.scope, merged))
async def add_messages(self, messages: list[Message], context: AddMessagesContext | None = None) -> Any:
"""Ingest raw conversation turns for Mem0 server-side extraction (``infer=True``).
A Strands ``Message.content`` is a list of content blocks (a text block is
``{"text": "..."}``); Mem0 keeps only ``{"type": "text"}`` parts, so the raw
blocks would be dropped. We render each turn's text blocks to a string and
skip turns that render empty (a pure tool-use / tool-result turn), so nothing
silently no-ops.
"""
payload: list[dict[str, str]] = []
for message in messages:
text = self._render_content(message.get("content"))
if text:
payload.append({"role": message["role"], "content": text})
if not payload:
return None
return await asyncio.to_thread(lambda: self.client.store_messages(payload, self.scope))
@staticmethod
def _render_content(content: Any) -> str:
"""Flatten a Strands message ``content`` to plain text.
Accepts either a string or a list of content blocks; joins the text of
every ``{"text": ...}`` block and ignores tool-use / image / other blocks.
"""
if isinstance(content, str):
return content
if isinstance(content, list):
return "\n".join(part["text"] for part in content if isinstance(part, dict) and part.get("text"))
return ""
def _merge_metadata(self, metadata: dict[str, Any] | None) -> dict[str, Any] | None:
"""Merge per-call metadata over the store's default metadata."""
if self.metadata and metadata:
return {**self.metadata, **metadata}
return metadata or self.metadata
@staticmethod
def _to_entry(memory: dict[str, Any]) -> MemoryEntry:
"""Map a Mem0 memory dict to a Strands :class:`~strands.memory.MemoryEntry`."""
content = memory.get("memory") or memory.get("content") or ""
metadata: dict[str, Any] = {}
for key in ("id", "score", "categories", "created_at", "updated_at", *_SCOPE_FIELDS):
value = memory.get(key)
if value is not None:
metadata[key] = value
extra = memory.get("metadata")
if isinstance(extra, dict):
metadata.update(extra)
return MemoryEntry(content=content, metadata=metadata or None)
@@ -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 strands_mem0.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 strands_mem0 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("strands_mem0.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