mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: PR1 — New session and context provider types (side-by-side) (#3763)
* PR1: Add core context provider types and tests New types in _sessions.py (no changes to existing code): - SessionContext: per-invocation state with extend_messages/get_messages/ extend_instructions/extend_tools and read-only response property - _ContextProviderBase: base class with before_run/after_run hooks - _HistoryProviderBase: storage base with load/store flags, abstract get_messages/save_messages, default before_run/after_run - AgentSession: lightweight session with state dict, to_dict/from_dict - InMemoryHistoryProvider: built-in provider storing in session.state 35 unit tests covering all classes and configuration flags. * feat: keyword-only params, stateless InMemoryHistoryProvider, deep serialization - Make before_run/after_run parameters keyword-only - InMemoryHistoryProvider stores ChatMessage objects directly (no per-cycle serialization) - Deep serialization via to_dict/from_dict only at session boundary - State type registry for automatic deserialization of registered types - Updated tests for new serialization approach * feat: add new-pattern provider implementations for external packages - _RedisContextProvider(BaseContextProvider) - Redis search/vector context - _RedisHistoryProvider(BaseHistoryProvider) - Redis-backed message storage - _Mem0ContextProvider(BaseContextProvider) - Mem0 semantic memory - _AzureAISearchContextProvider(BaseContextProvider) - Azure AI Search (semantic + agentic) All use temporary _ prefix names for side-by-side coexistence with existing providers. Will be renamed in PR2 when old ContextProvider/ChatMessageStore are removed. * test: add tests for new-pattern provider implementations - 32 tests for _RedisContextProvider and _RedisHistoryProvider - 29 tests for _Mem0ContextProvider - 17 tests for _AzureAISearchContextProvider * fix: address PR review comments and CI failures - Move module docstring before imports in _sessions.py (review comment) - Import TYPE_CHECKING unconditionally in Redis _context_provider.py (NameError on Python <3.12) - Fix Mem0 test_init_auto_creates_client_when_none to patch at class level * feat: add source attribution to extend_messages Set attribution marker in additional_properties for each message added via extend_messages(), matching the tool attribution pattern. Uses setdefault to preserve any existing attribution. * refactor: make attribution value a dict with source_id key * add attribution and use sets for filters * Add source_type to message attribution and copy messages in extend_messages - SessionContext.extend_messages now accepts source as str or object with source_id attribute; when an object is passed, its class name is recorded as source_type in the attribution dict - Messages are shallow-copied before attribution is added so callers' original objects are never mutated - Filter framework-internal keys (attribution) from A2A wire metadata to prevent leaking internal state over the wire * fix: correct mypy type: ignore comment from union-attr to attr-defined * set attribution to _attribution * adjusted naming of bools
This commit is contained in:
committed by
GitHub
Unverified
parent
ccff3d3452
commit
ac0e6b0ee1
@@ -8,6 +8,7 @@ import os
|
||||
if os.environ.get("MEM0_TELEMETRY") is None:
|
||||
os.environ["MEM0_TELEMETRY"] = "false"
|
||||
|
||||
from ._context_provider import _Mem0ContextProvider
|
||||
from ._provider import Mem0Provider
|
||||
|
||||
try:
|
||||
@@ -17,5 +18,6 @@ except importlib.metadata.PackageNotFoundError:
|
||||
|
||||
__all__ = [
|
||||
"Mem0Provider",
|
||||
"_Mem0ContextProvider",
|
||||
"__version__",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,193 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
"""New-pattern Mem0 context provider using BaseContextProvider.
|
||||
|
||||
This module provides ``_Mem0ContextProvider``, a side-by-side implementation of
|
||||
:class:`Mem0Provider` built on the new :class:`BaseContextProvider` hooks pattern.
|
||||
It will be renamed to ``Mem0ContextProvider`` in PR2 when the old class is removed.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import sys
|
||||
from contextlib import AbstractAsyncContextManager
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from agent_framework import ChatMessage
|
||||
from agent_framework._sessions import AgentSession, BaseContextProvider, SessionContext
|
||||
from agent_framework.exceptions import ServiceInitializationError
|
||||
from mem0 import AsyncMemory, AsyncMemoryClient
|
||||
|
||||
if sys.version_info >= (3, 11):
|
||||
from typing import NotRequired, Self, TypedDict # pragma: no cover
|
||||
else:
|
||||
from typing_extensions import NotRequired, Self, TypedDict # pragma: no cover
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from agent_framework._agents import SupportsAgentRun
|
||||
|
||||
|
||||
class _MemorySearchResponse_v1_1(TypedDict):
|
||||
results: list[dict[str, Any]]
|
||||
relations: NotRequired[list[dict[str, Any]]]
|
||||
|
||||
|
||||
_MemorySearchResponse_v2 = list[dict[str, Any]]
|
||||
|
||||
|
||||
class _Mem0ContextProvider(BaseContextProvider):
|
||||
"""Mem0 context provider using the new BaseContextProvider hooks pattern.
|
||||
|
||||
Integrates Mem0 for persistent semantic memory, searching and storing
|
||||
memories via the Mem0 API. This is the new-pattern equivalent of
|
||||
:class:`Mem0Provider`.
|
||||
|
||||
Note:
|
||||
This class uses a temporary ``_`` prefix to coexist with the existing
|
||||
:class:`Mem0Provider`. It will be renamed to ``Mem0ContextProvider``
|
||||
in PR2.
|
||||
"""
|
||||
|
||||
DEFAULT_CONTEXT_PROMPT = "## Memories\nConsider the following memories when answering user questions:"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
source_id: str,
|
||||
mem0_client: AsyncMemory | AsyncMemoryClient | None = None,
|
||||
api_key: str | None = None,
|
||||
application_id: str | None = None,
|
||||
agent_id: str | None = None,
|
||||
user_id: str | None = None,
|
||||
*,
|
||||
context_prompt: str | None = None,
|
||||
) -> None:
|
||||
"""Initialize the Mem0 context provider.
|
||||
|
||||
Args:
|
||||
source_id: Unique identifier for this provider instance.
|
||||
mem0_client: A pre-created Mem0 MemoryClient or None to create a default client.
|
||||
api_key: The API key for authenticating with the Mem0 API.
|
||||
application_id: The application ID for scoping memories.
|
||||
agent_id: The agent ID for scoping memories.
|
||||
user_id: The user ID for scoping memories.
|
||||
context_prompt: The prompt to prepend to retrieved memories.
|
||||
"""
|
||||
super().__init__(source_id)
|
||||
should_close_client = False
|
||||
if mem0_client is None:
|
||||
mem0_client = AsyncMemoryClient(api_key=api_key)
|
||||
should_close_client = True
|
||||
|
||||
self.api_key = api_key
|
||||
self.application_id = application_id
|
||||
self.agent_id = agent_id
|
||||
self.user_id = user_id
|
||||
self.context_prompt = context_prompt or self.DEFAULT_CONTEXT_PROMPT
|
||||
self.mem0_client = mem0_client
|
||||
self._should_close_client = should_close_client
|
||||
|
||||
async def __aenter__(self) -> Self:
|
||||
"""Async context manager entry."""
|
||||
if self.mem0_client and isinstance(self.mem0_client, AbstractAsyncContextManager):
|
||||
await self.mem0_client.__aenter__()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type: type[BaseException] | None, exc_val: BaseException | None, exc_tb: Any) -> None:
|
||||
"""Async context manager exit."""
|
||||
if self._should_close_client and self.mem0_client and isinstance(self.mem0_client, AbstractAsyncContextManager):
|
||||
await self.mem0_client.__aexit__(exc_type, exc_val, exc_tb)
|
||||
|
||||
# -- Hooks pattern ---------------------------------------------------------
|
||||
|
||||
async def before_run(
|
||||
self,
|
||||
*,
|
||||
agent: SupportsAgentRun,
|
||||
session: AgentSession,
|
||||
context: SessionContext,
|
||||
state: dict[str, Any],
|
||||
) -> None:
|
||||
"""Search Mem0 for relevant memories and add to the session context."""
|
||||
self._validate_filters()
|
||||
input_text = "\n".join(msg.text for msg in context.input_messages if msg and msg.text and msg.text.strip())
|
||||
if not input_text.strip():
|
||||
return
|
||||
|
||||
filters = self._build_filters(session_id=context.session_id)
|
||||
|
||||
search_response: _MemorySearchResponse_v1_1 | _MemorySearchResponse_v2 = await self.mem0_client.search( # type: ignore[misc]
|
||||
query=input_text,
|
||||
filters=filters,
|
||||
)
|
||||
|
||||
if isinstance(search_response, list):
|
||||
memories = search_response
|
||||
elif isinstance(search_response, dict) and "results" in search_response:
|
||||
memories = search_response["results"]
|
||||
else:
|
||||
memories = [search_response]
|
||||
|
||||
line_separated_memories = "\n".join(memory.get("memory", "") for memory in memories)
|
||||
if line_separated_memories:
|
||||
context.extend_messages(
|
||||
self.source_id,
|
||||
[ChatMessage(role="user", text=f"{self.context_prompt}\n{line_separated_memories}")],
|
||||
)
|
||||
|
||||
async def after_run(
|
||||
self,
|
||||
*,
|
||||
agent: SupportsAgentRun,
|
||||
session: AgentSession,
|
||||
context: SessionContext,
|
||||
state: dict[str, Any],
|
||||
) -> None:
|
||||
"""Store request/response messages to Mem0 for future retrieval."""
|
||||
self._validate_filters()
|
||||
|
||||
messages_to_store: list[ChatMessage] = list(context.input_messages)
|
||||
if context.response and context.response.messages:
|
||||
messages_to_store.extend(context.response.messages)
|
||||
|
||||
def get_role_value(role: Any) -> str:
|
||||
return role.value if hasattr(role, "value") else str(role)
|
||||
|
||||
messages: list[dict[str, str]] = [
|
||||
{"role": get_role_value(message.role), "content": message.text}
|
||||
for message in messages_to_store
|
||||
if get_role_value(message.role) in {"user", "assistant", "system"} and message.text and message.text.strip()
|
||||
]
|
||||
|
||||
if messages:
|
||||
await self.mem0_client.add( # type: ignore[misc]
|
||||
messages=messages,
|
||||
user_id=self.user_id,
|
||||
agent_id=self.agent_id,
|
||||
run_id=context.session_id,
|
||||
metadata={"application_id": self.application_id},
|
||||
)
|
||||
|
||||
# -- Internal methods ------------------------------------------------------
|
||||
|
||||
def _validate_filters(self) -> None:
|
||||
"""Validates that at least one filter is provided."""
|
||||
if not self.agent_id and not self.user_id and not self.application_id:
|
||||
raise ServiceInitializationError(
|
||||
"At least one of the filters: agent_id, user_id, or application_id is required."
|
||||
)
|
||||
|
||||
def _build_filters(self, *, session_id: str | None = None) -> dict[str, Any]:
|
||||
"""Build search filters from initialization parameters."""
|
||||
filters: dict[str, Any] = {}
|
||||
if self.user_id:
|
||||
filters["user_id"] = self.user_id
|
||||
if self.agent_id:
|
||||
filters["agent_id"] = self.agent_id
|
||||
if session_id:
|
||||
filters["run_id"] = session_id
|
||||
if self.application_id:
|
||||
filters["app_id"] = self.application_id
|
||||
return filters
|
||||
|
||||
|
||||
__all__ = ["_Mem0ContextProvider"]
|
||||
@@ -0,0 +1,352 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
# pyright: reportPrivateUsage=false
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
from agent_framework import AgentResponse, ChatMessage
|
||||
from agent_framework._sessions import AgentSession, SessionContext
|
||||
from agent_framework.exceptions import ServiceInitializationError
|
||||
|
||||
from agent_framework_mem0._context_provider import _Mem0ContextProvider
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_mem0_client() -> AsyncMock:
|
||||
"""Create a mock Mem0 AsyncMemoryClient."""
|
||||
from mem0 import AsyncMemoryClient
|
||||
|
||||
mock_client = AsyncMock(spec=AsyncMemoryClient)
|
||||
mock_client.add = AsyncMock()
|
||||
mock_client.search = AsyncMock()
|
||||
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
|
||||
mock_client.__aexit__ = AsyncMock()
|
||||
return mock_client
|
||||
|
||||
|
||||
# -- Initialization tests ------------------------------------------------------
|
||||
|
||||
|
||||
class TestInit:
|
||||
"""Test _Mem0ContextProvider initialization."""
|
||||
|
||||
def test_init_with_all_params(self, mock_mem0_client: AsyncMock) -> None:
|
||||
provider = _Mem0ContextProvider(
|
||||
source_id="mem0",
|
||||
mem0_client=mock_mem0_client,
|
||||
api_key="key-123",
|
||||
application_id="app1",
|
||||
agent_id="agent1",
|
||||
user_id="user1",
|
||||
context_prompt="Custom prompt",
|
||||
)
|
||||
assert provider.source_id == "mem0"
|
||||
assert provider.api_key == "key-123"
|
||||
assert provider.application_id == "app1"
|
||||
assert provider.agent_id == "agent1"
|
||||
assert provider.user_id == "user1"
|
||||
assert provider.context_prompt == "Custom prompt"
|
||||
assert provider.mem0_client is mock_mem0_client
|
||||
assert provider._should_close_client is False
|
||||
|
||||
def test_init_default_context_prompt(self, mock_mem0_client: AsyncMock) -> None:
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
assert provider.context_prompt == _Mem0ContextProvider.DEFAULT_CONTEXT_PROMPT
|
||||
|
||||
def test_init_auto_creates_client_when_none(self) -> None:
|
||||
"""When no client is provided, a default AsyncMemoryClient is created and flagged for closing."""
|
||||
with (
|
||||
patch("mem0.client.main.AsyncMemoryClient.__init__", return_value=None) as mock_init,
|
||||
patch("mem0.client.main.AsyncMemoryClient._validate_api_key", return_value=None),
|
||||
):
|
||||
provider = _Mem0ContextProvider(source_id="mem0", api_key="test-key", user_id="u1")
|
||||
mock_init.assert_called_once_with(api_key="test-key")
|
||||
assert provider._should_close_client is True
|
||||
|
||||
def test_provided_client_not_flagged_for_close(self, mock_mem0_client: AsyncMock) -> None:
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
assert provider._should_close_client is False
|
||||
|
||||
|
||||
# -- before_run tests ----------------------------------------------------------
|
||||
|
||||
|
||||
class TestBeforeRun:
|
||||
"""Test before_run hook."""
|
||||
|
||||
async def test_memories_added_to_context(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Mocked mem0 search returns memories → messages added to context with prompt."""
|
||||
mock_mem0_client.search.return_value = [
|
||||
{"memory": "User likes Python"},
|
||||
{"memory": "User prefers dark mode"},
|
||||
]
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="Hello")], session_id="s1")
|
||||
|
||||
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
mock_mem0_client.search.assert_awaited_once()
|
||||
assert "mem0" in ctx.context_messages
|
||||
added = ctx.context_messages["mem0"]
|
||||
assert len(added) == 1
|
||||
assert "User likes Python" in added[0].text # type: ignore[operator]
|
||||
assert "User prefers dark mode" in added[0].text # type: ignore[operator]
|
||||
assert provider.context_prompt in added[0].text # type: ignore[operator]
|
||||
|
||||
async def test_empty_input_skips_search(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Empty input messages → no search performed."""
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="")], session_id="s1")
|
||||
|
||||
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
mock_mem0_client.search.assert_not_awaited()
|
||||
assert "mem0" not in ctx.context_messages
|
||||
|
||||
async def test_empty_search_results_no_messages(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Empty search results → no messages added."""
|
||||
mock_mem0_client.search.return_value = []
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="test")], session_id="s1")
|
||||
|
||||
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
assert "mem0" not in ctx.context_messages
|
||||
|
||||
async def test_validates_filters_before_search(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Raises ServiceInitializationError when no filters."""
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client)
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="test")], session_id="s1")
|
||||
|
||||
with pytest.raises(ServiceInitializationError, match="At least one of the filters"):
|
||||
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
async def test_v1_1_response_format(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Search response in v1.1 dict format with 'results' key."""
|
||||
mock_mem0_client.search.return_value = {"results": [{"memory": "remembered fact"}]}
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="test")], session_id="s1")
|
||||
|
||||
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
added = ctx.context_messages["mem0"]
|
||||
assert "remembered fact" in added[0].text # type: ignore[operator]
|
||||
|
||||
async def test_search_query_combines_input_messages(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Multiple input messages are joined for the search query."""
|
||||
mock_mem0_client.search.return_value = []
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(
|
||||
input_messages=[
|
||||
ChatMessage(role="user", text="Hello"),
|
||||
ChatMessage(role="user", text="World"),
|
||||
],
|
||||
session_id="s1",
|
||||
)
|
||||
|
||||
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
call_kwargs = mock_mem0_client.search.call_args.kwargs
|
||||
assert call_kwargs["query"] == "Hello\nWorld"
|
||||
|
||||
|
||||
# -- after_run tests -----------------------------------------------------------
|
||||
|
||||
|
||||
class TestAfterRun:
|
||||
"""Test after_run hook."""
|
||||
|
||||
async def test_stores_input_and_response(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Stores input+response messages to mem0 via client.add."""
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="question")], session_id="s1")
|
||||
ctx._response = AgentResponse(messages=[ChatMessage(role="assistant", text="answer")])
|
||||
|
||||
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
mock_mem0_client.add.assert_awaited_once()
|
||||
call_kwargs = mock_mem0_client.add.call_args.kwargs
|
||||
assert call_kwargs["messages"] == [
|
||||
{"role": "user", "content": "question"},
|
||||
{"role": "assistant", "content": "answer"},
|
||||
]
|
||||
assert call_kwargs["user_id"] == "u1"
|
||||
assert call_kwargs["run_id"] == "s1"
|
||||
|
||||
async def test_only_stores_user_assistant_system(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Only stores user/assistant/system messages with text."""
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(
|
||||
input_messages=[
|
||||
ChatMessage(role="user", text="hello"),
|
||||
ChatMessage(role="tool", text="tool output"),
|
||||
],
|
||||
session_id="s1",
|
||||
)
|
||||
ctx._response = AgentResponse(messages=[ChatMessage(role="assistant", text="reply")])
|
||||
|
||||
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
call_kwargs = mock_mem0_client.add.call_args.kwargs
|
||||
roles = [m["role"] for m in call_kwargs["messages"]]
|
||||
assert "tool" not in roles
|
||||
assert roles == ["user", "assistant"]
|
||||
|
||||
async def test_skips_empty_messages(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Skips messages with empty text."""
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(
|
||||
input_messages=[
|
||||
ChatMessage(role="user", text=""),
|
||||
ChatMessage(role="user", text=" "),
|
||||
],
|
||||
session_id="s1",
|
||||
)
|
||||
ctx._response = AgentResponse(messages=[])
|
||||
|
||||
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
mock_mem0_client.add.assert_not_awaited()
|
||||
|
||||
async def test_uses_session_id_as_run_id(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Uses session_id as run_id."""
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="hi")], session_id="my-session")
|
||||
ctx._response = AgentResponse(messages=[ChatMessage(role="assistant", text="hey")])
|
||||
|
||||
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
assert mock_mem0_client.add.call_args.kwargs["run_id"] == "my-session"
|
||||
|
||||
async def test_validates_filters(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Raises ServiceInitializationError when no filters."""
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client)
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="hi")], session_id="s1")
|
||||
ctx._response = AgentResponse(messages=[ChatMessage(role="assistant", text="hey")])
|
||||
|
||||
with pytest.raises(ServiceInitializationError, match="At least one of the filters"):
|
||||
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
async def test_stores_with_application_id_metadata(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""application_id is passed in metadata."""
|
||||
provider = _Mem0ContextProvider(
|
||||
source_id="mem0", mem0_client=mock_mem0_client, user_id="u1", application_id="app1"
|
||||
)
|
||||
session = AgentSession(session_id="test-session")
|
||||
ctx = SessionContext(input_messages=[ChatMessage(role="user", text="hi")], session_id="s1")
|
||||
ctx._response = AgentResponse(messages=[])
|
||||
|
||||
await provider.after_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
|
||||
|
||||
assert mock_mem0_client.add.call_args.kwargs["metadata"] == {"application_id": "app1"}
|
||||
|
||||
|
||||
# -- _validate_filters tests --------------------------------------------------
|
||||
|
||||
|
||||
class TestValidateFilters:
|
||||
"""Test _validate_filters method."""
|
||||
|
||||
def test_raises_when_no_filters(self, mock_mem0_client: AsyncMock) -> None:
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client)
|
||||
with pytest.raises(ServiceInitializationError, match="At least one of the filters"):
|
||||
provider._validate_filters()
|
||||
|
||||
def test_passes_with_user_id(self, mock_mem0_client: AsyncMock) -> None:
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
provider._validate_filters() # should not raise
|
||||
|
||||
def test_passes_with_agent_id(self, mock_mem0_client: AsyncMock) -> None:
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, agent_id="a1")
|
||||
provider._validate_filters()
|
||||
|
||||
def test_passes_with_application_id(self, mock_mem0_client: AsyncMock) -> None:
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, application_id="app1")
|
||||
provider._validate_filters()
|
||||
|
||||
|
||||
# -- _build_filters tests -----------------------------------------------------
|
||||
|
||||
|
||||
class TestBuildFilters:
|
||||
"""Test _build_filters method."""
|
||||
|
||||
def test_user_id_only(self, mock_mem0_client: AsyncMock) -> None:
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
assert provider._build_filters() == {"user_id": "u1"}
|
||||
|
||||
def test_all_params(self, mock_mem0_client: AsyncMock) -> None:
|
||||
provider = _Mem0ContextProvider(
|
||||
source_id="mem0",
|
||||
mem0_client=mock_mem0_client,
|
||||
user_id="u1",
|
||||
agent_id="a1",
|
||||
application_id="app1",
|
||||
)
|
||||
assert provider._build_filters(session_id="sess1") == {
|
||||
"user_id": "u1",
|
||||
"agent_id": "a1",
|
||||
"run_id": "sess1",
|
||||
"app_id": "app1",
|
||||
}
|
||||
|
||||
def test_excludes_none_values(self, mock_mem0_client: AsyncMock) -> None:
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
filters = provider._build_filters()
|
||||
assert "agent_id" not in filters
|
||||
assert "run_id" not in filters
|
||||
assert "app_id" not in filters
|
||||
|
||||
def test_session_id_mapped_to_run_id(self, mock_mem0_client: AsyncMock) -> None:
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
filters = provider._build_filters(session_id="s99")
|
||||
assert filters["run_id"] == "s99"
|
||||
|
||||
def test_empty_when_no_params(self, mock_mem0_client: AsyncMock) -> None:
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client)
|
||||
assert provider._build_filters() == {}
|
||||
|
||||
|
||||
# -- Context manager tests -----------------------------------------------------
|
||||
|
||||
|
||||
class TestContextManager:
|
||||
"""Test __aenter__/__aexit__ delegation."""
|
||||
|
||||
async def test_aenter_delegates_to_client(self, mock_mem0_client: AsyncMock) -> None:
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
result = await provider.__aenter__()
|
||||
assert result is provider
|
||||
mock_mem0_client.__aenter__.assert_awaited_once()
|
||||
|
||||
async def test_aexit_closes_auto_created_client(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Auto-created clients (_should_close_client=True) are closed on exit."""
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
provider._should_close_client = True
|
||||
await provider.__aexit__(None, None, None)
|
||||
mock_mem0_client.__aexit__.assert_awaited_once()
|
||||
|
||||
async def test_aexit_does_not_close_provided_client(self, mock_mem0_client: AsyncMock) -> None:
|
||||
"""Provided clients (_should_close_client=False) are NOT closed on exit."""
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
assert provider._should_close_client is False
|
||||
await provider.__aexit__(None, None, None)
|
||||
mock_mem0_client.__aexit__.assert_not_awaited()
|
||||
|
||||
async def test_async_with_syntax(self, mock_mem0_client: AsyncMock) -> None:
|
||||
provider = _Mem0ContextProvider(source_id="mem0", mem0_client=mock_mem0_client, user_id="u1")
|
||||
async with provider as p:
|
||||
assert p is provider
|
||||
Reference in New Issue
Block a user