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
This commit is contained in:
Gourav Mittal 2026-10-01 15:15:18 -07:00
parent 2c0fd473b2
commit 6d37d3379d
8 changed files with 616 additions and 17 deletions

View file

@ -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": [
{

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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