This commit is contained in:
gourav-realm 2026-10-05 10:47:27 -07:00 • committed by GitHub
commit 1a57f90e8d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 1967 additions and 0 deletions

View file

@ -13553,6 +13553,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": [
{
@ -13723,6 +13735,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": [
{
@ -14115,6 +14139,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": [
{
@ -14157,6 +14193,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": [
{
@ -14361,6 +14409,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": [
{

View file

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

View file

@ -0,0 +1,74 @@
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,
pii=settings.pii,
pii_mask=settings.pii_mask,
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)})
)

View file

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

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

@ -0,0 +1,294 @@
"""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
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.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
from litellm.types.proxy.guardrails.guardrail_hooks.realmlabs import (
RealmLabsChatMessage,
RealmLabsGuardrailConfigModel,
RealmLabsGuardrailRequest,
RealmLabsGuardrailResponse,
RealmLabsPIISpan,
)
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)
_MESSAGE_ADAPTER: Final = TypeAdapter(Mapping[str, object])
_MESSAGES_ADAPTER: Final = TypeAdapter(tuple[object, ...])
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 and masks PII 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,
pii: bool | None = None,
pii_mask: bool | 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.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
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`` and replies on ``post_call``."""
return [GuardrailEventHooks.pre_call, GuardrailEventHooks.post_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
@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=self.pii,
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))
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(
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 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,
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:
"""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 = _messages_for_scan(inputs, request_data, input_type)
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 is checked before PII, so a hazardous prompt is rejected rather than masked and forwarded.
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",
_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,
)
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)}
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))

View file

@ -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"
@ -1188,6 +1192,7 @@ class LitellmParams( # pyright: ignore[reportIncompatibleVariableOverride] # o
GraySwanGuardrailConfigModel,
NomaGuardrailConfigModel,
PromptGuardConfigModel,
RealmLabsGuardrailConfigModel,
XecGuardConfigModel,
ToolPermissionGuardrailConfigModel,
ZscalerAIGuardConfigModel,

View file

@ -0,0 +1,156 @@
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 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):
"""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.",
)
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.",
)
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.'
),
)
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=(
"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"

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

@ -0,0 +1,884 @@
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.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,
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 Choices, GenericGuardrailAPIInputs, Message, ModelResponse
_API_KEY = "mls_gr_test"
_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])
@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,
pii_mask: bool | 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,
pii_mask=pii_mask,
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_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("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"),
],
)
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] = [
{"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": True,
"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": True,
"enable_thinking": False,
}
@pytest.mark.asyncio
@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:
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", "response"])
@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", "response"])
@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"
),
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",
),
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(
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", "response"])
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], ids=["mask-pii", "block-hazard"])
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": [{"type": "name", "text": "Alex", "future_field": {"source": "test"}}],
},
)
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) == {"texts": ["Hello [name]."]}
@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"], "pii": False, "enable_thinking": True},
["hazard_prompt", "dispute"],
True,
id="top-level",
),
pytest.param(
{"optional_params": {"probes": "all", "pii": False, "enable_thinking": True}},
"all",
True,
id="nested",
),
pytest.param(
{
"probes": ["hazard_prompt"],
"pii": True,
"enable_thinking": True,
"optional_params": {"probes": [], "pii": False, "enable_thinking": False},
},
[],
False,
id="nested-false-and-empty-list-win",
),
pytest.param(
{
"probes": "all",
"pii": False,
"enable_thinking": True,
"optional_params": {"probes": None, "pii": None, "enable_thinking": None},
},
"all",
True,
id="nested-null-falls-back",
),
pytest.param(
{"probes": "all", "pii": False, "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({"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",
[
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(
"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(
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(
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
@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")
]
)
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

View file

@ -36345,6 +36345,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.
@ -36417,6 +36422,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).
@ -36578,6 +36588,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.
@ -36596,6 +36611,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
@ -36669,6 +36689,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