From 9ca7d7340d4d47c53512daa98d559b623635e926 Mon Sep 17 00:00:00 2001 From: Alexander Grattan Date: Mon, 26 Jan 2026 12:04:33 -0500 Subject: [PATCH] fix: protect other guardrail translations --- .../chat/guardrail_translation/handler.py | 58 +++--- .../chat/guardrail_translation/handler.py | 21 +- .../guardrail_translation/handler.py | 6 +- .../test_anthropic_guardrail_handler.py | 187 ++++++++++++++++++ .../test_openai_guardrail_handler.py | 148 ++++++++++++++ 5 files changed, 385 insertions(+), 35 deletions(-) create mode 100644 tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 9d50cc4d92d..40cd978ff9c 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -35,7 +35,6 @@ from litellm.types.llms.openai import ( from litellm.types.utils import ( ChatCompletionMessageToolCall, GenericGuardrailAPIInputs, - ModelResponse, ) if TYPE_CHECKING: @@ -84,9 +83,9 @@ class AnthropicMessagesHandler(BaseTranslation): texts_to_check: List[str] = [] images_to_check: List[str] = [] - tools_to_check: List[ChatCompletionToolParam] = ( - chat_completion_compatible_request.get("tools", []) - ) + tools_to_check: List[ + ChatCompletionToolParam + ] = chat_completion_compatible_request.get("tools", []) task_mappings: List[Tuple[int, Optional[int]]] = [] # Track (message_index, content_index) for each text # content_index is None for string content, int for list content @@ -278,7 +277,10 @@ class AnthropicMessagesHandler(BaseTranslation): if hasattr(content_block, "model_dump"): block_dict = content_block.model_dump() else: - block_dict = {"type": block_type, "text": getattr(content_block, "text", None)} + block_dict = { + "type": block_type, + "text": getattr(content_block, "text", None), + } else: continue @@ -346,30 +348,35 @@ class AnthropicMessagesHandler(BaseTranslation): """ has_ended = self._check_streaming_has_ended(responses_so_far) if has_ended: - # build the model response from the responses_so_far - model_response = cast( - ModelResponse, + model_response = ( AnthropicPassthroughLoggingHandler._build_complete_streaming_response( all_chunks=responses_so_far, litellm_logging_obj=cast("LiteLLMLoggingObj", litellm_logging_obj), model="", - ), + ) ) - tool_calls_list = cast(Optional[List[ChatCompletionMessageToolCall]], model_response.choices[0].message.tool_calls) # type: ignore - string_so_far = model_response.choices[0].message.content # type: ignore - guardrail_inputs = GenericGuardrailAPIInputs() - if string_so_far: - guardrail_inputs["texts"] = [string_so_far] - if tool_calls_list: - guardrail_inputs["tool_calls"] = tool_calls_list - _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( # allow rejecting the response, if invalid - inputs=guardrail_inputs, - request_data={}, - input_type="response", - logging_obj=litellm_logging_obj, - ) + # Check if model_response is valid and has choices before accessing + if ( + model_response is not None + and hasattr(model_response, "choices") + and model_response.choices + ): + tool_calls_list = cast(Optional[List[ChatCompletionMessageToolCall]], model_response.choices[0].message.tool_calls) # type: ignore + string_so_far = model_response.choices[0].message.content # type: ignore + guardrail_inputs = GenericGuardrailAPIInputs() + if string_so_far: + guardrail_inputs["texts"] = [string_so_far] + if tool_calls_list: + guardrail_inputs["tool_calls"] = tool_calls_list + + _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( # allow rejecting the response, if invalid + inputs=guardrail_inputs, + request_data={}, + input_type="response", + logging_obj=litellm_logging_obj, + ) return responses_so_far string_so_far = self.get_streaming_string_so_far(responses_so_far) @@ -552,7 +559,7 @@ class AnthropicMessagesHandler(BaseTranslation): response_content = response.get("content", []) else: response_content = getattr(response, "content", None) or [] - + if not response_content: return False for content_block in response_content: @@ -636,7 +643,10 @@ class AnthropicMessagesHandler(BaseTranslation): if isinstance(content_block, dict): if content_block.get("type") == "text": cast(Dict[str, Any], content_block)["text"] = guardrail_response - elif hasattr(content_block, "type") and getattr(content_block, "type", None) == "text": + elif ( + hasattr(content_block, "type") + and getattr(content_block, "type", None) == "text" + ): # Update Pydantic object's text attribute if hasattr(content_block, "text"): content_block.text = guardrail_response diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index d0ed3f165cc..c86de1448f9 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -21,7 +21,13 @@ from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation from litellm.main import stream_chunk_builder from litellm.types.llms.openai import ChatCompletionToolParam -from litellm.types.utils import Choices, GenericGuardrailAPIInputs, ModelResponse, ModelResponseStream, StreamingChoices +from litellm.types.utils import ( + Choices, + GenericGuardrailAPIInputs, + ModelResponse, + ModelResponseStream, + StreamingChoices, +) if TYPE_CHECKING: from litellm.integrations.custom_guardrail import CustomGuardrail @@ -80,9 +86,9 @@ class OpenAIChatCompletionsHandler(BaseTranslation): if tool_calls_to_check: inputs["tool_calls"] = tool_calls_to_check # type: ignore if messages: - inputs["structured_messages"] = ( - messages # pass the openai /chat/completions messages to the guardrail, as-is - ) + inputs[ + "structured_messages" + ] = messages # pass the openai /chat/completions messages to the guardrail, as-is # Pass tools (function definitions) to the guardrail tools = data.get("tools") if tools: @@ -355,14 +361,17 @@ class OpenAIChatCompletionsHandler(BaseTranslation): # check if the stream has ended has_stream_ended = False for chunk in responses_so_far: - if chunk.choices[0].finish_reason is not None: + if chunk.choices and chunk.choices[0].finish_reason is not None: has_stream_ended = True break if has_stream_ended: # convert to model response model_response = cast( - ModelResponse, stream_chunk_builder(chunks=responses_so_far, logging_obj=litellm_logging_obj) + ModelResponse, + stream_chunk_builder( + chunks=responses_so_far, logging_obj=litellm_logging_obj + ), ) # run process_output_response await self.process_output_response( diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index 598c91bb128..1164717f272 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -311,9 +311,7 @@ class OpenAIResponsesHandler(BaseTranslation): return response if not response_output: - verbose_proxy_logger.debug( - "OpenAI Responses API: Empty output in response" - ) + verbose_proxy_logger.debug("OpenAI Responses API: Empty output in response") return response # Step 1: Extract all text content and tool calls from response output @@ -485,11 +483,9 @@ class OpenAIResponsesHandler(BaseTranslation): # Check if it's an OutputText with text if isinstance(content_item, OutputText): if content_item.text: - return True elif isinstance(content_item, dict): if content_item.get("text"): - return True return False diff --git a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py new file mode 100644 index 00000000000..f8632105da5 --- /dev/null +++ b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -0,0 +1,187 @@ +""" +Unit tests for Anthropic Messages Guardrail Translation Handler + +Tests the handler's ability to process streaming output for Anthropic Messages API +with guardrail transformations, specifically testing edge cases with empty choices. +""" + +import os +import sys +from typing import Any, List, Literal, Optional +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert( + 0, os.path.abspath("../../../../../../..") +) # Adds the parent directory to the system path + +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.llms.anthropic.chat.guardrail_translation.handler import ( + AnthropicMessagesHandler, +) +from litellm.types.utils import GenericGuardrailAPIInputs + + +class MockPassThroughGuardrail(CustomGuardrail): + """Mock guardrail that passes through without blocking - for testing streaming fallback behavior""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + """Simply return inputs unchanged""" + return inputs + + +class TestAnthropicMessagesHandlerStreamingOutputProcessing: + """Test streaming output processing functionality""" + + @pytest.mark.asyncio + async def test_process_output_streaming_response_empty_model_response(self): + """Test that streaming response with None model_response doesn't raise error + + This test verifies the fix for the bug where accessing model_response.choices[0] + would raise an error when _build_complete_streaming_response returns None. + """ + handler = AnthropicMessagesHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Mock _check_streaming_has_ended to return True (stream ended) + # and _build_complete_streaming_response to return None + with patch.object( + handler, "_check_streaming_has_ended", return_value=True + ), patch( + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + return_value=None, + ): + responses_so_far = [b"data: some chunk"] + + # This should not raise an error + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=MagicMock(), + ) + + # Should return the responses unchanged + assert result == responses_so_far + + @pytest.mark.asyncio + async def test_process_output_streaming_response_empty_choices(self): + """Test that streaming response with empty choices doesn't raise IndexError + + This test verifies the fix for the bug where accessing model_response.choices[0] + would raise IndexError when the response has an empty choices list. + """ + from litellm.types.utils import ModelResponse + + handler = AnthropicMessagesHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Create a mock response with empty choices + mock_response = ModelResponse( + id="msg_123", + created=1234567890, + model="claude-3", + object="chat.completion", + choices=[], # Empty choices + ) + + # Mock _check_streaming_has_ended to return True (stream ended) + # and _build_complete_streaming_response to return the mock response + with patch.object( + handler, "_check_streaming_has_ended", return_value=True + ), patch( + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + return_value=mock_response, + ): + responses_so_far = [b"data: some chunk"] + + # This should not raise IndexError + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=MagicMock(), + ) + + # Should return the responses unchanged + assert result == responses_so_far + + @pytest.mark.asyncio + async def test_process_output_streaming_response_with_valid_choices(self): + """Test that streaming response with valid choices still works correctly""" + from litellm.types.utils import Choices, Message, ModelResponse + + handler = AnthropicMessagesHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Create a mock response with valid choices + mock_response = ModelResponse( + id="msg_123", + created=1234567890, + model="claude-3", + object="chat.completion", + choices=[ + Choices( + finish_reason="stop", + index=0, + message=Message( + content="Hello world", + role="assistant", + ), + ) + ], + ) + + # Mock _check_streaming_has_ended to return True (stream ended) + # and _build_complete_streaming_response to return the mock response + with patch.object( + handler, "_check_streaming_has_ended", return_value=True + ), patch( + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + return_value=mock_response, + ): + responses_so_far = [b"data: some chunk"] + + # This should process successfully + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=MagicMock(), + ) + + # Should return the responses + assert result == responses_so_far + + @pytest.mark.asyncio + async def test_process_output_streaming_response_stream_not_ended(self): + """Test that streaming response falls back to text processing when stream hasn't ended""" + handler = AnthropicMessagesHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Mock _check_streaming_has_ended to return False (stream not ended) + with patch.object( + handler, "_check_streaming_has_ended", return_value=False + ), patch.object( + handler, "get_streaming_string_so_far", return_value="partial text" + ): + responses_so_far = [b"data: some chunk"] + + # This should process successfully using text-based guardrail + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=MagicMock(), + ) + + # Should return the responses + assert result == responses_so_far + + +if __name__ == "__main__": + # Run the tests + pytest.main([__file__, "-v"]) diff --git a/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py b/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py index 6c0195d2831..1f5f53d0f0c 100644 --- a/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py +++ b/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py @@ -733,6 +733,154 @@ class TestOpenAIChatCompletionsHandlerToolCallsOutput: assert response.choices[0].finish_reason == "tool_calls" +class MockPassThroughGuardrail(CustomGuardrail): + """Mock guardrail that passes through without blocking - for testing streaming fallback behavior""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + """Simply return inputs unchanged""" + return inputs + + +class TestOpenAIChatCompletionsHandlerStreamingOutput: + """Test streaming output processing functionality""" + + @pytest.mark.asyncio + async def test_process_output_streaming_response_empty_choices(self): + """Test that streaming response with empty choices doesn't raise IndexError + + This test verifies the fix for the bug where accessing chunk.choices[0] + would raise IndexError when a streaming chunk has an empty choices list. + """ + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + handler = OpenAIChatCompletionsHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Create a streaming chunk with empty choices + chunk_with_empty_choices = ModelResponseStream( + id="chatcmpl-123", + created=1234567890, + model="gpt-4", + object="chat.completion.chunk", + choices=[], # Empty choices - this was causing the IndexError + ) + + responses_so_far = [chunk_with_empty_choices] + + # This should not raise IndexError + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=None, + ) + + # Should return the responses unchanged + assert result == responses_so_far + + @pytest.mark.asyncio + async def test_process_output_streaming_response_with_valid_choices(self): + """Test that streaming response with valid choices still works correctly""" + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + handler = OpenAIChatCompletionsHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Create streaming chunks with valid choices + chunk1 = ModelResponseStream( + id="chatcmpl-123", + created=1234567890, + model="gpt-4", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + index=0, + delta=Delta(content="Hello"), + finish_reason=None, + ) + ], + ) + + chunk2 = ModelResponseStream( + id="chatcmpl-123", + created=1234567890, + model="gpt-4", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + index=0, + delta=Delta(content=" world"), + finish_reason="stop", + ) + ], + ) + + responses_so_far = [chunk1, chunk2] + + # This should process successfully + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=None, + ) + + # Should return the responses + assert result == responses_so_far + + @pytest.mark.asyncio + async def test_process_output_streaming_response_mixed_empty_and_valid_choices_no_finish(self): + """Test streaming response with mix of empty and valid choices chunks (stream not finished) + + This tests the has_stream_ended check when iterating through chunks with mixed choices. + The stream hasn't finished yet (no finish_reason), so it won't trigger stream_chunk_builder. + """ + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + handler = OpenAIChatCompletionsHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Mix of chunks - some with empty choices, some with valid choices + # Stream hasn't finished (no finish_reason) + chunk_empty = ModelResponseStream( + id="chatcmpl-123", + created=1234567890, + model="gpt-4", + object="chat.completion.chunk", + choices=[], + ) + + chunk_valid = ModelResponseStream( + id="chatcmpl-123", + created=1234567890, + model="gpt-4", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + index=0, + delta=Delta(content="Hello"), + finish_reason=None, # Stream not finished + ) + ], + ) + + responses_so_far = [chunk_empty, chunk_valid] + + # This should not raise IndexError when checking has_stream_ended + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=None, + ) + + # Should return the responses + assert result == responses_so_far + + if __name__ == "__main__": # Run the tests pytest.main([__file__, "-v"])