mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(guardrails): read the rewrite from texts when a guardrail echoes every row back unchanged
This commit is contained in:
parent
37447c98f7
commit
2923c4ac55
2 changed files with 49 additions and 10 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue