diff --git a/python/packages/main/agent_framework/_middleware.py b/python/packages/main/agent_framework/_middleware.py index 9b0e1debfc..59f6dd5c29 100644 --- a/python/packages/main/agent_framework/_middleware.py +++ b/python/packages/main/agent_framework/_middleware.py @@ -15,8 +15,8 @@ if TYPE_CHECKING: TAgent = TypeVar("TAgent", bound="AgentProtocol") __all__ = [ - "AgentInvocationContext", "AgentMiddleware", + "AgentRunContext", "FunctionInvocationContext", "FunctionMiddleware", "MiddlewareType", @@ -24,7 +24,7 @@ __all__ = [ ] -class AgentInvocationContext: +class AgentRunContext: """Context object for agent middleware invocations. Attributes: @@ -88,8 +88,8 @@ class AgentMiddleware(ABC): @abstractmethod async def process( self, - context: AgentInvocationContext, - next: Callable[[AgentInvocationContext], Awaitable[None]], + context: AgentRunContext, + next: Callable[[AgentRunContext], Awaitable[None]], ) -> None: """Process an agent invocation. @@ -136,9 +136,7 @@ class FunctionMiddleware(ABC): # Pure function type definitions for convenience -AgentMiddlewareCallable = Callable[ - [AgentInvocationContext, Callable[[AgentInvocationContext], Awaitable[None]]], Awaitable[None] -] +AgentMiddlewareCallable = Callable[[AgentRunContext, Callable[[AgentRunContext], Awaitable[None]]], Awaitable[None]] FunctionMiddlewareCallable = Callable[ [FunctionInvocationContext, Callable[[FunctionInvocationContext], Awaitable[None]]], Awaitable[None] @@ -156,8 +154,8 @@ class AgentMiddlewareWrapper(AgentMiddleware): async def process( self, - context: AgentInvocationContext, - next: Callable[[AgentInvocationContext], Awaitable[None]], + context: AgentRunContext, + next: Callable[[AgentRunContext], Awaitable[None]], ) -> None: await self.func(context, next) @@ -202,8 +200,8 @@ class AgentMiddlewarePipeline: self, agent: "AgentProtocol", messages: list[ChatMessage], - context: AgentInvocationContext, - final_handler: Callable[[AgentInvocationContext], Awaitable[AgentRunResponse]], + context: AgentRunContext, + final_handler: Callable[[AgentRunContext], Awaitable[AgentRunResponse]], ) -> AgentRunResponse | None: """Execute the agent middleware pipeline for non-streaming. @@ -227,10 +225,10 @@ class AgentMiddlewarePipeline: # Store the final result result_container: dict[str, AgentRunResponse | None] = {"response": None} - def create_next_handler(index: int) -> Callable[[AgentInvocationContext], Awaitable[None]]: + def create_next_handler(index: int) -> Callable[[AgentRunContext], Awaitable[None]]: if index >= len(self._middlewares): - async def final_wrapper(c: AgentInvocationContext) -> None: + async def final_wrapper(c: AgentRunContext) -> None: result_container["response"] = await final_handler(c) return final_wrapper @@ -238,7 +236,7 @@ class AgentMiddlewarePipeline: middleware = self._middlewares[index] next_handler = create_next_handler(index + 1) - async def current_handler(c: AgentInvocationContext) -> None: + async def current_handler(c: AgentRunContext) -> None: await middleware.process(c, next_handler) return current_handler @@ -253,8 +251,8 @@ class AgentMiddlewarePipeline: self, agent: "AgentProtocol", messages: list[ChatMessage], - context: AgentInvocationContext, - final_handler: Callable[[AgentInvocationContext], AsyncIterable[AgentRunResponseUpdate]], + context: AgentRunContext, + final_handler: Callable[[AgentRunContext], AsyncIterable[AgentRunResponseUpdate]], ) -> AsyncIterable[AgentRunResponseUpdate]: """Execute the agent middleware pipeline for streaming. @@ -280,10 +278,10 @@ class AgentMiddlewarePipeline: # Store the final result result_container: dict[str, AsyncIterable[AgentRunResponseUpdate] | None] = {"response_stream": None} - def create_next_handler(index: int) -> Callable[[AgentInvocationContext], Awaitable[None]]: + def create_next_handler(index: int) -> Callable[[AgentRunContext], Awaitable[None]]: if index >= len(self._middlewares): - async def final_wrapper(c: AgentInvocationContext) -> None: # noqa: RUF029 + async def final_wrapper(c: AgentRunContext) -> None: # noqa: RUF029 result_container["response_stream"] = final_handler(c) return final_wrapper @@ -291,7 +289,7 @@ class AgentMiddlewarePipeline: middleware = self._middlewares[index] next_handler = create_next_handler(index + 1) - async def current_handler(c: AgentInvocationContext) -> None: + async def current_handler(c: AgentRunContext) -> None: await middleware.process(c, next_handler) return current_handler @@ -434,10 +432,10 @@ def use_agent_middleware(agent_class: type[TAgent]) -> type[TAgent]: params = list(sig.parameters.values()) if len(params) >= 1: first_param = params[0] - # Check if first parameter is AgentInvocationContext or FunctionInvocationContext + # Check if first parameter is AgentRunContext or FunctionInvocationContext if ( hasattr(first_param.annotation, "__name__") - and first_param.annotation.__name__ == "AgentInvocationContext" + and first_param.annotation.__name__ == "AgentRunContext" ): agent_middlewares.append(middleware) # type: ignore elif ( @@ -497,13 +495,13 @@ def use_agent_middleware(agent_class: type[TAgent]) -> type[TAgent]: # Execute with middleware if available if self._agent_middleware_pipeline.has_middlewares: - context = AgentInvocationContext( + context = AgentRunContext( agent=self, # type: ignore[arg-type] messages=normalized_messages, is_streaming=False, ) - async def _execute_handler(ctx: AgentInvocationContext) -> AgentRunResponse: + 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( @@ -555,13 +553,13 @@ def use_agent_middleware(agent_class: type[TAgent]) -> type[TAgent]: # Execute with middleware if available if self._agent_middleware_pipeline.has_middlewares: - context = AgentInvocationContext( + context = AgentRunContext( agent=self, # type: ignore[arg-type] messages=normalized_messages, is_streaming=True, ) - async def _execute_stream_handler(ctx: AgentInvocationContext) -> AsyncIterable[AgentRunResponseUpdate]: + async def _execute_stream_handler(ctx: AgentRunContext) -> AsyncIterable[AgentRunResponseUpdate]: async for update in original_run_stream(self, ctx.messages, thread=thread, **kwargs): # type: ignore[misc] yield update diff --git a/python/samples/getting_started/middleware/class_based_middleware.py b/python/samples/getting_started/middleware/class_based_middleware.py index 7a1bb324d0..e8e501d560 100644 --- a/python/samples/getting_started/middleware/class_based_middleware.py +++ b/python/samples/getting_started/middleware/class_based_middleware.py @@ -6,7 +6,7 @@ from collections.abc import Awaitable, Callable from random import randint from typing import Annotated -from agent_framework import AgentInvocationContext, AgentMiddleware, FunctionInvocationContext, FunctionMiddleware +from agent_framework import AgentMiddleware, AgentRunContext, FunctionInvocationContext, FunctionMiddleware from agent_framework.foundry import FoundryChatClient from azure.identity.aio import AzureCliCredential from pydantic import Field @@ -25,8 +25,8 @@ class SecurityAgentMiddleware(AgentMiddleware): async def process( self, - context: AgentInvocationContext, - next: Callable[[AgentInvocationContext], Awaitable[None]], + context: AgentRunContext, + next: Callable[[AgentRunContext], Awaitable[None]], ) -> None: # Check for potential security violations in the query # Look at the last user message diff --git a/python/samples/getting_started/middleware/function_based_middleware.py b/python/samples/getting_started/middleware/function_based_middleware.py index 5c31f79936..da0cfd6985 100644 --- a/python/samples/getting_started/middleware/function_based_middleware.py +++ b/python/samples/getting_started/middleware/function_based_middleware.py @@ -7,7 +7,7 @@ from random import randint from typing import Annotated from agent_framework import ( - AgentInvocationContext, + AgentRunContext, FunctionInvocationContext, ) from agent_framework.foundry import FoundryChatClient @@ -24,8 +24,8 @@ def get_weather( async def security_agent_middleware( - context: AgentInvocationContext, - next: Callable[[AgentInvocationContext], Awaitable[None]], + context: AgentRunContext, + next: Callable[[AgentRunContext], Awaitable[None]], ) -> None: """Agent middleware that checks for security violations.""" # Check for potential security violations in the query