mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: [BREAKING]: removed display_name, renamed context_providers, middleware and AggregateContextProvider (#3139)
* removed display_name, renamed context_providers, middleware and AggregateContextProvider * fixes * fixed test * testfix * removed mistakenly put back test * updated new test * rename middlewares to middleware * middleware fixes
This commit is contained in:
committed by
GitHub
Unverified
parent
ef44fb4960
commit
203fb7b1c4
@@ -221,11 +221,6 @@ class MockAgent(AgentProtocol):
|
||||
"""Returns the name of the agent."""
|
||||
return "Name"
|
||||
|
||||
@property
|
||||
def display_name(self) -> str:
|
||||
"""Returns the name of the agent."""
|
||||
return "Display Name"
|
||||
|
||||
@property
|
||||
def description(self) -> str | None:
|
||||
return "Description"
|
||||
|
||||
@@ -13,7 +13,6 @@ from agent_framework import (
|
||||
AgentRunResponse,
|
||||
AgentRunResponseUpdate,
|
||||
AgentThread,
|
||||
AggregateContextProvider,
|
||||
ChatAgent,
|
||||
ChatClientProtocol,
|
||||
ChatMessage,
|
||||
@@ -66,7 +65,6 @@ async def test_chat_client_agent_init(chat_client: ChatClientProtocol) -> None:
|
||||
assert agent.id == agent_id
|
||||
assert agent.name is None
|
||||
assert agent.description == "Test"
|
||||
assert agent.display_name == agent_id # Display name defaults to id if name is None
|
||||
|
||||
|
||||
async def test_chat_client_agent_init_with_name(chat_client: ChatClientProtocol) -> None:
|
||||
@@ -76,7 +74,6 @@ async def test_chat_client_agent_init_with_name(chat_client: ChatClientProtocol)
|
||||
assert agent.id == agent_id
|
||||
assert agent.name == "Test Agent"
|
||||
assert agent.description == "Test"
|
||||
assert agent.display_name == "Test Agent" # Display name is the name if present
|
||||
|
||||
|
||||
async def test_chat_client_agent_run(chat_client: ChatClientProtocol) -> None:
|
||||
@@ -255,7 +252,7 @@ class MockContextProvider(ContextProvider):
|
||||
async def test_chat_agent_context_providers_model_invoking(chat_client: ChatClientProtocol) -> None:
|
||||
"""Test that context providers' invoking is called during agent run."""
|
||||
mock_provider = MockContextProvider(messages=[ChatMessage(role=Role.SYSTEM, text="Test context instructions")])
|
||||
agent = ChatAgent(chat_client=chat_client, context_providers=mock_provider)
|
||||
agent = ChatAgent(chat_client=chat_client, context_provider=mock_provider)
|
||||
|
||||
await agent.run("Hello")
|
||||
|
||||
@@ -272,7 +269,7 @@ async def test_chat_agent_context_providers_thread_created(chat_client_base: Cha
|
||||
)
|
||||
]
|
||||
|
||||
agent = ChatAgent(chat_client=chat_client_base, context_providers=mock_provider)
|
||||
agent = ChatAgent(chat_client=chat_client_base, context_provider=mock_provider)
|
||||
|
||||
await agent.run("Hello")
|
||||
|
||||
@@ -283,7 +280,7 @@ async def test_chat_agent_context_providers_thread_created(chat_client_base: Cha
|
||||
async def test_chat_agent_context_providers_messages_adding(chat_client: ChatClientProtocol) -> None:
|
||||
"""Test that context providers' invoked is called during agent run."""
|
||||
mock_provider = MockContextProvider()
|
||||
agent = ChatAgent(chat_client=chat_client, context_providers=mock_provider)
|
||||
agent = ChatAgent(chat_client=chat_client, context_provider=mock_provider)
|
||||
|
||||
await agent.run("Hello")
|
||||
|
||||
@@ -295,7 +292,7 @@ async def test_chat_agent_context_providers_messages_adding(chat_client: ChatCli
|
||||
async def test_chat_agent_context_instructions_in_messages(chat_client: ChatClientProtocol) -> None:
|
||||
"""Test that AI context instructions are included in messages."""
|
||||
mock_provider = MockContextProvider(messages=[ChatMessage(role="system", text="Context-specific instructions")])
|
||||
agent = ChatAgent(chat_client=chat_client, instructions="Agent instructions", context_providers=mock_provider)
|
||||
agent = ChatAgent(chat_client=chat_client, instructions="Agent instructions", context_provider=mock_provider)
|
||||
|
||||
# We need to test the _prepare_thread_and_messages method directly
|
||||
_, _, messages = await agent._prepare_thread_and_messages( # type: ignore[reportPrivateUsage]
|
||||
@@ -314,7 +311,7 @@ async def test_chat_agent_context_instructions_in_messages(chat_client: ChatClie
|
||||
async def test_chat_agent_no_context_instructions(chat_client: ChatClientProtocol) -> None:
|
||||
"""Test behavior when AI context has no instructions."""
|
||||
mock_provider = MockContextProvider()
|
||||
agent = ChatAgent(chat_client=chat_client, instructions="Agent instructions", context_providers=mock_provider)
|
||||
agent = ChatAgent(chat_client=chat_client, instructions="Agent instructions", context_provider=mock_provider)
|
||||
|
||||
_, _, messages = await agent._prepare_thread_and_messages( # type: ignore[reportPrivateUsage]
|
||||
thread=None, input_messages=[ChatMessage(role=Role.USER, text="Hello")]
|
||||
@@ -329,7 +326,7 @@ async def test_chat_agent_no_context_instructions(chat_client: ChatClientProtoco
|
||||
async def test_chat_agent_run_stream_context_providers(chat_client: ChatClientProtocol) -> None:
|
||||
"""Test that context providers work with run_stream method."""
|
||||
mock_provider = MockContextProvider(messages=[ChatMessage(role=Role.SYSTEM, text="Stream context instructions")])
|
||||
agent = ChatAgent(chat_client=chat_client, context_providers=mock_provider)
|
||||
agent = ChatAgent(chat_client=chat_client, context_provider=mock_provider)
|
||||
|
||||
# Collect all stream updates
|
||||
updates: list[AgentRunResponseUpdate] = []
|
||||
@@ -343,44 +340,6 @@ async def test_chat_agent_run_stream_context_providers(chat_client: ChatClientPr
|
||||
assert mock_provider.invoked_called
|
||||
|
||||
|
||||
async def test_chat_agent_multiple_context_providers(chat_client: ChatClientProtocol) -> None:
|
||||
"""Test that multiple context providers work together."""
|
||||
provider1 = MockContextProvider(messages=[ChatMessage(role=Role.SYSTEM, text="First provider instructions")])
|
||||
provider2 = MockContextProvider(messages=[ChatMessage(role=Role.SYSTEM, text="Second provider instructions")])
|
||||
|
||||
agent = ChatAgent(chat_client=chat_client, context_providers=[provider1, provider2])
|
||||
|
||||
await agent.run("Hello")
|
||||
|
||||
# Both providers should be called
|
||||
assert provider1.invoking_called
|
||||
assert not provider1.thread_created_called
|
||||
assert provider1.invoked_called
|
||||
|
||||
assert provider2.invoking_called
|
||||
assert not provider2.thread_created_called
|
||||
assert provider2.invoked_called
|
||||
|
||||
|
||||
async def test_chat_agent_aggregate_context_provider_combines_instructions() -> None:
|
||||
"""Test that AggregateContextProvider combines instructions from multiple providers."""
|
||||
provider1 = MockContextProvider(messages=[ChatMessage(role=Role.SYSTEM, text="First instruction")])
|
||||
provider2 = MockContextProvider(messages=[ChatMessage(role=Role.SYSTEM, text="Second instruction")])
|
||||
|
||||
aggregate = AggregateContextProvider()
|
||||
aggregate.providers.append(provider1)
|
||||
aggregate.providers.append(provider2)
|
||||
|
||||
# Test invoking combines instructions
|
||||
result = await aggregate.invoking([ChatMessage(role=Role.USER, text="Test")])
|
||||
|
||||
assert result.messages
|
||||
assert isinstance(result.messages[0], ChatMessage)
|
||||
assert isinstance(result.messages[1], ChatMessage)
|
||||
assert result.messages[0].text == "First instruction"
|
||||
assert result.messages[1].text == "Second instruction"
|
||||
|
||||
|
||||
async def test_chat_agent_context_providers_with_thread_service_id(chat_client_base: ChatClientProtocol) -> None:
|
||||
"""Test context providers with service-managed thread."""
|
||||
mock_provider = MockContextProvider()
|
||||
@@ -391,7 +350,7 @@ async def test_chat_agent_context_providers_with_thread_service_id(chat_client_b
|
||||
)
|
||||
]
|
||||
|
||||
agent = ChatAgent(chat_client=chat_client_base, context_providers=mock_provider)
|
||||
agent = ChatAgent(chat_client=chat_client_base, context_provider=mock_provider)
|
||||
|
||||
# Use existing service-managed thread
|
||||
thread = agent.get_new_thread(service_thread_id="existing-thread-id")
|
||||
|
||||
@@ -2,10 +2,9 @@
|
||||
|
||||
from collections.abc import MutableSequence
|
||||
from typing import Any
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
from agent_framework import ChatMessage, Role, TextContent
|
||||
from agent_framework._memory import AggregateContextProvider, Context, ContextProvider
|
||||
from agent_framework import ChatMessage, Role
|
||||
from agent_framework._memory import Context, ContextProvider
|
||||
|
||||
|
||||
class MockContextProvider(ContextProvider):
|
||||
@@ -45,252 +44,50 @@ class MockContextProvider(ContextProvider):
|
||||
return context
|
||||
|
||||
|
||||
class TestAggregateContextProvider:
|
||||
"""Tests for AggregateContextProvider class."""
|
||||
class TestContext:
|
||||
"""Tests for Context class."""
|
||||
|
||||
def test_init_with_no_providers(self) -> None:
|
||||
"""Test initialization with no providers."""
|
||||
aggregate = AggregateContextProvider()
|
||||
assert aggregate.providers == []
|
||||
def test_context_default_values(self) -> None:
|
||||
"""Test Context has correct default values."""
|
||||
context = Context()
|
||||
assert context.instructions is None
|
||||
assert context.messages == []
|
||||
assert context.tools == []
|
||||
|
||||
def test_init_with_none_providers(self) -> None:
|
||||
"""Test initialization with None providers."""
|
||||
aggregate = AggregateContextProvider(None)
|
||||
assert aggregate.providers == []
|
||||
def test_context_with_values(self) -> None:
|
||||
"""Test Context can be initialized with values."""
|
||||
messages = [ChatMessage(role=Role.USER, text="Test message")]
|
||||
context = Context(instructions="Test instructions", messages=messages)
|
||||
assert context.instructions == "Test instructions"
|
||||
assert len(context.messages) == 1
|
||||
assert context.messages[0].text == "Test message"
|
||||
|
||||
def test_init_with_providers(self) -> None:
|
||||
"""Test initialization with providers."""
|
||||
provider1 = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions 1")])
|
||||
provider2 = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions 2")])
|
||||
provider3 = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions 3")])
|
||||
providers = [provider1, provider2, provider3]
|
||||
|
||||
aggregate = AggregateContextProvider(providers)
|
||||
assert len(aggregate.providers) == 3
|
||||
assert aggregate.providers[0] is provider1
|
||||
assert aggregate.providers[1] is provider2
|
||||
assert aggregate.providers[2] is provider3
|
||||
|
||||
def test_add_provider(self) -> None:
|
||||
"""Test adding a provider."""
|
||||
aggregate = AggregateContextProvider()
|
||||
provider = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions")])
|
||||
|
||||
aggregate.add(provider)
|
||||
assert len(aggregate.providers) == 1
|
||||
assert aggregate.providers[0] is provider
|
||||
|
||||
def test_add_multiple_providers(self) -> None:
|
||||
"""Test adding multiple providers."""
|
||||
aggregate = AggregateContextProvider()
|
||||
provider1 = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions 1")])
|
||||
provider2 = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions 2")])
|
||||
|
||||
aggregate.add(provider1)
|
||||
aggregate.add(provider2)
|
||||
|
||||
assert len(aggregate.providers) == 2
|
||||
assert aggregate.providers[0] is provider1
|
||||
assert aggregate.providers[1] is provider2
|
||||
|
||||
async def test_thread_created_with_no_providers(self) -> None:
|
||||
"""Test thread_created with no providers."""
|
||||
aggregate = AggregateContextProvider()
|
||||
|
||||
# Should not raise an exception
|
||||
await aggregate.thread_created("thread-123")
|
||||
|
||||
async def test_thread_created_with_providers(self) -> None:
|
||||
"""Test thread_created calls all providers."""
|
||||
provider1 = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions 1")])
|
||||
provider2 = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions 2")])
|
||||
aggregate = AggregateContextProvider([provider1, provider2])
|
||||
|
||||
thread_id = "thread-123"
|
||||
await aggregate.thread_created(thread_id)
|
||||
|
||||
assert provider1.thread_created_called
|
||||
assert provider1.thread_created_thread_id == thread_id
|
||||
assert provider2.thread_created_called
|
||||
assert provider2.thread_created_thread_id == thread_id
|
||||
|
||||
async def test_thread_created_with_none_thread_id(self) -> None:
|
||||
"""Test thread_created with None thread_id."""
|
||||
provider = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions")])
|
||||
aggregate = AggregateContextProvider([provider])
|
||||
|
||||
await aggregate.thread_created(None)
|
||||
class TestContextProvider:
|
||||
"""Tests for ContextProvider class."""
|
||||
|
||||
async def test_thread_created(self) -> None:
|
||||
"""Test thread_created is called."""
|
||||
provider = MockContextProvider()
|
||||
await provider.thread_created("test-thread-id")
|
||||
assert provider.thread_created_called
|
||||
assert provider.thread_created_thread_id is None
|
||||
|
||||
async def test_messages_adding_with_no_providers(self) -> None:
|
||||
"""Test invoked with no providers."""
|
||||
aggregate = AggregateContextProvider()
|
||||
message = ChatMessage(text="Hello", role=Role.USER)
|
||||
|
||||
# Should not raise an exception
|
||||
await aggregate.invoked(message)
|
||||
|
||||
async def test_messages_adding_with_single_message(self) -> None:
|
||||
"""Test invoked with a single message."""
|
||||
provider1 = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions 1")])
|
||||
provider2 = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions 2")])
|
||||
aggregate = AggregateContextProvider([provider1, provider2])
|
||||
|
||||
message = ChatMessage(text="Hello", role=Role.USER)
|
||||
await aggregate.invoked(message)
|
||||
|
||||
assert provider1.invoked_called
|
||||
assert provider1.new_messages == message
|
||||
assert provider2.invoked_called
|
||||
assert provider2.new_messages == message
|
||||
|
||||
async def test_messages_adding_with_message_sequence(self) -> None:
|
||||
"""Test invoked with a sequence of messages."""
|
||||
provider = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions")])
|
||||
aggregate = AggregateContextProvider([provider])
|
||||
|
||||
messages = [
|
||||
ChatMessage(text="Hello", role=Role.USER),
|
||||
ChatMessage(text="Hi there", role=Role.ASSISTANT),
|
||||
]
|
||||
await aggregate.invoked(messages)
|
||||
assert provider.thread_created_thread_id == "test-thread-id"
|
||||
|
||||
async def test_invoked(self) -> None:
|
||||
"""Test invoked is called."""
|
||||
provider = MockContextProvider()
|
||||
message = ChatMessage(role=Role.USER, text="Test message")
|
||||
await provider.invoked(message)
|
||||
assert provider.invoked_called
|
||||
assert provider.new_messages == messages
|
||||
|
||||
async def test_model_invoking_with_no_providers(self) -> None:
|
||||
"""Test invoking with no providers."""
|
||||
aggregate = AggregateContextProvider()
|
||||
message = ChatMessage(text="Hello", role=Role.USER)
|
||||
|
||||
context = await aggregate.invoking(message)
|
||||
|
||||
assert isinstance(context, Context)
|
||||
assert not context.messages
|
||||
|
||||
async def test_model_invoking_with_single_provider(self) -> None:
|
||||
"""Test invoking with a single provider."""
|
||||
provider = MockContextProvider(messages=[ChatMessage(role="user", text="Test instructions")])
|
||||
aggregate = AggregateContextProvider([provider])
|
||||
|
||||
message = [ChatMessage(text="Hello", role=Role.USER)]
|
||||
context = await aggregate.invoking(message)
|
||||
assert provider.new_messages == message
|
||||
|
||||
async def test_invoking(self) -> None:
|
||||
"""Test invoking is called and returns context."""
|
||||
provider = MockContextProvider(messages=[ChatMessage(role=Role.USER, text="Context message")])
|
||||
message = ChatMessage(role=Role.USER, text="Test message")
|
||||
context = await provider.invoking(message)
|
||||
assert provider.invoking_called
|
||||
assert provider.model_invoking_messages == message
|
||||
assert isinstance(context, Context)
|
||||
|
||||
assert context.messages
|
||||
assert isinstance(context.messages[0].contents[0], TextContent)
|
||||
assert context.messages[0].text == "Test instructions"
|
||||
|
||||
async def test_model_invoking_with_multiple_providers(self) -> None:
|
||||
"""Test invoking combines contexts from multiple providers."""
|
||||
provider1 = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions 1")])
|
||||
provider2 = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions 2")])
|
||||
provider3 = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions 3")])
|
||||
aggregate = AggregateContextProvider([provider1, provider2, provider3])
|
||||
|
||||
messages = [ChatMessage(text="Hello", role=Role.USER)]
|
||||
context = await aggregate.invoking(messages)
|
||||
|
||||
assert provider1.invoking_called
|
||||
assert provider1.model_invoking_messages == messages
|
||||
assert provider2.invoking_called
|
||||
assert provider2.model_invoking_messages == messages
|
||||
assert provider3.invoking_called
|
||||
assert provider3.model_invoking_messages == messages
|
||||
|
||||
assert isinstance(context, Context)
|
||||
|
||||
assert context.messages
|
||||
assert isinstance(context.messages[0].contents[0], TextContent)
|
||||
assert isinstance(context.messages[1].contents[0], TextContent)
|
||||
assert isinstance(context.messages[2].contents[0], TextContent)
|
||||
assert context.messages[0].text == "Instructions 1"
|
||||
assert context.messages[1].text == "Instructions 2"
|
||||
assert context.messages[2].text == "Instructions 3"
|
||||
|
||||
async def test_model_invoking_with_none_instructions(self) -> None:
|
||||
"""Test invoking filters out None instructions."""
|
||||
provider1 = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions 1")])
|
||||
provider2 = MockContextProvider(messages=None) # None instructions
|
||||
provider3 = MockContextProvider(messages=[ChatMessage(role="user", text="Instructions 3")])
|
||||
aggregate = AggregateContextProvider([provider1, provider2, provider3])
|
||||
|
||||
message = ChatMessage(text="Hello", role=Role.USER)
|
||||
context = await aggregate.invoking(message)
|
||||
|
||||
assert isinstance(context, Context)
|
||||
assert context.messages
|
||||
assert isinstance(context.messages[0].contents[0], TextContent)
|
||||
assert isinstance(context.messages[1].contents[0], TextContent)
|
||||
assert context.messages[0].text == "Instructions 1"
|
||||
assert context.messages[1].text == "Instructions 3"
|
||||
|
||||
async def test_model_invoking_with_all_none_instructions(self) -> None:
|
||||
"""Test invoking when all providers return None instructions."""
|
||||
provider1 = MockContextProvider(None)
|
||||
provider2 = MockContextProvider(None)
|
||||
aggregate = AggregateContextProvider([provider1, provider2])
|
||||
|
||||
message = ChatMessage(text="Hello", role=Role.USER)
|
||||
context = await aggregate.invoking(message)
|
||||
|
||||
assert isinstance(context, Context)
|
||||
assert not context.messages
|
||||
|
||||
async def test_model_invoking_with_mutable_sequence(self) -> None:
|
||||
"""Test invoking with MutableSequence of messages."""
|
||||
provider = MockContextProvider(messages=[ChatMessage(role="user", text="Test instructions")])
|
||||
aggregate = AggregateContextProvider([provider])
|
||||
|
||||
messages = [ChatMessage(text="Hello", role=Role.USER)]
|
||||
context = await aggregate.invoking(messages)
|
||||
|
||||
assert provider.invoking_called
|
||||
assert provider.model_invoking_messages == messages
|
||||
assert isinstance(context, Context)
|
||||
assert context.messages
|
||||
assert isinstance(context.messages[0].contents[0], TextContent)
|
||||
assert context.messages[0].text == "Test instructions"
|
||||
|
||||
async def test_async_methods_concurrent_execution(self) -> None:
|
||||
"""Test that async methods execute providers concurrently."""
|
||||
# Use AsyncMock to verify concurrent execution
|
||||
provider1 = Mock(spec=ContextProvider)
|
||||
provider1.thread_created = AsyncMock()
|
||||
provider1.invoked = AsyncMock()
|
||||
provider1.invoking = AsyncMock(return_value=Context(messages=[ChatMessage(role="user", text="Test 1")]))
|
||||
|
||||
provider2 = Mock(spec=ContextProvider)
|
||||
provider2.thread_created = AsyncMock()
|
||||
provider2.invoked = AsyncMock()
|
||||
provider2.invoking = AsyncMock(return_value=Context(messages=[ChatMessage(role="user", text="Test 2")]))
|
||||
|
||||
aggregate = AggregateContextProvider([provider1, provider2])
|
||||
|
||||
# Test thread_created
|
||||
await aggregate.thread_created("thread-123")
|
||||
provider1.thread_created.assert_called_once_with("thread-123")
|
||||
provider2.thread_created.assert_called_once_with("thread-123")
|
||||
|
||||
# Test invoked
|
||||
message = ChatMessage(text="Hello", role=Role.USER)
|
||||
await aggregate.invoked(message)
|
||||
provider1.invoked.assert_called_once_with(
|
||||
request_messages=message, response_messages=None, invoke_exception=None
|
||||
)
|
||||
provider2.invoked.assert_called_once_with(
|
||||
request_messages=message, response_messages=None, invoke_exception=None
|
||||
)
|
||||
|
||||
# Test invoking
|
||||
context = await aggregate.invoking(message)
|
||||
provider1.invoking.assert_called_once_with(message)
|
||||
provider2.invoking.assert_called_once_with(message)
|
||||
assert context.messages
|
||||
assert context.messages[0].text == "Test 1"
|
||||
assert context.messages[1].text == "Test 2"
|
||||
assert context.messages is not None
|
||||
assert len(context.messages) == 1
|
||||
assert context.messages[0].text == "Context message"
|
||||
|
||||
@@ -148,7 +148,7 @@ class TestAgentMiddlewarePipeline:
|
||||
context.terminate = True
|
||||
|
||||
def test_init_empty(self) -> None:
|
||||
"""Test AgentMiddlewarePipeline initialization with no middlewares."""
|
||||
"""Test AgentMiddlewarePipeline initialization with no middleware."""
|
||||
pipeline = AgentMiddlewarePipeline()
|
||||
assert not pipeline.has_middlewares
|
||||
|
||||
@@ -457,7 +457,7 @@ class TestFunctionMiddlewarePipeline:
|
||||
assert execution_order == ["handler"]
|
||||
|
||||
def test_init_empty(self) -> None:
|
||||
"""Test FunctionMiddlewarePipeline initialization with no middlewares."""
|
||||
"""Test FunctionMiddlewarePipeline initialization with no middleware."""
|
||||
pipeline = FunctionMiddlewarePipeline()
|
||||
assert not pipeline.has_middlewares
|
||||
|
||||
@@ -539,7 +539,7 @@ class TestChatMiddlewarePipeline:
|
||||
context.terminate = True
|
||||
|
||||
def test_init_empty(self) -> None:
|
||||
"""Test ChatMiddlewarePipeline initialization with no middlewares."""
|
||||
"""Test ChatMiddlewarePipeline initialization with no middleware."""
|
||||
pipeline = ChatMiddlewarePipeline()
|
||||
assert not pipeline.has_middlewares
|
||||
|
||||
@@ -979,7 +979,7 @@ class TestMultipleMiddlewareOrdering:
|
||||
"""Test cases for multiple middleware execution order."""
|
||||
|
||||
async def test_agent_middleware_execution_order(self, mock_agent: AgentProtocol) -> None:
|
||||
"""Test that multiple agent middlewares execute in registration order."""
|
||||
"""Test that multiple agent middleware execute in registration order."""
|
||||
execution_order: list[str] = []
|
||||
|
||||
class FirstMiddleware(AgentMiddleware):
|
||||
@@ -1006,8 +1006,8 @@ class TestMultipleMiddlewareOrdering:
|
||||
await next(context)
|
||||
execution_order.append("third_after")
|
||||
|
||||
middlewares = [FirstMiddleware(), SecondMiddleware(), ThirdMiddleware()]
|
||||
pipeline = AgentMiddlewarePipeline(middlewares) # type: ignore
|
||||
middleware = [FirstMiddleware(), SecondMiddleware(), ThirdMiddleware()]
|
||||
pipeline = AgentMiddlewarePipeline(middleware) # type: ignore
|
||||
messages = [ChatMessage(role=Role.USER, text="test")]
|
||||
context = AgentRunContext(agent=mock_agent, messages=messages)
|
||||
|
||||
@@ -1030,7 +1030,7 @@ class TestMultipleMiddlewareOrdering:
|
||||
assert execution_order == expected_order
|
||||
|
||||
async def test_function_middleware_execution_order(self, mock_function: AIFunction[Any, Any]) -> None:
|
||||
"""Test that multiple function middlewares execute in registration order."""
|
||||
"""Test that multiple function middleware execute in registration order."""
|
||||
execution_order: list[str] = []
|
||||
|
||||
class FirstMiddleware(FunctionMiddleware):
|
||||
@@ -1053,8 +1053,8 @@ class TestMultipleMiddlewareOrdering:
|
||||
await next(context)
|
||||
execution_order.append("second_after")
|
||||
|
||||
middlewares = [FirstMiddleware(), SecondMiddleware()]
|
||||
pipeline = FunctionMiddlewarePipeline(middlewares) # type: ignore
|
||||
middleware = [FirstMiddleware(), SecondMiddleware()]
|
||||
pipeline = FunctionMiddlewarePipeline(middleware) # type: ignore
|
||||
arguments = FunctionTestArgs(name="test")
|
||||
context = FunctionInvocationContext(function=mock_function, arguments=arguments)
|
||||
|
||||
@@ -1069,7 +1069,7 @@ class TestMultipleMiddlewareOrdering:
|
||||
assert execution_order == expected_order
|
||||
|
||||
async def test_chat_middleware_execution_order(self, mock_chat_client: Any) -> None:
|
||||
"""Test that multiple chat middlewares execute in registration order."""
|
||||
"""Test that multiple chat middleware execute in registration order."""
|
||||
execution_order: list[str] = []
|
||||
|
||||
class FirstChatMiddleware(ChatMiddleware):
|
||||
@@ -1090,8 +1090,8 @@ class TestMultipleMiddlewareOrdering:
|
||||
await next(context)
|
||||
execution_order.append("third_after")
|
||||
|
||||
middlewares = [FirstChatMiddleware(), SecondChatMiddleware(), ThirdChatMiddleware()]
|
||||
pipeline = ChatMiddlewarePipeline(middlewares) # type: ignore
|
||||
middleware = [FirstChatMiddleware(), SecondChatMiddleware(), ThirdChatMiddleware()]
|
||||
pipeline = ChatMiddlewarePipeline(middleware) # type: ignore
|
||||
messages = [ChatMessage(role=Role.USER, text="test")]
|
||||
chat_options = ChatOptions()
|
||||
context = ChatContext(chat_client=mock_chat_client, messages=messages, chat_options=chat_options)
|
||||
@@ -1542,7 +1542,7 @@ class TestMiddlewareExecutionControl:
|
||||
assert context.result is None
|
||||
|
||||
async def test_multiple_middlewares_early_stop(self, mock_agent: AgentProtocol) -> None:
|
||||
"""Test that when first middleware doesn't call next(), subsequent middlewares are not called."""
|
||||
"""Test that when first middleware doesn't call next(), subsequent middleware are not called."""
|
||||
execution_order: list[str] = []
|
||||
|
||||
class FirstMiddleware(AgentMiddleware):
|
||||
@@ -1641,7 +1641,7 @@ class TestMiddlewareExecutionControl:
|
||||
assert context.result is None
|
||||
|
||||
async def test_multiple_chat_middlewares_early_stop(self, mock_chat_client: Any) -> None:
|
||||
"""Test that when first chat middleware doesn't call next(), subsequent middlewares are not called."""
|
||||
"""Test that when first chat middleware doesn't call next(), subsequent middleware are not called."""
|
||||
execution_order: list[str] = []
|
||||
|
||||
class FirstChatMiddleware(ChatMiddleware):
|
||||
|
||||
@@ -418,7 +418,7 @@ class TestChatAgentMultipleMiddlewareOrdering:
|
||||
"""Test cases for multiple middleware execution order with ChatAgent."""
|
||||
|
||||
async def test_multiple_agent_middleware_execution_order(self, chat_client: "MockChatClient") -> None:
|
||||
"""Test that multiple agent middlewares execute in correct order with ChatAgent."""
|
||||
"""Test that multiple agent middleware execute in correct order with ChatAgent."""
|
||||
execution_order: list[str] = []
|
||||
|
||||
class OrderedMiddleware(AgentMiddleware):
|
||||
@@ -432,12 +432,12 @@ class TestChatAgentMultipleMiddlewareOrdering:
|
||||
await next(context)
|
||||
execution_order.append(f"{self.name}_after")
|
||||
|
||||
# Create multiple middlewares
|
||||
# Create multiple middleware
|
||||
middleware1 = OrderedMiddleware("first")
|
||||
middleware2 = OrderedMiddleware("second")
|
||||
middleware3 = OrderedMiddleware("third")
|
||||
|
||||
# Create ChatAgent with multiple middlewares
|
||||
# Create ChatAgent with multiple middleware
|
||||
agent = ChatAgent(chat_client=chat_client, middleware=[middleware1, middleware2, middleware3])
|
||||
|
||||
# Execute the agent
|
||||
@@ -453,7 +453,7 @@ class TestChatAgentMultipleMiddlewareOrdering:
|
||||
assert execution_order == expected_order
|
||||
|
||||
async def test_mixed_middleware_types_with_chat_agent(self, chat_client: "MockChatClient") -> None:
|
||||
"""Test mixed class and function-based middlewares with ChatAgent."""
|
||||
"""Test mixed class and function-based middleware with ChatAgent."""
|
||||
execution_order: list[str] = []
|
||||
|
||||
class ClassAgentMiddleware(AgentMiddleware):
|
||||
@@ -507,8 +507,8 @@ class TestChatAgentMultipleMiddlewareOrdering:
|
||||
assert response is not None
|
||||
assert chat_client.call_count == 1
|
||||
|
||||
# Verify that agent middlewares were executed in correct order
|
||||
# (Function middlewares won't execute since no functions are called)
|
||||
# Verify that agent middleware were executed in correct order
|
||||
# (Function middleware won't execute since no functions are called)
|
||||
expected_order = ["class_agent_before", "function_agent_before", "function_agent_after", "class_agent_after"]
|
||||
assert execution_order == expected_order
|
||||
|
||||
@@ -999,7 +999,7 @@ class TestRunLevelMiddleware:
|
||||
# Clear execution log
|
||||
execution_log.clear()
|
||||
|
||||
# Fourth run with both run middlewares - should see both
|
||||
# Fourth run with both run middleware - should see both
|
||||
await agent.run("Test message 4", middleware=[run_middleware1, run_middleware2])
|
||||
assert execution_log == ["run1_start", "run2_start", "run2_end", "run1_end"]
|
||||
|
||||
|
||||
@@ -342,7 +342,6 @@ def test_agent_decorator_with_valid_class():
|
||||
def __init__(self):
|
||||
self.id = "test_agent_id"
|
||||
self.name = "test_agent"
|
||||
self.display_name = "Test Agent"
|
||||
self.description = "Test agent description"
|
||||
|
||||
async def run(self, messages=None, *, thread=None, **kwargs):
|
||||
@@ -384,7 +383,6 @@ def test_agent_decorator_with_partial_methods():
|
||||
def __init__(self):
|
||||
self.id = "test_agent_id"
|
||||
self.name = "test_agent"
|
||||
self.display_name = "Test Agent"
|
||||
|
||||
async def run(self, messages=None, *, thread=None, **kwargs):
|
||||
return Mock()
|
||||
@@ -406,7 +404,6 @@ def mock_chat_agent():
|
||||
def __init__(self):
|
||||
self.id = "test_agent_id"
|
||||
self.name = "test_agent"
|
||||
self.display_name = "Test Agent"
|
||||
self.description = "Test agent description"
|
||||
self.chat_options = ChatOptions(model_id="TestModel")
|
||||
|
||||
@@ -441,10 +438,10 @@ async def test_agent_instrumentation_enabled(
|
||||
spans = span_exporter.get_finished_spans()
|
||||
assert len(spans) == 1
|
||||
span = spans[0]
|
||||
assert span.name == "invoke_agent Test Agent"
|
||||
assert span.name == "invoke_agent test_agent"
|
||||
assert span.attributes[OtelAttr.OPERATION.value] == OtelAttr.AGENT_INVOKE_OPERATION
|
||||
assert span.attributes[OtelAttr.AGENT_ID] == "test_agent_id"
|
||||
assert span.attributes[OtelAttr.AGENT_NAME] == "Test Agent"
|
||||
assert span.attributes[OtelAttr.AGENT_NAME] == "test_agent"
|
||||
assert span.attributes[OtelAttr.AGENT_DESCRIPTION] == "Test agent description"
|
||||
assert span.attributes[SpanAttributes.LLM_REQUEST_MODEL] == "TestModel"
|
||||
assert span.attributes[OtelAttr.INPUT_TOKENS] == 15
|
||||
@@ -469,10 +466,10 @@ async def test_agent_streaming_response_with_diagnostics_enabled_via_decorator(
|
||||
spans = span_exporter.get_finished_spans()
|
||||
assert len(spans) == 1
|
||||
span = spans[0]
|
||||
assert span.name == "invoke_agent Test Agent"
|
||||
assert span.name == "invoke_agent test_agent"
|
||||
assert span.attributes[OtelAttr.OPERATION.value] == OtelAttr.AGENT_INVOKE_OPERATION
|
||||
assert span.attributes[OtelAttr.AGENT_ID] == "test_agent_id"
|
||||
assert span.attributes[OtelAttr.AGENT_NAME] == "Test Agent"
|
||||
assert span.attributes[OtelAttr.AGENT_NAME] == "test_agent"
|
||||
assert span.attributes[OtelAttr.AGENT_DESCRIPTION] == "Test agent description"
|
||||
assert span.attributes[SpanAttributes.LLM_REQUEST_MODEL] == "TestModel"
|
||||
if enable_sensitive_data:
|
||||
|
||||
@@ -38,7 +38,7 @@ class _CountingAgent(BaseAgent):
|
||||
) -> AgentRunResponse:
|
||||
self.call_count += 1
|
||||
return AgentRunResponse(
|
||||
messages=[ChatMessage(role=Role.ASSISTANT, text=f"Response #{self.call_count}: {self.display_name}")]
|
||||
messages=[ChatMessage(role=Role.ASSISTANT, text=f"Response #{self.call_count}: {self.name}")]
|
||||
)
|
||||
|
||||
async def run_stream( # type: ignore[override]
|
||||
@@ -49,7 +49,7 @@ class _CountingAgent(BaseAgent):
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterable[AgentRunResponseUpdate]:
|
||||
self.call_count += 1
|
||||
yield AgentRunResponseUpdate(contents=[TextContent(text=f"Response #{self.call_count}: {self.display_name}")])
|
||||
yield AgentRunResponseUpdate(contents=[TextContent(text=f"Response #{self.call_count}: {self.name}")])
|
||||
|
||||
|
||||
async def test_agent_executor_checkpoint_stores_and_restores_state() -> None:
|
||||
|
||||
@@ -78,7 +78,7 @@ class _RecordingAgent(BaseAgent):
|
||||
text_handoff: bool = False,
|
||||
extra_properties: dict[str, object] | None = None,
|
||||
) -> None:
|
||||
super().__init__(id=name, name=name, display_name=name)
|
||||
super().__init__(id=name, name=name)
|
||||
self._agent_name = name
|
||||
self.handoff_to = handoff_to
|
||||
self.calls: list[list[ChatMessage]] = []
|
||||
@@ -102,7 +102,7 @@ class _RecordingAgent(BaseAgent):
|
||||
reply = ChatMessage(
|
||||
role=Role.ASSISTANT,
|
||||
contents=contents,
|
||||
author_name=self.display_name,
|
||||
author_name=self.name,
|
||||
additional_properties=additional_properties,
|
||||
)
|
||||
return AgentRunResponse(messages=[reply])
|
||||
|
||||
@@ -35,7 +35,7 @@ class _EchoAgent(BaseAgent):
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AgentRunResponse:
|
||||
return AgentRunResponse(messages=[ChatMessage(role=Role.ASSISTANT, text=f"{self.display_name} reply")])
|
||||
return AgentRunResponse(messages=[ChatMessage(role=Role.ASSISTANT, text=f"{self.name} reply")])
|
||||
|
||||
async def run_stream( # type: ignore[override]
|
||||
self,
|
||||
@@ -45,7 +45,7 @@ class _EchoAgent(BaseAgent):
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterable[AgentRunResponseUpdate]:
|
||||
# Minimal async generator with one assistant update
|
||||
yield AgentRunResponseUpdate(contents=[TextContent(text=f"{self.display_name} reply")])
|
||||
yield AgentRunResponseUpdate(contents=[TextContent(text=f"{self.name} reply")])
|
||||
|
||||
|
||||
class _SummarizerExec(Executor):
|
||||
|
||||
@@ -60,7 +60,7 @@ class _KwargsCapturingAgent(BaseAgent):
|
||||
**kwargs: Any,
|
||||
) -> AgentRunResponse:
|
||||
self.captured_kwargs.append(dict(kwargs))
|
||||
return AgentRunResponse(messages=[ChatMessage(role=Role.ASSISTANT, text=f"{self.display_name} response")])
|
||||
return AgentRunResponse(messages=[ChatMessage(role=Role.ASSISTANT, text=f"{self.name} response")])
|
||||
|
||||
async def run_stream(
|
||||
self,
|
||||
@@ -70,7 +70,7 @@ class _KwargsCapturingAgent(BaseAgent):
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterable[AgentRunResponseUpdate]:
|
||||
self.captured_kwargs.append(dict(kwargs))
|
||||
yield AgentRunResponseUpdate(contents=[TextContent(text=f"{self.display_name} response")])
|
||||
yield AgentRunResponseUpdate(contents=[TextContent(text=f"{self.name} response")])
|
||||
|
||||
|
||||
class _EchoAgent(BaseAgent):
|
||||
@@ -83,7 +83,7 @@ class _EchoAgent(BaseAgent):
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AgentRunResponse:
|
||||
return AgentRunResponse(messages=[ChatMessage(role=Role.ASSISTANT, text=f"{self.display_name} reply")])
|
||||
return AgentRunResponse(messages=[ChatMessage(role=Role.ASSISTANT, text=f"{self.name} reply")])
|
||||
|
||||
async def run_stream(
|
||||
self,
|
||||
@@ -92,7 +92,7 @@ class _EchoAgent(BaseAgent):
|
||||
thread: AgentThread | None = None,
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterable[AgentRunResponseUpdate]:
|
||||
yield AgentRunResponseUpdate(contents=[TextContent(text=f"{self.display_name} reply")])
|
||||
yield AgentRunResponseUpdate(contents=[TextContent(text=f"{self.name} reply")])
|
||||
|
||||
|
||||
# region Sequential Builder Tests
|
||||
|
||||
Reference in New Issue
Block a user