Renaming and updates to logic

This commit is contained in:
Dmytro Struk
2025-09-17 15:35:01 -07:00
Unverified
parent 89eeb58bdc
commit dd534a9e92
3 changed files with 47 additions and 90 deletions
@@ -32,10 +32,10 @@ class AgentRunContext:
messages: The messages being sent to the agent.
is_streaming: Whether this is a streaming invocation.
metadata: Metadata dictionary for sharing data between agent middleware.
response: Agent execution response. Can be set before calling next() to override execution,
or observed after calling next() to see the actual execution result.
For non-streaming: should be AgentRunResponse
For streaming: should be AsyncIterable[AgentRunResponseUpdate]
result: Agent execution result. Can be observed after calling next()
to see the actual execution result or can be set to override the execution result.
For non-streaming: should be AgentRunResponse
For streaming: should be AsyncIterable[AgentRunResponseUpdate]
"""
def __init__(
@@ -57,7 +57,7 @@ class AgentRunContext:
self.messages = messages
self.is_streaming = is_streaming
self.metadata = metadata or {}
self.response: AgentRunResponse | AsyncIterable[AgentRunResponseUpdate] | None = None
self.result: AgentRunResponse | AsyncIterable[AgentRunResponseUpdate] | None = None
class FunctionInvocationContext:
@@ -67,8 +67,8 @@ class FunctionInvocationContext:
function: The function being invoked.
arguments: The validated arguments for the function.
metadata: Metadata dictionary for sharing data between function middleware.
result: Function execution result. Can be set before calling next() to override execution,
or observed after calling next() to see the actual execution result.
result: Function execution result. Can be observed after calling next()
to see the actual execution result or can be set to override the execution result.
"""
def __init__(
@@ -104,7 +104,7 @@ class AgentMiddleware(ABC):
Args:
context: Agent invocation context containing agent, messages, and metadata.
Use context.is_streaming to determine if this is a streaming call.
Middleware can set context.response to override execution, or observe
Middleware can set context.result to override execution, or observe
the actual execution result after calling next().
For non-streaming: AgentRunResponse
For streaming: AsyncIterable[AgentRunResponseUpdate]
@@ -113,8 +113,8 @@ class AgentMiddleware(ABC):
Note:
Middleware should not return anything. All data manipulation should happen
within the context object. Set context.response to override execution,
or observe context.response after calling next() for actual results.
within the context object. Set context.result to override execution,
or observe context.result after calling next() for actual results.
"""
...
@@ -239,14 +239,10 @@ class AgentMiddlewarePipeline:
if index >= len(self._middlewares):
async def final_wrapper(c: AgentRunContext) -> None:
# If response was set before calling next(), skip execution
if c.response is not None and isinstance(c.response, AgentRunResponse):
result_container["response"] = c.response
return
# Execute actual handler and populate context for observability
result = await final_handler(c)
result_container["response"] = result
c.response = result
result_container["result"] = result
c.result = result
return final_wrapper
@@ -256,22 +252,22 @@ class AgentMiddlewarePipeline:
async def current_handler(c: AgentRunContext) -> None:
await middleware.process(c, next_handler)
# After middleware execution, check if response was overridden
if c.response is not None and isinstance(c.response, AgentRunResponse):
result_container["response"] = c.response
if c.result is not None and isinstance(c.result, AgentRunResponse):
result_container["result"] = c.result
return current_handler
first_handler = create_next_handler(0)
await first_handler(context)
# Return the response from result container or overridden response
if context.response is not None and isinstance(context.response, AgentRunResponse):
return context.response
# Return the result from result container or overridden result
if context.result is not None and isinstance(context.result, AgentRunResponse):
return context.result
# If no response was set (next() not called), return empty AgentRunResponse
response = result_container["response"]
# If no result was set (next() not called), return empty AgentRunResponse
response = result_container.get("result")
if response is None:
return AgentRunResponse(messages=[])
return AgentRunResponse()
return response
async def execute_stream(
@@ -303,21 +299,16 @@ class AgentMiddlewarePipeline:
return
# Store the final result
result_container: dict[str, AsyncIterable[AgentRunResponseUpdate] | None] = {"response_stream": None}
result_container: dict[str, AsyncIterable[AgentRunResponseUpdate] | None] = {"result_stream": None}
def create_next_handler(index: int) -> Callable[[AgentRunContext], Awaitable[None]]:
if index >= len(self._middlewares):
async def final_wrapper(c: AgentRunContext) -> None: # noqa: RUF029
# If response was set before calling next(), skip execution
if c.response is not None and hasattr(c.response, "__aiter__"):
result_container["response_stream"] = c.response # type: ignore
return
# Execute actual handler and populate context for observability
result = final_handler(c)
result_container["response_stream"] = result
c.response = result
result_container["result_stream"] = result
c.result = result
return final_wrapper
@@ -332,18 +323,18 @@ class AgentMiddlewarePipeline:
first_handler = create_next_handler(0)
await first_handler(context)
# Yield from the response stream in result container or overridden response
if context.response is not None and hasattr(context.response, "__aiter__"):
async for update in context.response: # type: ignore
# Yield from the result stream in result container or overridden result
if context.result is not None and hasattr(context.result, "__aiter__"):
async for update in context.result: # type: ignore
yield update
return
response_stream = result_container["response_stream"]
if response_stream is None:
# If no response stream was set (next() not called), yield nothing
result_stream = result_container["result_stream"]
if result_stream is None:
# If no result stream was set (next() not called), yield nothing
return
async for update in response_stream:
async for update in result_stream:
yield update
@property
@@ -555,14 +546,14 @@ def use_agent_middleware(agent_class: type[TAgent]) -> type[TAgent]:
async def _execute_handler(ctx: AgentRunContext) -> AgentRunResponse:
return await original_run(self, ctx.messages, thread=thread, **kwargs) # type: ignore
response = await self._agent_middleware_pipeline.execute(
result = await self._agent_middleware_pipeline.execute(
self, # type: ignore[arg-type]
normalized_messages,
context,
_execute_handler,
)
return response if response else AgentRunResponse()
return result if result else AgentRunResponse()
# No middleware, execute directly
return await original_run(self, normalized_messages, thread=thread, **kwargs) # type: ignore[return-value]
@@ -778,7 +778,7 @@ class TestMiddlewareExecutionControl:
assert isinstance(result, AgentRunResponse)
assert result.messages == [] # Empty response
assert not handler_called
assert context.response is None
assert context.result is None
async def test_agent_middleware_no_next_no_streaming_execution(self, mock_agent: AgentProtocol) -> None:
"""Test that when agent middleware doesn't call next(), no streaming execution happens."""
@@ -810,7 +810,7 @@ class TestMiddlewareExecutionControl:
# Verify no execution happened and no updates were yielded
assert len(updates) == 0
assert not handler_called
assert context.response is None
assert context.result is None
async def test_function_middleware_no_next_no_execution(self, mock_function: AIFunction[Any, Any]) -> None:
"""Test that when function middleware doesn't call next(), no execution happens."""
@@ -884,40 +884,6 @@ class TestMiddlewareExecutionControl:
assert result.messages == [] # Empty response
assert not handler_called
async def test_agent_middleware_pre_execution_override_with_next(self, mock_agent: AgentProtocol) -> None:
"""Test that middleware can override response before calling next() - this skips handler execution."""
class PreOverrideMiddleware(AgentMiddleware):
async def process(
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
# Set override first
context.response = AgentRunResponse(
messages=[ChatMessage(role=Role.ASSISTANT, text="pre-override response")]
)
# Then call next() to continue middleware pipeline
await next(context)
middleware = PreOverrideMiddleware()
pipeline = AgentMiddlewarePipeline([middleware])
messages = [ChatMessage(role=Role.USER, text="test")]
context = AgentRunContext(agent=mock_agent, messages=messages)
handler_called = False
async def final_handler(ctx: AgentRunContext) -> AgentRunResponse:
nonlocal handler_called
handler_called = True
# This should not be called when response is pre-set
return AgentRunResponse(messages=[ChatMessage(role=Role.ASSISTANT, text="original response")])
result = await pipeline.execute(mock_agent, messages, context, final_handler)
# Verify pre-override worked and handler was NOT called (because response was already set)
assert result is not None
assert result.messages[0].text == "pre-override response"
assert not handler_called
async def test_function_middleware_pre_execution_override_with_next(
self, mock_function: AIFunction[Any, Any]
) -> None:
@@ -48,7 +48,7 @@ class TestResultOverrideMiddleware:
) -> None:
# Execute the pipeline first, then override the response
await next(context)
context.response = override_response
context.result = override_response
middleware = ResponseOverrideMiddleware()
pipeline = AgentMiddlewarePipeline([middleware])
@@ -84,7 +84,7 @@ class TestResultOverrideMiddleware:
) -> None:
# Execute the pipeline first, then override the response stream
await next(context)
context.response = override_stream()
context.result = override_stream()
middleware = StreamResponseOverrideMiddleware()
pipeline = AgentMiddlewarePipeline([middleware])
@@ -148,7 +148,7 @@ class TestResultOverrideMiddleware:
await next(context)
# Then conditionally override based on content
if any("special" in msg.text for msg in context.messages if msg.text):
context.response = AgentRunResponse(
context.result = AgentRunResponse(
messages=[ChatMessage(role=Role.ASSISTANT, text="Special response from middleware!")]
)
@@ -187,7 +187,7 @@ class TestResultOverrideMiddleware:
await next(context)
# Then conditionally override based on content
if any("custom stream" in msg.text for msg in context.messages if msg.text):
context.response = custom_stream()
context.result = custom_stream()
# Create ChatAgent with override middleware
middleware = ChatAgentStreamOverrideMiddleware()
@@ -246,7 +246,7 @@ class TestResultOverrideMiddleware:
assert isinstance(no_execute_result, AgentRunResponse)
assert no_execute_result.messages == [] # Empty response
assert not handler_called
assert no_execute_context.response is None
assert no_execute_context.result is None
# Reset for next test
handler_called = False
@@ -320,15 +320,15 @@ class TestResultObservability:
self, context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
# Context should be empty before next()
assert context.response is None
assert context.result is None
# Call next to execute
await next(context)
# Context should now contain the response for observability
assert context.response is not None
assert isinstance(context.response, AgentRunResponse)
observed_responses.append(context.response)
assert context.result is not None
assert isinstance(context.result, AgentRunResponse)
observed_responses.append(context.result)
middleware = ObservabilityMiddleware()
pipeline = AgentMiddlewarePipeline([middleware])
@@ -391,12 +391,12 @@ class TestResultObservability:
await next(context)
# Now observe and conditionally override
assert context.response is not None
assert isinstance(context.response, AgentRunResponse)
assert context.result is not None
assert isinstance(context.result, AgentRunResponse)
if "modify" in context.response.messages[0].text:
if "modify" in context.result.messages[0].text:
# Override after observing
context.response = AgentRunResponse(
context.result = AgentRunResponse(
messages=[ChatMessage(role=Role.ASSISTANT, text="modified after execution")]
)