From 2c0fd473b22cb7eb5b28d86cc4b728a3ebf43d6b Mon Sep 17 00:00:00 2001 From: Gourav Mittal Date: Thu, 1 Oct 2026 15:13:28 -0700 Subject: [PATCH 1/4] feat(guardrails): add RealmLabs request hazard screening Screen request messages through /guardrail and block hazard scores above the configured threshold. Validate service responses and apply the configured failure policy to invalid replies and transport errors Resolve explicit nested settings before top-level settings and defaults Include provider registration, generated API schemas, and request tests --- litellm/proxy/_lazy_openapi_snapshot.json | 42 ++ .../guardrail_hooks/realmlabs/__init__.py | 72 +++ .../guardrail_hooks/realmlabs/realmlabs.py | 208 +++++++ litellm/types/guardrails.py | 5 + .../guardrails/guardrail_hooks/realmlabs.py | 122 ++++ .../guardrail_hooks/realmlabs/__init__.py | 0 .../realmlabs/test_realmlabs.py | 559 ++++++++++++++++++ ui/litellm-dashboard/src/lib/http/schema.d.ts | 15 + 8 files changed, 1023 insertions(+) create mode 100644 litellm/proxy/guardrails/guardrail_hooks/realmlabs/__init__.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/realmlabs/realmlabs.py create mode 100644 litellm/types/proxy/guardrails/guardrail_hooks/realmlabs.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/__init__.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_realmlabs.py diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 3346b0c9ff8..e09cceb501c 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -13498,6 +13498,18 @@ "description": "If True, will not raise an exception when the guardrail is blocked. Useful for OpenWebUI where exceptions can end the chat flow.", "title": "Disable Exception On Block" }, + "enable_thinking": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "description": "Whether MLS should render the chat template in thinking mode. Defaults to False.", + "title": "Enable Thinking" + }, "end_session_after_n_fails": { "anyOf": [ { @@ -13668,6 +13680,18 @@ "description": "Enable hallucination detection to detect factual inaccuracies.", "title": "Hallucinations Check" }, + "hazard_threshold": { + "anyOf": [ + { + "type": "number" + }, + { + "type": "null" + } + ], + "description": "Block the request when the hazard_prompt probe scores strictly above this value. Defaults to 0.703, the threshold MLS reports for that probe. Note the probe also responds to instruction-style phrasing such as \"repeat this back verbatim\", so raise this if benign traffic is being blocked.", + "title": "Hazard Threshold" + }, "include_evidence": { "anyOf": [ { @@ -14306,6 +14330,24 @@ "description": "Optional per-entity minimum confidence scores for Presidio detections. Entities below the threshold are ignored.", "title": "Presidio Score Thresholds" }, + "probes": { + "anyOf": [ + { + "items": { + "type": "string" + }, + "type": "array" + }, + { + "type": "string" + }, + { + "type": "null" + } + ], + "description": "Which classifier probes to run: a list of probe names, or \"all\". Defaults to [\"hazard_prompt\"] - the only probe whose score this guardrail enforces. An unknown probe name makes MLS return 404.", + "title": "Probes" + }, "project_id": { "anyOf": [ { diff --git a/litellm/proxy/guardrails/guardrail_hooks/realmlabs/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/__init__.py new file mode 100644 index 00000000000..4d06c5af79e --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/__init__.py @@ -0,0 +1,72 @@ +from types import MappingProxyType +from typing import TYPE_CHECKING, Final + +from pydantic import BaseModel + +from litellm.types.guardrails import GuardrailEventHooks, Mode, SupportedGuardrailIntegrations +from litellm.types.proxy.guardrails.guardrail_hooks.realmlabs import RealmLabsGuardrailOptionalParams + +from .realmlabs import RealmLabsGuardrail + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + +__all__ = ("RealmLabsGuardrail",) + + +def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail") -> RealmLabsGuardrail: + """Build the guardrail from its ``config.yaml`` entry and register it; found via the registries below.""" + import litellm + + settings: Final = _resolved_params(litellm_params) + _realmlabs_callback: Final = RealmLabsGuardrail( + api_key=litellm_params.api_key, + api_base=litellm_params.api_base, + probes=settings.probes, + hazard_threshold=settings.hazard_threshold, + block_on_error=settings.block_on_error, + enable_thinking=settings.enable_thinking, + timeout=settings.timeout, + guardrail_name=guardrail["guardrail_name"], + event_hook=_coerce_event_hook(litellm_params.mode), + default_on=litellm_params.default_on or False, + ) + litellm.logging_callback_manager.add_litellm_callback( # pyright: ignore[reportUnknownMemberType] # callback manager is untyped + _realmlabs_callback + ) + return _realmlabs_callback + + +guardrail_initializer_registry: Final = { + SupportedGuardrailIntegrations.REALMLABS.value: initialize_guardrail, +} + +guardrail_class_registry: Final = { + SupportedGuardrailIntegrations.REALMLABS.value: RealmLabsGuardrail, +} + + +def _coerce_event_hook( + mode: str | list[str] | Mode, # mutable-ok: mirrors LitellmParams.mode +) -> GuardrailEventHooks | list[GuardrailEventHooks] | Mode: # mutable-ok: mirrors CustomGuardrail's event_hook + """Convert the ``mode`` strings from ``config.yaml`` into the enum values ``CustomGuardrail`` expects.""" + if isinstance(mode, Mode): + return mode + if isinstance(mode, list): + return [GuardrailEventHooks(item) for item in mode] + return GuardrailEventHooks(mode) + + +def _resolved_params(litellm_params: "LitellmParams") -> RealmLabsGuardrailOptionalParams: + """Resolve explicit nested values, then top-level values, then RealmLabs defaults. + + Excluding unset fields prevents another guardrail's parsed defaults from overriding RealmLabs settings. + Null nested values fall back to the top level; false, zero, and empty lists remain explicit overrides. + """ + top_level: Final = RealmLabsGuardrailOptionalParams.model_validate(litellm_params.model_dump(exclude_unset=True)) + nested: Final = litellm_params.optional_params + if not isinstance(nested, BaseModel): + return top_level + return RealmLabsGuardrailOptionalParams.model_validate( + MappingProxyType({**top_level.model_dump(), **nested.model_dump(exclude_unset=True, exclude_none=True)}) + ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/realmlabs/realmlabs.py b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/realmlabs.py new file mode 100644 index 00000000000..9ccc142f013 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/realmlabs.py @@ -0,0 +1,208 @@ +"""RealmLabs MLS request hazard guardrail.""" + +from __future__ import annotations + +import json +from collections.abc import Mapping, Sequence +from dataclasses import dataclass +from typing import TYPE_CHECKING, Final, Literal + +from httpx import HTTPError +from pydantic import TypeAdapter, ValidationError + +from litellm._logging import verbose_proxy_logger +from litellm.exceptions import GuardrailRaisedException +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, # pyright: ignore[reportUnknownVariableType] # decorator is untyped in custom_guardrail +) +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # helper is untyped in http_handler + httpxSpecialProvider, +) +from litellm.secret_managers.main import get_secret_str +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.proxy.guardrails.guardrail_hooks.realmlabs import ( + RealmLabsChatMessage, + RealmLabsGuardrailConfigModel, + RealmLabsGuardrailRequest, + RealmLabsGuardrailResponse, +) + +if TYPE_CHECKING: + from httpx import Response as HttpxResponse + + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + from litellm.types.guardrails import Mode + from litellm.types.utils import GenericGuardrailAPIInputs + +_DEFAULT_API_BASE: Final = "https://mls.realmlabs.ai" +_GUARDRAIL_ENDPOINT: Final = "/guardrail" +_HAZARD_PROBE: Final = "hazard_prompt" +_DEFAULT_HAZARD_THRESHOLD: Final = 0.703 +_DEFAULT_TIMEOUT: Final = 15.0 + +_RESPONSE_ADAPTER: Final = TypeAdapter(RealmLabsGuardrailResponse) + + +class RealmLabsMissingCredentials(Exception): + """Raised at startup when no MLS API key is configured.""" + + +@dataclass(frozen=True, slots=True) +class _InvalidResponse: + reason: str + + +class RealmLabsGuardrail(CustomGuardrail): + """Blocks hazardous prompts using the RealmLabs MLS endpoint.""" + + def __init__( + self, + api_key: str | None = None, + api_base: str | None = None, + probes: Sequence[str] | str | None = None, + hazard_threshold: float | None = None, + block_on_error: bool | None = None, + enable_thinking: bool | None = None, + timeout: float | None = None, + guardrail_name: str | None = None, + event_hook: ( # mutable-ok: same type as CustomGuardrail + GuardrailEventHooks | list[GuardrailEventHooks] | Mode | None + ) = None, + default_on: bool = False, + ) -> None: + """Resolve each setting from its argument, then the ``REALMLABS_*`` env vars, then the module defaults.""" + self.api_key = api_key or get_secret_str("REALMLABS_API_KEY") + if not self.api_key: + raise RealmLabsMissingCredentials( + "RealmLabs API key is required. Set REALMLABS_API_KEY in the environment " + "or pass api_key in the guardrail config." + ) + + self.api_base = (api_base or get_secret_str("REALMLABS_API_BASE") or _DEFAULT_API_BASE).rstrip("/") + self.probes: Sequence[str] | str = (_HAZARD_PROBE,) if probes is None else probes + self.hazard_threshold = _DEFAULT_HAZARD_THRESHOLD if hazard_threshold is None else hazard_threshold + self.block_on_error = False if block_on_error is None else block_on_error + self.enable_thinking = False if enable_thinking is None else enable_thinking + self.timeout = _DEFAULT_TIMEOUT if timeout is None else timeout + self.async_handler: AsyncHTTPHandler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback, + ) + super().__init__( # pyright: ignore[reportUnknownMemberType] # CustomGuardrail.__init__ is untyped + guardrail_name=guardrail_name, + supported_event_hooks=self.get_supported_event_hooks(), + event_hook=event_hook, + default_on=default_on, + ) + + @staticmethod + def get_config_model() -> type[RealmLabsGuardrailConfigModel] | None: + """Config model the admin UI uses to render and validate this guardrail's settings.""" + return RealmLabsGuardrailConfigModel + + @classmethod + def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: # mutable-ok: base class returns a list + """Inspect requests on ``pre_call``.""" + return [GuardrailEventHooks.pre_call] + + @staticmethod + def _hazard_score(response: RealmLabsGuardrailResponse) -> float | None: + """Score of the first ``hazard_prompt`` verdict, unless MLS reports a role mismatch.""" + for result in response["results"]: + if result["probe"] == _HAZARD_PROBE: + return None if result.get("role_mismatch") else result["prob"] + return None + + def _build_request(self, messages: Sequence[Mapping[str, object]]) -> RealmLabsGuardrailRequest: + return RealmLabsGuardrailRequest( + messages=messages, + probes=self.probes, + pii=False, + enable_thinking=self.enable_thinking, + ) + + def _parse_response(self, content: bytes) -> RealmLabsGuardrailResponse | _InvalidResponse: + try: + result: Final = _RESPONSE_ADAPTER.validate_json(content) + except ValidationError as exc: + return _InvalidResponse(exc.json(include_input=False, include_context=False, include_url=False)) + + return result + + async def _call_mls( + self, messages: Sequence[Mapping[str, object]] + ) -> RealmLabsGuardrailResponse | _InvalidResponse: + endpoint: Final = f"{self.api_base}{_GUARDRAIL_ENDPOINT}" + verbose_proxy_logger.debug( + "RealmLabs MLS: %s msgs=%d probes=%s", + endpoint, + len(messages), + self.probes, + ) + response: Final[HttpxResponse] = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped + url=endpoint, + headers={ + "Authorization": f"Bearer {self.api_key}", + "Content-Type": "application/json", + }, + content=json.dumps(self._build_request(messages)), + timeout=self.timeout, + ) + response.raise_for_status() + return self._parse_response(response.content) + + def _handle_mls_error(self, inputs: GenericGuardrailAPIInputs, reason: str) -> GenericGuardrailAPIInputs: + verbose_proxy_logger.error("RealmLabs MLS error: %s", reason) + if self.block_on_error: + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"RealmLabs MLS error (block_on_error=True): {reason}", + ) + return inputs + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: Mapping[str, object], + input_type: Literal["request", "response"], + logging_obj: LiteLLMLoggingObj | None = None, + ) -> GenericGuardrailAPIInputs: + """Screen requests for hazardous content through LiteLLM's unified guardrail layer.""" + if input_type != "request": + return inputs + texts: Final = tuple(inputs.get("texts") or ()) + messages: Final = tuple(inputs.get("structured_messages") or ()) or tuple( + RealmLabsChatMessage(role="user", content=text) for text in texts + ) + if not messages: + return inputs + + try: + result: Final = await self._call_mls(messages) + except (HTTPError, TypeError, ValueError) as exc: + return self._handle_mls_error(inputs, str(exc)) + + if isinstance(result, _InvalidResponse): + return self._handle_mls_error(inputs, f"Invalid RealmLabs guardrail response: {result.reason}") + + hazard_score: Final = self._hazard_score(result) + if hazard_score is not None and hazard_score > self.hazard_threshold: + verbose_proxy_logger.warning( + "RealmLabs MLS blocked request: %s=%s > %s", + _HAZARD_PROBE, + hazard_score, + self.hazard_threshold, + ) + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=( + f"Blocked by RealmLabs {_HAZARD_PROBE} probe: " + f"score={hazard_score} exceeds threshold={self.hazard_threshold}" + ), + blocked_content=True, + ) + + return inputs diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 46026c12d24..ed416e8718c 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -53,6 +53,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.qohash import ( from litellm.types.proxy.guardrails.guardrail_hooks.qualifire import ( QualifireGuardrailConfigModel, ) +from litellm.types.proxy.guardrails.guardrail_hooks.realmlabs import ( + RealmLabsGuardrailConfigModel, +) from litellm.types.proxy.guardrails.guardrail_hooks.repelloai import ( RepelloAIGuardrailConfigModel, ) @@ -120,6 +123,7 @@ class SupportedGuardrailIntegrations(Enum): MCP_SECURITY = "mcp_security" ONYX = "onyx" PROMPTGUARD = "promptguard" + REALMLABS = "realmlabs" XECGUARD = "xecguard" PROMPT_SECURITY = "prompt_security" GENERIC_GUARDRAIL_API = "generic_guardrail_api" @@ -1187,6 +1191,7 @@ class LitellmParams( # pyright: ignore[reportIncompatibleVariableOverride] # o GraySwanGuardrailConfigModel, NomaGuardrailConfigModel, PromptGuardConfigModel, + RealmLabsGuardrailConfigModel, XecGuardConfigModel, ToolPermissionGuardrailConfigModel, ZscalerAIGuardConfigModel, diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/realmlabs.py b/litellm/types/proxy/guardrails/guardrail_hooks/realmlabs.py new file mode 100644 index 00000000000..9c7c8072b33 --- /dev/null +++ b/litellm/types/proxy/guardrails/guardrail_hooks/realmlabs.py @@ -0,0 +1,122 @@ +from collections.abc import Mapping, Sequence +from typing import Annotated + +from pydantic import BaseModel, Field +from typing_extensions import ReadOnly, Required, TypedDict + +from .base import GuardrailConfigModel + + +class RealmLabsChatMessage(TypedDict): + """A plain-text chat turn sent to MLS.""" + + role: ReadOnly[str] + content: ReadOnly[str] + + +class RealmLabsGuardrailRequest(TypedDict): + messages: ReadOnly[Sequence[Mapping[str, object]]] + probes: ReadOnly[Sequence[str] | str] + pii: ReadOnly[bool] + enable_thinking: ReadOnly[bool] + + +class RealmLabsProbeResult(TypedDict, total=False): + """One classifier probe verdict. + + ``prob`` is compared against the guardrail's own ``hazard_threshold``, not the ``threshold`` MLS reports. + """ + + probe: ReadOnly[Required[Annotated[str, Field(strict=True, min_length=1)]]] + prob: ReadOnly[Required[Annotated[float, Field(strict=True, ge=0, le=1, allow_inf_nan=False)]]] + threshold: ReadOnly[float | None] + decision: ReadOnly[bool | None] + role_mismatch: ReadOnly[Annotated[bool, Field(strict=True)] | None] + + +class RealmLabsGuardrailResponse(TypedDict, total=False): + """Response body of ``POST {api_base}/guardrail``. MLS is stateless, so it carries no turn id.""" + + results: ReadOnly[Required[Sequence[RealmLabsProbeResult]]] + focal_role: ReadOnly[str | None] + + +class RealmLabsGuardrailOptionalParams(BaseModel): + """Nested tuning settings; explicitly supplied non-null values override the top-level settings.""" + + probes: Sequence[str] | str | None = Field( + default=None, + description="Classifier probes to run: a list of names or 'all'. Overrides top-level probes when supplied.", + ) + hazard_threshold: float | None = Field( + default=None, + description="Block hazard scores strictly above this value. Overrides top-level hazard_threshold when supplied.", + ) + block_on_error: bool | None = Field( + default=None, + description="Whether to block when MLS fails. Overrides top-level block_on_error when supplied.", + ) + + enable_thinking: bool | None = Field( + default=False, + description="Whether MLS should render the chat template in thinking mode. Defaults to False.", + ) + + timeout: float | None = Field( + default=15.0, + description="Timeout in seconds for the MLS request. Defaults to 15.", + ) + + +class RealmLabsGuardrailConfigModel(GuardrailConfigModel[RealmLabsGuardrailOptionalParams]): + """Settings accepted under ``litellm_params`` for ``guardrail: realmlabs``.""" + + api_key: str | None = Field( + default=None, + description=( + "API key for the RealmLabs MLS guardrail endpoint, sent as a bearer token. " + "If not provided, the REALMLABS_API_KEY environment variable is used." + ), + ) + api_base: str | None = Field( + default=None, + description=( + "Base URL of the RealmLabs MLS deployment. The /guardrail path is " + "appended automatically. Defaults to https://mls.realmlabs.ai, and falls " + "back to the REALMLABS_API_BASE environment variable." + ), + ) + probes: list[str] | str | None = Field( # mutable-ok: the admin UI renders only list[...] fields as array inputs + default=None, + description=( + 'Which classifier probes to run: a list of probe names, or "all". ' + 'Defaults to ["hazard_prompt"] - the only probe whose score this ' + "guardrail enforces. An unknown probe name makes MLS return 404." + ), + ) + hazard_threshold: float | None = Field( + default=None, + description=( + "Block the request when the hazard_prompt probe scores strictly above this " + "value. Defaults to 0.703, the threshold MLS reports for that probe. Note " + 'the probe also responds to instruction-style phrasing such as "repeat this ' + 'back verbatim", so raise this if benign traffic is being blocked.' + ), + ) + block_on_error: bool | None = Field( + default=None, + description=( + "Whether to block the request when MLS is unreachable or returns an " + "unreadable response. Defaults to False (fail open), so an MLS outage does " + "not take the gateway down with it. Set to True to fail closed." + ), + ) + enable_thinking: bool | None = Field( + default=None, + description="Whether MLS should render the chat template in thinking mode. Defaults to False.", + ) + + @staticmethod + def ui_friendly_name() -> str: + """Name the admin UI shows for this guardrail.""" + return "RealmLabs MLS" diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/__init__.py b/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_realmlabs.py b/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_realmlabs.py new file mode 100644 index 00000000000..a5ac4c6de2e --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_realmlabs.py @@ -0,0 +1,559 @@ +import json +from collections.abc import Mapping +from http import HTTPStatus +from typing import Final, Literal + +import httpx +import pytest +import respx +from pydantic import TypeAdapter + +import litellm +from litellm.exceptions import GuardrailRaisedException +from litellm.proxy.guardrails.guardrail_hooks.realmlabs import guardrail_initializer_registry +from litellm.proxy.guardrails.guardrail_hooks.realmlabs.realmlabs import ( + RealmLabsGuardrail, + RealmLabsMissingCredentials, +) +from litellm.types.guardrails import GuardrailEventHooks, LitellmParams +from litellm.types.llms.openai import AllMessageValues +from litellm.types.proxy.guardrails.guardrail_hooks.realmlabs import RealmLabsGuardrailOptionalParams +from litellm.types.utils import GenericGuardrailAPIInputs + +_API_KEY = "mls_gr_test" +_API_BASE = "https://mls.example.test" +_URL = f"{_API_BASE}/guardrail" +_DEFAULT_THRESHOLD = 0.703 + + +_JSON_OBJECT = TypeAdapter(dict[str, object]) + + +@pytest.fixture(autouse=True) +def fresh_httpx_client(monkeypatch: pytest.MonkeyPatch) -> None: + """Route MLS calls through a fresh httpx client that ``respx`` can intercept, and clear ``REALMLABS_*`` env vars.""" + + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", None) + monkeypatch.delenv("REALMLABS_API_KEY", raising=False) + monkeypatch.delenv("REALMLABS_API_BASE", raising=False) + + +def _configured_guardrail(settings: Mapping[str, object]) -> RealmLabsGuardrail: + """Build a guardrail through the real config model and initializer, keeping each test's settings visible.""" + + params: Final = LitellmParams.model_validate( + { + "guardrail": "realmlabs", + "mode": "pre_call", + "api_key": _API_KEY, + "api_base": _API_BASE, + **settings, + } + ) + return guardrail_initializer_registry["realmlabs"]( + params, {"guardrail_name": "rl-config", "litellm_params": params} + ) + + +def _guardrail( + hazard_threshold: float | None = None, + block_on_error: bool | None = None, + event_hook: GuardrailEventHooks = GuardrailEventHooks.pre_call, +) -> RealmLabsGuardrail: + """Guardrail pointed at the fake MLS URL; unset arguments keep the guardrail's defaults.""" + + return RealmLabsGuardrail( + api_key=_API_KEY, + api_base=_API_BASE, + hazard_threshold=hazard_threshold, + block_on_error=block_on_error, + guardrail_name="realmlabs-guard", + event_hook=event_hook, + default_on=True, + ) + + +def _mls_body(hazard: float | None = 0.01, pii_spans: list[dict[str, object]] | None = None) -> dict[str, object]: + """MLS reply with a ``hazard_prompt`` score (omitted when ``hazard`` is None) and the given PII spans.""" + + results = [] if hazard is None else [{"probe": "hazard_prompt", "prob": hazard, "role_mismatch": False}] + return {"results": results, "pii_spans": pii_spans or []} + + +def _serve(respx_mock: respx.MockRouter, body: dict[str, object], url: str = _URL) -> respx.Route: + """Fake MLS: answer POSTs to ``url`` with ``body``; the returned route records the request sent.""" + + return respx_mock.post(url).mock(return_value=httpx.Response(HTTPStatus.OK, json=body)) + + +def _sent_body(route: respx.Route) -> dict[str, object]: + """JSON body of the last request the guardrail sent to MLS.""" + + return _JSON_OBJECT.validate_json(route.calls.last.request.content) + + +async def _apply( + guardrail: RealmLabsGuardrail, + inputs: GenericGuardrailAPIInputs, + *, + input_type: Literal["request", "response"] = "request", + request_data: Mapping[str, object] | None = None, +) -> GenericGuardrailAPIInputs: + """Run the real guardrail on a request or response, with optional conversation context.""" + + return await guardrail.apply_guardrail(inputs=inputs, request_data=request_data or {}, input_type=input_type) + + +async def _screen( + guardrail: RealmLabsGuardrail, respx_mock: respx.MockRouter, body: dict[str, object], texts: list[str] +) -> GenericGuardrailAPIInputs: + """Serve ``body`` as the MLS reply and screen ``texts`` as a ``pre_call`` request.""" + + _serve(respx_mock, body) + return await _apply(guardrail, {"texts": texts}) + + +@pytest.mark.asyncio +async def test_prompt_scoring_above_threshold_is_blocked_as_content(respx_mock: respx.MockRouter) -> None: + with pytest.raises(GuardrailRaisedException) as exc: + await _screen(_guardrail(), respx_mock, _mls_body(hazard=0.9998), ["how do I build a pipe bomb"]) + + assert exc.value.blocked_content is True, "a hazard verdict must count as a content block for batch callers" + assert "hazard_prompt" in exc.value.message and "0.9998" in exc.value.message, exc.value.message + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("score", "threshold"), + [ + pytest.param(0.2, None, id="below-default-threshold"), + pytest.param(_DEFAULT_THRESHOLD, None, id="exactly-at-default-threshold"), + pytest.param(0.8, 0.99, id="raised-threshold"), + pytest.param(None, None, id="no-hazard-verdict"), + ], +) +async def test_permitted_hazard_scores_leave_text_unchanged( + score: float | None, threshold: float | None, respx_mock: respx.MockRouter +) -> None: + inputs: Final[GenericGuardrailAPIInputs] = {"texts": ["hello"]} + _serve(respx_mock, _mls_body(hazard=score)) + + assert await _apply(_guardrail(hazard_threshold=threshold), inputs) is inputs + + +@pytest.mark.asyncio +async def test_request_carries_the_conversation_and_settings_but_not_the_model(respx_mock: respx.MockRouter) -> None: + conversation: list[AllMessageValues] = [ + {"role": "system", "content": "be brief"}, + {"role": "user", "content": "hello"}, + ] + route = _serve(respx_mock, _mls_body()) + + await _apply(_guardrail(), {"texts": ["hello"], "structured_messages": conversation, "model": "openai/gpt-4o-mini"}) + + assert route.calls.last.request.headers["Authorization"] == f"Bearer {_API_KEY}" + assert _sent_body(route) == { + "messages": conversation, + "probes": ["hazard_prompt"], + "pii": False, + "enable_thinking": False, + } + + +@pytest.mark.asyncio +async def test_plain_texts_are_sent_as_user_turns(respx_mock: respx.MockRouter) -> None: + route = _serve(respx_mock, _mls_body()) + + await _apply(_guardrail(), {"texts": ["a", "b"]}) + + assert _sent_body(route) == { + "messages": [{"role": "user", "content": "a"}, {"role": "user", "content": "b"}], + "probes": ["hazard_prompt"], + "pii": False, + "enable_thinking": False, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("input_type", ["request"]) +async def test_empty_inputs_are_returned_without_calling_mls( + input_type: Literal["request", "response"], respx_mock: respx.MockRouter +) -> None: + route = _serve(respx_mock, _mls_body()) + inputs: GenericGuardrailAPIInputs = {"texts": []} + + assert await _apply(_guardrail(), inputs, input_type=input_type) is inputs + assert route.call_count == 0, "empty input must not be billed as an MLS call" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("input_type", ["request"]) +@pytest.mark.parametrize("block_on_error", [None, True], ids=["default-fail-open", "fail-closed"]) +@pytest.mark.parametrize( + ("mls_reply", "reason"), + [ + pytest.param(httpx.ConnectError("connection refused"), "connection refused", id="connection-error"), + pytest.param( + httpx.Response(HTTPStatus.SERVICE_UNAVAILABLE), + f"{HTTPStatus.SERVICE_UNAVAILABLE.value} {HTTPStatus.SERVICE_UNAVAILABLE.phrase}", + id="server-error", + ), + pytest.param( + httpx.Response(HTTPStatus.OK, json={"results": "not a list"}), + "Invalid RealmLabs guardrail response", + id="invalid-body", + ), + ], +) +async def test_mls_failures_follow_the_error_policy( + mls_reply: httpx.Response | httpx.HTTPError, + reason: str, + block_on_error: bool | None, + input_type: Literal["request", "response"], + respx_mock: respx.MockRouter, + caplog: pytest.LogCaptureFixture, +) -> None: + route: Final = respx_mock.post(_URL) + if isinstance(mls_reply, httpx.HTTPError): + route.mock(side_effect=mls_reply) + else: + route.mock(return_value=mls_reply) + inputs: Final[GenericGuardrailAPIInputs] = {"texts": ["hello"]} + guardrail: Final = _guardrail(block_on_error=block_on_error) + + if block_on_error: + with pytest.raises(GuardrailRaisedException) as exc: + await _apply(guardrail, inputs, input_type=input_type) + assert exc.value.blocked_content is False + assert reason in exc.value.message + else: + assert await _apply(guardrail, inputs, input_type=input_type) is inputs + assert reason in caplog.text + + +@pytest.mark.asyncio +@pytest.mark.parametrize("input_type", ["request"]) +@pytest.mark.parametrize("block_on_error", [False, True]) +@pytest.mark.parametrize( + "body", + [ + pytest.param({}, id="empty-object"), + pytest.param({"choices": [{"message": {"content": "hello"}}]}, id="chat-envelope"), + pytest.param({"pii_spans": []}, id="missing-results"), + pytest.param( + {"results": [{"probe": "hazard_prompt", "probability": 0.99}], "pii_spans": []}, id="renamed-score" + ), + pytest.param({"results": [{"prob": 0.99}], "pii_spans": []}, id="missing-probe"), + pytest.param( + {"results": [{"probe": "hazard_prompt", "prob": 0.99, "role_mismatch": "true"}], "pii_spans": []}, + id="invalid-role-mismatch", + ), + ], +) +async def test_incomplete_verdicts_follow_the_error_policy( + body: dict[str, object], + block_on_error: bool, + input_type: Literal["request", "response"], + respx_mock: respx.MockRouter, + caplog: pytest.LogCaptureFixture, +) -> None: + _serve(respx_mock, body) + inputs: Final[GenericGuardrailAPIInputs] = {"texts": ["Hello Alex."]} + guardrail: Final = _guardrail(block_on_error=block_on_error) + + if block_on_error: + with pytest.raises(GuardrailRaisedException) as exc: + await _apply(guardrail, inputs, input_type=input_type) + assert exc.value.blocked_content is False, exc.value.message + assert "Invalid RealmLabs guardrail response" in exc.value.message, exc.value.message + else: + assert await _apply(guardrail, inputs, input_type=input_type) is inputs + assert "Invalid RealmLabs guardrail response" in caplog.text, caplog.text + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "prob", + [ + pytest.param(None, id="null"), + pytest.param("0.99", id="string"), + pytest.param(True, id="boolean"), + pytest.param(-0.1, id="negative"), + pytest.param(1.1, id="above-one"), + pytest.param(float("nan"), id="nan"), + pytest.param(float("inf"), id="infinity"), + ], +) +async def test_invalid_probabilities_cannot_bypass_fail_closed(prob: object, respx_mock: respx.MockRouter) -> None: + body: Final = {"results": [{"probe": "hazard_prompt", "prob": prob}], "pii_spans": []} + respx_mock.post(_URL).mock(return_value=httpx.Response(HTTPStatus.OK, content=json.dumps(body))) + + with pytest.raises(GuardrailRaisedException) as exc: + await _apply(_guardrail(block_on_error=True), {"texts": ["hello"]}) + + assert exc.value.blocked_content is False, exc.value.message + assert "Invalid RealmLabs guardrail response" in exc.value.message, exc.value.message + + +@pytest.mark.asyncio +@pytest.mark.parametrize("input_type", ["request"]) +async def test_empty_verdict_arrays_are_valid_in_fail_closed_mode( + input_type: Literal["request", "response"], respx_mock: respx.MockRouter +) -> None: + _serve(respx_mock, {"results": [], "pii_spans": []}) + inputs: Final[GenericGuardrailAPIInputs] = {"texts": ["hello"]} + + assert await _apply(_guardrail(block_on_error=True), inputs, input_type=input_type) is inputs + + +@pytest.mark.asyncio +@pytest.mark.parametrize("hazard", [False, True]) +async def test_additional_response_fields_preserve_policy_enforcement( + hazard: bool, respx_mock: respx.MockRouter +) -> None: + _serve( + respx_mock, + { + "future_metadata": {"version": 2}, + "results": [{"probe": "hazard_prompt", "prob": 1.0 if hazard else 0.0, "future_field": [1, 2]}], + "pii_spans": [], + }, + ) + guardrail: Final = _guardrail(hazard_threshold=0.5, block_on_error=True) + inputs: Final[GenericGuardrailAPIInputs] = {"texts": ["Hello Alex."]} + + if hazard: + with pytest.raises(GuardrailRaisedException) as exc: + await _apply(guardrail, inputs) + assert exc.value.blocked_content is True, exc.value.message + assert "hazard_prompt" in exc.value.message, exc.value.message + else: + assert await _apply(guardrail, inputs) is inputs + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("api_key", "api_base", "expected_key", "expected_base"), + [ + pytest.param(None, None, "env_key", "https://env.example.test", id="environment-fallback"), + pytest.param(_API_KEY, f"{_API_BASE}/", _API_KEY, _API_BASE, id="config-overrides-environment"), + ], +) +async def test_credentials_resolve_config_before_environment( + api_key: str | None, + api_base: str | None, + expected_key: str, + expected_base: str, + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setenv("REALMLABS_API_KEY", "env_key") + monkeypatch.setenv("REALMLABS_API_BASE", "https://env.example.test/") + route: Final = _serve(respx_mock, _mls_body(), f"{expected_base}/guardrail") + + await _apply(RealmLabsGuardrail(api_key=api_key, api_base=api_base), {"texts": ["hello"]}) + + assert route.call_count == 1 + assert route.calls.last.request.headers["Authorization"] == f"Bearer {expected_key}" + + +def test_missing_api_key_is_rejected_at_startup() -> None: + with pytest.raises(RealmLabsMissingCredentials): + RealmLabsGuardrail() + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("settings", "expected_probes", "expected_thinking"), + [ + pytest.param( + {"probes": ["hazard_prompt", "dispute"], "enable_thinking": True}, + ["hazard_prompt", "dispute"], + True, + id="top-level", + ), + pytest.param( + {"optional_params": {"probes": "all", "enable_thinking": True}}, + "all", + True, + id="nested", + ), + pytest.param( + { + "probes": ["hazard_prompt"], + "enable_thinking": True, + "optional_params": {"probes": [], "enable_thinking": False}, + }, + [], + False, + id="nested-false-and-empty-list-win", + ), + pytest.param( + { + "probes": "all", + "enable_thinking": True, + "optional_params": {"probes": None, "enable_thinking": None}, + }, + "all", + True, + id="nested-null-falls-back", + ), + pytest.param( + {"probes": "all", "enable_thinking": True, "optional_params": {}}, + "all", + True, + id="empty-options-keep-top-level", + ), + ], +) +async def test_config_yaml_settings_reach_mls( + settings: dict[str, object], + expected_probes: list[str] | str, + expected_thinking: bool, + respx_mock: respx.MockRouter, +) -> None: + guardrail: Final = _configured_guardrail(settings) + route: Final = _serve(respx_mock, _mls_body()) + + await _apply(guardrail, {"texts": ["hello"]}) + + assert route.calls.last.request.headers["Authorization"] == f"Bearer {_API_KEY}" + assert _sent_body(route) == { + "messages": [{"role": "user", "content": "hello"}], + "probes": expected_probes, + "pii": False, + "enable_thinking": expected_thinking, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "settings", + [ + pytest.param({"hazard_threshold": 0.0}, id="top-level-zero"), + pytest.param({"optional_params": {"hazard_threshold": 0.0}}, id="nested-zero"), + pytest.param({"hazard_threshold": 0.99, "optional_params": {"hazard_threshold": 0.0}}, id="nested-zero-wins"), + ], +) +async def test_configured_zero_hazard_threshold_blocks_a_positive_score( + settings: dict[str, object], respx_mock: respx.MockRouter +) -> None: + guardrail: Final = _configured_guardrail(settings) + _serve(respx_mock, _mls_body(hazard=0.01)) + + with pytest.raises(GuardrailRaisedException) as exc: + await _apply(guardrail, {"texts": ["hello"]}) + + assert exc.value.blocked_content is True + assert "threshold=0.0" in exc.value.message, exc.value.message + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "settings", + [ + pytest.param({"hazard_threshold": 0.99, "optional_params": {}}, id="omitted-nested"), + pytest.param({"hazard_threshold": 0.99, "optional_params": {"hazard_threshold": None}}, id="null-nested"), + pytest.param({"hazard_threshold": 0.0, "optional_params": {"hazard_threshold": 0.99}}, id="nested-wins"), + ], +) +async def test_configured_higher_hazard_threshold_allows_the_request( + settings: dict[str, object], respx_mock: respx.MockRouter +) -> None: + guardrail: Final = _configured_guardrail(settings) + _serve(respx_mock, _mls_body(hazard=0.9)) + inputs: Final[GenericGuardrailAPIInputs] = {"texts": ["hello"]} + + assert await _apply(guardrail, inputs) is inputs + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "settings", + [ + pytest.param({"block_on_error": True}, id="top-level"), + pytest.param({"optional_params": {"block_on_error": True}}, id="nested"), + pytest.param({"block_on_error": False, "optional_params": {"block_on_error": True}}, id="nested-wins"), + pytest.param({"block_on_error": True, "optional_params": {"block_on_error": None}}, id="null-nested"), + ], +) +async def test_configured_fail_closed_blocks_an_mls_outage( + settings: dict[str, object], respx_mock: respx.MockRouter +) -> None: + guardrail: Final = _configured_guardrail(settings) + respx_mock.post(_URL).mock(side_effect=httpx.ConnectError("connection refused")) + + with pytest.raises(GuardrailRaisedException) as exc: + await _apply(guardrail, {"texts": ["hello"]}) + + assert exc.value.blocked_content is False + assert "block_on_error=True" in exc.value.message, exc.value.message + + +@pytest.mark.asyncio +async def test_nested_fail_open_overrides_top_level_fail_closed(respx_mock: respx.MockRouter) -> None: + guardrail: Final = _configured_guardrail({"block_on_error": True, "optional_params": {"block_on_error": False}}) + respx_mock.post(_URL).mock(side_effect=httpx.ConnectError("connection refused")) + inputs: Final[GenericGuardrailAPIInputs] = {"texts": ["hello"]} + + assert await _apply(guardrail, inputs) is inputs + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("settings", "expected_timeout"), + [ + pytest.param({"optional_params": None}, None, id="default-no-options"), + pytest.param({"optional_params": {}}, None, id="default-empty-options"), + pytest.param({"optional_params": {"enable_thinking": True}}, None, id="default-thinking-enabled"), + pytest.param({"optional_params": {"timeout": None}}, None, id="default-null-timeout"), + pytest.param({"timeout": 2.0, "optional_params": None}, 2.0, id="top-level-no-options"), + pytest.param({"timeout": 2.0, "optional_params": {}}, 2.0, id="top-level-empty-options"), + pytest.param( + {"timeout": 2.0, "optional_params": {"enable_thinking": True}}, 2.0, id="top-level-thinking-enabled" + ), + pytest.param({"timeout": 2.0, "optional_params": {"timeout": None}}, 2.0, id="top-level-null-timeout"), + pytest.param({"timeout": 0.0, "optional_params": None}, 0.0, id="zero-no-options"), + pytest.param({"timeout": 0.0, "optional_params": {}}, 0.0, id="zero-empty-options"), + pytest.param({"timeout": 0.0, "optional_params": {"enable_thinking": True}}, 0.0, id="zero-thinking-enabled"), + pytest.param({"timeout": 0.0, "optional_params": {"timeout": None}}, 0.0, id="zero-null-timeout"), + pytest.param({"timeout": 2.0, "optional_params": {"timeout": 3.0}}, 3.0, id="nested-overrides-top-level"), + pytest.param({"timeout": 2.0, "optional_params": {"timeout": 10.0}}, 10.0, id="explicit-nested-ten-seconds"), + ], +) +async def test_configured_timeout_reaches_mls( + settings: dict[str, object], expected_timeout: float | None, respx_mock: respx.MockRouter +) -> None: + guardrail: Final = _configured_guardrail(settings) + route: Final = _serve(respx_mock, _mls_body()) + + await _apply(guardrail, {"texts": ["hello"]}) + + expected: Final = RealmLabsGuardrailOptionalParams().timeout if expected_timeout is None else expected_timeout + extensions: Final = _JSON_OBJECT.validate_python(route.calls.last.request.extensions) + timeouts: Final = _JSON_OBJECT.validate_python(extensions["timeout"]) + assert timeouts["read"] == expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mismatch_fields", [{}, {"role_mismatch": None}], ids=["omitted", "null"]) +async def test_missing_role_mismatch_still_enforces_request_hazard( + mismatch_fields: dict[str, object], respx_mock: respx.MockRouter +) -> None: + _serve(respx_mock, {"results": [{"probe": "hazard_prompt", "prob": 0.99, **mismatch_fields}], "pii_spans": []}) + + with pytest.raises(GuardrailRaisedException) as exc: + await _apply(_guardrail(), {"texts": ["a hazardous request"]}) + + assert exc.value.blocked_content is True + assert "hazard_prompt" in exc.value.message, exc.value.message + + +@pytest.mark.asyncio +async def test_response_text_is_returned_without_calling_mls(respx_mock: respx.MockRouter) -> None: + route: Final = _serve(respx_mock, _mls_body()) + inputs: Final[GenericGuardrailAPIInputs] = {"texts": ["hello"]} + + assert await _apply(_guardrail(), inputs, input_type="response") is inputs + assert route.call_count == 0 diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index cc0b081b670..b6b7beaf86f 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -36234,6 +36234,11 @@ export interface components { * @default false */ disable_exception_on_block: boolean | null; + /** + * Enable Thinking + * @description Whether MLS should render the chat template in thinking mode. Defaults to False. + */ + enable_thinking?: boolean | null; /** * End Session After N Fails * @description For /v1/realtime sessions: automatically close the session after this many guardrail violations. @@ -36306,6 +36311,11 @@ export interface components { * @description Enable hallucination detection to detect factual inaccuracies. */ hallucinations_check?: boolean | null; + /** + * Hazard Threshold + * @description Block the request when the hazard_prompt probe scores strictly above this value. Defaults to 0.703, the threshold MLS reports for that probe. Note the probe also responds to instruction-style phrasing such as "repeat this back verbatim", so raise this if benign traffic is being blocked. + */ + hazard_threshold?: number | null; /** * Include Evidence * @description Include detailed evidence payloads in responses (sets `plr_evidence` header). @@ -36558,6 +36568,11 @@ export interface components { presidio_score_thresholds?: { [key: string]: number; } | null; + /** + * Probes + * @description Which classifier probes to run: a list of probe names, or "all". Defaults to ["hazard_prompt"] - the only probe whose score this guardrail enforces. An unknown probe name makes MLS return 404. + */ + probes?: string[] | string | null; /** * Project Id * @description Project ID for the Lakera AI project From 6d37d3379d2c1ebc283051f7d7247bdec02ad9d0 Mon Sep 17 00:00:00 2001 From: Gourav Mittal Date: Thu, 1 Oct 2026 15:15:18 -0700 Subject: [PATCH 2/4] feat(guardrails): add RealmLabs PII masking and blocking Detect PII in requests and either mask literal matches or block the request Merge overlapping matches and replace spans against the original text to avoid partial exposure or repeated masking of inserted labels Validate PII results and include masking settings, configuration precedence, generated API schemas, and behavioral tests --- litellm/proxy/_lazy_openapi_snapshot.json | 24 ++ .../guardrail_hooks/realmlabs/__init__.py | 2 + .../guardrail_hooks/realmlabs/pii_masking.py | 106 +++++++++ .../guardrail_hooks/realmlabs/realmlabs.py | 46 +++- .../guardrails/guardrail_hooks/realmlabs.py | 34 +++ .../realmlabs/test_pii_masking.py | 223 ++++++++++++++++++ .../realmlabs/test_realmlabs.py | 188 ++++++++++++++- ui/litellm-dashboard/src/lib/http/schema.d.ts | 10 + 8 files changed, 616 insertions(+), 17 deletions(-) create mode 100644 litellm/proxy/guardrails/guardrail_hooks/realmlabs/pii_masking.py create mode 100644 tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_pii_masking.py diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index e09cceb501c..4ea6c83f401 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -14084,6 +14084,18 @@ "description": "Controls Pillar session persistence (sets `plr_persist` header). Set to False to disable persistence.", "title": "Persist Session" }, + "pii": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "description": "Whether to run MLS's PII detection head. Defaults to True.", + "title": "Pii" + }, "pii_check": { "anyOf": [ { @@ -14126,6 +14138,18 @@ "description": "Configuration for PII entity types and actions", "title": "Pii Entities Config" }, + "pii_mask": { + "anyOf": [ + { + "type": "boolean" + }, + { + "type": "null" + } + ], + "description": "What to do with detected PII. True (default) rewrites each span as its type in brackets, e.g. \"My name is Alex\" -> \"My name is [name]\", and lets the request through. False blocks the request instead.", + "title": "Pii Mask" + }, "policy_id": { "anyOf": [ { diff --git a/litellm/proxy/guardrails/guardrail_hooks/realmlabs/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/__init__.py index 4d06c5af79e..2b162b09659 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/realmlabs/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/__init__.py @@ -24,6 +24,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" api_base=litellm_params.api_base, probes=settings.probes, hazard_threshold=settings.hazard_threshold, + pii=settings.pii, + pii_mask=settings.pii_mask, block_on_error=settings.block_on_error, enable_thinking=settings.enable_thinking, timeout=settings.timeout, diff --git a/litellm/proxy/guardrails/guardrail_hooks/realmlabs/pii_masking.py b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/pii_masking.py new file mode 100644 index 00000000000..d20292cf860 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/pii_masking.py @@ -0,0 +1,106 @@ +"""Mask literal PII matches after merging overlaps in the original text.""" + +from __future__ import annotations + +import re +from collections.abc import Iterable, Iterator, Mapping, Sequence +from dataclasses import dataclass +from heapq import merge +from itertools import groupby +from types import MappingProxyType +from typing import TYPE_CHECKING, Final + +if TYPE_CHECKING: + from litellm.types.proxy.guardrails.guardrail_hooks.realmlabs import RealmLabsPIISpan + + +@dataclass(frozen=True, slots=True) +class _PIIMatch: + start: int + end: int + entity_type: str + + +def _unique_pii_values(spans: Sequence[RealmLabsPIISpan]) -> Iterator[tuple[str, str]]: + """Yield one label per literal value; conflicting types become ``pii``.""" + + pairs: Final = ( + (value, entity_type) for span in spans if (value := span.get("text")) and (entity_type := span.get("type")) + ) + for value, detections in groupby(sorted(frozenset(pairs)), key=lambda pair: pair[0]): + entity_types = tuple(entity_type for _, entity_type in detections) + yield value, entity_types[0] if len(entity_types) == 1 else "pii" + + +def _matches_for_value(text: str, value: str, entity_type: str) -> Iterator[_PIIMatch]: + """Yield literal occurrences in position order, including overlapping occurrences.""" + + length: Final = len(value) + start = text.find(value) # rebind-ok: advancing search cursor finds overlaps without rescanning earlier positions + while start != -1: + yield _PIIMatch(start, start + length, entity_type) + start = text.find(value, start + 1) + + +def _matches_for_group(text: str, values: Sequence[str], entity_types: Mapping[str, str]) -> Iterator[_PIIMatch]: + """Search values sharing their first character together, longest first at each position.""" + + if len(values) == 1: + yield from _matches_for_value(text, values[0], entity_types[values[0]]) + return + + alternatives: Final = "|".join(re.escape(value) for value in sorted(values, key=len, reverse=True)) + pattern: Final = re.compile(alternatives) + match = pattern.search(text) # rebind-ok: advance the search cursor while retaining overlaps + while match is not None: + yield _PIIMatch(match.start(), match.end(), entity_types[match.group()]) + match = pattern.search(text, match.start() + 1) + + +def _merged_pii_matches(matches: Iterable[_PIIMatch]) -> Iterator[_PIIMatch]: + """Merge ordered overlaps; keep an enclosing type, otherwise label the union ``pii``.""" + + remaining: Final = iter(matches) + region = next(remaining, None) # rebind-ok: keep one pending region while consuming the ordered stream + if region is None: + return + + for match in remaining: + if match.start >= region.end: + yield region + region = match + elif match.end > region.end: + region = _PIIMatch(region.start, match.end, "pii") + + yield region + + +def _masked_parts(text: str, regions: Iterable[_PIIMatch]) -> Iterator[str]: + """Yield unchanged gaps and one mask per region, then the trailing text.""" + + previous_end = 0 # rebind-ok: rendering cursor tracks the next unchanged slice without storing all regions + for region in regions: + yield text[previous_end : region.start] + yield f"[{region.entity_type}]" + previous_end = region.end + + yield text[previous_end:] + + +def mask_pii_in_text(text: str, spans: Sequence[RealmLabsPIISpan]) -> str: + """Find original-text matches, merge overlaps, and render each masked region once. + + MLS offsets describe its rendering of the whole conversation, so local positions come from literal text. + Group values by their first character to reduce repeated scans while preserving a literal regex prefix. + Each group emits its longest match at a given start; shorter matches at that start are fully contained. + With G groups and M emitted matches, ordering costs O(M log(G + 1)) time and O(G) space. Overlap merging + is O(M). Pattern preparation, searches, and output assembly have their own costs; regex search time + depends on the values and input text. + """ + + placeholders: Final = {f"[{label}]": label for span in spans if (label := span.get("type"))} + entity_types: Final = MappingProxyType({"[pii]": "pii", **placeholders, **dict(_unique_pii_values(spans))}) + groups: Final = groupby(sorted(entity_types), key=lambda value: value[0]) + streams: Final = (_matches_for_group(text, tuple(values), entity_types) for _, values in groups) + matches: Final = merge(*streams, key=lambda match: (match.start, -match.end)) + return "".join(_masked_parts(text, _merged_pii_matches(matches))) diff --git a/litellm/proxy/guardrails/guardrail_hooks/realmlabs/realmlabs.py b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/realmlabs.py index 9ccc142f013..32aa9cedde2 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/realmlabs/realmlabs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/realmlabs.py @@ -1,4 +1,4 @@ -"""RealmLabs MLS request hazard guardrail.""" +"""RealmLabs MLS request guardrail with hazard screening and PII protection.""" from __future__ import annotations @@ -20,6 +20,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # helper is untyped in http_handler httpxSpecialProvider, ) +from litellm.proxy.guardrails.guardrail_hooks.realmlabs.pii_masking import mask_pii_in_text from litellm.secret_managers.main import get_secret_str from litellm.types.guardrails import GuardrailEventHooks from litellm.types.proxy.guardrails.guardrail_hooks.realmlabs import ( @@ -27,6 +28,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.realmlabs import ( RealmLabsGuardrailConfigModel, RealmLabsGuardrailRequest, RealmLabsGuardrailResponse, + RealmLabsPIISpan, ) if TYPE_CHECKING: @@ -56,7 +58,7 @@ class _InvalidResponse: class RealmLabsGuardrail(CustomGuardrail): - """Blocks hazardous prompts using the RealmLabs MLS endpoint.""" + """Blocks hazardous prompts and masks PII using the RealmLabs MLS endpoint.""" def __init__( self, @@ -64,6 +66,8 @@ class RealmLabsGuardrail(CustomGuardrail): api_base: str | None = None, probes: Sequence[str] | str | None = None, hazard_threshold: float | None = None, + pii: bool | None = None, + pii_mask: bool | None = None, block_on_error: bool | None = None, enable_thinking: bool | None = None, timeout: float | None = None, @@ -84,6 +88,8 @@ class RealmLabsGuardrail(CustomGuardrail): self.api_base = (api_base or get_secret_str("REALMLABS_API_BASE") or _DEFAULT_API_BASE).rstrip("/") self.probes: Sequence[str] | str = (_HAZARD_PROBE,) if probes is None else probes self.hazard_threshold = _DEFAULT_HAZARD_THRESHOLD if hazard_threshold is None else hazard_threshold + self.pii = True if pii is None else pii + self.pii_mask = True if pii_mask is None else pii_mask self.block_on_error = False if block_on_error is None else block_on_error self.enable_thinking = False if enable_thinking is None else enable_thinking self.timeout = _DEFAULT_TIMEOUT if timeout is None else timeout @@ -115,11 +121,18 @@ class RealmLabsGuardrail(CustomGuardrail): return None if result.get("role_mismatch") else result["prob"] return None + @staticmethod + def _span_types(spans: Sequence[RealmLabsPIISpan]) -> str: + """Distinct span types in first-seen order, for block messages and logs; ``"unknown"`` if none.""" + entity_types: Final = (span.get("type") for span in spans) + unique_types: Final = tuple(dict.fromkeys(entity_type for entity_type in entity_types if entity_type)) + return ", ".join(unique_types) or "unknown" + def _build_request(self, messages: Sequence[Mapping[str, object]]) -> RealmLabsGuardrailRequest: return RealmLabsGuardrailRequest( messages=messages, probes=self.probes, - pii=False, + pii=self.pii, enable_thinking=self.enable_thinking, ) @@ -129,6 +142,10 @@ class RealmLabsGuardrail(CustomGuardrail): except ValidationError as exc: return _InvalidResponse(exc.json(include_input=False, include_context=False, include_url=False)) + if any(not span.get("type") for span in result["pii_spans"]): + return _InvalidResponse("pii_spans must include a nonempty type") + if self.pii_mask and any(not span.get("text") for span in result["pii_spans"]): + return _InvalidResponse("pii_spans must include nonempty text when masking is enabled") return result async def _call_mls( @@ -136,10 +153,11 @@ class RealmLabsGuardrail(CustomGuardrail): ) -> RealmLabsGuardrailResponse | _InvalidResponse: endpoint: Final = f"{self.api_base}{_GUARDRAIL_ENDPOINT}" verbose_proxy_logger.debug( - "RealmLabs MLS: %s msgs=%d probes=%s", + "RealmLabs MLS: %s msgs=%d probes=%s pii=%s", endpoint, len(messages), self.probes, + self.pii, ) response: Final[HttpxResponse] = await self.async_handler.post( # pyright: ignore[reportUnknownMemberType] # AsyncHTTPHandler.post is untyped url=endpoint, @@ -170,7 +188,7 @@ class RealmLabsGuardrail(CustomGuardrail): input_type: Literal["request", "response"], logging_obj: LiteLLMLoggingObj | None = None, ) -> GenericGuardrailAPIInputs: - """Screen requests for hazardous content through LiteLLM's unified guardrail layer.""" + """Screen requests for hazardous content and PII through LiteLLM's unified guardrail layer.""" if input_type != "request": return inputs texts: Final = tuple(inputs.get("texts") or ()) @@ -188,6 +206,7 @@ class RealmLabsGuardrail(CustomGuardrail): if isinstance(result, _InvalidResponse): return self._handle_mls_error(inputs, f"Invalid RealmLabs guardrail response: {result.reason}") + # Hazard is checked before PII, so a hazardous prompt is rejected rather than masked and forwarded. hazard_score: Final = self._hazard_score(result) if hazard_score is not None and hazard_score > self.hazard_threshold: verbose_proxy_logger.warning( @@ -205,4 +224,19 @@ class RealmLabsGuardrail(CustomGuardrail): blocked_content=True, ) - return inputs + spans: Final = result["pii_spans"] + if not spans: + return inputs + + if not self.pii_mask: + raise GuardrailRaisedException( + guardrail_name=self.guardrail_name, + message=f"Blocked by RealmLabs: PII detected in the {input_type} ({self._span_types(spans)})", + blocked_content=True, + ) + + masked_texts: Final = tuple(mask_pii_in_text(text, spans) for text in texts) + if masked_texts == texts: + return inputs + verbose_proxy_logger.debug("RealmLabs MLS masked PII types in the %s: %s", input_type, self._span_types(spans)) + return {**inputs, "texts": list(masked_texts)} diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/realmlabs.py b/litellm/types/proxy/guardrails/guardrail_hooks/realmlabs.py index 9c7c8072b33..97d8aa13ae1 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/realmlabs.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/realmlabs.py @@ -34,11 +34,25 @@ class RealmLabsProbeResult(TypedDict, total=False): role_mismatch: ReadOnly[Annotated[bool, Field(strict=True)] | None] +class RealmLabsPIISpan(TypedDict, total=False): + """One detected PII span. + + ``start``/``end`` index MLS's rendering of the whole conversation, not a single message, so they are not + used for masking; ``text`` is matched within each message instead. + """ + + type: ReadOnly[Annotated[str, Field(strict=True, min_length=1)]] + text: ReadOnly[str | None] + start: ReadOnly[int | None] + end: ReadOnly[int | None] + + class RealmLabsGuardrailResponse(TypedDict, total=False): """Response body of ``POST {api_base}/guardrail``. MLS is stateless, so it carries no turn id.""" results: ReadOnly[Required[Sequence[RealmLabsProbeResult]]] focal_role: ReadOnly[str | None] + pii_spans: ReadOnly[Required[Sequence[RealmLabsPIISpan]]] class RealmLabsGuardrailOptionalParams(BaseModel): @@ -52,6 +66,14 @@ class RealmLabsGuardrailOptionalParams(BaseModel): default=None, description="Block hazard scores strictly above this value. Overrides top-level hazard_threshold when supplied.", ) + pii: bool | None = Field( + default=None, + description="Whether to run PII detection. Overrides top-level pii when supplied.", + ) + pii_mask: bool | None = Field( + default=None, + description="Mask detected PII when true, otherwise block. Overrides top-level pii_mask when supplied.", + ) block_on_error: bool | None = Field( default=None, description="Whether to block when MLS fails. Overrides top-level block_on_error when supplied.", @@ -103,6 +125,18 @@ class RealmLabsGuardrailConfigModel(GuardrailConfigModel[RealmLabsGuardrailOptio 'back verbatim", so raise this if benign traffic is being blocked.' ), ) + pii: bool | None = Field( + default=None, + description=("Whether to run MLS's PII detection head. Defaults to True."), + ) + pii_mask: bool | None = Field( + default=None, + description=( + "What to do with detected PII. True (default) rewrites each span as its " + 'type in brackets, e.g. "My name is Alex" -> "My name is [name]", and ' + "lets the request through. False blocks the request instead." + ), + ) block_on_error: bool | None = Field( default=None, description=( diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_pii_masking.py b/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_pii_masking.py new file mode 100644 index 00000000000..48ed9f1ecba --- /dev/null +++ b/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_pii_masking.py @@ -0,0 +1,223 @@ +from collections.abc import Sequence +from typing import Final + +import pytest + +from litellm.proxy.guardrails.guardrail_hooks.realmlabs.pii_masking import mask_pii_in_text +from litellm.types.proxy.guardrails.guardrail_hooks.realmlabs import RealmLabsPIISpan + + +@pytest.mark.parametrize( + ("text", "spans", "expected"), + [ + pytest.param( + "Contact Ann Smith Jones today", + [{"type": "name", "text": "Ann Smith"}, {"type": "name", "text": "Smith Jones"}], + "Contact [pii] today", + id="partial-overlap", + ), + pytest.param( + "Contact Ann Smith Jones Lee today", + [ + {"type": "name", "text": "Ann Smith"}, + {"type": "name", "text": "Smith Jones"}, + {"type": "name", "text": "Jones Lee"}, + ], + "Contact [pii] today", + id="chain-of-overlaps", + ), + pytest.param( + "Contact Ann Smith at Ann.Smith@example.com", + [ + {"type": "name", "text": "Smith"}, + {"type": "name", "text": "Ann Smith"}, + {"type": "email", "text": "Ann.Smith@example.com"}, + ], + "Contact [name] at [email]", + id="contained-match-starts-later", + ), + pytest.param( + "Ann Smith Jones", + [ + {"type": "name", "text": "Ann Smith"}, + {"type": "name", "text": "Smith Jones"}, + {"type": "name", "text": "Ann Smith Jones"}, + ], + "[name]", + id="enclosing-match-keeps-its-label", + ), + pytest.param( + "AnnAnn.Smith@example.com", + [{"type": "name", "text": "Ann"}, {"type": "email", "text": "Ann.Smith@example.com"}], + "[name][email]", + id="adjacent-matches-stay-separate", + ), + pytest.param( + "ababa", + [{"type": "name", "text": "aba"}], + "[pii]", + id="overlapping-occurrences-of-one-value", + ), + pytest.param( + "ababaca", + [{"type": "name", "text": "ababa"}, {"type": "name", "text": "abaca"}], + "[pii]", + id="overlap-between-values-with-the-same-first-character", + ), + pytest.param( + "Contact Ann today", + [{"type": "name", "text": "Ann"}, {"type": "name", "text": "Ann"}], + "Contact [name] today", + id="duplicate-detections", + ), + pytest.param( + "Contact Ann today", + [{"type": "name", "text": "Ann"}, {"type": "username", "text": "Ann"}], + "Contact [pii] today", + id="same-match-with-conflicting-types", + ), + pytest.param( + "Ann1 Ann Annn Ann[1] Ann+", + [{"type": "name", "text": "Ann[1]"}, {"type": "username", "text": "Ann+"}], + "Ann1 Ann Annn [name] [username]", + id="regex-punctuation-is-literal", + ), + pytest.param( + "👋 Éva / Éva@example.com", + [{"type": "name", "text": "Éva"}, {"type": "email", "text": "Éva@example.com"}], + "👋 [name] / [email]", + id="unicode-positions-and-shorter-match-fallback", + ), + pytest.param( + "Ann Anna", + [ + {"type": "name", "text": "Ann"}, + {"type": "username", "text": "Ann"}, + {"type": "name", "text": "Anna"}, + ], + "[pii] [name]", + id="conflicting-inner-types-do-not-change-enclosing-type", + ), + pytest.param( + "[name] Alex [Alex] [unknown]", + ({"type": "name", "text": "[name"}, {"type": "name", "text": "Alex"}), + "[name] [name] [[name]] [unknown]", + id="real-pii-and-arbitrary-brackets-are-not-exempt", + ), + pytest.param( + "Alex[name]", + ({"type": "name", "text": "Alex[na"},), + "[pii]", + id="overlap-starts-before-placeholder", + ), + pytest.param( + "[name]Alex", + ({"type": "name", "text": "me]Alex"},), + "[pii]", + id="overlap-ends-after-placeholder", + ), + pytest.param( + "[name]Alex[name]", + ({"type": "name", "text": "me]Alex[na"},), + "[pii]", + id="overlap-joins-two-placeholders", + ), + pytest.param( + "Alex[name]Jones", + ({"type": "name", "text": "Alex[name]Jones"},), + "[name]", + id="real-pii-encloses-placeholder", + ), + pytest.param( + "[name]alex@example.com", + ({"type": "name", "text": "[name"}, {"type": "email", "text": "alex@example.com"}), + "[name][email]", + id="adjacent-pii-stays-separate", + ), + pytest.param( + "[name] [pii] name", + ({"type": "name", "text": "name"}, {"type": "username", "text": "name"}), + "[name] [pii] [pii]", + id="conflicting-types-preserve-existing-placeholders", + ), + pytest.param( + "[pii] [name", + ({"type": "name", "text": "pii"}, {"type": "name", "text": "[name"}), + "[pii] [name]", + id="generic-placeholder-and-incomplete-brackets", + ), + pytest.param( + "👋 [custom.type+] Éva", + ({"type": "custom.type+", "text": "type+"}, {"type": "name", "text": "Éva"}), + "👋 [custom.type+] [name]", + id="free-form-types-and-unicode", + ), + pytest.param( + "alex@example.com email", + ({"type": "email", "text": "alex@example.com"}, {"type": "name", "text": "email"}), + "[email] [name]", + id="inserted-labels-are-not-remasked", + ), + pytest.param("Ann is here", (), "Ann is here", id="no-detections"), + pytest.param( + "Ann is here", + ( + RealmLabsPIISpan(), + {"text": "Ann"}, + {"type": "name", "text": None}, + {"type": "name", "text": ""}, + {"type": "", "text": "Ann"}, + ), + "Ann is here", + id="incomplete-detections", + ), + pytest.param( + "Hello [name]! [name][name]", + ({"type": "name", "text": "[name"},), + "Hello [name]! [name][name]", + id="preserve-placeholder-missing-closing-bracket", + ), + pytest.param( + "Hello [name]! [name][name]", + ({"type": "name", "text": "name"},), + "Hello [name]! [name][name]", + id="preserve-placeholder-label-only", + ), + pytest.param( + "Hello [name]! [name][name]", + ({"type": "name", "text": "name]"},), + "Hello [name]! [name][name]", + id="preserve-placeholder-missing-opening-bracket", + ), + pytest.param( + "Hello [name]! [name][name]", + ({"type": "name", "text": "[name]"},), + "Hello [name]! [name][name]", + id="preserve-placeholder-whole-placeholder", + ), + pytest.param( + "Hello [name]! [name][name]", + ({"type": "name", "text": "nam"},), + "Hello [name]! [name][name]", + id="preserve-placeholder-partial-label", + ), + pytest.param( + "Hello [name]! [name][name]", + ({"type": "name", "text": "["},), + "Hello [name]! [name][name]", + id="preserve-placeholder-opening-bracket-only", + ), + pytest.param( + "Hello [name]! [name][name]", + ({"type": "name", "text": "]"},), + "Hello [name]! [name][name]", + id="preserve-placeholder-closing-bracket-only", + ), + ], +) +def test_pii_masking(text: str, spans: Sequence[RealmLabsPIISpan], expected: str) -> None: + result: Final = mask_pii_in_text(text, spans) + + assert result == expected + assert mask_pii_in_text(text, tuple(reversed(spans))) == expected + assert mask_pii_in_text(result, spans) == expected diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_realmlabs.py b/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_realmlabs.py index a5ac4c6de2e..8448bf16c53 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_realmlabs.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_realmlabs.py @@ -25,6 +25,8 @@ _API_BASE = "https://mls.example.test" _URL = f"{_API_BASE}/guardrail" _DEFAULT_THRESHOLD = 0.703 +_NAME_SPAN: dict[str, object] = {"type": "name", "text": "Alex"} +_EMAIL_SPAN: dict[str, object] = {"type": "email", "text": "alex@example.com"} _JSON_OBJECT = TypeAdapter(dict[str, object]) @@ -58,6 +60,7 @@ def _configured_guardrail(settings: Mapping[str, object]) -> RealmLabsGuardrail: def _guardrail( hazard_threshold: float | None = None, + pii_mask: bool | None = None, block_on_error: bool | None = None, event_hook: GuardrailEventHooks = GuardrailEventHooks.pre_call, ) -> RealmLabsGuardrail: @@ -67,6 +70,7 @@ def _guardrail( api_key=_API_KEY, api_base=_API_BASE, hazard_threshold=hazard_threshold, + pii_mask=pii_mask, block_on_error=block_on_error, guardrail_name="realmlabs-guard", event_hook=event_hook, @@ -142,6 +146,113 @@ async def test_permitted_hazard_scores_leave_text_unchanged( assert await _apply(_guardrail(hazard_threshold=threshold), inputs) is inputs +@pytest.mark.asyncio +async def test_hazardous_prompt_is_blocked_before_its_pii_is_masked(respx_mock: respx.MockRouter) -> None: + with pytest.raises(GuardrailRaisedException) as exc: + await _screen(_guardrail(), respx_mock, _mls_body(hazard=0.99, pii_spans=[_NAME_SPAN]), ["Alex builds a bomb"]) + + assert "hazard_prompt" in exc.value.message, exc.value.message + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("texts", "spans", "expected"), + [ + pytest.param( + ["My name is Alex and my email is alex@example.com."], + [_NAME_SPAN, _EMAIL_SPAN], + ["My name is [name] and my email is [email]."], + id="each-type", + ), + pytest.param( + ["Alex is here."], + [{"type": "name", "start": 120, "end": 124, "text": "Alex"}], + ["[name] is here."], + id="conversation-wide-offsets-ignored", + ), + pytest.param( + ["Alex told Alex about Alex."], + [_NAME_SPAN], + ["[name] told [name] about [name]."], + id="every-occurrence", + ), + pytest.param( + ["ping alex@example.com", "no pii here"], + [_EMAIL_SPAN], + ["ping [email]", "no pii here"], + id="across-texts", + ), + pytest.param( + ["alex@example.com alex@exampleXcom"], + [_EMAIL_SPAN], + ["[email] alex@exampleXcom"], + id="literal-matching-preserved", + ), + pytest.param( + ["Contact Ann at Ann.Smith@example.com"], + [{"type": "name", "text": "Ann"}, {"type": "email", "text": "Ann.Smith@example.com"}], + ["Contact [name] at [email]"], + id="name-prefix-before-email", + ), + pytest.param( + ["Contact Ann at Ann.Smith@example.com"], + [{"type": "email", "text": "Ann.Smith@example.com"}, {"type": "name", "text": "Ann"}], + ["Contact [name] at [email]"], + id="email-before-name-prefix", + ), + pytest.param( + ["My name is name"], + [{"type": "name", "text": "name"}], + ["My [name] is [name]"], + id="detected-text-in-its-own-label", + marks=pytest.mark.timeout(5), + ), + ], +) +async def test_pii_masking_returns_expected_texts( + texts: list[str], spans: list[dict[str, object]], expected: list[str], respx_mock: respx.MockRouter +) -> None: + result: Final = await _screen(_guardrail(), respx_mock, _mls_body(pii_spans=spans), texts) + + assert result == {"texts": expected} + + +@pytest.mark.asyncio +async def test_span_from_another_turn_leaves_the_prompt_as_is(respx_mock: respx.MockRouter) -> None: + inputs: GenericGuardrailAPIInputs = {"texts": ["nothing sensitive"]} + _serve(respx_mock, _mls_body(pii_spans=[{"type": "name", "text": "Bob"}])) + + assert await _apply(_guardrail(), inputs) is inputs + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("input_type", "texts", "spans", "expected_types"), + [ + pytest.param( + "request", ["Alex alex@example.com"], [_NAME_SPAN, _EMAIL_SPAN], "name, email", id="request-types" + ), + pytest.param("request", ["Hello Alex."], [{"type": "name"}], "name", id="missing-pii-text"), + pytest.param("request", ["Hello Alex."], [{"type": "name", "text": None}], "name", id="null-pii-text"), + pytest.param("request", ["Hello Alex."], [{"type": "name", "text": ""}], "name", id="empty-pii-text"), + ], +) +async def test_detected_pii_blocks_as_content_when_masking_is_disabled( + input_type: Literal["request", "response"], + texts: list[str], + spans: list[dict[str, object]], + expected_types: str, + respx_mock: respx.MockRouter, +) -> None: + _serve(respx_mock, _mls_body(pii_spans=spans)) + + with pytest.raises(GuardrailRaisedException) as exc: + await _apply(_guardrail(pii_mask=False, block_on_error=True), {"texts": texts}, input_type=input_type) + + assert exc.value.blocked_content is True + assert f"PII detected in the {input_type} ({expected_types})" in exc.value.message + + @pytest.mark.asyncio async def test_request_carries_the_conversation_and_settings_but_not_the_model(respx_mock: respx.MockRouter) -> None: conversation: list[AllMessageValues] = [ @@ -156,7 +267,7 @@ async def test_request_carries_the_conversation_and_settings_but_not_the_model(r assert _sent_body(route) == { "messages": conversation, "probes": ["hazard_prompt"], - "pii": False, + "pii": True, "enable_thinking": False, } @@ -170,7 +281,7 @@ async def test_plain_texts_are_sent_as_user_turns(respx_mock: respx.MockRouter) assert _sent_body(route) == { "messages": [{"role": "user", "content": "a"}, {"role": "user", "content": "b"}], "probes": ["hazard_prompt"], - "pii": False, + "pii": True, "enable_thinking": False, } @@ -234,13 +345,14 @@ async def test_mls_failures_follow_the_error_policy( @pytest.mark.asyncio @pytest.mark.parametrize("input_type", ["request"]) -@pytest.mark.parametrize("block_on_error", [False, True]) +@pytest.mark.parametrize("block_on_error", [False, True], ids=["fail-open", "fail-closed"]) @pytest.mark.parametrize( "body", [ pytest.param({}, id="empty-object"), pytest.param({"choices": [{"message": {"content": "hello"}}]}, id="chat-envelope"), pytest.param({"pii_spans": []}, id="missing-results"), + pytest.param({"results": []}, id="missing-pii-spans"), pytest.param( {"results": [{"probe": "hazard_prompt", "probability": 0.99}], "pii_spans": []}, id="renamed-score" ), @@ -249,6 +361,10 @@ async def test_mls_failures_follow_the_error_policy( {"results": [{"probe": "hazard_prompt", "prob": 0.99, "role_mismatch": "true"}], "pii_spans": []}, id="invalid-role-mismatch", ), + pytest.param(_mls_body(pii_spans=[{"text": "Alex"}]), id="missing-pii-type"), + pytest.param(_mls_body(pii_spans=[{"type": "name"}]), id="missing-pii-text"), + pytest.param(_mls_body(pii_spans=[{"type": "name", "text": None}]), id="null-pii-text"), + pytest.param(_mls_body(pii_spans=[{"type": "name", "text": ""}]), id="empty-pii-text"), ], ) async def test_incomplete_verdicts_follow_the_error_policy( @@ -308,7 +424,7 @@ async def test_empty_verdict_arrays_are_valid_in_fail_closed_mode( @pytest.mark.asyncio -@pytest.mark.parametrize("hazard", [False, True]) +@pytest.mark.parametrize("hazard", [False, True], ids=["mask-pii", "block-hazard"]) async def test_additional_response_fields_preserve_policy_enforcement( hazard: bool, respx_mock: respx.MockRouter ) -> None: @@ -317,7 +433,7 @@ async def test_additional_response_fields_preserve_policy_enforcement( { "future_metadata": {"version": 2}, "results": [{"probe": "hazard_prompt", "prob": 1.0 if hazard else 0.0, "future_field": [1, 2]}], - "pii_spans": [], + "pii_spans": [{"type": "name", "text": "Alex", "future_field": {"source": "test"}}], }, ) guardrail: Final = _guardrail(hazard_threshold=0.5, block_on_error=True) @@ -329,7 +445,7 @@ async def test_additional_response_fields_preserve_policy_enforcement( assert exc.value.blocked_content is True, exc.value.message assert "hazard_prompt" in exc.value.message, exc.value.message else: - assert await _apply(guardrail, inputs) is inputs + assert await _apply(guardrail, inputs) == {"texts": ["Hello [name]."]} @pytest.mark.asyncio @@ -368,13 +484,13 @@ def test_missing_api_key_is_rejected_at_startup() -> None: ("settings", "expected_probes", "expected_thinking"), [ pytest.param( - {"probes": ["hazard_prompt", "dispute"], "enable_thinking": True}, + {"probes": ["hazard_prompt", "dispute"], "pii": False, "enable_thinking": True}, ["hazard_prompt", "dispute"], True, id="top-level", ), pytest.param( - {"optional_params": {"probes": "all", "enable_thinking": True}}, + {"optional_params": {"probes": "all", "pii": False, "enable_thinking": True}}, "all", True, id="nested", @@ -382,8 +498,9 @@ def test_missing_api_key_is_rejected_at_startup() -> None: pytest.param( { "probes": ["hazard_prompt"], + "pii": True, "enable_thinking": True, - "optional_params": {"probes": [], "enable_thinking": False}, + "optional_params": {"probes": [], "pii": False, "enable_thinking": False}, }, [], False, @@ -392,15 +509,16 @@ def test_missing_api_key_is_rejected_at_startup() -> None: pytest.param( { "probes": "all", + "pii": False, "enable_thinking": True, - "optional_params": {"probes": None, "enable_thinking": None}, + "optional_params": {"probes": None, "pii": None, "enable_thinking": None}, }, "all", True, id="nested-null-falls-back", ), pytest.param( - {"probes": "all", "enable_thinking": True, "optional_params": {}}, + {"probes": "all", "pii": False, "enable_thinking": True, "optional_params": {}}, "all", True, id="empty-options-keep-top-level", @@ -468,6 +586,39 @@ async def test_configured_higher_hazard_threshold_allows_the_request( assert await _apply(guardrail, inputs) is inputs +@pytest.mark.asyncio +@pytest.mark.parametrize( + "settings", + [ + pytest.param({"pii_mask": False}, id="top-level"), + pytest.param({"optional_params": {"pii_mask": False}}, id="nested"), + pytest.param({"pii_mask": True, "optional_params": {"pii_mask": False}}, id="nested-false-wins"), + pytest.param({"pii_mask": False, "optional_params": {"pii_mask": None}}, id="null-nested"), + ], +) +async def test_configured_masking_disabled_blocks_detected_pii( + settings: dict[str, object], respx_mock: respx.MockRouter +) -> None: + guardrail: Final = _configured_guardrail(settings) + _serve(respx_mock, _mls_body(pii_spans=[_NAME_SPAN])) + + with pytest.raises(GuardrailRaisedException) as exc: + await _apply(guardrail, {"texts": ["Hello Alex."]}) + + assert exc.value.blocked_content is True + assert "PII detected in the request (name)" in exc.value.message, exc.value.message + + +@pytest.mark.asyncio +async def test_nested_masking_enabled_overrides_top_level_blocking(respx_mock: respx.MockRouter) -> None: + guardrail: Final = _configured_guardrail({"pii_mask": False, "optional_params": {"pii_mask": True}}) + _serve(respx_mock, _mls_body(pii_spans=[_NAME_SPAN])) + + result: Final = await _apply(guardrail, {"texts": ["Hello Alex."]}) + + assert result == {"texts": ["Hello [name]."]}, result + + @pytest.mark.asyncio @pytest.mark.parametrize( "settings", @@ -536,6 +687,21 @@ async def test_configured_timeout_reaches_mls( assert timeouts["read"] == expected +@pytest.mark.asyncio +async def test_role_mismatch_skips_request_hazard_but_still_masks_pii(respx_mock: respx.MockRouter) -> None: + _serve( + respx_mock, + { + "results": [{"probe": "hazard_prompt", "prob": 0.99, "role_mismatch": True}], + "pii_spans": [_NAME_SPAN], + }, + ) + + result: Final = await _apply(_guardrail(), {"texts": ["Hello Alex."]}) + + assert result == {"texts": ["Hello [name]."]}, result + + @pytest.mark.asyncio @pytest.mark.parametrize("mismatch_fields", [{}, {"role_mismatch": None}], ids=["omitted", "null"]) async def test_missing_role_mismatch_still_enforces_request_hazard( diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index b6b7beaf86f..05e7c78d495 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -36477,6 +36477,11 @@ export interface components { * @description Controls Pillar session persistence (sets `plr_persist` header). Set to False to disable persistence. */ persist_session?: boolean | null; + /** + * Pii + * @description Whether to run MLS's PII detection head. Defaults to True. + */ + pii?: boolean | null; /** * Pii Check * @description Enable PII (Personally Identifiable Information) detection. @@ -36495,6 +36500,11 @@ export interface components { pii_entities_config?: { [key: string]: components["schemas"]["PiiAction"]; } | null; + /** + * Pii Mask + * @description What to do with detected PII. True (default) rewrites each span as its type in brackets, e.g. "My name is Alex" -> "My name is [name]", and lets the request through. False blocks the request instead. + */ + pii_mask?: boolean | null; /** * Policy Id * @description Policy ID for Zscaler AI Guard. Can also be set via ZSCALER_AI_GUARD_POLICY_ID environment variable From b1a7fcc614243f2b60111486c7556a09f664abf2 Mon Sep 17 00:00:00 2001 From: Gourav Mittal Date: Thu, 1 Oct 2026 15:16:39 -0700 Subject: [PATCH 3/4] feat(guardrails): scan RealmLabs responses with conversation context Scan assistant replies with the request conversation as context and apply the PII masking or blocking policy to the current response text Preserve message roles and use shared text extraction for conversation history. Keep hazard blocking scoped to requests and test response handling, context selection, and service failure policies --- .../guardrail_hooks/realmlabs/realmlabs.py | 72 ++++++- .../realmlabs/test_realmlabs.py | 179 +++++++++++++++++- 2 files changed, 231 insertions(+), 20 deletions(-) diff --git a/litellm/proxy/guardrails/guardrail_hooks/realmlabs/realmlabs.py b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/realmlabs.py index 32aa9cedde2..98d1282b235 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/realmlabs/realmlabs.py +++ b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/realmlabs.py @@ -1,4 +1,12 @@ -"""RealmLabs MLS request guardrail with hazard screening and PII protection.""" +"""RealmLabs MLS guardrail. + +On ``pre_call`` MLS checks the request. A matching ``hazard_prompt`` verdict above ``hazard_threshold`` blocks +the request. On ``post_call`` MLS checks the assistant reply with the request's conversation as context. +Both hooks mask detected PII as ``[type]``, or block when ``pii_mask`` is False. + +Conversation context comes from Chat Completions-style ``request_data.messages``; nothing is kept between +hooks. Streaming uses LiteLLM's existing delivery settings, which do not enable text rewrites by default. +""" from __future__ import annotations @@ -20,6 +28,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # helper is untyped in http_handler httpxSpecialProvider, ) +from litellm.proxy.guardrails.guardrail_hooks.content_text import content_to_text from litellm.proxy.guardrails.guardrail_hooks.realmlabs.pii_masking import mask_pii_in_text from litellm.secret_managers.main import get_secret_str from litellm.types.guardrails import GuardrailEventHooks @@ -46,6 +55,8 @@ _DEFAULT_HAZARD_THRESHOLD: Final = 0.703 _DEFAULT_TIMEOUT: Final = 15.0 _RESPONSE_ADAPTER: Final = TypeAdapter(RealmLabsGuardrailResponse) +_MESSAGE_ADAPTER: Final = TypeAdapter(Mapping[str, object]) +_MESSAGES_ADAPTER: Final = TypeAdapter(tuple[object, ...]) class RealmLabsMissingCredentials(Exception): @@ -110,8 +121,8 @@ class RealmLabsGuardrail(CustomGuardrail): @classmethod def get_supported_event_hooks(cls) -> list[GuardrailEventHooks]: # mutable-ok: base class returns a list - """Inspect requests on ``pre_call``.""" - return [GuardrailEventHooks.pre_call] + """Inspect requests on ``pre_call`` and replies on ``post_call``.""" + return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call] @staticmethod def _hazard_score(response: RealmLabsGuardrailResponse) -> float | None: @@ -188,13 +199,13 @@ class RealmLabsGuardrail(CustomGuardrail): input_type: Literal["request", "response"], logging_obj: LiteLLMLoggingObj | None = None, ) -> GenericGuardrailAPIInputs: - """Screen requests for hazardous content and PII through LiteLLM's unified guardrail layer.""" - if input_type != "request": - return inputs + """Inspect request or response text through LiteLLM's unified guardrail layer. + + Enforce hazard scores only on requests with a matching role, then mask or block PII on either side. + Only the current hook's texts are rewritten. MLS failures pass through unless ``block_on_error`` is set. + """ texts: Final = tuple(inputs.get("texts") or ()) - messages: Final = tuple(inputs.get("structured_messages") or ()) or tuple( - RealmLabsChatMessage(role="user", content=text) for text in texts - ) + messages: Final = _messages_for_scan(inputs, request_data, input_type) if not messages: return inputs @@ -207,7 +218,7 @@ class RealmLabsGuardrail(CustomGuardrail): return self._handle_mls_error(inputs, f"Invalid RealmLabs guardrail response: {result.reason}") # Hazard is checked before PII, so a hazardous prompt is rejected rather than masked and forwarded. - hazard_score: Final = self._hazard_score(result) + hazard_score: Final = self._hazard_score(result) if input_type == "request" else None if hazard_score is not None and hazard_score > self.hazard_threshold: verbose_proxy_logger.warning( "RealmLabs MLS blocked request: %s=%s > %s", @@ -240,3 +251,44 @@ class RealmLabsGuardrail(CustomGuardrail): return inputs verbose_proxy_logger.debug("RealmLabs MLS masked PII types in the %s: %s", input_type, self._span_types(spans)) return {**inputs, "texts": list(masked_texts)} + + +def _conversation_message(message: object) -> RealmLabsChatMessage | None: + """Extract a plain-text turn without carrying images or unrelated message fields to MLS.""" + + if not isinstance(message, dict): + return None + row: Final = _MESSAGE_ADAPTER.validate_python(message) + role: Final = row.get("role") + text: Final = content_to_text(row.get("content")) + if isinstance(role, str) and role and text: + return RealmLabsChatMessage(role=role, content=text) + return None + + +def _conversation_messages(request_data: Mapping[str, object]) -> tuple[RealmLabsChatMessage, ...]: + """Preserve roles and text from request history, skipping empty or non-message entries.""" + + raw_messages: Final = request_data.get("messages") + if not isinstance(raw_messages, list): + return () + + messages: Final = _MESSAGES_ADAPTER.validate_python(raw_messages) + return tuple(message for item in messages if (message := _conversation_message(item)) is not None) + + +def _messages_for_scan( + inputs: GenericGuardrailAPIInputs, + request_data: Mapping[str, object], + input_type: Literal["request", "response"], +) -> Sequence[Mapping[str, object]]: + """Prefer scoped request messages; append response texts as assistant turns to the request history.""" + + if input_type == "request" and (structured_messages := inputs.get("structured_messages")): + return tuple(structured_messages) + + conversation: Final = _conversation_messages(request_data) + texts: Final = inputs.get("texts") or () + if input_type == "request": + return conversation or tuple(RealmLabsChatMessage(role="user", content=text) for text in texts) + return (*conversation, *(RealmLabsChatMessage(role="assistant", content=text) for text in texts)) diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_realmlabs.py b/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_realmlabs.py index 8448bf16c53..a87a0db3a46 100644 --- a/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_realmlabs.py +++ b/tests/unit/proxy/guardrails/guardrail_hooks/realmlabs/test_realmlabs.py @@ -10,6 +10,7 @@ from pydantic import TypeAdapter import litellm from litellm.exceptions import GuardrailRaisedException +from litellm.llms.openai.chat.guardrail_translation.handler import OpenAIChatCompletionsHandler from litellm.proxy.guardrails.guardrail_hooks.realmlabs import guardrail_initializer_registry from litellm.proxy.guardrails.guardrail_hooks.realmlabs.realmlabs import ( RealmLabsGuardrail, @@ -18,7 +19,7 @@ from litellm.proxy.guardrails.guardrail_hooks.realmlabs.realmlabs import ( from litellm.types.guardrails import GuardrailEventHooks, LitellmParams from litellm.types.llms.openai import AllMessageValues from litellm.types.proxy.guardrails.guardrail_hooks.realmlabs import RealmLabsGuardrailOptionalParams -from litellm.types.utils import GenericGuardrailAPIInputs +from litellm.types.utils import Choices, GenericGuardrailAPIInputs, Message, ModelResponse _API_KEY = "mls_gr_test" _API_BASE = "https://mls.example.test" @@ -232,6 +233,7 @@ async def test_span_from_another_turn_leaves_the_prompt_as_is(respx_mock: respx. pytest.param( "request", ["Alex alex@example.com"], [_NAME_SPAN, _EMAIL_SPAN], "name, email", id="request-types" ), + pytest.param("response", ["alex@example.com"], [_EMAIL_SPAN], "email", id="response-type"), pytest.param("request", ["Hello Alex."], [{"type": "name"}], "name", id="missing-pii-text"), pytest.param("request", ["Hello Alex."], [{"type": "name", "text": None}], "name", id="null-pii-text"), pytest.param("request", ["Hello Alex."], [{"type": "name", "text": ""}], "name", id="empty-pii-text"), @@ -287,7 +289,7 @@ async def test_plain_texts_are_sent_as_user_turns(respx_mock: respx.MockRouter) @pytest.mark.asyncio -@pytest.mark.parametrize("input_type", ["request"]) +@pytest.mark.parametrize("input_type", ["request", "response"]) async def test_empty_inputs_are_returned_without_calling_mls( input_type: Literal["request", "response"], respx_mock: respx.MockRouter ) -> None: @@ -299,7 +301,7 @@ async def test_empty_inputs_are_returned_without_calling_mls( @pytest.mark.asyncio -@pytest.mark.parametrize("input_type", ["request"]) +@pytest.mark.parametrize("input_type", ["request", "response"]) @pytest.mark.parametrize("block_on_error", [None, True], ids=["default-fail-open", "fail-closed"]) @pytest.mark.parametrize( ("mls_reply", "reason"), @@ -344,7 +346,7 @@ async def test_mls_failures_follow_the_error_policy( @pytest.mark.asyncio -@pytest.mark.parametrize("input_type", ["request"]) +@pytest.mark.parametrize("input_type", ["request", "response"]) @pytest.mark.parametrize("block_on_error", [False, True], ids=["fail-open", "fail-closed"]) @pytest.mark.parametrize( "body", @@ -413,7 +415,7 @@ async def test_invalid_probabilities_cannot_bypass_fail_closed(prob: object, res @pytest.mark.asyncio -@pytest.mark.parametrize("input_type", ["request"]) +@pytest.mark.parametrize("input_type", ["request", "response"]) async def test_empty_verdict_arrays_are_valid_in_fail_closed_mode( input_type: Literal["request", "response"], respx_mock: respx.MockRouter ) -> None: @@ -687,6 +689,38 @@ async def test_configured_timeout_reaches_mls( assert timeouts["read"] == expected +@pytest.mark.asyncio +@pytest.mark.parametrize( + "use_structured_messages", [False, True], ids=["request-fallback", "structured-takes-priority"] +) +async def test_request_uses_structured_messages_before_conversation_fallback( + use_structured_messages: bool, respx_mock: respx.MockRouter +) -> None: + conversation: Final[list[AllMessageValues]] = [ + {"role": "system", "content": "Be brief."}, + {"role": "user", "content": "Hello."}, + ] + inputs: Final[GenericGuardrailAPIInputs] = ( + {"texts": ["Hello."], "structured_messages": conversation} if use_structured_messages else {"texts": ["Hello."]} + ) + request_data: Final = { + "messages": [{"role": "user", "content": "Outside the selected scope."}] + if use_structured_messages + else conversation + } + route: Final = _serve(respx_mock, _mls_body()) + + result: Final = await _apply(_guardrail(), inputs, request_data=request_data) + + assert result is inputs + assert _sent_body(route) == { + "messages": conversation, + "probes": ["hazard_prompt"], + "pii": True, + "enable_thinking": False, + } + + @pytest.mark.asyncio async def test_role_mismatch_skips_request_hazard_but_still_masks_pii(respx_mock: respx.MockRouter) -> None: _serve( @@ -717,9 +751,134 @@ async def test_missing_role_mismatch_still_enforces_request_hazard( @pytest.mark.asyncio -async def test_response_text_is_returned_without_calling_mls(respx_mock: respx.MockRouter) -> None: - route: Final = _serve(respx_mock, _mls_body()) - inputs: Final[GenericGuardrailAPIInputs] = {"texts": ["hello"]} +@pytest.mark.parametrize("name", ["Alex", "[name]"], ids=["unmasked-name", "existing-placeholder"]) +async def test_post_call_masks_the_reply_in_conversation_context(name: str, respx_mock: respx.MockRouter) -> None: + guardrail: Final = _guardrail(event_hook=GuardrailEventHooks.post_call) + conversation: Final[list[AllMessageValues]] = [ + {"role": "system", "content": "Be brief."}, + {"role": "user", "content": "My name is Alex."}, + {"role": "assistant", "content": "Hello!"}, + {"role": "user", "content": "What is my name?"}, + ] + route: Final = _serve(respx_mock, _mls_body(pii_spans=[{"type": "name", "text": name.removesuffix("]")}])) + response: Final = ModelResponse( + choices=[ + Choices(index=0, message=Message(content=f"Your name is {name}.", role="assistant"), finish_reason="stop") + ] + ) - assert await _apply(_guardrail(), inputs, input_type="response") is inputs - assert route.call_count == 0 + result: Final = await OpenAIChatCompletionsHandler().process_output_response( # pyright: ignore[reportUnknownMemberType] # upstream request_data parameter uses an unparameterized dict + response=response, guardrail_to_apply=guardrail, request_data={"messages": conversation} + ) + + assert result.choices == [ + Choices(index=0, message=Message(content="Your name is [name].", role="assistant"), finish_reason="stop") + ], result.choices + assert conversation == [ + {"role": "system", "content": "Be brief."}, + {"role": "user", "content": "My name is Alex."}, + {"role": "assistant", "content": "Hello!"}, + {"role": "user", "content": "What is my name?"}, + ], conversation + assert _sent_body(route) == { + "messages": [*conversation, {"role": "assistant", "content": f"Your name is {name}."}], + "probes": ["hazard_prompt"], + "pii": True, + "enable_thinking": False, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("input_type", ["request", "response"]) +async def test_conversation_text_parts_use_blank_lines_without_changing_internal_paragraphs( + input_type: Literal["request", "response"], respx_mock: respx.MockRouter +) -> None: + conversation: Final = [ + {"role": "system", "content": "Be brief."}, + { + "role": "user", + "content": [ + {"type": "text", "text": "Describe this.\n\nKeep this paragraph."}, + {"type": "image_url", "image_url": {"url": "https://example.test/image.png"}}, + {"type": "text", "text": "In one sentence."}, + ], + }, + {"role": "assistant", "content": None}, + {"role": "user", "content": [{"type": "image_url", "image_url": {"url": "https://example.test/other.png"}}]}, + {"content": "No role."}, + None, + ] + route: Final = _serve(respx_mock, _mls_body()) + inputs: Final[GenericGuardrailAPIInputs] = {"texts": ["A landscape."]} + + result: Final = await _apply(_guardrail(), inputs, input_type=input_type, request_data={"messages": conversation}) + + expected_history: Final = [ + {"role": "system", "content": "Be brief."}, + {"role": "user", "content": "Describe this.\n\nKeep this paragraph.\n\nIn one sentence."}, + ] + expected_messages: Final = ( + [*expected_history, {"role": "assistant", "content": "A landscape."}] + if input_type == "response" + else expected_history + ) + + assert result is inputs + assert _sent_body(route) == { + "messages": expected_messages, + "probes": ["hazard_prompt"], + "pii": True, + "enable_thinking": False, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("request_data", [{}, {"messages": []}, {"messages": None}], ids=["missing", "empty", "null"]) +async def test_response_without_history_is_still_scanned_as_assistant_text( + request_data: dict[str, object], respx_mock: respx.MockRouter +) -> None: + route: Final = _serve(respx_mock, _mls_body(pii_spans=[_NAME_SPAN])) + + result: Final = await _apply( + _guardrail(), {"texts": ["Alex was here."]}, input_type="response", request_data=request_data + ) + + assert result == {"texts": ["[name] was here."]}, result + assert _sent_body(route) == { + "messages": [{"role": "assistant", "content": "Alex was here."}], + "probes": ["hazard_prompt"], + "pii": True, + "enable_thinking": False, + } + + +@pytest.mark.asyncio +@pytest.mark.parametrize("role_mismatch", [False, True, None], ids=["matching-role", "mismatched-role", "null-role"]) +async def test_response_ignores_hazard_scores_and_merges_overlapping_pii( + role_mismatch: bool | None, respx_mock: respx.MockRouter +) -> None: + _serve( + respx_mock, + { + "results": [{"probe": "hazard_prompt", "prob": 0.99, "role_mismatch": role_mismatch}], + "pii_spans": [{"type": "name", "text": "Ann"}, {"type": "email", "text": "Ann.Smith@example.com"}], + }, + ) + inputs: Final[GenericGuardrailAPIInputs] = {"texts": ["Contact Ann at Ann.Smith@example.com.", "No PII here."]} + + result: Final = await _apply(_guardrail(), inputs, input_type="response") + + assert result == {"texts": ["Contact [name] at [email].", "No PII here."]}, result + assert inputs == {"texts": ["Contact Ann at Ann.Smith@example.com.", "No PII here."]}, inputs + + +@pytest.mark.asyncio +async def test_pii_only_in_history_leaves_the_reply_unchanged(respx_mock: respx.MockRouter) -> None: + _serve(respx_mock, _mls_body(pii_spans=[_NAME_SPAN])) + inputs: Final[GenericGuardrailAPIInputs] = {"texts": ["Hello there."]} + conversation: Final = [{"role": "user", "content": "My name is Alex."}] + + result: Final = await _apply(_guardrail(), inputs, input_type="response", request_data={"messages": conversation}) + + assert result is inputs + assert conversation == [{"role": "user", "content": "My name is Alex."}], conversation From 39851d88a30c6e7b3cb6ab67c144b1582efa7b68 Mon Sep 17 00:00:00 2001 From: Gourav Mittal Date: Thu, 1 Oct 2026 15:17:13 -0700 Subject: [PATCH 4/4] docs(guardrails): document RealmLabs configuration and usage Document request and response screening, configuration precedence, failure policies, and the guardrail API contract. Add a runnable proxy configuration example and explain where future request and response changes belong --- .../guardrail_hooks/realmlabs/README.md | 106 ++++++++++++++++++ .../realmlabs/example_config.yaml | 28 +++++ 2 files changed, 134 insertions(+) create mode 100644 litellm/proxy/guardrails/guardrail_hooks/realmlabs/README.md create mode 100644 litellm/proxy/guardrails/guardrail_hooks/realmlabs/example_config.yaml diff --git a/litellm/proxy/guardrails/guardrail_hooks/realmlabs/README.md b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/README.md new file mode 100644 index 00000000000..a210ee1aaa8 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/README.md @@ -0,0 +1,106 @@ +# RealmLabs MLS Guardrail Integration + +Checks prompts and model replies with RealmLabs MLS. It blocks hazardous prompts and masks detected PII as `[type]`, or blocks PII when masking is disabled + +## Configuration + +Add the guardrail to your proxy configuration. The [complete example](example_config.yaml) also configures Claude Haiku 4.5 as the model + +```yaml +guardrails: + - guardrail_name: realmlabs-guard + litellm_params: + guardrail: realmlabs + mode: [pre_call, post_call] + api_key: os.environ/REALMLABS_API_KEY + default_on: true + probes: [hazard_prompt] + hazard_threshold: 0.703 + pii: true + pii_mask: true + block_on_error: false + optional_params: + enable_thinking: false + timeout: 15 +``` + +### Credentials and endpoint + +| Setting | Meaning | +| --- | --- | +| `api_key` | Required MLS guardrail bearer token; falls back to `REALMLABS_API_KEY` when omitted | +| `api_base` | MLS base URL; falls back to `REALMLABS_API_BASE`, then `https://mls.realmlabs.ai` | + +The integration appends `/guardrail` to the base URL. This route uses a guardrail API key, separate from the token for `/llm/*` routes + +### Tuning parameters + +| Setting | Default | Meaning | +| --- | --- | --- | +| `probes` | `[hazard_prompt]` | Probe names or `all`; only the `hazard_prompt` score is enforced | +| `hazard_threshold` | `0.703` | Block request scores strictly above this value | +| `pii` | `true` | Ask MLS to detect PII | +| `pii_mask` | `true` | Mask detected PII; `false` blocks instead | +| `block_on_error` | `false` | Allow MLS failures through; `true` blocks them | +| `enable_thinking` | `false` | Ask MLS to use thinking mode in its chat template | +| `timeout` | `15` seconds | HTTP timeout for each MLS call, separate from the model request timeout | + +All seven tuning parameters support top-level values under `litellm_params` or nested `optional_params`. An explicit, non-null nested value wins, then the top-level value, then the RealmLabs default. False, zero, and empty lists remain explicit overrides + +Current configuration parsing limits nested `optional_params.timeout` to 1–60 seconds. For a value outside that range, set top-level `timeout` and omit the nested timeout + +## Usage examples + +With the proxy running using the example config, send a non-streaming Chat Completions request. Set `LITELLM_MASTER_KEY` in the client terminal to your proxy key + +```bash +curl --silent --show-error --include http://localhost:4000/v1/chat/completions \ + --header "Authorization: Bearer ${LITELLM_MASTER_KEY}" \ + --header 'Content-Type: application/json' \ + --data '{ + "model": "claude-haiku-4-5", + "messages": [{"role": "user", "content": "My name is Alex and my email is alex@example.com. What do you know about me?"}], + "max_tokens": 100 + }' +``` + +If MLS permits the prompt and detects both values, the model receives `My name is [name] and my email is [email]. What do you know about me?`. Detections and model replies can vary + +| MLS verdict | Result | +| --- | --- | +| No blocking hazard or detected PII | Text passes through unchanged | +| PII detected with `pii_mask: true` | Matching text is masked; overlapping matches are merged | +| PII detected with `pii_mask: false` | The request or reply is blocked | +| Request hazard score exceeds the threshold | The request is blocked before the model is called | + +The example's `default_on: true` applies the guardrail automatically. For opt-in use, set it to `false` and include `"guardrails": ["realmlabs-guard"]` in each request that should be checked + +## Supported event hooks + +| Hook | Behavior | +| --- | --- | +| `pre_call` | Checks the request, enforcing hazard before masking or blocking PII | +| `post_call` | Checks the reply with conversation context, masking or blocking PII; hazard scores do not block replies | + +A hazard verdict marked `role_mismatch` is not enforced. Conversation context comes from Chat Completions-style `messages`. Streaming uses LiteLLM's existing delivery settings, which do not enable text rewrites by default + +## Error handling + +Hazard and PII policy blocks raise `GuardrailRaisedException` with HTTP status 400. A hazard block includes `Blocked by RealmLabs hazard_prompt probe` and the score and threshold + +An MLS connection failure, timeout, HTTP error, or unreadable response follows `block_on_error`: the default `false` passes text through unchanged; `true` raises an error. A successful completion alone does not prove MLS returned an allow verdict + +Responses must include `results` and `pii_spans` arrays, which may be empty. Each probe result needs a name and a finite probability between zero and one. Missing verdict fields are errors, not clean results. When masking is enabled, a PII span without a nonempty type and text also follows `block_on_error`. When masking is disabled, spans without text still trigger a PII content block + +Request fields are defined by `RealmLabsGuardrailRequest` and assembled by `_build_request`. Response types and `_parse_response` define what blocking and masking consume. Additional response fields are ignored, including fields inside probe results and PII spans. Update these boundaries and their behavior tests when adding supported fields or changing the contract + +## Unit tests + +From the repository root with the development environment installed: + +```bash +LITELLM_LOCAL_MODEL_COST_MAP=True .venv/bin/python -m pytest \ + tests/unit/proxy/guardrails/guardrail_hooks/realmlabs -q +``` + +Tests cover policy decisions, overlapping PII, response scanning, MLS failures, and configuration precedence using simulated MLS HTTP responses diff --git a/litellm/proxy/guardrails/guardrail_hooks/realmlabs/example_config.yaml b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/example_config.yaml new file mode 100644 index 00000000000..4f50151a544 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/realmlabs/example_config.yaml @@ -0,0 +1,28 @@ +# Set ANTHROPIC_API_KEY and REALMLABS_API_KEY before starting the proxy +# Run: litellm --config litellm/proxy/guardrails/guardrail_hooks/realmlabs/example_config.yaml + +model_list: + - model_name: claude-haiku-4-5 + litellm_params: + model: anthropic/claude-haiku-4-5 + api_key: os.environ/ANTHROPIC_API_KEY + +guardrails: + - guardrail_name: realmlabs-guard + litellm_params: + guardrail: realmlabs + mode: [pre_call, post_call] + api_key: os.environ/REALMLABS_API_KEY + api_base: https://mls.realmlabs.ai + default_on: true + + probes: [hazard_prompt] + hazard_threshold: 0.703 + pii: true + pii_mask: true + block_on_error: false + + # All seven tuning settings can use either location; non-null nested values win + optional_params: + enable_thinking: false + timeout: 15 # seconds