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:
mateo-berri 2026-09-13 00:54:29 -07:00
parent a635d7be6a
commit 09314f239c
8 changed files with 244 additions and 145 deletions

View file

@ -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:])
]

View file

@ -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"),

View file

@ -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,
)

View file

@ -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,

View file

@ -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")),
)

View file

@ -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:

View file

@ -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"""

View file

@ -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]"]