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
This commit is contained in:
mateo-berri 2026-06-03 03:58:20 +00:00
parent 5685433706
commit 820e1248c3
No known key found for this signature in database
3 changed files with 45 additions and 29 deletions

View file

@ -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

View file

@ -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,
)

View file

@ -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",