From b2d14f9174b3f18d0529faaed163aa66ed64702b Mon Sep 17 00:00:00 2001 From: albertbausili Date: Mon, 28 Sep 2026 13:07:40 +0200 Subject: [PATCH] fix(guardrails): preserve per-choice tool inspection context --- .../chat/guardrail_translation/handler.py | 39 ++++++++++++++++--- .../test_unified_guardrail.py | 29 +++++++++++++- .../test_openai_guardrail_handler.py | 6 +-- 3 files changed, 63 insertions(+), 11 deletions(-) diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index 62a370adf71..d8b3a45c72a 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -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( diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index 98269add1ad..52a628412cd 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -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]: diff --git a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py index 29a6e2ef0eb..9cce5d8ff44 100644 --- a/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py +++ b/tests/unit/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py @@ -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: