mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Merge 5b61f8bdef into 65b0557f80
This commit is contained in:
commit
09cc03cceb
5 changed files with 494 additions and 10 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 types import MappingProxyType
|
||||
from typing import Any, Final
|
||||
|
||||
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``
|
||||
|
|
@ -307,3 +311,18 @@ 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 _null_free_object(pairs: Sequence[tuple[str, object]]) -> Mapping[str, object]:
|
||||
return MappingProxyType({key: item for key, item in pairs if item is not None})
|
||||
|
||||
|
||||
def _null_free(value: object) -> object:
|
||||
normalized: Final[object] = json.loads( # pyright: ignore[reportAny] # stdlib parse of our own json.dumps output
|
||||
json.dumps(value, default=to_jsonable_python), object_pairs_hook=_null_free_object
|
||||
)
|
||||
return normalized
|
||||
|
||||
|
||||
def same_json_ignoring_nulls(left: object, right: object) -> bool:
|
||||
return _null_free(left) == _null_free(right)
|
||||
|
|
|
|||
|
|
@ -6,11 +6,15 @@
|
|||
# Thank you users! We ❤️ you! - Krrish & Ishaan
|
||||
|
||||
import fnmatch
|
||||
import json
|
||||
import os
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
|
||||
|
||||
import httpx
|
||||
from pydantic import JsonValue
|
||||
from pydantic_core import to_jsonable_python
|
||||
from typing_extensions import TypeIs
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._version import version as litellm_version
|
||||
|
|
@ -20,9 +24,11 @@ from litellm.integrations.custom_guardrail import (
|
|||
log_guardrail_information,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.proxy.guardrails._content_utils import same_json_ignoring_nulls
|
||||
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 (
|
||||
|
|
@ -30,6 +36,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import
|
|||
GenericGuardrailAPIRequest,
|
||||
GenericGuardrailAPIResponse,
|
||||
GuardrailToolParam,
|
||||
structured_messages_from_json,
|
||||
)
|
||||
from litellm.types.utils import GenericGuardrailAPIInputs
|
||||
|
||||
|
|
@ -150,22 +157,52 @@ def _extract_inbound_headers(
|
|||
return None
|
||||
|
||||
|
||||
def _is_part_list(value: object) -> TypeIs[Sequence[object]]: # guard-ok: trivial isinstance narrowing
|
||||
return isinstance(value, list)
|
||||
|
||||
|
||||
def _as_posted_json(value: object) -> JsonValue:
|
||||
"""httpx encodes the body with the stdlib codec, so the rows an echo is compared with must go through it too"""
|
||||
posted: Final[JsonValue] = json.loads( # pyright: ignore[reportAny] # untyped stdlib parse of json.dumps output
|
||||
json.dumps(value, default=to_jsonable_python)
|
||||
)
|
||||
return posted
|
||||
|
||||
|
||||
def _row_as_sent(dumped: JsonValue, caller: Mapping[str, object]) -> JsonValue:
|
||||
"""The request model dumps a part 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_posted_json(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,
|
||||
returned_rows: Sequence[AllMessageValues],
|
||||
) -> tuple[AllMessageValues, ...] | None:
|
||||
"""The request model drops row keys its message types do not declare, so a
|
||||
row the server echoes back verbatim is restored to the original row object.
|
||||
row the server echoes back, null fields aside, is restored to the original row object.
|
||||
A server that echoes every row back unchanged has not rewritten anything
|
||||
per row, so its answer is read from texts, as it was before rows could be
|
||||
returned at all."""
|
||||
if original_rows is None or shown_rows is None or len(returned_rows) != len(original_rows):
|
||||
return tuple(returned_rows)
|
||||
if all(returned == shown for shown, returned in zip(shown_rows, returned_rows)):
|
||||
if all(same_json_ignoring_nulls(returned, shown) for shown, returned in zip(shown_rows, returned_rows)):
|
||||
return None
|
||||
return tuple(
|
||||
original if returned == shown else returned
|
||||
original if same_json_ignoring_nulls(returned, shown) else returned
|
||||
for original, shown, returned in zip(original_rows, shown_rows, returned_rows)
|
||||
)
|
||||
|
||||
|
|
@ -204,9 +241,12 @@ 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,
|
||||
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 []
|
||||
|
||||
|
|
@ -470,12 +510,14 @@ 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}
|
||||
|
||||
# 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=request_json,
|
||||
headers=headers,
|
||||
timeout=self.timeout,
|
||||
)
|
||||
|
|
@ -504,7 +546,7 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
images=images,
|
||||
tools=tools,
|
||||
structured_messages=structured_messages,
|
||||
shown_messages=guardrail_request.structured_messages,
|
||||
shown_messages=structured_messages_from_json(sent_messages),
|
||||
guardrail_response=guardrail_response,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -159,7 +159,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 +212,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,24 @@ 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, TypeAlias
|
||||
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.llms.openai.responses.guardrail_translation.handler import OpenAIResponsesHandler
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
|
||||
GenericGuardrailAPI,
|
||||
|
|
@ -22,6 +30,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 +731,386 @@ 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: TypeAlias = 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 _without_null_values(row: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]:
|
||||
return {key: value for key, value in row.items() if value is not None}
|
||||
|
||||
|
||||
def _echo_every_row_without_its_nulls_and_mask_texts(request_json: Mapping[str, JsonValue]) -> Mapping[str, JsonValue]:
|
||||
rows_without_nulls: Final[JsonValue] = json.loads( # pyright: ignore[reportAny] # stdlib parse of a JSON copy
|
||||
json.dumps(request_json["structured_messages"]), object_hook=_without_null_values
|
||||
)
|
||||
return {**_echo_every_row_and_mask_texts(request_json), "structured_messages": rows_without_nulls}
|
||||
|
||||
|
||||
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 _echo_first_row_without_its_nulls_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": [
|
||||
_without_null_values(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"]
|
||||
|
||||
|
||||
async def _llm_bound_responses_input(guardrail: GenericGuardrailAPI, input_items: list[JsonValue]) -> object:
|
||||
data: Final = await OpenAIResponsesHandler().process_input_messages(
|
||||
data={"model": "gpt-5.6", "input": input_items}, guardrail_to_apply=guardrail
|
||||
)
|
||||
return data["input"]
|
||||
|
||||
|
||||
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_echo_that_drops_null_fields_still_applies_the_masked_texts(self) -> None:
|
||||
guardrail: Final = _guardrail_answering(_echo_every_row_without_its_nulls_and_mask_texts)
|
||||
tool_use_turn: Final[AllAnthropicMessageValues] = {
|
||||
"role": "assistant",
|
||||
"content": [
|
||||
{"type": "text", "text": "calling"},
|
||||
{"type": "tool_use", "id": "t1", "name": "f", "input": {}},
|
||||
],
|
||||
}
|
||||
tool_result_turn: Final[AllAnthropicMessageValues] = {
|
||||
"role": "user",
|
||||
"content": [{"type": "tool_result", "tool_use_id": "t1", "content": "r"}],
|
||||
}
|
||||
|
||||
llm_bound: Final = await _llm_bound_anthropic_messages(
|
||||
guardrail, [{"role": "user", "content": f"my ssn is {_SSN}"}, tool_use_turn, tool_result_turn]
|
||||
)
|
||||
|
||||
assert llm_bound == [
|
||||
{"role": "user", "content": "my ssn is [SSN]"},
|
||||
tool_use_turn,
|
||||
tool_result_turn,
|
||||
], "an echo without the null fields the rows were posted with must leave the rewrite to texts"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_every_responses_input_text_row_echoed_applies_the_masked_texts(self) -> None:
|
||||
guardrail: Final = _guardrail_answering(_echo_every_row_and_mask_texts)
|
||||
|
||||
llm_bound: Final = await _llm_bound_responses_input(
|
||||
guardrail, [{"role": "user", "content": [{"type": "input_text", "text": f"my ssn is {_SSN}"}]}]
|
||||
)
|
||||
|
||||
assert llm_bound == [{"role": "user", "content": [{"type": "input_text", "text": "my ssn is [SSN]"}]}], (
|
||||
"an unchanged echo of every input_text 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
|
||||
async def test_a_row_echoed_without_its_null_fields_is_restored_to_the_callers_row(self) -> None:
|
||||
guardrail: Final = _guardrail_answering(_echo_first_row_without_its_nulls_and_mask_the_rest)
|
||||
tool_call_turn: Final[AllMessageValues] = {
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [{"id": "c1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}}],
|
||||
}
|
||||
|
||||
llm_bound: Final = await _llm_bound_messages(
|
||||
guardrail, [tool_call_turn, {"role": "tool", "tool_call_id": "c1", "content": f"ssn {_SSN}"}]
|
||||
)
|
||||
|
||||
assert llm_bound == [
|
||||
tool_call_turn,
|
||||
{"role": "tool", "tool_call_id": "c1", "content": "ssn [SSN]"},
|
||||
], "a row echoed without the null fields it was posted with must stay the caller's row"
|
||||
|
||||
@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_unchanged_echo_of_a_rejected_part_holding_a_tuple_applies_the_masked_texts(self) -> None:
|
||||
guardrail: Final = _guardrail_answering(_echo_every_row_and_mask_texts)
|
||||
tagged_part: Final = {"type": "guarded_text", "text": "keep this guarded", "tags": ("a", "b")}
|
||||
|
||||
llm_bound: Final = await _llm_bound_messages(
|
||||
guardrail, [{"role": "user", "content": [{"type": "text", "text": f"my ssn is {_SSN}"}, tagged_part]}]
|
||||
)
|
||||
|
||||
assert llm_bound == [
|
||||
{"role": "user", "content": [{"type": "text", "text": "my ssn is [SSN]"}, tagged_part]}
|
||||
], "the rows an echo is compared with must equal the JSON the guardrail received"
|
||||
|
||||
@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,5 +1,7 @@
|
|||
"""Tests for the shared guardrail content extraction helpers."""
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.guardrails._content_utils import (
|
||||
apply_redacted_messages_back,
|
||||
build_inspection_messages,
|
||||
|
|
@ -7,6 +9,7 @@ from litellm.proxy.guardrails._content_utils import (
|
|||
is_non_conversational_call_type,
|
||||
is_string_batch_input,
|
||||
iter_message_text,
|
||||
same_json_ignoring_nulls,
|
||||
walk_user_text,
|
||||
)
|
||||
|
||||
|
|
@ -741,3 +744,33 @@ 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
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("left", "right", "same"),
|
||||
[
|
||||
({"role": "assistant", "thinking_blocks": None}, {"role": "assistant"}, True),
|
||||
(
|
||||
{"content": [{"type": "text", "text": "x", "cache_control": None}]},
|
||||
{"content": [{"type": "text", "text": "x"}]},
|
||||
True,
|
||||
),
|
||||
({"content": ("a", "b")}, {"content": ["a", "b"]}, True),
|
||||
({"content": "x"}, {"content": "y"}, False),
|
||||
({"role": "user", "name": "a"}, {"role": "user"}, False),
|
||||
([None], [], False),
|
||||
],
|
||||
ids=[
|
||||
"dropped_null_key",
|
||||
"dropped_nested_null_key",
|
||||
"tuple_as_list",
|
||||
"changed_value",
|
||||
"dropped_set_key",
|
||||
"null_list_item",
|
||||
],
|
||||
)
|
||||
def test_same_json_ignoring_nulls_treats_only_a_dropped_null_field_as_no_change(
|
||||
left: object, right: object, same: bool
|
||||
) -> None:
|
||||
assert same_json_ignoring_nulls(left, right) is same
|
||||
assert same_json_ignoring_nulls(right, left) is same
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue