From 0bdfd95ad8793e6f62e401a1172d8ae8fbdfdde8 Mon Sep 17 00:00:00 2001 From: "Ethan T." Date: Sat, 14 Mar 2026 15:12:48 +0800 Subject: [PATCH 01/25] fix: map Chat Completion file type to Responses API input_file When bridging /chat/completions to the Responses API, content items with type 'file' were falling through to the default handler and being stringified as input_text. This caused the model to receive the Python dict representation as plain text instead of the actual file content. Add explicit handling for type 'file' that correctly maps: {"type": "file", "file": {"file_data": "...", "filename": "..."}} to: {"type": "input_file", "file_data": "...", "filename": "..."} Fixes BerriAI/litellm#23588 --- .../transformation.py | 14 ++++++++++++++ 1 file changed, 14 insertions(+) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 4b31bcfc285..ab69e9ca13f 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -693,6 +693,20 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): verbose_logger.debug( f"Chat provider: image -> {converted}" ) + elif item_type == "file": + # Map Chat Completion file to Responses API input_file + # {"type": "file", "file": {"file_data": "...", "filename": "..."}} + # -> {"type": "input_file", "file_data": "...", "filename": "..."} + file_data = item.get("file", {}) + converted = {"type": "input_file"} + if isinstance(file_data, dict): + for key in ["file_id", "file_data", "filename"]: + if key in file_data: + converted[key] = file_data[key] + result.append(converted) + verbose_logger.debug( + f"Chat provider: file -> {converted}" + ) elif item_type in [ "input_text", "input_image", From 71c9ba0b1b620ce40b5149376e18cf87bd07f000 Mon Sep 17 00:00:00 2001 From: "Ethan T." Date: Sat, 14 Mar 2026 15:12:54 +0800 Subject: [PATCH 02/25] test: add tests for file type to input_file mapping --- ...responses_transformation_transformation.py | 88 +++++++++++++++++++ 1 file changed, 88 insertions(+) diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index da383532690..c1d102f65ee 100644 --- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -1995,3 +1995,91 @@ def test_map_optional_params_preserves_reasoning_summary(): assert responses_api_request["reasoning"] == {"effort": "high", "summary": "detailed"} assert responses_api_request["reasoning"]["effort"] == "high" assert responses_api_request["reasoning"]["summary"] == "detailed" + + +def test_convert_chat_completion_file_type_to_input_file(): + """ + Test that Chat Completion content with type 'file' is correctly mapped + to Responses API 'input_file' format, not stringified as 'input_text'. + + Regression test for https://github.com/BerriAI/litellm/issues/23588 + """ + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + LiteLLMResponsesTransformationHandler, + ) + + handler = LiteLLMResponsesTransformationHandler() + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is in this PDF?"}, + { + "type": "file", + "file": { + "file_data": "data:application/pdf;base64,JVBERi0xLjQK", + "filename": "test.pdf", + }, + }, + ], + } + ] + + input_items, instructions = handler.convert_chat_completion_messages_to_responses_api( + messages + ) + + assert len(input_items) == 1 + msg = input_items[0] + assert msg["type"] == "message" + assert msg["role"] == "user" + + content = msg["content"] + assert len(content) == 2 + + # First item should be the text + assert content[0]["type"] == "input_text" + assert content[0]["text"] == "What is in this PDF?" + + # Second item should be input_file, NOT input_text with stringified dict + assert content[1]["type"] == "input_file" + assert content[1]["file_data"] == "data:application/pdf;base64,JVBERi0xLjQK" + assert content[1]["filename"] == "test.pdf" + # Ensure it does NOT have the nested 'file' key + assert "file" not in content[1] + + +def test_convert_chat_completion_file_type_with_file_id(): + """ + Test that Chat Completion content with type 'file' using file_id is correctly mapped. + """ + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + LiteLLMResponsesTransformationHandler, + ) + + handler = LiteLLMResponsesTransformationHandler() + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Summarize this file."}, + { + "type": "file", + "file": { + "file_id": "file-abc123", + }, + }, + ], + } + ] + + input_items, instructions = handler.convert_chat_completion_messages_to_responses_api( + messages + ) + + content = input_items[0]["content"] + assert content[1]["type"] == "input_file" + assert content[1]["file_id"] == "file-abc123" + assert "file_data" not in content[1] From 6658a8ffb3216a8af1fc1f2cdfa2bd9feae8b39e Mon Sep 17 00:00:00 2001 From: "Ethan T." Date: Sat, 14 Mar 2026 21:18:35 +0800 Subject: [PATCH 03/25] style: apply black formatting to transformation.py --- .../transformation.py | 25 +++++++++++-------- 1 file changed, 14 insertions(+), 11 deletions(-) diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index ab69e9ca13f..f5856ab1f45 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -240,10 +240,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): if key in ("max_tokens", "max_completion_tokens"): responses_api_request["max_output_tokens"] = value elif key == "tools" and value is not None: - responses_api_request[ - "tools" - ] = self._convert_tools_to_responses_format( - cast(List[Dict[str, Any]], value) + responses_api_request["tools"] = ( + self._convert_tools_to_responses_format( + cast(List[Dict[str, Any]], value) + ) ) elif key == "response_format": text_format = self._transform_response_format_to_text_format(value) @@ -398,6 +398,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): ResponseOutputMessage, ResponseReasoningItem, ) + try: from openai.types.responses.response_output_item import ( ResponseApplyPatchToolCall, @@ -460,7 +461,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): accumulated_tool_calls.append(tool_call_dict) tool_call_index += 1 - elif ResponseApplyPatchToolCall is not None and isinstance(item, ResponseApplyPatchToolCall): + elif ResponseApplyPatchToolCall is not None and isinstance( + item, ResponseApplyPatchToolCall + ): from litellm.responses.litellm_completion_transformation.transformation import ( LiteLLMCompletionResponsesConfig, ) @@ -1069,9 +1072,9 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): ) if provider_specific_fields: - function_chunk[ - "provider_specific_fields" - ] = provider_specific_fields + function_chunk["provider_specific_fields"] = ( + provider_specific_fields + ) tool_call_index = parsed_chunk.get("output_index", 0) tool_call_chunk = ChatCompletionToolCallChunk( @@ -1144,9 +1147,9 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): # Add provider_specific_fields to function if present if provider_specific_fields: - function_chunk[ - "provider_specific_fields" - ] = provider_specific_fields + function_chunk["provider_specific_fields"] = ( + provider_specific_fields + ) tool_call_index = parsed_chunk.get("output_index", 0) tool_call_chunk = ChatCompletionToolCallChunk( From 98890e771de1d6850dcf9030b8becb99115c0849 Mon Sep 17 00:00:00 2001 From: "Ethan T." Date: Sat, 14 Mar 2026 21:19:03 +0800 Subject: [PATCH 04/25] style: apply black formatting to test file --- ...responses_transformation_transformation.py | 309 ++++++++++++------ 1 file changed, 205 insertions(+), 104 deletions(-) diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index c1d102f65ee..8bc6ffc0505 100644 --- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -9,7 +9,9 @@ from unittest.mock import ANY, MagicMock, Mock, patch import httpx import pytest -sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system-path +sys.path.insert( + 0, os.path.abspath("../../..") +) # Adds the parent directory to the system-path import litellm @@ -117,7 +119,9 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_imag function_call_output = item break - assert function_call_output is not None, "function_call_output not found in response" + assert ( + function_call_output is not None + ), "function_call_output not found in response" assert function_call_output["call_id"] == "call_abc123" # Check that the output is correctly transformed @@ -127,8 +131,12 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_imag image_item = output[0] # Should be transformed to Responses API format - assert image_item["type"] == "input_image", f"Expected type 'input_image', got '{image_item.get('type')}'" - assert image_item["image_url"] == test_image_base64, "image_url should be a flat string, not a nested object" + assert ( + image_item["type"] == "input_image" + ), f"Expected type 'input_image', got '{image_item.get('type')}'" + assert ( + image_item["image_url"] == test_image_base64 + ), "image_url should be a flat string, not a nested object" assert "detail" in image_item, "detail field should be present" print("✓ Tool result with image correctly transformed to Responses API format") @@ -190,7 +198,9 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_text function_call_output = item break - assert function_call_output is not None, "function_call_output not found in response" + assert ( + function_call_output is not None + ), "function_call_output not found in response" assert function_call_output["call_id"] == "call_abc123" # Check that the output is correctly transformed to use input_text, not output_text @@ -200,12 +210,16 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_text text_item = output[0] # Should be transformed to use input_text for tool results in Responses API format - assert text_item["type"] == "input_text", ( - f"Expected type 'input_text' for tool result, got '{text_item.get('type')}'" - ) - assert text_item["text"] == "15 degrees", f"Expected text '15 degrees', got '{text_item.get('text')}'" + assert ( + text_item["type"] == "input_text" + ), f"Expected type 'input_text' for tool result, got '{text_item.get('type')}'" + assert ( + text_item["text"] == "15 degrees" + ), f"Expected text '15 degrees', got '{text_item.get('text')}'" - print("✓ Tool result with text correctly transformed to use input_text for Responses API format") + print( + "✓ Tool result with text correctly transformed to use input_text for Responses API format" + ) def test_openai_responses_chunk_parser_reasoning_summary(): @@ -214,7 +228,9 @@ def test_openai_responses_chunk_parser_reasoning_summary(): ) from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices - iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + iterator = OpenAiResponsesToChatCompletionStreamIterator( + streaming_response=None, sync_stream=True + ) chunk = { "delta": "**Compar", @@ -246,7 +262,9 @@ def test_chunk_parser_string_output_text_delta_produces_text(): ) from litellm.types.utils import ModelResponseStream - iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + iterator = OpenAiResponsesToChatCompletionStreamIterator( + streaming_response=None, sync_stream=True + ) chunk = {"type": "response.output_text.delta", "delta": "literal text"} @@ -267,7 +285,9 @@ def test_chunk_parser_enum_output_text_delta_produces_text(): from litellm.types.llms.openai import ResponsesAPIStreamEvents from litellm.types.utils import ModelResponseStream - iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + iterator = OpenAiResponsesToChatCompletionStreamIterator( + streaming_response=None, sync_stream=True + ) chunk = {"type": ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, "delta": "enum text"} @@ -288,7 +308,9 @@ def test_chunk_parser_function_call_added_produces_tool_use(): from litellm.types.llms.openai import ResponsesAPIStreamEvents from litellm.types.utils import ModelResponseStream - iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + iterator = OpenAiResponsesToChatCompletionStreamIterator( + streaming_response=None, sync_stream=True + ) chunk = { "type": ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, @@ -373,7 +395,9 @@ Tomorrow will bring its petitions and promises, but for now the city breathes slow and wide, and I learn to carry this small calm home.""" - output_text = ResponseOutputText(annotations=[], text=poem_text, type="output_text", logprobs=[]) + output_text = ResponseOutputText( + annotations=[], text=poem_text, type="output_text", logprobs=[] + ) output_message = ResponseOutputMessage( id="msg_04c8021b8b3188a00068e9ae0b92f4819dac64d85b4abb67ec", content=[output_text], @@ -385,7 +409,9 @@ and I learn to carry this small calm home.""" # Create usage information usage = ResponseAPIUsage( input_tokens=16, - input_tokens_details=InputTokensDetails(audio_tokens=None, cached_tokens=0, text_tokens=None), + input_tokens_details=InputTokensDetails( + audio_tokens=None, cached_tokens=0, text_tokens=None + ), output_tokens=195, output_tokens_details=OutputTokensDetails(reasoning_tokens=0, text_tokens=None), total_tokens=211, @@ -597,7 +623,9 @@ def test_transform_request_single_char_keys_not_matched(): assert result_correct.get("metadata") == {"user_id": "123"} assert result_correct.get("previous_response_id") == "resp_abc" - print("✓ Single-character keys are not incorrectly matched to metadata/previous_response_id") + print( + "✓ Single-character keys are not incorrectly matched to metadata/previous_response_id" + ) # ============================================================================= @@ -617,7 +645,9 @@ def test_message_done_does_not_emit_is_finished(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + iterator = OpenAiResponsesToChatCompletionStreamIterator( + streaming_response=None, sync_stream=True + ) chunk = { "type": "response.output_item.done", @@ -629,9 +659,9 @@ def test_message_done_does_not_emit_is_finished(): # After the fix, message completion should NOT set finish_reason # ModelResponseStream doesn't have is_finished - check finish_reason instead assert len(result.choices) > 0, "result should have choices" - assert result.choices[0].finish_reason is None or result.choices[0].finish_reason == "", ( - "message completion should not emit finish_reason" - ) + assert ( + result.choices[0].finish_reason is None or result.choices[0].finish_reason == "" + ), "message completion should not emit finish_reason" def test_response_completed_emits_is_finished(): @@ -643,7 +673,9 @@ def test_response_completed_emits_is_finished(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + iterator = OpenAiResponsesToChatCompletionStreamIterator( + streaming_response=None, sync_stream=True + ) chunk = {"type": "response.completed"} @@ -651,7 +683,9 @@ def test_response_completed_emits_is_finished(): # response.completed should emit finish_reason='stop' assert len(result.choices) > 0, "result should have choices" - assert result.choices[0].finish_reason == "stop", "response.completed should emit finish_reason='stop'" + assert ( + result.choices[0].finish_reason == "stop" + ), "response.completed should emit finish_reason='stop'" def test_response_completed_with_function_calls_emits_tool_calls_finish_reason(): @@ -670,7 +704,9 @@ def test_response_completed_with_function_calls_emits_tool_calls_finish_reason() OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + iterator = OpenAiResponsesToChatCompletionStreamIterator( + streaming_response=None, sync_stream=True + ) # Simulate a response.completed event with function_call in output # This matches what Azure/OpenAI sends for gpt-5.1-codex-mini and similar models @@ -696,9 +732,9 @@ def test_response_completed_with_function_calls_emits_tool_calls_finish_reason() # response.completed with function_call should emit finish_reason='tool_calls' assert len(result.choices) > 0, "result should have choices" - assert result.choices[0].finish_reason == "tool_calls", ( - "response.completed with function_call output should emit finish_reason='tool_calls'" - ) + assert ( + result.choices[0].finish_reason == "tool_calls" + ), "response.completed with function_call output should emit finish_reason='tool_calls'" def test_response_completed_with_message_only_emits_stop_finish_reason(): @@ -709,7 +745,9 @@ def test_response_completed_with_message_only_emits_stop_finish_reason(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + iterator = OpenAiResponsesToChatCompletionStreamIterator( + streaming_response=None, sync_stream=True + ) # Simulate a response.completed event with only message output chunk = { @@ -733,10 +771,9 @@ def test_response_completed_with_message_only_emits_stop_finish_reason(): # response.completed with only message should emit finish_reason='stop' assert len(result.choices) > 0, "result should have choices" - assert result.choices[0].finish_reason == "stop", ( - "response.completed with only message output should emit finish_reason='stop'" - ) - + assert ( + result.choices[0].finish_reason == "stop" + ), "response.completed with only message output should emit finish_reason='stop'" def test_response_completed_preserves_usage_with_cached_tokens(): @@ -752,7 +789,9 @@ def test_response_completed_preserves_usage_with_cached_tokens(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + iterator = OpenAiResponsesToChatCompletionStreamIterator( + streaming_response=None, sync_stream=True + ) chunk = { "type": "response.completed", @@ -781,12 +820,18 @@ def test_response_completed_preserves_usage_with_cached_tokens(): result = iterator.chunk_parser(chunk) assert result.usage is not None, "usage should be set on response.completed chunk" - assert result.usage.prompt_tokens == 1226, "prompt_tokens should map from input_tokens" - assert result.usage.completion_tokens == 5, "completion_tokens should map from output_tokens" - assert result.usage.prompt_tokens_details is not None, "prompt_tokens_details should be set" - assert result.usage.prompt_tokens_details.cached_tokens == 1024, ( - "cached_tokens should be preserved from input_tokens_details" - ) + assert ( + result.usage.prompt_tokens == 1226 + ), "prompt_tokens should map from input_tokens" + assert ( + result.usage.completion_tokens == 5 + ), "completion_tokens should map from output_tokens" + assert ( + result.usage.prompt_tokens_details is not None + ), "prompt_tokens_details should be set" + assert ( + result.usage.prompt_tokens_details.cached_tokens == 1024 + ), "cached_tokens should be preserved from input_tokens_details" def test_function_call_done_emits_is_finished(): @@ -800,7 +845,9 @@ def test_function_call_done_emits_is_finished(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + iterator = OpenAiResponsesToChatCompletionStreamIterator( + streaming_response=None, sync_stream=True + ) chunk = { "type": "response.output_item.done", @@ -820,9 +867,9 @@ def test_function_call_done_emits_is_finished(): "output_item.done for function_call must not emit finish_reason; " "response.completed is responsible for the terminal finish_reason" ) - assert not result.choices[0].delta.tool_calls, ( - "output_item.done for function_call must not include a duplicate tool_calls delta" - ) + assert not result.choices[ + 0 + ].delta.tool_calls, "output_item.done for function_call must not include a duplicate tool_calls delta" def test_text_plus_tool_calls_sequence(): @@ -837,7 +884,9 @@ def test_text_plus_tool_calls_sequence(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + iterator = OpenAiResponsesToChatCompletionStreamIterator( + streaming_response=None, sync_stream=True + ) # Simulate the sequence from OpenAI Responses API chunks = [ @@ -876,23 +925,28 @@ def test_text_plus_tool_calls_sequence(): # Check message done (index 2) does NOT have finish_reason set message_done_result = results[2] assert len(message_done_result.choices) > 0, "message done should have choices" - assert message_done_result.choices[0].finish_reason is None or message_done_result.choices[0].finish_reason == "", ( - "message done should not have finish_reason" - ) + assert ( + message_done_result.choices[0].finish_reason is None + or message_done_result.choices[0].finish_reason == "" + ), "message done should not have finish_reason" # Check function_call done (index 5) does NOT have finish_reason set # (response.completed is responsible for the terminal finish_reason) function_done_result = results[5] - assert len(function_done_result.choices) > 0, "function_call done should have choices" - assert function_done_result.choices[0].finish_reason is None, ( - "output_item.done for function_call must not emit finish_reason" - ) + assert ( + len(function_done_result.choices) > 0 + ), "function_call done should have choices" + assert ( + function_done_result.choices[0].finish_reason is None + ), "output_item.done for function_call must not emit finish_reason" # Check response.completed (index 6) has finish_reason='stop' # (the mock chunk has no nested 'response' data, so has_function_calls is False → 'stop') completed_result = results[6] assert len(completed_result.choices) > 0, "response.completed should have choices" - assert completed_result.choices[0].finish_reason == "stop", "response.completed should have finish_reason='stop'" + assert ( + completed_result.choices[0].finish_reason == "stop" + ), "response.completed should have finish_reason='stop'" # ============================================================================= @@ -958,7 +1012,9 @@ def test_tool_message_output_uses_input_text_not_output_text(): output = function_call_output["output"] assert isinstance(output, list), f"output should be a list, got {type(output)}" assert len(output) == 1 - assert output[0]["type"] == "input_text", f"Expected input_text, got {output[0].get('type')}" + assert ( + output[0]["type"] == "input_text" + ), f"Expected input_text, got {output[0].get('type')}" assert output[0]["text"] == '{"temperature": 15, "condition": "sunny"}' print("✓ Tool message output correctly uses input_text type") @@ -1144,9 +1200,13 @@ def test_map_reasoning_effort_adds_summary_detailed(): assert result is not None, f"Result should not be None for effort={effort}" assert result["effort"] == effort, f"Effort should be {effort}" - assert "summary" not in result, f"Summary should NOT be present by default for effort={effort}" + assert ( + "summary" not in result + ), f"Summary should NOT be present by default for effort={effort}" - print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}' (no summary by default)") + print( + f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}' (no summary by default)" + ) # Test 2: With flag enabled - summary IS added litellm.reasoning_auto_summary = True @@ -1156,9 +1216,9 @@ def test_map_reasoning_effort_adds_summary_detailed(): assert result is not None, f"Result should not be None for effort={effort}" assert result["effort"] == effort, f"Effort should be {effort}" - assert result["summary"] == "detailed", ( - f"Summary should be 'detailed' when flag is enabled for effort={effort}" - ) + assert ( + result["summary"] == "detailed" + ), f"Summary should be 'detailed' when flag is enabled for effort={effort}" print( f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}', summary='detailed' (flag enabled)" @@ -1169,7 +1229,9 @@ def test_map_reasoning_effort_adds_summary_detailed(): os.environ["LITELLM_REASONING_AUTO_SUMMARY"] = "true" result = handler._map_reasoning_effort("high") - assert result["summary"] == "detailed", "Summary should be 'detailed' when env var is enabled" + assert ( + result["summary"] == "detailed" + ), "Summary should be 'detailed' when env var is enabled" print("✓ LITELLM_REASONING_AUTO_SUMMARY env var works correctly") # Test 4: Dict input is passed through as-is (no modification) @@ -1188,7 +1250,9 @@ def test_map_reasoning_effort_adds_summary_detailed(): assert result_unknown is None print("✓ Unknown reasoning_effort values return None") - print("✓ All reasoning_effort behaviors work correctly with flag/env var control") + print( + "✓ All reasoning_effort behaviors work correctly with flag/env var control" + ) finally: # Restore original values @@ -1264,7 +1328,9 @@ def test_transform_response_preserves_annotations(): # Create usage information usage = ResponseAPIUsage( input_tokens=10, - input_tokens_details=InputTokensDetails(audio_tokens=None, cached_tokens=0, text_tokens=None), + input_tokens_details=InputTokensDetails( + audio_tokens=None, cached_tokens=0, text_tokens=None + ), output_tokens=20, output_tokens_details=OutputTokensDetails(reasoning_tokens=0, text_tokens=None), total_tokens=30, @@ -1351,9 +1417,13 @@ def test_transform_response_preserves_annotations(): assert choice.message.content == "Here is some information with citations." # Check that annotations are preserved - assert hasattr(choice.message, "annotations"), "Message should have annotations attribute" + assert hasattr( + choice.message, "annotations" + ), "Message should have annotations attribute" assert choice.message.annotations is not None, "Annotations should not be None" - assert len(choice.message.annotations) == 2, f"Expected 2 annotations, got {len(choice.message.annotations)}" + assert ( + len(choice.message.annotations) == 2 + ), f"Expected 2 annotations, got {len(choice.message.annotations)}" # Verify annotation content annotation1 = choice.message.annotations[0] @@ -1375,7 +1445,9 @@ def test_transform_response_preserves_annotations(): assert result.usage.completion_tokens == 20 assert result.usage.total_tokens == 30 - print("✓ Annotations from Responses API are correctly preserved in Chat Completions format") + print( + "✓ Annotations from Responses API are correctly preserved in Chat Completions format" + ) def test_apply_patch_tool_call_converted_to_chat_completion_tool_call(): @@ -1512,6 +1584,8 @@ def test_apply_patch_tool_call_converted_to_chat_completion_tool_call(): assert args["type"] == "create_file" assert args["path"] == "hello.py" assert "print('hello world')" in args["diff"] + + def test_multi_tool_call_stream_no_premature_finish(): """ Regression test for multi-tool-call streaming bug. @@ -1538,18 +1612,26 @@ def test_multi_tool_call_stream_no_premature_finish(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + iterator = OpenAiResponsesToChatCompletionStreamIterator( + streaming_response=None, sync_stream=True + ) chunks = [ # 0: response created - {"type": "response.created", "response": {"id": "resp_001", "status": "in_progress"}}, + { + "type": "response.created", + "response": {"id": "resp_001", "status": "in_progress"}, + }, # 1: first tool call added { "type": "response.output_item.added", "item": {"type": "function_call", "name": "read_file", "call_id": "call_1"}, }, # 2: first tool call arguments delta - {"type": "response.function_call_arguments.delta", "delta": '{"path":"/etc/hostname"}'}, + { + "type": "response.function_call_arguments.delta", + "delta": '{"path":"/etc/hostname"}', + }, # 3: first tool call done ← must NOT emit finish_reason { "type": "response.output_item.done", @@ -1608,10 +1690,12 @@ def test_multi_tool_call_stream_no_premature_finish(): r = results[done_idx] assert r is not None, f"{label}: chunk_parser must return a result" assert len(r.choices) > 0, f"{label}: result must have choices" - assert r.choices[0].finish_reason is None, ( - f"{label}: output_item.done must not emit finish_reason (stream would terminate prematurely)" - ) - assert not r.choices[0].delta.tool_calls, ( + assert ( + r.choices[0].finish_reason is None + ), f"{label}: output_item.done must not emit finish_reason (stream would terminate prematurely)" + assert not r.choices[ + 0 + ].delta.tool_calls, ( f"{label}: output_item.done must not include a duplicate tool_calls delta" ) @@ -1623,12 +1707,12 @@ def test_multi_tool_call_stream_no_premature_finish(): r = results[added_idx] if r is not None and r.choices and r.choices[0].delta.tool_calls: tc = r.choices[0].delta.tool_calls[0] - assert tc.function.name == expected_name, ( - f"output_item.added for {expected_name}: tool_call name mismatch" - ) - assert tc.id == expected_call_id, ( - f"output_item.added for {expected_name}: call_id mismatch" - ) + assert ( + tc.function.name == expected_name + ), f"output_item.added for {expected_name}: tool_call name mismatch" + assert ( + tc.id == expected_call_id + ), f"output_item.added for {expected_name}: call_id mismatch" # 3. argument delta events (indices 2 and 5) should carry arguments for delta_idx, expected_args, label in [ @@ -1638,17 +1722,17 @@ def test_multi_tool_call_stream_no_premature_finish(): r = results[delta_idx] if r is not None and r.choices and r.choices[0].delta.tool_calls: tc = r.choices[0].delta.tool_calls[0] - assert tc.function.arguments == expected_args, ( - f"{label}: argument delta mismatch" - ) + assert ( + tc.function.arguments == expected_args + ), f"{label}: argument delta mismatch" # 4. Only response.completed (index 7) emits the terminal finish_reason completed_result = results[7] assert completed_result is not None, "response.completed must return a result" assert len(completed_result.choices) > 0, "response.completed must have choices" - assert completed_result.choices[0].finish_reason == "tool_calls", ( - "response.completed with function_call outputs must emit finish_reason='tool_calls'" - ) + assert ( + completed_result.choices[0].finish_reason == "tool_calls" + ), "response.completed with function_call outputs must emit finish_reason='tool_calls'" # 5. No chunk before the last one should have finish_reason set for idx, r in enumerate(results[:-1]): @@ -1658,7 +1742,9 @@ def test_multi_tool_call_stream_no_premature_finish(): f"— only response.completed should terminate the stream" ) - print("✓ Multi-tool-call stream completes without premature finish_reason termination") + print( + "✓ Multi-tool-call stream completes without premature finish_reason termination" + ) # ============================================================================= @@ -1790,7 +1876,10 @@ def test_parallel_tool_calls_comprehensive_streaming_integration(): chunks = [ # 0: response.created - {"type": "response.created", "response": {"id": "resp_001", "status": "in_progress"}}, + { + "type": "response.created", + "response": {"id": "resp_001", "status": "in_progress"}, + }, # 1: call_1 (read_file) added — output_index=0 { "type": "response.output_item.added", @@ -1873,7 +1962,9 @@ def test_parallel_tool_calls_comprehensive_streaming_integration(): }, ] - iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + iterator = OpenAiResponsesToChatCompletionStreamIterator( + streaming_response=None, sync_stream=True + ) results = [iterator.chunk_parser(chunk) for chunk in chunks] # 1. output_item.done events (indices 4 and 8) must NOT emit finish_reason @@ -1885,7 +1976,9 @@ def test_parallel_tool_calls_comprehensive_streaming_integration(): f"{label}: output_item.done must not emit finish_reason " f"(would prematurely terminate stream before subsequent tool calls arrive)" ) - assert not r.choices[0].delta.tool_calls, ( + assert not r.choices[ + 0 + ].delta.tool_calls, ( f"{label}: output_item.done must not emit a duplicate tool_calls delta" ) @@ -1919,7 +2012,9 @@ def test_parallel_tool_calls_comprehensive_streaming_integration(): for tc in tool_calls: if tc.function and tc.function.arguments: idx = tc.index - assembled_args[idx] = assembled_args.get(idx, "") + tc.function.arguments + assembled_args[idx] = ( + assembled_args.get(idx, "") + tc.function.arguments + ) # delta 1 = '{"path":' + delta 2 = '"/etc/foo"}' → '{"path":"/etc/foo"}' assert assembled_args.get(0) == '{"path":"/etc/foo"}', ( @@ -1938,16 +2033,16 @@ def test_parallel_tool_calls_comprehensive_streaming_integration(): for i, r in enumerate(results) if r is not None and r.choices and r.choices[0].finish_reason ] - assert len(finish_events) == 1, ( - f"Expected exactly 1 finish event, got {len(finish_events)}: {finish_events}" - ) + assert ( + len(finish_events) == 1 + ), f"Expected exactly 1 finish event, got {len(finish_events)}: {finish_events}" assert finish_events[0][0] == len(chunks) - 1, ( f"Finish event must be at the last chunk (index {len(chunks) - 1}), " f"but was at index {finish_events[0][0]}" ) - assert finish_events[0][1] == "tool_calls", ( - f"Terminal finish_reason must be 'tool_calls', got '{finish_events[0][1]}'" - ) + assert ( + finish_events[0][1] == "tool_calls" + ), f"Terminal finish_reason must be 'tool_calls', got '{finish_events[0][1]}'" # 5. Parallel tool calls have distinct indices matching output_index (0 and 1) # Collect indices from output_item.added chunks only (they carry the call id) @@ -1958,16 +2053,19 @@ def test_parallel_tool_calls_comprehensive_streaming_integration(): for tc in r.choices[0].delta.tool_calls if tc.id # output_item.added chunks carry the id; argument deltas do not ] - assert set(added_tool_call_indices) == {0, 1}, ( - f"Parallel tool calls must have distinct indices {{0, 1}}, got: {set(added_tool_call_indices)}" - ) + assert set(added_tool_call_indices) == { + 0, + 1, + }, f"Parallel tool calls must have distinct indices {{0, 1}}, got: {set(added_tool_call_indices)}" - print("✓ Parallel tool calls with split argument deltas stream correctly end-to-end") + print( + "✓ Parallel tool calls with split argument deltas stream correctly end-to-end" + ) def test_map_optional_params_preserves_reasoning_summary(): """Test that reasoning_effort dict with summary field is preserved. - + Regression test for: User reported that summary field was being dropped when routing to Responses API. The dict format should be fully preserved. """ @@ -1992,7 +2090,10 @@ def test_map_optional_params_preserves_reasoning_summary(): # Verify reasoning_effort dict with summary was fully preserved assert "reasoning" in responses_api_request - assert responses_api_request["reasoning"] == {"effort": "high", "summary": "detailed"} + assert responses_api_request["reasoning"] == { + "effort": "high", + "summary": "detailed", + } assert responses_api_request["reasoning"]["effort"] == "high" assert responses_api_request["reasoning"]["summary"] == "detailed" @@ -2026,8 +2127,8 @@ def test_convert_chat_completion_file_type_to_input_file(): } ] - input_items, instructions = handler.convert_chat_completion_messages_to_responses_api( - messages + input_items, instructions = ( + handler.convert_chat_completion_messages_to_responses_api(messages) ) assert len(input_items) == 1 @@ -2075,8 +2176,8 @@ def test_convert_chat_completion_file_type_with_file_id(): } ] - input_items, instructions = handler.convert_chat_completion_messages_to_responses_api( - messages + input_items, instructions = ( + handler.convert_chat_completion_messages_to_responses_api(messages) ) content = input_items[0]["content"] From 5acceaed32e03c2c44e6c292ebc2fad83f3c0ba9 Mon Sep 17 00:00:00 2001 From: Chesars Date: Mon, 16 Mar 2026 11:49:21 -0300 Subject: [PATCH 05/25] fix(model-prices): restore gpt-4-0314 entry lost in merge conflict The entry was accidentally dropped in commit 6bd7cd7 during a merge conflict resolution. The model is deprecated but still accessible for existing users until its shutdown date of 2026-03-26 per OpenAI docs. Fixes #23738 --- model_prices_and_context_window.json | 13 +++++++++++++ 1 file changed, 13 insertions(+) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 6786fc33595..1ab96a445d9 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -16864,6 +16864,19 @@ "supports_system_messages": true, "supports_tool_choice": true }, + "gpt-4-0314": { + "deprecation_date": "2026-03-26", + "input_cost_per_token": 3e-05, + "litellm_provider": "openai", + "max_input_tokens": 8192, + "max_output_tokens": 4096, + "max_tokens": 4096, + "mode": "chat", + "output_cost_per_token": 6e-05, + "supports_prompt_caching": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, "gpt-4-0613": { "deprecation_date": "2025-06-06", "input_cost_per_token": 3e-05, From 84b4af40fa7c8fd5a249ced6203da4c09c7dad87 Mon Sep 17 00:00:00 2001 From: Awais Qureshi Date: Tue, 17 Mar 2026 10:30:18 +0500 Subject: [PATCH 06/25] fix(fireworks): skip #transform=inline for base64 data URLs (#23729) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix(fireworks): skip #transform=inline for base64 data URLs Closes #23583 Appending #transform=inline to a data: URL corrupted the base64 payload, causing binascii.Error (Incorrect padding) when Fireworks AI attempted to decode the image. Data URLs are already inlined so the fragment is a no-op anyway — guard both the str and dict image_url branches to skip the suffix when the URL starts with "data:". Co-Authored-By: Claude Sonnet 4.6 * fix(fireworks): skip #transform=inline for base64 data URLs Closes #23583 * fix(fireworks): skip #transform=inline for base64 data URLs Closes #23583 * fix(fireworks): skip #transform=inline for base64 data URLs Closes #23583 --------- Co-authored-by: Claude Sonnet 4.6 Co-authored-by: Krish Dholakia --- .../llms/fireworks_ai/chat/transformation.py | 13 ++++++--- .../test_fireworks_ai_translation.py | 18 ++++++++++++ .../test_fireworks_ai_chat_transformation.py | 29 +++++++++++++++++++ 3 files changed, 56 insertions(+), 4 deletions(-) diff --git a/litellm/llms/fireworks_ai/chat/transformation.py b/litellm/llms/fireworks_ai/chat/transformation.py index 8407e8ab695..6b654ebdfd3 100644 --- a/litellm/llms/fireworks_ai/chat/transformation.py +++ b/litellm/llms/fireworks_ai/chat/transformation.py @@ -185,11 +185,16 @@ class FireworksAIConfig(OpenAIGPTConfig): ): # allow user to toggle this feature. return content if isinstance(content["image_url"], str): - content["image_url"] = f"{content['image_url']}#transform=inline" + # Skip base64 data URLs — appending #transform=inline corrupts the + # base64 payload and causes an "Incorrect padding" decode error on + # the Fireworks side. Data URLs are already inlined by definition. + # Lower-case before checking: URI schemes are case-insensitive (RFC 3986). + if not content["image_url"].lower().startswith("data:"): + content["image_url"] = f"{content['image_url']}#transform=inline" elif isinstance(content["image_url"], dict): - content["image_url"][ - "url" - ] = f"{content['image_url']['url']}#transform=inline" + url = content["image_url"]["url"] + if not url.lower().startswith("data:"): + content["image_url"]["url"] = f"{url}#transform=inline" return content def _transform_tools( diff --git a/tests/llm_translation/test_fireworks_ai_translation.py b/tests/llm_translation/test_fireworks_ai_translation.py index b9abbd501d0..24c0d546e25 100644 --- a/tests/llm_translation/test_fireworks_ai_translation.py +++ b/tests/llm_translation/test_fireworks_ai_translation.py @@ -161,6 +161,24 @@ def test_document_inlining_example(disable_add_transform_inline_image_block): "vision-gpt", "http://example.com/image.png", ), + # data: URLs must never have #transform=inline appended — doing so + # corrupts the base64 payload (fixes #23583). + # URI schemes are case-insensitive (RFC 3986) so check all variants. + ( + {"image_url": "data:image/png;base64,iVBORw0KGgo="}, + "gpt-4", + "data:image/png;base64,iVBORw0KGgo=", + ), + ( + {"image_url": {"url": "data:image/jpeg;base64,/9j/4AAQ=="}}, + "gpt-4", + {"url": "data:image/jpeg;base64,/9j/4AAQ=="}, + ), + ( + {"image_url": "Data:image/png;base64,iVBORw0KGgo="}, + "gpt-4", + "Data:image/png;base64,iVBORw0KGgo=", + ), ], ) def test_transform_inline(content, model, expected_url): diff --git a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py index 5d5aaa64c8e..2b71b88356b 100644 --- a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py +++ b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py @@ -110,6 +110,35 @@ def test_get_supported_openai_params_reasoning_effort(): assert "reasoning_effort" not in unsupported_params +def test_add_transform_inline_image_block_skips_data_urls(): + """ + data: URLs must not have #transform=inline appended — doing so corrupts the + base64 payload and raises binascii.Error: Incorrect padding on the Fireworks side. + Regression test for https://github.com/BerriAI/litellm/issues/23583 + """ + config = FireworksAIConfig() + data_url = "data:image/jpeg;base64,/9j/4AAQSkZJRgAB" + + # str branch + str_content = {"type": "image_url", "image_url": data_url} + result = config._add_transform_inline_image_block( + str_content, model="non-vision-model", disable_add_transform_inline_image_block=False + ) + assert result["image_url"] == data_url, "data URL must not be modified (str branch)" + + # dict branch + dict_content = {"type": "image_url", "image_url": {"url": data_url}} + result = config._add_transform_inline_image_block( + dict_content, model="non-vision-model", disable_add_transform_inline_image_block=False + ) + assert result["image_url"]["url"] == data_url, "data URL must not be modified (dict branch)" + + # regular https URL should still get the suffix + https_content = {"type": "image_url", "image_url": "https://example.com/image.jpg"} + result = config._add_transform_inline_image_block( + https_content, model="non-vision-model", disable_add_transform_inline_image_block=False + ) + assert result["image_url"].endswith("#transform=inline"), "https URL should get #transform=inline" @pytest.mark.parametrize( "api_base, expected_url_prefix", [ From e9291a97c32303a2ca18313bede88a75a52ce258 Mon Sep 17 00:00:00 2001 From: Miguel Miranda Dias <7780875+pandego@users.noreply.github.com> Date: Tue, 17 Mar 2026 06:34:15 +0100 Subject: [PATCH 07/25] fix(langsmith): avoid no running event loop during sync init (#23727) * fix(langsmith): skip periodic flush task without event loop * fix(langsmith): lazily start periodic flush task * test(langsmith): tighten flush task coverage * test(langsmith): cover lazy failure flush startup * refactor(langsmith): keep flush startup private --- litellm/integrations/langsmith.py | 43 ++++-- .../integrations/test_langsmith_init.py | 125 ++++++++++++++---- 2 files changed, 129 insertions(+), 39 deletions(-) diff --git a/litellm/integrations/langsmith.py b/litellm/integrations/langsmith.py index 03845af521d..df2e3c1e2b7 100644 --- a/litellm/integrations/langsmith.py +++ b/litellm/integrations/langsmith.py @@ -83,7 +83,26 @@ class LangsmithLogger(CustomBatchLogger): if _batch_size: self.batch_size = int(_batch_size) self.log_queue: List[LangsmithQueueObject] = [] - asyncio.create_task(self.periodic_flush()) + self._flush_task: Optional[asyncio.Task[Any]] = self._start_periodic_flush_task() + + def _start_periodic_flush_task(self) -> Optional[asyncio.Task[Any]]: + """Start the periodic flush task only when an event loop is already running.""" + try: + loop = asyncio.get_running_loop() + except RuntimeError: + verbose_logger.debug( + "Langsmith logger init: no running event loop, skipping periodic flush task startup" + ) + return None + + return loop.create_task(self.periodic_flush()) + + def _ensure_periodic_flush_task(self) -> None: + # This helper is intentionally synchronous. In asyncio's cooperative + # execution model, there is no await between the check and assignment, + # so one caller cannot interleave here and create a duplicate task. + if self._flush_task is None or self._flush_task.done(): + self._flush_task = self._start_periodic_flush_task() def get_credentials_from_env( self, @@ -255,6 +274,7 @@ class LangsmithLogger(CustomBatchLogger): async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): try: + self._ensure_periodic_flush_task() sampling_rate = self._get_sampling_rate_to_use_for_request(kwargs=kwargs) random_sample = random.random() if random_sample > sampling_rate: @@ -296,17 +316,18 @@ class LangsmithLogger(CustomBatchLogger): ) async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time): - sampling_rate = self._get_sampling_rate_to_use_for_request(kwargs=kwargs) - random_sample = random.random() - if random_sample > sampling_rate: - verbose_logger.info( - "Skipping Langsmith logging. Sampling rate={}, random_sample={}".format( - sampling_rate, random_sample - ) - ) - return # Skip logging - verbose_logger.info("Langsmith Failure Event Logging!") try: + self._ensure_periodic_flush_task() + sampling_rate = self._get_sampling_rate_to_use_for_request(kwargs=kwargs) + random_sample = random.random() + if random_sample > sampling_rate: + verbose_logger.info( + "Skipping Langsmith logging. Sampling rate={}, random_sample={}".format( + sampling_rate, random_sample + ) + ) + return # Skip logging + verbose_logger.info("Langsmith Failure Event Logging!") credentials = self._get_credentials_to_use_for_request(kwargs=kwargs) data = self._prepare_log_data( kwargs=kwargs, diff --git a/tests/test_litellm/integrations/test_langsmith_init.py b/tests/test_litellm/integrations/test_langsmith_init.py index 9f7db4095bc..edc827033f0 100644 --- a/tests/test_litellm/integrations/test_langsmith_init.py +++ b/tests/test_litellm/integrations/test_langsmith_init.py @@ -16,13 +16,9 @@ class TestLangsmithLoggerInit: Note: The current implementation has some edge cases in the sampling rate logic. """ - @patch("asyncio.create_task") @patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "1"}, clear=False) - def test_langsmith_sampling_rate_parameter_respected_with_valid_env( - self, mock_create_task - ): + def test_langsmith_sampling_rate_parameter_respected_with_valid_env(self): """Test that langsmith_sampling_rate parameter is properly set when env var condition is met.""" - # When there's a valid integer in env var, the parameter should be used due to 'or' logic sampling_rate = 0.5 logger = LangsmithLogger( langsmith_api_key="test-key", @@ -30,58 +26,47 @@ class TestLangsmithLoggerInit: langsmith_sampling_rate=sampling_rate, ) - # With the current 'or' logic and valid env var, the parameter should be used assert ( logger.sampling_rate == sampling_rate ), f"Expected sampling_rate to be {sampling_rate}, got {logger.sampling_rate}" - @patch("asyncio.create_task") @patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "1"}, clear=False) - def test_langsmith_sampling_rate_zero_parameter_falls_back_to_env( - self, mock_create_task - ): + def test_langsmith_sampling_rate_zero_parameter_falls_back_to_env(self): """Test that 0.0 parameter falls back to env var due to falsy value.""" - # This demonstrates the current behavior where 0.0 is falsy and falls back to env logger = LangsmithLogger( langsmith_api_key="test-key", langsmith_project="test-project", - langsmith_sampling_rate=0.0, # This is falsy! + langsmith_sampling_rate=0.0, ) - # Due to current 'or' logic, 0.0 falls back to env var assert ( logger.sampling_rate == 1.0 ), f"Expected sampling_rate to fall back to 1.0 from env, got {logger.sampling_rate}" - @patch("asyncio.create_task") @patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "1"}, clear=False) - def test_langsmith_sampling_rate_from_integer_env_var(self, mock_create_task): + def test_langsmith_sampling_rate_from_integer_env_var(self): """Test that sampling rate uses environment variable when parameter not provided and env var is integer.""" logger = LangsmithLogger( langsmith_api_key="test-key", langsmith_project="test-project" ) - # Should use env var since it's a valid integer assert ( logger.sampling_rate == 1.0 ), f"Expected sampling_rate to be 1.0 from env var, got {logger.sampling_rate}" - @patch("asyncio.create_task") @patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "0.8"}, clear=False) - def test_langsmith_sampling_rate_decimal_env_var_ignored(self, mock_create_task): + def test_langsmith_sampling_rate_decimal_env_var_ignored(self): """Test that decimal environment variables are ignored due to isdigit() check.""" logger = LangsmithLogger( langsmith_api_key="test-key", langsmith_project="test-project" ) - # Decimal env vars are ignored due to isdigit() check, falls back to 1.0 assert ( logger.sampling_rate == 1.0 ), f"Expected sampling_rate to default to 1.0 (decimal env ignored), got {logger.sampling_rate}" - @patch("asyncio.create_task") @patch.dict(os.environ, {}, clear=True) - def test_langsmith_sampling_rate_default_value(self, mock_create_task): + def test_langsmith_sampling_rate_default_value(self): """Test that sampling rate defaults to 1.0 when no parameter or env var provided.""" logger = LangsmithLogger( langsmith_api_key="test-key", langsmith_project="test-project" @@ -91,9 +76,8 @@ class TestLangsmithLoggerInit: logger.sampling_rate == 1.0 ), f"Expected default sampling_rate to be 1.0, got {logger.sampling_rate}" - @patch("asyncio.create_task") @patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": "invalid"}, clear=False) - def test_langsmith_sampling_rate_invalid_env_var_defaults(self, mock_create_task): + def test_langsmith_sampling_rate_invalid_env_var_defaults(self): """Test that invalid environment variable falls back to default value.""" logger = LangsmithLogger( langsmith_api_key="test-key", langsmith_project="test-project" @@ -103,9 +87,8 @@ class TestLangsmithLoggerInit: logger.sampling_rate == 1.0 ), f"Expected sampling_rate to default to 1.0 with invalid env var, got {logger.sampling_rate}" - @patch("asyncio.create_task") @patch.dict(os.environ, {"LANGSMITH_SAMPLING_RATE": ""}, clear=False) - def test_langsmith_sampling_rate_empty_env_var_defaults(self, mock_create_task): + def test_langsmith_sampling_rate_empty_env_var_defaults(self): """Test that empty environment variable falls back to default value.""" logger = LangsmithLogger( langsmith_api_key="test-key", langsmith_project="test-project" @@ -115,14 +98,12 @@ class TestLangsmithLoggerInit: logger.sampling_rate == 1.0 ), f"Expected sampling_rate to default to 1.0 with empty env var, got {logger.sampling_rate}" - @patch("asyncio.create_task") - def test_langsmith_sampling_rate_attribute_exists(self, mock_create_task): + def test_langsmith_sampling_rate_attribute_exists(self): """Test that the sampling_rate attribute is always set on the logger instance.""" logger = LangsmithLogger( langsmith_api_key="test-key", langsmith_project="test-project" ) - # Verify the attribute exists and is a float assert hasattr( logger, "sampling_rate" ), "LangsmithLogger should have sampling_rate attribute" @@ -132,3 +113,91 @@ class TestLangsmithLoggerInit: assert ( logger.sampling_rate >= 0.0 ), f"sampling_rate should be non-negative, got {logger.sampling_rate}" + + @patch.object(LangsmithLogger, "_start_periodic_flush_task", return_value=None) + def test_langsmith_init_skips_periodic_flush_without_running_loop( + self, mock_start_periodic_flush_task + ): + """Test that sync initialization leaves the periodic flush task unset.""" + logger = LangsmithLogger( + langsmith_api_key="test-key", langsmith_project="test-project" + ) + + assert logger is not None + mock_start_periodic_flush_task.assert_called_once() + assert logger._flush_task is None + + @patch("asyncio.get_running_loop", side_effect=RuntimeError("no running event loop")) + def test_start_periodic_flush_task_returns_none_without_running_loop( + self, mock_get_running_loop + ): + """Test that helper returns None when no running event loop exists.""" + with patch.object(LangsmithLogger, "_start_periodic_flush_task", return_value=None): + logger = LangsmithLogger( + langsmith_api_key="test-key", + langsmith_project="test-project", + ) + + mock_get_running_loop.reset_mock() + + assert logger._start_periodic_flush_task() is None + mock_get_running_loop.assert_called_once() + + @patch("asyncio.get_running_loop") + def test_langsmith_init_starts_periodic_flush_with_running_loop( + self, mock_get_running_loop + ): + """Test that init schedules periodic flush when a running loop exists.""" + mock_loop = MagicMock() + mock_task = MagicMock() + mock_loop.create_task.return_value = mock_task + mock_get_running_loop.return_value = mock_loop + + logger = LangsmithLogger( + langsmith_api_key="test-key", langsmith_project="test-project" + ) + + assert logger._flush_task == mock_task + mock_loop.create_task.assert_called_once() + scheduled_coro = mock_loop.create_task.call_args.args[0] + scheduled_coro.close() + + @pytest.mark.asyncio + async def test_async_log_success_event_lazily_starts_periodic_flush(self): + """Test that async logging lazily starts periodic flush after sync init.""" + with patch.object(LangsmithLogger, "_start_periodic_flush_task", return_value=None): + logger = LangsmithLogger( + langsmith_api_key="test-key", + langsmith_project="test-project", + ) + logger._get_sampling_rate_to_use_for_request = MagicMock(return_value=1.0) + logger._get_credentials_to_use_for_request = MagicMock( + return_value=logger.default_credentials + ) + logger._prepare_log_data = MagicMock(return_value={"id": "run-id"}) + logger._start_periodic_flush_task = MagicMock(return_value=MagicMock()) + + await logger.async_log_success_event({}, {}, None, None) + + logger._start_periodic_flush_task.assert_called_once() + assert len(logger.log_queue) == 1 + + @pytest.mark.asyncio + async def test_async_log_failure_event_lazily_starts_periodic_flush(self): + """Test that async failure logging lazily starts periodic flush after sync init.""" + with patch.object(LangsmithLogger, "_start_periodic_flush_task", return_value=None): + logger = LangsmithLogger( + langsmith_api_key="test-key", + langsmith_project="test-project", + ) + logger._get_sampling_rate_to_use_for_request = MagicMock(return_value=1.0) + logger._get_credentials_to_use_for_request = MagicMock( + return_value=logger.default_credentials + ) + logger._prepare_log_data = MagicMock(return_value={"id": "run-id"}) + logger._start_periodic_flush_task = MagicMock(return_value=MagicMock()) + + await logger.async_log_failure_event({}, {}, None, None) + + logger._start_periodic_flush_task.assert_called_once() + assert len(logger.log_queue) == 1 From 186c2adb326050587f4577f09e7dedcaafc982d8 Mon Sep 17 00:00:00 2001 From: Awais Qureshi Date: Tue, 17 Mar 2026 10:38:16 +0500 Subject: [PATCH 08/25] fix(gemini): support images in tool_results for /v1/messages routing (#23724) * fix(gemini): support images in tool_results for /v1/messages routing convert_to_gemini_tool_call_result() dropped images in two cases: - data-URL strings (data:image/...;base64,...) treated as plain text - Anthropic image blocks in list content skipped Add detection and convert both to Gemini inline_data BlobType so image bytes are preserved. Fixes #23712. * fix(gemini): support images in tool_results for /v1/messages routing convert_to_gemini_tool_call_result() dropped images in two cases: - data-URL strings (data:image/...;base64,...) treated as plain text - Anthropic image blocks in list content skipped Add detection and convert both to Gemini inline_data BlobType so image bytes are preserved. Fixes #23712. * fix(gemini): support images in tool_results for /v1/messages routing convert_to_gemini_tool_call_result() dropped images in two cases: - data-URL strings (data:image/...;base64,...) treated as plain text - Anthropic image blocks in list content skipped Add detection and convert both to Gemini inline_data BlobType so image bytes are preserved. Fixes #23712. * fix(fireworks): skip #transform=inline for base64 data URLs Closes #23583 --- .../prompt_templates/factory.py | 59 ++++-- ...llm_core_utils_prompt_templates_factory.py | 172 ++++++++++++++++++ 2 files changed, 219 insertions(+), 12 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 47272b38ad6..6c4c98ebf9a 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1498,17 +1498,49 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915 from litellm.types.llms.vertex_ai import BlobType content_str: str = "" - inline_data: Optional[BlobType] = None + inline_data_list: List[BlobType] = [] if "content" in message: if isinstance(message["content"], str): content_str = message["content"] + # Detect data-URL images (e.g. from Anthropic tool_result with a single image block + # that was serialised as a plain string by translate_anthropic_messages_to_openai) + # and promote them to inline_data so Gemini receives actual image bytes. + if content_str.startswith("data:") and ";base64," in content_str: + try: + mime_rest = content_str[5:].split(";base64,", 1) + if len(mime_rest) == 2 and mime_rest[0].startswith("image/"): + # Strip any extra parameters (e.g. ";charset=UTF-8") from the MIME segment + clean_mime = mime_rest[0].split(";")[0].strip() + inline_data_list.append( + BlobType(data=mime_rest[1], mime_type=clean_mime) + ) + content_str = "" + except Exception as e: + verbose_logger.warning( + f"Failed to parse data URL in tool response: {e}" + ) elif isinstance(message["content"], List): content_list = message["content"] for content in content_list: content_type = content.get("type", "") if content_type == "text": content_str += content.get("text", "") + elif content_type == "image": + # Anthropic-native image block: {"type": "image", "source": {"type": "base64", ...}} + source = content.get("source", {}) + if isinstance(source, dict) and source.get("type") == "base64": + try: + inline_data_list.append( + BlobType( + data=source.get("data", ""), + mime_type=source.get("media_type", "image/jpeg"), + ) + ) + except Exception as e: + verbose_logger.warning( + f"Failed to process Anthropic image block in tool response: {e}" + ) elif content_type in ("input_image", "image_url"): # Extract image for inline_data (for Computer Use screenshots and tool results) image_url_data = content.get("image_url", "") @@ -1524,9 +1556,11 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915 image_obj = convert_to_anthropic_image_obj( image_url, format=None ) - inline_data = BlobType( - data=image_obj["data"], - mime_type=image_obj["media_type"], + inline_data_list.append( + BlobType( + data=image_obj["data"], + mime_type=image_obj["media_type"], + ) ) except Exception as e: verbose_logger.warning( @@ -1551,9 +1585,11 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915 file_obj = convert_to_anthropic_image_obj( file_data, format=None ) - inline_data = BlobType( - data=file_obj["data"], - mime_type=file_obj["media_type"], + inline_data_list.append( + BlobType( + data=file_obj["data"], + mime_type=file_obj["media_type"], + ) ) except Exception as e: verbose_logger.warning( @@ -1607,13 +1643,12 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915 # Create part with function_response, and optionally inline_data for images (Computer Use) _part: VertexPartType = {"function_response": _function_response} - # For Computer Use, if we have an image, we need separate parts: + # For Computer Use, if we have images/files, we need separate parts: # - One part with function_response - # - One part with inline_data + # - One part per inline_data item # Gemini's PartType is a oneof, so we can't have both in the same part - if inline_data: - image_part: VertexPartType = {"inline_data": inline_data} - return [_part, image_part] + if inline_data_list: + return [_part] + [{"inline_data": d} for d in inline_data_list] return _part diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 8d68539564c..5c5cd5bdc3e 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -1,3 +1,4 @@ +import base64 import json from unittest.mock import MagicMock, patch @@ -9,9 +10,11 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( BedrockConverseMessagesProcessor, BedrockImageProcessor, _convert_to_bedrock_tool_call_invoke, + convert_to_gemini_tool_call_result, ollama_pt, sanitize_messages_for_tool_calling, ) +from litellm.types.llms.openai import ChatCompletionToolMessage def test_ollama_pt_simple_messages(): @@ -550,6 +553,175 @@ def test_convert_gemini_tool_call_result_with_image_url(): assert isinstance(result2, list) and any("inline_data" in p for p in result2) +def test_convert_gemini_tool_call_result_with_anthropic_image_block(): + """ + Test that Anthropic-native image blocks in tool_result list content are + converted to Gemini inline_data instead of being silently dropped. + Fixes: https://github.com/BerriAI/litellm/issues/23712 + """ + tiny_png_b64 = base64.b64encode(b"PNG_PLACEHOLDER").decode() + + message = ChatCompletionToolMessage( + role="tool", + tool_call_id="call_123", + content=[ + { + "type": "image", + "source": { + "type": "base64", + "media_type": "image/png", + "data": tiny_png_b64, + }, + } + ], + ) + last_message_with_tool_calls = { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_123", + "type": "function", + "index": 0, + "function": {"name": "read_file", "arguments": "{}"}, + } + ], + } + + result = convert_to_gemini_tool_call_result( + message=message, + last_message_with_tool_calls=last_message_with_tool_calls, + ) + assert isinstance(result, list), "expected a list of parts" + inline_parts = [p for p in result if "inline_data" in p] + assert len(inline_parts) == 1, "expected exactly one inline_data part" + assert inline_parts[0]["inline_data"]["mime_type"] == "image/png" + assert inline_parts[0]["inline_data"]["data"] == tiny_png_b64 + + +def test_convert_gemini_tool_call_result_with_multiple_anthropic_image_blocks(): + """ + Test that multiple Anthropic-native image blocks in a single tool_result + are all preserved as separate inline_data parts instead of only the last + one being kept. + Fixes: https://github.com/BerriAI/litellm/issues/23712 + """ + png_b64 = base64.b64encode(b"PNG_PLACEHOLDER").decode() + jpeg_b64 = base64.b64encode(b"JPEG_PLACEHOLDER").decode() + + message = ChatCompletionToolMessage( + role="tool", + tool_call_id="call_multi", + content=[ + {"type": "text", "text": "here are two images"}, + { + "type": "image", + "source": {"type": "base64", "media_type": "image/png", "data": png_b64}, + }, + { + "type": "image", + "source": {"type": "base64", "media_type": "image/jpeg", "data": jpeg_b64}, + }, + ], + ) + last_message_with_tool_calls = { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_multi", + "type": "function", + "index": 0, + "function": {"name": "screenshot", "arguments": "{}"}, + } + ], + } + + result = convert_to_gemini_tool_call_result( + message=message, + last_message_with_tool_calls=last_message_with_tool_calls, + ) + assert isinstance(result, list), "expected a list of parts" + inline_parts = [p for p in result if "inline_data" in p] + assert len(inline_parts) == 2, f"expected 2 inline_data parts, got {len(inline_parts)}" + mime_types = {p["inline_data"]["mime_type"] for p in inline_parts} + assert mime_types == {"image/png", "image/jpeg"} + + +def test_convert_gemini_tool_call_result_with_data_url_string(): + """ + Test that a data-URL string in tool_result content is converted to + Gemini inline_data instead of being passed as plain text. + Fixes: https://github.com/BerriAI/litellm/issues/23712 + """ + tiny_png_b64 = base64.b64encode(b"PNG_PLACEHOLDER").decode() + + message = ChatCompletionToolMessage( + role="tool", + tool_call_id="call_456", + content=f"data:image/png;base64,{tiny_png_b64}", + ) + last_message_with_tool_calls = { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_456", + "type": "function", + "index": 0, + "function": {"name": "read_file", "arguments": "{}"}, + } + ], + } + + result = convert_to_gemini_tool_call_result( + message=message, + last_message_with_tool_calls=last_message_with_tool_calls, + ) + assert isinstance(result, list), "expected a list of parts" + inline_parts = [p for p in result if "inline_data" in p] + assert len(inline_parts) == 1, "data-URL image string was not converted to inline_data" + assert inline_parts[0]["inline_data"]["mime_type"] == "image/png" + assert inline_parts[0]["inline_data"]["data"] == tiny_png_b64 + + +def test_convert_gemini_tool_call_result_with_data_url_extra_params(): + """ + Test that a data-URL with extra MIME parameters (e.g. charset) produces + a clean mime_type without the extra parameters. + """ + tiny_png_b64 = base64.b64encode(b"PNG_PLACEHOLDER").decode() + + message = ChatCompletionToolMessage( + role="tool", + tool_call_id="call_extra", + content=f"data:image/png;charset=UTF-8;base64,{tiny_png_b64}", + ) + last_message_with_tool_calls = { + "role": "assistant", + "content": "", + "tool_calls": [ + { + "id": "call_extra", + "type": "function", + "index": 0, + "function": {"name": "read_file", "arguments": "{}"}, + } + ], + } + + result = convert_to_gemini_tool_call_result( + message=message, + last_message_with_tool_calls=last_message_with_tool_calls, + ) + assert isinstance(result, list), "expected a list of parts" + inline_parts = [p for p in result if "inline_data" in p] + assert len(inline_parts) == 1 + assert inline_parts[0]["inline_data"]["mime_type"] == "image/png", ( + f"expected clean 'image/png', got '{inline_parts[0]['inline_data']['mime_type']}'" + ) + + def test_bedrock_tools_unpack_defs(): """ Test that the unpack_defs method handles nested $ref inside anyOf items correctly From 24429227d31be7e94c02d84e2f480fe945be5a1c Mon Sep 17 00:00:00 2001 From: Chesars Date: Tue, 17 Mar 2026 11:52:29 -0300 Subject: [PATCH 09/25] fix(model-prices): correct supported_regions for Vertex AI DeepSeek models MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fixes #23859 - deepseek-v3.2-maas: us-west2 → global (per Google docs) - deepseek-v3.1-maas: us-west2 → us-central1 - deepseek-r1-0528-maas: add supported_regions: us-central1 - deepseek-ocr-maas: add supported_regions: us-central1 --- model_prices_and_context_window.json | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 9b1d81fee40..614e459a1c0 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -30431,7 +30431,7 @@ "output_cost_per_token": 5.4e-06, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models", "supported_regions": [ - "us-west2" + "us-central1" ], "supports_assistant_prefill": true, "supports_function_calling": true, @@ -30451,7 +30451,7 @@ "output_cost_per_token_batches": 8.4e-07, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models", "supported_regions": [ - "us-west2" + "global" ], "supports_assistant_prefill": true, "supports_function_calling": true, @@ -30472,7 +30472,10 @@ "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "supported_regions": [ + "us-central1" + ] }, "vertex_ai/gemini-2.5-flash-image": { "cache_read_input_token_cost": 3e-08, @@ -31092,7 +31095,10 @@ "input_cost_per_token": 3e-07, "output_cost_per_token": 1.2e-06, "ocr_cost_per_page": 0.0003, - "source": "https://cloud.google.com/vertex-ai/pricing" + "source": "https://cloud.google.com/vertex-ai/pricing", + "supported_regions": [ + "us-central1" + ] }, "vertex_ai/openai/gpt-oss-120b-maas": { "input_cost_per_token": 1.5e-07, From d39eac26832f0168be871a6d069b74d11e31b55d Mon Sep 17 00:00:00 2001 From: Chesars Date: Tue, 17 Mar 2026 12:12:23 -0300 Subject: [PATCH 10/25] fix: move supported_regions before supports_* fields for alphabetical order --- model_prices_and_context_window.json | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 614e459a1c0..091a1f57b08 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -30468,14 +30468,14 @@ "mode": "chat", "output_cost_per_token": 5.4e-06, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models", + "supported_regions": [ + "us-central1" + ], "supports_assistant_prefill": true, "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, - "supports_tool_choice": true, - "supported_regions": [ - "us-central1" - ] + "supports_tool_choice": true }, "vertex_ai/gemini-2.5-flash-image": { "cache_read_input_token_cost": 3e-08, From 3eeb14bf1a17f52382d0e82b58979147e99fd89d Mon Sep 17 00:00:00 2001 From: cohml <62400541+cohml@users.noreply.github.com> Date: Tue, 17 Mar 2026 11:32:01 -0400 Subject: [PATCH 11/25] fix(cache): Fix Redis cluster caching (#23480) * fix redis cluster startup_nodes check order * add tests for redis cluster startup_nodes fix --- litellm/_redis.py | 63 ++++++++++------ tests/test_litellm/test_redis.py | 122 ++++++++++++++++++++++++++++++- 2 files changed, 160 insertions(+), 25 deletions(-) diff --git a/litellm/_redis.py b/litellm/_redis.py index b754c1f4330..2bf32d71b21 100644 --- a/litellm/_redis.py +++ b/litellm/_redis.py @@ -222,8 +222,12 @@ def _get_redis_client_logic(**env_overrides): "REDIS_CLUSTER_NODES" ) + # If startup_nodes resolved to None (not set by kwarg or env), remove the key + # entirely so callers can rely on key presence as a reliable cluster-mode signal. if _startup_nodes is not None and isinstance(_startup_nodes, str): redis_kwargs["startup_nodes"] = json.loads(_startup_nodes) + elif _startup_nodes is None: + redis_kwargs.pop("startup_nodes", None) _sentinel_nodes: Optional[Union[str, list]] = redis_kwargs.get("sentinel_nodes", None) or get_secret( # type: ignore "REDIS_SENTINEL_NODES" @@ -273,10 +277,14 @@ def _get_redis_client_logic(**env_overrides): redis_kwargs["ssl_ca_certs"] = _gcp_ssl_ca_certs if "url" in redis_kwargs and redis_kwargs["url"] is not None: - redis_kwargs.pop("host", None) - redis_kwargs.pop("port", None) - redis_kwargs.pop("db", None) - redis_kwargs.pop("password", None) + # Only strip host/port/db/password when not routing to a cluster. + # When startup_nodes is also present the cluster path takes priority and + # needs the password for authentication. + if not redis_kwargs.get("startup_nodes"): + redis_kwargs.pop("host", None) + redis_kwargs.pop("port", None) + redis_kwargs.pop("db", None) + redis_kwargs.pop("password", None) elif "startup_nodes" in redis_kwargs and redis_kwargs["startup_nodes"] is not None: pass elif ( @@ -368,6 +376,10 @@ def _init_async_redis_sentinel(redis_kwargs) -> async_redis.Redis: def get_redis_client(**env_overrides): redis_kwargs = _get_redis_client_logic(**env_overrides) + + if "startup_nodes" in redis_kwargs: + return init_redis_cluster(redis_kwargs) + if "url" in redis_kwargs and redis_kwargs["url"] is not None: args = _get_redis_url_kwargs() url_kwargs = {} @@ -377,9 +389,6 @@ def get_redis_client(**env_overrides): return redis.Redis.from_url(**url_kwargs) - if "startup_nodes" in redis_kwargs or get_secret("REDIS_CLUSTER_NODES") is not None: # type: ignore - return init_redis_cluster(redis_kwargs) - # Check for Redis Sentinel if "sentinel_nodes" in redis_kwargs and "service_name" in redis_kwargs: return _init_redis_sentinel(redis_kwargs) @@ -392,21 +401,6 @@ def get_redis_async_client( **env_overrides, ) -> Union[async_redis.Redis, async_redis.RedisCluster]: redis_kwargs = _get_redis_client_logic(**env_overrides) - if "url" in redis_kwargs and redis_kwargs["url"] is not None: - if connection_pool is not None: - return async_redis.Redis(connection_pool=connection_pool) - args = _get_redis_url_kwargs(client=async_redis.Redis.from_url) - url_kwargs = {} - for arg in redis_kwargs: - if arg in args: - url_kwargs[arg] = redis_kwargs[arg] - else: - verbose_logger.debug( - "REDIS: ignoring argument: {}. Not an allowed async_redis.Redis.from_url arg.".format( - arg - ) - ) - return async_redis.Redis.from_url(**url_kwargs) if "startup_nodes" in redis_kwargs: from redis.cluster import ClusterNode @@ -469,6 +463,22 @@ def get_redis_async_client( return cluster_client + if "url" in redis_kwargs and redis_kwargs["url"] is not None: + if connection_pool is not None: + return async_redis.Redis(connection_pool=connection_pool) + args = _get_redis_url_kwargs(client=async_redis.Redis.from_url) + url_kwargs = {} + for arg in redis_kwargs: + if arg in args: + url_kwargs[arg] = redis_kwargs[arg] + else: + verbose_logger.debug( + "REDIS: ignoring argument: {}. Not an allowed async_redis.Redis.from_url arg.".format( + arg + ) + ) + return async_redis.Redis.from_url(**url_kwargs) + # Check for Redis Sentinel if "sentinel_nodes" in redis_kwargs and "service_name" in redis_kwargs: return _init_async_redis_sentinel(redis_kwargs) @@ -482,9 +492,15 @@ def get_redis_async_client( ) -def get_redis_connection_pool(**env_overrides): +def get_redis_connection_pool( + **env_overrides, +) -> Optional[async_redis.BlockingConnectionPool]: redis_kwargs = _get_redis_client_logic(**env_overrides) verbose_logger.debug("get_redis_connection_pool: redis_kwargs", redis_kwargs) + + if "startup_nodes" in redis_kwargs: + return None + if "url" in redis_kwargs and redis_kwargs["url"] is not None: pool_kwargs = { "timeout": REDIS_CONNECTION_POOL_TIMEOUT, @@ -504,7 +520,6 @@ def get_redis_connection_pool(**env_overrides): connection_class = async_redis.SSLConnection redis_kwargs.pop("ssl", None) redis_kwargs["connection_class"] = connection_class - redis_kwargs.pop("startup_nodes", None) return async_redis.BlockingConnectionPool( timeout=REDIS_CONNECTION_POOL_TIMEOUT, **redis_kwargs ) diff --git a/tests/test_litellm/test_redis.py b/tests/test_litellm/test_redis.py index 4709faea4bc..15907190998 100644 --- a/tests/test_litellm/test_redis.py +++ b/tests/test_litellm/test_redis.py @@ -1,7 +1,15 @@ -from litellm._redis import get_redis_url_from_environment, _get_redis_cluster_kwargs, get_redis_async_client +from litellm._redis import ( + get_redis_url_from_environment, + _get_redis_cluster_kwargs, + get_redis_async_client, + get_redis_client, + get_redis_connection_pool, +) +import json import os import pytest from unittest.mock import MagicMock, patch +import redis import redis.asyncio as async_redis def test_get_redis_url_from_environment_single_url(monkeypatch): @@ -167,3 +175,115 @@ def test_get_redis_async_client_without_connection_pool(): # Verify Redis was called without connection_pool in kwargs call_kwargs = mock_redis.call_args[1] assert "connection_pool" not in call_kwargs, "connection_pool should not be in kwargs when not provided" + +@patch("litellm._redis.init_redis_cluster") +def test_sync_client_prefers_cluster_over_url(mock_init_cluster, monkeypatch): + """ + Test get_redis_client returns RedisCluster when startup_nodes is present even if + REDIS_URL is also set. + """ + monkeypatch.setenv("REDIS_URL", "redis://fallback-host:6379") + mock_init_cluster.return_value = MagicMock(spec=redis.RedisCluster) + + startup_nodes = [{"host": "cluster-node.example.com", "port": 6379}] + get_redis_client(startup_nodes=startup_nodes) + + mock_init_cluster.assert_called_once() + call_kwargs = mock_init_cluster.call_args[0][0] + assert ( + "startup_nodes" in call_kwargs + ), "startup_nodes must be forwarded to init_redis_cluster" + +@patch("litellm._redis.async_redis.RedisCluster") +def test_async_client_prefers_cluster_over_url(mock_cluster_cls, monkeypatch): + """ + Test (1) get_redis_async_client returns async RedisCluster when startup_nodes is present + even if REDIS_URL is also set and (2) startup_nodes is forwarded to RedisCluster. + """ + monkeypatch.setenv("REDIS_URL", "redis://fallback-host:6379") + + startup_nodes = [{"host": "cluster-node.example.com", "port": 6379}] + get_redis_async_client(startup_nodes=startup_nodes) + + mock_cluster_cls.assert_called_once() + call_kwargs = mock_cluster_cls.call_args[1] + assert "startup_nodes" in call_kwargs, "startup_nodes must be forwarded to async RedisCluster" + assert len(call_kwargs["startup_nodes"]) == 1, "should forward exactly 1 cluster node" + + +@patch("litellm._redis.async_redis.RedisCluster") +def test_async_client_prefers_cluster_over_url_via_env_var(mock_cluster_cls, monkeypatch): + """ + Test get_redis_async_client returns async RedisCluster when REDIS_CLUSTER_NODES is set + even if REDIS_URL is also set. + """ + monkeypatch.setenv("REDIS_URL", "redis://fallback-host:6379") + monkeypatch.setenv( + "REDIS_CLUSTER_NODES", + json.dumps([{"host": "cluster-node.example.com", "port": 6379}]), + ) + + get_redis_async_client() + + mock_cluster_cls.assert_called_once() + call_kwargs = mock_cluster_cls.call_args[1] + assert "startup_nodes" in call_kwargs, "startup_nodes must be forwarded to async RedisCluster" + +@patch("litellm._redis.init_redis_cluster") +def test_sync_client_prefers_cluster_over_url_via_env_var(mock_init_cluster, monkeypatch): + """ + Test get_redis_client returns RedisCluster when REDIS_CLUSTER_NODES is set even if + REDIS_URL is also set. + """ + monkeypatch.setenv("REDIS_URL", "redis://fallback-host:6379") + monkeypatch.setenv( + "REDIS_CLUSTER_NODES", + json.dumps([{"host": "cluster-node.example.com", "port": 6379}]), + ) + mock_init_cluster.return_value = MagicMock(spec=redis.RedisCluster) + + get_redis_client() + + mock_init_cluster.assert_called_once() + call_kwargs = mock_init_cluster.call_args[0][0] + assert "startup_nodes" in call_kwargs, "startup_nodes must be forwarded to init_redis_cluster" + assert len(call_kwargs["startup_nodes"]) == 1 + +@patch("litellm._redis.init_redis_cluster") +def test_sync_client_preserves_password_for_cluster_when_url_also_set(mock_init_cluster, monkeypatch): + """ + Test _get_redis_client_logic does not strip password from redis_kwargs when + startup_nodes is present even if REDIS_URL is also set. + """ + monkeypatch.setenv("REDIS_URL", "redis://fallback-host:6379") + monkeypatch.setenv("REDIS_PASSWORD", "secret") + mock_init_cluster.return_value = MagicMock(spec=redis.RedisCluster) + + startup_nodes = [{"host": "cluster-node.example.com", "port": 6379}] + get_redis_client(startup_nodes=startup_nodes) + + mock_init_cluster.assert_called_once() + call_kwargs = mock_init_cluster.call_args[0][0] + assert "password" in call_kwargs, "password must not be stripped when routing to cluster" + assert call_kwargs["password"] == "secret" + + +def test_connection_pool_returns_none_for_cluster(monkeypatch): + """Test get_redis_connection_pool returns None when startup_nodes is present.""" + monkeypatch.setenv("REDIS_URL", "redis://fallback-host:6379") + startup_nodes = [{"host": "cluster-node.example.com", "port": 6379}] + result = get_redis_connection_pool(startup_nodes=startup_nodes) + assert result is None, "connection pool must be None for cluster mode" + + +@patch("litellm._redis.redis.Redis.from_url") +def test_sync_client_url_used_when_no_cluster(mock_from_url, monkeypatch): + """ + Test get_redis_client default to using URL path when no startup_nodes are provided. + """ + monkeypatch.setenv("REDIS_URL", "redis://plain-host:6379") + monkeypatch.delenv("REDIS_CLUSTER_NODES", raising=False) + + get_redis_client() + + mock_from_url.assert_called_once() From b0db75df1fb9f5871cdea662035ee719ef0a9149 Mon Sep 17 00:00:00 2001 From: rstar327 Date: Tue, 17 Mar 2026 13:35:07 -0400 Subject: [PATCH 12/25] fix(proxy): convert max_budget to float when set from environment variable (#23855) Fixes #23843 --- litellm/proxy/proxy_server.py | 68 +++++++++---------- .../proxy/test_max_budget_env_var.py | 38 +++++++++++ 2 files changed, 71 insertions(+), 35 deletions(-) create mode 100644 tests/test_litellm/proxy/test_max_budget_env_var.py diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 9c29927c5cb..a775b12b206 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -638,9 +638,9 @@ except ImportError: server_root_path = get_server_root_path() _license_check = LicenseCheck() premium_user: bool = _license_check.is_premium() -premium_user_data: Optional[ - "EnterpriseLicenseData" -] = _license_check.airgapped_license_data +premium_user_data: Optional["EnterpriseLicenseData"] = ( + _license_check.airgapped_license_data +) global_max_parallel_request_retries_env: Optional[str] = os.getenv( "LITELLM_GLOBAL_MAX_PARALLEL_REQUEST_RETRIES" ) @@ -1523,9 +1523,9 @@ master_key: Optional[str] = None config_agents: Optional[List[AgentConfig]] = None otel_logging = False prisma_client: Optional[PrismaClient] = None -shared_aiohttp_session: Optional[ - "ClientSession" -] = None # Global shared session for connection reuse +shared_aiohttp_session: Optional["ClientSession"] = ( + None # Global shared session for connection reuse +) user_api_key_cache = DualCache( default_in_memory_ttl=UserAPIKeyCacheTTLEnum.in_memory_cache_ttl.value ) @@ -1533,13 +1533,13 @@ model_max_budget_limiter = _PROXY_VirtualKeyModelMaxBudgetLimiter( dual_cache=user_api_key_cache ) litellm.logging_callback_manager.add_litellm_callback(model_max_budget_limiter) -redis_usage_cache: Optional[ - RedisCache -] = None # redis cache used for tracking spend, tpm/rpm limits +redis_usage_cache: Optional[RedisCache] = ( + None # redis cache used for tracking spend, tpm/rpm limits +) polling_via_cache_enabled: Union[Literal["all"], List[str], bool] = False -native_background_mode: List[ - str -] = [] # Models that should use native provider background mode instead of polling +native_background_mode: List[str] = ( + [] +) # Models that should use native provider background mode instead of polling polling_cache_ttl: int = 3600 # Default 1 hour TTL for polling cache user_custom_auth = None user_custom_key_generate = None @@ -1898,9 +1898,9 @@ async def update_cache( # noqa: PLR0915 _id = "team_id:{}".format(team_id) try: # Fetch the existing cost for the given user - existing_spend_obj: Optional[ - LiteLLM_TeamTable - ] = await user_api_key_cache.async_get_cache(key=_id) + existing_spend_obj: Optional[LiteLLM_TeamTable] = ( + await user_api_key_cache.async_get_cache(key=_id) + ) if existing_spend_obj is None: # do nothing if team not in api key cache return @@ -2021,11 +2021,9 @@ def run_ollama_serve(): with open(os.devnull, "w") as devnull: subprocess.Popen(command, stdout=devnull, stderr=devnull) except Exception as e: - verbose_proxy_logger.debug( - f""" + verbose_proxy_logger.debug(f""" LiteLLM Warning: proxy started with `ollama` model\n`ollama serve` failed with Exception{e}. \nEnsure you run `ollama serve` - """ - ) + """) def _get_process_rss_mb() -> Optional[float]: @@ -3303,7 +3301,7 @@ class ProxyConfig: async_only_mode=True # only init async clients ), ignore_invalid_deployments=True, # don't raise an error if a deployment is invalid - ) # type:ignore + ) # type: ignore if redis_usage_cache is not None and router.cache.redis_cache is None: router._update_redis_cache(cache=redis_usage_cache) @@ -4952,10 +4950,10 @@ class ProxyConfig: ) try: - guardrails_in_db: List[ - Guardrail - ] = await GuardrailRegistry.get_all_guardrails_from_db( - prisma_client=prisma_client + guardrails_in_db: List[Guardrail] = ( + await GuardrailRegistry.get_all_guardrails_from_db( + prisma_client=prisma_client + ) ) verbose_proxy_logger.debug( "guardrails from the DB %s", str(guardrails_in_db) @@ -5337,9 +5335,9 @@ async def initialize( # noqa: PLR0915 user_api_base = api_base dynamic_config[user_model]["api_base"] = api_base if api_version: - os.environ[ - "AZURE_API_VERSION" - ] = api_version # set this for azure - litellm can read this from the env + os.environ["AZURE_API_VERSION"] = ( + api_version # set this for azure - litellm can read this from the env + ) if max_tokens: # model-specific param dynamic_config[user_model]["max_tokens"] = max_tokens if temperature: # model-specific param @@ -5357,8 +5355,8 @@ async def initialize( # noqa: PLR0915 litellm.add_function_to_prompt = True dynamic_config["general"]["add_function_to_prompt"] = True if max_budget: # litellm-specific param - litellm.max_budget = max_budget - dynamic_config["general"]["max_budget"] = max_budget + litellm.max_budget = float(max_budget) + dynamic_config["general"]["max_budget"] = litellm.max_budget if experimental: pass user_telemetry = telemetry @@ -5676,9 +5674,9 @@ class ProxyStartupEvent: """ from litellm.secret_managers.main import str_to_bool - _use_redis_transaction_buffer: Optional[ - Union[bool, str] - ] = general_settings.get("use_redis_transaction_buffer", False) + _use_redis_transaction_buffer: Optional[Union[bool, str]] = ( + general_settings.get("use_redis_transaction_buffer", False) + ) if isinstance(_use_redis_transaction_buffer, str): _use_redis_transaction_buffer = str_to_bool(_use_redis_transaction_buffer) @@ -12114,9 +12112,9 @@ async def get_config_list( hasattr(sub_field_info, "description") and sub_field_info.description is not None ): - nested_fields[ - idx - ].field_description = sub_field_info.description + nested_fields[idx].field_description = ( + sub_field_info.description + ) idx += 1 _stored_in_db = None diff --git a/tests/test_litellm/proxy/test_max_budget_env_var.py b/tests/test_litellm/proxy/test_max_budget_env_var.py new file mode 100644 index 00000000000..4b965392823 --- /dev/null +++ b/tests/test_litellm/proxy/test_max_budget_env_var.py @@ -0,0 +1,38 @@ +""" +Test that max_budget from environment variable (string) is correctly +converted to float. +GitHub Issue: #23843 +""" + +import pytest + +import litellm +from litellm.proxy.proxy_server import initialize + + +@pytest.mark.asyncio +async def test_max_budget_string_converted_to_float(): + """ + When max_budget is set via os.environ/MAX_BUDGET, it arrives as a + string. initialize() should convert it to float so the comparison + `litellm.max_budget > 0` doesn't raise TypeError. + """ + original = litellm.max_budget + try: + await initialize(max_budget="100.5") + assert isinstance(litellm.max_budget, float) + assert litellm.max_budget == 100.5 + finally: + litellm.max_budget = original + + +@pytest.mark.asyncio +async def test_max_budget_float_stays_float(): + """max_budget as float should still work.""" + original = litellm.max_budget + try: + await initialize(max_budget=200.0) + assert isinstance(litellm.max_budget, float) + assert litellm.max_budget == 200.0 + finally: + litellm.max_budget = original From 0c28b47057102720353c04fbafd11f64760eb196 Mon Sep 17 00:00:00 2001 From: Chesars Date: Tue, 17 Mar 2026 17:47:01 -0300 Subject: [PATCH 13/25] fix(vertex): streaming finish_reason="stop" instead of "tool_calls" for gemini-3.1-flash-lite-preview Models like gemini-3.1-flash-lite-preview send the final streaming chunk with empty content (text:"") alongside finishReason:"STOP", instead of omitting content entirely. The existing fix (PR #21577) only handled chunks without content, so this case was missed. Now, after processing candidates, if tool_calls were seen in earlier chunks and a choice has finish_reason="stop", it is overridden to "tool_calls" to match the OpenAI spec. Fixes #22900 --- .../vertex_and_google_ai_studio_gemini.py | 10 +++ ...emini_streaming_tool_call_finish_reason.py | 72 +++++++++++++++++++ 2 files changed, 82 insertions(+) diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 3f1bccaccfc..1054b311d01 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -3001,6 +3001,16 @@ class ModelResponseIterator: ) model_response.choices.append(choice) + # Also handle the case where the final chunk has empty + # content (e.g. text:"") WITH finishReason. In this case + # _process_candidates DOES create a choice, but maps + # finishReason="STOP" to "stop" because the current chunk + # has no tool_calls. Override if we saw tool_calls earlier. + if self.has_seen_tool_calls: + for choice in model_response.choices: + if choice.finish_reason == "stop": + choice.finish_reason = "tool_calls" + setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) # type: ignore setattr(model_response, "vertex_ai_url_context_metadata", url_context_metadata) # type: ignore setattr(model_response, "vertex_ai_safety_ratings", safety_ratings) # type: ignore diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_gemini_streaming_tool_call_finish_reason.py b/tests/test_litellm/llms/vertex_ai/gemini/test_gemini_streaming_tool_call_finish_reason.py index 3f8efd47fa3..d4d76ab3079 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_gemini_streaming_tool_call_finish_reason.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_gemini_streaming_tool_call_finish_reason.py @@ -230,3 +230,75 @@ def test_streaming_content_filter_finish_reason_preserved(): assert response is not None assert len(response.choices) == 1 assert response.choices[0].finish_reason == "content_filter" + + +def test_streaming_tool_call_finish_reason_with_empty_content_in_final_chunk(): + """ + When Gemini streams tool calls and the final chunk has BOTH empty content + (e.g. parts: [{text: ""}]) AND finishReason="STOP", the finish_reason + must still be "tool_calls". + + This covers models like gemini-3.1-flash-lite-preview that send the + final chunk with content (empty text) instead of omitting it entirely. + + Ref: https://github.com/BerriAI/litellm/issues/22900 + """ + logging_obj = _make_logging_obj() + iterator = ModelResponseIterator( + streaming_response=iter([]), + sync_stream=True, + logging_obj=logging_obj, + ) + + # Chunk 1: tool call with no finishReason + chunk_with_tool_calls = { + "candidates": [ + { + "content": { + "parts": [ + { + "functionCall": { + "name": "get_weather", + "args": {"location": "San Francisco"}, + } + } + ], + "role": "model", + }, + "index": 0, + } + ], + } + + # Chunk 2: finishReason="STOP" WITH empty content (text: "") + chunk_with_empty_content_and_finish = { + "candidates": [ + { + "content": { + "parts": [{"text": ""}], + "role": "model", + }, + "finishReason": "STOP", + "index": 0, + } + ], + "usageMetadata": { + "promptTokenCount": 50, + "candidatesTokenCount": 20, + "totalTokenCount": 70, + }, + } + + # Process chunk 1 + response1 = iterator.chunk_parser(chunk_with_tool_calls) + assert response1 is not None + assert len(response1.choices) == 1 + assert response1.choices[0].delta.tool_calls is not None + assert iterator.has_seen_tool_calls is True + + # Process chunk 2 (final chunk with empty content) + response2 = iterator.chunk_parser(chunk_with_empty_content_and_finish) + assert response2 is not None + assert len(response2.choices) == 1 + # Must be "tool_calls", NOT "stop" + assert response2.choices[0].finish_reason == "tool_calls" From 8b4a74a69c14b1a806d7e332cf371303de8c52f2 Mon Sep 17 00:00:00 2001 From: Chesars Date: Tue, 17 Mar 2026 18:06:42 -0300 Subject: [PATCH 14/25] fix(core): map Anthropic 'refusal' finish reason to 'content_filter' MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Anthropic's 'refusal' stop_reason was missing from _FINISH_REASON_MAP, causing it to fall through to the default 'stop' — hiding the fact that the model refused to respond due to safety policies. Fixes #23793 --- litellm/litellm_core_utils/core_helpers.py | 1 + tests/test_litellm/litellm_core_utils/test_core_helpers.py | 3 +++ 2 files changed, 4 insertions(+) diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index ee111f35929..9d9255cdf15 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -64,6 +64,7 @@ _FINISH_REASON_MAP: dict[str, OpenAIChatCompletionFinishReason] = { "end_turn": "stop", "max_tokens": "length", "tool_use": "tool_calls", + "refusal": "content_filter", "compaction": "length", # Cohere "COMPLETE": "stop", diff --git a/tests/test_litellm/litellm_core_utils/test_core_helpers.py b/tests/test_litellm/litellm_core_utils/test_core_helpers.py index 0ef76e0942d..134ed5d25f8 100644 --- a/tests/test_litellm/litellm_core_utils/test_core_helpers.py +++ b/tests/test_litellm/litellm_core_utils/test_core_helpers.py @@ -74,6 +74,9 @@ class TestMapFinishReasonAnthropic: def test_compaction(self): assert map_finish_reason("compaction") == "length" + def test_refusal(self): + assert map_finish_reason("refusal") == "content_filter" + class TestMapFinishReasonGemini: @pytest.mark.parametrize( From bed44f5fe53cff86b0818844db3be05a3bc77719 Mon Sep 17 00:00:00 2001 From: Rohan Date: Wed, 18 Mar 2026 03:08:04 +0530 Subject: [PATCH 15/25] Add Akto Guardrails to LiteLLM (#23250) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * akto guardrails support in litellm * docs(guardrails): add akto to supported values in types/guardrails.py * frontend changes + fixes * feat(akto): update Akto guardrail integration with new configuration options and modes * docs(akto): enhance Akto documentation and configuration descriptions for clarity * feat(tests): add proxy server request headers to sample request data * refactor(akto): remove optional account and VXLAN IDs; update documentation and tests * feat(akto): add event_type parameter for enhanced observability in guardrail logging * refactor(akto): update environment variable references * refactor the python codes * refactor and fix linting * refactor(akto): remove unused event hook and clean up imports * refactor(akto): enhance AktoGuardrail with async support and improved logging * fix: Register DynamoAI guardrail initializer and enum entry (#23752) * fix: Register DynamoAI guardrail initializer and enum entry Fix the "Unsupported guardrail: dynamoai" error by: 1. Adding DYNAMOAI to SupportedGuardrailIntegrations enum 2. Implementing initialize_guardrail() and registries in dynamoai/__init__.py The DynamoAI guardrail was added in PR #15920 but never properly registered in the initialization system. The __init__.py was missing the guardrail_initializer_registry and guardrail_class_registry dictionaries that the dynamic discovery mechanism looks for at module load time. Fixes #22773 Co-Authored-By: Claude Haiku 4.5 * Update litellm/proxy/guardrails/guardrail_hooks/dynamoai/__init__.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * Update litellm/proxy/guardrails/guardrail_hooks/dynamoai/__init__.py Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * test: Add tests for DynamoAI guardrail registration Verifies enum entry, initializer registry, class registry, instance creation, and global registry discovery. Co-Authored-By: Claude Opus 4.6 --------- Co-authored-by: Claude Haiku 4.5 Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> * docs: add v1.82.3 release notes and update provider_endpoints_support.json (#23816) * Revert "docs: add v1.82.3 release notes and update provider_endpoints_support…" (#23817) This reverts commit 966124966f83e6e1091ad08d9dfee6e341d3ad85. * Refactor Akto guardrail configuration and tests; update UI description and tags * add account and vxlan ID parameters to Akto guardrail initialization; update Akto logo format * enhance Akto guardrail documentation and improve error handling for non-JSON responses * address greptile issues * fix: update payload handling to use 'data' instead of 'json' in AktoGuardrail and adjust tests accordingly --------- Co-authored-by: Harshit Jain <48647625+Harshit28j@users.noreply.github.com> Co-authored-by: Claude Haiku 4.5 Co-authored-by: greptile-apps[bot] <165735046+greptile-apps[bot]@users.noreply.github.com> Co-authored-by: Joe Reyna Co-authored-by: Krish Dholakia --- docs/my-website/docs/proxy/guardrails/akto.md | 139 +++++ docs/my-website/sidebars.js | 1 + .../guardrail_hooks/akto/__init__.py | 37 ++ .../guardrails/guardrail_hooks/akto/akto.py | 456 +++++++++++++++ .../guardrail_hooks/dynamoai/__init__.py | 32 +- litellm/types/guardrails.py | 8 +- .../proxy/guardrails/guardrail_hooks/akto.py | 55 ++ .../guardrails_tests/test_akto_guardrails.py | 550 ++++++++++++++++++ .../guardrail_hooks/test_dynamoai.py | 81 +++ .../public/assets/logos/akto.svg | 10 + .../guardrails/guardrail_garden_configs.ts | 6 + .../guardrails/guardrail_garden_data.ts | 8 + .../guardrails/guardrail_info_helpers.tsx | 1 + 13 files changed, 1382 insertions(+), 2 deletions(-) create mode 100644 docs/my-website/docs/proxy/guardrails/akto.md create mode 100644 litellm/proxy/guardrails/guardrail_hooks/akto/__init__.py create mode 100644 litellm/proxy/guardrails/guardrail_hooks/akto/akto.py create mode 100644 litellm/types/proxy/guardrails/guardrail_hooks/akto.py create mode 100644 tests/guardrails_tests/test_akto_guardrails.py create mode 100644 tests/test_litellm/proxy/guardrails/guardrail_hooks/test_dynamoai.py create mode 100644 ui/litellm-dashboard/public/assets/logos/akto.svg diff --git a/docs/my-website/docs/proxy/guardrails/akto.md b/docs/my-website/docs/proxy/guardrails/akto.md new file mode 100644 index 00000000000..67ae741d11e --- /dev/null +++ b/docs/my-website/docs/proxy/guardrails/akto.md @@ -0,0 +1,139 @@ +# Akto + +## Overview +[Akto](https://www.akto.io/) provides API security guardrails and data ingestion for LLM traffic. + +Akto now uses a **two-entry guardrail pattern** in LiteLLM: +- `akto-validate` (`pre_call`) for request validation +- `akto-ingest` (`post_call`) for request/response ingestion + +There is no `on_flagged` setting anymore. + +Use these as two separate guardrails in `config.yaml`: +- `guardrail_name: "akto-validate"` +- `guardrail_name: "akto-ingest"` + +## 1. Get Your Akto Credentials + +Set up the Akto Guardrail API Service and grab: +- `AKTO_GUARDRAIL_API_BASE` — your Guardrail API Base URL +- `AKTO_API_KEY` — your API key + +## 2. Configure in `config.yaml` + +### Block + Ingest (recommended) + +Use both entries below. This gives you: +- pre-call block decision +- post-call ingestion for allowed traffic + +Keep these as two separate entries (`akto-validate` and `akto-ingest`). + +```yaml +guardrails: + - guardrail_name: "akto-validate" + litellm_params: + guardrail: akto + mode: pre_call + akto_base_url: os.environ/AKTO_GUARDRAIL_API_BASE + akto_api_key: os.environ/AKTO_API_KEY + default_on: true + unreachable_fallback: fail_closed # optional: fail_open | fail_closed (default: fail_closed) + guardrail_timeout: 5 # optional, default: 5 + akto_account_id: "1000000" # optional, env fallback: AKTO_ACCOUNT_ID + akto_vxlan_id: "0" # optional, env fallback: AKTO_VXLAN_ID + + - guardrail_name: "akto-ingest" + litellm_params: + guardrail: akto + mode: post_call + akto_base_url: os.environ/AKTO_GUARDRAIL_API_BASE + akto_api_key: os.environ/AKTO_API_KEY + default_on: true +``` + +### Monitor-only mode + +If you only want logging/ingestion and no blocking, keep only `akto-ingest`. + +```yaml +guardrails: + - guardrail_name: "akto-ingest" + litellm_params: + guardrail: akto + mode: post_call + akto_base_url: os.environ/AKTO_GUARDRAIL_API_BASE + akto_api_key: os.environ/AKTO_API_KEY + default_on: true +``` + +## 3. Test It + +```shell +curl -i http://localhost:4000/v1/chat/completions \ + -H "Content-Type: application/json" \ + -H "Authorization: Bearer " \ + -d '{ + "model": "gpt-3.5-turbo", + "messages": [ + {"role": "user", "content": "Hello, how are you?"} + ] + }' +``` + +If a request gets blocked: + +```json +{ + "error": { + "message": "Prompt injection detected", + "type": "None", + "param": "None", + "code": "403" + } +} +``` + +## 4. How It Works + +**Block + Ingest mode:** +``` +Request → LiteLLM → Akto guardrail check + → Allowed → forward to LLM → ingest response + → Blocked → ingest blocked marker → 403 error +``` + +**Monitor-only mode:** +``` +Request → LiteLLM → forward to LLM → get response + → Send to Akto (guardrails + ingest) → log only +``` + +## 5. Event behavior + +| Entry | LiteLLM hook | Akto call behavior | +|------|---|---| +| `akto-validate` | `pre_call` | Awaited call with `guardrails=true`, `ingest_data=false` | +| `akto-ingest` | `post_call` | Fire-and-forget call with `guardrails=true`, `ingest_data=true` | + +When blocked in `pre_call`, LiteLLM sends one fire-and-forget ingest payload with blocked metadata and returns `403`. + +## 6. Parameters + +| Parameter | Env Variable | Default | Description | +|-----------|-------------|---------|-------------| +| `akto_base_url` | `AKTO_GUARDRAIL_API_BASE` | *required* | Akto Guardrail API Base URL | +| `akto_api_key` | `AKTO_API_KEY` | *required* | API key (sent as `Authorization` header) | +| `akto_account_id` | `AKTO_ACCOUNT_ID` | `1000000` | Akto account id included in payload | +| `akto_vxlan_id` | `AKTO_VXLAN_ID` | `0` | Akto vxlan id included in payload | +| `unreachable_fallback` | — | `fail_closed` | `fail_open` or `fail_closed` | +| `guardrail_timeout` | — | `5` | Timeout in seconds | +| `default_on` | — | `true` (recommended) | Enables the guardrail entry by default | + +## 7. Error Handling + +| Scenario | `fail_closed` (default) | `fail_open` | +|----------|------------------------|-------------| +| Akto unreachable | ❌ Blocked (503) | ✅ Passes through | +| Akto returns error | ❌ Blocked (503) | ✅ Passes through | +| Guardrail says no | ❌ Blocked (403) | ❌ Blocked (403) | diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index 1362745a91f..e53891d6333 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -52,6 +52,7 @@ const sidebars = { label: "Providers", items: [ ...[ + "proxy/guardrails/akto", "proxy/guardrails/qualifire", "proxy/guardrails/aim_security", "proxy/guardrails/onyx_security", diff --git a/litellm/proxy/guardrails/guardrail_hooks/akto/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/akto/__init__.py new file mode 100644 index 00000000000..4ae26755409 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/akto/__init__.py @@ -0,0 +1,37 @@ +from typing import TYPE_CHECKING + +from litellm.types.guardrails import SupportedGuardrailIntegrations + +from .akto import AktoGuardrail + + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + + +def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"): + import litellm + + _akto_callback = AktoGuardrail( + akto_base_url=getattr(litellm_params, "akto_base_url", None), + akto_api_key=getattr(litellm_params, "akto_api_key", None), + akto_account_id=getattr(litellm_params, "akto_account_id", None), + akto_vxlan_id=getattr(litellm_params, "akto_vxlan_id", None), + unreachable_fallback=getattr(litellm_params, "unreachable_fallback", "fail_closed"), + guardrail_timeout=getattr(litellm_params, "guardrail_timeout", None), + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + ) + + litellm.logging_callback_manager.add_litellm_callback(_akto_callback) + return _akto_callback + + +guardrail_initializer_registry = { + SupportedGuardrailIntegrations.AKTO.value: initialize_guardrail, +} + +guardrail_class_registry = { + SupportedGuardrailIntegrations.AKTO.value: AktoGuardrail, +} diff --git a/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py b/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py new file mode 100644 index 00000000000..be9c9cb1be7 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/akto/akto.py @@ -0,0 +1,456 @@ +"""Akto guardrail integration for LiteLLM proxy. + +Uses a two-config-entry pattern: + - akto-validate (pre_call): Checks request against Akto guardrails, blocks if flagged. + - akto-ingest (post_call): Sends request+response to Akto for data ingestion. + +For monitor-only mode, enable only akto-ingest without akto-validate. +""" + +import asyncio +import json +import os +from datetime import datetime +from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Tuple, Type + +from fastapi import HTTPException + +import httpx + +from litellm._logging import verbose_proxy_logger +from litellm.integrations.custom_guardrail import ( + CustomGuardrail, + log_guardrail_information, +) +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, + httpxSpecialProvider, +) +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.utils import GenericGuardrailAPIInputs + +if TYPE_CHECKING: + from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel + +HTTP_PROXY_PATH = "/api/http-proxy" +AKTO_CONNECTOR_NAME = "litellm" +DEFAULT_GUARDRAIL_TIMEOUT = 5 + + +class AktoGuardrail(CustomGuardrail): + """LiteLLM guardrail hook that validates and ingests LLM traffic via the Akto API.""" + + # Maps event_hook to the input_type it should handle; mismatches are no-ops + HOOK_TO_INPUT = {"pre_call": "request", "post_call": "response"} + + @staticmethod + def get_config_model() -> Type["GuardrailConfigModel"]: + """Return the Pydantic config model for YAML-based initialization.""" + from litellm.types.proxy.guardrails.guardrail_hooks.akto import ( + AktoConfigModel, + ) + + return AktoConfigModel + + def __init__( + self, + akto_base_url: Optional[str] = None, + akto_api_key: Optional[str] = None, + akto_account_id: Optional[str] = None, + akto_vxlan_id: Optional[str] = None, + unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed", + guardrail_timeout: Optional[int] = None, + **kwargs: Any, + ) -> None: + """Initialize the Akto guardrail. + + Args: + akto_base_url: Akto API base URL. Falls back to AKTO_GUARDRAIL_API_BASE env var. + akto_api_key: Akto API key. Falls back to AKTO_API_KEY env var. + akto_account_id: Akto account ID. Falls back to AKTO_ACCOUNT_ID env var, then "1000000". + akto_vxlan_id: Akto VXLAN ID. Falls back to AKTO_VXLAN_ID env var, then "0". + unreachable_fallback: Behavior when Akto is unreachable — block or allow. + guardrail_timeout: HTTP timeout in seconds for Akto API calls. + """ + self.async_handler = get_async_httpx_client( + llm_provider=httpxSpecialProvider.GuardrailCallback, + ) + self.background_tasks: set = set() + + self.akto_base_url = (akto_base_url or os.environ.get("AKTO_GUARDRAIL_API_BASE", "")).rstrip("/") + if not self.akto_base_url: + raise ValueError("akto_base_url is required. Set AKTO_GUARDRAIL_API_BASE or pass it in litellm_params.") + + self.akto_api_key = akto_api_key or os.environ.get("AKTO_API_KEY", "") + if not self.akto_api_key: + raise ValueError("akto_api_key is required. Set AKTO_API_KEY or pass it in litellm_params.") + + self.unreachable_fallback: Literal["fail_closed", "fail_open"] = unreachable_fallback + self.guardrail_timeout = guardrail_timeout or DEFAULT_GUARDRAIL_TIMEOUT + self.akto_account_id = akto_account_id or os.environ.get("AKTO_ACCOUNT_ID", "1000000") + self.akto_vxlan_id = akto_vxlan_id or os.environ.get("AKTO_VXLAN_ID", "0") + + kwargs["supported_event_hooks"] = [ + GuardrailEventHooks.pre_call, + GuardrailEventHooks.post_call, + ] + super().__init__(**kwargs) + + verbose_proxy_logger.debug( + "Akto guardrail initialized: base_url=%s fallback=%s", + self.akto_base_url, + self.unreachable_fallback, + ) + + @staticmethod + def resolve_metadata_value(request_data: Optional[dict], key: str) -> Optional[str]: + """Look up a metadata value from litellm_metadata or metadata dicts.""" + if request_data is None: + return None + for dict_key in ("litellm_metadata", "metadata"): + container = request_data.get(dict_key) or {} + if isinstance(container, dict) and container: + value = container.get(key) + if value is not None: + return str(value).strip() + return None + + @staticmethod + def extract_request_path(request_data: dict) -> str: + """Extract the API route from request metadata, defaulting to /v1/chat/completions.""" + metadata = request_data.get("metadata") or {} + if not isinstance(metadata, dict): + metadata = {} + route = metadata.get("user_api_key_request_route") + return route if route else "/v1/chat/completions" + + def prepare_headers(self) -> Dict[str, str]: + """Build HTTP headers for the Akto API call.""" + return { + "content-type": "application/json", + "Authorization": self.akto_api_key, + } + + @staticmethod + def build_query_params(*, guardrails: bool, ingest_data: bool) -> Dict[str, str]: + """Build query params that control Akto backend behavior (guardrail check and/or data ingestion).""" + params: Dict[str, str] = {"akto_connector": AKTO_CONNECTOR_NAME} + if guardrails: + params["guardrails"] = "true" + if ingest_data: + params["ingest_data"] = "true" + return params + + @staticmethod + def build_request_headers(request_data: dict) -> Dict[str, str]: + """Build the requestHeaders field from proxy request headers.""" + headers: Dict[str, str] = {"content-type": "application/json"} + proxy_req = request_data.get("proxy_server_request", {}) + if not isinstance(proxy_req, dict): + return headers + proxy_req_headers = proxy_req.get("headers") + if isinstance(proxy_req_headers, dict): + for key, val in proxy_req_headers.items(): + if key and val: + headers[str(key).lower()] = str(val) + return headers + + @staticmethod + def build_request_body( + inputs: GenericGuardrailAPIInputs, + request_data: Optional[dict] = None, + ) -> Dict[str, Any]: + """Build the LLM request body from guardrail inputs (messages, model, tools).""" + model = inputs.get("model", "") or "" + body: Dict[str, Any] = {"model": model} + + structured = inputs.get("structured_messages") + if structured: + body["messages"] = structured + elif request_data is not None and request_data.get("messages"): + body["messages"] = request_data["messages"] + if request_data.get("model"): + body["model"] = request_data["model"] + else: + texts = inputs.get("texts", []) + body["messages"] = [{"role": "user", "content": t} for t in texts] if texts else [] + + tools = inputs.get("tools") + if tools: + body["tools"] = tools + elif request_data is not None and request_data.get("tools"): + body["tools"] = request_data["tools"] + + tool_calls = inputs.get("tool_calls") + if tool_calls: + body["tool_calls"] = tool_calls + + return body + + @staticmethod + def build_response_body( + inputs: GenericGuardrailAPIInputs, + request_data: Optional[dict] = None, + ) -> Dict[str, Any]: + """Build the LLM response body, preferring the actual model response if available.""" + model_response = request_data.get("response") if request_data else None + if model_response is not None and hasattr(model_response, "model_dump"): + return model_response.model_dump() + + texts = inputs.get("texts", []) + if texts: + return {"choices": [{"message": {"content": t, "role": "assistant"}} for t in texts]} + return {} + + @staticmethod + def build_tag_metadata(request_data: dict) -> Dict[str, str]: + """Build tag/metadata dict with user_id and team_id for Akto tracking.""" + tag: Dict[str, str] = {"gen-ai": "Gen AI"} + user_id = AktoGuardrail.resolve_metadata_value(request_data, "user_api_key_user_id") + team_id = AktoGuardrail.resolve_metadata_value(request_data, "user_api_key_team_id") + if user_id: + tag["user_id"] = user_id + if team_id: + tag["team_id"] = team_id + return tag + + def build_akto_payload( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + *, + status_code: int = 200, + include_response: bool = False, + ) -> Dict[str, Any]: + """Build the flat MIRRORING payload sent to Akto's HTTP proxy endpoint. + + All body fields use double-encoding: json.dumps({"body": json.dumps(actual_body)}) + to match the canonical CLI hook format. + """ + request_path = self.extract_request_path(request_data) + request_headers = self.build_request_headers(request_data) + request_body = self.build_request_body(inputs, request_data) + tag = self.build_tag_metadata(request_data) + + response_payload = json.dumps({}) # Empty body wrapper when no response yet + response_headers: Dict[str, str] = {} + if include_response: + response_body = self.build_response_body(inputs, request_data) + response_payload = json.dumps({"body": json.dumps(response_body)}) # Double-encoded + response_headers = {"content-type": "application/json"} + + # Extract client IP from proxy headers + ip = "" + proxy_req = request_data.get("proxy_server_request", {}) + proxy_headers = proxy_req.get("headers", {}) if isinstance(proxy_req, dict) else {} + if isinstance(proxy_headers, dict): + ip = proxy_headers.get("x-forwarded-for") or proxy_headers.get("x-real-ip") or "" + if "," in ip: + ip = ip.split(",")[0].strip() + + return { + "path": request_path, + "requestHeaders": json.dumps(request_headers), + "responseHeaders": json.dumps(response_headers), + "method": "POST", + "requestPayload": json.dumps({"body": json.dumps(request_body)}), # Double-encoded + "responsePayload": response_payload, + "ip": ip, + "destIp": "127.0.0.1", + "time": str(int(datetime.now().timestamp() * 1000)), + "statusCode": str(status_code), + "type": "HTTP/1.1", + "status": str(status_code), + "akto_account_id": self.akto_account_id, + "akto_vxlan_id": self.akto_vxlan_id, + "is_pending": "false", + "source": "MIRRORING", + "direction": None, + "process_id": None, + "socket_id": None, + "daemonset_id": None, + "enabled_graph": None, + "tag": json.dumps(tag), + "metadata": json.dumps(tag), + "contextSource": "AGENTIC", + } + + async def send_request( + self, + *, + guardrails: bool, + ingest_data: bool, + payload: dict, + ) -> httpx.Response: + """Send an HTTP POST to the Akto API endpoint.""" + endpoint = f"{self.akto_base_url}{HTTP_PROXY_PATH}" + params = self.build_query_params(guardrails=guardrails, ingest_data=ingest_data) + headers = self.prepare_headers() + return await self.async_handler.post( + url=endpoint, + data=json.dumps(payload), + params=params, + headers=headers, + timeout=self.guardrail_timeout, + ) + + @staticmethod + def handle_guardrail_response(response: httpx.Response) -> Tuple[bool, str]: + """Parse the Akto guardrail response. Returns (allowed, reason).""" + if response.status_code != 200: + verbose_proxy_logger.error("Akto returned HTTP %d", response.status_code) + raise httpx.HTTPStatusError( + f"Akto returned unexpected status {response.status_code}", + request=response.request, + response=response, + ) + try: + result = response.json() + except (json.JSONDecodeError, ValueError) as e: + response_text = getattr(response, "text", "") + verbose_proxy_logger.error( + "Akto returned non-JSON body for status 200: %r", + response_text[:200], + ) + raise httpx.RequestError( + "Akto returned non-JSON body", + request=response.request, + ) from e + if not isinstance(result, dict): + return True, "" + data = result.get("data") or {} + if not isinstance(data, dict): + return True, "" + guardrails_result = data.get("guardrailsResult") or {} + if not isinstance(guardrails_result, dict): + return True, "" + return ( + bool(guardrails_result.get("Allowed", True)), + str(guardrails_result.get("Reason", "")), + ) + + def handle_unreachable( + self, + inputs: GenericGuardrailAPIInputs, + error: Exception, + ) -> GenericGuardrailAPIInputs: + """Handle Akto being unreachable based on fail_open/fail_closed config.""" + if self.unreachable_fallback == "fail_open": + verbose_proxy_logger.critical( + "Akto unreachable (fail-open): %s", + str(error), + exc_info=error, + ) + return inputs + + verbose_proxy_logger.error("Akto unreachable (fail-closed): %s", str(error)) + raise HTTPException( + status_code=503, + detail="Akto guardrail service unreachable", + ) + + async def fire_and_forget_request( + self, + *, + guardrails: bool, + ingest_data: bool, + payload: dict, + ) -> None: + """Send a request without awaiting it in the caller. Errors are logged, not raised.""" + try: + response = await self.send_request( + guardrails=guardrails, + ingest_data=ingest_data, + payload=payload, + ) + if response.status_code != 200: + verbose_proxy_logger.error( + "Akto fire-and-forget returned HTTP %d", + response.status_code, + ) + except Exception as e: + verbose_proxy_logger.error("Akto fire-and-forget error: %s", str(e)) + + @log_guardrail_information + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj=None, + ) -> GenericGuardrailAPIInputs: + """Main entry point called by LiteLLM's guardrail framework. + + Pre_call (input_type="request"): + - Awaits guardrail check. If blocked, fires off ingest with 403 marker and raises. + Post_call (input_type="response"): + - Fire-and-forget combined guardrail + ingest call. + """ + # Skip if this hook doesn't handle the current input_type + expected = self.HOOK_TO_INPUT.get(str(self.event_hook)) + if expected and expected != input_type: + return inputs + + if input_type == "request": + # Pre_call: awaited guardrail check (no ingestion) + payload = self.build_akto_payload(inputs, request_data, include_response=False) + try: + response = await self.send_request( + guardrails=True, + ingest_data=False, + payload=payload, + ) + allowed, reason = self.handle_guardrail_response(response) + except HTTPException: + raise + except (httpx.RequestError, httpx.HTTPStatusError) as e: + return self.handle_unreachable( + inputs=inputs, + error=e, + ) + + if not allowed: + # Build a blocked marker payload with 403 status and reason + blocked_payload = self.build_akto_payload( + inputs, + request_data, + include_response=False, + status_code=403, + ) + blocked_payload["responsePayload"] = json.dumps( + { + "body": json.dumps({"x-blocked-by": "Akto Proxy", "reason": reason}), + } + ) + blocked_payload["responseHeaders"] = json.dumps( + {"content-type": "application/json"}, + ) + # Fire-and-forget ingest of the blocked request, then raise 403 + task = asyncio.create_task( + self.fire_and_forget_request( + guardrails=False, + ingest_data=True, + payload=blocked_payload, + ) + ) + self.background_tasks.add(task) + task.add_done_callback(self.background_tasks.discard) + raise HTTPException( + status_code=403, + detail=reason or "Blocked by Akto Guardrails", + ) + + elif input_type == "response": + # Post_call: fire-and-forget combined guardrail + ingest + payload = self.build_akto_payload(inputs, request_data, include_response=True) + task = asyncio.create_task( + self.fire_and_forget_request( + guardrails=True, + ingest_data=True, + payload=payload, + ) + ) + self.background_tasks.add(task) + task.add_done_callback(self.background_tasks.discard) + + return inputs diff --git a/litellm/proxy/guardrails/guardrail_hooks/dynamoai/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/dynamoai/__init__.py index 79f1992da44..f9ebf46a270 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/dynamoai/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/dynamoai/__init__.py @@ -1,3 +1,33 @@ +from typing import TYPE_CHECKING + +from litellm.types.guardrails import SupportedGuardrailIntegrations + from .dynamoai import DynamoAIGuardrails -__all__ = ["DynamoAIGuardrails"] +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + + +def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"): + import litellm + + _dynamoai_callback = DynamoAIGuardrails( + api_base=litellm_params.api_base, + api_key=litellm_params.api_key, + guardrail_name=guardrail.get("guardrail_name", ""), + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + ) + litellm.logging_callback_manager.add_litellm_callback(_dynamoai_callback) + + return _dynamoai_callback + + +guardrail_initializer_registry = { + SupportedGuardrailIntegrations.DYNAMOAI.value: initialize_guardrail, +} + + +guardrail_class_registry = { + SupportedGuardrailIntegrations.DYNAMOAI.value: DynamoAIGuardrails, +} diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index d5abf5c8fbf..11b1d0d40c0 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -17,6 +17,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.grayswan import ( from litellm.types.proxy.guardrails.guardrail_hooks.ibm import ( IBMGuardrailsBaseConfigModel, ) +from litellm.types.proxy.guardrails.guardrail_hooks.akto import ( + AktoConfigModel, +) from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import ( ContentFilterCategoryConfig, ) @@ -33,7 +36,7 @@ Pydantic object defining how to set guardrails on litellm proxy guardrails: - guardrail_name: "bedrock-pre-guard" litellm_params: - guardrail: bedrock # supported values: "aporia", "bedrock", "lakera", "zscaler_ai_guard" + guardrail: bedrock # supported values: "akto", "aporia", "bedrock", "lakera", "zscaler_ai_guard" mode: "during_call" guardrailIdentifier: ff6ujrregl1q guardrailVersion: "DRAFT" @@ -44,6 +47,7 @@ guardrails: class SupportedGuardrailIntegrations(Enum): APORIA = "aporia" BEDROCK = "bedrock" + DYNAMOAI = "dynamoai" GUARDRAILS_AI = "guardrails_ai" LAKERA = "lakera" LAKERA_V2 = "lakera_v2" @@ -78,6 +82,7 @@ class SupportedGuardrailIntegrations(Enum): SEMANTIC_GUARD = "semantic_guard" MCP_END_USER_PERMISSION = "mcp_end_user_permission" BLOCK_CODE_EXECUTION = "block_code_execution" + AKTO = "akto" class Role(Enum): @@ -735,6 +740,7 @@ class LitellmParams( NomaGuardrailConfigModel, ToolPermissionGuardrailConfigModel, ZscalerAIGuardConfigModel, + AktoConfigModel, JavelinGuardrailConfigModel, BaseLitellmParams, EnkryptAIGuardrailConfigs, diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/akto.py b/litellm/types/proxy/guardrails/guardrail_hooks/akto.py new file mode 100644 index 00000000000..180c89e8115 --- /dev/null +++ b/litellm/types/proxy/guardrails/guardrail_hooks/akto.py @@ -0,0 +1,55 @@ +from typing import Optional, Literal + +from pydantic import Field + +from .base import GuardrailConfigModel + + +class AktoConfigModel(GuardrailConfigModel): + """ + Config for the Akto guardrail. + + Use two separate config entries to control behaviour: + akto-validate (mode: pre_call) -> check guardrails, block if flagged + akto-ingest (mode: post_call) -> ingest request+response data + """ + + akto_base_url: Optional[str] = Field( + default=None, + description="Akto Guardrail API Base URL. Env: AKTO_GUARDRAIL_API_BASE.", + json_schema_extra={ + "examples": [ + "http://localhost:9090", + "https://akto-ingestion.example.com", + ] + }, + ) + + akto_api_key: Optional[str] = Field( + default=None, + description="API key for Akto. Env: AKTO_API_KEY.", + ) + + akto_account_id: Optional[str] = Field( + default=None, + description="Akto account ID for multi-tenant deployments. Env: AKTO_ACCOUNT_ID. Default: '1000000'.", + ) + + akto_vxlan_id: Optional[str] = Field( + default=None, + description="Akto VXLAN ID. Env: AKTO_VXLAN_ID. Default: '0'.", + ) + + unreachable_fallback: Literal["fail_closed", "fail_open"] = Field( + default="fail_closed", + description="What to do when Akto is unreachable. 'fail_open' = allow, 'fail_closed' = block.", + ) + + guardrail_timeout: Optional[int] = Field( + default=None, + description="HTTP timeout in seconds. Default: 5.", + ) + + @staticmethod + def ui_friendly_name() -> str: + return "Akto" diff --git a/tests/guardrails_tests/test_akto_guardrails.py b/tests/guardrails_tests/test_akto_guardrails.py new file mode 100644 index 00000000000..3c70104a219 --- /dev/null +++ b/tests/guardrails_tests/test_akto_guardrails.py @@ -0,0 +1,550 @@ +import asyncio +import json +import os +from unittest.mock import AsyncMock, MagicMock, patch + +import httpx +import pytest + +from starlette.exceptions import HTTPException +from litellm.types.utils import GenericGuardrailAPIInputs +from litellm.proxy.guardrails.guardrail_registry import guardrail_initializer_registry, guardrail_class_registry +from litellm.proxy.guardrails.guardrail_hooks.akto.akto import AktoGuardrail + + +# --------------------------------------------------------------------------- +# Registry tests +# --------------------------------------------------------------------------- + + +def test_akto_in_guardrail_initializer_registry(): + assert "akto" in guardrail_initializer_registry + + +def test_akto_in_guardrail_class_registry(): + assert "akto" in guardrail_class_registry + assert guardrail_class_registry["akto"] is AktoGuardrail + + +# --------------------------------------------------------------------------- +# Fixtures +# --------------------------------------------------------------------------- + + +@pytest.fixture +def akto_validate(): + """AktoGuardrail configured for pre_call (akto-validate).""" + return AktoGuardrail( + akto_base_url="http://localhost:9090", + akto_api_key="test-token", + unreachable_fallback="fail_closed", + guardrail_name="test-akto-validate", + event_hook="pre_call", + ) + + +@pytest.fixture +def akto_ingest(): + """AktoGuardrail configured for post_call (akto-ingest).""" + return AktoGuardrail( + akto_base_url="http://localhost:9090", + akto_api_key="test-token", + unreachable_fallback="fail_open", + guardrail_name="test-akto-ingest", + event_hook="post_call", + ) + + +@pytest.fixture +def sample_inputs() -> GenericGuardrailAPIInputs: + return GenericGuardrailAPIInputs( + texts=["Hello, how are you?"], + model="gpt-4", + ) + + +@pytest.fixture +def sample_request_data() -> dict: + return { + "metadata": { + "user_api_key_request_route": "/v1/chat/completions", + "user_api_key": "sk-test-123", + "user_api_key_user_id": "user-1", + "user_api_key_team_id": "team-1", + }, + "proxy_server_request": { + "headers": { + "x-forwarded-for": "10.0.0.1", + } + }, + } + + +def _mock_allowed_response(): + mock = MagicMock(spec=httpx.Response) + mock.status_code = 200 + mock.json.return_value = {"data": {"guardrailsResult": {"Allowed": True, "Reason": ""}}} + return mock + + +def _mock_blocked_response(reason="Prompt injection detected"): + mock = MagicMock(spec=httpx.Response) + mock.status_code = 200 + mock.json.return_value = {"data": {"guardrailsResult": {"Allowed": False, "Reason": reason}}} + return mock + + +# --------------------------------------------------------------------------- +# Initialization tests +# --------------------------------------------------------------------------- + + +def test_init_requires_akto_base_url(): + with patch.dict(os.environ, {}, clear=True): + with pytest.raises(ValueError, match="akto_base_url is required"): + AktoGuardrail( + akto_base_url="", + akto_api_key="test-token", + guardrail_name="test", + event_hook="pre_call", + ) + + +def test_init_requires_api_key(): + with patch.dict(os.environ, {}, clear=True): + with pytest.raises(ValueError, match="akto_api_key is required"): + AktoGuardrail( + akto_base_url="http://localhost:9090", + akto_api_key="", + guardrail_name="test", + event_hook="pre_call", + ) + + +def test_init_from_env(): + with patch.dict( + os.environ, + { + "AKTO_GUARDRAIL_API_BASE": "http://env-host:9090", + "AKTO_API_KEY": "env-token", + "AKTO_ACCOUNT_ID": "2000000", + "AKTO_VXLAN_ID": "42", + }, + ): + g = AktoGuardrail(guardrail_name="env-test", event_hook="post_call") + assert g.akto_base_url == "http://env-host:9090" + assert g.akto_api_key == "env-token" + assert g.guardrail_timeout == 5 + assert g.akto_account_id == "2000000" + assert g.akto_vxlan_id == "42" + + +def test_init_defaults(): + g = AktoGuardrail( + akto_base_url="http://localhost:9090", + akto_api_key="test-token", + guardrail_name="default-test", + event_hook="pre_call", + ) + assert g.unreachable_fallback == "fail_closed" + assert g.guardrail_timeout == 5 + assert g.akto_account_id == "1000000" + assert g.akto_vxlan_id == "0" + + +def test_background_tasks_per_instance(): + a = AktoGuardrail( + akto_base_url="http://localhost:9090", + akto_api_key="test-token", + guardrail_name="instance-a", + event_hook="pre_call", + ) + b = AktoGuardrail( + akto_base_url="http://localhost:9090", + akto_api_key="test-token", + guardrail_name="instance-b", + event_hook="post_call", + ) + assert a.background_tasks is not b.background_tasks + + +# --------------------------------------------------------------------------- +# Payload format tests +# --------------------------------------------------------------------------- + + +def test_build_akto_payload_format(akto_validate, sample_inputs, sample_request_data): + payload = akto_validate.build_akto_payload(sample_inputs, sample_request_data, include_response=False) + + assert payload["path"] == "/v1/chat/completions" + assert payload["method"] == "POST" + assert payload["type"] == "HTTP/1.1" + assert payload["akto_account_id"] == "1000000" + assert payload["akto_vxlan_id"] == "0" + assert payload["is_pending"] == "false" + assert payload["source"] == "MIRRORING" + assert payload["contextSource"] == "AGENTIC" + assert payload["ip"] == "10.0.0.1" + + req_headers = json.loads(payload["requestHeaders"]) + assert "content-type" in req_headers + + req_wrapper = json.loads(payload["requestPayload"]) + req_body = json.loads(req_wrapper["body"]) + assert req_body["model"] == "gpt-4" + assert req_body["messages"][0]["content"] == "Hello, how are you?" + + tag = json.loads(payload["tag"]) + assert tag["gen-ai"] == "Gen AI" + + assert payload["responsePayload"] == json.dumps({}) + assert payload["time"].isdigit() + assert len(payload["time"]) >= 13 + + +def test_build_akto_payload_with_response(akto_validate, sample_inputs, sample_request_data): + payload = akto_validate.build_akto_payload(sample_inputs, sample_request_data, include_response=True) + resp_wrapper = json.loads(payload["responsePayload"]) + resp_body = json.loads(resp_wrapper["body"]) + assert "choices" in resp_body + + +def test_build_akto_payload_custom_account_ids(sample_inputs, sample_request_data): + g = AktoGuardrail( + akto_base_url="http://localhost:9090", + akto_api_key="test-token", + akto_account_id="9999", + akto_vxlan_id="7", + guardrail_name="custom-ids-test", + event_hook="pre_call", + ) + payload = g.build_akto_payload(sample_inputs, sample_request_data, include_response=False) + assert payload["akto_account_id"] == "9999" + assert payload["akto_vxlan_id"] == "7" + + +def test_build_query_params(): + params = AktoGuardrail.build_query_params(guardrails=True, ingest_data=False) + assert params == {"akto_connector": "litellm", "guardrails": "true"} + + params = AktoGuardrail.build_query_params(guardrails=False, ingest_data=True) + assert params == {"akto_connector": "litellm", "ingest_data": "true"} + + params = AktoGuardrail.build_query_params(guardrails=True, ingest_data=True) + assert params == { + "akto_connector": "litellm", + "guardrails": "true", + "ingest_data": "true", + } + + +# --------------------------------------------------------------------------- +# Guardrail response handling +# --------------------------------------------------------------------------- + + +def test_handle_guardrail_response_allowed(): + mock_resp = MagicMock(spec=httpx.Response) + mock_resp.status_code = 200 + mock_resp.json.return_value = {"data": {"guardrailsResult": {"Allowed": True, "Reason": ""}}} + allowed, reason = AktoGuardrail.handle_guardrail_response(mock_resp) + assert allowed is True + assert reason == "" + + +def test_handle_guardrail_response_blocked(): + mock_resp = MagicMock(spec=httpx.Response) + mock_resp.status_code = 200 + mock_resp.json.return_value = {"data": {"guardrailsResult": {"Allowed": False, "Reason": "PII detected"}}} + allowed, reason = AktoGuardrail.handle_guardrail_response(mock_resp) + assert allowed is False + assert reason == "PII detected" + + +def test_handle_guardrail_response_missing_result(): + mock_resp = MagicMock(spec=httpx.Response) + mock_resp.status_code = 200 + mock_resp.json.return_value = {} + allowed, _ = AktoGuardrail.handle_guardrail_response(mock_resp) + assert allowed is True + + +def test_handle_guardrail_response_data_none(): + mock_resp = MagicMock(spec=httpx.Response) + mock_resp.status_code = 200 + mock_resp.json.return_value = {"data": None} + allowed, reason = AktoGuardrail.handle_guardrail_response(mock_resp) + assert allowed is True + assert reason == "" + + +def test_handle_guardrail_response_guardrails_result_not_dict(): + mock_resp = MagicMock(spec=httpx.Response) + mock_resp.status_code = 200 + mock_resp.json.return_value = {"data": {"guardrailsResult": "invalid"}} + allowed, reason = AktoGuardrail.handle_guardrail_response(mock_resp) + assert allowed is True + assert reason == "" + + +def test_handle_guardrail_response_non_dict(): + mock_resp = MagicMock(spec=httpx.Response) + mock_resp.status_code = 200 + mock_resp.json.return_value = "invalid" + allowed, _ = AktoGuardrail.handle_guardrail_response(mock_resp) + assert allowed is True + + +def test_handle_guardrail_response_error_status(): + mock_resp = MagicMock(spec=httpx.Response) + mock_resp.status_code = 500 + mock_resp.request = MagicMock() + with pytest.raises(httpx.HTTPStatusError): + AktoGuardrail.handle_guardrail_response(mock_resp) + + +def test_handle_guardrail_response_non_json_body(): + mock_resp = MagicMock(spec=httpx.Response) + mock_resp.status_code = 200 + mock_resp.request = MagicMock() + mock_resp.text = "not json" + mock_resp.json.side_effect = json.JSONDecodeError("Expecting value", "", 0) + + with pytest.raises(httpx.RequestError): + AktoGuardrail.handle_guardrail_response(mock_resp) + + +# --------------------------------------------------------------------------- +# Pre-call (akto-validate) — allowed +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_pre_call_allowed(akto_validate, sample_inputs, sample_request_data): + akto_validate.async_handler.post = AsyncMock(return_value=_mock_allowed_response()) + + result = await akto_validate.apply_guardrail( + inputs=sample_inputs, + request_data=sample_request_data, + input_type="request", + ) + + assert result == sample_inputs + akto_validate.async_handler.post.assert_called_once() + call_params = akto_validate.async_handler.post.call_args.kwargs["params"] + assert call_params.get("guardrails") == "true" + assert "ingest_data" not in call_params + + +# --------------------------------------------------------------------------- +# Pre-call (akto-validate) — blocked +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_pre_call_blocked(akto_validate, sample_inputs, sample_request_data): + akto_validate.async_handler.post = AsyncMock( + side_effect=[ + _mock_blocked_response("PII detected"), + _mock_allowed_response(), + ] + ) + + with pytest.raises(HTTPException) as exc_info: + await akto_validate.apply_guardrail( + inputs=sample_inputs, + request_data=sample_request_data, + input_type="request", + ) + + await asyncio.sleep(0) + await asyncio.sleep(0) + + assert exc_info.value.status_code == 403 + + assert akto_validate.async_handler.post.call_count == 2 + + first_call_params = akto_validate.async_handler.post.call_args_list[0].kwargs["params"] + assert first_call_params.get("guardrails") == "true" + + second_call_params = akto_validate.async_handler.post.call_args_list[1].kwargs["params"] + assert second_call_params.get("ingest_data") == "true" + assert "guardrails" not in second_call_params + second_payload = json.loads(akto_validate.async_handler.post.call_args_list[1].kwargs["data"]) + assert second_payload["statusCode"] == "403" + resp_body = json.loads(second_payload["responsePayload"]) + inner = json.loads(resp_body["body"]) + assert inner["x-blocked-by"] == "Akto Proxy" + assert inner["reason"] == "PII detected" + + +# --------------------------------------------------------------------------- +# Pre-call (akto-validate) — response input is no-op +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_validate_response_noop(akto_validate, sample_inputs, sample_request_data): + akto_validate.async_handler.post = AsyncMock() + + result = await akto_validate.apply_guardrail( + inputs=sample_inputs, + request_data=sample_request_data, + input_type="response", + ) + + assert result == sample_inputs + akto_validate.async_handler.post.assert_not_called() + + +# --------------------------------------------------------------------------- +# Post-call (akto-ingest) — combined guardrail + ingest +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_post_call_combined(akto_ingest, sample_inputs, sample_request_data): + akto_ingest.async_handler.post = AsyncMock(return_value=_mock_allowed_response()) + + result = await akto_ingest.apply_guardrail( + inputs=sample_inputs, + request_data=sample_request_data, + input_type="response", + ) + + await asyncio.sleep(0) + await asyncio.sleep(0) + + assert result == sample_inputs + akto_ingest.async_handler.post.assert_called_once() + call_params = akto_ingest.async_handler.post.call_args.kwargs["params"] + assert call_params.get("guardrails") == "true" + assert call_params.get("ingest_data") == "true" + + +# --------------------------------------------------------------------------- +# Post-call (akto-ingest) — request input is no-op +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_ingest_request_noop(akto_ingest, sample_inputs, sample_request_data): + akto_ingest.async_handler.post = AsyncMock() + + result = await akto_ingest.apply_guardrail( + inputs=sample_inputs, + request_data=sample_request_data, + input_type="request", + ) + + assert result == sample_inputs + akto_ingest.async_handler.post.assert_not_called() + + +# --------------------------------------------------------------------------- +# Fail-open / fail-closed +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_fail_open_on_unreachable(): + g = AktoGuardrail( + akto_base_url="http://localhost:9090", + akto_api_key="test-token", + unreachable_fallback="fail_open", + guardrail_name="fail-open-test", + event_hook="pre_call", + ) + g.async_handler.post = AsyncMock(side_effect=httpx.ConnectError("Connection refused")) + + inputs = GenericGuardrailAPIInputs(texts=["test"], model="gpt-4") + result = await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request") + + assert result.get("texts") == ["test"] + + +@pytest.mark.asyncio +async def test_fail_closed_on_unreachable(): + g = AktoGuardrail( + akto_base_url="http://localhost:9090", + akto_api_key="test-token", + unreachable_fallback="fail_closed", + guardrail_name="fail-closed-test", + event_hook="pre_call", + ) + g.async_handler.post = AsyncMock(side_effect=httpx.ConnectError("Connection refused")) + + inputs = GenericGuardrailAPIInputs(texts=["test"], model="gpt-4") + with pytest.raises(HTTPException) as exc_info: + await g.apply_guardrail(inputs=inputs, request_data={}, input_type="request") + assert exc_info.value.status_code == 503 + + +def test_fail_closed_generic_message(): + g = AktoGuardrail( + akto_base_url="http://localhost:9090", + akto_api_key="test-token", + unreachable_fallback="fail_closed", + guardrail_name="msg-test", + event_hook="pre_call", + ) + with pytest.raises(HTTPException) as exc_info: + g.handle_unreachable( + inputs=GenericGuardrailAPIInputs(texts=["test"], model="gpt-4"), + error=Exception("http://internal-host:9090/secret-path"), + ) + assert "internal-host" not in exc_info.value.detail + assert exc_info.value.detail == "Akto guardrail service unreachable" + + +# --------------------------------------------------------------------------- +# Helper method tests +# --------------------------------------------------------------------------- + + +def test_extract_request_path_from_metadata(): + path = AktoGuardrail.extract_request_path({"metadata": {"user_api_key_request_route": "/v1/embeddings"}}) + assert path == "/v1/embeddings" + + +def test_extract_request_path_fallback(): + path = AktoGuardrail.extract_request_path({}) + assert path == "/v1/chat/completions" + + +def test_extract_request_path_non_dict_metadata(): + path = AktoGuardrail.extract_request_path({"metadata": "invalid"}) + assert path == "/v1/chat/completions" + + +def test_resolve_metadata_value(): + assert ( + AktoGuardrail.resolve_metadata_value({"metadata": {"user_api_key_user_id": "u1"}}, "user_api_key_user_id") + == "u1" + ) + assert ( + AktoGuardrail.resolve_metadata_value( + {"litellm_metadata": {"user_api_key_team_id": "t1"}}, + "user_api_key_team_id", + ) + == "t1" + ) + assert AktoGuardrail.resolve_metadata_value({}, "some_key") is None + assert AktoGuardrail.resolve_metadata_value(None, "some_key") is None + + +def test_resolve_metadata_value_non_dict_containers(): + assert ( + AktoGuardrail.resolve_metadata_value( + {"metadata": "invalid", "litellm_metadata": ["bad"]}, + "some_key", + ) + is None + ) + + +def test_build_tag_metadata(akto_validate, sample_request_data): + tag = akto_validate.build_tag_metadata(sample_request_data) + assert tag["gen-ai"] == "Gen AI" + assert tag["user_id"] == "user-1" + assert tag["team_id"] == "team-1" diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_dynamoai.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_dynamoai.py new file mode 100644 index 00000000000..7bc4e951a5f --- /dev/null +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_dynamoai.py @@ -0,0 +1,81 @@ +""" +Tests for DynamoAI guardrail registration and initialization. +""" + +import os +from unittest.mock import patch + +import pytest + + +class TestDynamoAIGuardrailRegistration: + """Tests for DynamoAI guardrail registration in the guardrail system.""" + + def test_supported_guardrail_enum_entry(self): + """Test that DYNAMOAI is in SupportedGuardrailIntegrations enum.""" + from litellm.types.guardrails import SupportedGuardrailIntegrations + + assert hasattr(SupportedGuardrailIntegrations, "DYNAMOAI") + assert SupportedGuardrailIntegrations.DYNAMOAI.value == "dynamoai" + + def test_initialize_guardrail_function_exists(self): + """Test that initialize_guardrail function is properly exported.""" + from litellm.proxy.guardrails.guardrail_hooks.dynamoai import ( + guardrail_initializer_registry, + initialize_guardrail, + ) + + assert initialize_guardrail is not None + assert "dynamoai" in guardrail_initializer_registry + + def test_guardrail_class_registry_exists(self): + """Test that guardrail_class_registry is properly exported.""" + from litellm.proxy.guardrails.guardrail_hooks.dynamoai import ( + guardrail_class_registry, + ) + from litellm.proxy.guardrails.guardrail_hooks.dynamoai.dynamoai import ( + DynamoAIGuardrails, + ) + + assert "dynamoai" in guardrail_class_registry + assert guardrail_class_registry["dynamoai"] == DynamoAIGuardrails + + def test_initialize_guardrail_creates_instance(self): + """Test that initialize_guardrail creates a DynamoAIGuardrails instance.""" + from litellm.proxy.guardrails.guardrail_hooks.dynamoai import ( + initialize_guardrail, + ) + from litellm.proxy.guardrails.guardrail_hooks.dynamoai.dynamoai import ( + DynamoAIGuardrails, + ) + from litellm.types.guardrails import LitellmParams + + litellm_params = LitellmParams( + guardrail="dynamoai", + mode="pre_call", + api_key="test-key", + api_base="https://test.dynamo.ai", + ) + + guardrail = { + "guardrail_name": "test-dynamoai-guard", + } + + with patch( + "litellm.logging_callback_manager.add_litellm_callback" + ) as mock_add: + result = initialize_guardrail(litellm_params, guardrail) + + assert isinstance(result, DynamoAIGuardrails) + assert result.api_key == "test-key" + assert result.api_base == "https://test.dynamo.ai" + assert result.guardrail_name == "test-dynamoai-guard" + mock_add.assert_called_once_with(result) + + def test_dynamoai_in_global_registry(self): + """Test that dynamoai is discoverable in the global guardrail registry.""" + from litellm.proxy.guardrails.guardrail_registry import ( + guardrail_initializer_registry, + ) + + assert "dynamoai" in guardrail_initializer_registry diff --git a/ui/litellm-dashboard/public/assets/logos/akto.svg b/ui/litellm-dashboard/public/assets/logos/akto.svg new file mode 100644 index 00000000000..cdea32535f2 --- /dev/null +++ b/ui/litellm-dashboard/public/assets/logos/akto.svg @@ -0,0 +1,10 @@ + + + + + + + + + + diff --git a/ui/litellm-dashboard/src/components/guardrails/guardrail_garden_configs.ts b/ui/litellm-dashboard/src/components/guardrails/guardrail_garden_configs.ts index 7a1b5314d33..e42ecaef579 100644 --- a/ui/litellm-dashboard/src/components/guardrails/guardrail_garden_configs.ts +++ b/ui/litellm-dashboard/src/components/guardrails/guardrail_garden_configs.ts @@ -264,4 +264,10 @@ export const GUARDRAIL_PRESETS: Record = { mode: "pre_call", defaultOn: false, }, + akto: { + provider: "Akto", + guardrailNameSuggestion: "Akto Guardrail", + mode: "pre_call", + defaultOn: false, + }, }; diff --git a/ui/litellm-dashboard/src/components/guardrails/guardrail_garden_data.ts b/ui/litellm-dashboard/src/components/guardrails/guardrail_garden_data.ts index 53ccb32c184..b06400ce508 100644 --- a/ui/litellm-dashboard/src/components/guardrails/guardrail_garden_data.ts +++ b/ui/litellm-dashboard/src/components/guardrails/guardrail_garden_data.ts @@ -373,6 +373,14 @@ export const PARTNER_GUARDRAIL_CARDS: GuardrailCardInfo[] = [ logo: `${ASSET_PREFIX}pillar.jpeg`, tags: ["Monitoring", "Safety"], }, + { + id: "akto", + name: "Akto Guardrail", + description: "AI security platform from Akto.io with automatic monitoring and guardrails for AI/ML applications.", + category: "partner", + logo: `${ASSET_PREFIX}akto.svg`, + tags: ["Security", "Safety", "Monitoring"], + }, ]; export const ALL_CARDS = [...LITELLM_CONTENT_FILTER_CARDS, ...PARTNER_GUARDRAIL_CARDS]; diff --git a/ui/litellm-dashboard/src/components/guardrails/guardrail_info_helpers.tsx b/ui/litellm-dashboard/src/components/guardrails/guardrail_info_helpers.tsx index d957be4306b..c78835dae04 100644 --- a/ui/litellm-dashboard/src/components/guardrails/guardrail_info_helpers.tsx +++ b/ui/litellm-dashboard/src/components/guardrails/guardrail_info_helpers.tsx @@ -125,6 +125,7 @@ export const guardrailLogoMap: Record = { EnkryptAI: `${asset_logos_folder}enkrypt_ai.avif`, "Prompt Security": `${asset_logos_folder}prompt_security.png`, "LiteLLM Content Filter": `${asset_logos_folder}litellm_logo.jpg`, + "Akto": `${asset_logos_folder}akto.svg`, }; export const getGuardrailLogoAndName = (guardrailValue: string): { logo: string; displayName: string } => { From 20f8d413e59098bad5410e2d737e54c6beba40bd Mon Sep 17 00:00:00 2001 From: Chesars Date: Tue, 17 Mar 2026 19:10:19 -0300 Subject: [PATCH 16/25] fix(anthropic): preserve cache_control on file-type content blocks Fixes #23873 --- .../prompt_templates/factory.py | 11 ++-- ...llm_core_utils_prompt_templates_factory.py | 53 +++++++++++++++++++ 2 files changed, 60 insertions(+), 4 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index ea1f81f9b36..82afe54b809 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -2441,11 +2441,14 @@ def anthropic_messages_pt( # noqa: PLR0915 elif m.get("type", "") == "document": user_content.append(cast(AnthropicMessagesDocumentParam, m)) elif m.get("type", "") == "file": - user_content.append( - anthropic_process_openai_file_message( - cast(ChatCompletionFileObject, m) - ) + _file_content_element = anthropic_process_openai_file_message( + cast(ChatCompletionFileObject, m) ) + _file_content_element = add_cache_control_to_content( + anthropic_content_element=_file_content_element, + original_content_element=dict(m), + ) + user_content.append(_file_content_element) elif isinstance(user_message_types_block["content"], str): _anthropic_content_text_element: AnthropicMessagesTextParam = { "type": "text", diff --git a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index 8d68539564c..438674bfb1a 100644 --- a/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/test_litellm/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -2027,3 +2027,56 @@ def test_sanitize_messages_combined_case_a_and_case_d(): ) finally: litellm.modify_params = original + + +def test_anthropic_messages_pt_file_block_preserves_cache_control(): + """ + Test that cache_control is preserved on file-type content blocks + when translated to Anthropic document params. + Regression test for https://github.com/BerriAI/litellm/issues/23873 + """ + from litellm.litellm_core_utils.prompt_templates.factory import ( + anthropic_messages_pt, + ) + + messages = [ + { + "role": "user", + "content": [ + { + "type": "file", + "file": { + "filename": "doc.pdf", + "file_data": "data:application/pdf;base64,JVBERi0xLjQ=", + }, + "cache_control": {"type": "ephemeral"}, + }, + { + "type": "text", + "text": "Summarize this document.", + "cache_control": {"type": "ephemeral"}, + }, + ], + } + ] + + result = anthropic_messages_pt( + messages, model="claude-sonnet-4-20250514", llm_provider="anthropic" + ) + + content_blocks = result[0]["content"] + assert len(content_blocks) == 2 + + # Document block (from file) should preserve cache_control + doc_block = content_blocks[0] + assert doc_block["type"] == "document" + assert "cache_control" in doc_block, ( + "cache_control was dropped from file/document block" + ) + assert doc_block["cache_control"]["type"] == "ephemeral" + + # Text block should also preserve cache_control + text_block = content_blocks[1] + assert text_block["type"] == "text" + assert "cache_control" in text_block + assert text_block["cache_control"]["type"] == "ephemeral" From 8f015e2db2ebb16598bf15e14e0e9d9d8509a2d4 Mon Sep 17 00:00:00 2001 From: Chesars Date: Tue, 17 Mar 2026 19:14:08 -0300 Subject: [PATCH 17/25] fix(vertex): respect vertex_count_tokens_location for Claude count_tokens MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The count_tokens handler unconditionally overrode vertex_location to us-central1 for Claude models, ignoring the user-configured vertex_count_tokens_location parameter. Also, us-central1 is no longer a supported region — Google now supports us-east5, europe-west1, and asia-southeast1. Now vertex_count_tokens_location takes precedence, vertex_location is used as fallback, and us-east5 is the default only when neither is set. Fixes #23872 --- .../count_tokens/handler.py | 12 +- .../count_tokens/__init__.py | 0 .../test_count_tokens_location.py | 164 ++++++++++++++++++ 3 files changed, 172 insertions(+), 4 deletions(-) create mode 100644 tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/__init__.py create mode 100644 tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/test_count_tokens_location.py diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py index c6914ac3d6b..079a691395a 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py @@ -105,12 +105,16 @@ class VertexAIPartnerModelsTokenCounter(VertexBase): # Extract Vertex AI credentials and settings vertex_credentials = self.get_vertex_ai_credentials(litellm_params) vertex_project = self.get_vertex_ai_project(litellm_params) - vertex_location = self.get_vertex_ai_location(litellm_params) + vertex_location = ( + litellm_params.get("vertex_count_tokens_location") + or self.get_vertex_ai_location(litellm_params) + ) - # Map empty location/cluade models to a supported region for count-tokens endpoint + # Default Claude models to us-east5 for count-tokens endpoint when no location is set + # Supported regions: us-east5, europe-west1, asia-southeast1 # https://docs.cloud.google.com/vertex-ai/generative-ai/docs/partner-models/claude/count-tokens - if not vertex_location or "claude" in model.lower(): - vertex_location = "us-central1" + if not vertex_location and "claude" in model.lower(): + vertex_location = "us-east5" # Get access token and resolved project ID access_token, project_id = await self._ensure_access_token_async( diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/__init__.py b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/test_count_tokens_location.py b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/test_count_tokens_location.py new file mode 100644 index 00000000000..6487ea25f21 --- /dev/null +++ b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/test_count_tokens_location.py @@ -0,0 +1,164 @@ +""" +Tests for Vertex AI partner models count_tokens location resolution. + +Ref: https://github.com/BerriAI/litellm/issues/23872 +""" +import pytest + +from litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens.handler import ( + VertexAIPartnerModelsTokenCounter, +) + + +@pytest.fixture +def counter(): + return VertexAIPartnerModelsTokenCounter() + + +class TestCountTokensLocationResolution: + """Verify that vertex_count_tokens_location is respected in handle_count_tokens_request.""" + + def _build_litellm_params( + self, + vertex_location=None, + vertex_count_tokens_location=None, + ): + params = {} + if vertex_location is not None: + params["vertex_location"] = vertex_location + if vertex_count_tokens_location is not None: + params["vertex_count_tokens_location"] = vertex_count_tokens_location + return params + + @pytest.mark.asyncio + async def test_count_tokens_location_overrides_vertex_location(self, counter, monkeypatch): + """vertex_count_tokens_location should take precedence over vertex_location.""" + captured = {} + + async def fake_ensure_access_token(self, credentials, project_id, custom_llm_provider): + return "fake-token", "fake-project" + + def fake_build_endpoint(self, model, project_id, vertex_location, api_base=None): + captured["vertex_location"] = vertex_location + return "https://fake-endpoint" + + monkeypatch.setattr( + VertexAIPartnerModelsTokenCounter, "_ensure_access_token_async", fake_ensure_access_token + ) + monkeypatch.setattr( + VertexAIPartnerModelsTokenCounter, "_build_count_tokens_endpoint", fake_build_endpoint + ) + + # Mock the HTTP call to avoid real network requests + class FakeResponse: + status_code = 200 + def json(self): + return {"input_tokens": 10} + def raise_for_status(self): + pass + + class FakeClient: + async def post(self, url, headers=None, json=None, **kwargs): + return FakeResponse() + + import litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens.handler as handler_mod + monkeypatch.setattr(handler_mod, "get_async_httpx_client", lambda **kwargs: FakeClient()) + + litellm_params = self._build_litellm_params( + vertex_location="us-east5", + vertex_count_tokens_location="europe-west1", + ) + + await counter.handle_count_tokens_request( + model="claude-sonnet-4-6", + request_data={"messages": [{"role": "user", "content": "hi"}]}, + litellm_params=litellm_params, + ) + + assert captured["vertex_location"] == "europe-west1" + + @pytest.mark.asyncio + async def test_claude_without_count_tokens_location_defaults_to_us_east5(self, counter, monkeypatch): + """Claude models without any location should default to us-east5.""" + captured = {} + + async def fake_ensure_access_token(self, credentials, project_id, custom_llm_provider): + return "fake-token", "fake-project" + + def fake_build_endpoint(self, model, project_id, vertex_location, api_base=None): + captured["vertex_location"] = vertex_location + return "https://fake-endpoint" + + monkeypatch.setattr( + VertexAIPartnerModelsTokenCounter, "_ensure_access_token_async", fake_ensure_access_token + ) + monkeypatch.setattr( + VertexAIPartnerModelsTokenCounter, "_build_count_tokens_endpoint", fake_build_endpoint + ) + + class FakeResponse: + status_code = 200 + def json(self): + return {"input_tokens": 10} + def raise_for_status(self): + pass + + class FakeClient: + async def post(self, url, headers=None, json=None, **kwargs): + return FakeResponse() + + import litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens.handler as handler_mod + monkeypatch.setattr(handler_mod, "get_async_httpx_client", lambda **kwargs: FakeClient()) + + litellm_params = self._build_litellm_params() # no location at all + + await counter.handle_count_tokens_request( + model="claude-sonnet-4-6", + request_data={"messages": [{"role": "user", "content": "hi"}]}, + litellm_params=litellm_params, + ) + + assert captured["vertex_location"] == "us-east5" + + @pytest.mark.asyncio + async def test_claude_with_vertex_location_uses_it(self, counter, monkeypatch): + """Claude models with vertex_location but no count_tokens_location should use vertex_location.""" + captured = {} + + async def fake_ensure_access_token(self, credentials, project_id, custom_llm_provider): + return "fake-token", "fake-project" + + def fake_build_endpoint(self, model, project_id, vertex_location, api_base=None): + captured["vertex_location"] = vertex_location + return "https://fake-endpoint" + + monkeypatch.setattr( + VertexAIPartnerModelsTokenCounter, "_ensure_access_token_async", fake_ensure_access_token + ) + monkeypatch.setattr( + VertexAIPartnerModelsTokenCounter, "_build_count_tokens_endpoint", fake_build_endpoint + ) + + class FakeResponse: + status_code = 200 + def json(self): + return {"input_tokens": 10} + def raise_for_status(self): + pass + + class FakeClient: + async def post(self, url, headers=None, json=None, **kwargs): + return FakeResponse() + + import litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens.handler as handler_mod + monkeypatch.setattr(handler_mod, "get_async_httpx_client", lambda **kwargs: FakeClient()) + + litellm_params = self._build_litellm_params(vertex_location="asia-southeast1") + + await counter.handle_count_tokens_request( + model="claude-sonnet-4-6", + request_data={"messages": [{"role": "user", "content": "hi"}]}, + litellm_params=litellm_params, + ) + + assert captured["vertex_location"] == "asia-southeast1" From 9afc4697258340100244e1759165790c792ebc47 Mon Sep 17 00:00:00 2001 From: Chesars Date: Tue, 17 Mar 2026 23:04:51 -0300 Subject: [PATCH 18/25] fix(mistral): preserve diarization segments in transcription response MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Fixes #23890 — Mistral's Voxtral transcription with `diarize=true` returns `segments` (with speaker_id, timestamps) and `language`, but these fields were dropped when mapping the response to TranscriptionResponse. --- .../audio_transcription/transformation.py | 7 +++ ...tral_audio_transcription_transformation.py | 44 +++++++++++++++++++ 2 files changed, 51 insertions(+) diff --git a/litellm/llms/mistral/audio_transcription/transformation.py b/litellm/llms/mistral/audio_transcription/transformation.py index 4d294063499..8c6d604acb4 100644 --- a/litellm/llms/mistral/audio_transcription/transformation.py +++ b/litellm/llms/mistral/audio_transcription/transformation.py @@ -148,5 +148,12 @@ class MistralAudioTranscriptionConfig(BaseAudioTranscriptionConfig): text = response_json.get("text") or "" response = TranscriptionResponse(text=text) + + # Preserve Mistral-specific fields (e.g. diarization segments) + if "segments" in response_json: + response["segments"] = response_json["segments"] + if "language" in response_json: + response["language"] = response_json["language"] + response._hidden_params = response_json return response diff --git a/tests/test_litellm/llms/mistral/audio_transcription/test_mistral_audio_transcription_transformation.py b/tests/test_litellm/llms/mistral/audio_transcription/test_mistral_audio_transcription_transformation.py index 7ef50dede0c..4ca3e8ae0c7 100644 --- a/tests/test_litellm/llms/mistral/audio_transcription/test_mistral_audio_transcription_transformation.py +++ b/tests/test_litellm/llms/mistral/audio_transcription/test_mistral_audio_transcription_transformation.py @@ -158,6 +158,50 @@ def test_mistral_audio_transcription_response_transform(): assert response.text == "Four score and seven years ago..." +def test_mistral_audio_transcription_response_transform_diarized(): + """Test that diarized responses preserve segments and language.""" + config = MistralAudioTranscriptionConfig() + + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = { + "model": "voxtral-mini-latest", + "text": "Hello, how are you? I am fine.", + "language": None, + "segments": [ + { + "text": "Hello, how are you?", + "start": 0.3, + "end": 2.1, + "speaker_id": "speaker_1", + "type": "transcription_segment", + }, + { + "text": "I am fine.", + "start": 2.5, + "end": 3.8, + "speaker_id": "speaker_2", + "type": "transcription_segment", + }, + ], + "usage": { + "prompt_audio_seconds": 4, + "prompt_tokens": 5, + "total_tokens": 50, + "completion_tokens": 20, + }, + } + + response = config.transform_audio_transcription_response(mock_response) + + assert isinstance(response, TranscriptionResponse) + assert response.text == "Hello, how are you? I am fine." + assert response["segments"] is not None + assert len(response["segments"]) == 2 + assert response["segments"][0]["speaker_id"] == "speaker_1" + assert response["segments"][1]["speaker_id"] == "speaker_2" + assert response["language"] is None + + def test_mistral_audio_transcription_response_transform_empty(): config = MistralAudioTranscriptionConfig() From cb15296693d5eba6ea38c2c9658c6ef37580a88c Mon Sep 17 00:00:00 2001 From: Chesars Date: Tue, 17 Mar 2026 23:06:39 -0300 Subject: [PATCH 19/25] fix(azure): auto-route gpt-5.4+ tools+reasoning to Responses API Azure GPT-5.4+ models now get the same auto-routing treatment as OpenAI when both `reasoning_effort` and `tools` are used in `litellm.completion()`. Previously, `reasoning_effort` was silently dropped for Azure; now the request is bridged to the Responses API which supports both parameters. Fixes #23914 --- docs/my-website/docs/reasoning_content.md | 4 +-- .../llms/azure/chat/gpt_5_transformation.py | 11 ++---- litellm/main.py | 23 +++++++------ .../chat/test_azure_gpt5_transformation.py | 9 ++--- tests/test_litellm/test_main.py | 34 +++++++++++++++++++ 5 files changed, 57 insertions(+), 24 deletions(-) diff --git a/docs/my-website/docs/reasoning_content.md b/docs/my-website/docs/reasoning_content.md index 8bf59f66a33..fb77f53a56c 100644 --- a/docs/my-website/docs/reasoning_content.md +++ b/docs/my-website/docs/reasoning_content.md @@ -594,9 +594,9 @@ Expected Response :::tip gpt-5.4: reasoning_effort + function tools -LiteLLM drops `reasoning_effort` from `gpt-5.4` requests to `litellm.completion()` that include tools, since that combination is supported in the Responses API. +When `gpt-5.4+` requests to `litellm.completion()` include both `reasoning_effort` and `tools`, LiteLLM **automatically routes** the request through the Responses API bridge. This works for both **OpenAI** (`openai/gpt-5.4`) and **Azure** (`azure/gpt-5.4`) providers — no extra configuration needed. -If you need reasoning **and** tools together, use `openai/responses/gpt-5.4` to route through the Responses API instead. See [Responses API Bridge](/docs/providers/openai#openai-chat-completion-to-responses-api-bridge) for details. +You can also route explicitly via `openai/responses/gpt-5.4` or `azure/responses/gpt-5.4`. See [Responses API Bridge](/docs/providers/openai#openai-chat-completion-to-responses-api-bridge) for details. ::: diff --git a/litellm/llms/azure/chat/gpt_5_transformation.py b/litellm/llms/azure/chat/gpt_5_transformation.py index 6310df9cecc..bc7483bf64d 100644 --- a/litellm/llms/azure/chat/gpt_5_transformation.py +++ b/litellm/llms/azure/chat/gpt_5_transformation.py @@ -131,14 +131,9 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config): if result_effort == "none" and not supports_none: result.pop("reasoning_effort") - # Azure Chat Completions: gpt-5.4+ does not support tools + reasoning together. - # Drop reasoning_effort when both are present (OpenAI routes to Responses API; Azure does not). - if self.is_model_gpt_5_4_plus_model(model): - has_tools = bool( - non_default_params.get("tools") or optional_params.get("tools") - ) - if has_tools and result_effort not in (None, "none"): - result.pop("reasoning_effort", None) + # Azure gpt-5.4+ with tools + reasoning_effort is now routed to the + # Responses API bridge (same as OpenAI), so we no longer need to drop + # reasoning_effort here. See: responses_api_bridge_check() in main.py. return result diff --git a/litellm/main.py b/litellm/main.py index 722b4a7aaec..e74a34a1ff5 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -955,16 +955,6 @@ def responses_api_bridge_check( model_info["mode"] = "responses" model = model.replace("responses/", "") - # OpenAI gpt-5.4+ chat-completions calls with both tools + reasoning_effort - # must be bridged to Responses API. - if ( - custom_llm_provider == "openai" - and OpenAIGPT5Config.is_model_gpt_5_4_plus_model(model) - and tools - and reasoning_effort is not None - ): - model_info["mode"] = "responses" - model = model.replace("responses/", "") except Exception as e: verbose_logger.debug("Error getting model info: {}".format(e)) @@ -974,6 +964,19 @@ def responses_api_bridge_check( model = model.replace("responses/", "") mode = "responses" model_info["mode"] = mode + + # OpenAI/Azure gpt-5.4+ chat-completions calls with both tools + reasoning_effort + # must be bridged to Responses API. + if ( + custom_llm_provider in ("openai", "azure") + and OpenAIGPT5Config.is_model_gpt_5_4_plus_model(model) + and tools + and reasoning_effort is not None + and model_info.get("mode") != "responses" + ): + model_info["mode"] = "responses" + model = model.replace("responses/", "") + return model_info, model diff --git a/tests/test_litellm/llms/azure/chat/test_azure_gpt5_transformation.py b/tests/test_litellm/llms/azure/chat/test_azure_gpt5_transformation.py index 635359563ba..28ccf7ffa8f 100644 --- a/tests/test_litellm/llms/azure/chat/test_azure_gpt5_transformation.py +++ b/tests/test_litellm/llms/azure/chat/test_azure_gpt5_transformation.py @@ -192,10 +192,11 @@ def test_azure_gpt5_1_series_temperature_handling(config: AzureOpenAIGPT5Config) assert params["temperature"] == 0.6 -def test_azure_gpt5_4_drops_reasoning_effort_when_tools_present(config: AzureOpenAIGPT5Config): - """Azure Chat Completions: gpt-5.4+ drops reasoning_effort when tools are present. +def test_azure_gpt5_4_preserves_reasoning_effort_when_tools_present(config: AzureOpenAIGPT5Config): + """Azure GPT-5.4+ no longer drops reasoning_effort when tools are present. - OpenAI routes tools+reasoning to Responses API; Azure does not, so we drop reasoning_effort. + Both OpenAI and Azure now route tools+reasoning to the Responses API bridge, + so reasoning_effort must be preserved in map_openai_params. """ tools = [{"type": "function", "function": {"name": "test", "description": "test"}}] params = config.map_openai_params( @@ -205,7 +206,7 @@ def test_azure_gpt5_4_drops_reasoning_effort_when_tools_present(config: AzureOpe drop_params=False, api_version="2024-05-01-preview", ) - assert "reasoning_effort" not in params + assert params.get("reasoning_effort") == "high" assert params["tools"] == tools diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index 6ac988b2c21..ce5873f5063 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -661,6 +661,40 @@ def test_responses_api_bridge_check_gpt_5_5_tools_plus_reasoning_routes_to_respo assert model_info.get("mode") == "responses" +def test_responses_api_bridge_check_azure_gpt_5_4_tools_plus_reasoning_routes_to_responses(): + """Azure gpt-5.4 with both tools and reasoning_effort should route to Responses API.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.4", + custom_llm_provider="azure", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort="high", + ) + + assert model == "gpt-5.4" + assert model_info.get("mode") == "responses" + + +def test_responses_api_bridge_check_azure_gpt_5_4_tools_without_reasoning_stays_chat(): + """Azure gpt-5.4 with tools only should not be force-routed to Responses API.""" + from litellm.main import responses_api_bridge_check + + with patch("litellm.main._get_model_info_helper") as mock_get_model_info: + mock_get_model_info.return_value = {"max_tokens": 128000} + model_info, model = responses_api_bridge_check( + model="gpt-5.4", + custom_llm_provider="azure", + tools=[{"type": "function", "function": {"name": "get_capital"}}], + reasoning_effort=None, + ) + + assert model == "gpt-5.4" + assert model_info.get("mode") != "responses" + + def test_responses_api_bridge_check_gpt_5_4_tools_without_reasoning_stays_chat(): """gpt-5.4 with tools only should not be force-routed to Responses API.""" from litellm.main import responses_api_bridge_check From 8828f002bea2d9842fb95fae78c90d60bc7ce41d Mon Sep 17 00:00:00 2001 From: Chesars Date: Tue, 17 Mar 2026 23:21:24 -0300 Subject: [PATCH 20/25] fix(gemini): pass model to context caching URL builder for custom api_base _get_token_and_url_context_caching() was hardcoding model=None when calling _check_custom_proxy(), which raises ValueError when api_base is set because Gemini proxy URLs need the model name: {api_base}/models/{model}:cachedContents Fixes #23846 --- .../vertex_ai_context_caching.py | 5 ++- .../test_vertex_ai_context_caching.py | 39 ++++++++++++++++++- 2 files changed, 42 insertions(+), 2 deletions(-) diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index db6be9499a2..c2e064d6561 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -51,6 +51,7 @@ class ContextCachingEndpoints(VertexBase): vertex_project: Optional[str], vertex_location: Optional[str], vertex_auth_header: Optional[str], + model: Optional[str] = None, ) -> Tuple[Optional[str], str]: """ Internal function. Returns the token and url for the call. @@ -89,7 +90,7 @@ class ContextCachingEndpoints(VertexBase): stream=None, auth_header=auth_header, url=url, - model=None, + model=model, vertex_project=vertex_project, vertex_location=vertex_location, vertex_api_version="v1beta1" @@ -342,6 +343,7 @@ class ContextCachingEndpoints(VertexBase): vertex_project=vertex_project, vertex_location=vertex_location, vertex_auth_header=vertex_auth_header, + model=model, ) headers = { @@ -488,6 +490,7 @@ class ContextCachingEndpoints(VertexBase): vertex_project=vertex_project, vertex_location=vertex_location, vertex_auth_header=vertex_auth_header, + model=model, ) headers = { diff --git a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py index 3f8cbf12361..11ccd34804a 100644 --- a/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py +++ b/tests/test_litellm/llms/vertex_ai/context_caching/test_vertex_ai_context_caching.py @@ -1317,4 +1317,41 @@ class TestVertexAIGlobalLocation: # Assert correct URL format for global with beta API expected_url = "https://aiplatform.googleapis.com/v1beta1/projects/test-project/locations/global/cachedContents" assert url == expected_url, f"Expected {expected_url}, got {url}" - assert "global-aiplatform" not in url, "URL should not contain 'global-aiplatform' prefix" \ No newline at end of file + assert "global-aiplatform" not in url, "URL should not contain 'global-aiplatform' prefix" + + def test_gemini_context_caching_with_custom_api_base_passes_model(self): + """Gemini context caching with custom api_base must pass model to _check_custom_proxy. + + Regression test for https://github.com/BerriAI/litellm/issues/23846 + Previously model was hardcoded to None, causing ValueError when api_base was set. + """ + caching = ContextCachingEndpoints() + + auth_header, url = caching._get_token_and_url_context_caching( + gemini_api_key="test-key", + custom_llm_provider="gemini", + api_base="https://my-proxy.example.com", + vertex_project=None, + vertex_location=None, + vertex_auth_header=None, + model="gemini-1.5-pro", + ) + + assert "models/gemini-1.5-pro" in url + assert url.startswith("https://my-proxy.example.com/") + + def test_gemini_context_caching_without_api_base_ignores_model(self): + """Without custom api_base, model param is not needed (default URL is used).""" + caching = ContextCachingEndpoints() + + auth_header, url = caching._get_token_and_url_context_caching( + gemini_api_key="test-key", + custom_llm_provider="gemini", + api_base=None, + vertex_project=None, + vertex_location=None, + vertex_auth_header=None, + ) + + assert "generativelanguage.googleapis.com" in url + assert "cachedContents" in url \ No newline at end of file From ff536e664aec4e6b22c7836f67434232f52b729a Mon Sep 17 00:00:00 2001 From: Chesars Date: Wed, 18 Mar 2026 00:38:35 -0300 Subject: [PATCH 21/25] fix(gemini): propagate model to check_cache/async_check_cache for custom api_base check_and_create_cache calls check_cache first (to avoid duplicates), which also needs model for the URL when api_base is set. Without this, the full flow still raises ValueError before reaching the create step. --- .../vertex_ai/context_caching/vertex_ai_context_caching.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py index c2e064d6561..b677cf3b1ec 100644 --- a/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py +++ b/litellm/llms/vertex_ai/context_caching/vertex_ai_context_caching.py @@ -110,6 +110,7 @@ class ContextCachingEndpoints(VertexBase): vertex_project: Optional[str], vertex_location: Optional[str], vertex_auth_header: Optional[str], + model: Optional[str] = None, ) -> Optional[str]: """ Checks if content already cached. @@ -129,6 +130,7 @@ class ContextCachingEndpoints(VertexBase): vertex_project=vertex_project, vertex_location=vertex_location, vertex_auth_header=vertex_auth_header, + model=model, ) page_token: Optional[str] = None @@ -202,6 +204,7 @@ class ContextCachingEndpoints(VertexBase): vertex_project: Optional[str], vertex_location: Optional[str], vertex_auth_header: Optional[str], + model: Optional[str] = None, ) -> Optional[str]: """ Checks if content already cached. @@ -221,6 +224,7 @@ class ContextCachingEndpoints(VertexBase): vertex_project=vertex_project, vertex_location=vertex_location, vertex_auth_header=vertex_auth_header, + model=model, ) page_token: Optional[str] = None @@ -379,6 +383,7 @@ class ContextCachingEndpoints(VertexBase): vertex_project=vertex_project, vertex_location=vertex_location, vertex_auth_header=vertex_auth_header, + model=model, ) if google_cache_name: return non_cached_messages, optional_params, google_cache_name @@ -523,6 +528,7 @@ class ContextCachingEndpoints(VertexBase): vertex_project=vertex_project, vertex_location=vertex_location, vertex_auth_header=vertex_auth_header, + model=model, ) if google_cache_name: From aaf860c19ba491424f16474b11d398fa1c791dc4 Mon Sep 17 00:00:00 2001 From: Chesars Date: Wed, 18 Mar 2026 00:56:11 -0300 Subject: [PATCH 22/25] docs: add Azure custom deployment name guidance for auto-routing --- docs/my-website/docs/reasoning_content.md | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) diff --git a/docs/my-website/docs/reasoning_content.md b/docs/my-website/docs/reasoning_content.md index fb77f53a56c..313228d511c 100644 --- a/docs/my-website/docs/reasoning_content.md +++ b/docs/my-website/docs/reasoning_content.md @@ -598,6 +598,23 @@ When `gpt-5.4+` requests to `litellm.completion()` include both `reasoning_effor You can also route explicitly via `openai/responses/gpt-5.4` or `azure/responses/gpt-5.4`. See [Responses API Bridge](/docs/providers/openai#openai-chat-completion-to-responses-api-bridge) for details. +**Azure custom deployment names:** Auto-routing relies on the deployment name matching the `gpt-5.4*` pattern. If you use a custom deployment name (e.g. `"my-reasoning-model"`), enable routing via: + +**SDK:** +```python +litellm.completion(model="azure/responses/my-reasoning-model", ...) +``` + +**Proxy config:** +```yaml +model_list: + - model_name: my-reasoning-model + litellm_params: + model: azure/my-reasoning-model + model_info: + mode: responses +``` + ::: ## OpenAI Responses API - Auto-Summary Control From 1284e4ebe599cb1585f634c25ed1b82b2a9f8ab6 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 19 Mar 2026 16:30:02 +0530 Subject: [PATCH 23/25] Fix cicd fialing tests --- .../vertex_ai_partner_models/count_tokens/handler.py | 11 +++++------ .../chat/test_fireworks_ai_chat_transformation.py | 6 +++--- 2 files changed, 8 insertions(+), 9 deletions(-) diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py index 079a691395a..ceb924b9b0f 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py @@ -105,16 +105,15 @@ class VertexAIPartnerModelsTokenCounter(VertexBase): # Extract Vertex AI credentials and settings vertex_credentials = self.get_vertex_ai_credentials(litellm_params) vertex_project = self.get_vertex_ai_project(litellm_params) - vertex_location = ( - litellm_params.get("vertex_count_tokens_location") - or self.get_vertex_ai_location(litellm_params) - ) + vertex_location_raw = self.get_vertex_ai_location(litellm_params) # Default Claude models to us-east5 for count-tokens endpoint when no location is set # Supported regions: us-east5, europe-west1, asia-southeast1 # https://docs.cloud.google.com/vertex-ai/generative-ai/docs/partner-models/claude/count-tokens - if not vertex_location and "claude" in model.lower(): - vertex_location = "us-east5" + if not vertex_location_raw or "claude" in model.lower(): + vertex_location: str = "us-central1" + else: + vertex_location = vertex_location_raw # Get access token and resolved project ID access_token, project_id = await self._ensure_access_token_async( diff --git a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py index 2b71b88356b..29265bb4b42 100644 --- a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py +++ b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py @@ -122,21 +122,21 @@ def test_add_transform_inline_image_block_skips_data_urls(): # str branch str_content = {"type": "image_url", "image_url": data_url} result = config._add_transform_inline_image_block( - str_content, model="non-vision-model", disable_add_transform_inline_image_block=False + str_content, model="gpt-4", disable_add_transform_inline_image_block=False ) assert result["image_url"] == data_url, "data URL must not be modified (str branch)" # dict branch dict_content = {"type": "image_url", "image_url": {"url": data_url}} result = config._add_transform_inline_image_block( - dict_content, model="non-vision-model", disable_add_transform_inline_image_block=False + dict_content, model="gpt-4", disable_add_transform_inline_image_block=False ) assert result["image_url"]["url"] == data_url, "data URL must not be modified (dict branch)" # regular https URL should still get the suffix https_content = {"type": "image_url", "image_url": "https://example.com/image.jpg"} result = config._add_transform_inline_image_block( - https_content, model="non-vision-model", disable_add_transform_inline_image_block=False + https_content, model="gpt-4", disable_add_transform_inline_image_block=False ) assert result["image_url"].endswith("#transform=inline"), "https URL should get #transform=inline" @pytest.mark.parametrize( From 6146196c6ad938bdfaab41cb20b9c6a90f36b5a4 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Thu, 19 Mar 2026 16:43:39 +0530 Subject: [PATCH 24/25] Fix tests --- .../count_tokens/handler.py | 20 ++++++++++++++----- 1 file changed, 15 insertions(+), 5 deletions(-) diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py index ceb924b9b0f..82076ff3600 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/count_tokens/handler.py @@ -105,15 +105,25 @@ class VertexAIPartnerModelsTokenCounter(VertexBase): # Extract Vertex AI credentials and settings vertex_credentials = self.get_vertex_ai_credentials(litellm_params) vertex_project = self.get_vertex_ai_project(litellm_params) + + # Check for count_tokens specific location override + vertex_count_tokens_location = litellm_params.get("vertex_count_tokens_location") vertex_location_raw = self.get_vertex_ai_location(litellm_params) - - # Default Claude models to us-east5 for count-tokens endpoint when no location is set + + # Determine final location with precedence: + # 1. vertex_count_tokens_location (if provided) + # 2. vertex_location (if provided) + # 3. Default to us-east5 for Claude models when no location is set # Supported regions: us-east5, europe-west1, asia-southeast1 # https://docs.cloud.google.com/vertex-ai/generative-ai/docs/partner-models/claude/count-tokens - if not vertex_location_raw or "claude" in model.lower(): - vertex_location: str = "us-central1" - else: + if vertex_count_tokens_location: + vertex_location: str = vertex_count_tokens_location + elif vertex_location_raw: vertex_location = vertex_location_raw + elif "claude" in model.lower(): + vertex_location = "us-east5" + else: + vertex_location = "us-east5" # Get access token and resolved project ID access_token, project_id = await self._ensure_access_token_async( From 509d2e9ac3b3bbdc4a2c40840029d067857a66cc Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Sat, 21 Mar 2026 11:30:29 -0700 Subject: [PATCH 25/25] Fix PR review issues: gpt-4-0314 prompt caching, case-insensitive data URL check, test I/O mocking - Remove incorrect supports_prompt_caching from gpt-4-0314 (predates the feature) - Make data-URL detection case-insensitive in Gemini tool call result conversion - Mock show_banner/generate_feedback_box in max_budget tests to prevent real I/O Co-Authored-By: Claude Opus 4.6 --- .../prompt_templates/factory.py | 2 +- ...odel_prices_and_context_window_backup.json | 32 ++++++++++++--- model_prices_and_context_window.json | 1 - .../proxy/test_max_budget_env_var.py | 41 ++++++++++++------- 4 files changed, 54 insertions(+), 22 deletions(-) diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 0e1fa05b9d4..d29ca1649ff 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -1506,7 +1506,7 @@ def convert_to_gemini_tool_call_result( # noqa: PLR0915 # Detect data-URL images (e.g. from Anthropic tool_result with a single image block # that was serialised as a plain string by translate_anthropic_messages_to_openai) # and promote them to inline_data so Gemini receives actual image bytes. - if content_str.startswith("data:") and ";base64," in content_str: + if content_str[:5].lower() == "data:" and ";base64," in content_str: try: mime_rest = content_str[5:].split(";base64,", 1) if len(mime_rest) == 2 and mime_rest[0].startswith("image/"): diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 879dd42be47..b2fabb4936f 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -6152,7 +6152,8 @@ "max_query_tokens": 4096, "max_tokens": 32768, "mode": "rerank", - "output_cost_per_token": 0.0 + "output_cost_per_token": 0.0, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-cohere-rerank-4-0-in-microsoft-foundry/4477076" }, "azure_ai/cohere-rerank-v4.0-fast": { "input_cost_per_query": 0.002, @@ -6163,7 +6164,8 @@ "max_query_tokens": 4096, "max_tokens": 32768, "mode": "rerank", - "output_cost_per_token": 0.0 + "output_cost_per_token": 0.0, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-cohere-rerank-4-0-in-microsoft-foundry/4477076" }, "azure_ai/deepseek-v3.2": { "input_cost_per_token": 5.8e-07, @@ -6173,6 +6175,7 @@ "max_tokens": 163840, "mode": "chat", "output_cost_per_token": 1.68e-06, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-deepseek-v3-2-and-deepseek-v3-2-speciale-in-microsoft-foundry/4477549", "supports_assistant_prefill": true, "supports_function_calling": true, "supports_prompt_caching": true, @@ -6187,6 +6190,7 @@ "max_tokens": 163840, "mode": "chat", "output_cost_per_token": 1.68e-06, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-deepseek-v3-2-and-deepseek-v3-2-speciale-in-microsoft-foundry/4477549", "supports_assistant_prefill": true, "supports_function_calling": true, "supports_prompt_caching": true, @@ -16936,6 +16940,18 @@ "supports_system_messages": true, "supports_tool_choice": true }, + "gpt-4-0314": { + "deprecation_date": "2026-03-26", + "input_cost_per_token": 3e-05, + "litellm_provider": "openai", + "max_input_tokens": 8192, + "max_output_tokens": 4096, + "max_tokens": 4096, + "mode": "chat", + "output_cost_per_token": 6e-05, + "supports_system_messages": true, + "supports_tool_choice": true + }, "gpt-4-0613": { "deprecation_date": "2025-06-06", "input_cost_per_token": 3e-05, @@ -30506,7 +30522,7 @@ "output_cost_per_token": 5.4e-06, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models", "supported_regions": [ - "us-west2" + "us-central1" ], "supports_assistant_prefill": true, "supports_function_calling": true, @@ -30526,7 +30542,7 @@ "output_cost_per_token_batches": 8.4e-07, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models", "supported_regions": [ - "us-west2" + "global" ], "supports_assistant_prefill": true, "supports_function_calling": true, @@ -30543,6 +30559,9 @@ "mode": "chat", "output_cost_per_token": 5.4e-06, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing#partner-models", + "supported_regions": [ + "us-central1" + ], "supports_assistant_prefill": true, "supports_function_calling": true, "supports_prompt_caching": true, @@ -31167,7 +31186,10 @@ "input_cost_per_token": 3e-07, "output_cost_per_token": 1.2e-06, "ocr_cost_per_page": 0.0003, - "source": "https://cloud.google.com/vertex-ai/pricing" + "source": "https://cloud.google.com/vertex-ai/pricing", + "supported_regions": [ + "us-central1" + ] }, "vertex_ai/openai/gpt-oss-120b-maas": { "input_cost_per_token": 1.5e-07, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index ec96dcb2152..b2fabb4936f 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -16949,7 +16949,6 @@ "max_tokens": 4096, "mode": "chat", "output_cost_per_token": 6e-05, - "supports_prompt_caching": true, "supports_system_messages": true, "supports_tool_choice": true }, diff --git a/tests/test_litellm/proxy/test_max_budget_env_var.py b/tests/test_litellm/proxy/test_max_budget_env_var.py index 4b965392823..90dfb81f3ae 100644 --- a/tests/test_litellm/proxy/test_max_budget_env_var.py +++ b/tests/test_litellm/proxy/test_max_budget_env_var.py @@ -4,10 +4,11 @@ converted to float. GitHub Issue: #23843 """ +from unittest.mock import patch + import pytest import litellm -from litellm.proxy.proxy_server import initialize @pytest.mark.asyncio @@ -17,22 +18,32 @@ async def test_max_budget_string_converted_to_float(): string. initialize() should convert it to float so the comparison `litellm.max_budget > 0` doesn't raise TypeError. """ - original = litellm.max_budget - try: - await initialize(max_budget="100.5") - assert isinstance(litellm.max_budget, float) - assert litellm.max_budget == 100.5 - finally: - litellm.max_budget = original + with patch("litellm.proxy.common_utils.banner.show_banner"), patch( + "litellm.proxy.proxy_server.generate_feedback_box" + ): + from litellm.proxy.proxy_server import initialize + + original = litellm.max_budget + try: + await initialize(max_budget="100.5") + assert isinstance(litellm.max_budget, float) + assert litellm.max_budget == 100.5 + finally: + litellm.max_budget = original @pytest.mark.asyncio async def test_max_budget_float_stays_float(): """max_budget as float should still work.""" - original = litellm.max_budget - try: - await initialize(max_budget=200.0) - assert isinstance(litellm.max_budget, float) - assert litellm.max_budget == 200.0 - finally: - litellm.max_budget = original + with patch("litellm.proxy.common_utils.banner.show_banner"), patch( + "litellm.proxy.proxy_server.generate_feedback_box" + ): + from litellm.proxy.proxy_server import initialize + + original = litellm.max_budget + try: + await initialize(max_budget=200.0) + assert isinstance(litellm.max_budget, float) + assert litellm.max_budget == 200.0 + finally: + litellm.max_budget = original