Python: [BREAKING] cleanup of thread API and serialization (#893)

* cleanup of threads and serialization

* fix for sliding window

* fix redis test

* updated from comments

* updated context provider and threads

* updated lock

* add asyncio default

* fix redis tests

* fix tests

* fix tests

* renamed to invoking

* fixed tests

* fix for instructions
This commit is contained in:
Eduard van Valkenburg
2025-09-29 18:22:34 +02:00
committed by GitHub
Unverified
parent bf5931932e
commit 10d10364a9
52 changed files with 1642 additions and 1411 deletions
@@ -4,7 +4,7 @@
from unittest.mock import AsyncMock, patch
import pytest
from agent_framework import ChatMessage, Context, Role, TextContent
from agent_framework import ChatMessage, Context, Role
from agent_framework.exceptions import ServiceInitializationError
from agent_framework.mem0 import Mem0Provider
@@ -173,27 +173,25 @@ class TestMem0ProviderThreadMethods:
assert "can only be used with one thread at a time" in str(exc_info.value)
async def test_messages_adding_sets_per_operation_thread_id(
self, mock_mem0_client: AsyncMock, sample_messages: list[ChatMessage]
) -> None:
"""Test that messages_adding sets per-operation thread ID."""
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.messages_adding("thread123", sample_messages)
await provider.thread_created("thread123")
assert provider._per_operation_thread_id == "thread123"
class TestMem0ProviderMessagesAdding:
"""Test messages_adding method."""
"""Test invoked method."""
async def test_messages_adding_fails_without_filters(self, mock_mem0_client: AsyncMock) -> None:
"""Test that messages_adding fails when no filters are provided."""
"""Test that invoked fails when no filters are provided."""
provider = Mem0Provider(mem0_client=mock_mem0_client)
message = ChatMessage(role=Role.USER, text="Hello!")
with pytest.raises(ServiceInitializationError) as exc_info:
await provider.messages_adding("thread123", message)
await provider.invoked(message)
assert "At least one of the filters" in str(exc_info.value)
@@ -202,7 +200,7 @@ class TestMem0ProviderMessagesAdding:
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
message = ChatMessage(role=Role.USER, text="Hello!")
await provider.messages_adding("thread123", message)
await provider.invoked(message)
mock_mem0_client.add.assert_called_once()
call_args = mock_mem0_client.add.call_args
@@ -215,7 +213,7 @@ class TestMem0ProviderMessagesAdding:
"""Test adding multiple messages."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
await provider.messages_adding("thread123", sample_messages)
await provider.invoked(sample_messages)
mock_mem0_client.add.assert_called_once()
call_args = mock_mem0_client.add.call_args
@@ -232,7 +230,7 @@ class TestMem0ProviderMessagesAdding:
"""Test adding messages with agent_id."""
provider = Mem0Provider(agent_id="agent123", mem0_client=mock_mem0_client)
await provider.messages_adding("thread123", sample_messages)
await provider.invoked(sample_messages)
call_args = mock_mem0_client.add.call_args
assert call_args.kwargs["agent_id"] == "agent123"
@@ -244,7 +242,7 @@ class TestMem0ProviderMessagesAdding:
"""Test adding messages with application_id in metadata."""
provider = Mem0Provider(user_id="user123", application_id="app123", mem0_client=mock_mem0_client)
await provider.messages_adding("thread123", sample_messages)
await provider.invoked(sample_messages)
call_args = mock_mem0_client.add.call_args
assert call_args.kwargs["metadata"] == {"application_id": "app123"}
@@ -261,7 +259,8 @@ class TestMem0ProviderMessagesAdding:
)
provider._per_operation_thread_id = "operation_thread"
await provider.messages_adding("operation_thread", sample_messages)
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"
@@ -277,7 +276,7 @@ class TestMem0ProviderMessagesAdding:
mem0_client=mock_mem0_client,
)
await provider.messages_adding("operation_thread", sample_messages)
await provider.invoked(sample_messages)
call_args = mock_mem0_client.add.call_args
assert call_args.kwargs["run_id"] == "base_thread"
@@ -291,7 +290,7 @@ class TestMem0ProviderMessagesAdding:
ChatMessage(role=Role.USER, text="Valid message"),
]
await provider.messages_adding("thread123", messages)
await provider.invoked(messages)
call_args = mock_mem0_client.add.call_args
# Should only include the valid message
@@ -305,26 +304,26 @@ class TestMem0ProviderMessagesAdding:
ChatMessage(role=Role.USER, text=" "),
]
await provider.messages_adding("thread123", messages)
await provider.invoked(messages)
mock_mem0_client.add.assert_not_called()
class TestMem0ProviderModelInvoking:
"""Test model_invoking method."""
"""Test invoking method."""
async def test_model_invoking_fails_without_filters(self, mock_mem0_client: AsyncMock) -> None:
"""Test that model_invoking fails when no filters are provided."""
"""Test that invoking fails when no filters are provided."""
provider = Mem0Provider(mem0_client=mock_mem0_client)
message = ChatMessage(role=Role.USER, text="What's the weather?")
with pytest.raises(ServiceInitializationError) as exc_info:
await provider.model_invoking(message)
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 model_invoking with a single message."""
"""Test invoking with a single message."""
provider = Mem0Provider(user_id="user123", mem0_client=mock_mem0_client)
message = ChatMessage(role=Role.USER, text="What's the weather?")
@@ -334,7 +333,7 @@ class TestMem0ProviderModelInvoking:
{"memory": "User lives in Seattle"},
]
context = await provider.model_invoking(message)
context = await provider.invoking(message)
mock_mem0_client.search.assert_called_once()
call_args = mock_mem0_client.search.call_args
@@ -347,39 +346,38 @@ class TestMem0ProviderModelInvoking:
"User likes outdoor activities\nUser lives in Seattle"
)
assert context.contents
assert isinstance(context.contents[0], TextContent)
assert context.contents[0].text == expected_instructions
assert context.messages
assert context.messages[0].text == expected_instructions
async def test_model_invoking_multiple_messages(
self, mock_mem0_client: AsyncMock, sample_messages: list[ChatMessage]
) -> None:
"""Test model_invoking with multiple messages."""
"""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.model_invoking(sample_messages)
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 model_invoking with agent_id."""
"""Test invoking with agent_id."""
provider = Mem0Provider(agent_id="agent123", mem0_client=mock_mem0_client)
message = ChatMessage(role=Role.USER, text="Hello")
mock_mem0_client.search.return_value = []
await provider.model_invoking(message)
await provider.invoking(message)
call_args = mock_mem0_client.search.call_args
assert call_args.kwargs["agent_id"] == "agent123"
assert call_args.kwargs["user_id"] is None
async def test_model_invoking_with_scope_to_per_operation_thread_id(self, mock_mem0_client: AsyncMock) -> None:
"""Test model_invoking with scope_to_per_operation_thread_id enabled."""
"""Test invoking with scope_to_per_operation_thread_id enabled."""
provider = Mem0Provider(
user_id="user123",
thread_id="base_thread",
@@ -391,7 +389,7 @@ class TestMem0ProviderModelInvoking:
mock_mem0_client.search.return_value = []
await provider.model_invoking(message)
await provider.invoking(message)
call_args = mock_mem0_client.search.call_args
assert call_args.kwargs["run_id"] == "operation_thread"
@@ -403,10 +401,10 @@ class TestMem0ProviderModelInvoking:
mock_mem0_client.search.return_value = []
context = await provider.model_invoking(message)
context = await provider.invoking(message)
assert isinstance(context, Context)
assert not context.contents
assert not context.messages
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."""
@@ -419,13 +417,13 @@ class TestMem0ProviderModelInvoking:
mock_mem0_client.search.return_value = []
await provider.model_invoking(messages)
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 model_invoking with custom context prompt."""
"""Test invoking with custom context prompt."""
custom_prompt = "## Custom Context\nRemember these details:"
provider = Mem0Provider(
user_id="user123",
@@ -436,12 +434,11 @@ class TestMem0ProviderModelInvoking:
mock_mem0_client.search.return_value = [{"memory": "Test memory"}]
context = await provider.model_invoking(message)
context = await provider.invoking(message)
expected_instructions = "## Custom Context\nRemember these details:\nTest memory"
assert context.contents
assert isinstance(context.contents[0], TextContent)
assert context.contents[0].text == expected_instructions
assert context.messages
assert context.messages[0].text == expected_instructions
class TestMem0ProviderValidation: