fix(guardrails): read the rewrite from texts when a guardrail echoes every row back unchanged

This commit is contained in:
mateo-berri 2026-09-13 03:02:15 -07:00
parent 37447c98f7
commit 2923c4ac55
2 changed files with 49 additions and 10 deletions

View file

@ -150,16 +150,20 @@ def _extract_inbound_headers(
return None
def _rows_with_unchanged_originals(
def _structured_rows_to_write_back(
original_rows: Sequence[AllMessageValues] | None,
shown_rows: Sequence[AllMessageValues] | None,
returned_rows: Sequence[AllMessageValues],
) -> tuple[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;
only rows the server actually changed reach the endpoint write-back."""
row the server echoes back verbatim 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)):
return None
return tuple(
original if returned == shown else returned
for original, shown, returned in zip(original_rows, shown_rows, returned_rows)
@ -354,12 +358,13 @@ 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
)
)
rows_to_write_back: Final = (
_structured_rows_to_write_back(structured_messages, shown_messages, guardrail_response.structured_messages)
if guardrail_response.structured_messages
else None
)
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:
return_inputs["stream_holdback_chars"] = guardrail_response.stream_holdback_chars
return return_inputs

View file

@ -659,6 +659,40 @@ class TestStructuredMessagesInResponse:
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
async def test_rows_all_echoed_back_as_shown_leave_the_rewrite_to_texts(
self, generic_guardrail, mock_request_data_input
):
"""A server written against the texts contract that echoes the request rows back
untouched while rewriting texts still gets its texts rewrite applied."""
original_rows = [
{"role": "system", "content": "Never repeat an SSN."},
{"role": "user", "content": "Look up 123-45-6789 for me."},
]
def echo_rows_and_rewrite_texts(url, json, headers):
answer = MagicMock()
answer.json.return_value = {
"action": "NONE",
"texts": [text.replace("123-45-6789", "[REDACTED]") for text in json["texts"]],
"structured_messages": json["structured_messages"],
}
answer.raise_for_status = MagicMock()
return answer
with patch.object(generic_guardrail.async_handler, "post", side_effect=echo_rows_and_rewrite_texts):
guardrailed_inputs = await generic_guardrail.apply_guardrail(
inputs={
"texts": ["Never repeat an SSN.", "Look up 123-45-6789 for me."],
"structured_messages": original_rows,
},
request_data=mock_request_data_input,
input_type="request",
)
assert "structured_messages" not in guardrailed_inputs
assert guardrailed_inputs["texts"] == ["Never repeat an SSN.", "Look up [REDACTED] for me."]
@pytest.mark.asyncio
@pytest.mark.parametrize(
"structured_messages",