mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
fix(guardrails): address review nits on generic_guardrail_api streaming
Validate streaming_sampling_rate >= 1 in the constructor and Pydantic optional_params (ge=1), and add /v1/responses streaming coverage through the unified post-call hook so Responses API usage is exercised alongside chat completions.
This commit is contained in:
parent
3a20cd7479
commit
572acf0037
3 changed files with 160 additions and 2 deletions
|
|
@ -233,6 +233,11 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
if streaming_end_of_stream_only is None
|
||||
else streaming_end_of_stream_only
|
||||
)
|
||||
if streaming_sampling_rate is not None and streaming_sampling_rate < 1:
|
||||
raise ValueError(
|
||||
"streaming_sampling_rate must be >= 1 "
|
||||
f"(got {streaming_sampling_rate})"
|
||||
)
|
||||
self.streaming_sampling_rate: int = (
|
||||
5 if streaming_sampling_rate is None else streaming_sampling_rate
|
||||
)
|
||||
|
|
|
|||
|
|
@ -54,11 +54,12 @@ class GenericGuardrailAPIOptionalParams(BaseModel):
|
|||
|
||||
streaming_sampling_rate: Optional[int] = Field(
|
||||
default=None,
|
||||
ge=1,
|
||||
description=(
|
||||
"When streaming_end_of_stream_only is False, the guardrail runs every Nth "
|
||||
"streamed chunk. Ignored when streaming_end_of_stream_only is True. "
|
||||
"Defaults to 5 in GenericGuardrailAPI.__init__ when None so unset "
|
||||
"optional_params does not shadow top-level litellm_params."
|
||||
"Must be >= 1 when set. Defaults to 5 in GenericGuardrailAPI.__init__ "
|
||||
"when None so unset optional_params does not shadow top-level litellm_params."
|
||||
),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -1068,6 +1068,42 @@ def _mock_guardrail_post_response(action: str = "NONE", texts=None, blocked_reas
|
|||
return mock_response
|
||||
|
||||
|
||||
def _make_responses_stream_events(text: str):
|
||||
"""Minimal /v1/responses SSE event sequence ending in response.completed."""
|
||||
return (
|
||||
{"type": "response.created", "response": {"id": "resp_test"}},
|
||||
{
|
||||
"type": "response.output_item.added",
|
||||
"item": {"type": "message", "id": "msg_test"},
|
||||
},
|
||||
{
|
||||
"type": "response.content_part.added",
|
||||
"part": {"type": "output_text", "text": ""},
|
||||
},
|
||||
{"type": "response.output_text.delta", "delta": text},
|
||||
{
|
||||
"type": "response.output_text.done",
|
||||
"text": text,
|
||||
},
|
||||
{
|
||||
"type": "response.completed",
|
||||
"response": {
|
||||
"id": "resp_test",
|
||||
"output": [
|
||||
{
|
||||
"type": "message",
|
||||
"id": "msg_test",
|
||||
"status": "completed",
|
||||
"role": "assistant",
|
||||
"content": [{"type": "output_text", "text": text}],
|
||||
}
|
||||
],
|
||||
"status": "completed",
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class TestGenericGuardrailAPIStreamingConfig:
|
||||
"""Streaming knobs on GenericGuardrailAPI and initialize_guardrail plumbing."""
|
||||
|
||||
|
|
@ -1091,6 +1127,26 @@ class TestGenericGuardrailAPIStreamingConfig:
|
|||
assert guardrail.streaming_end_of_stream_only is True
|
||||
assert guardrail.streaming_sampling_rate == 2
|
||||
|
||||
@pytest.mark.parametrize("invalid_rate", [0, -1, -5])
|
||||
def test_streaming_sampling_rate_rejects_non_positive(self, invalid_rate):
|
||||
with pytest.raises(ValueError, match="streaming_sampling_rate must be >= 1"):
|
||||
GenericGuardrailAPI(
|
||||
api_base="https://api.test.guardrail.com",
|
||||
guardrail_name="test-generic-guardrail",
|
||||
event_hook="post_call",
|
||||
streaming_sampling_rate=invalid_rate,
|
||||
)
|
||||
|
||||
def test_optional_params_streaming_sampling_rate_ge_one(self):
|
||||
from pydantic import ValidationError
|
||||
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
|
||||
GenericGuardrailAPIOptionalParams,
|
||||
)
|
||||
|
||||
with pytest.raises(ValidationError):
|
||||
GenericGuardrailAPIOptionalParams(streaming_sampling_rate=0)
|
||||
|
||||
def test_get_config_model(self):
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
|
||||
GenericGuardrailAPIConfigModel,
|
||||
|
|
@ -1522,3 +1578,99 @@ class TestGenericGuardrailAPIStreamingViaUnified:
|
|||
chunks_received += 1
|
||||
|
||||
assert chunks_received == 3
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_streaming_end_of_stream_only_calls_guardrail_once(self):
|
||||
"""/v1/responses path through unified hook; end-of-stream-only = one call."""
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
|
||||
guardrail = GenericGuardrailAPI(
|
||||
api_base="https://api.test.guardrail.com",
|
||||
guardrail_name="test-generic-guardrail",
|
||||
event_hook="post_call",
|
||||
streaming_end_of_stream_only=True,
|
||||
)
|
||||
unified_guardrail = UnifiedLLMGuardrails()
|
||||
|
||||
async def mock_responses_stream():
|
||||
for event in _make_responses_stream_events("Hello world"):
|
||||
yield event
|
||||
|
||||
mock_post = AsyncMock(
|
||||
return_value=_mock_guardrail_post_response(
|
||||
action="NONE", texts=["Hello world"]
|
||||
)
|
||||
)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", mock_post):
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test", request_route="/v1/responses"
|
||||
)
|
||||
request_data = {
|
||||
"input": "hi",
|
||||
"guardrail_to_apply": guardrail,
|
||||
"metadata": {"guardrails": ["test-generic-guardrail"]},
|
||||
}
|
||||
|
||||
events_received = 0
|
||||
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=mock_responses_stream(),
|
||||
request_data=request_data,
|
||||
):
|
||||
events_received += 1
|
||||
|
||||
assert events_received == 6
|
||||
assert mock_post.await_count == 1, (
|
||||
f"Expected exactly one guardrail call at end of /v1/responses stream, "
|
||||
f"got {mock_post.await_count}"
|
||||
)
|
||||
assert mock_post.await_args.kwargs["json"]["input_type"] == "response"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_responses_api_streaming_blocked_raises(self):
|
||||
"""Mid-stream BLOCKED on /v1/responses surfaces GuardrailRaisedException."""
|
||||
from litellm.exceptions import GuardrailRaisedException
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
UnifiedLLMGuardrails,
|
||||
)
|
||||
|
||||
guardrail = GenericGuardrailAPI(
|
||||
api_base="https://api.test.guardrail.com",
|
||||
guardrail_name="test-generic-guardrail",
|
||||
event_hook="post_call",
|
||||
streaming_sampling_rate=1,
|
||||
)
|
||||
unified_guardrail = UnifiedLLMGuardrails()
|
||||
|
||||
async def mock_responses_stream():
|
||||
for event in _make_responses_stream_events("blocked content"):
|
||||
yield event
|
||||
|
||||
mock_post = AsyncMock(
|
||||
return_value=_mock_guardrail_post_response(
|
||||
action="BLOCKED", blocked_reason="Responses content not allowed"
|
||||
)
|
||||
)
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", mock_post):
|
||||
user_api_key_dict = UserAPIKeyAuth(
|
||||
api_key="test", request_route="/v1/responses"
|
||||
)
|
||||
request_data = {
|
||||
"input": "hi",
|
||||
"guardrail_to_apply": guardrail,
|
||||
"metadata": {"guardrails": ["test-generic-guardrail"]},
|
||||
}
|
||||
|
||||
with pytest.raises(GuardrailRaisedException) as exc_info:
|
||||
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
response=mock_responses_stream(),
|
||||
request_data=request_data,
|
||||
):
|
||||
pass
|
||||
|
||||
assert "Responses content not allowed" in str(exc_info.value)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue