Removed context parameter from call_next (#3829)

This commit is contained in:
Dmytro Struk
2026-02-11 02:47:41 -08:00
committed by GitHub
Unverified
parent 38f22ef006
commit 1fdc4be88d
29 changed files with 451 additions and 583 deletions
+1 -1
View File
@@ -119,7 +119,7 @@ from agent_framework import Agent, AgentMiddleware, AgentContext
class LoggingMiddleware(AgentMiddleware):
async def process(self, context: AgentContext, call_next) -> None:
print(f"Input: {context.messages}")
await call_next(context)
await call_next()
print(f"Output: {context.result}")
agent = Agent(..., middleware=[LoggingMiddleware()])
@@ -145,7 +145,7 @@ class AgentContext:
context.metadata["start_time"] = time.time()
# Continue execution
await call_next(context)
await call_next()
# Access result after execution
print(f"Result: {context.result}")
@@ -229,7 +229,7 @@ class FunctionInvocationContext:
raise MiddlewareTermination("Validation failed")
# Continue execution
await call_next(context)
await call_next()
"""
def __init__(
@@ -293,7 +293,7 @@ class ChatContext:
context.metadata["input_tokens"] = self.count_tokens(context.messages)
# Continue execution
await call_next(context)
await call_next()
# Access result and count output tokens
if context.result:
@@ -365,7 +365,7 @@ class AgentMiddleware(ABC):
async def process(self, context: AgentContext, call_next):
for attempt in range(self.max_retries):
await call_next(context)
await call_next()
if context.result and not context.result.is_error:
break
print(f"Retry {attempt + 1}/{self.max_retries}")
@@ -379,7 +379,7 @@ class AgentMiddleware(ABC):
async def process(
self,
context: AgentContext,
call_next: Callable[[AgentContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
"""Process an agent invocation.
@@ -431,7 +431,7 @@ class FunctionMiddleware(ABC):
raise MiddlewareTermination()
# Execute function
await call_next(context)
await call_next()
# Cache result
if context.result:
@@ -446,7 +446,7 @@ class FunctionMiddleware(ABC):
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
"""Process a function invocation.
@@ -493,7 +493,7 @@ class ChatMiddleware(ABC):
context.messages.insert(0, Message(role="system", text=self.system_prompt))
# Continue execution
await call_next(context)
await call_next()
# Use with an agent
@@ -508,7 +508,7 @@ class ChatMiddleware(ABC):
async def process(
self,
context: ChatContext,
call_next: Callable[[ChatContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
"""Process a chat client request.
@@ -531,15 +531,13 @@ class ChatMiddleware(ABC):
# Pure function type definitions for convenience
AgentMiddlewareCallable = Callable[[AgentContext, Callable[[AgentContext], Awaitable[None]]], Awaitable[None]]
AgentMiddlewareCallable = Callable[[AgentContext, Callable[[], Awaitable[None]]], Awaitable[None]]
AgentMiddlewareTypes: TypeAlias = AgentMiddleware | AgentMiddlewareCallable
FunctionMiddlewareCallable = Callable[
[FunctionInvocationContext, Callable[[FunctionInvocationContext], Awaitable[None]]], Awaitable[None]
]
FunctionMiddlewareCallable = Callable[[FunctionInvocationContext, Callable[[], Awaitable[None]]], Awaitable[None]]
FunctionMiddlewareTypes: TypeAlias = FunctionMiddleware | FunctionMiddlewareCallable
ChatMiddlewareCallable = Callable[[ChatContext, Callable[[ChatContext], Awaitable[None]]], Awaitable[None]]
ChatMiddlewareCallable = Callable[[ChatContext, Callable[[], Awaitable[None]]], Awaitable[None]]
ChatMiddlewareTypes: TypeAlias = ChatMiddleware | ChatMiddlewareCallable
ChatAndFunctionMiddlewareTypes: TypeAlias = (
@@ -578,7 +576,7 @@ def agent_middleware(func: AgentMiddlewareCallable) -> AgentMiddlewareCallable:
@agent_middleware
async def logging_middleware(context: AgentContext, call_next):
print(f"Before: {context.agent.name}")
await call_next(context)
await call_next()
print(f"After: {context.result}")
@@ -611,7 +609,7 @@ def function_middleware(func: FunctionMiddlewareCallable) -> FunctionMiddlewareC
@function_middleware
async def logging_middleware(context: FunctionInvocationContext, call_next):
print(f"Calling: {context.function.name}")
await call_next(context)
await call_next()
print(f"Result: {context.result}")
@@ -644,7 +642,7 @@ def chat_middleware(func: ChatMiddlewareCallable) -> ChatMiddlewareCallable:
@chat_middleware
async def logging_middleware(context: ChatContext, call_next):
print(f"Messages: {len(context.messages)}")
await call_next(context)
await call_next()
print(f"Response: {context.result}")
@@ -666,10 +664,10 @@ class MiddlewareWrapper(Generic[ContextT]):
ContextT: The type of context object this middleware operates on.
"""
def __init__(self, func: Callable[[ContextT, Callable[[ContextT], Awaitable[None]]], Awaitable[None]]) -> None:
def __init__(self, func: Callable[[ContextT, Callable[[], Awaitable[None]]], Awaitable[None]]) -> None:
self.func = func
async def process(self, context: ContextT, call_next: Callable[[ContextT], Awaitable[None]]) -> None:
async def process(self, context: ContextT, call_next: Callable[[], Awaitable[None]]) -> None:
await self.func(context, call_next)
@@ -772,25 +770,25 @@ class AgentMiddlewarePipeline(BaseMiddlewarePipeline):
context.result = await context.result
return context.result
def create_next_handler(index: int) -> Callable[[AgentContext], Awaitable[None]]:
def create_next_handler(index: int) -> Callable[[], Awaitable[None]]:
if index >= len(self._middleware):
async def final_wrapper(c: AgentContext) -> None:
c.result = final_handler(c) # type: ignore[assignment]
if inspect.isawaitable(c.result):
c.result = await c.result
async def final_wrapper() -> None:
context.result = final_handler(context) # type: ignore[assignment]
if inspect.isawaitable(context.result):
context.result = await context.result
return final_wrapper
async def current_handler(c: AgentContext) -> None:
async def current_handler() -> None:
# MiddlewareTermination bubbles up to execute() to skip post-processing
await self._middleware[index].process(c, create_next_handler(index + 1))
await self._middleware[index].process(context, create_next_handler(index + 1))
return current_handler
first_handler = create_next_handler(0)
with contextlib.suppress(MiddlewareTermination):
await first_handler(context)
await first_handler()
if context.result and isinstance(context.result, ResponseStream):
for hook in context.stream_transform_hooks:
@@ -847,25 +845,25 @@ class FunctionMiddlewarePipeline(BaseMiddlewarePipeline):
if not self._middleware:
return await final_handler(context)
def create_next_handler(index: int) -> Callable[[FunctionInvocationContext], Awaitable[None]]:
def create_next_handler(index: int) -> Callable[[], Awaitable[None]]:
if index >= len(self._middleware):
async def final_wrapper(c: FunctionInvocationContext) -> None:
c.result = final_handler(c)
if inspect.isawaitable(c.result):
c.result = await c.result
async def final_wrapper() -> None:
context.result = final_handler(context)
if inspect.isawaitable(context.result):
context.result = await context.result
return final_wrapper
async def current_handler(c: FunctionInvocationContext) -> None:
async def current_handler() -> None:
# MiddlewareTermination bubbles up to execute() to skip post-processing
await self._middleware[index].process(c, create_next_handler(index + 1))
await self._middleware[index].process(context, create_next_handler(index + 1))
return current_handler
first_handler = create_next_handler(0)
# Don't suppress MiddlewareTermination - let it propagate to signal loop termination
await first_handler(context)
await first_handler()
return context.result
@@ -922,25 +920,25 @@ class ChatMiddlewarePipeline(BaseMiddlewarePipeline):
raise ValueError("Streaming agent middleware requires a ResponseStream result.")
return context.result
def create_next_handler(index: int) -> Callable[[ChatContext], Awaitable[None]]:
def create_next_handler(index: int) -> Callable[[], Awaitable[None]]:
if index >= len(self._middleware):
async def final_wrapper(c: ChatContext) -> None:
c.result = final_handler(c) # type: ignore[assignment]
if inspect.isawaitable(c.result):
c.result = await c.result
async def final_wrapper() -> None:
context.result = final_handler(context) # type: ignore[assignment]
if inspect.isawaitable(context.result):
context.result = await context.result
return final_wrapper
async def current_handler(c: ChatContext) -> None:
async def current_handler() -> None:
# MiddlewareTermination bubbles up to execute() to skip post-processing
await self._middleware[index].process(c, create_next_handler(index + 1))
await self._middleware[index].process(context, create_next_handler(index + 1))
return current_handler
first_handler = create_next_handler(0)
with contextlib.suppress(MiddlewareTermination):
await first_handler(context)
await first_handler()
if context.result and isinstance(context.result, ResponseStream):
for hook in context.stream_transform_hooks:
@@ -19,12 +19,10 @@ class TestAsToolKwargsPropagation:
captured_kwargs: dict[str, Any] = {}
@agent_middleware
async def capture_middleware(
context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def capture_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Capture kwargs passed to the sub-agent
captured_kwargs.update(context.kwargs)
await call_next(context)
await call_next()
# Setup mock response
client.responses = [
@@ -62,11 +60,9 @@ class TestAsToolKwargsPropagation:
captured_kwargs: dict[str, Any] = {}
@agent_middleware
async def capture_middleware(
context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def capture_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
captured_kwargs.update(context.kwargs)
await call_next(context)
await call_next()
# Setup mock response
client.responses = [
@@ -99,12 +95,10 @@ class TestAsToolKwargsPropagation:
captured_kwargs_list: list[dict[str, Any]] = []
@agent_middleware
async def capture_middleware(
context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def capture_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Capture kwargs at each level
captured_kwargs_list.append(dict(context.kwargs))
await call_next(context)
await call_next()
# Setup mock responses to trigger nested tool invocation: B calls tool C, then completes.
client.responses = [
@@ -162,11 +156,9 @@ class TestAsToolKwargsPropagation:
captured_kwargs: dict[str, Any] = {}
@agent_middleware
async def capture_middleware(
context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def capture_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
captured_kwargs.update(context.kwargs)
await call_next(context)
await call_next()
# Setup mock streaming responses
from agent_framework import ChatResponseUpdate
@@ -224,11 +216,9 @@ class TestAsToolKwargsPropagation:
captured_kwargs: dict[str, Any] = {}
@agent_middleware
async def capture_middleware(
context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def capture_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
captured_kwargs.update(context.kwargs)
await call_next(context)
await call_next()
# Setup mock response
client.responses = [
@@ -266,16 +256,14 @@ class TestAsToolKwargsPropagation:
call_count = 0
@agent_middleware
async def capture_middleware(
context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def capture_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
nonlocal call_count
call_count += 1
if call_count == 1:
first_call_kwargs.update(context.kwargs)
elif call_count == 2:
second_call_kwargs.update(context.kwargs)
await call_next(context)
await call_next()
# Setup mock responses for both calls
client.responses = [
@@ -318,11 +306,9 @@ class TestAsToolKwargsPropagation:
captured_kwargs: dict[str, Any] = {}
@agent_middleware
async def capture_middleware(
context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def capture_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
captured_kwargs.update(context.kwargs)
await call_next(context)
await call_next()
# Setup mock response
client.responses = [
@@ -2298,9 +2298,7 @@ async def test_streaming_error_recovery_resets_counter(chat_client_base: Support
class TerminateLoopMiddleware(FunctionMiddleware):
"""Middleware that raises MiddlewareTermination to exit the function calling loop."""
async def process(
self, context: FunctionInvocationContext, next_handler: Callable[[FunctionInvocationContext], Awaitable[None]]
) -> None:
async def process(self, context: FunctionInvocationContext, next_handler: Callable[[], Awaitable[None]]) -> None:
# Set result to a simple value - the framework will wrap it in FunctionResultContent
context.result = "terminated by middleware"
raise MiddlewareTermination
@@ -2355,14 +2353,12 @@ async def test_terminate_loop_single_function_call(chat_client_base: SupportsCha
class SelectiveTerminateMiddleware(FunctionMiddleware):
"""Only terminates for terminating_function."""
async def process(
self, context: FunctionInvocationContext, next_handler: Callable[[FunctionInvocationContext], Awaitable[None]]
) -> None:
async def process(self, context: FunctionInvocationContext, next_handler: Callable[[], Awaitable[None]]) -> None:
if context.function.name == "terminating_function":
# Set result to a simple value - the framework will wrap it in FunctionResultContent
context.result = "terminated by middleware"
raise MiddlewareTermination
await next_handler(context)
await next_handler()
async def test_terminate_loop_multiple_function_calls_one_terminates(chat_client_base: SupportsChatGetResponse):
@@ -135,12 +135,12 @@ class TestAgentMiddlewarePipeline:
"""Test cases for AgentMiddlewarePipeline."""
class PreNextTerminateMiddleware(AgentMiddleware):
async def process(self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
raise MiddlewareTermination
class PostNextTerminateMiddleware(AgentMiddleware):
async def process(self, context: AgentContext, call_next: Any) -> None:
await call_next(context)
await call_next()
raise MiddlewareTermination
def test_init_empty(self) -> None:
@@ -157,8 +157,8 @@ class TestAgentMiddlewarePipeline:
def test_init_with_function_middleware(self) -> None:
"""Test AgentMiddlewarePipeline initialization with function-based middleware."""
async def test_middleware(context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]) -> None:
await call_next(context)
async def test_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
await call_next()
pipeline = AgentMiddlewarePipeline(test_middleware)
assert pipeline.has_middlewares
@@ -185,11 +185,9 @@ class TestAgentMiddlewarePipeline:
def __init__(self, name: str):
self.name = name
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append(f"{self.name}_before")
await call_next(context)
await call_next()
execution_order.append(f"{self.name}_after")
middleware = OrderTrackingMiddleware("test")
@@ -238,11 +236,9 @@ class TestAgentMiddlewarePipeline:
def __init__(self, name: str):
self.name = name
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append(f"{self.name}_before")
await call_next(context)
await call_next()
execution_order.append(f"{self.name}_after")
middleware = StreamOrderTrackingMiddleware("test")
@@ -367,12 +363,10 @@ class TestAgentMiddlewarePipeline:
captured_thread = None
class ThreadCapturingMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
nonlocal captured_thread
captured_thread = context.thread
await call_next(context)
await call_next()
middleware = ThreadCapturingMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
@@ -394,12 +388,10 @@ class TestAgentMiddlewarePipeline:
captured_thread = "not_none" # Use string to distinguish from None
class ThreadCapturingMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
nonlocal captured_thread
captured_thread = context.thread
await call_next(context)
await call_next()
middleware = ThreadCapturingMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
@@ -425,7 +417,7 @@ class TestFunctionMiddlewarePipeline:
class PostNextTerminateFunctionMiddleware(FunctionMiddleware):
async def process(self, context: FunctionInvocationContext, call_next: Any) -> None:
await call_next(context)
await call_next()
raise MiddlewareTermination
async def test_execute_with_pre_next_termination(self, mock_function: FunctionTool[Any, Any]) -> None:
@@ -482,10 +474,8 @@ class TestFunctionMiddlewarePipeline:
def test_init_with_function_middleware(self) -> None:
"""Test FunctionMiddlewarePipeline initialization with function-based middleware."""
async def test_middleware(
context: FunctionInvocationContext, call_next: Callable[[FunctionInvocationContext], Awaitable[None]]
) -> None:
await call_next(context)
async def test_middleware(context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]) -> None:
await call_next()
pipeline = FunctionMiddlewarePipeline(test_middleware)
assert pipeline.has_middlewares
@@ -515,10 +505,10 @@ class TestFunctionMiddlewarePipeline:
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
execution_order.append(f"{self.name}_before")
await call_next(context)
await call_next()
execution_order.append(f"{self.name}_after")
middleware = OrderTrackingFunctionMiddleware("test")
@@ -541,12 +531,12 @@ class TestChatMiddlewarePipeline:
"""Test cases for ChatMiddlewarePipeline."""
class PreNextTerminateChatMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
raise MiddlewareTermination
class PostNextTerminateChatMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
await call_next(context)
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
await call_next()
raise MiddlewareTermination
def test_init_empty(self) -> None:
@@ -563,8 +553,8 @@ class TestChatMiddlewarePipeline:
def test_init_with_function_middleware(self) -> None:
"""Test ChatMiddlewarePipeline initialization with function-based middleware."""
async def test_middleware(context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
await call_next(context)
async def test_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
await call_next()
pipeline = ChatMiddlewarePipeline(test_middleware)
assert pipeline.has_middlewares
@@ -592,9 +582,9 @@ class TestChatMiddlewarePipeline:
def __init__(self, name: str):
self.name = name
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append(f"{self.name}_before")
await call_next(context)
await call_next()
execution_order.append(f"{self.name}_after")
middleware = OrderTrackingChatMiddleware("test")
@@ -644,9 +634,9 @@ class TestChatMiddlewarePipeline:
def __init__(self, name: str):
self.name = name
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append(f"{self.name}_before")
await call_next(context)
await call_next()
execution_order.append(f"{self.name}_after")
middleware = StreamOrderTrackingChatMiddleware("test")
@@ -774,12 +764,10 @@ class TestClassBasedMiddleware:
metadata_updates: list[str] = []
class MetadataAgentMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
context.metadata["before"] = True
metadata_updates.append("before")
await call_next(context)
await call_next()
context.metadata["after"] = True
metadata_updates.append("after")
@@ -807,11 +795,11 @@ class TestClassBasedMiddleware:
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
context.metadata["before"] = True
metadata_updates.append("before")
await call_next(context)
await call_next()
context.metadata["after"] = True
metadata_updates.append("after")
@@ -839,12 +827,10 @@ class TestFunctionBasedMiddleware:
"""Test function-based agent middleware."""
execution_order: list[str] = []
async def test_agent_middleware(
context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def test_agent_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("function_before")
context.metadata["function_middleware"] = True
await call_next(context)
await call_next()
execution_order.append("function_after")
pipeline = AgentMiddlewarePipeline(test_agent_middleware)
@@ -866,11 +852,11 @@ class TestFunctionBasedMiddleware:
execution_order: list[str] = []
async def test_function_middleware(
context: FunctionInvocationContext, call_next: Callable[[FunctionInvocationContext], Awaitable[None]]
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
) -> None:
execution_order.append("function_before")
context.metadata["function_middleware"] = True
await call_next(context)
await call_next()
execution_order.append("function_after")
pipeline = FunctionMiddlewarePipeline(test_function_middleware)
@@ -896,18 +882,14 @@ class TestMixedMiddleware:
execution_order: list[str] = []
class ClassMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("class_before")
await call_next(context)
await call_next()
execution_order.append("class_after")
async def function_middleware(
context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def function_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("function_before")
await call_next(context)
await call_next()
execution_order.append("function_after")
pipeline = AgentMiddlewarePipeline(ClassMiddleware(), function_middleware)
@@ -931,17 +913,17 @@ class TestMixedMiddleware:
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
execution_order.append("class_before")
await call_next(context)
await call_next()
execution_order.append("class_after")
async def function_middleware(
context: FunctionInvocationContext, call_next: Callable[[FunctionInvocationContext], Awaitable[None]]
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
) -> None:
execution_order.append("function_before")
await call_next(context)
await call_next()
execution_order.append("function_after")
pipeline = FunctionMiddlewarePipeline(ClassMiddleware(), function_middleware)
@@ -962,16 +944,14 @@ class TestMixedMiddleware:
execution_order: list[str] = []
class ClassChatMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("class_before")
await call_next(context)
await call_next()
execution_order.append("class_after")
async def function_chat_middleware(
context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]
) -> None:
async def function_chat_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("function_before")
await call_next(context)
await call_next()
execution_order.append("function_after")
pipeline = ChatMiddlewarePipeline(ClassChatMiddleware(), function_chat_middleware)
@@ -997,27 +977,21 @@ class TestMultipleMiddlewareOrdering:
execution_order: list[str] = []
class FirstMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("first_before")
await call_next(context)
await call_next()
execution_order.append("first_after")
class SecondMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("second_before")
await call_next(context)
await call_next()
execution_order.append("second_after")
class ThirdMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("third_before")
await call_next(context)
await call_next()
execution_order.append("third_after")
middleware = [FirstMiddleware(), SecondMiddleware(), ThirdMiddleware()]
@@ -1051,20 +1025,20 @@ class TestMultipleMiddlewareOrdering:
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
execution_order.append("first_before")
await call_next(context)
await call_next()
execution_order.append("first_after")
class SecondMiddleware(FunctionMiddleware):
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
execution_order.append("second_before")
await call_next(context)
await call_next()
execution_order.append("second_after")
middleware = [FirstMiddleware(), SecondMiddleware()]
@@ -1087,21 +1061,21 @@ class TestMultipleMiddlewareOrdering:
execution_order: list[str] = []
class FirstChatMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("first_before")
await call_next(context)
await call_next()
execution_order.append("first_after")
class SecondChatMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("second_before")
await call_next(context)
await call_next()
execution_order.append("second_after")
class ThirdChatMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("third_before")
await call_next(context)
await call_next()
execution_order.append("third_after")
middleware = [FirstChatMiddleware(), SecondChatMiddleware(), ThirdChatMiddleware()]
@@ -1136,9 +1110,7 @@ class TestContextContentValidation:
"""Test that agent context contains expected data."""
class ContextValidationMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Verify context has all expected attributes
assert hasattr(context, "agent")
assert hasattr(context, "messages")
@@ -1156,7 +1128,7 @@ class TestContextContentValidation:
# Add custom metadata
context.metadata["validated"] = True
await call_next(context)
await call_next()
middleware = ContextValidationMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
@@ -1178,7 +1150,7 @@ class TestContextContentValidation:
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
# Verify context has all expected attributes
assert hasattr(context, "function")
@@ -1194,7 +1166,7 @@ class TestContextContentValidation:
# Add custom metadata
context.metadata["validated"] = True
await call_next(context)
await call_next()
middleware = ContextValidationMiddleware()
pipeline = FunctionMiddlewarePipeline(middleware)
@@ -1213,7 +1185,7 @@ class TestContextContentValidation:
"""Test that chat context contains expected data."""
class ChatContextValidationMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Verify context has all expected attributes
assert hasattr(context, "client")
assert hasattr(context, "messages")
@@ -1235,7 +1207,7 @@ class TestContextContentValidation:
# Add custom metadata
context.metadata["validated"] = True
await call_next(context)
await call_next()
middleware = ChatContextValidationMiddleware()
pipeline = ChatMiddlewarePipeline(middleware)
@@ -1260,11 +1232,9 @@ class TestStreamingScenarios:
streaming_flags: list[bool] = []
class StreamingFlagMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
streaming_flags.append(context.stream)
await call_next(context)
await call_next()
middleware = StreamingFlagMiddleware()
pipeline = AgentMiddlewarePipeline(middleware)
@@ -1302,11 +1272,9 @@ class TestStreamingScenarios:
chunks_processed: list[str] = []
class StreamProcessingMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
chunks_processed.append("before_stream")
await call_next(context)
await call_next()
chunks_processed.append("after_stream")
middleware = StreamProcessingMiddleware()
@@ -1345,9 +1313,9 @@ class TestStreamingScenarios:
streaming_flags: list[bool] = []
class ChatStreamingFlagMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
streaming_flags.append(context.stream)
await call_next(context)
await call_next()
middleware = ChatStreamingFlagMiddleware()
pipeline = ChatMiddlewarePipeline(middleware)
@@ -1386,9 +1354,9 @@ class TestStreamingScenarios:
chunks_processed: list[str] = []
class ChatStreamProcessingMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
chunks_processed.append("before_stream")
await call_next(context)
await call_next()
chunks_processed.append("after_stream")
middleware = ChatStreamProcessingMiddleware()
@@ -1436,24 +1404,22 @@ class FunctionTestArgs(BaseModel):
class TestAgentMiddleware(AgentMiddleware):
"""Test implementation of AgentMiddleware."""
async def process(self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]) -> None:
await call_next(context)
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
await call_next()
class TestFunctionMiddleware(FunctionMiddleware):
"""Test implementation of FunctionMiddleware."""
async def process(
self, context: FunctionInvocationContext, call_next: Callable[[FunctionInvocationContext], Awaitable[None]]
) -> None:
await call_next(context)
async def process(self, context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]) -> None:
await call_next()
class TestChatMiddleware(ChatMiddleware):
"""Test implementation of ChatMiddleware."""
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
await call_next(context)
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
await call_next()
class MockFunctionArgs(BaseModel):
@@ -1469,9 +1435,7 @@ class TestMiddlewareExecutionControl:
"""Test that when agent middleware doesn't call next(), no execution happens."""
class NoNextMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Don't call next() - this should prevent any execution
pass
@@ -1498,9 +1462,7 @@ class TestMiddlewareExecutionControl:
"""Test that when agent middleware doesn't call next(), no streaming execution happens."""
class NoNextStreamingMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Don't call next() - this should prevent any execution
pass
@@ -1537,7 +1499,7 @@ class TestMiddlewareExecutionControl:
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
# Don't call next() - this should prevent any execution
pass
@@ -1566,18 +1528,14 @@ class TestMiddlewareExecutionControl:
execution_order: list[str] = []
class FirstMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("first")
# Don't call next() - this should stop the pipeline
class SecondMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("second")
await call_next(context)
await call_next()
pipeline = AgentMiddlewarePipeline(FirstMiddleware(), SecondMiddleware())
messages = [Message(role="user", text="test")]
@@ -1601,7 +1559,7 @@ class TestMiddlewareExecutionControl:
"""Test that when chat middleware doesn't call next(), no execution happens."""
class NoNextChatMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Don't call next() - this should prevent any execution
pass
@@ -1629,7 +1587,7 @@ class TestMiddlewareExecutionControl:
"""Test that when chat middleware doesn't call next(), no streaming execution happens."""
class NoNextStreamingChatMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Don't call next() - this should prevent any execution
pass
@@ -1670,14 +1628,14 @@ class TestMiddlewareExecutionControl:
execution_order: list[str] = []
class FirstChatMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("first")
# Don't call next() - this should stop the pipeline
class SecondChatMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("second")
await call_next(context)
await call_next()
pipeline = ChatMiddlewarePipeline(FirstChatMiddleware(), SecondChatMiddleware())
messages = [Message(role="user", text="test")]
@@ -43,11 +43,9 @@ class TestResultOverrideMiddleware:
override_response = AgentResponse(messages=[Message(role="assistant", text="overridden response")])
class ResponseOverrideMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Execute the pipeline first, then override the response
await call_next(context)
await call_next()
context.result = override_response
middleware = ResponseOverrideMiddleware()
@@ -79,11 +77,9 @@ class TestResultOverrideMiddleware:
yield AgentResponseUpdate(contents=[Content.from_text(text=" stream")])
class StreamResponseOverrideMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Execute the pipeline first, then override the response stream
await call_next(context)
await call_next()
context.result = ResponseStream(override_stream())
middleware = StreamResponseOverrideMiddleware()
@@ -115,10 +111,10 @@ class TestResultOverrideMiddleware:
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
# Execute the pipeline first, then override the result
await call_next(context)
await call_next()
context.result = override_result
middleware = ResultOverrideMiddleware()
@@ -145,11 +141,9 @@ class TestResultOverrideMiddleware:
mock_chat_client = MockChatClient()
class ChatAgentResponseOverrideMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Always call next() first to allow execution
await call_next(context)
await call_next()
# Then conditionally override based on content
if any("special" in msg.text for msg in context.messages if msg.text):
context.result = AgentResponse(
@@ -184,15 +178,13 @@ class TestResultOverrideMiddleware:
yield AgentResponseUpdate(contents=[Content.from_text(text=" response!")])
class ChatAgentStreamOverrideMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], 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())
return # Don't call next() - we're overriding the entire result
# Normal case - let the agent handle it
await call_next(context)
await call_next()
# Create Agent with override middleware
middleware = ChatAgentStreamOverrideMiddleware()
@@ -223,12 +215,10 @@ class TestResultOverrideMiddleware:
"""Test that when agent middleware conditionally doesn't call next(), no execution happens."""
class ConditionalNoNextMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], 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 call_next(context)
await call_next()
# Otherwise, don't call next() - no execution should happen
middleware = ConditionalNoNextMiddleware()
@@ -269,13 +259,13 @@ class TestResultOverrideMiddleware:
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
# Only call next() if argument name contains "execute"
args = context.arguments
assert isinstance(args, FunctionTestArgs)
if "execute" in args.name:
await call_next(context)
await call_next()
# Otherwise, don't call next() - no execution should happen
middleware = ConditionalNoNextFunctionMiddleware()
@@ -318,14 +308,12 @@ class TestResultObservability:
observed_responses: list[AgentResponse] = []
class ObservabilityMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Context should be empty before next()
assert context.result is None
# Call next to execute
await call_next(context)
await call_next()
# Context should now contain the response for observability
assert context.result is not None
@@ -355,13 +343,13 @@ class TestResultObservability:
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
# Context should be empty before next()
assert context.result is None
# Call next to execute
await call_next(context)
await call_next()
# Context should now contain the result for observability
assert context.result is not None
@@ -386,11 +374,9 @@ class TestResultObservability:
"""Test that middleware can override response after observing execution."""
class PostExecutionOverrideMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Call next to execute first
await call_next(context)
await call_next()
# Now observe and conditionally override
assert context.result is not None
@@ -423,10 +409,10 @@ class TestResultObservability:
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
# Call next to execute first
await call_next(context)
await call_next()
# Now observe and conditionally override
assert context.result is not None
@@ -44,11 +44,9 @@ class TestChatAgentClassBasedMiddleware:
def __init__(self, name: str):
self.name = name
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append(f"{self.name}_before")
await call_next(context)
await call_next()
execution_order.append(f"{self.name}_after")
# Create Agent with middleware
@@ -76,9 +74,9 @@ class TestChatAgentClassBasedMiddleware:
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
await call_next(context)
await call_next()
middleware = TrackingFunctionMiddleware()
Agent(client=client, middleware=[middleware])
@@ -96,10 +94,10 @@ class TestChatAgentClassBasedMiddleware:
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
execution_order.append(f"{self.name}_before")
await call_next(context)
await call_next()
execution_order.append(f"{self.name}_after")
middleware = TrackingFunctionMiddleware("function_middleware")
@@ -122,13 +120,11 @@ class TestChatAgentFunctionBasedMiddleware:
execution_order: list[str] = []
class PreTerminationMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("middleware_before")
raise MiddlewareTermination
# Code after raise is unreachable
await call_next(context)
await call_next()
execution_order.append("middleware_after")
# Create Agent with terminating middleware
@@ -153,11 +149,9 @@ class TestChatAgentFunctionBasedMiddleware:
execution_order: list[str] = []
class PostTerminationMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("middleware_before")
await call_next(context)
await call_next()
execution_order.append("middleware_after")
context.terminate = True
@@ -193,12 +187,12 @@ class TestChatAgentFunctionBasedMiddleware:
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
execution_order.append("middleware_before")
context.terminate = True
# We call next() but since terminate=True, subsequent middleware and handler should not execute
await call_next(context)
await call_next()
execution_order.append("middleware_after")
Agent(client=client, middleware=[PreTerminationFunctionMiddleware()], tools=[])
@@ -211,10 +205,10 @@ class TestChatAgentFunctionBasedMiddleware:
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
execution_order.append("middleware_before")
await call_next(context)
await call_next()
execution_order.append("middleware_after")
context.terminate = True
@@ -224,11 +218,9 @@ class TestChatAgentFunctionBasedMiddleware:
"""Test function-based agent middleware with Agent."""
execution_order: list[str] = []
async def tracking_agent_middleware(
context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def tracking_agent_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("agent_function_before")
await call_next(context)
await call_next()
execution_order.append("agent_function_after")
# Create Agent with function middleware
@@ -252,9 +244,9 @@ class TestChatAgentFunctionBasedMiddleware:
"""Test function-based function middleware with Agent."""
async def tracking_function_middleware(
context: FunctionInvocationContext, call_next: Callable[[FunctionInvocationContext], Awaitable[None]]
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
) -> None:
await call_next(context)
await call_next()
Agent(client=client, middleware=[tracking_function_middleware])
@@ -265,10 +257,10 @@ class TestChatAgentFunctionBasedMiddleware:
execution_order: list[str] = []
async def tracking_function_middleware(
context: FunctionInvocationContext, call_next: Callable[[FunctionInvocationContext], Awaitable[None]]
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
) -> None:
execution_order.append("function_function_before")
await call_next(context)
await call_next()
execution_order.append("function_function_after")
agent = Agent(client=chat_client_base, middleware=[tracking_function_middleware])
@@ -290,12 +282,10 @@ class TestChatAgentStreamingMiddleware:
streaming_flags: list[bool] = []
class StreamingTrackingMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("middleware_before")
streaming_flags.append(context.stream)
await call_next(context)
await call_next()
execution_order.append("middleware_after")
# Create Agent with middleware
@@ -334,11 +324,9 @@ class TestChatAgentStreamingMiddleware:
streaming_flags: list[bool] = []
class FlagTrackingMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
streaming_flags.append(context.stream)
await call_next(context)
await call_next()
# Create Agent with middleware
middleware = FlagTrackingMiddleware()
@@ -368,11 +356,9 @@ class TestChatAgentMultipleMiddlewareOrdering:
def __init__(self, name: str):
self.name = name
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append(f"{self.name}_before")
await call_next(context)
await call_next()
execution_order.append(f"{self.name}_after")
# Create multiple middleware
@@ -400,35 +386,31 @@ class TestChatAgentMultipleMiddlewareOrdering:
execution_order: list[str] = []
class ClassAgentMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("class_agent_before")
await call_next(context)
await call_next()
execution_order.append("class_agent_after")
async def function_agent_middleware(
context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def function_agent_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("function_agent_before")
await call_next(context)
await call_next()
execution_order.append("function_agent_after")
class ClassFunctionMiddleware(FunctionMiddleware):
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
execution_order.append("class_function_before")
await call_next(context)
await call_next()
execution_order.append("class_function_after")
async def function_function_middleware(
context: FunctionInvocationContext, call_next: Callable[[FunctionInvocationContext], Awaitable[None]]
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
) -> None:
execution_order.append("function_function_before")
await call_next(context)
await call_next()
execution_order.append("function_function_after")
agent = Agent(
@@ -447,25 +429,21 @@ class TestChatAgentMultipleMiddlewareOrdering:
execution_order: list[str] = []
class ClassAgentMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("class_agent_before")
await call_next(context)
await call_next()
execution_order.append("class_agent_after")
async def function_agent_middleware(
context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def function_agent_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("function_agent_before")
await call_next(context)
await call_next()
execution_order.append("function_agent_after")
async def function_function_middleware(
context: FunctionInvocationContext, call_next: Callable[[FunctionInvocationContext], Awaitable[None]]
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
) -> None:
execution_order.append("function_function_before")
await call_next(context)
await call_next()
execution_order.append("function_function_after")
agent = Agent(
@@ -521,10 +499,10 @@ class TestChatAgentFunctionMiddlewareWithTools:
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
execution_order.append(f"{self.name}_before")
await call_next(context)
await call_next()
execution_order.append(f"{self.name}_after")
# Set up mock to return a function call first, then a regular response
@@ -583,10 +561,10 @@ class TestChatAgentFunctionMiddlewareWithTools:
execution_order: list[str] = []
async def tracking_function_middleware(
context: FunctionInvocationContext, call_next: Callable[[FunctionInvocationContext], Awaitable[None]]
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
) -> None:
execution_order.append("function_middleware_before")
await call_next(context)
await call_next()
execution_order.append("function_middleware_after")
# Set up mock to return a function call first, then a regular response
@@ -647,20 +625,20 @@ class TestChatAgentFunctionMiddlewareWithTools:
async def process(
self,
context: AgentContext,
call_next: Callable[[AgentContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
execution_order.append("agent_middleware_before")
await call_next(context)
await call_next()
execution_order.append("agent_middleware_after")
class TrackingFunctionMiddleware(FunctionMiddleware):
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
execution_order.append("function_middleware_before")
await call_next(context)
await call_next()
execution_order.append("function_middleware_after")
# Set up mock to return a function call first, then a regular response
@@ -728,7 +706,7 @@ class TestChatAgentFunctionMiddlewareWithTools:
@function_middleware
async def kwargs_middleware(
context: FunctionInvocationContext, call_next: Callable[[FunctionInvocationContext], Awaitable[None]]
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
) -> None:
nonlocal middleware_called
middleware_called = True
@@ -748,7 +726,7 @@ class TestChatAgentFunctionMiddlewareWithTools:
modified_kwargs["new_param"] = context.kwargs.get("new_param")
modified_kwargs["custom_param"] = context.kwargs.get("custom_param")
await call_next(context)
await call_next()
chat_client_base.run_responses = [
ChatResponse(
@@ -801,9 +779,9 @@ class TestMiddlewareDynamicRebuild:
self.name = name
self.execution_log = execution_log
async def process(self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
self.execution_log.append(f"{self.name}_start")
await call_next(context)
await call_next()
self.execution_log.append(f"{self.name}_end")
async def test_middleware_dynamic_rebuild_non_streaming(self, client: "MockChatClient") -> None:
@@ -924,9 +902,9 @@ class TestRunLevelMiddleware:
self.name = name
self.execution_log = execution_log
async def process(self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
self.execution_log.append(f"{self.name}_start")
await call_next(context)
await call_next()
self.execution_log.append(f"{self.name}_end")
async def test_run_level_middleware_isolation(self, client: "MockChatClient") -> None:
@@ -976,29 +954,25 @@ class TestRunLevelMiddleware:
def __init__(self, name: str):
self.name = name
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], 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"
await call_next(context)
await call_next()
execution_log.append(f"{self.name}_end")
class MetadataRunMiddleware(AgentMiddleware):
def __init__(self, name: str):
self.name = name
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_log.append(f"{self.name}_start")
# Read metadata set by agent middleware
for key, value in context.metadata.items():
metadata_log.append(f"{self.name}_reads_{key}:{value}")
# Set run-level metadata
context.metadata[f"{self.name}_key"] = f"{self.name}_value"
await call_next(context)
await call_next()
execution_log.append(f"{self.name}_end")
# Create agent with agent-level middleware
@@ -1049,12 +1023,10 @@ class TestRunLevelMiddleware:
def __init__(self, name: str):
self.name = name
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_log.append(f"{self.name}_start")
streaming_flags.append(context.stream)
await call_next(context)
await call_next()
execution_log.append(f"{self.name}_end")
# Create agent without agent-level middleware
@@ -1093,48 +1065,44 @@ class TestRunLevelMiddleware:
# Agent-level middleware
class AgentLevelAgentMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_log.append("agent_level_agent_start")
context.metadata["agent_level_agent"] = "processed"
await call_next(context)
await call_next()
execution_log.append("agent_level_agent_end")
class AgentLevelFunctionMiddleware(FunctionMiddleware):
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
execution_log.append("agent_level_function_start")
context.metadata["agent_level_function"] = "processed"
await call_next(context)
await call_next()
execution_log.append("agent_level_function_end")
# Run-level middleware
class RunLevelAgentMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_log.append("run_level_agent_start")
# Verify agent-level middleware metadata is available
assert "agent_level_agent" in context.metadata
context.metadata["run_level_agent"] = "processed"
await call_next(context)
await call_next()
execution_log.append("run_level_agent_end")
class RunLevelFunctionMiddleware(FunctionMiddleware):
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
execution_log.append("run_level_function_start")
# Verify agent-level function middleware metadata is available
assert "agent_level_function" in context.metadata
context.metadata["run_level_function"] = "processed"
await call_next(context)
await call_next()
execution_log.append("run_level_function_end")
# Create tool function for testing function middleware
@@ -1217,18 +1185,16 @@ class TestMiddlewareDecoratorLogic:
execution_order: list[str] = []
@agent_middleware
async def matching_agent_middleware(
context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def matching_agent_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("decorator_type_match_agent")
await call_next(context)
await call_next()
@function_middleware
async def matching_function_middleware(
context: FunctionInvocationContext, call_next: Callable[[FunctionInvocationContext], Awaitable[None]]
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
) -> None:
execution_order.append("decorator_type_match_function")
await call_next(context)
await call_next()
# Create tool function for testing function middleware
def custom_tool(message: str) -> str:
@@ -1282,7 +1248,7 @@ class TestMiddlewareDecoratorLogic:
context: FunctionInvocationContext, # Wrong type for @agent_middleware
call_next: Any,
) -> None:
await call_next(context)
await call_next()
agent = Agent(client=client, middleware=[mismatched_middleware])
await agent.run([Message(role="user", text="test")])
@@ -1294,12 +1260,12 @@ class TestMiddlewareDecoratorLogic:
@agent_middleware
async def decorator_only_agent(context: Any, call_next: Any) -> None: # No type annotation
execution_order.append("decorator_only_agent")
await call_next(context)
await call_next()
@function_middleware
async def decorator_only_function(context: Any, call_next: Any) -> None: # No type annotation
execution_order.append("decorator_only_function")
await call_next(context)
await call_next()
# Create tool function for testing function middleware
def custom_tool(message: str) -> str:
@@ -1346,16 +1312,16 @@ class TestMiddlewareDecoratorLogic:
execution_order: list[str] = []
# No decorator
async def type_only_agent(context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]) -> None:
async def type_only_agent(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("type_only_agent")
await call_next(context)
await call_next()
# No decorator
async def type_only_function(
context: FunctionInvocationContext, call_next: Callable[[FunctionInvocationContext], Awaitable[None]]
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
) -> None:
execution_order.append("type_only_function")
await call_next(context)
await call_next()
# Create tool function for testing function middleware
def custom_tool(message: str) -> str:
@@ -1399,7 +1365,7 @@ class TestMiddlewareDecoratorLogic:
"""Neither decorator nor parameter type specified - should throw exception."""
async def no_info_middleware(context: Any, call_next: Any) -> None: # No decorator, no type
await call_next(context)
await call_next()
# Should raise MiddlewareException
with pytest.raises(MiddlewareException, match="Cannot determine middleware type"):
@@ -1447,9 +1413,7 @@ class TestChatAgentThreadBehavior:
thread_states: list[dict[str, Any]] = []
class ThreadTrackingMiddleware(AgentMiddleware):
async def process(
self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def process(self, context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Capture state before next() call
thread_messages = []
if context.thread and context.thread.message_store:
@@ -1464,7 +1428,7 @@ class TestChatAgentThreadBehavior:
}
thread_states.append(before_state)
await call_next(context)
await call_next()
# Capture state after next() call
thread_messages_after = []
@@ -1560,9 +1524,9 @@ class TestChatAgentChatMiddleware:
execution_order: list[str] = []
class TrackingChatMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("chat_middleware_before")
await call_next(context)
await call_next()
execution_order.append("chat_middleware_after")
# Create Agent with chat middleware
@@ -1588,11 +1552,9 @@ class TestChatAgentChatMiddleware:
"""Test function-based chat middleware with Agent."""
execution_order: list[str] = []
async def tracking_chat_middleware(
context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]
) -> None:
async def tracking_chat_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("chat_middleware_before")
await call_next(context)
await call_next()
execution_order.append("chat_middleware_after")
# Create Agent with function-based chat middleware
@@ -1617,9 +1579,7 @@ class TestChatAgentChatMiddleware:
"""Test that chat middleware can modify messages before sending to model."""
@chat_middleware
async def message_modifier_middleware(
context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]
) -> None:
async def message_modifier_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Modify the first message by adding a prefix
if context.messages:
for idx, msg in enumerate(context.messages):
@@ -1628,7 +1588,7 @@ class TestChatAgentChatMiddleware:
original_text = msg.text or ""
context.messages[idx] = Message(role=msg.role, text=f"MODIFIED: {original_text}")
break
await call_next(context)
await call_next()
# Create Agent with message-modifying middleware
client = MockBaseChatClient()
@@ -1646,9 +1606,7 @@ class TestChatAgentChatMiddleware:
"""Test that chat middleware can override the response."""
@chat_middleware
async def response_override_middleware(
context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]
) -> None:
async def response_override_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Override the response without calling next()
context.result = ChatResponse(
messages=[Message(role="assistant", text="MiddlewareTypes overridden response")],
@@ -1675,15 +1633,15 @@ class TestChatAgentChatMiddleware:
execution_order: list[str] = []
@chat_middleware
async def first_middleware(context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def first_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("first_before")
await call_next(context)
await call_next()
execution_order.append("first_after")
@chat_middleware
async def second_middleware(context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def second_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("second_before")
await call_next(context)
await call_next()
execution_order.append("second_after")
# Create Agent with multiple chat middleware
@@ -1709,10 +1667,10 @@ class TestChatAgentChatMiddleware:
streaming_flags: list[bool] = []
class StreamingTrackingChatMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("streaming_chat_before")
streaming_flags.append(context.stream)
await call_next(context)
await call_next()
execution_order.append("streaming_chat_after")
# Create Agent with chat middleware
@@ -1749,13 +1707,13 @@ class TestChatAgentChatMiddleware:
execution_order: list[str] = []
class PreTerminationChatMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("middleware_before")
# Set a custom response since we're terminating
context.result = ChatResponse(messages=[Message(role="assistant", text="Terminated by middleware")])
raise MiddlewareTermination
# We call next() but since terminate=True, execution should stop
await call_next(context)
await call_next()
execution_order.append("middleware_after")
# Create Agent with terminating middleware
@@ -1777,9 +1735,9 @@ class TestChatAgentChatMiddleware:
execution_order: list[str] = []
class PostTerminationChatMiddleware(ChatMiddleware):
async def process(self, context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def process(self, context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("middleware_before")
await call_next(context)
await call_next()
execution_order.append("middleware_after")
context.terminate = True
@@ -1804,21 +1762,21 @@ class TestChatAgentChatMiddleware:
"""Test Agent with combined middleware types."""
execution_order: list[str] = []
async def agent_middleware(context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]) -> None:
async def agent_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("agent_middleware_before")
await call_next(context)
await call_next()
execution_order.append("agent_middleware_after")
async def chat_middleware(context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def chat_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("chat_middleware_before")
await call_next(context)
await call_next()
execution_order.append("chat_middleware_after")
async def function_middleware(
context: FunctionInvocationContext, call_next: Callable[[FunctionInvocationContext], Awaitable[None]]
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
) -> None:
execution_order.append("function_middleware_before")
await call_next(context)
await call_next()
execution_order.append("function_middleware_after")
# Create Agent with function middleware and tools
@@ -1842,9 +1800,7 @@ class TestChatAgentChatMiddleware:
modified_kwargs: dict[str, Any] = {}
@agent_middleware
async def kwargs_middleware(
context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
) -> None:
async def kwargs_middleware(context: AgentContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Capture the original kwargs
captured_kwargs.update(context.kwargs)
@@ -1856,7 +1812,7 @@ class TestChatAgentChatMiddleware:
# Store modified kwargs for verification
modified_kwargs.update(context.kwargs)
await call_next(context)
await call_next()
# Create Agent with agent middleware
client = MockBaseChatClient()
@@ -1895,10 +1851,10 @@ class TestChatAgentChatMiddleware:
# class TrackingMiddleware(AgentMiddleware):
# async def process(
# self, context: AgentContext, call_next: Callable[[AgentContext], Awaitable[None]]
# self, context: AgentContext, call_next: Callable[[], Awaitable[None]]
# ) -> None:
# execution_order.append("before")
# await call_next(context)
# await call_next()
# execution_order.append("after")
# @use_agent_middleware
@@ -32,10 +32,10 @@ class TestChatMiddleware:
async def process(
self,
context: ChatContext,
call_next: Callable[[ChatContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
execution_order.append("chat_middleware_before")
await call_next(context)
await call_next()
execution_order.append("chat_middleware_after")
# Add middleware to chat client
@@ -58,11 +58,9 @@ class TestChatMiddleware:
execution_order: list[str] = []
@chat_middleware
async def logging_chat_middleware(
context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]
) -> None:
async def logging_chat_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("function_middleware_before")
await call_next(context)
await call_next()
execution_order.append("function_middleware_after")
# Add middleware to chat client
@@ -84,14 +82,12 @@ class TestChatMiddleware:
"""Test that chat middleware can modify messages before sending to model."""
@chat_middleware
async def message_modifier_middleware(
context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]
) -> None:
async def message_modifier_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Modify the first message by adding a prefix
if context.messages and len(context.messages) > 0:
original_text = context.messages[0].text or ""
context.messages[0] = Message(role=context.messages[0].role, text=f"MODIFIED: {original_text}")
await call_next(context)
await call_next()
# Add middleware to chat client
chat_client_base.chat_middleware = [message_modifier_middleware]
@@ -110,9 +106,7 @@ class TestChatMiddleware:
"""Test that chat middleware can override the response."""
@chat_middleware
async def response_override_middleware(
context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]
) -> None:
async def response_override_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Override the response without calling next()
context.result = ChatResponse(
messages=[Message(role="assistant", text="MiddlewareTypes overridden response")],
@@ -138,15 +132,15 @@ class TestChatMiddleware:
execution_order: list[str] = []
@chat_middleware
async def first_middleware(context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def first_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("first_before")
await call_next(context)
await call_next()
execution_order.append("first_after")
@chat_middleware
async def second_middleware(context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def second_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("second_before")
await call_next(context)
await call_next()
execution_order.append("second_after")
# Add middleware to chat client (order should be preserved)
@@ -173,11 +167,9 @@ class TestChatMiddleware:
execution_order: list[str] = []
@chat_middleware
async def agent_level_chat_middleware(
context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]
) -> None:
async def agent_level_chat_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("agent_chat_middleware_before")
await call_next(context)
await call_next()
execution_order.append("agent_chat_middleware_after")
client = MockBaseChatClient()
@@ -205,15 +197,15 @@ class TestChatMiddleware:
execution_order: list[str] = []
@chat_middleware
async def first_middleware(context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def first_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("first_before")
await call_next(context)
await call_next()
execution_order.append("first_after")
@chat_middleware
async def second_middleware(context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def second_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("second_before")
await call_next(context)
await call_next()
execution_order.append("second_after")
# Create Agent with multiple chat middleware
@@ -240,9 +232,7 @@ class TestChatMiddleware:
execution_order: list[str] = []
@chat_middleware
async def streaming_middleware(
context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]
) -> None:
async def streaming_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_order.append("streaming_before")
# Verify it's a streaming context
assert context.stream is True
@@ -254,7 +244,7 @@ class TestChatMiddleware:
return update
context.stream_transform_hooks.append(upper_case_update)
await call_next(context)
await call_next()
execution_order.append("streaming_after")
# Add middleware to chat client
@@ -278,11 +268,9 @@ class TestChatMiddleware:
execution_count = {"count": 0}
@chat_middleware
async def counting_middleware(
context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]
) -> None:
async def counting_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
execution_count["count"] += 1
await call_next(context)
await call_next()
# First call with run-level middleware
messages = [Message(role="user", text="first message")]
@@ -310,7 +298,7 @@ class TestChatMiddleware:
modified_kwargs: dict[str, Any] = {}
@chat_middleware
async def kwargs_middleware(context: ChatContext, call_next: Callable[[ChatContext], Awaitable[None]]) -> None:
async def kwargs_middleware(context: ChatContext, call_next: Callable[[], Awaitable[None]]) -> None:
# Capture the original kwargs
captured_kwargs.update(context.kwargs)
@@ -322,7 +310,7 @@ class TestChatMiddleware:
# Store modified kwargs for verification
modified_kwargs.update(context.kwargs)
await call_next(context)
await call_next()
# Add middleware to chat client
chat_client_base.chat_middleware = [kwargs_middleware]
@@ -355,11 +343,11 @@ class TestChatMiddleware:
@function_middleware
async def test_function_middleware(
context: FunctionInvocationContext, call_next: Callable[[FunctionInvocationContext], Awaitable[None]]
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
) -> None:
nonlocal execution_order
execution_order.append(f"function_middleware_before_{context.function.name}")
await call_next(context)
await call_next()
execution_order.append(f"function_middleware_after_{context.function.name}")
# Define a simple tool function
@@ -421,10 +409,10 @@ class TestChatMiddleware:
@function_middleware
async def run_level_function_middleware(
context: FunctionInvocationContext, call_next: Callable[[FunctionInvocationContext], Awaitable[None]]
context: FunctionInvocationContext, call_next: Callable[[], Awaitable[None]]
) -> None:
execution_order.append("run_level_function_middleware_before")
await call_next(context)
await call_next()
execution_order.append("run_level_function_middleware_after")
# Define a simple tool function
@@ -207,7 +207,7 @@ def test_serialize(ollama_unit_test_env: dict[str, str]) -> None:
def test_chat_middleware(ollama_unit_test_env: dict[str, str]) -> None:
@chat_middleware
async def sample_middleware(context, call_next):
await call_next(context)
await call_next()
ollama_chat_client = OllamaChatClient(middleware=[sample_middleware])
assert len(ollama_chat_client.middleware) == 1
@@ -129,11 +129,11 @@ class _AutoHandoffMiddleware(FunctionMiddleware):
async def process(
self,
context: FunctionInvocationContext,
call_next: Callable[[FunctionInvocationContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None:
"""Intercept matching handoff tool calls and inject synthetic results."""
if context.function.name not in self._handoff_functions:
await call_next(context)
await call_next()
return
from agent_framework._middleware import MiddlewareTermination
@@ -65,7 +65,7 @@ class PurviewPolicyMiddleware(AgentMiddleware):
async def process(
self,
context: AgentContext,
call_next: Callable[[AgentContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None: # type: ignore[override]
resolved_user_id: str | None = None
try:
@@ -92,7 +92,7 @@ class PurviewPolicyMiddleware(AgentMiddleware):
if not self._settings.ignore_exceptions:
raise
await call_next(context)
await call_next()
try:
# Post (response) check only if we have a normal AgentResponse
@@ -162,7 +162,7 @@ class PurviewChatPolicyMiddleware(ChatMiddleware):
async def process(
self,
context: ChatContext,
call_next: Callable[[ChatContext], Awaitable[None]],
call_next: Callable[[], Awaitable[None]],
) -> None: # type: ignore[override]
resolved_user_id: str | None = None
try:
@@ -187,7 +187,7 @@ class PurviewChatPolicyMiddleware(ChatMiddleware):
if not self._settings.ignore_exceptions:
raise
await call_next(context)
await call_next()
try:
# Post (response) evaluation only if non-streaming and we have messages result shape
@@ -49,7 +49,7 @@ class TestPurviewChatPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
next_called = False
async def mock_next(ctx: ChatContext) -> None:
async def mock_next() -> None:
nonlocal next_called
next_called = True
@@ -57,7 +57,7 @@ class TestPurviewChatPolicyMiddleware:
def __init__(self):
self.messages = [Message(role="assistant", text="Hi there")]
ctx.result = Result()
chat_context.result = Result()
await middleware.process(chat_context, mock_next)
assert next_called
@@ -67,7 +67,7 @@ class TestPurviewChatPolicyMiddleware:
async def test_blocks_prompt(self, middleware: PurviewChatPolicyMiddleware, chat_context: ChatContext) -> None:
with patch.object(middleware._processor, "process_messages", return_value=(True, "user-123")):
async def mock_next(ctx: ChatContext) -> None: # should not run
async def mock_next() -> None: # should not run
raise AssertionError("next should not be called when prompt blocked")
with pytest.raises(MiddlewareTermination):
@@ -88,12 +88,12 @@ class TestPurviewChatPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", side_effect=side_effect):
async def mock_next(ctx: ChatContext) -> None:
async def mock_next() -> None:
class Result:
def __init__(self):
self.messages = [Message(role="assistant", text="Sensitive output")] # pragma: no cover
ctx.result = Result()
chat_context.result = Result()
await middleware.process(chat_context, mock_next)
assert call_state["count"] == 2
@@ -114,8 +114,8 @@ class TestPurviewChatPolicyMiddleware:
)
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
async def mock_next(ctx: ChatContext) -> None:
ctx.result = MagicMock()
async def mock_next() -> None:
streaming_context.result = MagicMock()
await middleware.process(streaming_context, mock_next)
assert mock_proc.call_count == 1
@@ -138,10 +138,10 @@ class TestPurviewChatPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
async def mock_next(ctx: ChatContext) -> None:
async def mock_next() -> None:
result = MagicMock()
result.messages = [Message(role="assistant", text="Response")]
ctx.result = result
chat_context.result = result
await middleware.process(chat_context, mock_next)
@@ -162,10 +162,10 @@ class TestPurviewChatPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
async def mock_next(ctx: ChatContext) -> None:
async def mock_next() -> None:
result = MagicMock()
result.messages = [Message(role="assistant", text="Response")]
ctx.result = result
chat_context.result = result
await middleware.process(chat_context, mock_next)
@@ -194,7 +194,7 @@ class TestPurviewChatPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
async def mock_next(ctx: ChatContext) -> None:
async def mock_next() -> None:
raise AssertionError("next should not be called")
# Should raise the exception
@@ -224,10 +224,10 @@ class TestPurviewChatPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", side_effect=side_effect):
async def mock_next(ctx: ChatContext) -> None:
async def mock_next() -> None:
result = MagicMock()
result.messages = [Message(role="assistant", text="OK")]
ctx.result = result
context.result = result
with pytest.raises(PurviewPaymentRequiredError):
await middleware.process(context, mock_next)
@@ -249,7 +249,7 @@ class TestPurviewChatPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
async def mock_next(ctx: ChatContext) -> None:
async def mock_next() -> None:
result = MagicMock()
result.messages = [Message(role="assistant", text="Response")]
context.result = result
@@ -265,9 +265,9 @@ class TestPurviewChatPolicyMiddleware:
"""Test middleware handles result that doesn't have messages attribute."""
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")):
async def mock_next(ctx: ChatContext) -> None:
async def mock_next() -> None:
# Set result to something without messages attribute
ctx.result = "Some string result"
chat_context.result = "Some string result"
await middleware.process(chat_context, mock_next)
@@ -289,7 +289,7 @@ class TestPurviewChatPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
async def mock_next(ctx: ChatContext) -> None:
async def mock_next() -> None:
result = MagicMock()
result.messages = [Message(role="assistant", text="Response")]
context.result = result
@@ -313,7 +313,7 @@ class TestPurviewChatPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", side_effect=ValueError("boom")):
async def mock_next(_: ChatContext) -> None:
async def mock_next() -> None:
raise AssertionError("next should not be called")
with pytest.raises(ValueError, match="boom"):
@@ -342,10 +342,10 @@ class TestPurviewChatPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", side_effect=side_effect):
async def mock_next(ctx: ChatContext) -> None:
async def mock_next() -> None:
result = MagicMock()
result.messages = [Message(role="assistant", text="OK")]
ctx.result = result
context.result = result
with pytest.raises(ValueError, match="post"):
await middleware.process(context, mock_next)
@@ -361,10 +361,10 @@ class TestPurviewChatPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
async def mock_next(ctx: ChatContext) -> None:
async def mock_next() -> None:
result = MagicMock()
result.messages = [Message(role="assistant", text="Hi")]
ctx.result = result
context.result = result
await middleware.process(context, mock_next)
@@ -382,10 +382,10 @@ class TestPurviewChatPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
async def mock_next(ctx: ChatContext) -> None:
async def mock_next() -> None:
result = MagicMock()
result.messages = [Message(role="assistant", text="Hi")]
ctx.result = result
context.result = result
await middleware.process(context, mock_next)
@@ -401,10 +401,10 @@ class TestPurviewChatPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
async def mock_next(ctx: ChatContext) -> None:
async def mock_next() -> None:
result = MagicMock()
result.messages = [Message(role="assistant", text="Response")]
ctx.result = result
context.result = result
await middleware.process(context, mock_next)
@@ -55,10 +55,10 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")):
next_called = False
async def mock_next(ctx: AgentContext) -> None:
async def mock_next() -> None:
nonlocal next_called
next_called = True
ctx.result = AgentResponse(messages=[Message(role="assistant", text="I'm good, thanks!")])
context.result = AgentResponse(messages=[Message(role="assistant", text="I'm good, thanks!")])
await middleware.process(context, mock_next)
@@ -74,7 +74,7 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", return_value=(True, "user-123")):
next_called = False
async def mock_next(ctx: AgentContext) -> None:
async def mock_next() -> None:
nonlocal next_called
next_called = True
@@ -101,8 +101,8 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
async def mock_next(ctx: AgentContext) -> None:
ctx.result = AgentResponse(
async def mock_next() -> None:
context.result = AgentResponse(
messages=[Message(role="assistant", text="Here's some sensitive information")]
)
@@ -125,8 +125,8 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")):
async def mock_next(ctx: AgentContext) -> None:
ctx.result = "Some non-standard result"
async def mock_next() -> None:
context.result = "Some non-standard result"
await middleware.process(context, mock_next)
@@ -142,8 +142,8 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_process:
async def mock_next(ctx: AgentContext) -> None:
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
async def mock_next() -> None:
context.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
await middleware.process(context, mock_next)
@@ -160,8 +160,8 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
async def mock_next(ctx: AgentContext) -> None:
ctx.result = AgentResponse(messages=[Message(role="assistant", text="streaming")])
async def mock_next() -> None:
context.result = AgentResponse(messages=[Message(role="assistant", text="streaming")])
await middleware.process(context, mock_next)
@@ -181,7 +181,7 @@ class TestPurviewPolicyMiddleware:
side_effect=PurviewPaymentRequiredError("Payment required"),
):
async def mock_next(_: AgentContext) -> None:
async def mock_next() -> None:
raise AssertionError("next should not be called")
with pytest.raises(PurviewPaymentRequiredError):
@@ -206,8 +206,8 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", side_effect=side_effect):
async def mock_next(ctx: AgentContext) -> None:
ctx.result = AgentResponse(messages=[Message(role="assistant", text="OK")])
async def mock_next() -> None:
context.result = AgentResponse(messages=[Message(role="assistant", text="OK")])
with pytest.raises(PurviewPaymentRequiredError):
await middleware.process(context, mock_next)
@@ -231,8 +231,8 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", side_effect=side_effect):
async def mock_next(ctx: AgentContext) -> None:
ctx.result = AgentResponse(messages=[Message(role="assistant", text="OK")])
async def mock_next() -> None:
context.result = AgentResponse(messages=[Message(role="assistant", text="OK")])
with pytest.raises(ValueError, match="Post-check blew up"):
await middleware.process(context, mock_next)
@@ -250,8 +250,8 @@ class TestPurviewPolicyMiddleware:
middleware._processor, "process_messages", side_effect=Exception("Pre-check error")
) as mock_process:
async def mock_next(ctx: AgentContext) -> None:
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
async def mock_next() -> None:
context.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
await middleware.process(context, mock_next)
@@ -280,8 +280,8 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
async def mock_next(ctx: AgentContext) -> None:
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
async def mock_next() -> None:
context.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
await middleware.process(context, mock_next)
@@ -306,8 +306,8 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
async def mock_next(ctx):
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
async def mock_next():
context.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
# Should not raise, just log
await middleware.process(context, mock_next)
@@ -330,7 +330,7 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
async def mock_next(ctx):
async def mock_next():
pass
# Should raise the exception
@@ -346,8 +346,8 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
async def mock_next(ctx: AgentContext) -> None:
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
async def mock_next() -> None:
context.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
await middleware.process(context, mock_next)
@@ -364,8 +364,8 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
async def mock_next(ctx: AgentContext) -> None:
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
async def mock_next() -> None:
context.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
await middleware.process(context, mock_next)
@@ -383,8 +383,8 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
async def mock_next(ctx: AgentContext) -> None:
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
async def mock_next() -> None:
context.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
await middleware.process(context, mock_next)
@@ -399,8 +399,8 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
async def mock_next(ctx: AgentContext) -> None:
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
async def mock_next() -> None:
context.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
await middleware.process(context, mock_next)
@@ -416,8 +416,8 @@ class TestPurviewPolicyMiddleware:
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
async def mock_next(ctx: AgentContext) -> None:
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
async def mock_next() -> None:
context.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
await middleware.process(context, mock_next)