Python: [BREAKING] PR2 — Wire context provider pipeline, remove old types, update all consumers (#3850)

* PR2: Wire context provider pipeline and update all internal consumers

- Replace AgentThread with AgentSession across all packages
- Replace ContextProvider with BaseContextProvider across all packages
- Replace context_provider param with context_providers (Sequence)
- Replace thread= with session= in run() signatures
- Replace get_new_thread() with create_session()
- Add get_session(service_session_id) to agent interface
- DurableAgentThread -> DurableAgentSession
- Remove _notify_thread_of_new_messages from WorkflowAgent
- Wire before_run/after_run context provider pipeline in RawAgent
- Auto-inject InMemoryHistoryProvider when no providers configured

* fix: update all tests for context provider pipeline, fix lazy-loaders, remove old test files

* refactor: update all sample files for context provider pipeline (AgentThread→AgentSession, ContextProvider→BaseContextProvider)

* fix: update remaining ag-ui references (client docstring, getting_started sample)

* fix: make get_session service_session_id keyword-only to avoid confusion with session_id

* refactor: rename _RunContext.thread_messages to session_messages

* refactor: remove _threads.py, _memory.py, and old provider files; migrate devui to use plain message lists

* rename: remove _new_ prefix from test files

* refactor: rewrite SlidingWindowChatMessageStore as SlidingWindowHistoryProvider(InMemoryHistoryProvider)

* fix: read full history from session state directly instead of reaching into provider internals

* fix: update stale .pyi stubs, sample imports, and README references for new provider types

* fix: remove stale message_store, _notify_thread_of_new_messages, and session_id.key references in samples

* refactor: merge context_providers and sessions sample folders into sessions, remove aggregate_context_provider

* refactor: UserInfoMemory stores state in session.state instead of instance attributes

* feat: add Pydantic BaseModel support to session state serialization

Pydantic models stored in session.state are now automatically serialized
via model_dump() and restored via model_validate() during to_dict()/from_dict()
round-trips. Models are auto-registered on first serialization; use
register_state_type() for cold-start deserialization.

Also export register_state_type as a public API.

* fix mem0

* Update sample README links and descriptions for session terminology

- Replace 'thread' with 'session' in sample descriptions across all READMEs
- Update file links for renamed samples (mem0_sessions, redis_sessions, etc.)
- Fix Threads section → Sessions section in main samples/README.md
- Update tools, middleware, workflows, durabletask, azure_functions READMEs
- Update architecture diagrams in concepts/tools/README.md
- Update migration guides (autogen, semantic-kernel)

* Fix broken Redis README link to renamed sample

* Fix Mem0 OSS client search: pass scoping params as direct kwargs

AsyncMemory (OSS) expects user_id/agent_id/run_id as direct kwargs,
while AsyncMemoryClient (Platform) expects them in a filters dict.
Adds tests for both client types.

Port of fix from #3844 to new Mem0ContextProvider.

* Fix rebase issues: restore missing _conversation_state.py and checkpoint decode logic

- Add back _conversation_state.py (encode/decode_chat_messages) lost in rebase
- Fix on_checkpoint_restore to decode cache/conversation with decode_chat_messages
- Fix on_checkpoint_restore to use decode_checkpoint_value for pending requests
- Add tests/workflow/__init__.py for relative import support
- Fix test_agent_executor checkpoint selection (checkpoints[1] not superstep)

* Add STORES_BY_DEFAULT ClassVar to skip redundant InMemoryHistoryProvider injection

Chat clients that store history server-side by default (OpenAI Responses API,
Azure AI Agent) now declare STORES_BY_DEFAULT = True. The agent checks this
during auto-injection and skips InMemoryHistoryProvider unless the user
explicitly sets store=False.

* Fix broken markdown links in azure_ai and redis READMEs

* Fix getting-started samples to use session API instead of removed thread/ContextProvider API

* updates to workflow as agent

* fix group chat import

* Rename Thread→Session throughout, fix service_session_id propagation, remove stale AGUIThread

- Fix: Propagate conversation_id from ChatResponse back to session.service_session_id
  in both streaming and non-streaming paths in _agents.py
- Rename AgentThreadException → AgentSessionException
- Remove stale AGUIThread from ag_ui lazy-loader
- Rename use_service_thread → use_service_session in ag-ui package
- Rename test functions from *_thread_* to *_session_*
- Rename sample files from *_thread* to *_session*
- Update docstrings and comments: thread → session
- Update _mcp.py kwargs filter: add 'session' alongside 'thread'
- Fix ContinuationToken docstring example: thread=thread → session=session
- Fix _clients.py docstring: 'Agent threads' → 'Agent sessions'

* Fix broken markdown links after thread→session file renames

* fix azure ai test
This commit is contained in:
Eduard van Valkenburg
2026-02-12 22:00:32 +01:00
committed by GitHub
Unverified
parent 0c67dbbce5
commit 1e350ea22f
312 changed files with 6669 additions and 11423 deletions
+1 -1
View File
@@ -27,5 +27,5 @@ Mem0's telemetry is **disabled by default** when using this package. If you want
import os
os.environ["MEM0_TELEMETRY"] = "true"
from agent_framework.mem0 import Mem0Provider
from agent_framework.mem0 import Mem0ContextProvider
```
@@ -8,8 +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
from ._context_provider import Mem0ContextProvider
try:
__version__ = importlib.metadata.version(__name__)
@@ -17,7 +16,6 @@ except importlib.metadata.PackageNotFoundError:
__version__ = "0.0.0" # Fallback for development mode
__all__ = [
"Mem0Provider",
"_Mem0ContextProvider",
"Mem0ContextProvider",
"__version__",
]
@@ -2,9 +2,8 @@
"""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.
This module provides ``Mem0ContextProvider``, built on the new
:class:`BaseContextProvider` hooks pattern.
"""
from __future__ import annotations
@@ -35,17 +34,11 @@ class _MemorySearchResponse_v1_1(TypedDict):
_MemorySearchResponse_v2 = list[dict[str, Any]]
class _Mem0ContextProvider(BaseContextProvider):
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.
memories via the Mem0 API.
"""
DEFAULT_CONTEXT_PROMPT = "## Memories\nConsider the following memories when answering user questions:"
@@ -115,9 +108,16 @@ class _Mem0ContextProvider(BaseContextProvider):
filters = self._build_filters(session_id=context.session_id)
# AsyncMemory (OSS) expects user_id/agent_id/run_id as direct kwargs
# AsyncMemoryClient (Platform) expects them in a filters dict
search_kwargs: dict[str, Any] = {"query": input_text}
if isinstance(self.mem0_client, AsyncMemory):
search_kwargs.update(filters)
else:
search_kwargs["filters"] = filters
search_response: _MemorySearchResponse_v1_1 | _MemorySearchResponse_v2 = await self.mem0_client.search( # type: ignore[misc]
query=input_text,
filters=filters,
**search_kwargs,
)
if isinstance(search_response, list):
@@ -190,4 +190,4 @@ class _Mem0ContextProvider(BaseContextProvider):
return filters
__all__ = ["_Mem0ContextProvider"]
__all__ = ["Mem0ContextProvider"]
@@ -1,239 +0,0 @@
# Copyright (c) Microsoft. All rights reserved.
from __future__ import annotations
import sys
from collections.abc import MutableSequence, Sequence
from contextlib import AbstractAsyncContextManager
from typing import Any
from agent_framework import Context, ContextProvider, Message
from agent_framework.exceptions import ServiceInitializationError
from mem0 import AsyncMemory, AsyncMemoryClient
if sys.version_info >= (3, 12):
from typing import override # type: ignore # pragma: no cover
else:
from typing_extensions import override # type: ignore[import] # pragma: no cover
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
# Type aliases for Mem0 search response formats (v1.1 and v2; v1 is deprecated, but matches the type definition for v2)
class MemorySearchResponse_v1_1(TypedDict):
results: list[dict[str, Any]]
relations: NotRequired[list[dict[str, Any]]]
MemorySearchResponse_v2 = list[dict[str, Any]]
class Mem0Provider(ContextProvider):
"""Mem0 Context Provider.
Note:
Mem0's telemetry is disabled by default when using this package.
To enable telemetry, set the environment variable ``MEM0_TELEMETRY=true`` before
importing this package.
"""
def __init__(
self,
mem0_client: AsyncMemory | AsyncMemoryClient | None = None,
api_key: str | None = None,
application_id: str | None = None,
agent_id: str | None = None,
thread_id: str | None = None,
user_id: str | None = None,
scope_to_per_operation_thread_id: bool = False,
context_prompt: str = ContextProvider.DEFAULT_CONTEXT_PROMPT,
) -> None:
"""Initializes a new instance of the Mem0Provider class.
Args:
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. If not
provided, it will attempt to use the MEM0_API_KEY environment variable.
application_id: The application ID for scoping memories or None.
agent_id: The agent ID for scoping memories or None.
thread_id: The thread ID for scoping memories or None.
user_id: The user ID for scoping memories or None.
scope_to_per_operation_thread_id: Whether to scope memories to per-operation thread ID.
context_prompt: The prompt to prepend to retrieved memories.
"""
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.thread_id = thread_id
self.user_id = user_id
self.scope_to_per_operation_thread_id = scope_to_per_operation_thread_id
self.context_prompt = context_prompt
self.mem0_client = mem0_client
self._per_operation_thread_id: str | None = None
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)
async def thread_created(self, thread_id: str | None = None) -> None:
"""Called when a new thread is created.
Args:
thread_id: The ID of the thread or None.
"""
self._validate_per_operation_thread_id(thread_id)
self._per_operation_thread_id = self._per_operation_thread_id or thread_id
@override
async def invoked(
self,
request_messages: Message | Sequence[Message],
response_messages: Message | Sequence[Message] | None = None,
invoke_exception: Exception | None = None,
**kwargs: Any,
) -> None:
self._validate_filters()
request_messages_list = [request_messages] if isinstance(request_messages, Message) else list(request_messages)
response_messages_list = (
[response_messages]
if isinstance(response_messages, Message)
else list(response_messages)
if response_messages
else []
)
messages_list = [*request_messages_list, *response_messages_list]
# Extract role value - it may be a Role enum or a string
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_list
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=self._per_operation_thread_id if self.scope_to_per_operation_thread_id else self.thread_id,
metadata={"application_id": self.application_id},
)
@override
async def invoking(self, messages: Message | MutableSequence[Message], **kwargs: Any) -> Context:
"""Called before invoking the AI model to provide context.
Args:
messages: List of new messages in the thread.
Keyword Args:
**kwargs: not used at present.
Returns:
Context: Context object containing instructions with memories.
"""
self._validate_filters()
messages_list = [messages] if isinstance(messages, Message) else list(messages)
input_text = "\n".join(msg.text for msg in messages_list if msg and msg.text and msg.text.strip())
# Validate input text is not empty before searching (possible for function approval responses)
if not input_text.strip():
return Context(messages=None)
# Build filters from init parameters
filters = self._build_filters()
search_response: MemorySearchResponse_v1_1 | MemorySearchResponse_v2 = await self.mem0_client.search( # type: ignore[misc]
query=input_text,
filters=filters,
)
# Depending on the API version, the response schema varies slightly
if isinstance(search_response, list):
memories = search_response
elif isinstance(search_response, dict) and "results" in search_response:
memories = search_response["results"]
else:
# Fallback for unexpected schema - return response as text as-is
memories = [search_response]
line_separated_memories = "\n".join(memory.get("memory", "") for memory in memories)
return Context(
messages=[Message(role="user", text=f"{self.context_prompt}\n{line_separated_memories}")]
if line_separated_memories
else None
)
def _validate_filters(self) -> None:
"""Validates that at least one filter is provided.
Raises:
ServiceInitializationError: If no filters are provided.
"""
if not self.agent_id and not self.user_id and not self.application_id and not self.thread_id:
raise ServiceInitializationError(
"At least one of the filters: agent_id, user_id, application_id, or thread_id is required."
)
def _build_filters(self) -> dict[str, Any]:
"""Build search filters from initialization parameters.
Returns:
Filter dictionary for mem0 v2 search API containing initialization parameters.
In the v2 API, filters holds the user_id, agent_id, run_id (thread_id), and app_id
(application_id) which are required for scoping memory search operations.
"""
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 self.scope_to_per_operation_thread_id and self._per_operation_thread_id:
filters["run_id"] = self._per_operation_thread_id
elif self.thread_id:
filters["run_id"] = self.thread_id
if self.application_id:
filters["app_id"] = self.application_id
return filters
def _validate_per_operation_thread_id(self, thread_id: str | None) -> None:
"""Validates that a new thread ID doesn't conflict with an existing one when scoped.
Args:
thread_id: The new thread ID or None.
Raises:
ValueError: If a new thread ID is provided when one already exists.
"""
if (
self.scope_to_per_operation_thread_id
and thread_id
and self._per_operation_thread_id
and thread_id != self._per_operation_thread_id
):
raise ValueError(
"Mem0Provider can only be used with one thread at a time when scope_to_per_operation_thread_id is True."
)
@@ -1,20 +1,16 @@
# Copyright (c) Microsoft. All rights reserved.
# pyright: reportPrivateUsage=false
import importlib
import os
import sys
from unittest.mock import AsyncMock
from __future__ import annotations
from unittest.mock import AsyncMock, patch
import pytest
from agent_framework import Content, Context, Message
from agent_framework import AgentResponse, Message
from agent_framework._sessions import AgentSession, SessionContext
from agent_framework.exceptions import ServiceInitializationError
from agent_framework.mem0 import Mem0Provider
def test_mem0_provider_import() -> None:
"""Test that Mem0Provider can be imported."""
assert Mem0Provider is not None
from agent_framework_mem0._context_provider import Mem0ContextProvider
@pytest.fixture
@@ -27,577 +23,385 @@ def mock_mem0_client() -> AsyncMock:
mock_client.search = AsyncMock()
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
mock_client.__aexit__ = AsyncMock()
mock_client.async_client = AsyncMock()
mock_client.async_client.aclose = AsyncMock()
return mock_client
@pytest.fixture
def sample_messages() -> list[Message]:
"""Create sample chat messages for testing."""
return [
Message(role="user", text="Hello, how are you?"),
Message(role="assistant", text="I'm doing well, thank you!"),
Message(role="system", text="You are a helpful assistant"),
]
def mock_oss_mem0_client() -> AsyncMock:
"""Create a mock Mem0 OSS AsyncMemory client."""
from mem0 import AsyncMemory
mock_client = AsyncMock(spec=AsyncMemory)
mock_client.add = AsyncMock()
mock_client.search = AsyncMock()
return mock_client
def test_init_with_all_ids(mock_mem0_client: AsyncMock) -> None:
"""Test initialization with all IDs provided."""
provider = Mem0Provider(
user_id="user123",
agent_id="agent123",
application_id="app123",
thread_id="thread123",
mem0_client=mock_mem0_client,
)
assert provider.user_id == "user123"
assert provider.agent_id == "agent123"
assert provider.application_id == "app123"
assert provider.thread_id == "thread123"
# -- Initialization tests ------------------------------------------------------
def test_init_without_filters_succeeds(mock_mem0_client: AsyncMock) -> None:
"""Test that initialization succeeds even without filters (validation happens during invocation)."""
provider = Mem0Provider(mem0_client=mock_mem0_client)
assert provider.user_id is None
assert provider.agent_id is None
assert provider.application_id is None
assert provider.thread_id is None
class TestInit:
"""Test Mem0ContextProvider initialization."""
def test_init_with_custom_context_prompt(mock_mem0_client: AsyncMock) -> None:
"""Test initialization with custom context prompt."""
custom_prompt = "## Custom Memories\nConsider these memories:"
provider = Mem0Provider(user_id="user123", context_prompt=custom_prompt, mem0_client=mock_mem0_client)
assert provider.context_prompt == custom_prompt
def test_init_with_scope_to_per_operation_thread_id(mock_mem0_client: AsyncMock) -> None:
"""Test initialization with scope_to_per_operation_thread_id enabled."""
provider = Mem0Provider(
user_id="user123",
scope_to_per_operation_thread_id=True,
mem0_client=mock_mem0_client,
)
assert provider.scope_to_per_operation_thread_id is True
def test_init_with_provided_client_should_not_close(mock_mem0_client: AsyncMock) -> None:
"""Test that provided client should not be closed by provider."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
assert provider._should_close_client is False
async def test_async_context_manager_entry(mock_mem0_client: AsyncMock) -> None:
"""Test async context manager entry returns self."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
async with provider as ctx:
assert ctx is provider
async def test_async_context_manager_exit_does_not_close_provided_client(mock_mem0_client: AsyncMock) -> None:
"""Test that async context manager does not close provided client."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
assert provider._should_close_client is False
async with provider:
pass
mock_mem0_client.__aexit__.assert_not_called()
class TestMem0ProviderThreadMethods:
"""Test thread lifecycle methods."""
async def test_thread_created_sets_per_operation_thread_id(self, mock_mem0_client: AsyncMock) -> None:
"""Test that thread_created sets per-operation thread ID."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
await provider.thread_created("thread123")
assert provider._per_operation_thread_id == "thread123"
async def test_thread_created_with_existing_thread_id(self, mock_mem0_client: AsyncMock) -> None:
"""Test thread_created when thread ID already exists."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
provider._per_operation_thread_id = "existing_thread"
await provider.thread_created("thread123")
# Should not overwrite existing thread ID
assert provider._per_operation_thread_id == "existing_thread"
async def test_thread_created_validation_with_scope_enabled(self, mock_mem0_client: AsyncMock) -> None:
"""Test thread_created validation when scope_to_per_operation_thread_id is enabled."""
provider = Mem0Provider(
user_id="user123",
scope_to_per_operation_thread_id=True,
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",
)
provider._per_operation_thread_id = "existing_thread"
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
with pytest.raises(ValueError) as exc_info:
await provider.thread_created("different_thread")
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
assert "can only be used with one thread at a time" in str(exc_info.value)
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
async def test_messages_adding_sets_per_operation_thread_id(self, mock_mem0_client: AsyncMock) -> None:
"""Test that invoked sets per-operation thread ID."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
await provider.thread_created("thread123")
assert provider._per_operation_thread_id == "thread123"
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
class TestMem0ProviderMessagesAdding:
"""Test invoked method."""
async def test_messages_adding_fails_without_filters(self, mock_mem0_client: AsyncMock) -> None:
"""Test that invoked fails when no filters are provided."""
provider = Mem0Provider(mem0_client=mock_mem0_client)
message = Message(role="user", text="Hello!")
with pytest.raises(ServiceInitializationError) as exc_info:
await provider.invoked(message)
assert "At least one of the filters" in str(exc_info.value)
async def test_messages_adding_single_message(self, mock_mem0_client: AsyncMock) -> None:
"""Test adding a single message."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
message = Message(role="user", text="Hello!")
await provider.invoked(message)
mock_mem0_client.add.assert_called_once()
call_args = mock_mem0_client.add.call_args
assert call_args.kwargs["messages"] == [{"role": "user", "content": "Hello!"}]
assert call_args.kwargs["user_id"] == "user123"
async def test_messages_adding_multiple_messages(
self, mock_mem0_client: AsyncMock, sample_messages: list[Message]
) -> None:
"""Test adding multiple messages."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
await provider.invoked(sample_messages)
mock_mem0_client.add.assert_called_once()
call_args = mock_mem0_client.add.call_args
expected_messages = [
{"role": "user", "content": "Hello, how are you?"},
{"role": "assistant", "content": "I'm doing well, thank you!"},
{"role": "system", "content": "You are a helpful assistant"},
]
assert call_args.kwargs["messages"] == expected_messages
async def test_messages_adding_with_agent_id(
self, mock_mem0_client: AsyncMock, sample_messages: list[Message]
) -> None:
"""Test adding messages with agent_id."""
provider = Mem0Provider(agent_id="agent123", mem0_client=mock_mem0_client)
await provider.invoked(sample_messages)
call_args = mock_mem0_client.add.call_args
assert call_args.kwargs["agent_id"] == "agent123"
assert call_args.kwargs["user_id"] is None
async def test_messages_adding_with_application_id(
self, mock_mem0_client: AsyncMock, sample_messages: list[Message]
) -> None:
"""Test adding messages with application_id in metadata."""
provider = Mem0Provider(user_id="user123", application_id="app123", mem0_client=mock_mem0_client)
await provider.invoked(sample_messages)
call_args = mock_mem0_client.add.call_args
assert call_args.kwargs["metadata"] == {"application_id": "app123"}
async def test_messages_adding_with_scope_to_per_operation_thread_id(
self, mock_mem0_client: AsyncMock, sample_messages: list[Message]
) -> None:
"""Test adding messages with scope_to_per_operation_thread_id enabled."""
provider = Mem0Provider(
user_id="user123",
thread_id="base_thread",
scope_to_per_operation_thread_id=True,
mem0_client=mock_mem0_client,
)
provider._per_operation_thread_id = "operation_thread"
await provider.thread_created(thread_id="operation_thread")
await provider.invoked(sample_messages)
call_args = mock_mem0_client.add.call_args
assert call_args.kwargs["run_id"] == "operation_thread"
async def test_messages_adding_without_scope_uses_base_thread_id(
self, mock_mem0_client: AsyncMock, sample_messages: list[Message]
) -> None:
"""Test adding messages without scope uses base thread_id."""
provider = Mem0Provider(
user_id="user123",
thread_id="base_thread",
scope_to_per_operation_thread_id=False,
mem0_client=mock_mem0_client,
)
await provider.invoked(sample_messages)
call_args = mock_mem0_client.add.call_args
assert call_args.kwargs["run_id"] == "base_thread"
async def test_messages_adding_filters_empty_messages(self, mock_mem0_client: AsyncMock) -> None:
"""Test that empty or invalid messages are filtered out."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
messages = [
Message(role="user", text=""), # Empty text
Message(role="user", text=" "), # Whitespace only
Message(role="user", text="Valid message"),
]
await provider.invoked(messages)
call_args = mock_mem0_client.add.call_args
# Should only include the valid message
assert call_args.kwargs["messages"] == [{"role": "user", "content": "Valid message"}]
async def test_messages_adding_skips_when_no_valid_messages(self, mock_mem0_client: AsyncMock) -> None:
"""Test that mem0 client is not called when no valid messages exist."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
messages = [
Message(role="user", text=""),
Message(role="user", text=" "),
]
await provider.invoked(messages)
mock_mem0_client.add.assert_not_called()
# -- before_run tests ----------------------------------------------------------
class TestMem0ProviderModelInvoking:
"""Test invoking method."""
class TestBeforeRun:
"""Test before_run hook."""
async def test_model_invoking_fails_without_filters(self, mock_mem0_client: AsyncMock) -> None:
"""Test that invoking fails when no filters are provided."""
provider = Mem0Provider(mem0_client=mock_mem0_client)
message = Message(role="user", text="What's the weather?")
with pytest.raises(ServiceInitializationError) as exc_info:
await provider.invoking(message)
assert "At least one of the filters" in str(exc_info.value)
async def test_model_invoking_single_message(self, mock_mem0_client: AsyncMock) -> None:
"""Test invoking with a single message."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
message = Message(role="user", text="What's the weather?")
# Mock search results
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 outdoor activities"},
{"memory": "User lives in Seattle"},
{"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=[Message(role="user", text="Hello")], session_id="s1")
context = await provider.invoking(message)
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
mock_mem0_client.search.assert_called_once()
call_args = mock_mem0_client.search.call_args
assert call_args.kwargs["query"] == "What's the weather?"
assert call_args.kwargs["filters"] == {"user_id": "user123"}
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]
assert isinstance(context, Context)
expected_instructions = (
"## Memories\nConsider the following memories when answering user questions:\n"
"User likes outdoor activities\nUser lives in Seattle"
)
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=[Message(role="user", text="")], session_id="s1")
assert context.messages
assert context.messages[0].text == expected_instructions
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
async def test_model_invoking_multiple_messages(
self, mock_mem0_client: AsyncMock, sample_messages: list[Message]
) -> None:
"""Test invoking with multiple messages."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
mock_mem0_client.search.return_value = [{"memory": "Previous conversation context"}]
await provider.invoking(sample_messages)
call_args = mock_mem0_client.search.call_args
expected_query = "Hello, how are you?\nI'm doing well, thank you!\nYou are a helpful assistant"
assert call_args.kwargs["query"] == expected_query
async def test_model_invoking_with_agent_id(self, mock_mem0_client: AsyncMock) -> None:
"""Test invoking with agent_id."""
provider = Mem0Provider(agent_id="agent123", mem0_client=mock_mem0_client)
message = Message(role="user", text="Hello")
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=[Message(role="user", text="test")], session_id="s1")
await provider.invoking(message)
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
call_args = mock_mem0_client.search.call_args
assert call_args.kwargs["filters"] == {"agent_id": "agent123"}
assert "mem0" not in ctx.context_messages
async def test_model_invoking_with_scope_to_per_operation_thread_id(self, mock_mem0_client: AsyncMock) -> None:
"""Test invoking with scope_to_per_operation_thread_id enabled."""
provider = Mem0Provider(
user_id="user123",
thread_id="base_thread",
scope_to_per_operation_thread_id=True,
mem0_client=mock_mem0_client,
)
provider._per_operation_thread_id = "operation_thread"
message = Message(role="user", text="Hello")
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=[Message(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=[Message(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 = []
await provider.invoking(message)
call_args = mock_mem0_client.search.call_args
assert call_args.kwargs["filters"] == {"user_id": "user123", "run_id": "operation_thread"}
async def test_model_invoking_no_memories_returns_none_instructions(self, mock_mem0_client: AsyncMock) -> None:
"""Test that no memories returns context with None instructions."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
message = Message(role="user", text="Hello")
mock_mem0_client.search.return_value = []
context = await provider.invoking(message)
assert isinstance(context, Context)
assert not context.messages
async def test_model_invoking_function_approval_response_returns_none_instructions(
self, mock_mem0_client: AsyncMock
) -> None:
"""Test invoking with function approval response content messages returns context with None instructions."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
function_call = Content.from_function_call(call_id="1", name="test_func", arguments='{"arg1": "value1"}')
message = Message(
role="user",
contents=[
Content.from_function_approval_response(
id="approval_1",
function_call=function_call,
approved=True,
)
provider = Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(
input_messages=[
Message(role="user", text="Hello"),
Message(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"
async def test_oss_client_passes_direct_kwargs(self, mock_oss_mem0_client: AsyncMock) -> None:
"""OSS AsyncMemory client should receive user_id as direct kwarg, not in filters."""
mock_oss_mem0_client.search.return_value = [{"memory": "User likes Python"}]
provider = Mem0ContextProvider(source_id="mem0", mem0_client=mock_oss_mem0_client, user_id="u1")
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", text="Hello")], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
call_kwargs = mock_oss_mem0_client.search.call_args.kwargs
assert call_kwargs["query"] == "Hello"
assert call_kwargs["user_id"] == "u1"
assert "filters" not in call_kwargs
async def test_oss_client_all_scoping_params(self, mock_oss_mem0_client: AsyncMock) -> None:
"""OSS client with all scoping parameters passes them as direct kwargs."""
mock_oss_mem0_client.search.return_value = []
provider = Mem0ContextProvider(
source_id="mem0", mem0_client=mock_oss_mem0_client, user_id="u1", agent_id="a1", application_id="app1"
)
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", text="Hello")], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
call_kwargs = mock_oss_mem0_client.search.call_args.kwargs
assert call_kwargs["user_id"] == "u1"
assert call_kwargs["agent_id"] == "a1"
assert "filters" not in call_kwargs
async def test_platform_client_passes_filters_dict(self, mock_mem0_client: AsyncMock) -> None:
"""Platform AsyncMemoryClient should receive scoping params in a filters dict."""
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=[Message(role="user", text="Hello")], session_id="s1")
context = await provider.invoking(message)
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
assert isinstance(context, Context)
assert not context.messages
call_kwargs = mock_mem0_client.search.call_args.kwargs
assert call_kwargs["query"] == "Hello"
assert "filters" in call_kwargs
assert call_kwargs["filters"]["user_id"] == "u1"
async def test_model_invoking_filters_empty_message_text(self, mock_mem0_client: AsyncMock) -> None:
"""Test that empty message text is filtered out from query."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
messages = [
Message(role="user", text=""),
Message(role="user", text="Valid message"),
Message(role="user", text=" "),
# -- 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=[Message(role="user", text="question")], session_id="s1")
ctx._response = AgentResponse(messages=[Message(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"
mock_mem0_client.search.return_value = []
await provider.invoking(messages)
call_args = mock_mem0_client.search.call_args
assert call_args.kwargs["query"] == "Valid message"
async def test_model_invoking_custom_context_prompt(self, mock_mem0_client: AsyncMock) -> None:
"""Test invoking with custom context prompt."""
custom_prompt = "## Custom Context\nRemember these details:"
provider = Mem0Provider(
user_id="user123",
context_prompt=custom_prompt,
mem0_client=mock_mem0_client,
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=[
Message(role="user", text="hello"),
Message(role="tool", text="tool output"),
],
session_id="s1",
)
message = Message(role="user", text="Hello")
ctx._response = AgentResponse(messages=[Message(role="assistant", text="reply")])
mock_mem0_client.search.return_value = [{"memory": "Test memory"}]
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
context = await provider.invoking(message)
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"]
expected_instructions = "## Custom Context\nRemember these details:\nTest memory"
assert context.messages
assert context.messages[0].text == expected_instructions
class TestMem0ProviderValidation:
"""Test validation methods."""
def test_validate_per_operation_thread_id_success(self, mock_mem0_client: AsyncMock) -> None:
"""Test successful validation of per-operation thread ID."""
provider = Mem0Provider(
user_id="user123",
scope_to_per_operation_thread_id=True,
mem0_client=mock_mem0_client,
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=[
Message(role="user", text=""),
Message(role="user", text=" "),
],
session_id="s1",
)
provider._per_operation_thread_id = "thread123"
ctx._response = AgentResponse(messages=[])
# Should not raise exception for same thread ID
provider._validate_per_operation_thread_id("thread123")
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
# Should not raise exception for None
provider._validate_per_operation_thread_id(None)
mock_mem0_client.add.assert_not_awaited()
def test_validate_per_operation_thread_id_failure(self, mock_mem0_client: AsyncMock) -> None:
"""Test validation failure for conflicting thread IDs."""
provider = Mem0Provider(
user_id="user123",
scope_to_per_operation_thread_id=True,
mem0_client=mock_mem0_client,
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=[Message(role="user", text="hi")], session_id="my-session")
ctx._response = AgentResponse(messages=[Message(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=[Message(role="user", text="hi")], session_id="s1")
ctx._response = AgentResponse(messages=[Message(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"
)
provider._per_operation_thread_id = "thread123"
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[Message(role="user", text="hi")], session_id="s1")
ctx._response = AgentResponse(messages=[])
with pytest.raises(ValueError) as exc_info:
provider._validate_per_operation_thread_id("different_thread")
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
assert "can only be used with one thread at a time" in str(exc_info.value)
assert mock_mem0_client.add.call_args.kwargs["metadata"] == {"application_id": "app1"}
def test_validate_per_operation_thread_id_disabled_scope(self, mock_mem0_client: AsyncMock) -> None:
"""Test that validation is skipped when scope is disabled."""
provider = Mem0Provider(
user_id="user123",
scope_to_per_operation_thread_id=False,
# -- _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",
)
provider._per_operation_thread_id = "thread123"
# Should not raise exception even with different thread ID
provider._validate_per_operation_thread_id("different_thread")
class TestMem0ProviderBuildFilters:
"""Test the _build_filters method."""
def test_build_filters_with_user_id_only(self, mock_mem0_client: AsyncMock) -> None:
"""Test building filters with only user_id."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
filters = provider._build_filters()
assert filters == {"user_id": "user123"}
def test_build_filters_with_all_parameters(self, mock_mem0_client: AsyncMock) -> None:
"""Test building filters with all initialization parameters."""
provider = Mem0Provider(
user_id="user123",
agent_id="agent456",
thread_id="thread789",
application_id="app999",
mem0_client=mock_mem0_client,
)
filters = provider._build_filters()
assert filters == {
"user_id": "user123",
"agent_id": "agent456",
"run_id": "thread789",
"app_id": "app999",
assert provider._build_filters(session_id="sess1") == {
"user_id": "u1",
"agent_id": "a1",
"run_id": "sess1",
"app_id": "app1",
}
def test_build_filters_excludes_none_values(self, mock_mem0_client: AsyncMock) -> None:
"""Test that None values are excluded from filters."""
provider = Mem0Provider(
user_id="user123",
agent_id=None,
thread_id=None,
application_id=None,
mem0_client=mock_mem0_client,
)
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 filters == {"user_id": "user123"}
assert "agent_id" not in filters
assert "run_id" not in filters
assert "app_id" not in filters
def test_build_filters_with_per_operation_thread_id(self, mock_mem0_client: AsyncMock) -> None:
"""Test that per-operation thread ID takes precedence over base thread_id."""
provider = Mem0Provider(
user_id="user123",
thread_id="base_thread",
scope_to_per_operation_thread_id=True,
mem0_client=mock_mem0_client,
)
provider._per_operation_thread_id = "operation_thread"
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"
filters = provider._build_filters()
assert filters == {
"user_id": "user123",
"run_id": "operation_thread", # Per-operation thread, not base_thread
}
def test_build_filters_uses_base_thread_when_no_per_operation(self, mock_mem0_client: AsyncMock) -> None:
"""Test that base thread_id is used when per-operation thread is not set."""
provider = Mem0Provider(
user_id="user123",
thread_id="base_thread",
scope_to_per_operation_thread_id=True,
mem0_client=mock_mem0_client,
)
# _per_operation_thread_id is None
filters = provider._build_filters()
assert filters == {
"user_id": "user123",
"run_id": "base_thread", # Falls back to base thread_id
}
def test_build_filters_returns_empty_dict_when_no_parameters(self, mock_mem0_client: AsyncMock) -> None:
"""Test that _build_filters returns an empty dict when no parameters are set."""
provider = Mem0Provider(mem0_client=mock_mem0_client)
filters = provider._build_filters()
assert filters == {}
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() == {}
class TestMem0Telemetry:
"""Test telemetry configuration for Mem0."""
# -- Context manager tests -----------------------------------------------------
def test_mem0_telemetry_disabled_by_default(self, monkeypatch: pytest.MonkeyPatch) -> None:
"""Test that MEM0_TELEMETRY is set to 'false' by default when importing the package."""
# Ensure MEM0_TELEMETRY is not set before importing the module under test
monkeypatch.delenv("MEM0_TELEMETRY", raising=False)
# Remove cached modules to force re-import and trigger module-level initialization
modules_to_remove = [key for key in sys.modules if key.startswith("agent_framework_mem0")]
for mod in modules_to_remove:
del sys.modules[mod]
class TestContextManager:
"""Test __aenter__/__aexit__ delegation."""
# Import (and reload) the module so that it can set MEM0_TELEMETRY when unset
import agent_framework_mem0
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()
importlib.reload(agent_framework_mem0)
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()
# The environment variable should be set to "false" after importing
assert os.environ.get("MEM0_TELEMETRY") == "false"
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()
def test_mem0_telemetry_respects_user_setting(self, monkeypatch: pytest.MonkeyPatch) -> None:
"""Test that user-set MEM0_TELEMETRY value is not overwritten."""
# Remove cached modules to force re-import
modules_to_remove = [key for key in sys.modules if key.startswith("agent_framework_mem0")]
for mod in modules_to_remove:
del sys.modules[mod]
# Set user preference before import
monkeypatch.setenv("MEM0_TELEMETRY", "true")
# Re-import the module
import agent_framework_mem0
importlib.reload(agent_framework_mem0)
# User setting should be preserved
assert os.environ.get("MEM0_TELEMETRY") == "true"
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
@@ -1,352 +0,0 @@
# 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, Message
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=[Message(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=[Message(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=[Message(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=[Message(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=[Message(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=[
Message(role="user", text="Hello"),
Message(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=[Message(role="user", text="question")], session_id="s1")
ctx._response = AgentResponse(messages=[Message(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=[
Message(role="user", text="hello"),
Message(role="tool", text="tool output"),
],
session_id="s1",
)
ctx._response = AgentResponse(messages=[Message(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=[
Message(role="user", text=""),
Message(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=[Message(role="user", text="hi")], session_id="my-session")
ctx._response = AgentResponse(messages=[Message(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=[Message(role="user", text="hi")], session_id="s1")
ctx._response = AgentResponse(messages=[Message(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=[Message(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