test(policy_engine): cover streaming pipeline gate branches

Adds regression tests for the modify_response block on the Anthropic route,
the gate with no iterator overrides, and the per-chunk hook skipping
pipeline-managed guardrails. Corrects the gate docstring: an allow releases
the chunks as the endpoint translation left them, not verbatim
This commit is contained in:
mateo-berri 2026-08-29 14:07:34 -07:00
parent 1bed9bae43
commit d51198fdeb
2 changed files with 115 additions and 2 deletions

View file

@ -3468,8 +3468,11 @@ class ProxyLogging:
pipeline allows it), then runs each pipeline's steps against the
assembled output through the endpoint guardrail translation, the same
machinery flat post_call guardrails use at end of stream. An allow
releases the buffered chunks verbatim; a block or modify_response
terminates with the translation's block chunks or the raised error.
releases the buffered chunks as that machinery left them (the
Responses and A2A translations write guardrail output back into the
final chunk, exactly as they do for flat guardrails); a block or
modify_response terminates with the translation's block chunks or the
raised error.
"""
buffered: Final[list[object]] = [] # mutable-ok: accumulates the stream before the pipeline verdict
async for item in response:

View file

@ -10,6 +10,7 @@ Covers ``_should_use_guardrail_load_balancing``, ``_execute_guardrail_hook``,
from __future__ import annotations
import asyncio
import json
from typing import Any, Dict, List
from unittest.mock import AsyncMock, MagicMock, patch
@ -1557,3 +1558,112 @@ async def test_streaming_iterator_hook_pipeline_withholds_unresolvable_response_
assert info.value.code == "500"
assert "withheld" in info.value.message
assert seen.get("count") is None
def _anthropic_sse_chunks() -> List[bytes]:
events = [
("message_start", {"type": "message_start", "message": {"id": "msg_1", "type": "message", "role": "assistant", "model": "m", "content": [], "stop_reason": None, "usage": {"input_tokens": 1, "output_tokens": 0}}}),
("content_block_start", {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}),
("content_block_delta", {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "hello world"}}),
("content_block_stop", {"type": "content_block_stop", "index": 0}),
("message_delta", {"type": "message_delta", "delta": {"stop_reason": "end_turn", "stop_sequence": None}, "usage": {"output_tokens": 2}}),
("message_stop", {"type": "message_stop"}),
]
return [f"event: {name}\ndata: {json.dumps(payload)}\n\n".encode() for name, payload in events]
@pytest.mark.asyncio
async def test_streaming_iterator_hook_pipeline_modify_response_emits_translated_block(
proxy_logging, make_user_api_key_auth, monkeypatch
):
seen: Dict[str, Any] = {}
monkeypatch.setattr(litellm, "callbacks", [_unified_stream_guardrail(seen, block=True)])
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
pipeline = GuardrailPipeline(
mode="post_call",
steps=[
PipelineStep(
guardrail="gr-post",
on_pass="allow",
on_fail="modify_response",
modify_response_message="content policy block",
)
],
)
data = _post_call_pipeline_data(stream=True)
data["metadata"]["_guardrail_pipelines"] = [("response-governance", pipeline)]
chunks = _anthropic_sse_chunks()
delivered = [
item
async for item in proxy_logging.async_post_call_streaming_iterator_hook(
user_api_key_dict=make_user_api_key_auth(request_route="/v1/messages"),
response=_async_chunk_iter(chunks),
request_data=data,
)
]
raw = b"".join(delivered).decode()
assert seen["count"] == 1
assert "content policy block" in raw
assert "hello world" not in raw
assert not any(item is chunk for item in delivered for chunk in chunks)
@pytest.mark.asyncio
async def test_streaming_iterator_hook_pipeline_gates_without_iterator_overrides(
proxy_logging, make_user_api_key_auth, monkeypatch
):
monkeypatch.setattr(litellm, "callbacks", [])
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
data = _post_call_pipeline_data(stream=True)
delivered: List[Any] = []
async def _drain() -> None:
async for item in proxy_logging.async_post_call_streaming_iterator_hook(
user_api_key_dict=make_user_api_key_auth(request_route="/v1/chat/completions"),
response=_async_chunk_iter(_stream_chunks()),
request_data=data,
):
delivered.append(item)
with pytest.raises(HTTPException) as info:
await _drain()
assert delivered == []
assert info.value.status_code == 400
assert info.value.detail["error"]["pipeline_context"]["step_results"] == [
{"guardrail": "gr-post", "outcome": "error", "action": "block"}
]
@pytest.mark.asyncio
async def test_per_chunk_streaming_hook_skips_pipeline_managed_guardrail(
proxy_logging, make_user_api_key_auth, monkeypatch
):
seen: Dict[str, Any] = {}
class RecordingGuardrail(CustomGuardrail):
async def async_post_call_streaming_hook(self, user_api_key_dict, response):
seen[self.guardrail_name] = seen.get(self.guardrail_name, 0) + 1
return None
managed = RecordingGuardrail(
guardrail_name="gr-post", event_hook=GuardrailEventHooks.post_call, default_on=True
)
free = RecordingGuardrail(
guardrail_name="gr-free", event_hook=GuardrailEventHooks.post_call, default_on=True
)
monkeypatch.setattr(litellm, "callbacks", [managed, free])
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", None, raising=False)
data = _post_call_pipeline_data(stream=True)
result = await proxy_logging.async_post_call_streaming_hook(
data=data,
response=_stream_chunks()[0],
user_api_key_dict=make_user_api_key_auth(request_route="/v1/chat/completions"),
)
assert result is not None
assert seen.get("gr-post") is None
assert seen["gr-free"] == 1