mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Renaming and updates to logic
This commit is contained in:
@@ -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")]
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user