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