From 820e1248c3de7d6d6456c0340d724a8a21e3b359 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 3 Jun 2026 03:58:20 +0000 Subject: [PATCH] Fix sensitive_data_routing event_hook typing and tighten coverage Normalize event_hook to a GuardrailEventHooks before passing it to the CustomGuardrail base so mypy is satisfied with the declared type. Drop the unused raw-config fallback in the initializer and read params directly from litellm_params, and add tests covering the config model exposure, non-dict messages, and the missing guardrail_name validation https://claude.ai/code/session_01R2hGPr2jc5kxnSYLRgqfMs --- .../sensitive_data_routing/__init__.py | 37 +++++-------------- .../sensitive_data_routing.py | 6 ++- .../test_sensitive_data_routing.py | 31 ++++++++++++++++ 3 files changed, 45 insertions(+), 29 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/sensitive_data_routing/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/sensitive_data_routing/__init__.py index c33c37a2f28..3e2e5ac2aeb 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/sensitive_data_routing/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/sensitive_data_routing/__init__.py @@ -1,6 +1,6 @@ """Sensitive Data Routing guardrail: reroutes requests with sensitive data to an on-premise model.""" -from typing import TYPE_CHECKING, Any, List +from typing import TYPE_CHECKING, List from litellm.types.guardrails import SupportedGuardrailIntegrations @@ -10,21 +10,6 @@ if TYPE_CHECKING: from litellm.types.guardrails import Guardrail, LitellmParams -def _get_param( - litellm_params: "LitellmParams", - guardrail: "Guardrail", - key: str, - default: Any = None, -) -> Any: - value = getattr(litellm_params, key, None) - if value is not None: - return value - raw = guardrail.get("litellm_params") - if isinstance(raw, dict) and key in raw: - return raw[key] - return default - - def initialize_guardrail( litellm_params: "LitellmParams", guardrail: "Guardrail", @@ -35,7 +20,7 @@ def initialize_guardrail( if not guardrail_name: raise ValueError("sensitive_data_routing guardrail requires a guardrail_name") - on_premise_model = _get_param(litellm_params, guardrail, "on_premise_model") + on_premise_model = getattr(litellm_params, "on_premise_model", None) if not on_premise_model: raise ValueError( "sensitive_data_routing guardrail requires 'on_premise_model' (the model_list " @@ -45,17 +30,13 @@ def initialize_guardrail( instance = SensitiveDataRoutingGuardrail( guardrail_name=guardrail_name, on_premise_model=on_premise_model, - prebuilt_patterns=_get_param(litellm_params, guardrail, "prebuilt_patterns"), - regex_patterns=_get_param(litellm_params, guardrail, "regex_patterns"), - keywords=_get_param(litellm_params, guardrail, "keywords"), - sticky_session=bool( - _get_param(litellm_params, guardrail, "sticky_session", True) - ), - session_ttl_seconds=int( - _get_param(litellm_params, guardrail, "session_ttl_seconds", 14400) - ), - event_hook=_get_param(litellm_params, guardrail, "mode"), - default_on=bool(_get_param(litellm_params, guardrail, "default_on", False)), + prebuilt_patterns=getattr(litellm_params, "prebuilt_patterns", None), + regex_patterns=getattr(litellm_params, "regex_patterns", None), + keywords=getattr(litellm_params, "keywords", None), + sticky_session=bool(getattr(litellm_params, "sticky_session", True)), + session_ttl_seconds=int(getattr(litellm_params, "session_ttl_seconds", 14400)), + event_hook=getattr(litellm_params, "mode", None), + default_on=bool(getattr(litellm_params, "default_on", False)), ) litellm.logging_callback_manager.add_litellm_callback(instance) return instance diff --git a/litellm/proxy/guardrails/guardrail_hooks/sensitive_data_routing/sensitive_data_routing.py b/litellm/proxy/guardrails/guardrail_hooks/sensitive_data_routing/sensitive_data_routing.py index e801fdfe011..120ff752f09 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/sensitive_data_routing/sensitive_data_routing.py +++ b/litellm/proxy/guardrails/guardrail_hooks/sensitive_data_routing/sensitive_data_routing.py @@ -49,7 +49,11 @@ class SensitiveDataRoutingGuardrail(CustomGuardrail): super().__init__( guardrail_name=guardrail_name or "sensitive_data_routing", supported_event_hooks=[GuardrailEventHooks.pre_call], - event_hook=event_hook or GuardrailEventHooks.pre_call, + event_hook=( + GuardrailEventHooks(event_hook) + if event_hook is not None + else GuardrailEventHooks.pre_call + ), default_on=default_on, **kwargs, ) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_sensitive_data_routing.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_sensitive_data_routing.py index cb31716bac1..d2aa9430879 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_sensitive_data_routing.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_sensitive_data_routing.py @@ -11,6 +11,9 @@ from litellm.proxy.guardrails.guardrail_hooks.sensitive_data_routing import ( initialize_guardrail, ) from litellm.types.guardrails import GuardrailEventHooks, LitellmParams +from litellm.types.proxy.guardrails.guardrail_hooks.sensitive_data_routing import ( + SensitiveDataRoutingConfigModel, +) ON_PREM = "on-prem-model" USER_KEY = UserAPIKeyAuth() @@ -87,6 +90,17 @@ async def test_keyword_match_is_case_insensitive(): assert result is not None and result["model"] == ON_PREM +@pytest.mark.asyncio +async def test_non_dict_messages_are_skipped(): + g = _make_guardrail() + data = { + "model": "gpt-4o", + "messages": ["not-a-dict", {"role": "user", "content": "ssn 123-45-6789"}], + } + result = await _hook(g, data) + assert result is not None and result["model"] == ON_PREM + + @pytest.mark.asyncio async def test_detects_text_in_content_parts_list(): g = _make_guardrail() @@ -180,6 +194,17 @@ def test_only_supports_pre_call_event_hook(): _make_guardrail(event_hook=GuardrailEventHooks.post_call) +def test_initializer_requires_guardrail_name(): + params = LitellmParams( + guardrail="sensitive_data_routing", mode="pre_call", on_premise_model=ON_PREM + ) + with pytest.raises(ValueError): + initialize_guardrail( + litellm_params=params, + guardrail={"guardrail_name": "", "litellm_params": {}}, + ) + + def test_initializer_requires_on_premise_model(): params = LitellmParams(guardrail="sensitive_data_routing", mode="pre_call") with pytest.raises(ValueError): @@ -189,6 +214,12 @@ def test_initializer_requires_on_premise_model(): ) +def test_config_model_is_exposed_for_ui(): + config_model = SensitiveDataRoutingGuardrail.get_config_model() + assert config_model is SensitiveDataRoutingConfigModel + assert config_model.ui_friendly_name() == "Sensitive Data Routing" + + def test_initializer_builds_guardrail_from_config(): params = LitellmParams( guardrail="sensitive_data_routing",