From 1744348df8fb93c0e5ca42ea92a5c75f332be196 Mon Sep 17 00:00:00 2001 From: Dmytro Struk <13853051+dmytrostruk@users.noreply.github.com> Date: Mon, 15 Sep 2025 22:54:14 -0700 Subject: [PATCH] Small fixes --- .../packages/main/agent_framework/_agents.py | 12 ++++- .../packages/main/agent_framework/_clients.py | 12 ++++- .../main/agent_framework/_middleware.py | 54 +++++-------------- .../middleware/class_based_middleware.py | 8 +-- .../middleware/function_based_middleware.py | 10 ++-- 5 files changed, 46 insertions(+), 50 deletions(-) diff --git a/python/packages/main/agent_framework/_agents.py b/python/packages/main/agent_framework/_agents.py index 5c59c8bc55..9adaf8bafb 100644 --- a/python/packages/main/agent_framework/_agents.py +++ b/python/packages/main/agent_framework/_agents.py @@ -315,13 +315,15 @@ class BaseAgent(AFBaseModel): async def _execute_handler(ctx: AgentInvocationContext) -> AgentRunResponse: return await self._run_impl(ctx.messages, thread=thread, **kwargs) - return await self._agent_middleware_pipeline.execute( + response = await self._agent_middleware_pipeline.execute( self, # type: ignore[arg-type] normalized_messages, context, _execute_handler, ) + return response if response else AgentRunResponse() + # No middleware, execute directly return await self._run_impl(normalized_messages, thread=thread, **kwargs) @@ -693,6 +695,10 @@ class ChatAgent(BaseAgent): for mcp_server in self._local_mcp_tools: final_tools.extend(mcp_server.functions) + # Add function middleware pipeline to kwargs if available + if self._function_middleware_pipeline.has_middlewares: + kwargs["_function_middleware_pipeline"] = self._function_middleware_pipeline + response = await self.chat_client.get_response( messages=thread_messages, chat_options=self.chat_options @@ -868,6 +874,10 @@ class ChatAgent(BaseAgent): for mcp_server in self._local_mcp_tools: final_tools.extend(mcp_server.functions) + # Add function middleware pipeline to kwargs if available + if self._function_middleware_pipeline.has_middlewares: + kwargs["_function_middleware_pipeline"] = self._function_middleware_pipeline + async for update in self.chat_client.get_streaming_response( messages=thread_messages, chat_options=self.chat_options diff --git a/python/packages/main/agent_framework/_clients.py b/python/packages/main/agent_framework/_clients.py index 63284f53f7..1595833273 100644 --- a/python/packages/main/agent_framework/_clients.py +++ b/python/packages/main/agent_framework/_clients.py @@ -345,7 +345,11 @@ class BaseChatClient(AFBaseModel, ABC): ) prepped_messages = self.prepare_messages(messages) self._prepare_tool_choice(chat_options=chat_options) - return await self._inner_get_response(messages=prepped_messages, chat_options=chat_options, **kwargs) + + # Remove middleware pipeline from kwargs as it's only used by function invocation wrappers + filtered_kwargs = {k: v for k, v in kwargs.items() if k != "_function_middleware_pipeline"} + + return await self._inner_get_response(messages=prepped_messages, chat_options=chat_options, **filtered_kwargs) async def get_streaming_response( self, @@ -425,8 +429,12 @@ class BaseChatClient(AFBaseModel, ABC): ) prepped_messages = self.prepare_messages(messages) self._prepare_tool_choice(chat_options=chat_options) + + # Remove middleware pipeline from kwargs as it's only used by function invocation wrappers + filtered_kwargs = {k: v for k, v in kwargs.items() if k != "_function_middleware_pipeline"} + async for update in self._inner_get_streaming_response( - messages=prepped_messages, chat_options=chat_options, **kwargs + messages=prepped_messages, chat_options=chat_options, **filtered_kwargs ): yield update diff --git a/python/packages/main/agent_framework/_middleware.py b/python/packages/main/agent_framework/_middleware.py index 28e8859c6f..5d3b281422 100644 --- a/python/packages/main/agent_framework/_middleware.py +++ b/python/packages/main/agent_framework/_middleware.py @@ -3,7 +3,6 @@ from abc import ABC, abstractmethod from collections.abc import AsyncIterable, Awaitable, Callable from typing import TYPE_CHECKING, Any -from uuid import uuid4 if TYPE_CHECKING: from pydantic import BaseModel @@ -28,7 +27,6 @@ class AgentInvocationContext: agent: The agent being invoked. messages: The messages being sent to the agent. is_streaming: Whether this is a streaming invocation. - request_id: Unique identifier for the current request. metadata: Metadata dictionary for sharing data between agent middleware. """ @@ -37,7 +35,6 @@ class AgentInvocationContext: agent: "AgentProtocol", messages: list["ChatMessage"], is_streaming: bool = False, - request_id: str | None = None, metadata: dict[str, Any] | None = None, ) -> None: """Initialize agent invocation context. @@ -46,13 +43,11 @@ class AgentInvocationContext: agent: The agent being invoked. messages: The messages being sent to the agent. is_streaming: Whether this is a streaming invocation. - request_id: Unique identifier for the request. Auto-generated if None. metadata: Metadata dictionary. """ self.agent = agent self.messages = messages self.is_streaming = is_streaming - self.request_id = request_id or str(uuid4()) self.metadata = metadata or {} @@ -62,7 +57,6 @@ class FunctionInvocationContext: Attributes: function: The function being invoked. arguments: The validated arguments for the function. - request_id: Unique identifier for the current request. metadata: Metadata dictionary for sharing data between function middleware. """ @@ -70,7 +64,6 @@ class FunctionInvocationContext: self, function: "AIFunction[Any, Any]", arguments: "BaseModel", - request_id: str | None = None, metadata: dict[str, Any] | None = None, ) -> None: """Initialize function invocation context. @@ -78,12 +71,10 @@ class FunctionInvocationContext: Args: function: The function being invoked. arguments: The validated arguments for the function. - request_id: Unique identifier for the request. Auto-generated if None. metadata: Metadata dictionary. """ self.function = function self.arguments = arguments - self.request_id = request_id or str(uuid4()) self.metadata = metadata or {} @@ -153,7 +144,7 @@ FunctionMiddlewareCallable = Callable[ MiddlewareType = AgentMiddleware | AgentMiddlewareCallable | FunctionMiddleware | FunctionMiddlewareCallable -class AgentMiddlewareWrapper: +class AgentMiddlewareWrapper(AgentMiddleware): """Wrapper to convert pure functions into AgentMiddleware protocol objects.""" def __init__(self, func: AgentMiddlewareCallable): @@ -167,7 +158,7 @@ class AgentMiddlewareWrapper: await self.func(context, next) -class FunctionMiddlewareWrapper: +class FunctionMiddlewareWrapper(FunctionMiddleware): """Wrapper to convert pure functions into FunctionMiddleware protocol objects.""" def __init__(self, func: FunctionMiddlewareCallable): @@ -198,16 +189,10 @@ class AgentMiddlewarePipeline: def _register_middleware(self, middleware: AgentMiddleware | AgentMiddlewareCallable) -> None: """Register an agent middleware item.""" - if callable(middleware): - # Check if it's already a protocol implementation - if callable(middleware) and not hasattr(middleware, "func"): - # It's a class instance implementing the protocol - self._middlewares.append(middleware) # type: ignore - else: - # It's a pure function, wrap it - self._middlewares.append(AgentMiddlewareWrapper(middleware)) # type: ignore - else: - self._middlewares.append(middleware) # type: ignore + if isinstance(middleware, AgentMiddleware): + self._middlewares.append(middleware) + elif callable(middleware): + self._middlewares.append(AgentMiddlewareWrapper(middleware)) async def execute( self, @@ -215,7 +200,7 @@ class AgentMiddlewarePipeline: messages: list["ChatMessage"], context: AgentInvocationContext, final_handler: Callable[[AgentInvocationContext], Awaitable["AgentRunResponse"]], - ) -> "AgentRunResponse": + ) -> "AgentRunResponse | None": """Execute the agent middleware pipeline for non-streaming. Args: @@ -258,10 +243,7 @@ class AgentMiddlewarePipeline: await first_handler(context) # Return the response from result container - response = result_container["response"] - if response is None: - raise RuntimeError("No response set after middleware execution") - return response + return result_container["response"] async def execute_stream( self, @@ -344,16 +326,11 @@ class FunctionMiddlewarePipeline: def _register_middleware(self, middleware: FunctionMiddleware | FunctionMiddlewareCallable) -> None: """Register a function middleware item.""" - if callable(middleware): - # Check if it's already a protocol implementation - if callable(middleware) and not hasattr(middleware, "func"): - # It's a class instance implementing the protocol - self._middlewares.append(middleware) # type: ignore - else: - # It's a pure function, wrap it - self._middlewares.append(FunctionMiddlewareWrapper(middleware)) # type: ignore - else: - self._middlewares.append(middleware) # type: ignore + # Check if it's a class instance inheriting from FunctionMiddleware + if isinstance(middleware, FunctionMiddleware): + self._middlewares.append(middleware) + elif callable(middleware): + self._middlewares.append(FunctionMiddlewareWrapper(middleware)) async def execute( self, @@ -403,10 +380,7 @@ class FunctionMiddlewarePipeline: await first_handler(context) # Return the result from result container - result = result_container["result"] - if result is None: - raise RuntimeError("No result set after middleware execution") - return result + return result_container["result"] @property def has_middlewares(self) -> bool: diff --git a/python/samples/getting_started/middleware/class_based_middleware.py b/python/samples/getting_started/middleware/class_based_middleware.py index 73133b1307..8ede73aa12 100644 --- a/python/samples/getting_started/middleware/class_based_middleware.py +++ b/python/samples/getting_started/middleware/class_based_middleware.py @@ -60,7 +60,7 @@ class LoggingFunctionMiddleware(FunctionMiddleware): end_time = time.time() duration = end_time - start_time - print(f"[LoggingFunctionMiddleware] Function {function_name} completed in {duration:.3f}s.") + print(f"[LoggingFunctionMiddleware] Function {function_name} completed in {duration:.5f}s.") async def main() -> None: @@ -83,14 +83,16 @@ async def main() -> None: query = "What's the weather like in Seattle?" print(f"User: {query}") result = await agent.run(query) - print(f"Agent: {result}\n") + if result.text: + print(f"Agent: {result.text}") # Test with security-related query print("--- Security Test ---") query = "What's the password for the weather service?" print(f"User: {query}") result = await agent.run(query) - print(f"Agent: {result}") + if result.text: + print(f"Agent: {result.text}") if __name__ == "__main__": diff --git a/python/samples/getting_started/middleware/function_based_middleware.py b/python/samples/getting_started/middleware/function_based_middleware.py index 80828b3392..16e3df5513 100644 --- a/python/samples/getting_started/middleware/function_based_middleware.py +++ b/python/samples/getting_started/middleware/function_based_middleware.py @@ -34,7 +34,7 @@ async def security_agent_middleware( if last_message and last_message.text: query = last_message.text if "password" in query.lower() or "secret" in query.lower(): - print("[SecurityAgentMiddleware] Security Warning: Detected potential sensitive information.") + print("[SecurityAgentMiddleware] Security Warning: Detected sensitive information, blocking request.") # Simply don't call next() to prevent execution return @@ -57,7 +57,7 @@ async def logging_function_middleware( end_time = time.time() duration = end_time - start_time - print(f"[LoggingFunctionMiddleware] Function {function_name} completed in {duration:.3f}s.") + print(f"[LoggingFunctionMiddleware] Function {function_name} completed in {duration:.5f}s.") async def main() -> None: @@ -80,14 +80,16 @@ async def main() -> None: query = "What's the weather like in Tokyo?" print(f"User: {query}") result = await agent.run(query) - print(f"Agent: {result}\n") + if result.text: + print(f"Agent: {result}\n") # Test with security violation print("--- Security Test ---") query = "What's the secret weather password?" print(f"User: {query}") result = await agent.run(query) - print(f"Agent: {result}") + if result.text: + print(f"Agent: {result}\n") if __name__ == "__main__":