mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(guardrails): hand per-message rewrites back as structured_messages
A guardrail that rewrites text per chat message now returns the rewritten rows as structured_messages instead of only texts, so the Responses and chat handlers write the rewrite back through the structured path. The generic guardrail API response accepts an optional structured_messages list, Prompt Security modify builds one from modified_messages, and rows a server echoes back exactly as shown are restored to the original row objects because the request model drops undeclared keys. Texts-only per-message answers keep the named rejection on both endpoints.
This commit is contained in:
parent
a635d7be6a
commit
09314f239c
8 changed files with 244 additions and 145 deletions
|
|
@ -2,7 +2,6 @@ from __future__ import annotations
|
|||
|
||||
import json
|
||||
from collections.abc import Callable, Iterator, Mapping, Sequence
|
||||
from itertools import accumulate
|
||||
from types import MappingProxyType
|
||||
from typing import Final, TypeVar, cast # noqa: TID251 # a rebuilt chat row has no typed constructor across roles
|
||||
|
||||
|
|
@ -390,7 +389,7 @@ def _part_with_text(part: object, text: str) -> object:
|
|||
return {**part, "text": text} # mutable-ok: content parts stay JSON-plain dicts
|
||||
|
||||
|
||||
def _content_with_slot_texts(content: Sequence[object], texts: Sequence[str]) -> list[object]:
|
||||
def _content_with_slot_texts(content: Sequence[object], texts: Sequence[str]) -> Sequence[object]:
|
||||
text_part_indices: Final = tuple(
|
||||
index for index, part in enumerate(content) if _content_part_text(part) is not None
|
||||
)
|
||||
|
|
@ -401,39 +400,18 @@ def _content_with_slot_texts(content: Sequence[object], texts: Sequence[str]) ->
|
|||
]
|
||||
|
||||
|
||||
def _message_with_slot_texts(message: AllMessageValues, texts: Sequence[str]) -> AllMessageValues:
|
||||
def message_with_slot_texts(message: AllMessageValues, texts: Sequence[str]) -> AllMessageValues | None:
|
||||
"""Swap one rewritten text into each text slot of a chat row, in order.
|
||||
|
||||
A slot is a string ``content`` or one list part carrying a string ``text``;
|
||||
images and other parts ride along untouched. Returns None unless the counts
|
||||
line up exactly, so a rewrite never lands on the wrong slot.
|
||||
"""
|
||||
if message_text_slot_count(message) != len(texts):
|
||||
return None
|
||||
content: Final = message.get("content")
|
||||
if not isinstance(content, (str, list)):
|
||||
return message
|
||||
rewritten_content: Final = texts[0] if isinstance(content, str) else _content_with_slot_texts(content, texts)
|
||||
rewritten: Final = {**message, "content": rewritten_content} # mutable-ok: chat rows stay JSON-plain dicts
|
||||
return cast("AllMessageValues", rewritten) # cast-ok: the same row with only its text slots swapped
|
||||
|
||||
|
||||
def message_with_slot_texts(message: AllMessageValues, texts: Sequence[str]) -> AllMessageValues | None:
|
||||
if message_text_slot_count(message) != len(texts):
|
||||
return None
|
||||
return _message_with_slot_texts(message, texts)
|
||||
|
||||
|
||||
def messages_with_slot_texts(
|
||||
messages: Sequence[AllMessageValues],
|
||||
texts: Sequence[str],
|
||||
) -> list[AllMessageValues] | None:
|
||||
"""Spread one flat list of rewritten texts over the messages' text slots, in order.
|
||||
|
||||
A slot is a string ``content`` or one list part carrying a string ``text``;
|
||||
images and other parts ride along untouched. A guardrail that answers one
|
||||
text per message it saw produces exactly this shape, which stops matching
|
||||
the endpoint's own per-text extraction as soon as the request carries
|
||||
instructions or tool items. Returns None unless the counts line up exactly,
|
||||
so a rewrite never lands on the wrong slot.
|
||||
"""
|
||||
slot_counts: Final = tuple(message_text_slot_count(message) for message in messages)
|
||||
if sum(slot_counts) != len(texts):
|
||||
return None
|
||||
offsets: Final = tuple(accumulate(slot_counts, initial=0))
|
||||
return [ # mutable-ok: guardrail rows travel as a list
|
||||
_message_with_slot_texts(message, texts[start:end])
|
||||
for message, start, end in zip(messages, offsets, offsets[1:])
|
||||
]
|
||||
|
|
|
|||
|
|
@ -53,7 +53,6 @@ from litellm.llms.base_llm.guardrail_translation.base_translation import (
|
|||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.utils import (
|
||||
blocked_responses_stream_usage,
|
||||
messages_with_slot_texts,
|
||||
stream_item_field,
|
||||
stream_item_fingerprint,
|
||||
stream_item_items,
|
||||
|
|
@ -396,20 +395,6 @@ def _patched_request_fields(
|
|||
)
|
||||
|
||||
|
||||
def _guardrailed_structured_messages(
|
||||
structured_messages: Sequence[AllMessageValues] | None,
|
||||
sent_text_count: int,
|
||||
guardrailed_inputs: GenericGuardrailAPIInputs,
|
||||
) -> Sequence[AllMessageValues] | None:
|
||||
returned: Final = guardrailed_inputs.get("structured_messages")
|
||||
if returned is not None and returned is not structured_messages:
|
||||
return returned
|
||||
rewritten_texts: Final = guardrailed_inputs.get("texts")
|
||||
if not structured_messages or rewritten_texts is None or len(rewritten_texts) == sent_text_count:
|
||||
return None
|
||||
return messages_with_slot_texts(structured_messages, rewritten_texts)
|
||||
|
||||
|
||||
def _patch_or_convert_request_fields(
|
||||
raw_input: object,
|
||||
instructions: object,
|
||||
|
|
@ -488,8 +473,7 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
form.chat_tools for form in LiteLLMCompletionResponsesConfig.responses_tools_to_chat_forms(original_tools)
|
||||
)
|
||||
extracted: Final = self._extract_guardrail_inputs(data, input_data, flattened_tool_groups)
|
||||
sent_texts: Final = extracted.inputs.get("texts")
|
||||
if not sent_texts:
|
||||
if not extracted.inputs.get("texts"):
|
||||
return data
|
||||
if structured_messages:
|
||||
extracted.inputs["structured_messages"] = structured_messages
|
||||
|
|
@ -502,9 +486,7 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
self._apply_guardrailed_tools_to_data(
|
||||
data, original_tools, flattened_tool_groups, guardrailed_inputs.get("tools")
|
||||
)
|
||||
written_back: Final = self._written_back_request_fields(
|
||||
data, structured_messages, len(sent_texts), guardrailed_inputs
|
||||
)
|
||||
written_back: Final = self._written_back_request_fields(data, structured_messages, guardrailed_inputs)
|
||||
if written_back is not None:
|
||||
data["input"] = list(written_back.input) # mutable-ok: JSON body
|
||||
if written_back.instructions is None:
|
||||
|
|
@ -571,11 +553,10 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
def _written_back_request_fields(
|
||||
data: Mapping[str, object],
|
||||
structured_messages: Sequence[AllMessageValues] | None,
|
||||
sent_text_count: int,
|
||||
guardrailed_inputs: GenericGuardrailAPIInputs,
|
||||
) -> _RequestFields | None:
|
||||
guardrailed: Final = _guardrailed_structured_messages(structured_messages, sent_text_count, guardrailed_inputs)
|
||||
if guardrailed is None:
|
||||
guardrailed: Final = guardrailed_inputs.get("structured_messages")
|
||||
if guardrailed is None or guardrailed is structured_messages:
|
||||
return None
|
||||
return _patch_or_convert_request_fields(
|
||||
data.get("input"),
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@
|
|||
|
||||
import fnmatch
|
||||
import os
|
||||
from collections.abc import Mapping
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import TYPE_CHECKING, Any, Final, Literal, Optional
|
||||
|
||||
import httpx
|
||||
|
|
@ -24,7 +24,7 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.types.guardrails import GuardrailEventHooks
|
||||
from litellm.types.llms.openai import ChatCompletionToolParam
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolParam
|
||||
from litellm.types.proxy.guardrails.guardrail_hooks.generic_guardrail_api import (
|
||||
GenericGuardrailAPIMetadata,
|
||||
GenericGuardrailAPIRequest,
|
||||
|
|
@ -150,6 +150,22 @@ def _extract_inbound_headers(
|
|||
return None
|
||||
|
||||
|
||||
def _rows_with_unchanged_originals(
|
||||
original_rows: Sequence[AllMessageValues] | None,
|
||||
shown_rows: Sequence[AllMessageValues] | None,
|
||||
returned_rows: Sequence[AllMessageValues],
|
||||
) -> tuple[AllMessageValues, ...]:
|
||||
"""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;
|
||||
only rows the server actually changed reach the endpoint write-back."""
|
||||
if original_rows is None or shown_rows is None or len(returned_rows) != len(original_rows):
|
||||
return tuple(returned_rows)
|
||||
return tuple(
|
||||
original if returned == shown else returned
|
||||
for original, shown, returned in zip(original_rows, shown_rows, returned_rows)
|
||||
)
|
||||
|
||||
|
||||
class GenericGuardrailAPI(CustomGuardrail):
|
||||
"""
|
||||
Generic Guardrail API integration for LiteLLM.
|
||||
|
|
@ -322,6 +338,8 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
texts: list,
|
||||
images: list[str] | None,
|
||||
tools: list[ChatCompletionToolParam] | None,
|
||||
structured_messages: Sequence[AllMessageValues] | None,
|
||||
shown_messages: Sequence[AllMessageValues] | None,
|
||||
guardrail_response: GenericGuardrailAPIResponse,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
# Action is NONE or no modifications needed
|
||||
|
|
@ -336,6 +354,12 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
return_inputs["tools"] = guardrail_response.tools
|
||||
elif tools:
|
||||
return_inputs["tools"] = tools
|
||||
if guardrail_response.structured_messages:
|
||||
return_inputs["structured_messages"] = list( # mutable-ok: guardrail inputs take a list
|
||||
_rows_with_unchanged_originals(
|
||||
structured_messages, shown_messages, guardrail_response.structured_messages
|
||||
)
|
||||
)
|
||||
if guardrail_response.stream_holdback_chars is not None:
|
||||
return_inputs["stream_holdback_chars"] = guardrail_response.stream_holdback_chars
|
||||
return return_inputs
|
||||
|
|
@ -473,6 +497,8 @@ class GenericGuardrailAPI(CustomGuardrail):
|
|||
texts=texts,
|
||||
images=images,
|
||||
tools=tools,
|
||||
structured_messages=structured_messages,
|
||||
shown_messages=guardrail_request.structured_messages,
|
||||
guardrail_response=guardrail_response,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -287,7 +287,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
structured_messages, modified_messages
|
||||
)
|
||||
if rewritten_messages is not None:
|
||||
inputs["structured_messages"] = rewritten_messages
|
||||
inputs["structured_messages"] = list(rewritten_messages) # mutable-ok: guardrail inputs take a list
|
||||
|
||||
return inputs
|
||||
|
||||
|
|
@ -298,7 +298,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
self,
|
||||
structured_messages: Sequence[AllMessageValues],
|
||||
modified_messages: Sequence[Mapping[str, object]],
|
||||
) -> list[AllMessageValues] | None:
|
||||
) -> tuple[AllMessageValues, ...] | None:
|
||||
sent_indices: Final = tuple(
|
||||
index for index, message in enumerate(structured_messages) if self._is_sent_to_protect(message)
|
||||
)
|
||||
|
|
@ -313,9 +313,7 @@ class PromptSecurityGuardrail(CustomGuardrail):
|
|||
)
|
||||
if len(replacements) != len(sent_indices):
|
||||
return None
|
||||
return [ # mutable-ok: guardrail inputs take a list
|
||||
replacements.get(index, message) for index, message in enumerate(structured_messages)
|
||||
]
|
||||
return tuple(replacements.get(index, message) for index, message in enumerate(structured_messages))
|
||||
|
||||
async def _apply_guardrail_on_response(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
from typing import Any, Final, Literal
|
||||
from collections.abc import Mapping, Sequence
|
||||
from typing import Any, Final, Literal, cast # noqa: TID251 # JSON chat rows have no typed constructor across roles
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from typing_extensions import TypedDict
|
||||
|
|
@ -158,12 +159,21 @@ def coerce_stream_holdback_value(value: Any) -> int:
|
|||
return 0
|
||||
|
||||
|
||||
def structured_messages_from_response(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):
|
||||
return None
|
||||
return cast("Sequence[AllMessageValues]", value) # cast-ok: JSON rows checked for a role, the same trust texts get
|
||||
|
||||
|
||||
class GenericGuardrailAPIResponse:
|
||||
"""Response model for the Generic Guardrail API"""
|
||||
|
||||
texts: list[str] | None
|
||||
images: list[str] | None
|
||||
tools: list[GuardrailToolParam] | None
|
||||
structured_messages: Sequence[AllMessageValues] | None
|
||||
action: str
|
||||
blocked_reason: str | None
|
||||
stream_holdback_chars: list[int] | None
|
||||
|
|
@ -176,12 +186,14 @@ class GenericGuardrailAPIResponse:
|
|||
images: list[str] | None = None,
|
||||
tools: list[GuardrailToolParam] | None = None,
|
||||
stream_holdback_chars: list[int] | None = None,
|
||||
structured_messages: Sequence[AllMessageValues] | None = None,
|
||||
) -> None:
|
||||
self.action = action
|
||||
self.blocked_reason = blocked_reason
|
||||
self.texts = texts
|
||||
self.images = images
|
||||
self.tools = tools
|
||||
self.structured_messages = structured_messages
|
||||
# Number of trailing chars, indexed the same as ``texts``, that the
|
||||
# framework must withhold from streaming emission until the next
|
||||
# processing round (word-boundary safety for text transformations).
|
||||
|
|
@ -200,4 +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")),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@ with guardrail transformations.
|
|||
import copy
|
||||
from collections.abc import Callable
|
||||
from typing import Any, List, Literal, Optional, Tuple
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import logging
|
||||
|
||||
|
|
@ -31,6 +31,7 @@ from litellm.llms.openai.responses.guardrail_translation.handler import (
|
|||
OpenAIResponsesHandler,
|
||||
)
|
||||
from litellm.llms.openai.responses.guardrail_translation.tool_merge import merge_guardrailed_tools
|
||||
from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import GenericGuardrailAPI
|
||||
from litellm.types.llms.openai import ChatCompletionToolCallChunk
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
|
|
@ -2342,103 +2343,94 @@ SSN = "123-45-6789"
|
|||
REDACTED_SSN = "<US_SSN>"
|
||||
|
||||
|
||||
def _slot_texts(message: dict) -> list[str]:
|
||||
content = message.get("content")
|
||||
if isinstance(content, str):
|
||||
return [content]
|
||||
if isinstance(content, list):
|
||||
return [part["text"] for part in content if isinstance(part, dict) and isinstance(part.get("text"), str)]
|
||||
return []
|
||||
def _redacted(value: object) -> object:
|
||||
if isinstance(value, str):
|
||||
return value.replace(SSN, REDACTED_SSN)
|
||||
if isinstance(value, list):
|
||||
return [{**part, "text": _redacted(part["text"])} if "text" in part else part for part in value]
|
||||
return value
|
||||
|
||||
|
||||
class PerMessageRedactionGuardrail(CustomGuardrail):
|
||||
"""Guardrail that answers one redacted text per message it was shown and hands
|
||||
back only texts, the way Prompt Security in modify mode and a generic guardrail
|
||||
API server that scans per message do."""
|
||||
def _per_message_guardrail_server(structured_messages_in_answer: bool) -> Callable[..., MagicMock]:
|
||||
"""Answers one redacted text per chat row it was shown, the way a guardrail
|
||||
that scans per message does, and optionally the rewritten rows themselves."""
|
||||
|
||||
def __init__(self, extra_texts: int = 0):
|
||||
super().__init__(guardrail_name="per-message-redactor")
|
||||
self.extra_texts = extra_texts
|
||||
def post(url: str, json: dict, headers: dict) -> MagicMock:
|
||||
rows = json["structured_messages"]
|
||||
answer: dict = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"texts": [_redacted(row["content"]) if isinstance(row.get("content"), str) else "" for row in rows],
|
||||
}
|
||||
if structured_messages_in_answer:
|
||||
answer["structured_messages"] = [{**row, "content": _redacted(row.get("content"))} for row in rows]
|
||||
response = MagicMock()
|
||||
response.json.return_value = answer
|
||||
response.raise_for_status = MagicMock()
|
||||
return response
|
||||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional[Any] = None,
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
messages = inputs.get("structured_messages") or []
|
||||
texts = [text.replace(SSN, REDACTED_SSN) for message in messages for text in _slot_texts(message)]
|
||||
return {**inputs, "texts": texts + ["junk"] * self.extra_texts}
|
||||
return post
|
||||
|
||||
|
||||
class TestPerMessageTextWriteBack:
|
||||
"""A guardrail that rewrites one text per message it saw must land on the
|
||||
instructions and the input items those messages came from, not be rejected."""
|
||||
def _per_message_redactor() -> GenericGuardrailAPI:
|
||||
return GenericGuardrailAPI(
|
||||
api_base="https://guardrail.test",
|
||||
guardrail_name="per-message-redactor",
|
||||
event_hook="pre_call",
|
||||
default_on=True,
|
||||
)
|
||||
|
||||
|
||||
def _tool_replay_request() -> dict:
|
||||
return {
|
||||
"model": "gpt-5.6",
|
||||
"instructions": "Never repeat the SSN " + SSN + " back.",
|
||||
"input": [
|
||||
{"role": "user", "content": "Look up " + SSN + " for me."},
|
||||
{"type": "function_call", "call_id": "call_1", "name": "lookup_customer", "arguments": '{"id": "42"}'},
|
||||
{"type": "function_call_output", "call_id": "call_1", "output": '{"ssn": "' + SSN + '"}'},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
class TestPerMessageRewriteWriteBack:
|
||||
"""A guardrail that rewrites per chat row hands the rows back as
|
||||
structured_messages, and the handler lands them on the instructions and the
|
||||
input items they came from; the same rewrite handed back as texts alone has
|
||||
no item to land on and is rejected by name instead of sent unrewritten."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_instructions_plus_tool_replay_gets_each_rewrite_in_place(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
function_call_item = {
|
||||
"type": "function_call",
|
||||
"call_id": "call_1",
|
||||
"name": "lookup_customer",
|
||||
"arguments": '{"query": "' + SSN + '"}',
|
||||
}
|
||||
data = {
|
||||
"model": "gpt-5.6",
|
||||
"instructions": "Never repeat the SSN " + SSN + " back.",
|
||||
"input": [
|
||||
{"role": "user", "content": "Look up " + SSN + " for me."},
|
||||
function_call_item,
|
||||
{"type": "function_call_output", "call_id": "call_1", "output": '{"ssn": "' + SSN + '"}'},
|
||||
],
|
||||
}
|
||||
async def test_structured_rows_land_on_instructions_and_tool_output(self):
|
||||
guardrail = _per_message_redactor()
|
||||
data = _tool_replay_request()
|
||||
function_call_item = data["input"][1]
|
||||
|
||||
result = await handler.process_input_messages(data, PerMessageRedactionGuardrail())
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(True)):
|
||||
result = await OpenAIResponsesHandler().process_input_messages(data, guardrail)
|
||||
|
||||
assert result["instructions"] == "Never repeat the SSN " + REDACTED_SSN + " back."
|
||||
assert [item.get("type", item.get("role")) for item in result["input"]] == [
|
||||
"user",
|
||||
"function_call",
|
||||
"function_call_output",
|
||||
]
|
||||
assert _slot_texts(result["input"][0]) == ["Look up " + REDACTED_SSN + " for me."]
|
||||
assert _texts(result["input"][0]) == ["Look up " + REDACTED_SSN + " for me."]
|
||||
assert result["input"][1] == function_call_item
|
||||
assert result["input"][2]["output"] == '{"ssn": "' + REDACTED_SSN + '"}'
|
||||
assert result["input"][2]["call_id"] == "call_1"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_string_input_with_instructions_keeps_the_two_apart(self):
|
||||
handler = OpenAIResponsesHandler()
|
||||
data = {
|
||||
"model": "gpt-5.6",
|
||||
"instructions": "Redact " + SSN + " everywhere.",
|
||||
"input": "My SSN is " + SSN + ".",
|
||||
assert result["input"][2] == {
|
||||
"type": "function_call_output",
|
||||
"call_id": "call_1",
|
||||
"output": '{"ssn": "' + REDACTED_SSN + '"}',
|
||||
}
|
||||
|
||||
result = await handler.process_input_messages(data, PerMessageRedactionGuardrail())
|
||||
|
||||
assert result["instructions"] == "Redact " + REDACTED_SSN + " everywhere."
|
||||
assert [_slot_texts(item) for item in result["input"]] == [["My SSN is " + REDACTED_SSN + "."]]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_count_matching_neither_texts_nor_messages_is_still_rejected(self):
|
||||
async def test_texts_only_per_message_answer_is_rejected_by_name(self):
|
||||
from litellm.proxy.policy_engine.pipeline_executor import UnappliableRequestRewrite
|
||||
|
||||
handler = OpenAIResponsesHandler()
|
||||
original_input = [
|
||||
{"role": "user", "content": "Look up " + SSN + " for me."},
|
||||
{"type": "function_call_output", "call_id": "call_1", "output": '{"ssn": "' + SSN + '"}'},
|
||||
]
|
||||
data = {"model": "gpt-5.6", "instructions": "Be terse.", "input": copy.deepcopy(original_input)}
|
||||
guardrail = _per_message_redactor()
|
||||
data = _tool_replay_request()
|
||||
original = copy.deepcopy(data)
|
||||
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await handler.process_input_messages(data, PerMessageRedactionGuardrail(extra_texts=1))
|
||||
with patch.object(guardrail.async_handler, "post", side_effect=_per_message_guardrail_server(False)):
|
||||
with pytest.raises(UnappliableRequestRewrite) as excinfo:
|
||||
await OpenAIResponsesHandler().process_input_messages(data, guardrail)
|
||||
|
||||
assert excinfo.value.guardrail_name == "per-message-redactor"
|
||||
assert data["input"] == original_input
|
||||
assert data["instructions"] == "Be terse."
|
||||
assert data["input"] == original["input"]
|
||||
assert data["instructions"] == original["instructions"]
|
||||
|
||||
|
||||
class TestProvenancePatching:
|
||||
|
|
|
|||
|
|
@ -582,6 +582,111 @@ class TestGuardrailActions:
|
|||
assert result_images is None
|
||||
|
||||
|
||||
class TestStructuredMessagesInResponse:
|
||||
"""A guardrail server that rewrites per chat row answers with the rewritten
|
||||
rows as structured_messages, which the endpoint handlers write back by row."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returned_rows_are_handed_back_as_structured_messages(
|
||||
self, generic_guardrail, mock_request_data_input
|
||||
):
|
||||
rewritten_rows = [
|
||||
{"role": "system", "content": "Never repeat an SSN."},
|
||||
{"role": "user", "content": "Look up [REDACTED] for me."},
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "[REDACTED]"}'},
|
||||
]
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"texts": ["Never repeat an SSN.", "Look up [REDACTED] for me.", '{"ssn": "[REDACTED]"}'],
|
||||
"structured_messages": rewritten_rows,
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
with patch.object(generic_guardrail.async_handler, "post", return_value=mock_response):
|
||||
guardrailed_inputs = await generic_guardrail.apply_guardrail(
|
||||
inputs={"texts": ["Look up 123-45-6789 for me."]},
|
||||
request_data=mock_request_data_input,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert guardrailed_inputs["structured_messages"] == rewritten_rows
|
||||
assert guardrailed_inputs["texts"] == mock_response.json.return_value["texts"]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rows_echoed_back_as_shown_keep_their_original_keys(
|
||||
self, generic_guardrail, mock_request_data_input
|
||||
):
|
||||
tool_call_row = {
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [
|
||||
{"id": "call_1", "type": "function", "function": {"name": "lookup", "arguments": "{}"}, "index": 0}
|
||||
],
|
||||
}
|
||||
original_rows = [
|
||||
{"role": "user", "content": "Look up 123-45-6789 for me.", "name": "pat"},
|
||||
tool_call_row,
|
||||
{"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "123-45-6789"}'},
|
||||
]
|
||||
|
||||
def echo_with_tool_output_redacted(url, json, headers):
|
||||
shown_rows = json["structured_messages"]
|
||||
assert "index" not in shown_rows[1]["tool_calls"][0]
|
||||
assert "name" not in shown_rows[0]
|
||||
answer = MagicMock()
|
||||
answer.json.return_value = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"texts": ["Look up 123-45-6789 for me."],
|
||||
"structured_messages": [
|
||||
shown_rows[0],
|
||||
shown_rows[1],
|
||||
{**shown_rows[2], "content": '{"ssn": "[REDACTED]"}'},
|
||||
],
|
||||
}
|
||||
answer.raise_for_status = MagicMock()
|
||||
return answer
|
||||
|
||||
with patch.object(generic_guardrail.async_handler, "post", side_effect=echo_with_tool_output_redacted):
|
||||
guardrailed_inputs = await generic_guardrail.apply_guardrail(
|
||||
inputs={"texts": ["Look up 123-45-6789 for me."], "structured_messages": original_rows},
|
||||
request_data=mock_request_data_input,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
returned_rows = guardrailed_inputs["structured_messages"]
|
||||
assert returned_rows[0] is original_rows[0]
|
||||
assert returned_rows[1] is tool_call_row
|
||||
assert returned_rows[2] == {"role": "tool", "tool_call_id": "call_1", "content": '{"ssn": "[REDACTED]"}'}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"structured_messages",
|
||||
[[], [{"content": "a row with no role"}], "not a list"],
|
||||
ids=["empty", "no_role", "not_a_list"],
|
||||
)
|
||||
async def test_rows_that_are_not_chat_messages_are_ignored(
|
||||
self, generic_guardrail, mock_request_data_input, structured_messages
|
||||
):
|
||||
mock_response = MagicMock()
|
||||
mock_response.json.return_value = {
|
||||
"action": "GUARDRAIL_INTERVENED",
|
||||
"texts": ["[REDACTED]"],
|
||||
"structured_messages": structured_messages,
|
||||
}
|
||||
mock_response.raise_for_status = MagicMock()
|
||||
|
||||
with patch.object(generic_guardrail.async_handler, "post", return_value=mock_response):
|
||||
guardrailed_inputs = await generic_guardrail.apply_guardrail(
|
||||
inputs={"texts": ["Look up 123-45-6789 for me."]},
|
||||
request_data=mock_request_data_input,
|
||||
input_type="request",
|
||||
)
|
||||
|
||||
assert "structured_messages" not in guardrailed_inputs
|
||||
assert guardrailed_inputs["texts"] == ["[REDACTED]"]
|
||||
|
||||
|
||||
class TestImageSupport:
|
||||
"""Test image handling in guardrail requests"""
|
||||
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
import asyncio
|
||||
import base64
|
||||
from collections.abc import Mapping, Sequence
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
|
|
@ -12,6 +13,7 @@ from litellm.proxy.guardrails.guardrail_hooks.prompt_security.prompt_security im
|
|||
PromptSecurityGuardrailMissingSecrets,
|
||||
)
|
||||
from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
||||
def test_prompt_security_guard_config(monkeypatch: pytest.MonkeyPatch):
|
||||
|
|
@ -174,7 +176,7 @@ async def test_apply_guardrail_modify_request(monkeypatch: pytest.MonkeyPatch):
|
|||
assert result["texts"] == ["User prompt with PII: SSN [REDACTED]"]
|
||||
|
||||
|
||||
def _modify_response(modified_messages: list) -> Response:
|
||||
def _modify_response(modified_messages: Sequence[Mapping[str, object]]) -> Response:
|
||||
mock_response = Response(
|
||||
json={"result": {"prompt": {"action": "modify", "modified_messages": modified_messages}}},
|
||||
status_code=200,
|
||||
|
|
@ -184,7 +186,7 @@ def _modify_response(modified_messages: list) -> Response:
|
|||
return mock_response
|
||||
|
||||
|
||||
def _tool_replay_messages() -> list:
|
||||
def _tool_replay_messages() -> list[AllMessageValues]:
|
||||
return [
|
||||
{"role": "system", "content": "Never echo an SSN like 123-45-6789."},
|
||||
{
|
||||
|
|
@ -224,7 +226,9 @@ async def test_modify_returns_structured_messages_with_tool_rows_kept(monkeypatc
|
|||
]
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", return_value=_modify_response(modified_messages)):
|
||||
result = await guardrail.apply_guardrail(inputs=inputs, request_data={"messages": messages}, input_type="request")
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs, request_data={"messages": messages}, input_type="request"
|
||||
)
|
||||
|
||||
assert result["structured_messages"] == [
|
||||
{"role": "system", "content": "Never echo an SSN like [REDACTED]."},
|
||||
|
|
@ -257,7 +261,9 @@ async def test_modify_with_unexpected_message_count_keeps_texts_only(monkeypatch
|
|||
modified_messages = [{"role": "user", "content": "Look up [REDACTED]"}]
|
||||
|
||||
with patch.object(guardrail.async_handler, "post", return_value=_modify_response(modified_messages)):
|
||||
result = await guardrail.apply_guardrail(inputs=inputs, request_data={"messages": messages}, input_type="request")
|
||||
result = await guardrail.apply_guardrail(
|
||||
inputs=inputs, request_data={"messages": messages}, input_type="request"
|
||||
)
|
||||
|
||||
assert result["structured_messages"] is messages
|
||||
assert result["texts"] == ["Look up [REDACTED]"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue