mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-10 22:41:41 +00:00
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:
parent
1bed9bae43
commit
d51198fdeb
2 changed files with 115 additions and 2 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue