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:
Eduard van Valkenburg
2026-01-13 03:24:07 +01:00
committed by GitHub
Unverified
parent ef44fb4960
commit 203fb7b1c4
80 changed files with 596 additions and 838 deletions
@@ -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"
+7 -48
View File
@@ -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")
+38 -241
View File
@@ -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