fix: BerriAI/litellm#36535 - [Bug]: sensitive_data_routing documented but not recognized as guardrail

Signed-off-by: pjdurden <prajjwalchittori1@gmail.com>
This commit is contained in:
pjdurden 2026-08-19 10:30:37 -05:00
parent a112ba5f63
commit 953754c42a
No known key found for this signature in database
9 changed files with 476 additions and 4 deletions

View file

@ -1205,7 +1205,8 @@ class SensitiveDataRouteException(Exception):
The proxy catches this exception and:
1. Reroutes the current request to the specified model
2. When sticky_session_routing is True, stores the routing decision in session
cache so all subsequent requests in the same session are routed to the same model
cache so all subsequent requests in the same session are routed to the same model,
for session_ttl_seconds when the guardrail sets one and the proxy-wide default otherwise
"""
def __init__(
@ -1216,11 +1217,13 @@ class SensitiveDataRouteException(Exception):
detection_info: dict[str, Any] | None = None,
message: str | None = None,
sticky_session_routing: bool = True,
session_ttl_seconds: int | None = None,
):
self.route_to_model = route_to_model
self.session_id = session_id
self.guardrail_name = guardrail_name
self.detection_info = detection_info or {}
self.sticky_session_routing = sticky_session_routing
self.session_ttl_seconds = session_ttl_seconds
self.message = message or f"Sensitive data detected by {guardrail_name}. Routing to model: {route_to_model}"
super().__init__(self.message)

View file

@ -0,0 +1,67 @@
"""Built-in Sensitive Data Routing guardrail: reroutes sensitive prompts to an on-premise model."""
from typing import TYPE_CHECKING, Final
from litellm.types.guardrails import GuardrailEventHooks, Mode, SupportedGuardrailIntegrations
from litellm.types.proxy.guardrails.guardrail_hooks.sensitive_data_routing import (
SensitiveDataRoutingGuardrailConfigModel,
)
from .sensitive_data_routing import SensitiveDataRoutingGuardrail
if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
def _to_event_hook(
mode: str | list[str] | Mode | None,
) -> GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None:
if mode is None or isinstance(mode, Mode):
return mode
if isinstance(mode, str):
return GuardrailEventHooks(mode)
return [GuardrailEventHooks(single_mode) for single_mode in mode]
def initialize_guardrail(
litellm_params: "LitellmParams",
guardrail: "Guardrail",
) -> SensitiveDataRoutingGuardrail:
"""Initialize the Sensitive Data Routing guardrail from config."""
import litellm
guardrail_name: Final = guardrail.get("guardrail_name")
if not guardrail_name:
raise ValueError("sensitive_data_routing guardrail requires a guardrail_name")
config: Final = SensitiveDataRoutingGuardrailConfigModel.model_validate(litellm_params.model_dump())
if not config.on_premise_model:
raise ValueError("sensitive_data_routing guardrail requires 'on_premise_model'")
instance: Final = SensitiveDataRoutingGuardrail(
on_premise_model=config.on_premise_model,
guardrail_name=guardrail_name,
prebuilt_patterns=config.prebuilt_patterns,
regex_patterns=config.regex_patterns,
keywords=config.keywords,
sticky_session=config.sticky_session,
session_ttl_seconds=config.session_ttl_seconds,
event_hook=_to_event_hook(litellm_params.mode),
default_on=bool(litellm_params.default_on),
)
litellm.logging_callback_manager.add_litellm_callback(instance)
return instance
guardrail_initializer_registry: Final = {
SupportedGuardrailIntegrations.SENSITIVE_DATA_ROUTING.value: initialize_guardrail,
}
guardrail_class_registry: Final = {
SupportedGuardrailIntegrations.SENSITIVE_DATA_ROUTING.value: SensitiveDataRoutingGuardrail,
}
__all__ = [
"SensitiveDataRoutingGuardrail",
"initialize_guardrail",
]

View file

@ -0,0 +1,144 @@
"""
Built-in Sensitive Data Routing guardrail.
Detects sensitive data with prebuilt regex patterns, custom regex and keyword matching,
then reroutes the request to an on-premise model instead of blocking or redacting it.
"""
import re
from collections.abc import Sequence
from dataclasses import dataclass
from re import Pattern
from typing import TYPE_CHECKING, Final, Literal, Optional
from litellm.exceptions import SensitiveDataRouteException
from litellm.integrations.custom_guardrail import (
CustomGuardrail,
get_session_id_from_request_data,
log_guardrail_information,
)
from litellm.proxy.guardrails.guardrail_hooks.litellm_content_filter.patterns import (
get_compiled_pattern,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
from litellm.types.proxy.guardrails.guardrail_hooks.sensitive_data_routing import (
DEFAULT_SESSION_TTL_SECONDS,
SensitiveDataRoutingGuardrailConfigModel,
)
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
SensitiveDataDetectorKind = Literal["prebuilt_pattern", "regex_pattern", "keyword"]
@dataclass(frozen=True, slots=True)
class SensitiveDataDetector:
kind: SensitiveDataDetectorKind
rule: str
matcher: Pattern[str]
def build_detectors(
prebuilt_patterns: Sequence[str] | None,
regex_patterns: Sequence[str] | None,
keywords: Sequence[str] | None,
) -> tuple[SensitiveDataDetector, ...]:
return (
*(
SensitiveDataDetector("prebuilt_pattern", name, get_compiled_pattern(name))
for name in prebuilt_patterns or ()
),
*(
SensitiveDataDetector("regex_pattern", pattern, re.compile(pattern, re.IGNORECASE))
for pattern in regex_patterns or ()
),
*(
SensitiveDataDetector("keyword", keyword, re.compile(re.escape(keyword), re.IGNORECASE))
for keyword in keywords or ()
),
)
class SensitiveDataRoutingGuardrail(CustomGuardrail):
"""
Reroutes a request to an on-premise model when sensitive data is detected.
Runs locally with no external API call, and never blocks or redacts: the prompt is
forwarded unchanged to the on-premise model. When sticky_session is enabled and the
request carries a session id, the whole session stays pinned to that model.
"""
def __init__(
self,
on_premise_model: str,
guardrail_name: str | None = None,
prebuilt_patterns: Sequence[str] | None = None,
regex_patterns: Sequence[str] | None = None,
keywords: Sequence[str] | None = None,
sticky_session: bool = True,
session_ttl_seconds: int = DEFAULT_SESSION_TTL_SECONDS,
event_hook: GuardrailEventHooks | list[GuardrailEventHooks] | None = None,
default_on: bool = False,
) -> None:
if not on_premise_model:
raise ValueError("sensitive_data_routing guardrail requires 'on_premise_model'")
detectors: Final = build_detectors(prebuilt_patterns, regex_patterns, keywords)
if not detectors:
raise ValueError(
"sensitive_data_routing guardrail requires at least one of "
"'prebuilt_patterns', 'regex_patterns' or 'keywords'"
)
super().__init__(
guardrail_name=guardrail_name,
supported_event_hooks=self.get_supported_event_hooks(),
event_hook=event_hook or GuardrailEventHooks.pre_call,
default_on=default_on,
)
self.on_premise_model = on_premise_model
self.detectors = detectors
self.sticky_session = sticky_session
self.session_ttl_seconds = session_ttl_seconds
@classmethod
def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]:
return [GuardrailEventHooks.pre_call]
@staticmethod
def get_config_model() -> type[GuardrailConfigModel] | None:
return SensitiveDataRoutingGuardrailConfigModel
def detect(self, texts: Sequence[str]) -> SensitiveDataDetector | None:
return next(
(detector for detector in self.detectors for text in texts if text and detector.matcher.search(text)),
None,
)
@log_guardrail_information
async def apply_guardrail(
self,
inputs: GenericGuardrailAPIInputs,
request_data: dict,
input_type: Literal["request", "response"],
logging_obj: Optional["LiteLLMLoggingObj"] = None,
) -> GenericGuardrailAPIInputs:
if input_type != "request":
return inputs
detected: Final = self.detect(inputs.get("texts") or ())
if detected is None:
return inputs
session_id: Final = get_session_id_from_request_data(request_data)
raise SensitiveDataRouteException(
route_to_model=self.on_premise_model,
session_id=session_id or "",
guardrail_name=self.guardrail_name,
detection_info={"detection_type": detected.kind, "rule": detected.rule},
sticky_session_routing=self.sticky_session and session_id is not None,
session_ttl_seconds=self.session_ttl_seconds,
)

View file

@ -116,6 +116,7 @@ class _PROXY_SensitiveDataRoutingHandler(CustomLogger):
model: str,
user_api_key_dict: UserAPIKeyAuth | None = None,
guardrail_name: str | None = None,
ttl: int | None = None,
) -> None:
"""
Store a routing override for a session.
@ -123,15 +124,18 @@ class _PROXY_SensitiveDataRoutingHandler(CustomLogger):
Called by guardrails when they detect sensitive data and want to
route the session to a specific model. The override is scoped to the
requesting principal so sessions from different tenants cannot collide.
A guardrail that configures its own session TTL passes it as ``ttl``;
otherwise the proxy-wide default applies.
"""
cache_key: Final = self._make_cache_key(session_id, self._resolve_tenant(user_api_key_dict))
effective_ttl: Final = ttl if ttl is not None else self.ttl
verbose_proxy_logger.info(
"SensitiveDataRoutingHandler: Setting session routing session_id=%s model=%s guardrail=%s ttl=%s",
session_id,
model,
guardrail_name,
self.ttl,
effective_ttl,
)
if self.internal_usage_cache.dual_cache.redis_cache is not None:
@ -139,7 +143,7 @@ class _PROXY_SensitiveDataRoutingHandler(CustomLogger):
await self.internal_usage_cache.dual_cache.redis_cache.async_set_cache(
key=cache_key,
value=model,
ttl=self.ttl,
ttl=effective_ttl,
)
except Exception as e:
verbose_proxy_logger.warning(
@ -150,7 +154,7 @@ class _PROXY_SensitiveDataRoutingHandler(CustomLogger):
await self.internal_usage_cache.async_set_cache(
key=cache_key,
value=model,
ttl=self.ttl,
ttl=effective_ttl,
litellm_parent_otel_span=None,
local_only=True,
)

View file

@ -1804,6 +1804,7 @@ class ProxyLogging:
model=exc.route_to_model,
user_api_key_dict=user_api_key_dict,
guardrail_name=exc.guardrail_name,
ttl=exc.session_ttl_seconds,
)
else:
verbose_proxy_logger.warning(

View file

@ -134,6 +134,7 @@ class SupportedGuardrailIntegrations(Enum):
HEADROOM = "headroom"
COMPRESR = "compresr"
STRAIKER = "straiker"
SENSITIVE_DATA_ROUTING = "sensitive_data_routing"
class Role(Enum):

View file

@ -0,0 +1,42 @@
"""Types for the built-in Sensitive Data Routing guardrail."""
from typing import Final
from pydantic import Field
from .base import GuardrailConfigModel
DEFAULT_SESSION_TTL_SECONDS: Final = 14400
class SensitiveDataRoutingGuardrailConfigModel(GuardrailConfigModel):
"""Configuration for the built-in Sensitive Data Routing guardrail."""
on_premise_model: str | None = Field(
default=None,
description="Model group (from model_list) to route the request to when sensitive data is detected.",
)
prebuilt_patterns: list[str] | None = Field(
default=None,
description="Prebuilt pattern names to match (e.g. us_ssn, credit_card, email).",
)
regex_patterns: list[str] | None = Field(
default=None,
description="Custom regular expressions; a match in any message reroutes the request.",
)
keywords: list[str] | None = Field(
default=None,
description="Case-insensitive keywords; a match in any message reroutes the request.",
)
sticky_session: bool = Field(
default=True,
description="Keep the whole session on the on-premise model after sensitive data is first detected.",
)
session_ttl_seconds: int = Field(
default=DEFAULT_SESSION_TTL_SECONDS,
description="How long a session stays pinned to the on-premise model after detection.",
)
@staticmethod
def ui_friendly_name() -> str:
return "Sensitive Data Routing"

View file

@ -0,0 +1,163 @@
"""Tests for the built-in Sensitive Data Routing guardrail."""
import pytest
from litellm.exceptions import SensitiveDataRouteException
from litellm.proxy.guardrails.guardrail_hooks.sensitive_data_routing import (
SensitiveDataRoutingGuardrail,
initialize_guardrail,
)
from litellm.proxy.guardrails.guardrail_registry import (
IN_MEMORY_GUARDRAIL_HANDLER,
guardrail_class_registry,
guardrail_initializer_registry,
)
from litellm.types.guardrails import GuardrailEventHooks, LitellmParams
DOCUMENTED_LITELLM_PARAMS = {
"guardrail": "sensitive_data_routing",
"mode": "pre_call",
"default_on": True,
"on_premise_model": "on-prem-model",
"prebuilt_patterns": ["us_ssn", "credit_card", "email"],
"regex_patterns": [r"project\s+titan"],
"keywords": ["confidential", "internal only"],
"sticky_session": True,
"session_ttl_seconds": 14400,
}
def make_guardrail(**overrides) -> SensitiveDataRoutingGuardrail:
litellm_params = LitellmParams(**{**DOCUMENTED_LITELLM_PARAMS, **overrides})
return initialize_guardrail(litellm_params, {"guardrail_name": "sensitive-data-routing"})
class TestSensitiveDataRoutingGuardrailRegistration:
def test_guardrail_type_is_registered(self):
"""Regression: `guardrail: sensitive_data_routing` used to raise "Unsupported guardrail"."""
assert "sensitive_data_routing" in guardrail_initializer_registry
assert guardrail_class_registry["sensitive_data_routing"] is SensitiveDataRoutingGuardrail
def test_documented_config_initializes_through_the_registry(self):
"""The config.yaml from the docs initializes instead of failing proxy startup."""
guardrail = IN_MEMORY_GUARDRAIL_HANDLER.initialize_guardrail(
{
"guardrail_name": "docs-sensitive-data-routing",
"litellm_params": dict(DOCUMENTED_LITELLM_PARAMS),
}
)
assert guardrail is not None
callback = IN_MEMORY_GUARDRAIL_HANDLER.guardrail_id_to_custom_guardrail[guardrail["guardrail_id"]]
assert isinstance(callback, SensitiveDataRoutingGuardrail)
assert callback.on_premise_model == "on-prem-model"
assert callback.event_hook == GuardrailEventHooks.pre_call
def test_on_premise_model_is_required(self):
with pytest.raises(ValueError, match="on_premise_model"):
make_guardrail(on_premise_model=None)
def test_at_least_one_detector_is_required(self):
with pytest.raises(ValueError, match="prebuilt_patterns"):
make_guardrail(prebuilt_patterns=None, regex_patterns=None, keywords=None)
def test_unknown_prebuilt_pattern_fails_fast(self):
with pytest.raises(ValueError, match="Unknown pattern name"):
make_guardrail(prebuilt_patterns=["not_a_real_pattern"])
class TestSensitiveDataRoutingGuardrailDetection:
@pytest.mark.asyncio
async def test_clean_request_is_untouched(self):
guardrail = make_guardrail()
inputs = {"texts": ["What is the capital of France?"]}
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={"metadata": {"session_id": "abc-123"}},
input_type="request",
)
assert result["texts"] == ["What is the capital of France?"]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"text, detection_type, rule",
[
("My SSN is 123-45-6789, summarize my record", "prebuilt_pattern", "us_ssn"),
("notes on project titan", "regex_pattern", r"project\s+titan"),
("This is INTERNAL ONLY", "keyword", "internal only"),
],
)
async def test_detection_reroutes_to_the_on_premise_model(self, text, detection_type, rule):
guardrail = make_guardrail()
with pytest.raises(SensitiveDataRouteException) as exc_info:
await guardrail.apply_guardrail(
inputs={"texts": [text]},
request_data={"model": "cloud-model", "metadata": {"session_id": "abc-123"}},
input_type="request",
)
exc = exc_info.value
assert exc.route_to_model == "on-prem-model"
assert exc.session_id == "abc-123"
assert exc.guardrail_name == "sensitive-data-routing"
assert exc.detection_info == {"detection_type": detection_type, "rule": rule}
assert exc.sticky_session_routing is True
assert exc.session_ttl_seconds == 14400
@pytest.mark.asyncio
async def test_detection_never_leaks_the_matched_value(self):
guardrail = make_guardrail()
with pytest.raises(SensitiveDataRouteException) as exc_info:
await guardrail.apply_guardrail(
inputs={"texts": ["My SSN is 123-45-6789"]},
request_data={"metadata": {"session_id": "abc-123"}},
input_type="request",
)
assert "123-45-6789" not in str(exc_info.value.detection_info)
@pytest.mark.asyncio
async def test_request_without_session_id_reroutes_without_pinning(self):
"""Docs: turns without a session id are still routed, but never pinned."""
guardrail = make_guardrail()
with pytest.raises(SensitiveDataRouteException) as exc_info:
await guardrail.apply_guardrail(
inputs={"texts": ["My SSN is 123-45-6789"]},
request_data={"model": "cloud-model"},
input_type="request",
)
assert exc_info.value.route_to_model == "on-prem-model"
assert exc_info.value.sticky_session_routing is False
@pytest.mark.asyncio
async def test_sticky_session_disabled_does_not_pin(self):
guardrail = make_guardrail(sticky_session=False)
with pytest.raises(SensitiveDataRouteException) as exc_info:
await guardrail.apply_guardrail(
inputs={"texts": ["My SSN is 123-45-6789"]},
request_data={"metadata": {"session_id": "abc-123"}},
input_type="request",
)
assert exc_info.value.sticky_session_routing is False
@pytest.mark.asyncio
async def test_responses_are_not_scanned(self):
"""The guardrail only picks the model, so it has nothing to do on the response."""
guardrail = make_guardrail()
inputs = {"texts": ["My SSN is 123-45-6789"]}
result = await guardrail.apply_guardrail(
inputs=inputs,
request_data={"metadata": {"session_id": "abc-123"}},
input_type="response",
)
assert result["texts"] == ["My SSN is 123-45-6789"]

View file

@ -749,6 +749,53 @@ class TestProxyHandleSensitiveDataRouteException:
== "on-premise-model"
)
@pytest.mark.asyncio
async def test_guardrail_session_ttl_overrides_the_proxy_default(
self, proxy_logging, routing_hook
):
proxy_logging.proxy_hook_mapping["sensitive_data_routing"] = routing_hook
exc = SensitiveDataRouteException(
route_to_model="on-premise-model",
session_id="sess-ttl",
guardrail_name="sensitive-data-routing",
sticky_session_routing=True,
session_ttl_seconds=14400,
)
await proxy_logging._handle_sensitive_data_route_exception(
exc,
{"model": "gpt-4", "metadata": {"session_id": "sess-ttl"}},
UserAPIKeyAuth(api_key="tenant-a"),
)
cache_key = routing_hook._make_cache_key("sess-ttl", "tenant-a")
assert routing_hook.internal_usage_cache._ttls[cache_key] == 14400
assert routing_hook.ttl == DEFAULT_SENSITIVE_ROUTING_TTL
@pytest.mark.asyncio
async def test_routing_without_guardrail_ttl_uses_the_proxy_default(
self, proxy_logging, routing_hook
):
proxy_logging.proxy_hook_mapping["sensitive_data_routing"] = routing_hook
exc = SensitiveDataRouteException(
route_to_model="on-premise-model",
session_id="sess-default-ttl",
guardrail_name="pii",
sticky_session_routing=True,
)
await proxy_logging._handle_sensitive_data_route_exception(
exc,
{"model": "gpt-4", "metadata": {"session_id": "sess-default-ttl"}},
UserAPIKeyAuth(api_key="tenant-a"),
)
cache_key = routing_hook._make_cache_key("sess-default-ttl", "tenant-a")
assert (
routing_hook.internal_usage_cache._ttls[cache_key]
== DEFAULT_SENSITIVE_ROUTING_TTL
)
@pytest.mark.asyncio
async def test_non_sticky_routing_does_not_persist_override(
self, proxy_logging, routing_hook