From 647a0c633f1591e92a4b7684691bcdfe9a3385c4 Mon Sep 17 00:00:00 2001 From: Marton Schneider Date: Tue, 23 Jun 2026 16:58:27 +0200 Subject: [PATCH] fix(guardrails): read nested streaming config from dict optional_params Guardrail API/UI delivers optional_params as a plain dict, so getattr was silently ignoring streaming_sampling_rate and streaming_end_of_stream_only. Handle both dict and model shapes in _get_config_value with regression tests. --- .../generic_guardrail_api/__init__.py | 6 +- .../test_generic_guardrail_api.py | 62 +++++++++++++++++++ 2 files changed, 67 insertions(+), 1 deletion(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py index 02f02a8a5eb..0323cba6e95 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py @@ -12,7 +12,11 @@ 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) + value = ( + optional_params.get(attribute_name) + if isinstance(optional_params, dict) + else getattr(optional_params, attribute_name, None) + ) if value is not None: return value return getattr(litellm_params, attribute_name, None) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py index 42fe56d1b8b..da8a6aa58a3 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py @@ -1247,6 +1247,68 @@ class TestGenericGuardrailAPIStreamingConfig: assert guardrail.streaming_end_of_stream_only is True assert guardrail.streaming_sampling_rate == 1 + def test_initialize_guardrail_dict_optional_params_streaming_wins(self): + """Guardrail API/UI delivers optional_params as a plain dict, not a model.""" + 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, + ) + litellm_params.streaming_end_of_stream_only = False # type: ignore[attr-defined] + litellm_params.streaming_sampling_rate = 9 # type: ignore[attr-defined] + # Plain dict mirrors how configs arrive from the guardrail API/UI. + litellm_params.optional_params = { # type: ignore[attr-defined] + "streaming_end_of_stream_only": True, + "streaming_sampling_rate": 1, + } + + guardrail_config = {"guardrail_name": "test-generic-streaming-dict-optional"} + + with patch( + "litellm.logging_callback_manager.add_litellm_callback" + ): + guardrail = initialize_guardrail(litellm_params, guardrail_config) + + assert guardrail.streaming_end_of_stream_only is True + assert guardrail.streaming_sampling_rate == 1 + + def test_initialize_guardrail_dict_optional_params_sibling_only_falls_through( + self, + ): + """Dict optional_params without streaming keys must not shadow top-level knobs.""" + 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, + ) + litellm_params.streaming_end_of_stream_only = True # type: ignore[attr-defined] + litellm_params.streaming_sampling_rate = 2 # type: ignore[attr-defined] + litellm_params.optional_params = { # type: ignore[attr-defined] + "additional_provider_specific_params": {"tenant": "acme"}, + } + + guardrail_config = {"guardrail_name": "test-generic-streaming-dict-sibling"} + + with patch( + "litellm.logging_callback_manager.add_litellm_callback" + ): + guardrail = initialize_guardrail(litellm_params, guardrail_config) + + assert guardrail.streaming_end_of_stream_only is True + assert guardrail.streaming_sampling_rate == 2 + class TestGenericGuardrailAPIStreamingViaUnified: """Streaming output checks routed through UnifiedLLMGuardrails."""