mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
5685433706
commit
820e1248c3
3 changed files with 45 additions and 29 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue