mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Removed context parameter from call_next (#3829)
This commit is contained in:
committed by
GitHub
Unverified
parent
38f22ef006
commit
1fdc4be88d
@@ -65,7 +65,7 @@ class PurviewPolicyMiddleware(AgentMiddleware):
|
||||
async def process(
|
||||
self,
|
||||
context: AgentContext,
|
||||
call_next: Callable[[AgentContext], Awaitable[None]],
|
||||
call_next: Callable[[], Awaitable[None]],
|
||||
) -> None: # type: ignore[override]
|
||||
resolved_user_id: str | None = None
|
||||
try:
|
||||
@@ -92,7 +92,7 @@ class PurviewPolicyMiddleware(AgentMiddleware):
|
||||
if not self._settings.ignore_exceptions:
|
||||
raise
|
||||
|
||||
await call_next(context)
|
||||
await call_next()
|
||||
|
||||
try:
|
||||
# Post (response) check only if we have a normal AgentResponse
|
||||
@@ -162,7 +162,7 @@ class PurviewChatPolicyMiddleware(ChatMiddleware):
|
||||
async def process(
|
||||
self,
|
||||
context: ChatContext,
|
||||
call_next: Callable[[ChatContext], Awaitable[None]],
|
||||
call_next: Callable[[], Awaitable[None]],
|
||||
) -> None: # type: ignore[override]
|
||||
resolved_user_id: str | None = None
|
||||
try:
|
||||
@@ -187,7 +187,7 @@ class PurviewChatPolicyMiddleware(ChatMiddleware):
|
||||
if not self._settings.ignore_exceptions:
|
||||
raise
|
||||
|
||||
await call_next(context)
|
||||
await call_next()
|
||||
|
||||
try:
|
||||
# Post (response) evaluation only if non-streaming and we have messages result shape
|
||||
|
||||
@@ -49,7 +49,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
next_called = False
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
nonlocal next_called
|
||||
next_called = True
|
||||
|
||||
@@ -57,7 +57,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
def __init__(self):
|
||||
self.messages = [Message(role="assistant", text="Hi there")]
|
||||
|
||||
ctx.result = Result()
|
||||
chat_context.result = Result()
|
||||
|
||||
await middleware.process(chat_context, mock_next)
|
||||
assert next_called
|
||||
@@ -67,7 +67,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
async def test_blocks_prompt(self, middleware: PurviewChatPolicyMiddleware, chat_context: ChatContext) -> None:
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(True, "user-123")):
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None: # should not run
|
||||
async def mock_next() -> None: # should not run
|
||||
raise AssertionError("next should not be called when prompt blocked")
|
||||
|
||||
with pytest.raises(MiddlewareTermination):
|
||||
@@ -88,12 +88,12 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=side_effect):
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
class Result:
|
||||
def __init__(self):
|
||||
self.messages = [Message(role="assistant", text="Sensitive output")] # pragma: no cover
|
||||
|
||||
ctx.result = Result()
|
||||
chat_context.result = Result()
|
||||
|
||||
await middleware.process(chat_context, mock_next)
|
||||
assert call_state["count"] == 2
|
||||
@@ -114,8 +114,8 @@ class TestPurviewChatPolicyMiddleware:
|
||||
)
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
ctx.result = MagicMock()
|
||||
async def mock_next() -> None:
|
||||
streaming_context.result = MagicMock()
|
||||
|
||||
await middleware.process(streaming_context, mock_next)
|
||||
assert mock_proc.call_count == 1
|
||||
@@ -138,10 +138,10 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [Message(role="assistant", text="Response")]
|
||||
ctx.result = result
|
||||
chat_context.result = result
|
||||
|
||||
await middleware.process(chat_context, mock_next)
|
||||
|
||||
@@ -162,10 +162,10 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [Message(role="assistant", text="Response")]
|
||||
ctx.result = result
|
||||
chat_context.result = result
|
||||
|
||||
await middleware.process(chat_context, mock_next)
|
||||
|
||||
@@ -194,7 +194,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
raise AssertionError("next should not be called")
|
||||
|
||||
# Should raise the exception
|
||||
@@ -224,10 +224,10 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=side_effect):
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [Message(role="assistant", text="OK")]
|
||||
ctx.result = result
|
||||
context.result = result
|
||||
|
||||
with pytest.raises(PurviewPaymentRequiredError):
|
||||
await middleware.process(context, mock_next)
|
||||
@@ -249,7 +249,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [Message(role="assistant", text="Response")]
|
||||
context.result = result
|
||||
@@ -265,9 +265,9 @@ class TestPurviewChatPolicyMiddleware:
|
||||
"""Test middleware handles result that doesn't have messages attribute."""
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")):
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
# Set result to something without messages attribute
|
||||
ctx.result = "Some string result"
|
||||
chat_context.result = "Some string result"
|
||||
|
||||
await middleware.process(chat_context, mock_next)
|
||||
|
||||
@@ -289,7 +289,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [Message(role="assistant", text="Response")]
|
||||
context.result = result
|
||||
@@ -313,7 +313,7 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=ValueError("boom")):
|
||||
|
||||
async def mock_next(_: ChatContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
raise AssertionError("next should not be called")
|
||||
|
||||
with pytest.raises(ValueError, match="boom"):
|
||||
@@ -342,10 +342,10 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=side_effect):
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [Message(role="assistant", text="OK")]
|
||||
ctx.result = result
|
||||
context.result = result
|
||||
|
||||
with pytest.raises(ValueError, match="post"):
|
||||
await middleware.process(context, mock_next)
|
||||
@@ -361,10 +361,10 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [Message(role="assistant", text="Hi")]
|
||||
ctx.result = result
|
||||
context.result = result
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -382,10 +382,10 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [Message(role="assistant", text="Hi")]
|
||||
ctx.result = result
|
||||
context.result = result
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -401,10 +401,10 @@ class TestPurviewChatPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: ChatContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
result = MagicMock()
|
||||
result.messages = [Message(role="assistant", text="Response")]
|
||||
ctx.result = result
|
||||
context.result = result
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
|
||||
@@ -55,10 +55,10 @@ class TestPurviewPolicyMiddleware:
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")):
|
||||
next_called = False
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
nonlocal next_called
|
||||
next_called = True
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="I'm good, thanks!")])
|
||||
context.result = AgentResponse(messages=[Message(role="assistant", text="I'm good, thanks!")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -74,7 +74,7 @@ class TestPurviewPolicyMiddleware:
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(True, "user-123")):
|
||||
next_called = False
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
nonlocal next_called
|
||||
next_called = True
|
||||
|
||||
@@ -101,8 +101,8 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(
|
||||
async def mock_next() -> None:
|
||||
context.result = AgentResponse(
|
||||
messages=[Message(role="assistant", text="Here's some sensitive information")]
|
||||
)
|
||||
|
||||
@@ -125,8 +125,8 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")):
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = "Some non-standard result"
|
||||
async def mock_next() -> None:
|
||||
context.result = "Some non-standard result"
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -142,8 +142,8 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_process:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
|
||||
async def mock_next() -> None:
|
||||
context.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -160,8 +160,8 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="streaming")])
|
||||
async def mock_next() -> None:
|
||||
context.result = AgentResponse(messages=[Message(role="assistant", text="streaming")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -181,7 +181,7 @@ class TestPurviewPolicyMiddleware:
|
||||
side_effect=PurviewPaymentRequiredError("Payment required"),
|
||||
):
|
||||
|
||||
async def mock_next(_: AgentContext) -> None:
|
||||
async def mock_next() -> None:
|
||||
raise AssertionError("next should not be called")
|
||||
|
||||
with pytest.raises(PurviewPaymentRequiredError):
|
||||
@@ -206,8 +206,8 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=side_effect):
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="OK")])
|
||||
async def mock_next() -> None:
|
||||
context.result = AgentResponse(messages=[Message(role="assistant", text="OK")])
|
||||
|
||||
with pytest.raises(PurviewPaymentRequiredError):
|
||||
await middleware.process(context, mock_next)
|
||||
@@ -231,8 +231,8 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=side_effect):
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="OK")])
|
||||
async def mock_next() -> None:
|
||||
context.result = AgentResponse(messages=[Message(role="assistant", text="OK")])
|
||||
|
||||
with pytest.raises(ValueError, match="Post-check blew up"):
|
||||
await middleware.process(context, mock_next)
|
||||
@@ -250,8 +250,8 @@ class TestPurviewPolicyMiddleware:
|
||||
middleware._processor, "process_messages", side_effect=Exception("Pre-check error")
|
||||
) as mock_process:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
|
||||
async def mock_next() -> None:
|
||||
context.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -280,8 +280,8 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
|
||||
async def mock_next() -> None:
|
||||
context.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -306,8 +306,8 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
|
||||
|
||||
async def mock_next(ctx):
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
|
||||
async def mock_next():
|
||||
context.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
|
||||
|
||||
# Should not raise, just log
|
||||
await middleware.process(context, mock_next)
|
||||
@@ -330,7 +330,7 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", side_effect=mock_process_messages):
|
||||
|
||||
async def mock_next(ctx):
|
||||
async def mock_next():
|
||||
pass
|
||||
|
||||
# Should raise the exception
|
||||
@@ -346,8 +346,8 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
|
||||
async def mock_next() -> None:
|
||||
context.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -364,8 +364,8 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
|
||||
async def mock_next() -> None:
|
||||
context.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -383,8 +383,8 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
|
||||
async def mock_next() -> None:
|
||||
context.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -399,8 +399,8 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
|
||||
async def mock_next() -> None:
|
||||
context.result = AgentResponse(messages=[Message(role="assistant", text="Hi")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
@@ -416,8 +416,8 @@ class TestPurviewPolicyMiddleware:
|
||||
|
||||
with patch.object(middleware._processor, "process_messages", return_value=(False, "user-123")) as mock_proc:
|
||||
|
||||
async def mock_next(ctx: AgentContext) -> None:
|
||||
ctx.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
|
||||
async def mock_next() -> None:
|
||||
context.result = AgentResponse(messages=[Message(role="assistant", text="Response")])
|
||||
|
||||
await middleware.process(context, mock_next)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user