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