Python: [BREAKING] Renamed AgentRunContext to AgentContext (#3714)

* Renamed AgentRunContext to AgentContext

* Update python/packages/core/AGENTS.md

Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>

---------

Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
This commit is contained in:
Dmytro Struk
2026-02-05 22:47:51 -08:00
committed by GitHub
Unverified
parent c609b14f63
commit 09f59b21ad
18 changed files with 219 additions and 319 deletions
@@ -6,7 +6,7 @@ from collections.abc import Awaitable, Callable
from typing import Any
from agent_framework import ChatAgent, ChatMessage, ChatResponse, Content, agent_middleware
from agent_framework._middleware import AgentRunContext
from agent_framework._middleware import AgentContext
from .conftest import MockChatClient
@@ -19,9 +19,7 @@ class TestAsToolKwargsPropagation:
captured_kwargs: dict[str, Any] = {}
@agent_middleware
async def capture_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def capture_middleware(context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
# Capture kwargs passed to the sub-agent
captured_kwargs.update(context.kwargs)
await next(context)
@@ -62,9 +60,7 @@ class TestAsToolKwargsPropagation:
captured_kwargs: dict[str, Any] = {}
@agent_middleware
async def capture_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def capture_middleware(context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
captured_kwargs.update(context.kwargs)
await next(context)
@@ -99,9 +95,7 @@ class TestAsToolKwargsPropagation:
captured_kwargs_list: list[dict[str, Any]] = []
@agent_middleware
async def capture_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def capture_middleware(context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
# Capture kwargs at each level
captured_kwargs_list.append(dict(context.kwargs))
await next(context)
@@ -162,9 +156,7 @@ class TestAsToolKwargsPropagation:
captured_kwargs: dict[str, Any] = {}
@agent_middleware
async def capture_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def capture_middleware(context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
captured_kwargs.update(context.kwargs)
await next(context)
@@ -224,9 +216,7 @@ class TestAsToolKwargsPropagation:
captured_kwargs: dict[str, Any] = {}
@agent_middleware
async def capture_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def capture_middleware(context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
captured_kwargs.update(context.kwargs)
await next(context)
@@ -266,9 +256,7 @@ class TestAsToolKwargsPropagation:
call_count = 0
@agent_middleware
async def capture_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def capture_middleware(context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
nonlocal call_count
call_count += 1
if call_count == 1:
@@ -318,9 +306,7 @@ class TestAsToolKwargsPropagation:
captured_kwargs: dict[str, Any] = {}
@agent_middleware
async def capture_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def capture_middleware(context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
captured_kwargs.update(context.kwargs)
await next(context)
@@ -18,9 +18,9 @@ from agent_framework import (
ResponseStream,
)
from agent_framework._middleware import (
AgentContext,
AgentMiddleware,
AgentMiddlewarePipeline,
AgentRunContext,
ChatContext,
ChatMiddleware,
ChatMiddlewarePipeline,
@@ -32,13 +32,13 @@ from agent_framework._middleware import (
from agent_framework._tools import FunctionTool
class TestAgentRunContext:
"""Test cases for AgentRunContext."""
class TestAgentContext:
"""Test cases for AgentContext."""
def test_init_with_defaults(self, mock_agent: AgentProtocol) -> None:
"""Test AgentRunContext initialization with default values."""
"""Test AgentContext initialization with default values."""
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages)
context = AgentContext(agent=mock_agent, messages=messages)
assert context.agent is mock_agent
assert context.messages == messages
@@ -46,10 +46,10 @@ class TestAgentRunContext:
assert context.metadata == {}
def test_init_with_custom_values(self, mock_agent: AgentProtocol) -> None:
"""Test AgentRunContext initialization with custom values."""
"""Test AgentContext initialization with custom values."""
messages = [ChatMessage(role="user", text="test")]
metadata = {"key": "value"}
context = AgentRunContext(agent=mock_agent, messages=messages, stream=True, metadata=metadata)
context = AgentContext(agent=mock_agent, messages=messages, stream=True, metadata=metadata)
assert context.agent is mock_agent
assert context.messages == messages
@@ -57,12 +57,12 @@ class TestAgentRunContext:
assert context.metadata == metadata
def test_init_with_thread(self, mock_agent: AgentProtocol) -> None:
"""Test AgentRunContext initialization with thread parameter."""
"""Test AgentContext initialization with thread parameter."""
from agent_framework import AgentThread
messages = [ChatMessage(role="user", text="test")]
thread = AgentThread()
context = AgentRunContext(agent=mock_agent, messages=messages, thread=thread)
context = AgentContext(agent=mock_agent, messages=messages, thread=thread)
assert context.agent is mock_agent
assert context.messages == messages
@@ -135,11 +135,11 @@ class TestAgentMiddlewarePipeline:
"""Test cases for AgentMiddlewarePipeline."""
class PreNextTerminateMiddleware(AgentMiddleware):
async def process(self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
raise MiddlewareTermination
class PostNextTerminateMiddleware(AgentMiddleware):
async def process(self, context: AgentRunContext, next: Any) -> None:
async def process(self, context: AgentContext, next: Any) -> None:
await next(context)
raise MiddlewareTermination
@@ -157,7 +157,7 @@ class TestAgentMiddlewarePipeline:
def test_init_with_function_middleware(self) -> None:
"""Test AgentMiddlewarePipeline initialization with function-based middleware."""
async def test_middleware(context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]) -> None:
async def test_middleware(context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
await next(context)
pipeline = AgentMiddlewarePipeline(test_middleware)
@@ -167,11 +167,11 @@ class TestAgentMiddlewarePipeline:
"""Test pipeline execution with no middleware."""
pipeline = AgentMiddlewarePipeline()
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages)
context = AgentContext(agent=mock_agent, messages=messages)
expected_response = AgentResponse(messages=[ChatMessage(role="assistant", text="response")])
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
return expected_response
result = await pipeline.execute(context, final_handler)
@@ -185,9 +185,7 @@ class TestAgentMiddlewarePipeline:
def __init__(self, name: str):
self.name = name
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append(f"{self.name}_before")
await next(context)
execution_order.append(f"{self.name}_after")
@@ -195,11 +193,11 @@ class TestAgentMiddlewarePipeline:
middleware = OrderTrackingMiddleware("test")
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages)
context = AgentContext(agent=mock_agent, messages=messages)
expected_response = AgentResponse(messages=[ChatMessage(role="assistant", text="response")])
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
execution_order.append("handler")
return expected_response
@@ -211,9 +209,9 @@ class TestAgentMiddlewarePipeline:
"""Test pipeline streaming execution with no middleware."""
pipeline = AgentMiddlewarePipeline()
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages, stream=True)
context = AgentContext(agent=mock_agent, messages=messages, stream=True)
async def final_handler(ctx: AgentRunContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def final_handler(ctx: AgentContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def _stream() -> AsyncIterable[AgentResponseUpdate]:
yield AgentResponseUpdate(contents=[Content.from_text(text="chunk1")])
yield AgentResponseUpdate(contents=[Content.from_text(text="chunk2")])
@@ -238,9 +236,7 @@ class TestAgentMiddlewarePipeline:
def __init__(self, name: str):
self.name = name
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append(f"{self.name}_before")
await next(context)
execution_order.append(f"{self.name}_after")
@@ -248,9 +244,9 @@ class TestAgentMiddlewarePipeline:
middleware = StreamOrderTrackingMiddleware("test")
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages, stream=True)
context = AgentContext(agent=mock_agent, messages=messages, stream=True)
async def final_handler(ctx: AgentRunContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def final_handler(ctx: AgentContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def _stream() -> AsyncIterable[AgentResponseUpdate]:
execution_order.append("handler_start")
yield AgentResponseUpdate(contents=[Content.from_text(text="chunk1")])
@@ -274,10 +270,10 @@ class TestAgentMiddlewarePipeline:
middleware = self.PreNextTerminateMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages)
context = AgentContext(agent=mock_agent, messages=messages)
execution_order: list[str] = []
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
# Handler should not be executed when terminated before next()
execution_order.append("handler")
return AgentResponse(messages=[ChatMessage(role="assistant", text="response")])
@@ -292,10 +288,10 @@ class TestAgentMiddlewarePipeline:
middleware = self.PostNextTerminateMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages)
context = AgentContext(agent=mock_agent, messages=messages)
execution_order: list[str] = []
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
execution_order.append("handler")
return AgentResponse(messages=[ChatMessage(role="assistant", text="response")])
@@ -310,10 +306,10 @@ class TestAgentMiddlewarePipeline:
middleware = self.PreNextTerminateMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages, stream=True)
context = AgentContext(agent=mock_agent, messages=messages, stream=True)
execution_order: list[str] = []
async def final_handler(ctx: AgentRunContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def final_handler(ctx: AgentContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def _stream() -> AsyncIterable[AgentResponseUpdate]:
# Handler should not be executed when terminated before next()
execution_order.append("handler_start")
@@ -338,10 +334,10 @@ class TestAgentMiddlewarePipeline:
middleware = self.PostNextTerminateMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages, stream=True)
context = AgentContext(agent=mock_agent, messages=messages, stream=True)
execution_order: list[str] = []
async def final_handler(ctx: AgentRunContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def final_handler(ctx: AgentContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def _stream() -> AsyncIterable[AgentResponseUpdate]:
execution_order.append("handler_start")
yield AgentResponseUpdate(contents=[Content.from_text(text="chunk1")])
@@ -367,9 +363,7 @@ class TestAgentMiddlewarePipeline:
captured_thread = None
class ThreadCapturingMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
nonlocal captured_thread
captured_thread = context.thread
await next(context)
@@ -378,11 +372,11 @@ class TestAgentMiddlewarePipeline:
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
thread = AgentThread()
context = AgentRunContext(agent=mock_agent, messages=messages, thread=thread)
context = AgentContext(agent=mock_agent, messages=messages, thread=thread)
expected_response = AgentResponse(messages=[ChatMessage(role="assistant", text="response")])
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
return expected_response
result = await pipeline.execute(context, final_handler)
@@ -394,9 +388,7 @@ class TestAgentMiddlewarePipeline:
captured_thread = "not_none" # Use string to distinguish from None
class ThreadCapturingMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
nonlocal captured_thread
captured_thread = context.thread
await next(context)
@@ -404,11 +396,11 @@ class TestAgentMiddlewarePipeline:
middleware = ThreadCapturingMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages, thread=None)
context = AgentContext(agent=mock_agent, messages=messages, thread=None)
expected_response = AgentResponse(messages=[ChatMessage(role="assistant", text="response")])
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
return expected_response
result = await pipeline.execute(context, final_handler)
@@ -774,9 +766,7 @@ class TestClassBasedMiddleware:
metadata_updates: list[str] = []
class MetadataAgentMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
context.metadata["before"] = True
metadata_updates.append("before")
await next(context)
@@ -786,9 +776,9 @@ class TestClassBasedMiddleware:
middleware = MetadataAgentMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages)
context = AgentContext(agent=mock_agent, messages=messages)
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
metadata_updates.append("handler")
return AgentResponse(messages=[ChatMessage(role="assistant", text="response")])
@@ -839,9 +829,7 @@ class TestFunctionBasedMiddleware:
"""Test function-based agent middleware."""
execution_order: list[str] = []
async def test_agent_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def test_agent_middleware(context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append("function_before")
context.metadata["function_middleware"] = True
await next(context)
@@ -849,9 +837,9 @@ class TestFunctionBasedMiddleware:
pipeline = AgentMiddlewarePipeline(test_agent_middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages)
context = AgentContext(agent=mock_agent, messages=messages)
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
execution_order.append("handler")
return AgentResponse(messages=[ChatMessage(role="assistant", text="response")])
@@ -896,25 +884,21 @@ class TestMixedMiddleware:
execution_order: list[str] = []
class ClassMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append("class_before")
await next(context)
execution_order.append("class_after")
async def function_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def function_middleware(context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append("function_before")
await next(context)
execution_order.append("function_after")
pipeline = AgentMiddlewarePipeline(ClassMiddleware(), function_middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages)
context = AgentContext(agent=mock_agent, messages=messages)
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
execution_order.append("handler")
return AgentResponse(messages=[ChatMessage(role="assistant", text="response")])
@@ -997,25 +981,19 @@ class TestMultipleMiddlewareOrdering:
execution_order: list[str] = []
class FirstMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append("first_before")
await next(context)
execution_order.append("first_after")
class SecondMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append("second_before")
await next(context)
execution_order.append("second_after")
class ThirdMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append("third_before")
await next(context)
execution_order.append("third_after")
@@ -1023,9 +1001,9 @@ class TestMultipleMiddlewareOrdering:
middleware = [FirstMiddleware(), SecondMiddleware(), ThirdMiddleware()]
pipeline = AgentMiddlewarePipeline(*middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages)
context = AgentContext(agent=mock_agent, messages=messages)
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
execution_order.append("handler")
return AgentResponse(messages=[ChatMessage(role="assistant", text="response")])
@@ -1136,9 +1114,7 @@ class TestContextContentValidation:
"""Test that agent context contains expected data."""
class ContextValidationMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
# Verify context has all expected attributes
assert hasattr(context, "agent")
assert hasattr(context, "messages")
@@ -1161,9 +1137,9 @@ class TestContextContentValidation:
middleware = ContextValidationMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages)
context = AgentContext(agent=mock_agent, messages=messages)
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
# Verify metadata was set by middleware
assert ctx.metadata.get("validated") is True
return AgentResponse(messages=[ChatMessage(role="assistant", text="response")])
@@ -1260,9 +1236,7 @@ class TestStreamingScenarios:
streaming_flags: list[bool] = []
class StreamingFlagMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
streaming_flags.append(context.stream)
await next(context)
@@ -1271,18 +1245,18 @@ class TestStreamingScenarios:
messages = [ChatMessage(role="user", text="test")]
# Test non-streaming
context = AgentRunContext(agent=mock_agent, messages=messages)
context = AgentContext(agent=mock_agent, messages=messages)
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
streaming_flags.append(ctx.stream)
return AgentResponse(messages=[ChatMessage(role="assistant", text="response")])
await pipeline.execute(context, final_handler)
# Test streaming
context_stream = AgentRunContext(agent=mock_agent, messages=messages, stream=True)
context_stream = AgentContext(agent=mock_agent, messages=messages, stream=True)
async def final_stream_handler(ctx: AgentRunContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def final_stream_handler(ctx: AgentContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def _stream() -> AsyncIterable[AgentResponseUpdate]:
streaming_flags.append(ctx.stream)
yield AgentResponseUpdate(contents=[Content.from_text(text="chunk")])
@@ -1302,9 +1276,7 @@ class TestStreamingScenarios:
chunks_processed: list[str] = []
class StreamProcessingMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
chunks_processed.append("before_stream")
await next(context)
chunks_processed.append("after_stream")
@@ -1312,9 +1284,9 @@ class TestStreamingScenarios:
middleware = StreamProcessingMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages, stream=True)
context = AgentContext(agent=mock_agent, messages=messages, stream=True)
async def final_stream_handler(ctx: AgentRunContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def final_stream_handler(ctx: AgentContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def _stream() -> AsyncIterable[AgentResponseUpdate]:
chunks_processed.append("stream_start")
yield AgentResponseUpdate(contents=[Content.from_text(text="chunk1")])
@@ -1436,7 +1408,7 @@ class FunctionTestArgs(BaseModel):
class TestAgentMiddleware(AgentMiddleware):
"""Test implementation of AgentMiddleware."""
async def process(self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
await next(context)
@@ -1469,20 +1441,18 @@ class TestMiddlewareExecutionControl:
"""Test that when agent middleware doesn't call next(), no execution happens."""
class NoNextMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
# Don't call next() - this should prevent any execution
pass
middleware = NoNextMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages)
context = AgentContext(agent=mock_agent, messages=messages)
handler_called = False
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
nonlocal handler_called
handler_called = True
return AgentResponse(messages=[ChatMessage(role="assistant", text="should not execute")])
@@ -1498,20 +1468,18 @@ class TestMiddlewareExecutionControl:
"""Test that when agent middleware doesn't call next(), no streaming execution happens."""
class NoNextStreamingMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
# Don't call next() - this should prevent any execution
pass
middleware = NoNextStreamingMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages, stream=True)
context = AgentContext(agent=mock_agent, messages=messages, stream=True)
handler_called = False
async def final_handler(ctx: AgentRunContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def final_handler(ctx: AgentContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def _stream() -> AsyncIterable[AgentResponseUpdate]:
nonlocal handler_called
handler_called = True
@@ -1566,26 +1534,22 @@ class TestMiddlewareExecutionControl:
execution_order: list[str] = []
class FirstMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append("first")
# Don't call next() - this should stop the pipeline
class SecondMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append("second")
await next(context)
pipeline = AgentMiddlewarePipeline(FirstMiddleware(), SecondMiddleware())
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages)
context = AgentContext(agent=mock_agent, messages=messages)
handler_called = False
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
nonlocal handler_called
handler_called = True
return AgentResponse(messages=[ChatMessage(role="assistant", text="should not execute")])
@@ -17,9 +17,9 @@ from agent_framework import (
ResponseStream,
)
from agent_framework._middleware import (
AgentContext,
AgentMiddleware,
AgentMiddlewarePipeline,
AgentRunContext,
FunctionInvocationContext,
FunctionMiddleware,
FunctionMiddlewarePipeline,
@@ -43,9 +43,7 @@ class TestResultOverrideMiddleware:
override_response = AgentResponse(messages=[ChatMessage(role="assistant", text="overridden response")])
class ResponseOverrideMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
# Execute the pipeline first, then override the response
await next(context)
context.result = override_response
@@ -53,11 +51,11 @@ class TestResultOverrideMiddleware:
middleware = ResponseOverrideMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages)
context = AgentContext(agent=mock_agent, messages=messages)
handler_called = False
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
nonlocal handler_called
handler_called = True
return AgentResponse(messages=[ChatMessage(role="assistant", text="original response")])
@@ -79,9 +77,7 @@ class TestResultOverrideMiddleware:
yield AgentResponseUpdate(contents=[Content.from_text(text=" stream")])
class StreamResponseOverrideMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
# Execute the pipeline first, then override the response stream
await next(context)
context.result = ResponseStream(override_stream())
@@ -89,9 +85,9 @@ class TestResultOverrideMiddleware:
middleware = StreamResponseOverrideMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages, stream=True)
context = AgentContext(agent=mock_agent, messages=messages, stream=True)
async def final_handler(ctx: AgentRunContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def final_handler(ctx: AgentContext) -> ResponseStream[AgentResponseUpdate, AgentResponse]:
async def _stream() -> AsyncIterable[AgentResponseUpdate]:
yield AgentResponseUpdate(contents=[Content.from_text(text="original")])
@@ -145,9 +141,7 @@ class TestResultOverrideMiddleware:
mock_chat_client = MockChatClient()
class ChatAgentResponseOverrideMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
# Always call next() first to allow execution
await next(context)
# Then conditionally override based on content
@@ -184,9 +178,7 @@ class TestResultOverrideMiddleware:
yield AgentResponseUpdate(contents=[Content.from_text(text=" response!")])
class ChatAgentStreamOverrideMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
# Check if we want to override BEFORE calling next to avoid creating unused streams
if any("custom stream" in msg.text for msg in context.messages if msg.text):
context.result = ResponseStream(custom_stream())
@@ -223,9 +215,7 @@ class TestResultOverrideMiddleware:
"""Test that when agent middleware conditionally doesn't call next(), no execution happens."""
class ConditionalNoNextMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
# Only call next() if message contains "execute"
if any("execute" in msg.text for msg in context.messages if msg.text):
await next(context)
@@ -236,14 +226,14 @@ class TestResultOverrideMiddleware:
handler_called = False
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
nonlocal handler_called
handler_called = True
return AgentResponse(messages=[ChatMessage(role="assistant", text="executed response")])
# Test case where next() is NOT called
no_execute_messages = [ChatMessage(role="user", text="Don't run this")]
no_execute_context = AgentRunContext(agent=mock_agent, messages=no_execute_messages, stream=False)
no_execute_context = AgentContext(agent=mock_agent, messages=no_execute_messages, stream=False)
no_execute_result = await pipeline.execute(no_execute_context, final_handler)
# When middleware doesn't call next(), result should be empty AgentResponse
@@ -255,7 +245,7 @@ class TestResultOverrideMiddleware:
# Test case where next() IS called
execute_messages = [ChatMessage(role="user", text="Please execute this")]
execute_context = AgentRunContext(agent=mock_agent, messages=execute_messages, stream=False)
execute_context = AgentContext(agent=mock_agent, messages=execute_messages, stream=False)
execute_result = await pipeline.execute(execute_context, final_handler)
assert execute_result is not None
@@ -318,9 +308,7 @@ class TestResultObservability:
observed_responses: list[AgentResponse] = []
class ObservabilityMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
# Context should be empty before next()
assert context.result is None
@@ -335,9 +323,9 @@ class TestResultObservability:
middleware = ObservabilityMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages, stream=False)
context = AgentContext(agent=mock_agent, messages=messages, stream=False)
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
return AgentResponse(messages=[ChatMessage(role="assistant", text="executed response")])
result = await pipeline.execute(context, final_handler)
@@ -386,9 +374,7 @@ class TestResultObservability:
"""Test that middleware can override response after observing execution."""
class PostExecutionOverrideMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
# Call next to execute first
await next(context)
@@ -405,9 +391,9 @@ class TestResultObservability:
middleware = PostExecutionOverrideMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
messages = [ChatMessage(role="user", text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages, stream=False)
context = AgentContext(agent=mock_agent, messages=messages, stream=False)
async def final_handler(ctx: AgentRunContext) -> AgentResponse:
async def final_handler(ctx: AgentContext) -> AgentResponse:
return AgentResponse(messages=[ChatMessage(role="assistant", text="response to modify")])
result = await pipeline.execute(context, final_handler)
@@ -6,9 +6,9 @@ from typing import Any
import pytest
from agent_framework import (
AgentContext,
AgentMiddleware,
AgentResponseUpdate,
AgentRunContext,
ChatAgent,
ChatClientProtocol,
ChatContext,
@@ -44,9 +44,7 @@ class TestChatAgentClassBasedMiddleware:
def __init__(self, name: str):
self.name = name
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append(f"{self.name}_before")
await next(context)
execution_order.append(f"{self.name}_after")
@@ -122,9 +120,7 @@ class TestChatAgentFunctionBasedMiddleware:
execution_order: list[str] = []
class PreTerminationMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append("middleware_before")
raise MiddlewareTermination
# Code after raise is unreachable
@@ -153,9 +149,7 @@ class TestChatAgentFunctionBasedMiddleware:
execution_order: list[str] = []
class PostTerminationMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append("middleware_before")
await next(context)
execution_order.append("middleware_after")
@@ -225,7 +219,7 @@ class TestChatAgentFunctionBasedMiddleware:
execution_order: list[str] = []
async def tracking_agent_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]
) -> None:
execution_order.append("agent_function_before")
await next(context)
@@ -290,9 +284,7 @@ class TestChatAgentStreamingMiddleware:
streaming_flags: list[bool] = []
class StreamingTrackingMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append("middleware_before")
streaming_flags.append(context.stream)
await next(context)
@@ -334,9 +326,7 @@ class TestChatAgentStreamingMiddleware:
streaming_flags: list[bool] = []
class FlagTrackingMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
streaming_flags.append(context.stream)
await next(context)
@@ -368,9 +358,7 @@ class TestChatAgentMultipleMiddlewareOrdering:
def __init__(self, name: str):
self.name = name
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append(f"{self.name}_before")
await next(context)
execution_order.append(f"{self.name}_after")
@@ -400,15 +388,13 @@ class TestChatAgentMultipleMiddlewareOrdering:
execution_order: list[str] = []
class ClassAgentMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append("class_agent_before")
await next(context)
execution_order.append("class_agent_after")
async def function_agent_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]
) -> None:
execution_order.append("function_agent_before")
await next(context)
@@ -447,15 +433,13 @@ class TestChatAgentMultipleMiddlewareOrdering:
execution_order: list[str] = []
class ClassAgentMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append("class_agent_before")
await next(context)
execution_order.append("class_agent_after")
async def function_agent_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]
) -> None:
execution_order.append("function_agent_before")
await next(context)
@@ -646,8 +630,8 @@ class TestChatAgentFunctionMiddlewareWithTools:
class TrackingAgentMiddleware(AgentMiddleware):
async def process(
self,
context: AgentRunContext,
next: Callable[[AgentRunContext], Awaitable[None]],
context: AgentContext,
next: Callable[[AgentContext], Awaitable[None]],
) -> None:
execution_order.append("agent_middleware_before")
await next(context)
@@ -801,7 +785,7 @@ class TestMiddlewareDynamicRebuild:
self.name = name
self.execution_log = execution_log
async def process(self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
self.execution_log.append(f"{self.name}_start")
await next(context)
self.execution_log.append(f"{self.name}_end")
@@ -924,7 +908,7 @@ class TestRunLevelMiddleware:
self.name = name
self.execution_log = execution_log
async def process(self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
self.execution_log.append(f"{self.name}_start")
await next(context)
self.execution_log.append(f"{self.name}_end")
@@ -976,9 +960,7 @@ class TestRunLevelMiddleware:
def __init__(self, name: str):
self.name = name
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_log.append(f"{self.name}_start")
# Set metadata to pass information to run middleware
context.metadata[f"{self.name}_key"] = f"{self.name}_value"
@@ -989,9 +971,7 @@ class TestRunLevelMiddleware:
def __init__(self, name: str):
self.name = name
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_log.append(f"{self.name}_start")
# Read metadata set by agent middleware
for key, value in context.metadata.items():
@@ -1049,9 +1029,7 @@ class TestRunLevelMiddleware:
def __init__(self, name: str):
self.name = name
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_log.append(f"{self.name}_start")
streaming_flags.append(context.stream)
await next(context)
@@ -1093,9 +1071,7 @@ class TestRunLevelMiddleware:
# Agent-level middleware
class AgentLevelAgentMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_log.append("agent_level_agent_start")
context.metadata["agent_level_agent"] = "processed"
await next(context)
@@ -1114,9 +1090,7 @@ class TestRunLevelMiddleware:
# Run-level middleware
class RunLevelAgentMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_log.append("run_level_agent_start")
# Verify agent-level middleware metadata is available
assert "agent_level_agent" in context.metadata
@@ -1218,7 +1192,7 @@ class TestMiddlewareDecoratorLogic:
@agent_middleware
async def matching_agent_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]
) -> None:
execution_order.append("decorator_type_match_agent")
await next(context)
@@ -1346,7 +1320,7 @@ class TestMiddlewareDecoratorLogic:
execution_order: list[str] = []
# No decorator
async def type_only_agent(context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]) -> None:
async def type_only_agent(context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append("type_only_agent")
await next(context)
@@ -1440,16 +1414,14 @@ class TestMiddlewareDecoratorLogic:
class TestChatAgentThreadBehavior:
"""Test cases for thread behavior in AgentRunContext across multiple runs."""
"""Test cases for thread behavior in AgentContext across multiple runs."""
async def test_agent_run_context_thread_behavior_across_multiple_runs(self, chat_client: "MockChatClient") -> None:
"""Test that AgentRunContext.thread property behaves correctly across multiple agent runs."""
async def test_agent_context_thread_behavior_across_multiple_runs(self, chat_client: "MockChatClient") -> None:
"""Test that AgentContext.thread property behaves correctly across multiple agent runs."""
thread_states: list[dict[str, Any]] = []
class ThreadTrackingMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
# Capture state before next() call
thread_messages = []
if context.thread and context.thread.message_store:
@@ -1804,9 +1776,7 @@ class TestChatAgentChatMiddleware:
"""Test ChatAgent with combined middleware types."""
execution_order: list[str] = []
async def agent_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def agent_middleware(context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
execution_order.append("agent_middleware_before")
await next(context)
execution_order.append("agent_middleware_after")
@@ -1844,9 +1814,7 @@ class TestChatAgentChatMiddleware:
modified_kwargs: dict[str, Any] = {}
@agent_middleware
async def kwargs_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
async def kwargs_middleware(context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]) -> None:
# Capture the original kwargs
captured_kwargs.update(context.kwargs)
@@ -1897,7 +1865,7 @@ class TestChatAgentChatMiddleware:
# class TrackingMiddleware(AgentMiddleware):
# async def process(
# self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
# self, context: AgentContext, next: Callable[[AgentContext], Awaitable[None]]
# ) -> None:
# execution_order.append("before")
# await next(context)