diff --git a/litellm/proxy/guardrails/_content_utils.py b/litellm/proxy/guardrails/_content_utils.py
index 7529fe99f52..d17252833c2 100644
--- a/litellm/proxy/guardrails/_content_utils.py
+++ b/litellm/proxy/guardrails/_content_utils.py
@@ -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
diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py
index e3511d46544..7deaaef77bb 100644
--- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py
+++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py
@@ -39,6 +39,11 @@ 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"),
+ max_messages=_get_config_value(litellm_params, optional_params, "max_messages"),
+ max_text_chars=_get_config_value(litellm_params, optional_params, "max_text_chars"),
+ strip_patterns=_get_config_value(litellm_params, optional_params, "strip_patterns"),
)
litellm.logging_callback_manager.add_litellm_callback(_generic_guardrail_api_callback)
diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/config_parsing.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/config_parsing.py
new file mode 100644
index 00000000000..10934c8e244
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/config_parsing.py
@@ -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
diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py
index 3d1a173635e..84bc5de99ea 100644
--- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py
+++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py
@@ -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,16 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import
)
from litellm.types.utils import GenericGuardrailAPIInputs
+from .payload_policy import (
+ PayloadLoss,
+ accepted_rewrites,
+ block_only_response,
+ 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 +165,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 +241,17 @@ 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,
+ max_messages: int | None = None,
+ max_text_chars: int | None = None,
+ strip_patterns: 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 +296,15 @@ 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,
+ max_messages=max_messages,
+ max_text_chars=max_text_chars,
+ strip_patterns=strip_patterns,
+ guardrail_name=kwargs.get("guardrail_name"),
+ )
+
# Set supported event hooks
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
@@ -345,24 +399,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 +542,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, guardrail_name=self.guardrail_name)
- # 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 +578,14 @@ class GenericGuardrailAPI(CustomGuardrail):
images=images,
tools=tools,
structured_messages=structured_messages,
- shown_messages=guardrail_request.structured_messages,
- guardrail_response=guardrail_response,
+ shown_messages=payload.sent_messages,
+ guardrail_response=block_only_response(
+ guardrail_response, payload, input_type=input_type, guardrail_name=self.guardrail_name
+ ),
+ 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)
diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/payload_policy.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/payload_policy.py
new file mode 100644
index 00000000000..ae69643aa4e
--- /dev/null
+++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/payload_policy.py
@@ -0,0 +1,512 @@
+"""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.
+"""
+
+import time
+from collections.abc import Callable, Iterable, Mapping, Sequence
+from dataclasses import dataclass
+from functools import reduce
+from itertools import accumulate, chain
+from types import MappingProxyType
+from typing import Final, Literal, assert_never
+
+import regex
+from pydantic import JsonValue
+
+from litellm._logging import verbose_proxy_logger
+from litellm.exceptions import GuardrailRaisedException
+from litellm.llms.base_llm.guardrail_translation.utils import (
+ message_slot_texts,
+ message_with_slot_texts,
+ 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]"
+
+MAX_STRIP_SUBSTITUTIONS: Final = 64
+
+MAX_STRIP_CALL_CHARS: Final = 100_000
+
+STRIP_TIMEOUT_SECONDS: Final = 0.1
+
+
+@dataclass(frozen=True, slots=True)
+class PayloadPolicy:
+ send_images: bool = True
+ exclude_fields: frozenset[str] = frozenset()
+ max_messages: int | None = None
+ max_text_chars: int | None = None
+ strip_patterns: tuple[regex.Pattern[str], ...] = ()
+
+ @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 shapes_text(self) -> bool:
+ return self.max_text_chars is not None or bool(self.strip_patterns)
+
+ @property
+ def lossy_options(self) -> tuple[str, ...]:
+ options: Final = (
+ ("send_images=False", not self.send_images),
+ (f"exclude_payload_fields={sorted(self.exclude_fields)}", bool(self.exclude_fields)),
+ (f"max_messages={self.max_messages}", self.max_messages is not None),
+ (f"max_text_chars={self.max_text_chars}", self.max_text_chars is not None),
+ ("strip_patterns", bool(self.strip_patterns)),
+ )
+ return tuple(option for option, is_set in options if is_set)
+
+ @property
+ def is_lossy(self) -> bool:
+ return bool(self.lossy_options)
+
+
+@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
+ text_shaped: 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,
+ max_messages: object,
+ max_text_chars: object,
+ strip_patterns: 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
+ ),
+ max_messages=_positive_int(max_messages, option_name="max_messages"),
+ max_text_chars=_positive_int(max_text_chars, option_name="max_text_chars"),
+ strip_patterns=_compile_strip_patterns(strip_patterns),
+ )
+ if policy.is_lossy:
+ verbose_proxy_logger.warning(
+ "Generic Guardrail API (%s): %s keep part of the request from the guardrail, so it can only enforce "
+ "on what it is sent, and it cannot rewrite what it did not see.",
+ guardrail_name,
+ ", ".join(policy.lossy_options),
+ )
+ return policy
+
+
+def _positive_int(value: object, *, option_name: str) -> int | None:
+ match value:
+ case None:
+ return None
+ case bool():
+ raise ValueError(f"{option_name} must be an int, got {value!r}")
+ case int() if value >= 1:
+ return value
+ case int():
+ raise ValueError(f"{option_name} must be >= 1 (got {value})")
+ case _:
+ raise ValueError(f"{option_name} must be an int, got {value!r}")
+
+
+def _compile_strip_patterns(raw: Sequence[str] | None) -> tuple[regex.Pattern[str], ...]:
+ try:
+ return tuple(regex.compile(pattern) for pattern in config_values(raw, option_name="strip_patterns"))
+ except regex.error as e:
+ raise ValueError(f"strip_patterns contains an invalid regex: {e}") from e
+
+
+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 _string_list(value: JsonValue) -> tuple[str, ...]:
+ return tuple(item for item in value if isinstance(item, str)) if isinstance(value, list) else ()
+
+
+def _windowed_messages(messages: JsonValue, max_messages: int | None) -> JsonValue:
+ if max_messages is None or not isinstance(messages, list) or len(messages) <= max_messages:
+ return messages
+ return messages[-max_messages:]
+
+
+def _strip(text: str, patterns: tuple[regex.Pattern[str], ...], deadline: float) -> str | None:
+ def strip_one(acc: str | None, pattern: regex.Pattern[str]) -> str | None:
+ remaining: Final = deadline - time.monotonic()
+ if acc is None or remaining <= 0:
+ return None
+ try:
+ return pattern.sub("", acc, count=MAX_STRIP_SUBSTITUTIONS, timeout=remaining)
+ except TimeoutError:
+ return None
+
+ return reduce(strip_one, patterns, text)
+
+
+def _stripped_fragments(
+ fragments: Iterable[str], policy: PayloadPolicy, guardrail_name: str | None
+) -> Mapping[str, str]:
+ if not policy.strip_patterns:
+ return MappingProxyType({})
+ distinct: Final = tuple(dict.fromkeys(fragments))
+ spent: Final = tuple(accumulate(len(fragment) for fragment in distinct))
+ within_budget: Final = tuple(fragment for fragment, total in zip(distinct, spent) if total <= MAX_STRIP_CALL_CHARS)
+ deadline: Final = time.monotonic() + STRIP_TIMEOUT_SECONDS
+ stripped: Final = tuple((fragment, _strip(fragment, policy.strip_patterns, deadline)) for fragment in within_budget)
+ unstripped: Final = len(distinct) - sum(1 for _, text in stripped if text is not None)
+ if unstripped:
+ verbose_proxy_logger.warning(
+ "Generic Guardrail API (%s): %d text(s) are sent unstripped, because strip_patterns only run on the "
+ "first %d characters of distinct text per guardrail call and for at most %s seconds.",
+ guardrail_name,
+ unstripped,
+ MAX_STRIP_CALL_CHARS,
+ STRIP_TIMEOUT_SECONDS,
+ )
+ return MappingProxyType({fragment: text for fragment, text in stripped if text is not None})
+
+
+def _text_shaper(stripped: Mapping[str, str], max_text_chars: int | None) -> Callable[[str], str]:
+ def shape(text: str) -> str:
+ kept: Final = stripped.get(text, text)
+ return kept if max_text_chars is None else kept[:max_text_chars]
+
+ return shape
+
+
+def _row_with_shaped_text(row: AllMessageValues, shape: Callable[[str], str]) -> AllMessageValues:
+ return message_with_slot_texts(row, tuple(shape(text) for text in message_slot_texts(row))) or row
+
+
+def shape_payload(
+ dumped: Mapping[str, JsonValue], policy: PayloadPolicy, *, guardrail_name: str | None
+) -> ShapedPayload:
+ omitted: Final = policy.omitted_fields
+ dumped_messages: Final = dumped.get("structured_messages")
+ retained_messages: Final = _windowed_messages(dumped_messages, policy.max_messages)
+ rows_windowed: Final = retained_messages is not dumped_messages
+ retained_rows: Final = structured_messages_from_json(retained_messages) or ()
+ dumped_texts: Final = dumped.get("texts")
+ texts: Final = _string_list(dumped_texts)
+ windowed_texts: Final = (
+ tuple(chain.from_iterable(message_slot_texts(row) for row in retained_rows)) if rows_windowed else texts
+ )
+ shape: Final = _text_shaper(
+ _stripped_fragments(
+ chain(windowed_texts, chain.from_iterable(message_slot_texts(row) for row in retained_rows)),
+ policy,
+ guardrail_name,
+ ),
+ policy.max_text_chars,
+ )
+ sent_texts: Final = tuple(shape(text) for text in windowed_texts) if policy.shapes_text else windowed_texts
+ texted_messages: Final = (
+ as_json_value([_row_with_shaped_text(row, shape) for row in retained_rows])
+ if policy.shapes_text and retained_rows
+ else retained_messages
+ )
+ sent_messages: Final = (
+ map_messages_image_urls(texted_messages, _omit_image) if policy.shapes_messages else texted_messages
+ )
+ shaped: Final = MappingProxyType(
+ {
+ **dumped,
+ "structured_messages": sent_messages,
+ "texts": None if dumped_texts is None else list(sent_texts), # mutable-ok: JSON texts is an array
+ }
+ )
+ 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=(
+ frozenset() if rows_windowed else _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,
+ text_shaped=rows_windowed or sent_texts != texts or texted_messages != retained_messages,
+ ),
+ )
+
+
+def _without_nulls(value: JsonValue) -> JsonValue:
+ if isinstance(value, dict):
+ return {key: _without_nulls(item) for key, item in value.items() if item is not None} # mutable-ok: JSON
+ if isinstance(value, list):
+ return [_without_nulls(item) for item in value] # mutable-ok: JSON array
+ return value
+
+
+def _rewrites(response: GenericGuardrailAPIResponse, body: Mapping[str, JsonValue]) -> bool:
+ returned: Final = (
+ ("texts", response.texts),
+ ("structured_messages", response.structured_messages),
+ ("images", response.images),
+ ("tools", response.tools),
+ )
+ return response.action == "GUARDRAIL_INTERVENED" or any(
+ value and _without_nulls(as_json_value(value)) != _without_nulls(body.get(field)) for field, value in returned
+ )
+
+
+def block_only_response(
+ response: GenericGuardrailAPIResponse,
+ payload: ShapedPayload,
+ *,
+ input_type: Literal["request", "response"],
+ guardrail_name: str | None,
+) -> GenericGuardrailAPIResponse:
+ if not payload.loss.text_shaped:
+ return response
+ if not _rewrites(response, payload.body):
+ return GenericGuardrailAPIResponse(action=response.action, stream_holdback_chars=response.stream_holdback_chars)
+ verbose_proxy_logger.warning(
+ "Generic Guardrail API (%s): the guardrail rewrote a %s that max_messages, max_text_chars or strip_patterns "
+ "shaped before it was sent. A rewrite of content it saw only in part cannot be applied, so the %s is "
+ "rejected. These options are for block-only guardrails.",
+ guardrail_name,
+ input_type,
+ input_type,
+ )
+ match input_type:
+ case "request":
+ raise unappliable_request_rewrite(guardrail_name)
+ case "response":
+ raise GuardrailRaisedException(
+ guardrail_name=guardrail_name,
+ message=f"Guardrail '{guardrail_name}' returned a rewrite that cannot be applied to this response",
+ should_wrap_with_default_message=False,
+ )
+ case _:
+ assert_never(input_type)
+
+
+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)
diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py
index 44e2cc2404f..d9f8fb4b3fe 100644
--- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py
+++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py
@@ -103,6 +103,77 @@ 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."
+ ),
+ )
+
+ max_messages: int | None = Field(
+ default=None,
+ ge=1,
+ description=(
+ "If set and a request has more than N structured_messages, only the last N are sent, "
+ "and texts is rebuilt from the text of those N messages. Calls without "
+ "structured_messages, such as embeddings, rerank or an LLM response, are not affected. "
+ "images and tool_calls are not windowed. Bounds payload size when the whole conversation "
+ "is re-sent every turn, but the system prompt and early turns fall out of the window. "
+ "For block-only or observe-only guardrails: on a windowed call BLOCKED still applies, "
+ "but any rewrite the guardrail returns fails the call. A failed request or response is "
+ "rejected with an error, and a failed stream is cut off after the chunks already sent."
+ ),
+ )
+
+ max_text_chars: int | None = Field(
+ default=None,
+ ge=1,
+ description=(
+ "If set, every text in texts and in structured_messages content is cut to this many "
+ "characters before sending, so a caller can put content the guardrail never sees after "
+ "the first N characters. For block-only or observe-only guardrails: when any text was "
+ "cut, BLOCKED still applies, but any rewrite the guardrail returns fails the call, with "
+ "the same errors as max_messages."
+ ),
+ )
+
+ strip_patterns: tuple[str, ...] | None = Field(
+ default=None,
+ description=(
+ "Regexes whose matches are removed from every text in texts and in structured_messages "
+ "content before sending, e.g. volatile boilerplate the guardrail does not need. Roles, "
+ "ids, tool calls, tools and metadata are never touched. A caller can hide content from "
+ "the guardrail by wrapping it in something a pattern matches. For block-only or "
+ "observe-only guardrails: when any text was stripped, BLOCKED still applies, but any "
+ "rewrite the guardrail returns fails the call, with the same errors as max_messages. "
+ "An invalid regex raises at init. Patterns use the regex package and run against "
+ "caller requests and LLM responses alike. Each pattern removes at most 64 matches per "
+ "text. Per guardrail call, only the first 100,000 characters of distinct text are "
+ "stripped and stripping stops after 0.1 seconds. A text past either limit is sent "
+ "unstripped in full with a warning. Stripping runs on the worker's event loop, so a slow "
+ "pattern blocks that worker, and every request on it, for up to 0.1 seconds per "
+ "guardrail call. Keep patterns linear-time: no nested quantifiers such as (a+)+ and no "
+ "lazy match up to a closing delimiter such as ."
+ ),
+ )
+
class GenericGuardrailAPIConfigModel(
GuardrailConfigModel[GenericGuardrailAPIOptionalParams],
@@ -159,7 +230,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 +283,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")),
)
diff --git a/pyproject.toml b/pyproject.toml
index 77a1a3fdb75..381ee30a5d0 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -34,6 +34,7 @@ dependencies = [
"pydantic-settings>=2.14.1,<3.0",
"jsonschema>=4.0.0,<5.0",
"boto3>=1.43.1,<2.0",
+ "regex>=2022.1.18",
]
[project.urls]
diff --git a/tests/code_coverage_tests/liccheck.ini b/tests/code_coverage_tests/liccheck.ini
index 8a3e880043b..0c5977bd3a4 100644
--- a/tests/code_coverage_tests/liccheck.ini
+++ b/tests/code_coverage_tests/liccheck.ini
@@ -178,3 +178,4 @@ hypothesis: >=6.165.10 # MPL 2.0 license
pytest-rerunfailures: >=15.1 # MPL 2.0 license
pytest-recording: >=0.13.4 # MIT license
expression: >=5.6.0 # MIT License - https://github.com/cognitedata/Expression/blob/main/LICENSE
+regex: >=2022.1.18 # Apache-2.0 AND CNRI-Python, both permissive and OSI approved - https://github.com/mrabarnett/mrab-regex/blob/hg/LICENSE.txt
diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py
index a5e79f84ef1..7cc958697cf 100644
--- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py
+++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py
@@ -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"""
diff --git a/tests/test_litellm/proxy/guardrails/test_content_utils.py b/tests/test_litellm/proxy/guardrails/test_content_utils.py
index 920ffc77095..af6f0fae17d 100644
--- a/tests/test_litellm/proxy/guardrails/test_content_utils.py
+++ b/tests/test_litellm/proxy/guardrails/test_content_utils.py
@@ -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
diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py
new file mode 100644
index 00000000000..e69de29bb2d
diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_payload_policy.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_payload_policy.py
new file mode 100644
index 00000000000..4b19061c4df
--- /dev/null
+++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_payload_policy.py
@@ -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)
diff --git a/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_text_shaping.py b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_text_shaping.py
new file mode 100644
index 00000000000..a7b89004d09
--- /dev/null
+++ b/tests/unit/proxy/guardrails/guardrail_hooks/generic_guardrail_api/test_text_shaping.py
@@ -0,0 +1,746 @@
+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.cohere.rerank.guardrail_translation.handler import CohereRerankHandler
+from litellm.llms.openai.chat.guardrail_translation.handler import OpenAIChatCompletionsHandler
+from litellm.llms.openai.embeddings.guardrail_translation.handler import OpenAIEmbeddingsHandler
+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 (
+ MAX_STRIP_CALL_CHARS,
+ MAX_STRIP_SUBSTITUTIONS,
+)
+from litellm.types.utils import Choices, Message, ModelResponse, ModelResponseStream
+
+SSN = "123-45-6789"
+TIMESTAMP = r"\d+"
+IMAGE_URL = "data:image/png;base64,PIXELS"
+
+
+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 __init__(self) -> None:
+ self.model_call_details: dict = {}
+
+
+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="text-shaping-test",
+ event_hook="pre_call",
+ default_on=True,
+ async_handler=endpoint,
+ **options,
+ )
+
+
+async def _apply(guardrail: GenericGuardrailAPI, inputs: dict, input_type: str = "request") -> dict:
+ return await guardrail.apply_guardrail(
+ inputs=inputs, request_data={}, input_type=input_type, logging_obj=_LoggingObj()
+ )
+
+
+async def _chat_request(guardrail: GenericGuardrailAPI, messages: list[dict]) -> list[dict]:
+ data = await OpenAIChatCompletionsHandler().process_input_messages(
+ data={"model": "gpt-x", "messages": copy.deepcopy(messages)}, guardrail_to_apply=guardrail
+ )
+ return data["messages"]
+
+
+def _user(content: object) -> dict:
+ return {"role": "user", "content": content}
+
+
+def _assistant(content: object) -> dict:
+ return {"role": "assistant", "content": content}
+
+
+def _tool_call_turn() -> dict:
+ return {
+ "role": "assistant",
+ "content": None,
+ "tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}],
+ }
+
+
+def _multimodal_turn() -> dict:
+ return _user(
+ [
+ {"type": "text", "text": "compare"},
+ {"type": "image_url", "image_url": {"url": IMAGE_URL}},
+ {"type": "text", "text": "these two"},
+ ]
+ )
+
+
+def _conversation() -> list[dict]:
+ return [
+ {"role": "system", "content": "rules 1"},
+ _user("hi"),
+ _assistant("noted 2"),
+ _user(f"my ssn is {SSN}"),
+ ]
+
+
+def _mask(text: str) -> str:
+ return text.replace(SSN, "[SSN]")
+
+
+def _masked_row(row: dict) -> dict:
+ content = row.get("content")
+ if isinstance(content, str):
+ return {**row, "content": _mask(content)}
+ if not isinstance(content, list):
+ return row
+ return {
+ **row,
+ "content": [
+ {**part, "text": _mask(part["text"])} if isinstance(part.get("text"), str) else part for part in content
+ ],
+ }
+
+
+def _mask_ssn(payload: dict) -> dict:
+ rows = payload.get("structured_messages")
+ return {
+ "action": "GUARDRAIL_INTERVENED",
+ "texts": [_mask(text) for text in payload.get("texts") or ()],
+ **({"structured_messages": [_masked_row(row) for row in rows]} if rows else {}),
+ }
+
+
+def _echo(payload: dict) -> dict:
+ return {"action": "NONE", "texts": payload.get("texts"), "structured_messages": payload.get("structured_messages")}
+
+
+@pytest.mark.asyncio
+async def test_defaults_send_the_whole_conversation_unchanged():
+ endpoint = _FakeGuardrailEndpoint()
+ long_text = "x 1 " * 20_000
+ messages = [*_conversation(), _user(long_text)]
+
+ await _chat_request(_guardrail(endpoint, max_messages=None, max_text_chars=None, strip_patterns=None), messages)
+
+ payload = endpoint.payloads[0]
+ assert payload["texts"] == ["rules 1", "hi", "noted 2", f"my ssn is {SSN}", long_text]
+ assert payload["structured_messages"] == messages
+
+
+@pytest.mark.asyncio
+async def test_max_messages_sends_the_last_rows_and_only_their_texts():
+ endpoint = _FakeGuardrailEndpoint()
+
+ await _chat_request(_guardrail(endpoint, max_messages=2), _conversation())
+
+ payload = endpoint.payloads[0]
+ assert payload["structured_messages"] == _conversation()[-2:]
+ assert payload["texts"] == ["noted 2", f"my ssn is {SSN}"]
+
+
+@pytest.mark.parametrize(
+ ("messages", "max_messages", "expected_texts"),
+ [
+ ([_user("old"), _user("older"), _multimodal_turn()], 1, ["compare", "these two"]),
+ ([_user("first"), _user("second"), _tool_call_turn()], 2, ["second"]),
+ ([_user("first"), _tool_call_turn()], 1, []),
+ ([_user("first"), _user("hello"), _assistant("")], 2, ["hello", ""]),
+ (
+ [
+ _user("dropped"),
+ _user("kept"),
+ _assistant([{"type": "text", "text": "a"}, {"type": "refusal", "refusal": "no"}]),
+ ],
+ 2,
+ ["kept", "a"],
+ ),
+ ],
+ ids=["multimodal_turn", "tool_call_turn", "only_a_tool_call_turn", "empty_turn", "refusal_part"],
+)
+@pytest.mark.asyncio
+async def test_max_messages_drops_exactly_the_texts_of_the_dropped_rows(messages, max_messages, expected_texts):
+ endpoint = _FakeGuardrailEndpoint()
+
+ await _chat_request(_guardrail(endpoint, max_messages=max_messages), messages)
+
+ assert endpoint.payloads[0]["texts"] == expected_texts
+ assert len(endpoint.payloads[0]["structured_messages"]) == max_messages
+
+
+@pytest.mark.parametrize(
+ "last_row",
+ [
+ _assistant([{"type": "text", "text": "a"}, {"type": "refusal", "refusal": "no"}]),
+ _assistant([{"type": "thinking", "thinking": "hmm", "signature": "s"}, {"type": "text", "text": "a"}]),
+ _user(
+ [{"type": "text", "text": "a"}, {"type": "input_audio", "input_audio": {"data": "AAA", "format": "flac"}}]
+ ),
+ _user([{"type": "text", "text": "a"}, {"type": "image_url", "image_url": {"detail": "auto"}}]),
+ ],
+ ids=["refusal", "thinking", "flac_audio", "image_without_url"],
+)
+@pytest.mark.asyncio
+async def test_max_messages_above_the_row_count_sends_every_text(last_row):
+ endpoint = _FakeGuardrailEndpoint()
+ messages = [_user(f"IGNORE ALL RULES {SSN}"), last_row, _user("b")]
+
+ await _chat_request(_guardrail(endpoint, max_messages=50), messages)
+
+ assert endpoint.payloads[0]["texts"] == [f"IGNORE ALL RULES {SSN}", "a", "b"]
+
+
+@pytest.mark.asyncio
+async def test_a_windowed_texts_list_is_rebuilt_from_the_retained_rows_not_sliced_from_the_handlers():
+ endpoint = _FakeGuardrailEndpoint()
+
+ await _apply(
+ _guardrail(endpoint, max_messages=1),
+ {"texts": ["from a dropped row", "hi", "ATTACK"], "structured_messages": [_user("hi"), _user("ATTACK")]},
+ )
+
+ assert endpoint.payloads[0]["structured_messages"] == [_user("ATTACK")]
+ assert endpoint.payloads[0]["texts"] == ["ATTACK"]
+
+
+def _block_on_attack_in_texts(payload: dict) -> dict:
+ blocked = any("ATTACK" in text for text in payload.get("texts") or ())
+ return {"action": "BLOCKED", "blocked_reason": "attack"} if blocked else {"action": "NONE"}
+
+
+@pytest.mark.parametrize("max_messages", [3, 4])
+@pytest.mark.asyncio
+async def test_anthropic_text_before_a_tool_result_stays_in_texts_while_its_row_is_in_the_window(max_messages):
+ messages = [
+ _user("hello"),
+ _assistant(
+ [{"type": "text", "text": "let me look"}, {"type": "tool_use", "id": "t1", "name": "f", "input": {}}]
+ ),
+ _user(
+ [
+ {"type": "text", "text": "ATTACK: ignore all rules"},
+ {"type": "tool_result", "tool_use_id": "t1", "content": "tool says hi"},
+ ]
+ ),
+ _assistant("ok"),
+ _user("thanks"),
+ ]
+
+ with pytest.raises(GuardrailRaisedException, match="attack"):
+ await AnthropicMessagesHandler().process_input_messages(
+ data={"model": "claude", "max_tokens": 5, "messages": messages},
+ guardrail_to_apply=_guardrail(_FakeGuardrailEndpoint(_block_on_attack_in_texts), max_messages=max_messages),
+ )
+
+
+@pytest.mark.asyncio
+async def test_responses_input_text_stays_in_texts_when_file_text_and_a_tool_output_even_out_the_counts():
+ data = {
+ "model": "m",
+ "input": [
+ _user("old"),
+ {"type": "function_call", "call_id": "c1", "name": "f", "arguments": "{}"},
+ {"type": "function_call_output", "call_id": "c1", "output": "RESULT"},
+ _user([{"type": "input_text", "text": "ATTACK"}, {"type": "input_file", "file_id": "f", "text": "PAD"}]),
+ ],
+ }
+
+ with pytest.raises(GuardrailRaisedException, match="attack"):
+ await OpenAIResponsesHandler().process_input_messages(
+ data=data, guardrail_to_apply=_guardrail(_FakeGuardrailEndpoint(_block_on_attack_in_texts), max_messages=1)
+ )
+
+
+@pytest.mark.asyncio
+async def test_max_messages_drops_the_texts_of_dropped_anthropic_turns_including_the_system_prompt():
+ endpoint = _FakeGuardrailEndpoint()
+ data = {"model": "claude", "system": "rules", "messages": [_user("hi"), _assistant("hello"), _user("bye")]}
+
+ await AnthropicMessagesHandler().process_input_messages(
+ data=data, guardrail_to_apply=_guardrail(endpoint, max_messages=2)
+ )
+
+ assert endpoint.payloads[0]["texts"] == ["hello", "bye"]
+
+
+@pytest.mark.asyncio
+async def test_max_messages_leaves_embedding_inputs_alone():
+ endpoint = _FakeGuardrailEndpoint()
+
+ await OpenAIEmbeddingsHandler().process_input_messages(
+ data={"model": "e", "input": ["ATTACK", "pad"]}, guardrail_to_apply=_guardrail(endpoint, max_messages=1)
+ )
+
+ assert endpoint.payloads[0]["texts"] == ["ATTACK", "pad"]
+
+
+@pytest.mark.asyncio
+async def test_max_messages_leaves_rerank_inputs_alone():
+ endpoint = _FakeGuardrailEndpoint()
+ data = {"model": "r", "query": "ATTACK query", "instruction": "rank", "documents": ["a"]}
+
+ await CohereRerankHandler().process_input_messages(
+ data=data, guardrail_to_apply=_guardrail(endpoint, max_messages=1)
+ )
+
+ assert "ATTACK query" in endpoint.payloads[0]["texts"]
+
+
+@pytest.mark.asyncio
+async def test_max_messages_leaves_every_choice_of_a_response_alone():
+ endpoint = _FakeGuardrailEndpoint()
+ response = ModelResponse(
+ choices=[
+ Choices(index=0, message=Message(role="assistant", content="LEAKED SECRET")),
+ Choices(index=1, message=Message(role="assistant", content="fine")),
+ ]
+ )
+
+ await OpenAIChatCompletionsHandler().process_output_response(
+ response=response, guardrail_to_apply=_guardrail(endpoint, max_messages=1)
+ )
+
+ assert endpoint.payloads[0]["texts"] == ["LEAKED SECRET", "fine"]
+
+
+@pytest.mark.parametrize(
+ "options",
+ [{"max_messages": 2}, {"max_text_chars": 12}, {"strip_patterns": [TIMESTAMP]}],
+ ids=["windowed", "truncated", "stripped"],
+)
+@pytest.mark.asyncio
+async def test_a_mask_of_a_shaped_request_fails_the_call(options):
+ endpoint = _FakeGuardrailEndpoint(_mask_ssn)
+
+ with pytest.raises(UnappliableRequestRewrite):
+ await _chat_request(_guardrail(endpoint, **options), _conversation())
+
+
+@pytest.mark.asyncio
+async def test_a_mask_of_a_shaped_response_fails_the_call():
+ response = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content=f"ssn {SSN} ..."))])
+
+ with pytest.raises(GuardrailRaisedException, match="cannot be applied to this response") as raised:
+ await OpenAIChatCompletionsHandler().process_output_response(
+ response=response, guardrail_to_apply=_guardrail(_FakeGuardrailEndpoint(_mask_ssn), max_text_chars=5)
+ )
+
+ assert raised.value.status_code == 400
+ assert raised.value.blocked_content is False
+
+
+def _long_answer() -> str:
+ return f"your ssn is {SSN}. " + "and more text " * 20
+
+
+@pytest.mark.asyncio
+async def test_an_echo_of_a_shaped_response_leaves_the_llm_output_whole():
+ response = ModelResponse(choices=[Choices(index=0, message=Message(role="assistant", content=_long_answer()))])
+
+ result = await OpenAIChatCompletionsHandler().process_output_response(
+ response=response, guardrail_to_apply=_guardrail(_FakeGuardrailEndpoint(_echo), max_text_chars=30)
+ )
+
+ assert result.choices[0].message.content == _long_answer()
+
+
+def _answer_chunks() -> list[ModelResponseStream]:
+ words = _long_answer().split(" ")
+ deltas = [
+ ModelResponseStream(choices=[{"index": 0, "delta": {"role": "assistant", "content": f"{word} "}}])
+ for word in words
+ ]
+ return [*deltas, ModelResponseStream(choices=[{"index": 0, "delta": {}, "finish_reason": "stop"}])]
+
+
+@pytest.mark.asyncio
+async def test_a_mask_of_a_shaped_stream_fails_the_stream():
+ with pytest.raises(GuardrailRaisedException, match="cannot be applied to this response"):
+ await OpenAIChatCompletionsHandler().process_output_streaming_response(
+ responses_so_far=_answer_chunks(),
+ guardrail_to_apply=_guardrail(_FakeGuardrailEndpoint(_mask_ssn), max_text_chars=30),
+ )
+
+
+@pytest.mark.asyncio
+async def test_an_echo_of_a_shaped_stream_passes_the_stream_through():
+ chunks = _answer_chunks()
+
+ result = await OpenAIChatCompletionsHandler().process_output_streaming_response(
+ responses_so_far=chunks, guardrail_to_apply=_guardrail(_FakeGuardrailEndpoint(_echo), max_text_chars=30)
+ )
+
+ assert [chunk.choices[0].delta.content for chunk in result] == [chunk.choices[0].delta.content for chunk in chunks]
+ assert "".join(chunk.choices[0].delta.content or "" for chunk in result) == f"{_long_answer()} "
+
+
+@pytest.mark.asyncio
+async def test_an_echo_of_a_shaped_payload_keeps_the_stream_holdback():
+ endpoint = _FakeGuardrailEndpoint(_answer({"action": "NONE", "texts": ["hel"], "stream_holdback_chars": [2]}))
+
+ result = await _apply(_guardrail(endpoint, max_text_chars=3), {"texts": ["hello"]}, input_type="response")
+
+ assert result == {"texts": ["hello"], "stream_holdback_chars": [2]}
+
+
+@pytest.mark.asyncio
+async def test_an_echo_without_null_fields_is_not_a_rewrite():
+ def echo_without_nulls(payload: dict) -> dict:
+ return json.loads(
+ json.dumps(_echo(payload)), object_hook=lambda obj: {k: v for k, v in obj.items() if v is not None}
+ )
+
+ messages = [
+ _user("hi " * 10),
+ _assistant([{"type": "text", "text": "calling"}, {"type": "tool_use", "id": "t1", "name": "f", "input": {}}]),
+ _user([{"type": "tool_result", "tool_use_id": "t1", "content": "r"}]),
+ ]
+ endpoint = _FakeGuardrailEndpoint(echo_without_nulls)
+
+ data = await AnthropicMessagesHandler().process_input_messages(
+ data={"model": "claude", "max_tokens": 5, "messages": copy.deepcopy(messages)},
+ guardrail_to_apply=_guardrail(endpoint, max_text_chars=5),
+ )
+
+ assert endpoint.payloads[0]["structured_messages"][1]["thinking_blocks"] is None
+ assert data["messages"] == messages
+
+
+@pytest.mark.asyncio
+async def test_windowing_combined_with_withheld_images_still_sends_and_blocks():
+ endpoint = _FakeGuardrailEndpoint(_answer({"action": "BLOCKED", "blocked_reason": "no"}))
+ messages = [_user("old"), _multimodal_turn(), _user("new")]
+
+ with pytest.raises(GuardrailRaisedException, match="no"):
+ await _apply(
+ _guardrail(endpoint, max_messages=2, send_images=False, fail_on_error=False),
+ {"texts": ["old", "compare", "these two", "new"], "structured_messages": messages},
+ )
+
+ assert endpoint.payloads[0]["texts"] == ["compare", "these two", "new"]
+ assert IMAGE_URL not in json.dumps(endpoint.payloads[0])
+
+
+@pytest.mark.parametrize(
+ "respond",
+ [
+ _answer({"action": "GUARDRAIL_INTERVENED"}),
+ _answer({"action": "NONE", "texts": ["[A]"]}),
+ _answer({"action": "NONE", "images": ["data:image/png;base64,OTHER"]}),
+ _answer({"action": "NONE", "tools": [{"type": "function", "function": {"name": "other"}}]}),
+ _answer({"action": "NONE", "structured_messages": [_user("[A]")]}),
+ ],
+ ids=["intervened", "texts", "images", "tools", "rows"],
+)
+@pytest.mark.asyncio
+async def test_any_returned_change_to_a_shaped_payload_fails_the_call(respond):
+ with pytest.raises(UnappliableRequestRewrite):
+ await _apply(
+ _guardrail(_FakeGuardrailEndpoint(respond), max_text_chars=3),
+ {"texts": ["hello"], "structured_messages": [_user("hello")], "images": [IMAGE_URL]},
+ )
+
+
+@pytest.mark.parametrize(
+ "options",
+ [{"max_messages": 2}, {"max_text_chars": 12}, {"strip_patterns": [TIMESTAMP]}],
+ ids=["windowed", "truncated", "stripped"],
+)
+@pytest.mark.asyncio
+async def test_an_echo_of_a_shaped_request_passes_the_callers_request_through(options):
+ endpoint = _FakeGuardrailEndpoint(_echo)
+
+ messages = await _chat_request(_guardrail(endpoint, **options), _conversation())
+
+ assert messages == _conversation()
+
+
+@pytest.mark.asyncio
+async def test_blocked_still_blocks_a_shaped_request():
+ endpoint = _FakeGuardrailEndpoint(_answer({"action": "BLOCKED", "blocked_reason": "no"}))
+
+ with pytest.raises(GuardrailRaisedException) as raised:
+ await _chat_request(
+ _guardrail(endpoint, max_messages=1, max_text_chars=3, strip_patterns=["ssn"]), _conversation()
+ )
+
+ assert raised.value.blocked_content is True
+
+
+@pytest.mark.parametrize(
+ "options",
+ [{"max_messages": 4}, {"max_text_chars": 100}, {"strip_patterns": [r"\[debug\]"]}],
+ ids=["window_covers_all", "nothing_too_long", "nothing_matches"],
+)
+@pytest.mark.asyncio
+async def test_a_mask_applies_when_the_configured_shaping_left_the_payload_untouched(options):
+ endpoint = _FakeGuardrailEndpoint(_mask_ssn)
+
+ messages = await _chat_request(_guardrail(endpoint, **options), _conversation())
+
+ assert messages == [*_conversation()[:-1], _user("my ssn is [SSN]")]
+
+
+@pytest.mark.asyncio
+async def test_max_text_chars_truncates_every_text():
+ endpoint = _FakeGuardrailEndpoint()
+
+ await _chat_request(_guardrail(endpoint, max_text_chars=4), [_user("abcdefgh"), _multimodal_turn(), _user("abc")])
+
+ payload = endpoint.payloads[0]
+ assert payload["texts"] == ["abcd", "comp", "thes", "abc"]
+ assert payload["structured_messages"] == [
+ _user("abcd"),
+ _user(
+ [
+ {"type": "text", "text": "comp"},
+ {"type": "image_url", "image_url": {"url": IMAGE_URL}},
+ {"type": "text", "text": "thes"},
+ ]
+ ),
+ _user("abc"),
+ ]
+
+
+@pytest.mark.asyncio
+async def test_strip_patterns_remove_matches_from_every_text():
+ endpoint = _FakeGuardrailEndpoint()
+
+ await _chat_request(_guardrail(endpoint, strip_patterns=[TIMESTAMP, r" \[debug\]"]), _conversation())
+
+ payload = endpoint.payloads[0]
+ assert payload["texts"] == ["rules ", "hi", "noted ", f"my ssn is {SSN}"]
+ assert [row["content"] for row in payload["structured_messages"]] == payload["texts"]
+
+
+@pytest.mark.asyncio
+async def test_strip_patterns_touch_only_text_never_roles_ids_tool_calls_or_tools():
+ tools = [{"type": "function", "function": {"name": "SECRET_lookup", "description": "SECRET"}}]
+ tool_call_turn = {
+ "role": "assistant",
+ "content": None,
+ "tool_calls": [
+ {"id": "call_SECRET", "type": "function", "function": {"name": "SECRET", "arguments": '"SECRET"'}}
+ ],
+ }
+ messages = [
+ _user("SECRET question"),
+ tool_call_turn,
+ {"role": "tool", "tool_call_id": "call_SECRET", "content": "SECRET answer"},
+ ]
+ endpoint = _FakeGuardrailEndpoint()
+
+ await _apply(
+ _guardrail(endpoint, strip_patterns=["SECRET"]),
+ {"texts": ["SECRET question", "SECRET answer"], "structured_messages": messages, "tools": tools},
+ )
+
+ payload = endpoint.payloads[0]
+ assert payload["texts"] == [" question", " answer"]
+ assert payload["tools"] == tools
+ assert payload["structured_messages"] == [
+ _user(" question"),
+ tool_call_turn,
+ {"role": "tool", "tool_call_id": "call_SECRET", "content": " answer"},
+ ]
+
+
+@pytest.mark.asyncio
+async def test_each_pattern_removes_a_bounded_number_of_matches_per_text():
+ endpoint = _FakeGuardrailEndpoint()
+
+ await _apply(_guardrail(endpoint, strip_patterns=["x"]), {"texts": ["x" * (MAX_STRIP_SUBSTITUTIONS + 3)]})
+
+ assert endpoint.payloads[0]["texts"] == ["xxx"]
+
+
+@pytest.mark.asyncio
+async def test_text_past_the_request_strip_budget_is_sent_unstripped(caplog):
+ first = "1" + "a" * (MAX_STRIP_CALL_CHARS - len("1") - 20)
+ second = "2" + "b" * 10
+ third = "3" + "c" * 20
+ endpoint = _FakeGuardrailEndpoint()
+
+ with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
+ await _apply(_guardrail(endpoint, strip_patterns=[TIMESTAMP]), {"texts": [first, second, third]})
+
+ assert endpoint.payloads[0]["texts"] == [first[len("1") :], "b" * 10, third]
+ assert any("text-shaping-test" in message and "sent unstripped" in message for message in caplog.messages)
+
+
+@pytest.mark.timeout(10)
+@pytest.mark.asyncio
+async def test_a_catastrophic_pattern_times_out_and_sends_the_text_unstripped(caplog):
+ backtracking = "PREFIX " + "a" * 40 + "!"
+ endpoint = _FakeGuardrailEndpoint()
+
+ with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
+ await _apply(
+ _guardrail(endpoint, strip_patterns=[r"PREFIX ", r"(a|aa)+$"]), {"texts": [backtracking, "PREFIX short"]}
+ )
+
+ assert endpoint.payloads[0]["texts"] == [backtracking, "PREFIX short"]
+ assert any("text-shaping-test" in message and "sent unstripped" in message for message in caplog.messages)
+
+
+@pytest.mark.asyncio
+async def test_a_text_repeated_in_texts_and_rows_is_charged_to_the_strip_budget_once():
+ text = "1" + "a" * (MAX_STRIP_CALL_CHARS * 2 // 3)
+ endpoint = _FakeGuardrailEndpoint()
+
+ await _apply(
+ _guardrail(endpoint, strip_patterns=[TIMESTAMP]),
+ {"texts": [text, text, "2b"], "structured_messages": [_user(text), _user(text), _user("2b")]},
+ )
+
+ stripped = text[len("1") :]
+ assert endpoint.payloads[0]["texts"] == [stripped, stripped, "b"]
+ assert endpoint.payloads[0]["structured_messages"] == [_user(stripped), _user(stripped), _user("b")]
+
+
+@pytest.mark.asyncio
+async def test_a_row_rewrite_fails_when_only_text_free_rows_were_windowed_out():
+ endpoint = _FakeGuardrailEndpoint(_mask_ssn)
+
+ with pytest.raises(UnappliableRequestRewrite):
+ await _chat_request(_guardrail(endpoint, max_messages=1), [_tool_call_turn(), _user(f"my ssn is {SSN}")])
+
+ assert endpoint.payloads[0]["texts"] == [f"my ssn is {SSN}"]
+
+
+@pytest.mark.asyncio
+async def test_a_row_rewrite_fails_when_only_row_text_was_stripped():
+ endpoint = _FakeGuardrailEndpoint(_mask_ssn)
+
+ with pytest.raises(UnappliableRequestRewrite):
+ await _apply(
+ _guardrail(endpoint, strip_patterns=[TIMESTAMP]),
+ {"texts": [f"ssn {SSN}"], "structured_messages": [_user(f"1ssn {SSN}")]},
+ )
+
+
+@pytest.mark.asyncio
+async def test_a_texts_rewrite_fails_when_only_row_text_was_stripped():
+ endpoint = _FakeGuardrailEndpoint(_answer({"action": "NONE", "texts": ["ssn [SSN]"]}))
+
+ with pytest.raises(UnappliableRequestRewrite):
+ await _apply(
+ _guardrail(endpoint, strip_patterns=[TIMESTAMP]),
+ {"texts": [f"ssn {SSN}"], "structured_messages": [_user(f"1ssn {SSN}")]},
+ )
+
+
+@pytest.mark.asyncio
+async def test_max_messages_keeps_the_handlers_texts_when_no_row_is_dropped():
+ endpoint = _FakeGuardrailEndpoint()
+
+ await _apply(_guardrail(endpoint, max_messages=5), {"texts": ["extra", "a"], "structured_messages": [_user("a")]})
+
+ assert endpoint.payloads[0]["texts"] == ["extra", "a"]
+
+
+@pytest.mark.asyncio
+async def test_combined_options_window_then_strip_then_truncate():
+ endpoint = _FakeGuardrailEndpoint()
+
+ await _chat_request(
+ _guardrail(endpoint, max_messages=2, strip_patterns=[TIMESTAMP], max_text_chars=6), _conversation()
+ )
+
+ payload = endpoint.payloads[0]
+ assert payload["texts"] == ["noted ", "my ssn"]
+ assert payload["structured_messages"] == [_assistant("noted "), _user("my ssn")]
+
+
+@pytest.mark.parametrize(
+ ("options", "message"),
+ [
+ ({"strip_patterns": ["("]}, "strip_patterns contains an invalid regex"),
+ ({"strip_patterns": "ssn"}, "strip_patterns must be a list of strings"),
+ ({"max_messages": 0}, "max_messages must be >= 1"),
+ ({"max_text_chars": 0}, "max_text_chars must be >= 1"),
+ ({"max_text_chars": 10.5}, "max_text_chars must be an int"),
+ ({"max_messages": True}, "max_messages must be an int"),
+ ({"max_messages": "3"}, "max_messages must be an int"),
+ ],
+ ids=["invalid_regex", "bare_string", "zero_messages", "zero_chars", "float", "bool", "string"],
+)
+def test_invalid_options_raise_at_init(options, message):
+ with pytest.raises(ValueError, match=message):
+ _guardrail(_FakeGuardrailEndpoint(), **options)
+
+
+@pytest.mark.parametrize(
+ ("options", "named"),
+ [
+ ({"max_messages": 3}, "max_messages=3"),
+ ({"max_text_chars": 100}, "max_text_chars=100"),
+ ({"strip_patterns": [TIMESTAMP]}, "strip_patterns"),
+ ],
+ ids=["max_messages", "max_text_chars", "strip_patterns"],
+)
+def test_text_shaping_options_warn_that_the_guardrail_only_enforces_on_what_it_is_sent(caplog, options, named):
+ with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
+ _guardrail(_FakeGuardrailEndpoint(), **options)
+
+ assert any("can only enforce on what it is sent" in message and named in message for message in caplog.messages)
+
+
+def test_unset_text_shaping_options_do_not_warn(caplog):
+ with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
+ _guardrail(_FakeGuardrailEndpoint(), max_messages=None, max_text_chars=None, strip_patterns=[])
+
+ assert not any("can only enforce on what it is sent" in message for message in caplog.messages)
+
+
+@pytest.mark.asyncio
+async def test_absent_texts_stay_absent_when_the_rows_are_windowed():
+ endpoint = _FakeGuardrailEndpoint()
+
+ await _apply(
+ _guardrail(endpoint, max_messages=1, max_text_chars=2), {"texts": None, "structured_messages": _conversation()}
+ )
+
+ assert endpoint.payloads[0]["texts"] is None
+ assert endpoint.payloads[0]["structured_messages"] == [_user("my")]
+
+
+@pytest.mark.asyncio
+async def test_shaping_never_mutates_the_callers_inputs():
+ messages = [_user("old 1"), _multimodal_turn()]
+ texts = ["old 1", "compare", "these two"]
+ snapshot = copy.deepcopy((messages, texts))
+
+ await _apply(
+ _guardrail(_FakeGuardrailEndpoint(), max_messages=1, max_text_chars=3, strip_patterns=[TIMESTAMP]),
+ {"texts": texts, "structured_messages": messages},
+ )
+
+ assert (messages, texts) == snapshot
diff --git a/uv.lock b/uv.lock
index 4b9dbaba39b..05a4a691410 100644
--- a/uv.lock
+++ b/uv.lock
@@ -4519,6 +4519,7 @@ dependencies = [
{ name = "pydantic-settings" },
{ name = "python-dotenv" },
{ name = "pyyaml" },
+ { name = "regex" },
{ name = "tiktoken" },
{ name = "tokenizers" },
]
@@ -4835,6 +4836,7 @@ requires-dist = [
{ name = "pyyaml", marker = "extra == 'cli'", specifier = ">=6.0.3,<7.0" },
{ name = "pyyaml", marker = "extra == 'proxy'", specifier = ">=6.0.3,<7.0" },
{ name = "redisvl", marker = "extra == 'extra-proxy'", specifier = ">=0.4.1,<1.0" },
+ { name = "regex", specifier = ">=2022.1.18" },
{ name = "requests", marker = "extra == 'cli'", specifier = ">=2.32.0,<3.0" },
{ name = "resend", marker = "extra == 'extra-proxy'", specifier = ">=2.23.0,<3.0" },
{ name = "restrictedpython", marker = "extra == 'proxy'", specifier = ">=8.5,<9.0" },