mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix: protect other guardrail translations
This commit is contained in:
parent
97e67b0e48
commit
9ca7d7340d
5 changed files with 385 additions and 35 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"])
|
||||
|
|
@ -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"])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue