mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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.
This commit is contained in:
parent
572acf0037
commit
647a0c633f
2 changed files with 67 additions and 1 deletions
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue