diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index d8b3a45c72a..661cffa9846 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -665,14 +665,28 @@ class OpenAIChatCompletionsHandler(BaseTranslation): 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, - ) + inspection_request_data: Final = request_data if request_data is not None else {} + try: + for inspection_response in inspection_responses: + inspection_request_data["response"] = inspection_response + inspection_request_data["responses"] = ( + [ + self._narrowed_to_choice(chunk, inspection_response.choices[0].index) + for chunk in responses_so_far + ] + if len(inspection_response.choices) == 1 + else responses_so_far + ) + 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=inspection_request_data, + ) + finally: + inspection_request_data["response"] = model_response + inspection_request_data["responses"] = responses_so_far if not deliver_ended_stream_rewrites: return await self._write_ended_stream_text_rewrites( @@ -807,22 +821,38 @@ 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) + request_data["response"] = assembled 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, + choice_rounds: Final = tuple( + ( + choice, + StreamTransformSink(), + [self._narrowed_to_choice(chunk, choice.index) for chunk in responses_so_far], ) + for choice in assembled.choices + ) + try: + for choice, choice_sink, choice_chunks in choice_rounds: + request_data["response"] = assembled.model_copy(update=MappingProxyType({"choices": [choice]})) + request_data["responses"] = choice_chunks + await self._process_streaming_transform( + responses_so_far=choice_chunks, + 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, + ) + finally: + request_data["response"] = assembled + request_data["responses"] = responses_so_far sink.mutated_text_per_choice = dict( - chain.from_iterable(choice_sink.mutated_text_per_choice.items() for _, choice_sink in choice_sinks) + chain.from_iterable( + choice_sink.mutated_text_per_choice.items() for _, choice_sink, _ in choice_rounds + ) ) sink.holdback_per_choice = dict( - chain.from_iterable(choice_sink.holdback_per_choice.items() for _, choice_sink in choice_sinks) + chain.from_iterable(choice_sink.holdback_per_choice.items() for _, choice_sink, _ in choice_rounds) ) return tool_calls: Final = chain.from_iterable(choice.message.tool_calls or () for choice in assembled.choices) 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 9cce5d8ff44..6ba912f25df 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 @@ -1399,6 +1399,45 @@ class TestOpenAIChatCompletionsHandlerStreamingOutput: for inputs in guardrail.seen_inputs ] == [[("call_1", '{"fruit": "persimmon"}')], [("call_2", '{"fruit": "durian"}')]] + @pytest.mark.asyncio + @pytest.mark.parametrize("transform", [False, True]) + async def test_each_choice_refreshes_shared_guardrail_response_context(self, transform: bool) -> None: + from fastapi import HTTPException + from litellm.llms.base_llm.guardrail_translation.base_translation import StreamTransformSink + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + class ContextGuardrail(CustomGuardrail): + async def apply_guardrail( + self, inputs: GenericGuardrailAPIInputs, request_data: dict[str, object], + input_type: Literal["request", "response"], logging_obj: object = None, + ) -> GenericGuardrailAPIInputs: + response: Final = request_data["response"] + assert isinstance(response, ModelResponse) + request_data["visited_choices"] = (*request_data.get("visited_choices", ()), response.choices[0].index) + if response.choices[0].message.content == "forbidden": + raise HTTPException(status_code=400, detail="later choice blocked") + return inputs + + chunks: Final = [ModelResponseStream(choices=[StreamingChoices( + index=index, delta=Delta(content=text, tool_calls=[{ + "index": 0, "id": f"call_{index}", "type": "function", + "function": {"name": "contact", "arguments": "{}"}, + }]), finish_reason="tool_calls", + ) for index, text in enumerate(("allowed", "forbidden"))])] + request_data: Final = {"metadata": {"trace": "retained"}} + with pytest.raises(HTTPException, match="later choice blocked"): + await OpenAIChatCompletionsHandler().process_output_streaming_response( + responses_so_far=chunks, guardrail_to_apply=ContextGuardrail(guardrail_name="context"), + request_data=request_data, deliver_ended_stream_rewrites=True, + stream_transform_sink=StreamTransformSink() if transform else None, + ) + assert request_data["visited_choices"] == (0, 1) + assert request_data["metadata"]["trace"] == "retained" + assert request_data["responses"] is chunks + restored: Final = request_data["response"] + assert isinstance(restored, ModelResponse) + assert tuple(choice.index for choice in restored.choices) == (0, 1) + @pytest.mark.asyncio async def test_deliver_ended_stream_tool_rewrites_keep_choice_indices(self) -> None: handler: Final = OpenAIChatCompletionsHandler()