Removed context parameter from call_next (#3829)

This commit is contained in:
Dmytro Struk
2026-02-11 02:47:41 -08:00
committed by GitHub
Unverified
parent 38f22ef006
commit 1fdc4be88d
29 changed files with 451 additions and 583 deletions
@@ -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)