From 6d37d3379d2c1ebc283051f7d7247bdec02ad9d0 Mon Sep 17 00:00:00 2001 From: Gourav Mittal Date: Thu, 1 Oct 2026 15:15:18 -0700 Subject: [PATCH] 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