mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 8b7b7c2ce0 into f285229b51
This commit is contained in:
commit
70a46512ee
10 changed files with 1608 additions and 16 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -39,6 +39,8 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
|
|||
streaming_end_of_stream_only=_get_config_value(litellm_params, optional_params, "streaming_end_of_stream_only"),
|
||||
streaming_sampling_rate=_get_config_value(litellm_params, optional_params, "streaming_sampling_rate"),
|
||||
streaming_transform_mode=_get_config_value(litellm_params, optional_params, "streaming_transform_mode"),
|
||||
send_images=_get_config_value(litellm_params, optional_params, "send_images"),
|
||||
exclude_payload_fields=_get_config_value(litellm_params, optional_params, "exclude_payload_fields"),
|
||||
)
|
||||
|
||||
litellm.logging_callback_manager.add_litellm_callback(_generic_guardrail_api_callback)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -11,6 +11,8 @@ from collections.abc import Mapping, Sequence
|
|||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
|
||||
|
||||
import httpx
|
||||
from pydantic import JsonValue
|
||||
from typing_extensions import TypeIs
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._version import version as litellm_version
|
||||
|
|
@ -19,10 +21,13 @@ from litellm.integrations.custom_guardrail import (
|
|||
CustomGuardrail,
|
||||
log_guardrail_information,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import UnappliableRequestRewrite
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy.guardrails._content_utils import as_json_value
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
|
||||
|
|
@ -33,6 +38,15 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import
|
|||
)
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
from .payload_policy import (
|
||||
PayloadLoss,
|
||||
accepted_rewrites,
|
||||
raise_if_intervention_was_refused,
|
||||
resolve_payload_policy,
|
||||
restore_unseen_rows,
|
||||
shape_payload,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
|
||||
|
|
@ -150,6 +164,28 @@ def _extract_inbound_headers(
|
|||
return None
|
||||
|
||||
|
||||
def _is_part_list(value: object) -> TypeIs[list[object]]: # guard-ok: trivial isinstance narrowing
|
||||
return isinstance(value, list)
|
||||
|
||||
|
||||
def _row_as_sent(dumped: JsonValue, caller: Mapping[str, object]) -> JsonValue:
|
||||
"""The request model validates list content lazily and dumps a list holding
|
||||
any part it rejects as [], so such a row is sent with the caller's content."""
|
||||
caller_content: Final = caller.get("content")
|
||||
if not isinstance(dumped, dict) or not _is_part_list(caller_content):
|
||||
return dumped
|
||||
dumped_content: Final = dumped.get("content")
|
||||
if isinstance(dumped_content, list) and len(dumped_content) == len(caller_content):
|
||||
return dumped
|
||||
return {**dumped, "content": as_json_value(caller_content)}
|
||||
|
||||
|
||||
def _rows_as_sent(dumped_rows: JsonValue, caller_rows: Sequence[Mapping[str, object]] | None) -> JsonValue:
|
||||
if caller_rows is None or not isinstance(dumped_rows, list):
|
||||
return dumped_rows
|
||||
return [_row_as_sent(dumped, caller) for dumped, caller in zip(dumped_rows, caller_rows, strict=True)]
|
||||
|
||||
|
||||
def _structured_rows_to_write_back(
|
||||
original_rows: Sequence[AllMessageValues] | None,
|
||||
shown_rows: Sequence[AllMessageValues] | None,
|
||||
|
|
@ -204,9 +240,14 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
streaming_end_of_stream_only: bool | None = None,
|
||||
streaming_sampling_rate: int | None = None,
|
||||
streaming_transform_mode: Literal["block_only", "incremental_diff"] | None = None,
|
||||
send_images: bool | None = None,
|
||||
exclude_payload_fields: Sequence[str] | None = None,
|
||||
async_handler: AsyncHTTPHandler | None = None,
|
||||
**kwargs,
|
||||
):
|
||||
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
|
||||
self.async_handler = async_handler or get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.GuardrailCallback
|
||||
)
|
||||
self.headers = headers or {}
|
||||
self.extra_headers = extra_headers or []
|
||||
|
||||
|
|
@ -251,6 +292,12 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
"block_only" if streaming_transform_mode is None else streaming_transform_mode
|
||||
)
|
||||
|
||||
self._payload_policy: Final = resolve_payload_policy(
|
||||
send_images=send_images,
|
||||
exclude_payload_fields=exclude_payload_fields,
|
||||
guardrail_name=kwargs.get("guardrail_name"),
|
||||
)
|
||||
|
||||
# Set supported event hooks
|
||||
kwargs.setdefault("supported_event_hooks", list(self.get_supported_event_hooks()))
|
||||
|
||||
|
|
@ -345,24 +392,42 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
structured_messages: Sequence[AllMessageValues] | None,
|
||||
shown_messages: Sequence[AllMessageValues] | None,
|
||||
guardrail_response: GenericGuardrailAPIResponse,
|
||||
loss: PayloadLoss,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
# Action is NONE or no modifications needed
|
||||
return_inputs: Final = GenericGuardrailAPIInputs(texts=texts)
|
||||
if guardrail_response.texts:
|
||||
return_inputs["texts"] = guardrail_response.texts
|
||||
if guardrail_response.images:
|
||||
return_inputs["images"] = guardrail_response.images
|
||||
name: Final = self.guardrail_name
|
||||
accepted: Final = accepted_rewrites(guardrail_response, loss, guardrail_name=name)
|
||||
if accepted.texts:
|
||||
return_inputs["texts"] = accepted.texts
|
||||
if accepted.images:
|
||||
return_inputs["images"] = accepted.images
|
||||
elif images:
|
||||
return_inputs["images"] = images
|
||||
if guardrail_response.tools:
|
||||
return_inputs["tools"] = guardrail_response.tools
|
||||
if accepted.tools:
|
||||
return_inputs["tools"] = accepted.tools
|
||||
elif tools:
|
||||
return_inputs["tools"] = tools
|
||||
rows_to_write_back: Final = (
|
||||
_structured_rows_to_write_back(structured_messages, shown_messages, guardrail_response.structured_messages)
|
||||
if guardrail_response.structured_messages
|
||||
restore_unseen_rows(
|
||||
rows=_structured_rows_to_write_back(structured_messages, shown_messages, accepted.rows),
|
||||
caller=structured_messages,
|
||||
sent=shown_messages,
|
||||
loss=loss,
|
||||
guardrail_name=name,
|
||||
)
|
||||
if accepted.rows
|
||||
else None
|
||||
)
|
||||
raise_if_intervention_was_refused(
|
||||
action=guardrail_response.action,
|
||||
accepted=accepted,
|
||||
original_texts=texts,
|
||||
original_images=images,
|
||||
original_tools=tools,
|
||||
rows_written_back=rows_to_write_back is not None,
|
||||
guardrail_name=name,
|
||||
)
|
||||
if rows_to_write_back is not None:
|
||||
return_inputs["structured_messages"] = list(rows_to_write_back) # mutable-ok: guardrail inputs take a list
|
||||
if guardrail_response.stream_holdback_chars is not None:
|
||||
|
|
@ -470,12 +535,15 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
)
|
||||
|
||||
headers: Final = self._build_request_headers()
|
||||
# The model's list content is a lazy iterator that this dump consumes, so it cannot be read again
|
||||
dumped: Final[Mapping[str, JsonValue]] = guardrail_request.model_dump(mode="json")
|
||||
sent_messages: Final = _rows_as_sent(dumped.get("structured_messages"), structured_messages)
|
||||
request_json: Final = {**dumped, "structured_messages": sent_messages} # mutable-ok: JSON POST body
|
||||
payload: Final = shape_payload(request_json, self._payload_policy)
|
||||
|
||||
# Make the API request
|
||||
# Use mode="json" to ensure all iterables are converted to lists
|
||||
response: Final = await self.async_handler.post(
|
||||
url=self.api_base,
|
||||
json=guardrail_request.model_dump(mode="json"),
|
||||
json=payload.body,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
|
|
@ -503,11 +571,12 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
images=images,
|
||||
tools=tools,
|
||||
structured_messages=structured_messages,
|
||||
shown_messages=guardrail_request.structured_messages,
|
||||
shown_messages=payload.sent_messages,
|
||||
guardrail_response=guardrail_response,
|
||||
loss=payload.loss,
|
||||
)
|
||||
|
||||
except GuardrailRaisedException:
|
||||
except (GuardrailRaisedException, UnappliableRequestRewrite):
|
||||
raise
|
||||
except Timeout as e:
|
||||
return self._handle_guardrail_request_error(e, inputs, input_type, logging_obj)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,308 @@
|
|||
"""What the Generic Guardrail API endpoint is sent, and what it may write back.
|
||||
|
||||
Shaping is lossy, so every shaped payload carries a ``PayloadLoss``. Caller content the
|
||||
guardrail did not see in full is never replaced by the guardrail's response.
|
||||
"""
|
||||
|
||||
from collections.abc import Mapping, Sequence
|
||||
from dataclasses import dataclass
|
||||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from pydantic import JsonValue
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import unappliable_request_rewrite
|
||||
from litellm.proxy.guardrails._content_utils import as_json_value, image_part_url, map_messages_image_urls
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
|
||||
GenericGuardrailAPIRequest,
|
||||
GenericGuardrailAPIResponse,
|
||||
GuardrailToolParam,
|
||||
structured_messages_from_json,
|
||||
)
|
||||
|
||||
from .config_parsing import config_values
|
||||
|
||||
PROTECTED_PAYLOAD_FIELDS: Final = frozenset({"input_type", "litellm_call_id"})
|
||||
|
||||
IMAGE_OMITTED_PLACEHOLDER: Final = "[omitted]"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PayloadPolicy:
|
||||
send_images: bool = True
|
||||
exclude_fields: frozenset[str] = frozenset()
|
||||
|
||||
@property
|
||||
def omitted_fields(self) -> frozenset[str]:
|
||||
return self.exclude_fields if self.send_images else self.exclude_fields | frozenset(("images",))
|
||||
|
||||
@property
|
||||
def shapes_messages(self) -> bool:
|
||||
return not self.send_images
|
||||
|
||||
@property
|
||||
def is_lossy(self) -> bool:
|
||||
return bool(self.omitted_fields)
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class PayloadLoss:
|
||||
altered_message_indices: frozenset[int] = frozenset()
|
||||
texts_omitted: bool = False
|
||||
messages_omitted: bool = False
|
||||
images_omitted: bool = False
|
||||
tools_omitted: bool = False
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ShapedPayload:
|
||||
body: dict[str, JsonValue] # mutable-ok: the HTTP client takes the POST body as a dict
|
||||
sent_messages: Sequence[AllMessageValues] | None
|
||||
loss: PayloadLoss
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class AcceptedRewrites:
|
||||
texts: list[str] | None # mutable-ok: guardrail inputs take a list
|
||||
images: list[str] | None # mutable-ok: guardrail inputs take a list
|
||||
tools: list[GuardrailToolParam] | None # mutable-ok: guardrail inputs take a list
|
||||
rows: Sequence[AllMessageValues] | None
|
||||
refused_any: bool
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _Unappliable:
|
||||
pass
|
||||
|
||||
|
||||
_UNAPPLIABLE: Final = _Unappliable()
|
||||
|
||||
|
||||
def resolve_payload_policy(
|
||||
*,
|
||||
send_images: object,
|
||||
exclude_payload_fields: Sequence[str] | None,
|
||||
guardrail_name: str | None,
|
||||
) -> PayloadPolicy:
|
||||
policy: Final = PayloadPolicy(
|
||||
send_images=_send_images(send_images),
|
||||
exclude_fields=_resolve_exclude_fields(
|
||||
config_values(exclude_payload_fields, option_name="exclude_payload_fields"), guardrail_name=guardrail_name
|
||||
),
|
||||
)
|
||||
if policy.is_lossy:
|
||||
verbose_proxy_logger.warning(
|
||||
"Generic Guardrail API (%s): %s are not sent to the guardrail, so it can only enforce on what it is "
|
||||
"sent, and it cannot rewrite what it did not see.",
|
||||
guardrail_name,
|
||||
sorted(policy.omitted_fields),
|
||||
)
|
||||
return policy
|
||||
|
||||
|
||||
def _send_images(value: object) -> bool:
|
||||
match value:
|
||||
case None:
|
||||
return True
|
||||
case bool():
|
||||
return value
|
||||
case _:
|
||||
raise ValueError(f"send_images must be a bool, got {value!r}")
|
||||
|
||||
|
||||
def _resolve_exclude_fields(raw: Sequence[str], *, guardrail_name: str | None) -> frozenset[str]:
|
||||
known: Final = frozenset(GenericGuardrailAPIRequest.model_fields)
|
||||
unknown: Final = tuple(field for field in raw if field not in known)
|
||||
if unknown:
|
||||
verbose_proxy_logger.warning(
|
||||
"Generic Guardrail API (%s): ignoring unknown exclude_payload_fields %s. Known fields: %s",
|
||||
guardrail_name,
|
||||
unknown,
|
||||
sorted(known),
|
||||
)
|
||||
protected: Final = tuple(field for field in raw if field in PROTECTED_PAYLOAD_FIELDS)
|
||||
if protected:
|
||||
verbose_proxy_logger.warning(
|
||||
"Generic Guardrail API (%s): exclude_payload_fields cannot drop %s, the guardrail needs them to "
|
||||
"interpret the payload; they are still sent.",
|
||||
guardrail_name,
|
||||
protected,
|
||||
)
|
||||
excludable: Final = known - PROTECTED_PAYLOAD_FIELDS
|
||||
return frozenset(field for field in raw if field in excludable)
|
||||
|
||||
|
||||
def _omit_image(_url: str) -> str:
|
||||
return IMAGE_OMITTED_PLACEHOLDER
|
||||
|
||||
|
||||
def _content(row: JsonValue) -> JsonValue:
|
||||
return row.get("content") if isinstance(row, dict) else None
|
||||
|
||||
|
||||
def _holds_placeholder(part: JsonValue) -> bool:
|
||||
return image_part_url(part) == IMAGE_OMITTED_PLACEHOLDER
|
||||
|
||||
|
||||
def _altered_message_indices(unshaped: JsonValue, sent: JsonValue) -> frozenset[int]:
|
||||
if not isinstance(unshaped, list) or not isinstance(sent, list):
|
||||
return frozenset()
|
||||
return frozenset(index for index, (row, sent_row) in enumerate(zip(unshaped, sent, strict=True)) if row != sent_row)
|
||||
|
||||
|
||||
def shape_payload(dumped: Mapping[str, JsonValue], policy: PayloadPolicy) -> ShapedPayload:
|
||||
omitted: Final = policy.omitted_fields
|
||||
dumped_messages: Final = dumped.get("structured_messages")
|
||||
sent_messages: Final = (
|
||||
map_messages_image_urls(dumped_messages, _omit_image) if policy.shapes_messages else dumped_messages
|
||||
)
|
||||
shaped: Final = MappingProxyType({**dumped, "structured_messages": sent_messages})
|
||||
return ShapedPayload(
|
||||
body={key: value for key, value in shaped.items() if key not in omitted}, # mutable-ok: JSON POST body
|
||||
sent_messages=structured_messages_from_json(sent_messages),
|
||||
loss=PayloadLoss(
|
||||
altered_message_indices=_altered_message_indices(dumped_messages, sent_messages),
|
||||
texts_omitted="texts" in omitted,
|
||||
messages_omitted="structured_messages" in omitted,
|
||||
images_omitted="images" in omitted,
|
||||
tools_omitted="tools" in omitted,
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _log_refused(field: str, detail: str, guardrail_name: str | None) -> None:
|
||||
verbose_proxy_logger.warning(
|
||||
"Generic Guardrail API (%s): ignoring the returned %s, %s.", guardrail_name, field, detail
|
||||
)
|
||||
|
||||
|
||||
def _refused(returned: object, *, field: str, omitted: bool, guardrail_name: str | None) -> bool:
|
||||
if not returned or not omitted:
|
||||
return False
|
||||
_log_refused(field, "it was not sent to the guardrail", guardrail_name)
|
||||
return True
|
||||
|
||||
|
||||
def accepted_rewrites(
|
||||
response: GenericGuardrailAPIResponse, loss: PayloadLoss, *, guardrail_name: str | None
|
||||
) -> AcceptedRewrites:
|
||||
refused: Final = MappingProxyType(
|
||||
{
|
||||
field: _refused(returned, field=field, omitted=omitted, guardrail_name=guardrail_name)
|
||||
for field, returned, omitted in (
|
||||
("texts", response.texts, loss.texts_omitted),
|
||||
("images", response.images, loss.images_omitted),
|
||||
("tools", response.tools, loss.tools_omitted),
|
||||
("structured_messages", response.structured_messages, loss.messages_omitted),
|
||||
)
|
||||
}
|
||||
)
|
||||
return AcceptedRewrites(
|
||||
texts=None if refused["texts"] else response.texts or None,
|
||||
images=None if refused["images"] else response.images or None,
|
||||
tools=None if refused["tools"] else response.tools or None,
|
||||
rows=None if refused["structured_messages"] else response.structured_messages or None,
|
||||
refused_any=any(refused.values()),
|
||||
)
|
||||
|
||||
|
||||
def _changes(accepted: object, original: object) -> bool:
|
||||
return accepted is not None and as_json_value(accepted) != as_json_value(original)
|
||||
|
||||
|
||||
def raise_if_intervention_was_refused(
|
||||
*,
|
||||
action: str,
|
||||
accepted: AcceptedRewrites,
|
||||
original_texts: Sequence[str],
|
||||
original_images: Sequence[str] | None,
|
||||
original_tools: Sequence[Mapping[str, object]] | None,
|
||||
rows_written_back: bool,
|
||||
guardrail_name: str | None,
|
||||
) -> None:
|
||||
"""A guardrail whose only real rewrite went to a field it was not sent would otherwise let the request
|
||||
through unchanged, so it is rejected instead. An echo of what the caller sent changes nothing."""
|
||||
applied: Final = rows_written_back or any(
|
||||
_changes(accepted_value, original_value)
|
||||
for accepted_value, original_value in (
|
||||
(accepted.texts, original_texts),
|
||||
(accepted.images, original_images),
|
||||
(accepted.tools, original_tools),
|
||||
)
|
||||
)
|
||||
if action == "GUARDRAIL_INTERVENED" and accepted.refused_any and not applied:
|
||||
raise unappliable_request_rewrite(guardrail_name)
|
||||
|
||||
|
||||
def _omitted_image_count(content: JsonValue) -> int:
|
||||
if not isinstance(content, list):
|
||||
return 0
|
||||
return sum(1 for part in content if _holds_placeholder(part))
|
||||
|
||||
|
||||
def _restored_part(returned: JsonValue, caller: JsonValue, sent: JsonValue) -> JsonValue | _Unappliable:
|
||||
if returned == sent:
|
||||
return caller
|
||||
return _UNAPPLIABLE if _holds_placeholder(sent) or _holds_placeholder(returned) else returned
|
||||
|
||||
|
||||
def _restored_content(returned: JsonValue, caller: JsonValue, sent: JsonValue) -> JsonValue | _Unappliable:
|
||||
if returned == sent:
|
||||
return caller
|
||||
if not isinstance(returned, list) or not isinstance(caller, list) or not isinstance(sent, list):
|
||||
return _UNAPPLIABLE
|
||||
if not len(returned) == len(caller) == len(sent):
|
||||
return _UNAPPLIABLE
|
||||
parts: Final = tuple(_restored_part(*aligned) for aligned in zip(returned, caller, sent))
|
||||
restored: Final = [part for part in parts if not isinstance(part, _Unappliable)] # mutable-ok: JSON array
|
||||
return restored if len(restored) == len(parts) else _UNAPPLIABLE
|
||||
|
||||
|
||||
def _restored_row(
|
||||
returned: AllMessageValues,
|
||||
caller: AllMessageValues,
|
||||
sent: AllMessageValues,
|
||||
*,
|
||||
altered: bool,
|
||||
) -> AllMessageValues | JsonValue | _Unappliable:
|
||||
if returned is caller:
|
||||
return caller
|
||||
returned_json: Final = as_json_value(returned)
|
||||
caller_content: Final = _content(as_json_value(caller))
|
||||
if not altered:
|
||||
adds_placeholder: Final = _omitted_image_count(_content(returned_json)) > _omitted_image_count(caller_content)
|
||||
return _UNAPPLIABLE if adds_placeholder else returned
|
||||
content: Final = _restored_content(_content(returned_json), caller_content, _content(as_json_value(sent)))
|
||||
if isinstance(content, _Unappliable) or not isinstance(returned_json, dict):
|
||||
return _UNAPPLIABLE
|
||||
return {**returned_json, "content": content} # mutable-ok: JSON row
|
||||
|
||||
|
||||
def restore_unseen_rows(
|
||||
*,
|
||||
rows: tuple[AllMessageValues, ...] | None,
|
||||
caller: Sequence[AllMessageValues] | None,
|
||||
sent: Sequence[AllMessageValues] | None,
|
||||
loss: PayloadLoss,
|
||||
guardrail_name: str | None,
|
||||
) -> tuple[AllMessageValues, ...] | None:
|
||||
"""At a row the guardrail saw only in part, accept its rewrite only where the parts line up, and put the
|
||||
caller's image back into each part that still holds the placeholder. Anything else blocks the request,
|
||||
since keeping the caller's row would silently drop the guardrail's rewrite."""
|
||||
altered: Final = loss.altered_message_indices
|
||||
if rows is None or not altered:
|
||||
return rows
|
||||
if caller is None or sent is None or not len(rows) == len(caller) == len(sent):
|
||||
raise unappliable_request_rewrite(guardrail_name)
|
||||
restored: Final = tuple(
|
||||
_restored_row(row, caller_row, sent_row, altered=index in altered)
|
||||
for index, (row, caller_row, sent_row) in enumerate(zip(rows, caller, sent))
|
||||
)
|
||||
accepted: Final = structured_messages_from_json(
|
||||
[row for row in restored if not isinstance(row, _Unappliable)] # mutable-ok: JSON array
|
||||
)
|
||||
if accepted is None or len(accepted) != len(restored):
|
||||
raise unappliable_request_rewrite(guardrail_name)
|
||||
return tuple(accepted)
|
||||
|
|
@ -103,6 +103,30 @@ class GenericGuardrailAPIOptionalParams(BaseModel):
|
|||
),
|
||||
)
|
||||
|
||||
send_images: bool | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"If False, the top-level images field is not sent, and every inline image_url part in "
|
||||
"structured_messages keeps its place but has its URL replaced by '[omitted]'. File, "
|
||||
"audio and video parts are still sent as they are. Saves payload size for guardrails "
|
||||
"that only inspect text. The guardrail cannot replace the caller's images: a rewritten "
|
||||
"message gets the caller's image back in each part still holding '[omitted]', and a "
|
||||
"rewrite that changes or moves an image part is rejected. Defaults to True in "
|
||||
"GenericGuardrailAPI.__init__ when None."
|
||||
),
|
||||
)
|
||||
|
||||
exclude_payload_fields: tuple[str, ...] | None = Field(
|
||||
default=None,
|
||||
description=(
|
||||
"Top-level guardrail request fields to leave out of the payload, e.g. "
|
||||
"['request_headers', 'tools'], for guardrails that do not use them. Unknown fields "
|
||||
"are ignored with a warning at init, and input_type and litellm_call_id are always "
|
||||
"sent. A field that is not sent (texts, structured_messages, images or tools) cannot "
|
||||
"be rewritten by the guardrail response."
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class GenericGuardrailAPIConfigModel(
|
||||
GuardrailConfigModel[GenericGuardrailAPIOptionalParams],
|
||||
|
|
@ -159,7 +183,7 @@ def coerce_stream_holdback_value(value: Any) -> int:
|
|||
return 0
|
||||
|
||||
|
||||
def structured_messages_from_response(value: object) -> Sequence[AllMessageValues] | None:
|
||||
def structured_messages_from_json(value: object) -> Sequence[AllMessageValues] | None:
|
||||
if not isinstance(value, list):
|
||||
return None
|
||||
if not all(isinstance(message, Mapping) and isinstance(message.get("role"), str) for message in value):
|
||||
|
|
@ -212,5 +236,5 @@ class GenericGuardrailAPIResponse:
|
|||
images=data.get("images"),
|
||||
tools=data.get("tools"),
|
||||
stream_holdback_chars=stream_holdback_chars,
|
||||
structured_messages=structured_messages_from_response(data.get("structured_messages")),
|
||||
structured_messages=structured_messages_from_json(data.get("structured_messages")),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
Loading…
Add table
Reference in a new issue