fix(guardrails): preserve per-choice tool inspection context

This commit is contained in:
albertbausili 2026-09-28 13:07:40 +02:00
parent 4a709fa924
commit b2d14f9174
3 changed files with 63 additions and 11 deletions

View file

@ -657,13 +657,22 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
model_response: Final = self._rebuild_ended_stream_per_choice(responses_so_far, litellm_logging_obj)
pre_guardrail_texts: Final = self._string_choice_contents(model_response)
pre_guardrail_tool_calls: Final = self._function_tool_call_shapes(model_response)
await self.process_output_response(
response=model_response,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=litellm_logging_obj,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
inspection_responses: Final = (
tuple(
model_response.model_copy(update=MappingProxyType({"choices": [choice]}))
for choice in model_response.choices
)
if pre_guardrail_tool_calls and len(model_response.choices) > 1
else (model_response,)
)
for inspection_response in inspection_responses:
await self.process_output_response(
response=inspection_response,
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=litellm_logging_obj,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
)
if not deliver_ended_stream_rewrites:
return
await self._write_ended_stream_text_rewrites(
@ -798,6 +807,24 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
inputs: Final = GenericGuardrailAPIInputs(texts=texts_to_check)
if self._streamed_tool_call_fingerprints(responses_so_far):
assembled: Final = self._rebuild_ended_stream_per_choice(responses_so_far, litellm_logging_obj)
if len(assembled.choices) > 1:
choice_sinks: Final = tuple((choice.index, StreamTransformSink()) for choice in assembled.choices)
for index, choice_sink in choice_sinks:
await self._process_streaming_transform(
responses_so_far=[self._narrowed_to_choice(chunk, index) for chunk in responses_so_far],
guardrail_to_apply=guardrail_to_apply,
litellm_logging_obj=litellm_logging_obj,
user_api_key_dict=user_api_key_dict,
request_data=request_data,
sink=choice_sink,
)
sink.mutated_text_per_choice = dict(
chain.from_iterable(choice_sink.mutated_text_per_choice.items() for _, choice_sink in choice_sinks)
)
sink.holdback_per_choice = dict(
chain.from_iterable(choice_sink.holdback_per_choice.items() for _, choice_sink in choice_sinks)
)
return
tool_calls: Final = chain.from_iterable(choice.message.tool_calls or () for choice in assembled.choices)
inputs["tool_calls"] = TypeAdapter(list[ChatCompletionToolCallChunk]).validate_python(
tuple(

View file

@ -962,9 +962,11 @@ class _ToolRedactingGuardrail(CustomGuardrail):
calls: Final = TypeAdapter(tuple[ChatCompletionMessageToolCall, ...]).validate_python(
inputs.get("tool_calls", ())
)
input_texts: Final = inputs.get("texts", ())
texts: Final = tuple(
"checked:" + text.replace("SECRET", "MASKED") if calls or not self.require_tool_context else text
for text in inputs.get("texts", ())
"checked:" + text.replace("SECRET", "MASKED")
if not self.require_tool_context or (calls and index == len(input_texts) - 1) else text
for index, text in enumerate(input_texts)
)
return {
**inputs,
@ -1021,6 +1023,29 @@ class TestStreamingTransform:
assert out[-1].usage.total_tokens == 18
assert all("SECRET" not in chunk.model_dump_json() for chunk in out)
@pytest.mark.asyncio
@pytest.mark.parametrize("tool_choice_index", [0, 1])
async def test_tool_dependent_text_rewrites_keep_choice_context(self, tool_choice_index: int) -> None:
chunks: Final = (
ModelResponseStream(choices=[StreamingChoices(
index=index, delta=Delta(content="SECRET" if index == tool_choice_index else "plain"),
) for index in (1, 0)]),
ModelResponseStream(choices=[StreamingChoices(
index=tool_choice_index,
delta=Delta(tool_calls=[{
"index": 0, "id": "call_context", "type": "function",
"function": {"name": "contact", "arguments": '{"contact":"SECRET"}'},
}]), finish_reason="tool_calls",
)]),
)
out: Final = await _drive_stream(UnifiedLLMGuardrails(), _ToolRedactingGuardrail(True), chunks)
for index in (0, 1):
text: Final = "".join(
choice.delta.content or "" for chunk in out for choice in chunk.choices if choice.index == index
)
assert text == ("checked:MASKED" if index == tool_choice_index else "plain")
assert all("SECRET" not in chunk.model_dump_json() for chunk in out)
@pytest.mark.asyncio
async def test_tool_rewrites_keep_completion_choices_separate(self) -> None:
async def response() -> AsyncIterator[ModelResponseStream]:

View file

@ -1395,9 +1395,9 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput:
)
assert [
(tool_call["id"], tool_call["function"]["arguments"])
for tool_call in guardrail.seen_inputs[-1]["tool_calls"]
] == [("call_1", '{"fruit": "persimmon"}'), ("call_2", '{"fruit": "durian"}')]
[(tool_call["id"], tool_call["function"]["arguments"]) for tool_call in inputs["tool_calls"]]
for inputs in guardrail.seen_inputs
] == [[("call_1", '{"fruit": "persimmon"}')], [("call_2", '{"fruit": "durian"}')]]
@pytest.mark.asyncio
async def test_deliver_ended_stream_tool_rewrites_keep_choice_indices(self) -> None: