Python: PR1 — New session and context provider types (side-by-side) (#3763)

* PR1: Add core context provider types and tests

New types in _sessions.py (no changes to existing code):
- SessionContext: per-invocation state with extend_messages/get_messages/
  extend_instructions/extend_tools and read-only response property
- _ContextProviderBase: base class with before_run/after_run hooks
- _HistoryProviderBase: storage base with load/store flags, abstract
  get_messages/save_messages, default before_run/after_run
- AgentSession: lightweight session with state dict, to_dict/from_dict
- InMemoryHistoryProvider: built-in provider storing in session.state

35 unit tests covering all classes and configuration flags.

* feat: keyword-only params, stateless InMemoryHistoryProvider, deep serialization

- Make before_run/after_run parameters keyword-only
- InMemoryHistoryProvider stores ChatMessage objects directly (no per-cycle serialization)
- Deep serialization via to_dict/from_dict only at session boundary
- State type registry for automatic deserialization of registered types
- Updated tests for new serialization approach

* feat: add new-pattern provider implementations for external packages

- _RedisContextProvider(BaseContextProvider) - Redis search/vector context
- _RedisHistoryProvider(BaseHistoryProvider) - Redis-backed message storage
- _Mem0ContextProvider(BaseContextProvider) - Mem0 semantic memory
- _AzureAISearchContextProvider(BaseContextProvider) - Azure AI Search (semantic + agentic)

All use temporary _ prefix names for side-by-side coexistence with existing providers.
Will be renamed in PR2 when old ContextProvider/ChatMessageStore are removed.

* test: add tests for new-pattern provider implementations

- 32 tests for _RedisContextProvider and _RedisHistoryProvider
- 29 tests for _Mem0ContextProvider
- 17 tests for _AzureAISearchContextProvider

* fix: address PR review comments and CI failures

- Move module docstring before imports in _sessions.py (review comment)
- Import TYPE_CHECKING unconditionally in Redis _context_provider.py (NameError on Python <3.12)
- Fix Mem0 test_init_auto_creates_client_when_none to patch at class level

* feat: add source attribution to extend_messages

Set attribution marker in additional_properties for each message
added via extend_messages(), matching the tool attribution pattern.
Uses setdefault to preserve any existing attribution.

* refactor: make attribution value a dict with source_id key

* add attribution and use sets for filters

* Add source_type to message attribution and copy messages in extend_messages

- SessionContext.extend_messages now accepts source as str or object with
  source_id attribute; when an object is passed, its class name is recorded
  as source_type in the attribution dict
- Messages are shallow-copied before attribution is added so callers'
  original objects are never mutated
- Filter framework-internal keys (attribution) from A2A wire metadata
  to prevent leaking internal state over the wire

* fix: correct mypy type: ignore comment from union-attr to attr-defined

* set attribution to _attribution

* adjusted naming of bools
This commit is contained in:
Eduard van Valkenburg
2026-02-10 22:19:15 +01:00
committed by GitHub
Unverified
parent ccff3d3452
commit ac0e6b0ee1
13 changed files with 3494 additions and 2 deletions
@@ -8,6 +8,7 @@ import os
if os.environ.get("MEM0_TELEMETRY") is None:
os.environ["MEM0_TELEMETRY"] = "false"
from ._context_provider import _Mem0ContextProvider
from ._provider import Mem0Provider
try:
@@ -17,5 +18,6 @@ except importlib.metadata.PackageNotFoundError:
__all__ = [
"Mem0Provider",
"_Mem0ContextProvider",
"__version__",
]
@@ -0,0 +1,193 @@
# Copyright (c) Microsoft. All rights reserved.
"""New-pattern Mem0 context provider using BaseContextProvider.
This module provides ``_Mem0ContextProvider``, a side-by-side implementation of
:class:`Mem0Provider` built on the new :class:`BaseContextProvider` hooks pattern.
It will be renamed to ``Mem0ContextProvider`` in PR2 when the old class is removed.
"""
from __future__ import annotations
import sys
from contextlib import AbstractAsyncContextManager
from typing import TYPE_CHECKING, Any
from agent_framework import ChatMessage
from agent_framework._sessions import AgentSession, BaseContextProvider, SessionContext
from agent_framework.exceptions import ServiceInitializationError
from mem0 import AsyncMemory, AsyncMemoryClient
if sys.version_info >= (3, 11):
from typing import NotRequired, Self, TypedDict # pragma: no cover
else:
from typing_extensions import NotRequired, Self, TypedDict # pragma: no cover
if TYPE_CHECKING:
from agent_framework._agents import SupportsAgentRun
class _MemorySearchResponse_v1_1(TypedDict):
results: list[dict[str, Any]]
relations: NotRequired[list[dict[str, Any]]]
_MemorySearchResponse_v2 = list[dict[str, Any]]
class _Mem0ContextProvider(BaseContextProvider):
"""Mem0 context provider using the new BaseContextProvider hooks pattern.
Integrates Mem0 for persistent semantic memory, searching and storing
memories via the Mem0 API. This is the new-pattern equivalent of
:class:`Mem0Provider`.
Note:
This class uses a temporary ``_`` prefix to coexist with the existing
:class:`Mem0Provider`. It will be renamed to ``Mem0ContextProvider``
in PR2.
"""
DEFAULT_CONTEXT_PROMPT = "## Memories\nConsider the following memories when answering user questions:"
def __init__(
self,
source_id: str,
mem0_client: AsyncMemory | AsyncMemoryClient | None = None,
api_key: str | None = None,
application_id: str | None = None,
agent_id: str | None = None,
user_id: str | None = None,
*,
context_prompt: str | None = None,
) -> None:
"""Initialize the Mem0 context provider.
Args:
source_id: Unique identifier for this provider instance.
mem0_client: A pre-created Mem0 MemoryClient or None to create a default client.
api_key: The API key for authenticating with the Mem0 API.
application_id: The application ID for scoping memories.
agent_id: The agent ID for scoping memories.
user_id: The user ID for scoping memories.
context_prompt: The prompt to prepend to retrieved memories.
"""
super().__init__(source_id)
should_close_client = False
if mem0_client is None:
mem0_client = AsyncMemoryClient(api_key=api_key)
should_close_client = True
self.api_key = api_key
self.application_id = application_id
self.agent_id = agent_id
self.user_id = user_id
self.context_prompt = context_prompt or self.DEFAULT_CONTEXT_PROMPT
self.mem0_client = mem0_client
self._should_close_client = should_close_client
async def __aenter__(self) -> Self:
"""Async context manager entry."""
if self.mem0_client and isinstance(self.mem0_client, AbstractAsyncContextManager):
await self.mem0_client.__aenter__()
return self
async def __aexit__(self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: Any) -> None:
"""Async context manager exit."""
if self._should_close_client and self.mem0_client and isinstance(self.mem0_client, AbstractAsyncContextManager):
await self.mem0_client.__aexit__(exc_type, exc_val, exc_tb)
# -- Hooks pattern ---------------------------------------------------------
async def before_run(
self,
*,
agent: SupportsAgentRun,
session: AgentSession,
context: SessionContext,
state: dict[str, Any],
) -> None:
"""Search Mem0 for relevant memories and add to the session context."""
self._validate_filters()
input_text = "\n".join(msg.text for msg in context.input_messages if msg and msg.text and msg.text.strip())
if not input_text.strip():
return
filters = self._build_filters(session_id=context.session_id)
search_response: _MemorySearchResponse_v1_1 | _MemorySearchResponse_v2 = await self.mem0_client.search( # type: ignore[misc]
query=input_text,
filters=filters,
)
if isinstance(search_response, list):
memories = search_response
elif isinstance(search_response, dict) and "results" in search_response:
memories = search_response["results"]
else:
memories = [search_response]
line_separated_memories = "\n".join(memory.get("memory", "") for memory in memories)
if line_separated_memories:
context.extend_messages(
self.source_id,
[ChatMessage(role="user", text=f"{self.context_prompt}\n{line_separated_memories}")],
)
async def after_run(
self,
*,
agent: SupportsAgentRun,
session: AgentSession,
context: SessionContext,
state: dict[str, Any],
) -> None:
"""Store request/response messages to Mem0 for future retrieval."""
self._validate_filters()
messages_to_store: list[ChatMessage] = list(context.input_messages)
if context.response and context.response.messages:
messages_to_store.extend(context.response.messages)
def get_role_value(role: Any) -> str:
return role.value if hasattr(role, "value") else str(role)
messages: list[dict[str, str]] = [
{"role": get_role_value(message.role), "content": message.text}
for message in messages_to_store
if get_role_value(message.role) in {"user", "assistant", "system"} and message.text and message.text.strip()
]
if messages:
await self.mem0_client.add( # type: ignore[misc]
messages=messages,
user_id=self.user_id,
agent_id=self.agent_id,
run_id=context.session_id,
metadata={"application_id": self.application_id},
)
# -- Internal methods ------------------------------------------------------
def _validate_filters(self) -> None:
"""Validates that at least one filter is provided."""
if not self.agent_id and not self.user_id and not self.application_id:
raise ServiceInitializationError(
"At least one of the filters: agent_id, user_id, or application_id is required."
)
def _build_filters(self, *, session_id: str | None = None) -> dict[str, Any]:
"""Build search filters from initialization parameters."""
filters: dict[str, Any] = {}
if self.user_id:
filters["user_id"] = self.user_id
if self.agent_id:
filters["agent_id"] = self.agent_id
if session_id:
filters["run_id"] = session_id
if self.application_id:
filters["app_id"] = self.application_id
return filters
__all__ = ["_Mem0ContextProvider"]
@@ -0,0 +1,352 @@
# Copyright (c) Microsoft. All rights reserved.
# pyright: reportPrivateUsage=false
from __future__ import annotations
from unittest.mock import AsyncMock, patch
import pytest
from agent_framework import AgentResponse, ChatMessage
from agent_framework._sessions import AgentSession, SessionContext
from agent_framework.exceptions import ServiceInitializationError
from agent_framework_mem0._context_provider import _Mem0ContextProvider
@pytest.fixture
def mock_mem0_client() -> AsyncMock:
"""Create a mock Mem0 AsyncMemoryClient."""
from mem0 import AsyncMemoryClient
mock_client = AsyncMock(spec=AsyncMemoryClient)
mock_client.add = AsyncMock()
mock_client.search = AsyncMock()
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
mock_client.__aexit__ = AsyncMock()
return mock_client
# -- Initialization tests ------------------------------------------------------
class TestInit:
"""Test _Mem0ContextProvider initialization."""
def test_init_with_all_params(self, mock_mem0_client: AsyncMock) -> None:
provider = _Mem0ContextProvider(
source_id="mem0",
mem0_client=mock_mem0_client,
api_key="key-123",
application_id="app1",
agent_id="agent1",
user_id="user1",
context_prompt="Custom prompt",
)
assert provider.source_id == "mem0"
assert provider.api_key == "key-123"
assert provider.application_id == "app1"
assert provider.agent_id == "agent1"
assert provider.user_id == "user1"
assert provider.context_prompt == "Custom prompt"
assert provider.mem0_client is mock_mem0_client
assert provider._should_close_client is False
def test_init_default_context_prompt(self, mock_mem0_client: AsyncMock) -> None:
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
assert provider.context_prompt == _Mem0ContextProvider.DEFAULT_CONTEXT_PROMPT
def test_init_auto_creates_client_when_none(self) -> None:
"""When no client is provided, a default AsyncMemoryClient is created and flagged for closing."""
with (
patch("mem0.client.main.AsyncMemoryClient.__init__", return_value=None) as mock_init,
patch("mem0.client.main.AsyncMemoryClient._validate_api_key", return_value=None),
):
provider = _Mem0ContextProvider(source_id="mem0", api_key="test-key", user_id="u1")
mock_init.assert_called_once_with(api_key="test-key")
assert provider._should_close_client is True
def test_provided_client_not_flagged_for_close(self, mock_mem0_client: AsyncMock) -> None:
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
assert provider._should_close_client is False
# -- before_run tests ----------------------------------------------------------
class TestBeforeRun:
"""Test before_run hook."""
async def test_memories_added_to_context(self, mock_mem0_client: AsyncMock) -> None:
"""Mocked mem0 search returns memories → messages added to context with prompt."""
mock_mem0_client.search.return_value = [
{"memory": "User likes Python"},
{"memory": "User prefers dark mode"},
]
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="Hello")], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
mock_mem0_client.search.assert_awaited_once()
assert "mem0" in ctx.context_messages
added = ctx.context_messages["mem0"]
assert len(added) == 1
assert "User likes Python" in added[0].text # type: ignore[operator]
assert "User prefers dark mode" in added[0].text # type: ignore[operator]
assert provider.context_prompt in added[0].text # type: ignore[operator]
async def test_empty_input_skips_search(self, mock_mem0_client: AsyncMock) -> None:
"""Empty input messages → no search performed."""
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="")], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
mock_mem0_client.search.assert_not_awaited()
assert "mem0" not in ctx.context_messages
async def test_empty_search_results_no_messages(self, mock_mem0_client: AsyncMock) -> None:
"""Empty search results → no messages added."""
mock_mem0_client.search.return_value = []
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="test")], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
assert "mem0" not in ctx.context_messages
async def test_validates_filters_before_search(self, mock_mem0_client: AsyncMock) -> None:
"""Raises ServiceInitializationError when no filters."""
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client)
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="test")], session_id="s1")
with pytest.raises(ServiceInitializationError, match="At least one of the filters"):
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
async def test_v1_1_response_format(self, mock_mem0_client: AsyncMock) -> None:
"""Search response in v1.1 dict format with 'results' key."""
mock_mem0_client.search.return_value = {"results": [{"memory": "remembered fact"}]}
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="test")], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
added = ctx.context_messages["mem0"]
assert "remembered fact" in added[0].text # type: ignore[operator]
async def test_search_query_combines_input_messages(self, mock_mem0_client: AsyncMock) -> None:
"""Multiple input messages are joined for the search query."""
mock_mem0_client.search.return_value = []
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(
input_messages=[
ChatMessage(role="user", text="Hello"),
ChatMessage(role="user", text="World"),
],
session_id="s1",
)
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
call_kwargs = mock_mem0_client.search.call_args.kwargs
assert call_kwargs["query"] == "Hello\nWorld"
# -- after_run tests -----------------------------------------------------------
class TestAfterRun:
"""Test after_run hook."""
async def test_stores_input_and_response(self, mock_mem0_client: AsyncMock) -> None:
"""Stores input+response messages to mem0 via client.add."""
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="question")], session_id="s1")
ctx._response = AgentResponse(messages=[ChatMessage(role="assistant", text="answer")])
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
mock_mem0_client.add.assert_awaited_once()
call_kwargs = mock_mem0_client.add.call_args.kwargs
assert call_kwargs["messages"] == [
{"role": "user", "content": "question"},
{"role": "assistant", "content": "answer"},
]
assert call_kwargs["user_id"] == "u1"
assert call_kwargs["run_id"] == "s1"
async def test_only_stores_user_assistant_system(self, mock_mem0_client: AsyncMock) -> None:
"""Only stores user/assistant/system messages with text."""
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(
input_messages=[
ChatMessage(role="user", text="hello"),
ChatMessage(role="tool", text="tool output"),
],
session_id="s1",
)
ctx._response = AgentResponse(messages=[ChatMessage(role="assistant", text="reply")])
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
call_kwargs = mock_mem0_client.add.call_args.kwargs
roles = [m["role"] for m in call_kwargs["messages"]]
assert "tool" not in roles
assert roles == ["user", "assistant"]
async def test_skips_empty_messages(self, mock_mem0_client: AsyncMock) -> None:
"""Skips messages with empty text."""
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(
input_messages=[
ChatMessage(role="user", text=""),
ChatMessage(role="user", text=" "),
],
session_id="s1",
)
ctx._response = AgentResponse(messages=[])
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
mock_mem0_client.add.assert_not_awaited()
async def test_uses_session_id_as_run_id(self, mock_mem0_client: AsyncMock) -> None:
"""Uses session_id as run_id."""
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="hi")], session_id="my-session")
ctx._response = AgentResponse(messages=[ChatMessage(role="assistant", text="hey")])
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
assert mock_mem0_client.add.call_args.kwargs["run_id"] == "my-session"
async def test_validates_filters(self, mock_mem0_client: AsyncMock) -> None:
"""Raises ServiceInitializationError when no filters."""
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client)
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="hi")], session_id="s1")
ctx._response = AgentResponse(messages=[ChatMessage(role="assistant", text="hey")])
with pytest.raises(ServiceInitializationError, match="At least one of the filters"):
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
async def test_stores_with_application_id_metadata(self, mock_mem0_client: AsyncMock) -> None:
"""application_id is passed in metadata."""
provider = _Mem0ContextProvider(
source_id="mem0", mem0_client=mock_mem0_client, user_id="u1", application_id="app1"
)
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="hi")], session_id="s1")
ctx._response = AgentResponse(messages=[])
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
assert mock_mem0_client.add.call_args.kwargs["metadata"] == {"application_id": "app1"}
# -- _validate_filters tests --------------------------------------------------
class TestValidateFilters:
"""Test _validate_filters method."""
def test_raises_when_no_filters(self, mock_mem0_client: AsyncMock) -> None:
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client)
with pytest.raises(ServiceInitializationError, match="At least one of the filters"):
provider._validate_filters()
def test_passes_with_user_id(self, mock_mem0_client: AsyncMock) -> None:
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
provider._validate_filters() # should not raise
def test_passes_with_agent_id(self, mock_mem0_client: AsyncMock) -> None:
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, agent_id="a1")
provider._validate_filters()
def test_passes_with_application_id(self, mock_mem0_client: AsyncMock) -> None:
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, application_id="app1")
provider._validate_filters()
# -- _build_filters tests -----------------------------------------------------
class TestBuildFilters:
"""Test _build_filters method."""
def test_user_id_only(self, mock_mem0_client: AsyncMock) -> None:
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
assert provider._build_filters() == {"user_id": "u1"}
def test_all_params(self, mock_mem0_client: AsyncMock) -> None:
provider = _Mem0ContextProvider(
source_id="mem0",
mem0_client=mock_mem0_client,
user_id="u1",
agent_id="a1",
application_id="app1",
)
assert provider._build_filters(session_id="sess1") == {
"user_id": "u1",
"agent_id": "a1",
"run_id": "sess1",
"app_id": "app1",
}
def test_excludes_none_values(self, mock_mem0_client: AsyncMock) -> None:
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
filters = provider._build_filters()
assert "agent_id" not in filters
assert "run_id" not in filters
assert "app_id" not in filters
def test_session_id_mapped_to_run_id(self, mock_mem0_client: AsyncMock) -> None:
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
filters = provider._build_filters(session_id="s99")
assert filters["run_id"] == "s99"
def test_empty_when_no_params(self, mock_mem0_client: AsyncMock) -> None:
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client)
assert provider._build_filters() == {}
# -- Context manager tests -----------------------------------------------------
class TestContextManager:
"""Test __aenter__/__aexit__ delegation."""
async def test_aenter_delegates_to_client(self, mock_mem0_client: AsyncMock) -> None:
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
result = await provider.__aenter__()
assert result is provider
mock_mem0_client.__aenter__.assert_awaited_once()
async def test_aexit_closes_auto_created_client(self, mock_mem0_client: AsyncMock) -> None:
"""Auto-created clients (_should_close_client=True) are closed on exit."""
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
provider._should_close_client = True
await provider.__aexit__(None, None, None)
mock_mem0_client.__aexit__.assert_awaited_once()
async def test_aexit_does_not_close_provided_client(self, mock_mem0_client: AsyncMock) -> None:
"""Provided clients (_should_close_client=False) are NOT closed on exit."""
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
assert provider._should_close_client is False
await provider.__aexit__(None, None, None)
mock_mem0_client.__aexit__.assert_not_awaited()
async def test_async_with_syntax(self, mock_mem0_client: AsyncMock) -> None:
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
async with provider as p:
assert p is provider