This commit is contained in:
Caduri 2026-09-30 16:55:11 -04:00 • committed by GitHub
commit 70a46512ee
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
10 changed files with 1608 additions and 16 deletions

View file

@ -8,9 +8,13 @@ skip the other shapes — these helpers normalise that so every hook sees
every text fragment.
"""
import json
from collections.abc import Callable, Iterator, Mapping, Sequence
from typing import Any, Final
from pydantic import JsonValue
from pydantic_core import to_jsonable_python
# Call types whose body carries free-form chat / prompt text that
# text-content guardrails (banned keywords, content moderation, secret
# detection, …) should inspect. The proxy ingress passes ``route_type``
@ -146,6 +150,51 @@ def iter_message_text(data: Mapping[str, object]) -> Iterator[str]:
yield from _iter_text_parts_in_content(message.get("content"))
def image_part_url(part: JsonValue) -> str | None:
"""The URL of a Chat Completions ``image_url`` part, in either its object or bare-string form."""
if not isinstance(part, dict) or part.get("type") != "image_url":
return None
image_url: Final = part.get("image_url")
if isinstance(image_url, str):
return image_url
url: Final = image_url.get("url") if isinstance(image_url, dict) else None
return url if isinstance(url, str) else None
def map_content_image_urls(content: JsonValue, transform: Callable[[str], str]) -> JsonValue:
"""Return a copy of ``content`` with every ``image_url`` part's URL replaced by ``transform(url)``."""
if not isinstance(content, list):
return content
return [_image_part_with_mapped_url(part, transform) for part in content] # mutable-ok: JSON content is an array
def _image_part_with_mapped_url(part: JsonValue, transform: Callable[[str], str]) -> JsonValue:
if not isinstance(part, dict) or part.get("type") != "image_url":
return part
image_url: Final = part.get("image_url")
if isinstance(image_url, str):
return {**part, "image_url": transform(image_url)} # mutable-ok: JSON parts are objects
if not isinstance(image_url, dict):
return part
url: Final = image_url.get("url")
if not isinstance(url, str):
return part
return {**part, "image_url": {**image_url, "url": transform(url)}} # mutable-ok: JSON parts are objects
def map_messages_image_urls(messages: JsonValue, transform: Callable[[str], str]) -> JsonValue:
"""Return a new message list with :func:`map_content_image_urls` applied to every ``content``."""
if not isinstance(messages, list):
return messages
return [_message_with_mapped_image_urls(message, transform) for message in messages] # mutable-ok: JSON array
def _message_with_mapped_image_urls(message: JsonValue, transform: Callable[[str], str]) -> JsonValue:
if not isinstance(message, dict) or not isinstance(message.get("content"), list):
return message
return {**message, "content": map_content_image_urls(message["content"], transform)} # mutable-ok: JSON object
def walk_user_text(data: dict[str, Any], visit: Callable[[str], str]) -> int:
"""Rewrite every text fragment in place via ``visit``.
@ -307,3 +356,12 @@ def build_inspection_messages(data: dict[str, Any]) -> list[dict[str, str]]:
role = message.get("role", "user") or "user"
flattened.append({"role": role, "content": text})
return flattened
def as_json_value(value: object) -> JsonValue:
"""Round-trips through the stdlib codec because pydantic's serializer turns anything nested
past 254 levels into "...", while this keeps about the depth the proxy's request parser accepts"""
parsed: Final[JsonValue] = json.loads( # pyright: ignore[reportAny] # untyped stdlib parse of json.dumps output
json.dumps(value, default=to_jsonable_python)
)
return parsed

View file

@ -39,6 +39,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
streaming_end_of_stream_only=_get_config_value(litellm_params, optional_params, "streaming_end_of_stream_only"),
streaming_sampling_rate=_get_config_value(litellm_params, optional_params, "streaming_sampling_rate"),
streaming_transform_mode=_get_config_value(litellm_params, optional_params, "streaming_transform_mode"),
send_images=_get_config_value(litellm_params, optional_params, "send_images"),
exclude_payload_fields=_get_config_value(litellm_params, optional_params, "exclude_payload_fields"),
)
litellm.logging_callback_manager.add_litellm_callback(_generic_guardrail_api_callback)

View file

@ -0,0 +1,15 @@
import re
from collections.abc import Sequence
def config_values(raw: Sequence[str] | None, *, option_name: str) -> tuple[str, ...]:
if isinstance(raw, str):
raise ValueError(f"{option_name} must be a list of strings, got the single string {raw!r}")
return tuple(raw or ())
def compile_patterns(raw: Sequence[str] | None, *, option_name: str) -> tuple[re.Pattern[str], ...]:
try:
return tuple(re.compile(pattern) for pattern in config_values(raw, option_name=option_name))
except re.error as e:
raise ValueError(f"{option_name} contains an invalid regex: {e}") from e

View file

@ -11,6 +11,8 @@ from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
import httpx
from pydantic import JsonValue
from typing_extensions import TypeIs
from litellm._logging import verbose_proxy_logger
from litellm._version import version as litellm_version
@ -19,10 +21,13 @@ from litellm.integrations.custom_guardrail import (
CustomGuardrail,
log_guardrail_information,
)
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
from litellm.llms.custom_httpx.http_handler import (
AsyncHTTPHandler,
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy.guardrails._content_utils import as_json_value
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
@ -33,6 +38,15 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import
)
from litellm.types.utils import GenericGuardrailAPIInputs
from .payload_policy import (
PayloadLoss,
accepted_rewrites,
raise_if_intervention_was_refused,
resolve_payload_policy,
restore_unseen_rows,
shape_payload,
)
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
@ -150,6 +164,28 @@ def _extract_inbound_headers(
return None
def _is_part_list(value: object) -> TypeIs[list[object]]: # guard-ok: trivial isinstance narrowing
return isinstance(value, list)
def _row_as_sent(dumped: JsonValue, caller: Mapping[str, object]) -> JsonValue:
"""The request model validates list content lazily and dumps a list holding
any part it rejects as [], so such a row is sent with the caller's content."""
caller_content: Final = caller.get("content")
if not isinstance(dumped, dict) or not _is_part_list(caller_content):
return dumped
dumped_content: Final = dumped.get("content")
if isinstance(dumped_content, list) and len(dumped_content) == len(caller_content):
return dumped
return {**dumped, "content": as_json_value(caller_content)}
def _rows_as_sent(dumped_rows: JsonValue, caller_rows: Sequence[Mapping[str, object]] | None) -> JsonValue:
if caller_rows is None or not isinstance(dumped_rows, list):
return dumped_rows
return [_row_as_sent(dumped, caller) for dumped, caller in zip(dumped_rows, caller_rows, strict=True)]
def _structured_rows_to_write_back(
original_rows: Sequence[AllMessageValues] | None,
shown_rows: Sequence[AllMessageValues] | None,
@ -204,9 +240,14 @@ class GenericGuardrailAPI(CustomGuardrail):
streaming_end_of_stream_only: bool | None = None,
streaming_sampling_rate: int | None = None,
streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = None,
send_images: bool | None = None,
exclude_payload_fields: Sequence[str] | None = None,
async_handler: AsyncHTTPHandler | None = None,
**kwargs,
):
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
self.async_handler = async_handler or get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
)
self.headers = headers or {}
self.extra_headers = extra_headers or []
@ -251,6 +292,12 @@ class GenericGuardrailAPI(CustomGuardrail):
"block_only" if streaming_transform_mode is None else streaming_transform_mode
)
self._payload_policy: Final = resolve_payload_policy(
send_images=send_images,
exclude_payload_fields=exclude_payload_fields,
guardrail_name=kwargs.get("guardrail_name"),
)
# Set supported event hooks
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
@ -345,24 +392,42 @@ class GenericGuardrailAPI(CustomGuardrail):
structured_messages: Sequence[AllMessageValues] | None,
shown_messages: Sequence[AllMessageValues] | None,
guardrail_response: GenericGuardrailAPIResponse,
loss: PayloadLoss,
) -> GenericGuardrailAPIInputs:
# Action is NONE or no modifications needed
return_inputs: Final = GenericGuardrailAPIInputs(texts=texts)
if guardrail_response.texts:
return_inputs["texts"] = guardrail_response.texts
if guardrail_response.images:
return_inputs["images"] = guardrail_response.images
name: Final = self.guardrail_name
accepted: Final = accepted_rewrites(guardrail_response, loss, guardrail_name=name)
if accepted.texts:
return_inputs["texts"] = accepted.texts
if accepted.images:
return_inputs["images"] = accepted.images
elif images:
return_inputs["images"] = images
if guardrail_response.tools:
return_inputs["tools"] = guardrail_response.tools
if accepted.tools:
return_inputs["tools"] = accepted.tools
elif tools:
return_inputs["tools"] = tools
rows_to_write_back: Final = (
_structured_rows_to_write_back(structured_messages, shown_messages, guardrail_response.structured_messages)
if guardrail_response.structured_messages
restore_unseen_rows(
rows=_structured_rows_to_write_back(structured_messages, shown_messages, accepted.rows),
caller=structured_messages,
sent=shown_messages,
loss=loss,
guardrail_name=name,
)
if accepted.rows
else None
)
raise_if_intervention_was_refused(
action=guardrail_response.action,
accepted=accepted,
original_texts=texts,
original_images=images,
original_tools=tools,
rows_written_back=rows_to_write_back is not None,
guardrail_name=name,
)
if rows_to_write_back is not None:
return_inputs["structured_messages"] = list(rows_to_write_back) # mutable-ok: guardrail inputs take a list
if guardrail_response.stream_holdback_chars is not None:
@ -470,12 +535,15 @@ class GenericGuardrailAPI(CustomGuardrail):
)
headers: Final = self._build_request_headers()
# The model's list content is a lazy iterator that this dump consumes, so it cannot be read again
dumped: Final[Mapping[str, JsonValue]] = guardrail_request.model_dump(mode="json")
sent_messages: Final = _rows_as_sent(dumped.get("structured_messages"), structured_messages)
request_json: Final = {**dumped, "structured_messages": sent_messages} # mutable-ok: JSON POST body
payload: Final = shape_payload(request_json, self._payload_policy)
# Make the API request
# Use mode="json" to ensure all iterables are converted to lists
response: Final = await self.async_handler.post(
url=self.api_base,
json=guardrail_request.model_dump(mode="json"),
json=payload.body,
headers=headers,
)
@ -503,11 +571,12 @@ class GenericGuardrailAPI(CustomGuardrail):
images=images,
tools=tools,
structured_messages=structured_messages,
shown_messages=guardrail_request.structured_messages,
shown_messages=payload.sent_messages,
guardrail_response=guardrail_response,
loss=payload.loss,
)
except GuardrailRaisedException:
except (GuardrailRaisedException, UnappliableRequestRewrite):
raise
except Timeout as e:
return self._handle_guardrail_request_error(e, inputs, input_type, logging_obj)

View file

@ -0,0 +1,308 @@
"""What the Generic Guardrail API endpoint is sent, and what it may write back.
Shaping is lossy, so every shaped payload carries a ``PayloadLoss``. Caller content the
guardrail did not see in full is never replaced by the guardrail's response.
"""
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from types import MappingProxyType
from typing import Final
from pydantic import JsonValue
from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.guardrail_translation.utils import unappliable_request_rewrite
from litellm.proxy.guardrails._content_utils import as_json_value, image_part_url, map_messages_image_urls
from litellm.types.llms.openai import AllMessageValues
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPIRequest,
GenericGuardrailAPIResponse,
GuardrailToolParam,
structured_messages_from_json,
)
from .config_parsing import config_values
PROTECTED_PAYLOAD_FIELDS: Final = frozenset({"input_type", "litellm_call_id"})
IMAGE_OMITTED_PLACEHOLDER: Final = "[omitted]"
@dataclass(frozen=True, slots=True)
class PayloadPolicy:
send_images: bool = True
exclude_fields: frozenset[str] = frozenset()
@property
def omitted_fields(self) -> frozenset[str]:
return self.exclude_fields if self.send_images else self.exclude_fields | frozenset(("images",))
@property
def shapes_messages(self) -> bool:
return not self.send_images
@property
def is_lossy(self) -> bool:
return bool(self.omitted_fields)
@dataclass(frozen=True, slots=True)
class PayloadLoss:
altered_message_indices: frozenset[int] = frozenset()
texts_omitted: bool = False
messages_omitted: bool = False
images_omitted: bool = False
tools_omitted: bool = False
@dataclass(frozen=True, slots=True)
class ShapedPayload:
body: dict[str, JsonValue] # mutable-ok: the HTTP client takes the POST body as a dict
sent_messages: Sequence[AllMessageValues] | None
loss: PayloadLoss
@dataclass(frozen=True, slots=True)
class AcceptedRewrites:
texts: list[str] | None # mutable-ok: guardrail inputs take a list
images: list[str] | None # mutable-ok: guardrail inputs take a list
tools: list[GuardrailToolParam] | None # mutable-ok: guardrail inputs take a list
rows: Sequence[AllMessageValues] | None
refused_any: bool
@dataclass(frozen=True, slots=True)
class _Unappliable:
pass
_UNAPPLIABLE: Final = _Unappliable()
def resolve_payload_policy(
*,
send_images: object,
exclude_payload_fields: Sequence[str] | None,
guardrail_name: str | None,
) -> PayloadPolicy:
policy: Final = PayloadPolicy(
send_images=_send_images(send_images),
exclude_fields=_resolve_exclude_fields(
config_values(exclude_payload_fields, option_name="exclude_payload_fields"), guardrail_name=guardrail_name
),
)
if policy.is_lossy:
verbose_proxy_logger.warning(
"Generic Guardrail API (%s): %s are not sent to the guardrail, so it can only enforce on what it is "
"sent, and it cannot rewrite what it did not see.",
guardrail_name,
sorted(policy.omitted_fields),
)
return policy
def _send_images(value: object) -> bool:
match value:
case None:
return True
case bool():
return value
case _:
raise ValueError(f"send_images must be a bool, got {value!r}")
def _resolve_exclude_fields(raw: Sequence[str], *, guardrail_name: str | None) -> frozenset[str]:
known: Final = frozenset(GenericGuardrailAPIRequest.model_fields)
unknown: Final = tuple(field for field in raw if field not in known)
if unknown:
verbose_proxy_logger.warning(
"Generic Guardrail API (%s): ignoring unknown exclude_payload_fields %s. Known fields: %s",
guardrail_name,
unknown,
sorted(known),
)
protected: Final = tuple(field for field in raw if field in PROTECTED_PAYLOAD_FIELDS)
if protected:
verbose_proxy_logger.warning(
"Generic Guardrail API (%s): exclude_payload_fields cannot drop %s, the guardrail needs them to "
"interpret the payload; they are still sent.",
guardrail_name,
protected,
)
excludable: Final = known - PROTECTED_PAYLOAD_FIELDS
return frozenset(field for field in raw if field in excludable)
def _omit_image(_url: str) -> str:
return IMAGE_OMITTED_PLACEHOLDER
def _content(row: JsonValue) -> JsonValue:
return row.get("content") if isinstance(row, dict) else None
def _holds_placeholder(part: JsonValue) -> bool:
return image_part_url(part) == IMAGE_OMITTED_PLACEHOLDER
def _altered_message_indices(unshaped: JsonValue, sent: JsonValue) -> frozenset[int]:
if not isinstance(unshaped, list) or not isinstance(sent, list):
return frozenset()
return frozenset(index for index, (row, sent_row) in enumerate(zip(unshaped, sent, strict=True)) if row != sent_row)
def shape_payload(dumped: Mapping[str, JsonValue], policy: PayloadPolicy) -> ShapedPayload:
omitted: Final = policy.omitted_fields
dumped_messages: Final = dumped.get("structured_messages")
sent_messages: Final = (
map_messages_image_urls(dumped_messages, _omit_image) if policy.shapes_messages else dumped_messages
)
shaped: Final = MappingProxyType({**dumped, "structured_messages": sent_messages})
return ShapedPayload(
body={key: value for key, value in shaped.items() if key not in omitted}, # mutable-ok: JSON POST body
sent_messages=structured_messages_from_json(sent_messages),
loss=PayloadLoss(
altered_message_indices=_altered_message_indices(dumped_messages, sent_messages),
texts_omitted="texts" in omitted,
messages_omitted="structured_messages" in omitted,
images_omitted="images" in omitted,
tools_omitted="tools" in omitted,
),
)
def _log_refused(field: str, detail: str, guardrail_name: str | None) -> None:
verbose_proxy_logger.warning(
"Generic Guardrail API (%s): ignoring the returned %s, %s.", guardrail_name, field, detail
)
def _refused(returned: object, *, field: str, omitted: bool, guardrail_name: str | None) -> bool:
if not returned or not omitted:
return False
_log_refused(field, "it was not sent to the guardrail", guardrail_name)
return True
def accepted_rewrites(
response: GenericGuardrailAPIResponse, loss: PayloadLoss, *, guardrail_name: str | None
) -> AcceptedRewrites:
refused: Final = MappingProxyType(
{
field: _refused(returned, field=field, omitted=omitted, guardrail_name=guardrail_name)
for field, returned, omitted in (
("texts", response.texts, loss.texts_omitted),
("images", response.images, loss.images_omitted),
("tools", response.tools, loss.tools_omitted),
("structured_messages", response.structured_messages, loss.messages_omitted),
)
}
)
return AcceptedRewrites(
texts=None if refused["texts"] else response.texts or None,
images=None if refused["images"] else response.images or None,
tools=None if refused["tools"] else response.tools or None,
rows=None if refused["structured_messages"] else response.structured_messages or None,
refused_any=any(refused.values()),
)
def _changes(accepted: object, original: object) -> bool:
return accepted is not None and as_json_value(accepted) != as_json_value(original)
def raise_if_intervention_was_refused(
*,
action: str,
accepted: AcceptedRewrites,
original_texts: Sequence[str],
original_images: Sequence[str] | None,
original_tools: Sequence[Mapping[str, object]] | None,
rows_written_back: bool,
guardrail_name: str | None,
) -> None:
"""A guardrail whose only real rewrite went to a field it was not sent would otherwise let the request
through unchanged, so it is rejected instead. An echo of what the caller sent changes nothing."""
applied: Final = rows_written_back or any(
_changes(accepted_value, original_value)
for accepted_value, original_value in (
(accepted.texts, original_texts),
(accepted.images, original_images),
(accepted.tools, original_tools),
)
)
if action == "GUARDRAIL_INTERVENED" and accepted.refused_any and not applied:
raise unappliable_request_rewrite(guardrail_name)
def _omitted_image_count(content: JsonValue) -> int:
if not isinstance(content, list):
return 0
return sum(1 for part in content if _holds_placeholder(part))
def _restored_part(returned: JsonValue, caller: JsonValue, sent: JsonValue) -> JsonValue | _Unappliable:
if returned == sent:
return caller
return _UNAPPLIABLE if _holds_placeholder(sent) or _holds_placeholder(returned) else returned
def _restored_content(returned: JsonValue, caller: JsonValue, sent: JsonValue) -> JsonValue | _Unappliable:
if returned == sent:
return caller
if not isinstance(returned, list) or not isinstance(caller, list) or not isinstance(sent, list):
return _UNAPPLIABLE
if not len(returned) == len(caller) == len(sent):
return _UNAPPLIABLE
parts: Final = tuple(_restored_part(*aligned) for aligned in zip(returned, caller, sent))
restored: Final = [part for part in parts if not isinstance(part, _Unappliable)] # mutable-ok: JSON array
return restored if len(restored) == len(parts) else _UNAPPLIABLE
def _restored_row(
returned: AllMessageValues,
caller: AllMessageValues,
sent: AllMessageValues,
*,
altered: bool,
) -> AllMessageValues | JsonValue | _Unappliable:
if returned is caller:
return caller
returned_json: Final = as_json_value(returned)
caller_content: Final = _content(as_json_value(caller))
if not altered:
adds_placeholder: Final = _omitted_image_count(_content(returned_json)) > _omitted_image_count(caller_content)
return _UNAPPLIABLE if adds_placeholder else returned
content: Final = _restored_content(_content(returned_json), caller_content, _content(as_json_value(sent)))
if isinstance(content, _Unappliable) or not isinstance(returned_json, dict):
return _UNAPPLIABLE
return {**returned_json, "content": content} # mutable-ok: JSON row
def restore_unseen_rows(
*,
rows: tuple[AllMessageValues, ...] | None,
caller: Sequence[AllMessageValues] | None,
sent: Sequence[AllMessageValues] | None,
loss: PayloadLoss,
guardrail_name: str | None,
) -> tuple[AllMessageValues, ...] | None:
"""At a row the guardrail saw only in part, accept its rewrite only where the parts line up, and put the
caller's image back into each part that still holds the placeholder. Anything else blocks the request,
since keeping the caller's row would silently drop the guardrail's rewrite."""
altered: Final = loss.altered_message_indices
if rows is None or not altered:
return rows
if caller is None or sent is None or not len(rows) == len(caller) == len(sent):
raise unappliable_request_rewrite(guardrail_name)
restored: Final = tuple(
_restored_row(row, caller_row, sent_row, altered=index in altered)
for index, (row, caller_row, sent_row) in enumerate(zip(rows, caller, sent))
)
accepted: Final = structured_messages_from_json(
[row for row in restored if not isinstance(row, _Unappliable)] # mutable-ok: JSON array
)
if accepted is None or len(accepted) != len(restored):
raise unappliable_request_rewrite(guardrail_name)
return tuple(accepted)

View file

@ -103,6 +103,30 @@ class GenericGuardrailAPIOptionalParams(BaseModel):
),
)
send_images: bool | None = Field(
default=None,
description=(
"If False, the top-level images field is not sent, and every inline image_url part in "
"structured_messages keeps its place but has its URL replaced by '[omitted]'. File, "
"audio and video parts are still sent as they are. Saves payload size for guardrails "
"that only inspect text. The guardrail cannot replace the caller's images: a rewritten "
"message gets the caller's image back in each part still holding '[omitted]', and a "
"rewrite that changes or moves an image part is rejected. Defaults to True in "
"GenericGuardrailAPI.__init__ when None."
),
)
exclude_payload_fields: tuple[str, ...] | None = Field(
default=None,
description=(
"Top-level guardrail request fields to leave out of the payload, e.g. "
"['request_headers', 'tools'], for guardrails that do not use them. Unknown fields "
"are ignored with a warning at init, and input_type and litellm_call_id are always "
"sent. A field that is not sent (texts, structured_messages, images or tools) cannot "
"be rewritten by the guardrail response."
),
)
class GenericGuardrailAPIConfigModel(
GuardrailConfigModel[GenericGuardrailAPIOptionalParams],
@ -159,7 +183,7 @@ def coerce_stream_holdback_value(value: Any) -> int:
return 0
def structured_messages_from_response(value: object) -> Sequence[AllMessageValues] | None:
def structured_messages_from_json(value: object) -> Sequence[AllMessageValues] | None:
if not isinstance(value, list):
return None
if not all(isinstance(message, Mapping) and isinstance(message.get("role"), str) for message in value):
@ -212,5 +236,5 @@ class GenericGuardrailAPIResponse:
images=data.get("images"),
tools=data.get("tools"),
stream_holdback_chars=stream_holdback_chars,
structured_messages=structured_messages_from_response(data.get("structured_messages")),
structured_messages=structured_messages_from_json(data.get("structured_messages")),
)

View file

@ -5,16 +5,23 @@ This test file tests the Generic Guardrail API implementation,
specifically focusing on metadata extraction and passing.
"""
import json
import os
from collections.abc import Callable, Mapping
from typing import Final
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
import pytest
from pydantic import JsonValue
import litellm
from litellm import ModelResponse
from litellm._version import version as litellm_version
from litellm.exceptions import GuardrailRaisedException, Timeout
from litellm.llms.anthropic.chat.guardrail_translation.handler import AnthropicMessagesHandler
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
from litellm.llms.openai.chat.guardrail_translation.handler import OpenAIChatCompletionsHandler
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
GenericGuardrailAPI,
@ -22,6 +29,8 @@ from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.generic_guardrail_api import (
_HEADER_PRESENT_PLACEHOLDER,
)
from litellm.types.llms.anthropic import AllAnthropicMessageValues
from litellm.types.llms.openai import AllMessageValues, ChatCompletionImageObject
from litellm.types.utils import Choices, Message
@ -721,6 +730,287 @@ class TestStructuredMessagesInResponse:
assert guardrailed_inputs["texts"] == ["[REDACTED]"]
_SSN: Final = "123-45-6789"
def _image_part() -> ChatCompletionImageObject:
return {"type": "image_url", "image_url": {"url": "data:image/png;base64,iVBORw0KGgo="}}
GuardrailAnswer = Callable[[Mapping[str, JsonValue]], Mapping[str, JsonValue]]
def _guardrail_answering(answer: GuardrailAnswer) -> GenericGuardrailAPI:
def serve(request: httpx.Request) -> httpx.Response:
return httpx.Response(200, json=answer(json.loads(request.content)))
return GenericGuardrailAPI(
api_base="https://guardrail.test/beta/litellm_basic_guardrail_api",
guardrail_name="pii-masker",
event_hook="pre_call",
default_on=True,
async_handler=AsyncHTTPHandler(transport=httpx.MockTransport(serve)),
)
def _masked(text: str) -> str:
return text.replace(_SSN, "[SSN]")
def _echo_every_row_and_mask_texts(request_json: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]:
return {
"action": "GUARDRAIL_INTERVENED",
"structured_messages": request_json["structured_messages"],
"texts": [_masked(text) for text in request_json["texts"]],
}
def _echo_first_row_and_mask_the_rest(request_json: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]:
first_row, *other_rows = request_json["structured_messages"]
return {
"action": "GUARDRAIL_INTERVENED",
"structured_messages": [first_row, *({**row, "content": _masked(row["content"])} for row in other_rows)],
}
def _masked_part(part: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]:
return {**part, "text": _masked(part["text"])} if part["type"] == "text" else part
def _masked_row(row: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]:
content: Final = row["content"]
if isinstance(content, str):
return {**row, "content": _masked(content)}
return {**row, "content": [_masked_part(part) for part in content]}
def _mask_every_row_and_text(request_json: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]:
return {
"action": "GUARDRAIL_INTERVENED",
"texts": [_masked(text) for text in request_json["texts"]],
"structured_messages": [_masked_row(row) for row in request_json["structured_messages"]],
}
def _guarded_text_part() -> Mapping[str, JsonValue]:
return {"type": "guarded_text", "text": "keep this guarded"}
def _pdf_document_part() -> Mapping[str, JsonValue]:
return {"type": "document", "source": {"type": "base64", "media_type": "application/pdf", "data": "JVBERi0="}}
def _nested(depth: int) -> JsonValue:
return {"leaf": "x"} if depth == 0 else {"nested": _nested(depth - 1)}
async def _llm_bound_messages(guardrail: GenericGuardrailAPI, messages: list[AllMessageValues]) -> object:
data: Final = await OpenAIChatCompletionsHandler().process_input_messages(
data={"model": "gpt-5.6", "messages": messages}, guardrail_to_apply=guardrail
)
return data["messages"]
async def _structured_messages_posted_for(rows: list[AllMessageValues]) -> list[JsonValue]:
posted: Final[list[JsonValue]] = [] # mutable-ok: records what the endpoint received
def record(request_json: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]:
posted.append(request_json["structured_messages"])
return {"action": "NONE"}
await _guardrail_answering(record).apply_guardrail(
inputs={"texts": ["hi"], "structured_messages": rows}, request_data={}, input_type="request"
)
return posted
async def _llm_bound_anthropic_messages(
guardrail: GenericGuardrailAPI, messages: list[AllAnthropicMessageValues]
) -> object:
data: Final = await AnthropicMessagesHandler().process_input_messages(
data={"model": "claude-opus-5-5", "max_tokens": 64, "messages": messages}, guardrail_to_apply=guardrail
)
return data["messages"]
class TestEchoedRowsReachingTheLLM:
@pytest.mark.asyncio
@pytest.mark.parametrize(
("messages", "expected"),
[
(
[{"role": "user", "content": [{"type": "text", "text": f"my ssn is {_SSN}"}, _image_part()]}],
[{"role": "user", "content": [{"type": "text", "text": "my ssn is [SSN]"}, _image_part()]}],
),
(
[
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": [{"type": "text", "text": f"ssn {_SSN}"}, _image_part()]},
{"role": "user", "content": f"again {_SSN}"},
],
[
{"role": "system", "content": "You are helpful."},
{"role": "user", "content": [{"type": "text", "text": "ssn [SSN]"}, _image_part()]},
{"role": "user", "content": "again [SSN]"},
],
),
(
[{"role": "user", "content": f"my ssn is {_SSN}"}],
[{"role": "user", "content": "my ssn is [SSN]"}],
),
],
ids=["multipart", "multipart_among_string_rows", "string_only"],
)
async def test_every_row_echoed_applies_the_masked_texts(
self, messages: list[AllMessageValues], expected: list[AllMessageValues]
) -> None:
guardrail: Final = _guardrail_answering(_echo_every_row_and_mask_texts)
llm_bound: Final = await _llm_bound_messages(guardrail, messages)
assert llm_bound == expected, "an unchanged echo of every row must leave the rewrite to texts"
@pytest.mark.asyncio
async def test_every_anthropic_content_block_row_echoed_applies_the_masked_texts(self) -> None:
guardrail: Final = _guardrail_answering(_echo_every_row_and_mask_texts)
llm_bound: Final = await _llm_bound_anthropic_messages(
guardrail,
[
{
"role": "user",
"content": [{"type": "text", "text": f"my ssn is {_SSN}"}, {"type": "text", "text": "ok"}],
}
],
)
assert llm_bound == [
{"role": "user", "content": [{"type": "text", "text": "my ssn is [SSN]"}, {"type": "text", "text": "ok"}]}
], "an unchanged echo of every content block row must leave the rewrite to texts"
@pytest.mark.asyncio
async def test_an_echoed_multipart_row_is_restored_to_the_callers_row(self) -> None:
guardrail: Final = _guardrail_answering(_echo_first_row_and_mask_the_rest)
llm_bound: Final = await _llm_bound_messages(
guardrail,
[
{"role": "user", "name": "pat", "content": [{"type": "text", "text": "what is this?"}, _image_part()]},
{"role": "user", "content": f"my ssn is {_SSN}"},
],
)
assert llm_bound == [
{"role": "user", "name": "pat", "content": [{"type": "text", "text": "what is this?"}, _image_part()]},
{"role": "user", "content": "my ssn is [SSN]"},
], "the echoed row must keep the caller's keys the request model drops"
@pytest.mark.asyncio
@pytest.mark.parametrize(
"unvalidated_part", [_guarded_text_part, _pdf_document_part], ids=["guarded_text", "pdf_document"]
)
async def test_a_row_holding_a_part_the_request_model_rejects_is_masked(
self, unvalidated_part: Callable[[], Mapping[str, JsonValue]]
) -> None:
guardrail: Final = _guardrail_answering(_mask_every_row_and_text)
llm_bound: Final = await _llm_bound_messages(
guardrail,
[
{"role": "user", "content": [{"type": "text", "text": f"my ssn is {_SSN}"}, unvalidated_part()]},
{"role": "user", "content": f"also {_SSN}"},
],
)
assert llm_bound == [
{"role": "user", "content": [{"type": "text", "text": "my ssn is [SSN]"}, unvalidated_part()]},
{"role": "user", "content": "also [SSN]"},
], "the guardrail must see the whole row so its masking of it reaches the LLM"
@pytest.mark.asyncio
async def test_an_echoed_row_holding_a_part_the_request_model_rejects_is_restored_to_the_callers_row(
self,
) -> None:
guardrail: Final = _guardrail_answering(_echo_first_row_and_mask_the_rest)
llm_bound: Final = await _llm_bound_messages(
guardrail,
[
{"role": "user", "name": "pat", "content": [{"type": "text", "text": "hi"}, _guarded_text_part()]},
{"role": "user", "content": f"my ssn is {_SSN}"},
],
)
assert llm_bound == [
{"role": "user", "name": "pat", "content": [{"type": "text", "text": "hi"}, _guarded_text_part()]},
{"role": "user", "content": "my ssn is [SSN]"},
], "the echoed row must keep the caller's keys the request model drops"
@pytest.mark.asyncio
async def test_rows_the_request_model_accepts_are_posted_as_it_dumps_them(self) -> None:
posted: Final = await _structured_messages_posted_for(
[
{
"role": "user",
"name": "pat",
"content": [
{"type": "text", "text": "hi"},
{**_image_part(), "cache_control": {"type": "ephemeral"}},
],
},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_1",
"type": "function",
"function": {"name": "lookup", "arguments": "{}"},
"index": 0,
}
],
},
{"role": "tool", "tool_call_id": "call_1", "content": "done"},
]
)
assert posted == [
[
{"role": "user", "content": [{"type": "text", "text": "hi"}, _image_part()]},
{
"role": "assistant",
"content": None,
"tool_calls": [
{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}
],
},
{"role": "tool", "tool_call_id": "call_1", "content": "done"},
]
], "rows the request model dumps in full must reach the guardrail exactly as before"
@pytest.mark.asyncio
async def test_a_row_holding_a_part_the_request_model_rejects_is_posted_with_the_callers_content(self) -> None:
posted: Final = await _structured_messages_posted_for(
[{"role": "user", "name": "pat", "content": [{"type": "text", "text": "hi"}, _pdf_document_part()]}]
)
assert posted == [[{"role": "user", "content": [{"type": "text", "text": "hi"}, _pdf_document_part()]}]], (
"only the content the request model emptied is taken from the caller, the row keys stay as dumped"
)
@pytest.mark.asyncio
async def test_a_deeply_nested_part_the_proxy_accepts_is_posted_in_full(self) -> None:
deep_document: Final = {"type": "document", "source": _nested(300)}
posted: Final = await _structured_messages_posted_for(
[{"role": "user", "content": [{"type": "text", "text": "hi"}, deep_document]}]
)
assert posted == [[{"role": "user", "content": [{"type": "text", "text": "hi"}, deep_document]}]], (
"content nested as deep as the proxy's own JSON parser allows must still reach the guardrail"
)
class TestImageSupport:
"""Test image handling in guardrail requests"""

View file

@ -1,12 +1,20 @@
"""Tests for the shared guardrail content extraction helpers."""
import copy
import pytest
from litellm.proxy.guardrails._content_utils import (
apply_redacted_messages_back,
as_json_value,
build_inspection_messages,
has_non_string_content,
image_part_url,
is_non_conversational_call_type,
is_string_batch_input,
iter_message_text,
map_content_image_urls,
map_messages_image_urls,
walk_user_text,
)
@ -741,3 +749,96 @@ def test_is_non_conversational_call_type_defaults_to_inspecting_unknown_call_typ
"""A call type this module has never heard of must still be inspected —
failing closed is the point of the deny-list."""
assert is_non_conversational_call_type("some_future_call_type") is False
def _nested(depth: int) -> dict[str, object]:
return {"leaf": "x"} if depth == 0 else {"nested": _nested(depth - 1)}
def test_as_json_value_keeps_content_nested_past_the_pydantic_serializer_limit_and_decodes_bytes():
assert as_json_value([{"type": "document", "source": _nested(600), "data": b"raw"}, ("a", 1)]) == [
{"type": "document", "source": _nested(600), "data": "raw"},
["a", 1],
], "nothing may be truncated to '...' and non-JSON types must become their JSON form"
# ── map_content_image_urls / map_messages_image_urls ─────────────────────────────
def _omit(url: str) -> str:
return f"[omitted {len(url)}]"
def test_map_messages_image_urls_covers_every_image_part_shape():
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "describe"},
{"type": "image_url", "image_url": {"url": "data:AAA", "detail": "low"}},
{"type": "image_url", "image_url": "data:BBBB"},
],
}
]
assert map_messages_image_urls(messages, _omit) == [
{
"role": "user",
"content": [
{"type": "text", "text": "describe"},
{"type": "image_url", "image_url": {"url": "[omitted 8]", "detail": "low"}},
{"type": "image_url", "image_url": "[omitted 9]"},
],
}
]
def test_map_messages_image_urls_does_not_mutate_its_input():
messages = [
{"role": "user", "content": [{"type": "image_url", "image_url": {"url": "data:AAA"}}]},
{"role": "user", "content": [{"type": "image_url", "image_url": "data:BBB"}]},
]
snapshot = copy.deepcopy(messages)
map_messages_image_urls(messages, _omit)
assert messages == snapshot
def test_map_content_image_urls_passes_non_image_parts_through_by_identity():
text_part = {"type": "text", "text": "describe"}
audio_part = {"type": "input_audio", "input_audio": {"data": "AAA", "format": "wav"}}
file_id_image = {"type": "input_image", "file_id": "file-1"}
content = [text_part, "bare text", audio_part, file_id_image]
mapped = map_content_image_urls(content, _omit)
assert all(after is before for after, before in zip(mapped, content, strict=True))
def test_map_messages_image_urls_passes_messages_without_list_content_through_by_identity():
string_message = {"role": "user", "content": "just text"}
tool_call_message = {"role": "assistant", "content": None, "tool_calls": []}
messages = [string_message, tool_call_message]
mapped = map_messages_image_urls(messages, _omit)
assert mapped[0] is string_message
assert mapped[1] is tool_call_message
assert map_messages_image_urls(None, _omit) is None
@pytest.mark.parametrize(
("part", "url"),
[
({"type": "image_url", "image_url": {"url": "data:AAA", "detail": "low"}}, "data:AAA"),
({"type": "image_url", "image_url": "data:BBB"}, "data:BBB"),
({"type": "image_url", "image_url": {"detail": "low"}}, None),
({"type": "input_image", "image_url": "data:CCC"}, None),
({"type": "text", "text": "data:DDD"}, None),
("data:EEE", None),
],
ids=["object", "bare_string", "no_url", "responses_shape", "text", "bare_text"],
)
def test_image_part_url_reads_only_chat_image_url_parts(part, url):
assert image_part_url(part) == url

View file

@ -0,0 +1,725 @@
import copy
import json
import logging
from collections.abc import Callable
import httpx
import pytest
from litellm.exceptions import GuardrailRaisedException
from litellm.llms.anthropic.chat.guardrail_translation.handler import AnthropicMessagesHandler
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
from litellm.llms.openai.chat.guardrail_translation.handler import OpenAIChatCompletionsHandler
from litellm.llms.openai.responses.guardrail_translation.handler import OpenAIResponsesHandler
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api.payload_policy import (
IMAGE_OMITTED_PLACEHOLDER,
)
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPIRequest
SSN = "123-45-6789"
IMAGE_URL = "data:image/png;base64,SECRETPIXELS"
OTHER_IMAGE_URL = "data:image/png;base64,OTHERPIXELS"
def _answer(body: dict) -> Callable[[dict], dict]:
return lambda _payload: body
class _FakeGuardrailEndpoint:
def __init__(self, respond: Callable[[dict], dict] = _answer({"action": "NONE"})):
self.payloads: list[dict] = []
self._respond = respond
async def post(self, *, url: str, json: dict, headers: dict, **_kwargs: object) -> httpx.Response:
sent = _wire_copy(json)
self.payloads.append(sent)
return httpx.Response(200, json=self._respond(sent), request=httpx.Request("POST", url))
class _LoggingObj:
litellm_call_id = "call-abc"
litellm_trace_id = "trace-abc"
def _wire_copy(payload: dict) -> dict:
return json.loads(json.dumps(payload))
def _guardrail(endpoint: _FakeGuardrailEndpoint, **options: object) -> GenericGuardrailAPI:
return GenericGuardrailAPI(
api_base="https://guardrail.example",
guardrail_name="payload-policy-test",
event_hook="pre_call",
default_on=True,
async_handler=endpoint,
**options,
)
def _image_message() -> dict:
return {
"role": "user",
"content": [
{"type": "text", "text": f"my ssn is {SSN}"},
{"type": "image_url", "image_url": {"url": IMAGE_URL, "detail": "low"}},
{"type": "image_url", "image_url": OTHER_IMAGE_URL},
],
}
def _text_message() -> dict:
return {"role": "user", "content": f"also {SSN}"}
def _texts() -> list[str]:
return [f"my ssn is {SSN}", f"also {SSN}"]
def _masked_text(text: str) -> str:
return text.replace(SSN, "[SSN]")
def _masked_part(part: dict) -> dict:
return {**part, "text": _masked_text(part["text"])} if part.get("type") == "text" else part
def _masked_row(row: dict) -> dict:
content = row["content"]
if isinstance(content, str):
return {**row, "content": _masked_text(content)}
return {**row, "content": [_masked_part(part) for part in content]}
def _masking_guardrail(payload: dict) -> dict:
return {
"action": "GUARDRAIL_INTERVENED",
"texts": [_masked_text(text) for text in payload.get("texts", [])],
"structured_messages": [_masked_row(row) for row in payload.get("structured_messages") or []],
}
def _rewrite_rows(rewrite: Callable[[int, dict], dict]) -> Callable[[dict], dict]:
def respond(payload: dict) -> dict:
rows = payload["structured_messages"]
return {
"action": "GUARDRAIL_INTERVENED",
"structured_messages": [rewrite(i, row) for i, row in enumerate(rows)],
}
return respond
async def _apply(guardrail: GenericGuardrailAPI, inputs: dict) -> dict:
return await guardrail.apply_guardrail(
inputs=inputs,
request_data={"proxy_server_request": {"headers": {"user-agent": "curl/8"}}},
input_type="request",
logging_obj=_LoggingObj(),
)
@pytest.mark.asyncio
async def test_defaults_send_every_request_field_unchanged():
endpoint = _FakeGuardrailEndpoint()
tools = [{"type": "function", "function": {"name": "lookup"}}]
await _apply(
_guardrail(endpoint),
{"texts": ["describe this"], "images": [IMAGE_URL], "tools": tools, "structured_messages": [_image_message()]},
)
payload = endpoint.payloads[0]
assert set(payload) == set(GenericGuardrailAPIRequest.model_fields)
assert payload["images"] == [IMAGE_URL]
assert payload["texts"] == ["describe this"]
assert payload["tools"] == tools
assert payload["structured_messages"] == [_image_message()]
assert payload["request_headers"] == {"user-agent": "curl/8"}
@pytest.mark.asyncio
async def test_defaults_still_accept_every_rewrite():
rewritten_tools = [{"type": "function", "function": {"name": "safe_lookup"}}]
endpoint = _FakeGuardrailEndpoint(
_answer(
{
"action": "GUARDRAIL_INTERVENED",
"texts": ["[MASKED]"],
"images": [OTHER_IMAGE_URL],
"tools": rewritten_tools,
}
)
)
result = await _apply(
_guardrail(endpoint),
{"texts": ["my ssn is 123"], "images": [IMAGE_URL], "tools": [{"type": "function", "function": {"name": "x"}}]},
)
assert result == {"texts": ["[MASKED]"], "images": [OTHER_IMAGE_URL], "tools": rewritten_tools}
@pytest.mark.asyncio
async def test_send_images_false_withholds_every_image_but_keeps_the_parts():
endpoint = _FakeGuardrailEndpoint()
messages = [_image_message(), _text_message()]
await _apply(
_guardrail(endpoint, send_images=False),
{"texts": _texts(), "images": [IMAGE_URL], "structured_messages": messages},
)
payload = endpoint.payloads[0]
assert "SECRETPIXELS" not in json.dumps(payload)
assert "OTHERPIXELS" not in json.dumps(payload)
assert "images" not in payload
assert payload["texts"] == _texts()
assert payload["structured_messages"] == [
{
"role": "user",
"content": [
{"type": "text", "text": f"my ssn is {SSN}"},
{"type": "image_url", "image_url": {"url": IMAGE_OMITTED_PLACEHOLDER, "detail": "low"}},
{"type": "image_url", "image_url": IMAGE_OMITTED_PLACEHOLDER},
],
},
_text_message(),
]
assert messages == [_image_message(), _text_message()]
@pytest.mark.asyncio
async def test_send_images_false_withholds_the_image_in_a_row_the_request_model_cannot_validate():
endpoint = _FakeGuardrailEndpoint()
row = {
"role": "user",
"content": [
{"type": "guarded_text", "text": "hi"},
{"type": "image_url", "image_url": {"url": IMAGE_URL}},
],
}
await _apply(_guardrail(endpoint, send_images=False), {"texts": ["hi"], "structured_messages": [row]})
assert "SECRETPIXELS" not in json.dumps(endpoint.payloads[0])
assert endpoint.payloads[0]["structured_messages"] == [
{
"role": "user",
"content": [
{"type": "guarded_text", "text": "hi"},
{"type": "image_url", "image_url": {"url": IMAGE_OMITTED_PLACEHOLDER}},
],
}
]
@pytest.mark.asyncio
async def test_send_images_false_keeps_the_callers_images():
endpoint = _FakeGuardrailEndpoint(
_answer({"action": "GUARDRAIL_INTERVENED", "texts": ["[MASKED]"], "images": [OTHER_IMAGE_URL]})
)
result = await _apply(_guardrail(endpoint, send_images=False), {"texts": ["ssn 1"], "images": [IMAGE_URL]})
assert result == {"texts": ["[MASKED]"], "images": [IMAGE_URL]}
@pytest.mark.asyncio
async def test_send_images_false_masks_the_llm_bound_text_and_keeps_the_callers_image():
guardrail = _guardrail(_FakeGuardrailEndpoint(_masking_guardrail), send_images=False)
data = {"model": "gpt-x", "messages": [_image_message(), _text_message()]}
result = await OpenAIChatCompletionsHandler().process_input_messages(data=data, guardrail_to_apply=guardrail)
assert result["messages"] == [
{
"role": "user",
"content": [
{"type": "text", "text": "my ssn is [SSN]"},
{"type": "image_url", "image_url": {"url": IMAGE_URL, "detail": "low"}},
{"type": "image_url", "image_url": OTHER_IMAGE_URL},
],
},
{"role": "user", "content": "also [SSN]"},
]
@pytest.mark.asyncio
async def test_send_images_false_masks_the_responses_input_and_keeps_the_callers_input_image():
guardrail = _guardrail(_FakeGuardrailEndpoint(_masking_guardrail), send_images=False)
data = {
"model": "gpt-x",
"input": [
{
"role": "user",
"content": [
{"type": "input_text", "text": f"my ssn is {SSN}"},
{"type": "input_image", "image_url": IMAGE_URL},
],
},
{"role": "user", "content": f"also {SSN}"},
],
}
result = await OpenAIResponsesHandler().process_input_messages(data=data, guardrail_to_apply=guardrail)
assert result["input"] == [
{
"role": "user",
"content": [
{"type": "input_text", "text": "my ssn is [SSN]"},
{"type": "input_image", "image_url": IMAGE_URL, "detail": "auto"},
],
},
{"role": "user", "content": "also [SSN]"},
]
@pytest.mark.asyncio
async def test_send_images_false_masks_the_anthropic_messages_and_keeps_the_callers_base64_image():
endpoint = _FakeGuardrailEndpoint(_masking_guardrail)
guardrail = _guardrail(endpoint, send_images=False)
image_block = {"type": "image", "source": {"type": "base64", "media_type": "image/png", "data": "SECRETPIXELS"}}
data = {
"model": "claude-x",
"max_tokens": 5,
"system": f"sys {SSN}",
"messages": [
{"role": "user", "content": [{"type": "text", "text": f"my ssn is {SSN}"}, image_block]},
{"role": "assistant", "content": "ok"},
{"role": "user", "content": f"also {SSN}"},
],
}
result = await AnthropicMessagesHandler().process_input_messages(data=data, guardrail_to_apply=guardrail)
assert "SECRETPIXELS" not in json.dumps(endpoint.payloads[0])
assert result["system"] == [{"type": "text", "text": "sys [SSN]"}]
assert result["messages"] == [
{"role": "user", "content": [{"type": "text", "text": "my ssn is [SSN]"}, image_block]},
{"role": "assistant", "content": [{"type": "text", "text": "ok"}]},
{"role": "user", "content": [{"type": "text", "text": "also [SSN]"}]},
]
@pytest.mark.asyncio
async def test_a_rewritten_image_row_gets_the_callers_images_back():
messages = [_image_message(), _text_message()]
endpoint = _FakeGuardrailEndpoint(_rewrite_rows(lambda i, row: _masked_row(row) if i == 0 else row))
result = await _apply(_guardrail(endpoint, send_images=False), {"texts": _texts(), "structured_messages": messages})
assert result == {"texts": _texts(), "structured_messages": [_masked_row(_image_message()), messages[1]]}
assert result["structured_messages"][1] is messages[1]
@pytest.mark.asyncio
async def test_an_echoed_placeholder_row_is_restored_to_the_callers_row():
messages = [_image_message(), _text_message()]
endpoint = _FakeGuardrailEndpoint(_rewrite_rows(lambda i, row: _masked_row(row) if i == 1 else row))
result = await _apply(_guardrail(endpoint, send_images=False), {"texts": _texts(), "structured_messages": messages})
assert result == {"texts": _texts(), "structured_messages": [messages[0], _masked_row(_text_message())]}
assert result["structured_messages"][0] is messages[0]
_TEXT_PART = {"type": "text", "text": "my ssn is [SSN]"}
_OBJECT_PLACEHOLDER = {"type": "image_url", "image_url": {"url": IMAGE_OMITTED_PLACEHOLDER, "detail": "low"}}
_BARE_PLACEHOLDER = {"type": "image_url", "image_url": IMAGE_OMITTED_PLACEHOLDER}
def _with_content(content: object) -> Callable[[int, dict], dict]:
return lambda i, row: {**row, "content": content} if i == 0 else row
@pytest.mark.asyncio
@pytest.mark.parametrize(
"rewrite",
[
_with_content([_TEXT_PART, {"type": "image_url", "image_url": {"url": OTHER_IMAGE_URL}}, _BARE_PLACEHOLDER]),
_with_content([_TEXT_PART, {"type": "text", "text": "y"}, _BARE_PLACEHOLDER]),
_with_content([_BARE_PLACEHOLDER, _OBJECT_PLACEHOLDER, _BARE_PLACEHOLDER]),
_with_content([{"type": "text", "text": "x"}]),
_with_content("my ssn is [SSN]"),
],
ids=[
"image_part_changed",
"image_part_replaced_by_text",
"text_part_replaced_by_placeholder",
"parts_dropped",
"content_flattened",
],
)
async def test_a_rewrite_that_does_not_line_up_with_a_withheld_image_blocks_the_request(rewrite):
endpoint = _FakeGuardrailEndpoint(_rewrite_rows(rewrite))
with pytest.raises(UnappliableRequestRewrite):
await _apply(
_guardrail(endpoint, send_images=False),
{"texts": _texts(), "structured_messages": [_image_message(), _text_message()]},
)
@pytest.mark.asyncio
async def test_misaligned_rows_block_the_request_when_an_image_was_withheld():
endpoint = _FakeGuardrailEndpoint(
_answer({"action": "GUARDRAIL_INTERVENED", "structured_messages": [{"role": "user", "content": "[MASKED]"}]})
)
with pytest.raises(UnappliableRequestRewrite):
await _apply(
_guardrail(endpoint, send_images=False),
{"texts": _texts(), "structured_messages": [_text_message(), _image_message()]},
)
@pytest.mark.asyncio
async def test_a_placeholder_copied_to_a_row_without_images_blocks_the_request():
endpoint = _FakeGuardrailEndpoint(
lambda payload: {
"action": "GUARDRAIL_INTERVENED",
"structured_messages": [payload["structured_messages"][0], payload["structured_messages"][0]],
}
)
with pytest.raises(UnappliableRequestRewrite):
await _apply(
_guardrail(endpoint, send_images=False),
{"texts": _texts(), "structured_messages": [_image_message(), _text_message()]},
)
def _row_with_a_part_the_request_model_cannot_validate() -> dict:
return {
"role": "user",
"content": [
{"type": "text", "text": f"my ssn is {SSN}"},
{"type": "image_url", "image_url": {"url": IMAGE_URL}},
{"type": "input_audio", "input_audio": {"data": "AAAA", "format": "flac"}},
],
}
async def _chat_messages_reaching_the_llm(guardrail: GenericGuardrailAPI, messages: list[dict]) -> list[dict]:
data = {"model": "gpt-x", "messages": messages}
return (await OpenAIChatCompletionsHandler().process_input_messages(data=data, guardrail_to_apply=guardrail))[
"messages"
]
def _masked_row_with_a_part_the_request_model_cannot_validate() -> dict:
return {
"role": "user",
"content": [
{"type": "text", "text": "my ssn is [SSN]"},
{"type": "image_url", "image_url": {"url": IMAGE_URL}},
{"type": "input_audio", "input_audio": {"data": "AAAA", "format": "flac"}},
],
}
@pytest.mark.asyncio
async def test_a_row_with_a_part_the_request_model_cannot_validate_reaches_the_llm_masked():
guardrail = _guardrail(_FakeGuardrailEndpoint(_masking_guardrail), send_images=False)
messages = await _chat_messages_reaching_the_llm(
guardrail, [_row_with_a_part_the_request_model_cannot_validate(), _text_message()]
)
assert messages == [
_masked_row_with_a_part_the_request_model_cannot_validate(),
{"role": "user", "content": "also [SSN]"},
]
@pytest.mark.asyncio
async def test_a_row_with_a_part_the_request_model_cannot_validate_reaches_the_llm_masked_at_defaults():
guardrail = _guardrail(_FakeGuardrailEndpoint(_masking_guardrail))
messages = await _chat_messages_reaching_the_llm(
guardrail, [_row_with_a_part_the_request_model_cannot_validate(), _text_message()]
)
assert messages == [
_masked_row_with_a_part_the_request_model_cannot_validate(),
{"role": "user", "content": "also [SSN]"},
]
@pytest.mark.asyncio
async def test_defaults_still_write_back_rewritten_image_rows():
messages = [_image_message(), _text_message()]
endpoint = _FakeGuardrailEndpoint(_rewrite_rows(lambda _i, row: {**row, "content": "[MASKED]"}))
result = await _apply(_guardrail(endpoint), {"texts": _texts(), "structured_messages": messages})
assert result == {
"texts": _texts(),
"structured_messages": [{"role": "user", "content": "[MASKED]"}, {"role": "user", "content": "[MASKED]"}],
}
@pytest.mark.asyncio
async def test_exclude_payload_fields_drops_only_the_named_fields(caplog):
endpoint = _FakeGuardrailEndpoint()
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
guardrail = _guardrail(
endpoint, exclude_payload_fields=["request_headers", "litellm_version", "input_type", "not_a_field"]
)
await _apply(guardrail, {"texts": ["hello"]})
payload = endpoint.payloads[0]
assert set(GenericGuardrailAPIRequest.model_fields) - set(payload) == {"request_headers", "litellm_version"}
assert payload["input_type"] == "request"
assert payload["litellm_call_id"] == "call-abc"
assert payload["texts"] == ["hello"]
assert any("not_a_field" in message for message in caplog.messages)
assert any("input_type" in message and "still sent" in message for message in caplog.messages)
@pytest.mark.asyncio
async def test_litellm_call_id_cannot_be_excluded():
endpoint = _FakeGuardrailEndpoint()
await _apply(_guardrail(endpoint, exclude_payload_fields=["litellm_call_id"]), {"texts": ["hello"]})
assert endpoint.payloads[0]["litellm_call_id"] == "call-abc"
@pytest.mark.asyncio
async def test_an_exclude_config_naming_only_unknown_or_protected_fields_is_not_lossy(caplog):
endpoint = _FakeGuardrailEndpoint(_answer({"action": "GUARDRAIL_INTERVENED", "texts": ["[MASKED]"]}))
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
guardrail = _guardrail(endpoint, exclude_payload_fields=["not_a_field", "input_type"])
result = await _apply(guardrail, {"texts": ["ssn 123"]})
assert set(endpoint.payloads[0]) == set(GenericGuardrailAPIRequest.model_fields)
assert result == {"texts": ["[MASKED]"]}
assert not any("can only enforce on what it is sent" in message for message in caplog.messages)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("field", "inputs", "response", "expected"),
[
("texts", {"texts": ["ssn 1", "ssn 2"]}, {"texts": ["[M]", "[M]"]}, {"texts": ["ssn 1", "ssn 2"]}),
("texts", {"texts": []}, {"texts": ["INVENTED"]}, {"texts": []}),
(
"tools",
{"texts": ["hi"], "tools": [{"type": "function", "function": {"name": "run"}}]},
{"tools": [{"type": "function"}]},
{"texts": ["hi"], "tools": [{"type": "function", "function": {"name": "run"}}]},
),
(
"images",
{"texts": ["hi"], "images": [IMAGE_URL]},
{"images": [OTHER_IMAGE_URL]},
{"texts": ["hi"], "images": [IMAGE_URL]},
),
(
"structured_messages",
{"texts": ["hi"], "structured_messages": [{"role": "user", "content": "hi"}]},
{"structured_messages": [{"role": "user", "content": "INVENTED"}]},
{"texts": ["hi"]},
),
(
"structured_messages",
{"texts": ["hi"], "structured_messages": []},
{"structured_messages": [{"role": "user", "content": "INVENTED"}]},
{"texts": ["hi"]},
),
(
"structured_messages",
{"texts": ["hi"]},
{"structured_messages": [{"role": "user", "content": "INVENTED"}]},
{"texts": ["hi"]},
),
],
ids=[
"texts",
"texts_empty",
"tools",
"images",
"structured_messages",
"structured_messages_empty",
"structured_messages_none",
],
)
async def test_an_excluded_field_is_not_sent_and_its_rewrite_is_refused(caplog, field, inputs, response, expected):
endpoint = _FakeGuardrailEndpoint(_answer({"action": "NONE", **response}))
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
result = await _apply(_guardrail(endpoint, exclude_payload_fields=[field]), inputs)
assert field not in endpoint.payloads[0]
assert result == expected
assert any(f"ignoring the returned {field}" in message for message in caplog.messages)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("excluded", "response"),
[
(["texts"], {"texts": ["[MASKED]"]}),
(["texts"], {"texts": ["[MASKED]"], "structured_messages": [{"role": "user", "content": "ssn 1"}]}),
(
["texts"],
{
"texts": ["[MASKED]"],
"structured_messages": [{"role": "user", "content": "ssn 1"}],
"tools": [{"type": "function", "function": {"name": "run"}}],
"images": [IMAGE_URL],
},
),
(["images"], {"images": [OTHER_IMAGE_URL]}),
(["tools"], {"tools": [{"type": "function"}]}),
(["structured_messages"], {"structured_messages": [{"role": "user", "content": "[MASKED]"}]}),
],
ids=[
"texts",
"texts_with_rows_echoed",
"texts_with_every_other_field_echoed",
"images",
"tools",
"structured_messages",
],
)
async def test_an_intervention_made_only_through_fields_that_were_not_sent_blocks_the_request(excluded, response):
endpoint = _FakeGuardrailEndpoint(_answer({"action": "GUARDRAIL_INTERVENED", **response}))
inputs = {
"texts": ["ssn 1"],
"images": [IMAGE_URL],
"tools": [{"type": "function", "function": {"name": "run"}}],
"structured_messages": [{"role": "user", "content": "ssn 1"}],
}
with pytest.raises(UnappliableRequestRewrite):
await _apply(_guardrail(endpoint, exclude_payload_fields=excluded), inputs)
@pytest.mark.asyncio
async def test_an_echo_counts_as_no_change_whatever_container_the_caller_used():
endpoint = _FakeGuardrailEndpoint(
_answer({"action": "GUARDRAIL_INTERVENED", "texts": ["[MASKED]"], "images": [IMAGE_URL]})
)
with pytest.raises(UnappliableRequestRewrite):
await _apply(
_guardrail(endpoint, exclude_payload_fields=["texts"]), {"texts": ["ssn 1"], "images": (IMAGE_URL,)}
)
@pytest.mark.asyncio
async def test_a_chat_request_is_rejected_when_the_only_masking_went_to_texts_that_were_not_sent():
def echo_rows_and_tools_and_mask_texts(payload: dict) -> dict:
return {
"action": "GUARDRAIL_INTERVENED",
"texts": [_masked_text(row["content"]) for row in payload["structured_messages"]],
"structured_messages": payload["structured_messages"],
"tools": payload["tools"],
}
guardrail = _guardrail(_FakeGuardrailEndpoint(echo_rows_and_tools_and_mask_texts), exclude_payload_fields=["texts"])
data = {
"model": "gpt-x",
"messages": [_text_message()],
"tools": [{"type": "function", "function": {"name": "lookup", "parameters": {"type": "object"}}}],
}
with pytest.raises(UnappliableRequestRewrite):
await OpenAIChatCompletionsHandler().process_input_messages(data=data, guardrail_to_apply=guardrail)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("response", "expected"),
[
({"texts": ["[MASKED]"], "images": [OTHER_IMAGE_URL]}, {"texts": ["[MASKED]"], "images": [IMAGE_URL]}),
(
{"images": [OTHER_IMAGE_URL], "structured_messages": [{"role": "user", "content": "[MASKED]"}]},
{
"texts": ["ssn 1"],
"images": [IMAGE_URL],
"structured_messages": [{"role": "user", "content": "[MASKED]"}],
},
),
(
{"images": [OTHER_IMAGE_URL], "tools": [{"type": "function", "function": {"name": "safe"}}]},
{"texts": ["ssn 1"], "images": [IMAGE_URL], "tools": [{"type": "function", "function": {"name": "safe"}}]},
),
({}, {"texts": ["ssn 1"], "images": [IMAGE_URL]}),
],
ids=["texts_applied", "rows_applied", "tools_applied", "nothing_returned"],
)
async def test_an_intervention_that_does_not_depend_on_an_unsent_field_goes_through(response, expected):
endpoint = _FakeGuardrailEndpoint(_answer({"action": "GUARDRAIL_INTERVENED", **response}))
result = await _apply(
_guardrail(endpoint, exclude_payload_fields=["images"]),
{"texts": ["ssn 1"], "images": [IMAGE_URL], "structured_messages": [{"role": "user", "content": "ssn 1"}]},
)
assert result == expected
@pytest.mark.asyncio
async def test_the_responses_handler_rejects_rows_it_did_not_send_even_when_texts_are_echoed():
endpoint = _FakeGuardrailEndpoint(
_answer(
{
"action": "GUARDRAIL_INTERVENED",
"texts": ["hello"],
"structured_messages": [{"role": "user", "content": "INVENTED"}],
}
)
)
guardrail = _guardrail(endpoint, exclude_payload_fields=["structured_messages"])
data = {"model": "gpt-x", "input": [{"role": "user", "content": [{"type": "input_text", "text": "hello"}]}]}
original_input = copy.deepcopy(data["input"])
with pytest.raises(UnappliableRequestRewrite):
await OpenAIResponsesHandler().process_input_messages(data=data, guardrail_to_apply=guardrail)
assert data["input"] == original_input
@pytest.mark.asyncio
async def test_blocked_still_blocks_a_shaped_payload():
endpoint = _FakeGuardrailEndpoint(_answer({"action": "BLOCKED", "blocked_reason": "no"}))
guardrail = _guardrail(endpoint, send_images=False, exclude_payload_fields=["texts", "tools"])
with pytest.raises(GuardrailRaisedException):
await _apply(guardrail, {"texts": ["hello"], "structured_messages": [_image_message()]})
@pytest.mark.parametrize(
"options",
[{"send_images": False}, {"exclude_payload_fields": ["request_headers"]}],
ids=["send_images", "exclude_payload_fields"],
)
def test_lossy_options_warn_that_the_guardrail_only_enforces_on_what_it_is_sent(caplog, options):
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
_guardrail(_FakeGuardrailEndpoint(), **options)
assert any("can only enforce on what it is sent" in message for message in caplog.messages)
def test_defaults_do_not_warn(caplog):
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
_guardrail(_FakeGuardrailEndpoint(), send_images=True, exclude_payload_fields=[])
assert not any("can only enforce on what it is sent" in message for message in caplog.messages)
@pytest.mark.parametrize(
"options",
[{"send_images": "false"}, {"send_images": 0}, {"exclude_payload_fields": "tools"}],
ids=["send_images_string", "send_images_int", "exclude_payload_fields_string"],
)
def test_a_mistyped_option_fails_at_init(options):
with pytest.raises(ValueError, match=next(iter(options))):
_guardrail(_FakeGuardrailEndpoint(), **options)