feat(guardrails): expose streaming knobs on generic_guardrail_api

Wire streaming_end_of_stream_only and streaming_sampling_rate through
optional params, initialize_guardrail, and get_config_model so the
generic guardrail API participates in UnifiedLLMGuardrails streaming
checks with configurable cadence and end-of-stream-only mode.
This commit is contained in:
Marton Schneider 2026-06-21 15:17:25 +02:00
parent 84c1414aef
commit 53138180a9
4 changed files with 495 additions and 2 deletions

View file

@ -1,4 +1,4 @@
from typing import TYPE_CHECKING
from typing import TYPE_CHECKING, Any, Optional
from litellm.types.guardrails import SupportedGuardrailIntegrations
@ -8,9 +8,21 @@ if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
def _get_config_value(
litellm_params: Any, optional_params: Any, attribute_name: str
) -> Optional[Any]:
if optional_params is not None:
value = getattr(optional_params, attribute_name, None)
if value is not None:
return value
return getattr(litellm_params, attribute_name, None)
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
import litellm
optional_params = getattr(litellm_params, "optional_params", None)
_generic_guardrail_api_callback = GenericGuardrailAPI(
api_base=litellm_params.api_base,
api_key=litellm_params.api_key,
@ -25,6 +37,12 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
streaming_end_of_stream_only=_get_config_value(
litellm_params, optional_params, "streaming_end_of_stream_only"
),
streaming_sampling_rate=_get_config_value(
litellm_params, optional_params, "streaming_sampling_rate"
),
)
litellm.logging_callback_manager.add_litellm_callback(

View file

@ -7,7 +7,7 @@
import fnmatch
import os
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Set
from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Set, Type
import httpx
@ -32,6 +32,7 @@ from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
GUARDRAIL_NAME = "generic_guardrail_api"
@ -188,6 +189,8 @@ class GenericGuardrailAPI(CustomGuardrail):
additional_provider_specific_params: Optional[Dict[str, Any]] = None,
unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed",
extra_headers: Optional[list] = None,
streaming_end_of_stream_only: Optional[bool] = None,
streaming_sampling_rate: Optional[int] = None,
**kwargs,
):
self.async_handler = get_async_httpx_client(
@ -223,6 +226,17 @@ class GenericGuardrailAPI(CustomGuardrail):
unreachable_fallback
)
# Read by UnifiedLLMGuardrails.async_post_call_streaming_iterator_hook
# via getattr(guardrail_to_apply, "streaming_*", default).
self.streaming_end_of_stream_only: bool = (
False
if streaming_end_of_stream_only is None
else streaming_end_of_stream_only
)
self.streaming_sampling_rate: int = (
5 if streaming_sampling_rate is None else streaming_sampling_rate
)
# Set supported event hooks
if "supported_event_hooks" not in kwargs:
kwargs["supported_event_hooks"] = [
@ -511,3 +525,11 @@ class GenericGuardrailAPI(CustomGuardrail):
return self._handle_guardrail_request_error(
e, inputs, input_type, logging_obj, is_unreachable=False
)
@staticmethod
def get_config_model() -> Optional[Type["GuardrailConfigModel"]]:
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPIConfigModel,
)
return GenericGuardrailAPIConfigModel

View file

@ -39,6 +39,25 @@ class GenericGuardrailAPIOptionalParams(BaseModel):
),
)
streaming_end_of_stream_only: Optional[bool] = Field(
default=False,
description=(
"If False (default), the guardrail runs on sampled chunks during the stream "
"at the cadence set by streaming_sampling_rate, and an in-flight BLOCKED "
"stops further chunks from streaming. If True, the guardrail runs once at "
"end of stream over the assembled response; lower cost and latency, but "
"flagged content has already streamed to the client before the terminal block."
),
)
streaming_sampling_rate: Optional[int] = Field(
default=5,
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."
),
)
class GenericGuardrailAPIConfigModel(
GuardrailConfigModel[GenericGuardrailAPIOptionalParams],

View file

@ -1021,3 +1021,437 @@ class TestMultimodalSupport:
call_args = mock_post.call_args
json_payload = call_args.kwargs["json"]
assert isinstance(json_payload["structured_messages"], list)
def _make_stream_chunk(content: str, finish_reason=None):
"""Build a real ModelResponseStream so the handler's isinstance checks pass."""
import litellm
from litellm.types.utils import Delta, ModelResponseStream
return ModelResponseStream(
model="gpt-4",
choices=[
litellm.StreamingChoices(
index=0,
delta=Delta(role="assistant", content=content),
finish_reason=finish_reason,
)
],
)
def _make_assembled_model_response(content: str) -> ModelResponse:
import litellm
return ModelResponse(
id="mock-response",
model="gpt-4",
choices=[
litellm.Choices(
index=0,
message=litellm.Message(role="assistant", content=content),
finish_reason="stop",
)
],
)
def _mock_guardrail_post_response(action: str = "NONE", texts=None, blocked_reason=None):
mock_response = MagicMock()
payload = {"action": action}
if texts is not None:
payload["texts"] = texts
if blocked_reason is not None:
payload["blocked_reason"] = blocked_reason
mock_response.json.return_value = payload
mock_response.raise_for_status = MagicMock()
return mock_response
class TestGenericGuardrailAPIStreamingConfig:
"""Streaming knobs on GenericGuardrailAPI and initialize_guardrail plumbing."""
def test_streaming_defaults(self):
guardrail = GenericGuardrailAPI(
api_base="https://api.test.guardrail.com",
guardrail_name="test-generic-guardrail",
event_hook="post_call",
)
assert guardrail.streaming_end_of_stream_only is False
assert guardrail.streaming_sampling_rate == 5
def test_streaming_overrides(self):
guardrail = GenericGuardrailAPI(
api_base="https://api.test.guardrail.com",
guardrail_name="test-generic-guardrail",
event_hook="post_call",
streaming_end_of_stream_only=True,
streaming_sampling_rate=2,
)
assert guardrail.streaming_end_of_stream_only is True
assert guardrail.streaming_sampling_rate == 2
def test_get_config_model(self):
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPIConfigModel,
)
assert GenericGuardrailAPI.get_config_model() is GenericGuardrailAPIConfigModel
def test_initialize_guardrail_forwards_streaming_flags(self):
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
initialize_guardrail,
)
from litellm.types.guardrails import LitellmParams
litellm_params = LitellmParams(
guardrail="generic_guardrail_api",
mode="post_call",
api_base="https://api.test.guardrail.com",
default_on=False,
)
# LitellmParams uses extra="allow" on the base; set streaming knobs dynamically
litellm_params.streaming_end_of_stream_only = False # type: ignore[attr-defined]
litellm_params.streaming_sampling_rate = 3 # type: ignore[attr-defined]
guardrail_config = {"guardrail_name": "test-generic-streaming"}
with patch(
"litellm.logging_callback_manager.add_litellm_callback"
):
guardrail = initialize_guardrail(litellm_params, guardrail_config)
assert guardrail.streaming_end_of_stream_only is False
assert guardrail.streaming_sampling_rate == 3
class TestGenericGuardrailAPIStreamingViaUnified:
"""Streaming output checks routed through UnifiedLLMGuardrails."""
@pytest.mark.asyncio
async def test_streaming_safe_content_yields_all_chunks(self):
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",
)
unified_guardrail = UnifiedLLMGuardrails()
async def mock_stream():
chunks_data = ["Hello", " ", "world", "!", " Goodbye"]
for i, content in enumerate(chunks_data):
yield _make_stream_chunk(
content,
finish_reason="stop" if i == len(chunks_data) - 1 else None,
)
mock_post = AsyncMock(
return_value=_mock_guardrail_post_response(
action="NONE", texts=["Hello world! Goodbye"]
)
)
with (
patch.object(guardrail.async_handler, "post", mock_post),
patch(
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
return_value=_make_assembled_model_response("Hello world! Goodbye"),
),
):
user_api_key_dict = UserAPIKeyAuth(
api_key="test", request_route="/chat/completions"
)
request_data = {
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["test-generic-guardrail"]},
}
chunks_received = 0
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=mock_stream(),
request_data=request_data,
):
chunks_received += 1
assert chunks_received == 5
assert mock_post.await_count >= 1
@pytest.mark.asyncio
async def test_streaming_blocked_content_raises(self):
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_stream():
chunks_data = ["Hello", " ishaan", " here"]
for i, content in enumerate(chunks_data):
yield _make_stream_chunk(
content,
finish_reason="stop" if i == len(chunks_data) - 1 else None,
)
mock_post = AsyncMock(
return_value=_mock_guardrail_post_response(
action="BLOCKED", blocked_reason="Ishaan is not allowed"
)
)
with (
patch.object(guardrail.async_handler, "post", mock_post),
patch(
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
return_value=_make_assembled_model_response("Hello ishaan here"),
),
):
user_api_key_dict = UserAPIKeyAuth(
api_key="test", request_route="/chat/completions"
)
request_data = {
"messages": [{"role": "user", "content": "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_stream(),
request_data=request_data,
):
pass
assert "Ishaan is not allowed" in str(exc_info.value)
@pytest.mark.asyncio
async def test_streaming_default_uses_sampled_cadence(self):
"""Default samples every 5th chunk + final pass: 10 chunks → calls at 5, 10, and final = 3."""
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",
)
unified_guardrail = UnifiedLLMGuardrails()
async def mock_stream():
chunks_data = ["A", "B", "C", "D", "E", "F", "G", "H", "I", "J"]
for i, content in enumerate(chunks_data):
yield _make_stream_chunk(
content,
finish_reason="stop" if i == len(chunks_data) - 1 else None,
)
mock_post = AsyncMock(
return_value=_mock_guardrail_post_response(
action="NONE", texts=["ABCDEFGHIJ"]
)
)
with (
patch.object(guardrail.async_handler, "post", mock_post),
patch(
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
return_value=_make_assembled_model_response("ABCDEFGHIJ"),
),
):
user_api_key_dict = UserAPIKeyAuth(
api_key="test", request_route="/chat/completions"
)
request_data = {
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["test-generic-guardrail"]},
}
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=mock_stream(),
request_data=request_data,
):
pass
assert mock_post.await_count == 3, (
f"Expected 3 guardrail calls (2 sampled at chunks 5 / 10 + 1 final), "
f"got {mock_post.await_count}"
)
for call in mock_post.await_args_list:
assert call.kwargs["json"]["input_type"] == "response"
@pytest.mark.asyncio
async def test_streaming_end_of_stream_only_calls_guardrail_once(self):
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_stream():
chunks_data = ["A", "B", "C", "D", "E", "F", "G", "H", "I", "J"]
for i, content in enumerate(chunks_data):
yield _make_stream_chunk(
content,
finish_reason="stop" if i == len(chunks_data) - 1 else None,
)
mock_post = AsyncMock(
return_value=_mock_guardrail_post_response(
action="NONE", texts=["ABCDEFGHIJ"]
)
)
with (
patch.object(guardrail.async_handler, "post", mock_post),
patch(
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
return_value=_make_assembled_model_response("ABCDEFGHIJ"),
),
):
user_api_key_dict = UserAPIKeyAuth(
api_key="test", request_route="/chat/completions"
)
request_data = {
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["test-generic-guardrail"]},
}
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=mock_stream(),
request_data=request_data,
):
pass
assert mock_post.await_count == 1, (
f"Expected exactly one guardrail call at end of stream, "
f"got {mock_post.await_count}"
)
@pytest.mark.asyncio
async def test_streaming_sampling_rate_override(self):
"""sampling_rate=2 on 6 chunks → in-stream at 2,4,6 plus final = 4 calls."""
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=False,
streaming_sampling_rate=2,
)
unified_guardrail = UnifiedLLMGuardrails()
async def mock_stream():
chunks_data = ["A", "B", "C", "D", "E", "F"]
for i, content in enumerate(chunks_data):
yield _make_stream_chunk(
content,
finish_reason="stop" if i == len(chunks_data) - 1 else None,
)
mock_post = AsyncMock(
return_value=_mock_guardrail_post_response(action="NONE", texts=["ABCDEF"])
)
with (
patch.object(guardrail.async_handler, "post", mock_post),
patch(
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
return_value=_make_assembled_model_response("ABCDEF"),
),
):
user_api_key_dict = UserAPIKeyAuth(
api_key="test", request_route="/chat/completions"
)
request_data = {
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["test-generic-guardrail"]},
}
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=mock_stream(),
request_data=request_data,
):
pass
assert mock_post.await_count == 4, (
f"Expected 4 guardrail calls (3 sampled + 1 final aggregate), "
f"got {mock_post.await_count}"
)
@pytest.mark.asyncio
async def test_streaming_fail_open_on_unreachable_continues_stream(self):
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",
unreachable_fallback="fail_open",
streaming_end_of_stream_only=True,
)
unified_guardrail = UnifiedLLMGuardrails()
async def mock_stream():
for i, content in enumerate(["A", "B", "C"]):
yield _make_stream_chunk(
content, finish_reason="stop" if i == 2 else None
)
mock_post = AsyncMock(side_effect=httpx.ConnectError("connection refused"))
with (
patch.object(guardrail.async_handler, "post", mock_post),
patch(
"litellm.llms.openai.chat.guardrail_translation.handler.stream_chunk_builder",
return_value=_make_assembled_model_response("ABC"),
),
):
user_api_key_dict = UserAPIKeyAuth(
api_key="test", request_route="/chat/completions"
)
request_data = {
"messages": [{"role": "user", "content": "hi"}],
"guardrail_to_apply": guardrail,
"metadata": {"guardrails": ["test-generic-guardrail"]},
}
chunks_received = 0
async for _ in unified_guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=mock_stream(),
request_data=request_data,
):
chunks_received += 1
assert chunks_received == 3