mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
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:
committed by
GitHub
Unverified
parent
0c67dbbce5
commit
1e350ea22f
@@ -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
|
||||
Reference in New Issue
Block a user