From acdf9b64d9cb9d06325ed7fc5ba4f9fba6200c0e Mon Sep 17 00:00:00 2001 From: Shubham Pathak Date: Mon, 29 Sep 2025 12:04:44 +0530 Subject: [PATCH 01/28] Update request handling for original exceptions --- .../exception_mapping_utils.py | 20 +++++++++---------- 1 file changed, 10 insertions(+), 10 deletions(-) diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index 409f5ebe0bc..f9cebeb1765 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -1498,7 +1498,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"CohereException - {original_exception.message}", llm_provider="cohere", model=model, - request=original_exception.request, + request=getattr(original_exception, "request", None),, ) raise original_exception elif custom_llm_provider == "huggingface": @@ -1573,7 +1573,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"HuggingfaceException - {original_exception.message}", llm_provider="huggingface", model=model, - request=original_exception.request, + request=getattr(original_exception, "request", None),, ) elif custom_llm_provider == "ai21": if hasattr(original_exception, "message"): @@ -1632,7 +1632,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"AI21Exception - {original_exception.message}", llm_provider="ai21", model=model, - request=original_exception.request, + request=getattr(original_exception, "request", None),, ) elif custom_llm_provider == "nlp_cloud": if "detail" in error_str: @@ -1659,7 +1659,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"NLPCloudException - {error_str}", model=model, llm_provider="nlp_cloud", - request=original_exception.request, + request=getattr(original_exception, "request", None),, ) if hasattr( original_exception, "status_code" @@ -1719,7 +1719,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"NLPCloudException - {original_exception.message}", llm_provider="nlp_cloud", model=model, - request=original_exception.request, + request=getattr(original_exception, "request", None),, ) elif ( original_exception.status_code == 504 @@ -1739,7 +1739,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"NLPCloudException - {original_exception.message}", llm_provider="nlp_cloud", model=model, - request=original_exception.request, + request=getattr(original_exception, "request", None),, ) elif custom_llm_provider == "together_ai": try: @@ -1848,7 +1848,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"TogetherAIException - {original_exception.message}", llm_provider="together_ai", model=model, - request=original_exception.request, + request=getattr(original_exception, "request", None),, ) elif custom_llm_provider == "aleph_alpha": if ( @@ -1953,7 +1953,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"VLLMException - {original_exception.message}", llm_provider="vllm", model=model, - request=original_exception.request, + request=getattr(original_exception, "request", None),, ) elif custom_llm_provider == "azure" or custom_llm_provider == "azure_text": message = get_error_message(error_obj=original_exception) @@ -2208,7 +2208,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"APIError: {exception_provider} - {error_str}", llm_provider=custom_llm_provider, model=model, - request=original_exception.request, + request=getattr(original_exception, "request", None),, litellm_debug_info=extra_information, ) else: @@ -2243,7 +2243,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message="{} - {}".format(exception_provider, error_str), llm_provider=custom_llm_provider, model=model, - request=original_exception.request, + request=getattr(original_exception, "request", None),, ) else: raise APIConnectionError( From b8195f091bda03d0011977811c8649d2402bf446 Mon Sep 17 00:00:00 2001 From: Shubham Pathak Date: Mon, 29 Sep 2025 12:19:00 +0530 Subject: [PATCH 02/28] Fixed formatting --- .../exception_mapping_utils.py | 20 +++++++++---------- 1 file changed, 10 insertions(+), 10 deletions(-) diff --git a/litellm/litellm_core_utils/exception_mapping_utils.py b/litellm/litellm_core_utils/exception_mapping_utils.py index f9cebeb1765..c6d3637ffcb 100644 --- a/litellm/litellm_core_utils/exception_mapping_utils.py +++ b/litellm/litellm_core_utils/exception_mapping_utils.py @@ -1498,7 +1498,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"CohereException - {original_exception.message}", llm_provider="cohere", model=model, - request=getattr(original_exception, "request", None),, + request=getattr(original_exception, "request", None), ) raise original_exception elif custom_llm_provider == "huggingface": @@ -1573,7 +1573,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"HuggingfaceException - {original_exception.message}", llm_provider="huggingface", model=model, - request=getattr(original_exception, "request", None),, + request=getattr(original_exception, "request", None), ) elif custom_llm_provider == "ai21": if hasattr(original_exception, "message"): @@ -1632,7 +1632,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"AI21Exception - {original_exception.message}", llm_provider="ai21", model=model, - request=getattr(original_exception, "request", None),, + request=getattr(original_exception, "request", None), ) elif custom_llm_provider == "nlp_cloud": if "detail" in error_str: @@ -1659,7 +1659,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"NLPCloudException - {error_str}", model=model, llm_provider="nlp_cloud", - request=getattr(original_exception, "request", None),, + request=getattr(original_exception, "request", None), ) if hasattr( original_exception, "status_code" @@ -1719,7 +1719,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"NLPCloudException - {original_exception.message}", llm_provider="nlp_cloud", model=model, - request=getattr(original_exception, "request", None),, + request=getattr(original_exception, "request", None), ) elif ( original_exception.status_code == 504 @@ -1739,7 +1739,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"NLPCloudException - {original_exception.message}", llm_provider="nlp_cloud", model=model, - request=getattr(original_exception, "request", None),, + request=getattr(original_exception, "request", None), ) elif custom_llm_provider == "together_ai": try: @@ -1848,7 +1848,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"TogetherAIException - {original_exception.message}", llm_provider="together_ai", model=model, - request=getattr(original_exception, "request", None),, + request=getattr(original_exception, "request", None), ) elif custom_llm_provider == "aleph_alpha": if ( @@ -1953,7 +1953,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"VLLMException - {original_exception.message}", llm_provider="vllm", model=model, - request=getattr(original_exception, "request", None),, + request=getattr(original_exception, "request", None), ) elif custom_llm_provider == "azure" or custom_llm_provider == "azure_text": message = get_error_message(error_obj=original_exception) @@ -2208,7 +2208,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message=f"APIError: {exception_provider} - {error_str}", llm_provider=custom_llm_provider, model=model, - request=getattr(original_exception, "request", None),, + request=getattr(original_exception, "request", None), litellm_debug_info=extra_information, ) else: @@ -2243,7 +2243,7 @@ def exception_type( # type: ignore # noqa: PLR0915 message="{} - {}".format(exception_provider, error_str), llm_provider=custom_llm_provider, model=model, - request=getattr(original_exception, "request", None),, + request=getattr(original_exception, "request", None), ) else: raise APIConnectionError( From 58955a034890b3bf8072a1f23a500685a1250b56 Mon Sep 17 00:00:00 2001 From: Sameerlite Date: Mon, 29 Sep 2025 16:32:21 +0530 Subject: [PATCH 03/28] Ignore type param for gemini tools --- .../vertex_and_google_ai_studio_gemini.py | 59 ++++++++++--------- .../llms/vertex_ai/test_vertex.py | 55 +++++++++++++++++ 2 files changed, 87 insertions(+), 27 deletions(-) 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 dc3a6cf15e5..655da5c87a7 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 @@ -3,7 +3,6 @@ ## Initial implementation - covers gemini + image gen calls import json import time -from litellm._uuid import uuid from copy import deepcopy from functools import partial from typing import ( @@ -25,6 +24,7 @@ import litellm import litellm.litellm_core_utils import litellm.litellm_core_utils.litellm_logging from litellm import verbose_logger +from litellm._uuid import uuid from litellm.constants import ( DEFAULT_REASONING_EFFORT_DISABLE_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET, @@ -32,8 +32,8 @@ from litellm.constants import ( DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH, - DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO, DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE, + DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO, ) from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.llms.custom_httpx.http_handler import ( @@ -313,9 +313,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): return None for tool in value: - openai_function_object: Optional[ChatCompletionToolParamFunctionChunk] = ( - None - ) + openai_function_object: Optional[ + ChatCompletionToolParamFunctionChunk + ] = None if "function" in tool: # tools list _openai_function_object = ChatCompletionToolParamFunctionChunk( # type: ignore **tool["function"] @@ -335,6 +335,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): elif "name" in tool: # functions list openai_function_object = ChatCompletionToolParamFunctionChunk(**tool) # type: ignore + # Handle tools with 'type' field (OpenAI spec compliance) Ignore this field -> https://github.com/BerriAI/litellm/issues/14644#issuecomment-3342061838 + if "type" in tool: + del tool["type"] # type: ignore + tool_name = list(tool.keys())[0] if len(tool.keys()) == 1 else None if tool_name and ( tool_name == "codeExecution" or tool_name == "code_execution" @@ -437,7 +441,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): elif model and "gemini-2.5-pro" in model.lower(): budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO elif model and "gemini-2.5-flash" in model.lower(): - budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH + budget = ( + DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH + ) else: budget = DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET @@ -621,16 +627,16 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): elif param == "seed": optional_params["seed"] = value elif param == "reasoning_effort" and isinstance(value, str): - optional_params["thinkingConfig"] = ( - VertexGeminiConfig._map_reasoning_effort_to_thinking_budget( - value, model - ) + optional_params[ + "thinkingConfig" + ] = VertexGeminiConfig._map_reasoning_effort_to_thinking_budget( + value, model ) elif param == "thinking": - optional_params["thinkingConfig"] = ( - VertexGeminiConfig._map_thinking_param( - cast(AnthropicThinkingParam, value) - ) + optional_params[ + "thinkingConfig" + ] = VertexGeminiConfig._map_thinking_param( + cast(AnthropicThinkingParam, value) ) elif param == "modalities" and isinstance(value, list): response_modalities = self.map_response_modalities(value) @@ -1066,7 +1072,6 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): GenerateContentResponseBody, BidiGenerateContentServerMessage ], ) -> Usage: - if ( completion_response is not None and "usageMetadata" not in completion_response @@ -1502,28 +1507,28 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): ## ADD METADATA TO RESPONSE ## setattr(model_response, "vertex_ai_grounding_metadata", grounding_metadata) - model_response._hidden_params["vertex_ai_grounding_metadata"] = ( - grounding_metadata - ) + model_response._hidden_params[ + "vertex_ai_grounding_metadata" + ] = grounding_metadata setattr( model_response, "vertex_ai_url_context_metadata", url_context_metadata ) - model_response._hidden_params["vertex_ai_url_context_metadata"] = ( - url_context_metadata - ) + model_response._hidden_params[ + "vertex_ai_url_context_metadata" + ] = url_context_metadata setattr(model_response, "vertex_ai_safety_results", safety_ratings) - model_response._hidden_params["vertex_ai_safety_results"] = ( - safety_ratings # older approach - maintaining to prevent regressions - ) + model_response._hidden_params[ + "vertex_ai_safety_results" + ] = safety_ratings # older approach - maintaining to prevent regressions ## ADD CITATION METADATA ## setattr(model_response, "vertex_ai_citation_metadata", citation_metadata) - model_response._hidden_params["vertex_ai_citation_metadata"] = ( - citation_metadata # older approach - maintaining to prevent regressions - ) + model_response._hidden_params[ + "vertex_ai_citation_metadata" + ] = citation_metadata # older approach - maintaining to prevent regressions except Exception as e: raise VertexAIError( diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex.py b/tests/test_litellm/llms/vertex_ai/test_vertex.py index 7e683d1f54e..31d4fd1c198 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex.py @@ -199,6 +199,61 @@ def test_vertex_function_translation(tool, expect_parameters): ) +def test_vertex_tool_type_field_removal(): + """ + Test that the 'type' field is removed from tools during processing + to avoid issues with Vertex AI API while maintaining functionality. + """ + # Test with Google Search tool that has 'type' field + tools_with_type = [{"type": "google_search", "googleSearch": {}}] + + optional_params = get_optional_params( + model="gemini-1.5-pro", + custom_llm_provider="vertex_ai", + tools=tools_with_type, + ) + + # Verify the tool is processed correctly + assert "tools" in optional_params + assert len(optional_params["tools"]) == 1 + assert "googleSearch" in optional_params["tools"][0] + assert optional_params["tools"][0]["googleSearch"] == {} + + # Verify the 'type' field is not present in the final result + assert "type" not in optional_params["tools"][0] + + # Test with function tool that has 'type' field + function_tools_with_type = [ + { + "type": "function", + "function": { + "name": "test_function", + "description": "A test function", + "parameters": { + "type": "object", + "properties": {"param": {"type": "string"}} + } + } + } + ] + + optional_params_function = get_optional_params( + model="gemini-1.5-pro", + custom_llm_provider="vertex_ai", + tools=function_tools_with_type, + ) + + # Verify function tool is processed correctly + assert "tools" in optional_params_function + assert len(optional_params_function["tools"]) == 1 + assert "function_declarations" in optional_params_function["tools"][0] + assert len(optional_params_function["tools"][0]["function_declarations"]) == 1 + assert optional_params_function["tools"][0]["function_declarations"][0]["name"] == "test_function" + + # Verify the 'type' field is not present in the final result + assert "type" not in optional_params_function["tools"][0] + + def test_function_calling_with_gemini(): from litellm.llms.custom_httpx.http_handler import HTTPHandler From a44b9ebb3c80edd698cb543b014bbe632fac290a Mon Sep 17 00:00:00 2001 From: Ihsan Soydemir Date: Mon, 29 Sep 2025 14:06:10 +0200 Subject: [PATCH 04/28] fix(files): use extra_query for GET/DELETE in Files endpoints --- litellm/proxy/openai_files_endpoints/files_endpoints.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/litellm/proxy/openai_files_endpoints/files_endpoints.py b/litellm/proxy/openai_files_endpoints/files_endpoints.py index f2a7ccccdf9..f5de78832bc 100644 --- a/litellm/proxy/openai_files_endpoints/files_endpoints.py +++ b/litellm/proxy/openai_files_endpoints/files_endpoints.py @@ -31,6 +31,7 @@ from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing from litellm.proxy.common_utils.openai_endpoint_utils import ( get_custom_llm_provider_from_request_body, + get_custom_llm_provider_from_request_query, ) from litellm.proxy.utils import ProxyLogging, is_known_model from litellm.router import Router @@ -237,6 +238,7 @@ async def create_file( file_content = await file.read() custom_llm_provider = ( provider + or get_custom_llm_provider_from_request_query(request=request) or await get_custom_llm_provider_from_request_body(request=request) or "openai" ) @@ -425,6 +427,7 @@ async def get_file_content( custom_llm_provider = ( provider + or get_custom_llm_provider_from_request_query(request=request) or await get_custom_llm_provider_from_request_body(request=request) or "openai" ) @@ -591,6 +594,7 @@ async def get_file( try: custom_llm_provider = ( provider + or get_custom_llm_provider_from_request_query(request=request) or await get_custom_llm_provider_from_request_body(request=request) or "openai" ) @@ -733,6 +737,7 @@ async def delete_file( try: custom_llm_provider = ( provider + or get_custom_llm_provider_from_request_query(request=request) or await get_custom_llm_provider_from_request_body(request=request) or "openai" ) @@ -917,6 +922,7 @@ async def list_files( else: custom_llm_provider = ( provider + or get_custom_llm_provider_from_request_query(request=request) or await get_custom_llm_provider_from_request_body(request=request) or "openai" ) From 3f312d0527ba95158c935ba7756169d3eb3342cb Mon Sep 17 00:00:00 2001 From: Shubham Pathak Date: Mon, 29 Sep 2025 21:43:31 +0530 Subject: [PATCH 05/28] Added test --- ...est_exception_mapping_request_attribute.py | 269 ++++++++++++++++++ 1 file changed, 269 insertions(+) create mode 100644 tests/test_litellm/test_exception_mapping_request_attribute.py diff --git a/tests/test_litellm/test_exception_mapping_request_attribute.py b/tests/test_litellm/test_exception_mapping_request_attribute.py new file mode 100644 index 00000000000..aeb77527157 --- /dev/null +++ b/tests/test_litellm/test_exception_mapping_request_attribute.py @@ -0,0 +1,269 @@ + +""" +Unit tests for the exception mapping request attribute handling fix. + +This test verifies the fix for PR #15013 where getattr(original_exception, "request", None) +is used instead of original_exception.request to handle cases where exceptions don't have +a request attribute. + +The key fix is that accessing original_exception.request directly would raise AttributeError +if the exception doesn't have a request attribute, but getattr(original_exception, "request", None) +safely returns None instead. + +PR #15013 fixed 12 locations in exception_mapping_utils.py where direct access to .request +was replaced with getattr() calls: +- Line 1501: Cohere exception mapping +- Line 1574: HuggingFace exception mapping +- Line 1635: AI21 exception mapping +- Line 1660: NLP Cloud exception mapping +- Line 1720: NLP Cloud exception mapping (another case) +- Line 1740: NLP Cloud exception mapping (another case) +- Line 1851: Together AI exception mapping +- Line 1954: VLLM exception mapping +- Line 2209: Generic provider exception mapping +- Line 2244: Generic provider exception mapping (fallback) +- OpenRouter exception mapping (multiple locations) + +This test ensures that none of these code paths will raise AttributeError when an exception +object doesn't have a request attribute, which was the root cause of the bug. +""" + +import pytest +import httpx +from unittest.mock import patch + +import litellm +from litellm.litellm_core_utils.exception_mapping_utils import exception_type +from litellm.exceptions import APIError, APIConnectionError + + +class MockExceptionWithoutRequest: + """Mock exception that does NOT have a request attribute.""" + + def __init__(self, status_code=500, message="Test error"): + self.status_code = status_code + self.message = message + # Intentionally no request attribute + + +def test_exception_mapping_request_attribute_fix(): + """ + Test the core fix: getattr(original_exception, "request", None) should not raise AttributeError + even when the exception doesn't have a request attribute. + + This is the main test for PR #15013. + """ + + # Test case 1: Exception without request attribute should not cause AttributeError + mock_exception = MockExceptionWithoutRequest( + status_code=500, + message="Test error without request attribute" + ) + + # The test is that this should NOT raise an AttributeError about missing 'request' + try: + exception_type( + model="test-model", + custom_llm_provider="cohere", # Using cohere as it's one of the affected providers + original_exception=mock_exception, + completion_kwargs={}, + extra_kwargs={} + ) + # We expect some exception to be raised (the mapped exception), but not AttributeError + except AttributeError as e: + if "'request'" in str(e): + pytest.fail(f"The fix failed: Should not raise AttributeError about missing 'request' attribute: {e}") + else: + # If it's a different AttributeError, re-raise it + raise + except Exception: + # Any other exception is fine - we just want to ensure no AttributeError about 'request' + pass + + +def test_request_attribute_safety_with_getattr(): + """ + Test that the getattr approach works correctly for both cases: + 1. When request attribute exists + 2. When request attribute doesn't exist + """ + + # Case 1: Exception with request attribute + class MockExceptionWithRequest: + def __init__(self): + self.status_code = 500 + self.message = "Test error" + self.request = httpx.Request(method="POST", url="https://api.example.com") + + exception_with_request = MockExceptionWithRequest() + request_value = getattr(exception_with_request, "request", None) + assert request_value is not None + assert isinstance(request_value, httpx.Request) + + # Case 2: Exception without request attribute + exception_without_request = MockExceptionWithoutRequest() + request_value = getattr(exception_without_request, "request", None) + assert request_value is None # Should be None, not raise AttributeError + + +def test_providers_affected_by_fix(): + """ + Test that the specific providers mentioned in the PR changes handle missing request attributes correctly. + + The PR changes affected these provider-specific code paths: + - cohere: line 1501 + - huggingface: line 1574 + - ai21: line 1635 + - nlp_cloud: lines 1660, 1720, 1740 + - together_ai: line 1851 + - vllm: line 1954 + - generic providers: lines 2209, 2244 + """ + + providers_to_test = [ + "cohere", + "ai21", + "together_ai", + "vllm" + ] + + for provider in providers_to_test: + mock_exception = MockExceptionWithoutRequest( + status_code=500, + message=f"Test error for {provider}" + ) + + # The key test: this should not raise AttributeError about missing 'request' + try: + exception_type( + model=f"{provider}-test-model", + custom_llm_provider=provider, + original_exception=mock_exception, + completion_kwargs={}, + extra_kwargs={} + ) + except AttributeError as e: + if "'request'" in str(e): + pytest.fail(f"Provider {provider} failed: Should not raise AttributeError about missing 'request' attribute: {e}") + except Exception: + # Any other exception is expected and fine + pass + + +def test_huggingface_specific_case(): + """ + Test HuggingFace specific case which has its own handling logic. + """ + mock_exception = MockExceptionWithoutRequest( + status_code=400, + message="length limit exceeded" + ) + + try: + exception_type( + model="huggingface-model", + custom_llm_provider="huggingface", + original_exception=mock_exception, + completion_kwargs={}, + extra_kwargs={} + ) + except AttributeError as e: + if "'request'" in str(e): + pytest.fail(f"HuggingFace exception handling failed: Should not raise AttributeError about missing 'request' attribute: {e}") + except litellm.ContextWindowExceededError: + # Expected for "length limit exceeded" message + pass + except Exception: + # Other exceptions are fine + pass + + +def test_nlp_cloud_specific_case(): + """ + Test NLP Cloud specific case which had multiple lines changed in the PR. + """ + mock_exception = MockExceptionWithoutRequest( + status_code=504, + message="Gateway timeout" + ) + + try: + exception_type( + model="nlp-cloud-model", + custom_llm_provider="nlp_cloud", + original_exception=mock_exception, + completion_kwargs={}, + extra_kwargs={} + ) + except AttributeError as e: + if "'request'" in str(e): + pytest.fail(f"NLP Cloud exception handling failed: Should not raise AttributeError about missing 'request' attribute: {e}") + except Exception: + # Any other exception is expected + pass + + +def test_generic_fallback_case(): + """ + Test the generic fallback case at the end of exception_type function. + This tests the changes in lines 2209 and 2244 of the PR. + """ + mock_exception = MockExceptionWithoutRequest( + status_code=500, + message="Generic error" + ) + + try: + exception_type( + model="unknown-model", + custom_llm_provider="unknown_provider", + original_exception=mock_exception, + completion_kwargs={}, + extra_kwargs={} + ) + except AttributeError as e: + if "'request'" in str(e): + pytest.fail(f"Generic fallback failed: Should not raise AttributeError about missing 'request' attribute: {e}") + except APIConnectionError: + # Expected for generic fallback + pass + except Exception: + # Other exceptions might be fine too + pass + + +def test_openrouter_specific_case(): + """ + Test OpenRouter which also uses the request attribute in exception mapping. + """ + mock_exception = MockExceptionWithoutRequest( + status_code=500, + message="OpenRouter error" + ) + + try: + exception_type( + model="openrouter-model", + custom_llm_provider="openrouter", + original_exception=mock_exception, + completion_kwargs={}, + extra_kwargs={} + ) + except AttributeError as e: + if "'request'" in str(e): + pytest.fail(f"OpenRouter exception handling failed: Should not raise AttributeError about missing 'request' attribute: {e}") + except Exception: + # Other exceptions are expected + pass + + +if __name__ == "__main__": + # Run tests for manual verification + test_exception_mapping_request_attribute_fix() + test_request_attribute_safety_with_getattr() + test_providers_affected_by_fix() + test_huggingface_specific_case() + test_nlp_cloud_specific_case() + test_generic_fallback_case() + test_openrouter_specific_case() + print("All tests passed!") From ff2d19e4cab88c7ce0a417f47e5ff7be85da515a Mon Sep 17 00:00:00 2001 From: Zero Clover <13190004+ZeroClover@users.noreply.github.com> Date: Tue, 30 Sep 2025 03:12:32 +0800 Subject: [PATCH 06/28] feat: improve vertex AI/gemini api_base handling for proxy services (#15039) --- litellm/llms/vertex_ai/vertex_llm_base.py | 165 ++++++++-------- .../llms/vertex_ai/test_vertex_llm_base.py | 184 ++++++++++-------- 2 files changed, 193 insertions(+), 156 deletions(-) diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py index 6d194d41add..6769bc3fb28 100644 --- a/litellm/llms/vertex_ai/vertex_llm_base.py +++ b/litellm/llms/vertex_ai/vertex_llm_base.py @@ -46,14 +46,10 @@ class VertexBase: return "global" return vertex_region or "us-central1" - def load_auth( - self, credentials: Optional[VERTEX_CREDENTIALS_TYPES], project_id: Optional[str] - ) -> Tuple[Any, str]: + def load_auth(self, credentials: Optional[VERTEX_CREDENTIALS_TYPES], project_id: Optional[str]) -> Tuple[Any, str]: if credentials is not None: if isinstance(credentials, str): - verbose_logger.debug( - "Vertex: Loading vertex credentials from %s", credentials - ) + verbose_logger.debug("Vertex: Loading vertex credentials from %s", credentials) verbose_logger.debug( "Vertex: checking if credentials is a valid path, os.path.exists(%s)=%s, current dir %s", credentials, @@ -67,26 +63,18 @@ class VertexBase: else: json_obj = json.loads(credentials) except Exception: - raise Exception( - "Unable to load vertex credentials from environment. Got={}".format( - credentials - ) - ) + raise Exception("Unable to load vertex credentials from environment. Got={}".format(credentials)) elif isinstance(credentials, dict): json_obj = credentials else: - raise ValueError( - "Invalid credentials type: {}".format(type(credentials)) - ) + raise ValueError("Invalid credentials type: {}".format(type(credentials))) # Check if the JSON object contains Workload Identity Federation configuration if "type" in json_obj and json_obj["type"] == "external_account": # If environment_id key contains "aws" value it corresponds to an AWS config file credential_source = json_obj.get("credential_source", {}) environment_id = ( - credential_source.get("environment_id", "") - if isinstance(credential_source, dict) - else "" + credential_source.get("environment_id", "") if isinstance(credential_source, dict) else "" ) if isinstance(environment_id, str) and "aws" in environment_id: creds = self._credentials_from_identity_pool_with_aws(json_obj) @@ -123,9 +111,7 @@ class VertexBase: raise ValueError("Could not resolve project_id") if not isinstance(project_id, str): - raise TypeError( - f"Expected project_id to be a str but got {type(project_id)}" - ) + raise TypeError(f"Expected project_id to be a str but got {type(project_id)}") return creds, project_id @@ -143,16 +129,12 @@ class VertexBase: def _credentials_from_authorized_user(self, json_obj, scopes): import google.oauth2.credentials - return google.oauth2.credentials.Credentials.from_authorized_user_info( - json_obj, scopes=scopes - ) + return google.oauth2.credentials.Credentials.from_authorized_user_info(json_obj, scopes=scopes) def _credentials_from_service_account(self, json_obj, scopes): import google.oauth2.service_account - return google.oauth2.service_account.Credentials.from_service_account_info( - json_obj, scopes=scopes - ) + return google.oauth2.service_account.Credentials.from_service_account_info(json_obj, scopes=scopes) def _credentials_from_default_auth(self, scopes): import google.auth as google_auth @@ -162,9 +144,7 @@ class VertexBase: def get_default_vertex_location(self) -> str: return "us-central1" - def get_api_base( - self, api_base: Optional[str], vertex_location: Optional[str] - ) -> str: + def get_api_base(self, api_base: Optional[str], vertex_location: Optional[str]) -> str: if api_base: return api_base elif vertex_location == "global": @@ -214,9 +194,7 @@ class VertexBase: stream: Optional[bool], model: str, ) -> str: - api_base = self.get_api_base( - api_base=custom_api_base, vertex_location=vertex_location - ) + api_base = self.get_api_base(api_base=custom_api_base, vertex_location=vertex_location) default_api_base = VertexBase.create_vertex_url( vertex_location=vertex_location or "us-central1", vertex_project=vertex_project or project_id, @@ -278,6 +256,45 @@ class VertexBase: """ return False + def _is_complete_gemini_url(self, api_base: str, model: str) -> bool: + import re + + # If URL already contains /models/{model_name}, consider it complete + if re.search(r"/models/" + re.escape(model), api_base): + return True + # Or check if it contains /models/ path segment (generic detection) + if "/models/" in api_base: + return True + return False + + def _is_complete_vertex_url(self, api_base: str) -> bool: + import re + + # If contains Vertex AI full path pattern, consider it complete + complete_url_patterns = [ + r"/projects/[^/]+/locations/[^/]+/publishers/", # Partner models + r"/endpoints/\d+", # Model Garden endpoints + ] + + for pattern in complete_url_patterns: + if re.search(pattern, api_base): + return True + + return False + + def _extract_gemini_path(self, url: str) -> str: + from urllib.parse import urlparse + + parsed = urlparse(url) + # Return path without query parameters (Cloudflare may need different auth) + return parsed.path + + def _extract_vertex_path(self, url: str) -> str: + from urllib.parse import urlparse + + parsed = urlparse(url) + return parsed.path + def _check_custom_proxy( self, api_base: Optional[str], @@ -299,19 +316,32 @@ class VertexBase: if custom_llm_provider == "gemini": # For Gemini (Google AI Studio), construct the full path like other providers if model is None: - raise ValueError( - "Model parameter is required for Gemini custom API base URLs" - ) - url = "{}/models/{}:{}".format(api_base, model, endpoint) - if gemini_api_key is None: - raise ValueError( - "Missing gemini_api_key, please set `GEMINI_API_KEY`" - ) - auth_header = ( - gemini_api_key # cloudflare expects api key as bearer token - ) - else: - url = "{}:{}".format(api_base, endpoint) + raise ValueError("Model parameter is required for Gemini custom API base URLs") + + # Smart detection: is api_base a complete path or base URL? + if self._is_complete_gemini_url(api_base, model): + # Old behavior: user provided complete path, only append endpoint + url = "{}:{}".format(api_base, endpoint) + else: + # New behavior: user provided base URL, need to append full path + # Extract path from default url (/v1beta/models/{model}:{endpoint}) + path_with_endpoint = self._extract_gemini_path(url) + url = api_base.rstrip("/") + path_with_endpoint + + # Set auth_header only if gemini_api_key is provided + # Cloudflare AI Gateway can store API keys server-side + if gemini_api_key is not None: + auth_header = gemini_api_key + else: # vertex_ai + # Smart detection: is api_base a complete path or base URL? + if self._is_complete_vertex_url(api_base): + # Old behavior: user provided complete path, only append endpoint + url = "{}:{}".format(api_base, endpoint) + else: + # New behavior: user provided base URL, need to append full path + # Extract path from create_vertex_url() generated url + path_with_endpoint = self._extract_vertex_path(url) + url = api_base.rstrip("/") + path_with_endpoint if stream is True: url = url + "?alt=sse" @@ -354,9 +384,7 @@ class VertexBase: ) ### SET RUNTIME ENDPOINT ### - version: Literal["v1beta1", "v1"] = ( - "v1beta1" if should_use_v1beta1_features is True else "v1" - ) + version: Literal["v1beta1", "v1"] = "v1beta1" if should_use_v1beta1_features is True else "v1" url, endpoint = _get_vertex_url( mode=mode, model=model, @@ -403,8 +431,7 @@ class VertexBase: The original error if reauthentication fails """ verbose_logger.debug( - f"Handling reauthentication for project_id: {project_id}. " - f"Clearing cache and retrying once." + f"Handling reauthentication for project_id: {project_id}. Clearing cache and retrying once." ) # Clear the cached credentials @@ -451,20 +478,14 @@ class VertexBase: """ # Convert dict credentials to string for caching - cache_credentials = ( - json.dumps(credentials) if isinstance(credentials, dict) else credentials - ) + cache_credentials = json.dumps(credentials) if isinstance(credentials, dict) else credentials credential_cache_key = (cache_credentials, project_id) _credentials: Optional[GoogleCredentialsObject] = None - verbose_logger.debug( - f"Checking cached credentials for project_id: {project_id}" - ) + verbose_logger.debug(f"Checking cached credentials for project_id: {project_id}") if credential_cache_key in self._credentials_project_mapping: - verbose_logger.debug( - f"Cached credentials found for project_id: {project_id}." - ) + verbose_logger.debug(f"Cached credentials found for project_id: {project_id}.") # Retrieve both credentials and cached project_id cached_entry = self._credentials_project_mapping[credential_cache_key] verbose_logger.debug("cached_entry: %s", cached_entry) @@ -473,9 +494,7 @@ class VertexBase: else: # Backward compatibility with old cache format _credentials = cached_entry - credential_project_id = _credentials.quota_project_id or getattr( - _credentials, "project_id", None - ) + credential_project_id = _credentials.quota_project_id or getattr(_credentials, "project_id", None) verbose_logger.debug( "Using cached credentials for project_id: %s", credential_project_id, @@ -487,9 +506,7 @@ class VertexBase: ) try: - _credentials, credential_project_id = self.load_auth( - credentials=credentials, project_id=project_id - ) + _credentials, credential_project_id = self.load_auth(credentials=credentials, project_id=project_id) except Exception as e: verbose_logger.exception( f"Failed to load vertex credentials. Check to see if credentials containing partial/invalid information. Error: {str(e)}" @@ -510,11 +527,7 @@ class VertexBase: ## VALIDATE CREDENTIALS verbose_logger.debug(f"Validating credentials for project_id: {project_id}") - if ( - project_id is None - and credential_project_id is not None - and isinstance(credential_project_id, str) - ): + if project_id is None and credential_project_id is not None and isinstance(credential_project_id, str): project_id = credential_project_id # Update cache with resolved project_id for future lookups resolved_cache_key = (cache_credentials, project_id) @@ -530,9 +543,7 @@ class VertexBase: if _credentials.expired: try: - verbose_logger.debug( - f"Credentials expired, refreshing for project_id: {project_id}" - ) + verbose_logger.debug(f"Credentials expired, refreshing for project_id: {project_id}") self.refresh_auth(_credentials) self._credentials_project_mapping[credential_cache_key] = ( _credentials, @@ -553,9 +564,7 @@ class VertexBase: ## VALIDATION STEP if _credentials.token is None or not isinstance(_credentials.token, str): raise ValueError( - "Could not resolve credentials token. Got None or non-string token - {}".format( - _credentials.token - ) + "Could not resolve credentials token. Got None or non-string token - {}".format(_credentials.token) ) if project_id is None: @@ -585,9 +594,7 @@ class VertexBase: except Exception as e: raise e - def set_headers( - self, auth_header: Optional[str], extra_headers: Optional[dict] - ) -> dict: + def set_headers(self, auth_header: Optional[str], extra_headers: Optional[dict]) -> dict: headers = { "Content-Type": "application/json", } diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py b/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py index c85f6070084..e65db6614d6 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py @@ -708,9 +708,9 @@ class TestVertexBase: @pytest.mark.parametrize( "api_base, custom_llm_provider, gemini_api_key, endpoint, stream, auth_header, url, model, expected_auth_header, expected_url", [ - # Test case 1: Gemini with custom API base + # Test case 1: Gemini with custom API base (new behavior - appends full path) ( - "https://proxy.example.com/generativelanguage.googleapis.com/v1beta", + "https://proxy.example.com", "gemini", "test-api-key", "generateContent", @@ -719,11 +719,11 @@ class TestVertexBase: "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", "gemini-2.5-flash-lite", "test-api-key", - "https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent" + "https://proxy.example.com/v1beta/models/gemini-2.5-flash-lite:generateContent" ), - # Test case 2: Gemini with custom API base and streaming + # Test case 2: Gemini with custom API base and streaming (new behavior) ( - "https://proxy.example.com/generativelanguage.googleapis.com/v1beta", + "https://proxy.example.com", "gemini", "test-api-key", "generateContent", @@ -732,9 +732,22 @@ class TestVertexBase: "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", "gemini-2.5-flash-lite", "test-api-key", - "https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent?alt=sse" + "https://proxy.example.com/v1beta/models/gemini-2.5-flash-lite:generateContent?alt=sse" ), - # Test case 3: Non-Gemini provider with custom API base + # Test case 3: Gemini with complete URL (old behavior - backward compatibility) + ( + "https://proxy.example.com/v1beta/models/gemini-2.5-flash-lite", + "gemini", + "test-api-key", + "generateContent", + False, + None, + "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", + "gemini-2.5-flash-lite", + "test-api-key", + "https://proxy.example.com/v1beta/models/gemini-2.5-flash-lite:generateContent" + ), + # Test case 4: Vertex AI with base URL (new behavior - appends full path) ( "https://custom-vertex-api.com", "vertex_ai", @@ -745,9 +758,22 @@ class TestVertexBase: "https://aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:generateContent", "gemini-pro", "Bearer token123", - "https://custom-vertex-api.com:generateContent" + "https://custom-vertex-api.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:generateContent" ), - # Test case 4: No API base provided (should return original values) + # Test case 5: Vertex AI with complete URL (old behavior - backward compatibility) + ( + "https://custom-vertex-api.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro", + "vertex_ai", + None, + "generateContent", + False, + "Bearer token123", + "https://aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:generateContent", + "gemini-pro", + "Bearer token123", + "https://custom-vertex-api.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:generateContent" + ), + # Test case 6: No API base provided (should return original values) ( None, "gemini", @@ -760,80 +786,52 @@ class TestVertexBase: "Bearer token123", "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent" ), - # Test case 5: Gemini without API key (should raise ValueError) - ( - "https://proxy.example.com/generativelanguage.googleapis.com/v1beta", - "gemini", - None, - "generateContent", - False, - None, - "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", - "gemini-2.5-flash-lite", - None, # This should raise an exception - None - ), ], ) def test_check_custom_proxy( - self, - api_base, - custom_llm_provider, - gemini_api_key, - endpoint, - stream, - auth_header, - url, - model, - expected_auth_header, + self, + api_base, + custom_llm_provider, + gemini_api_key, + endpoint, + stream, + auth_header, + url, + model, + expected_auth_header, expected_url ): """Test the _check_custom_proxy method for handling custom API base URLs""" vertex_base = VertexBase() - - if custom_llm_provider == "gemini" and api_base and gemini_api_key is None: - # Test case 5: Should raise ValueError for Gemini without API key - with pytest.raises(ValueError, match="Missing gemini_api_key"): - vertex_base._check_custom_proxy( - api_base=api_base, - custom_llm_provider=custom_llm_provider, - gemini_api_key=gemini_api_key, - endpoint=endpoint, - stream=stream, - auth_header=auth_header, - url=url, - model=model, - ) - else: - # Test cases 1-4: Should work correctly - result_auth_header, result_url = vertex_base._check_custom_proxy( - api_base=api_base, - custom_llm_provider=custom_llm_provider, - gemini_api_key=gemini_api_key, - endpoint=endpoint, - stream=stream, - auth_header=auth_header, - url=url, - model=model, - ) - - assert result_auth_header == expected_auth_header, f"Expected auth_header {expected_auth_header}, got {result_auth_header}" - assert result_url == expected_url, f"Expected URL {expected_url}, got {result_url}" + + result_auth_header, result_url = vertex_base._check_custom_proxy( + api_base=api_base, + custom_llm_provider=custom_llm_provider, + gemini_api_key=gemini_api_key, + endpoint=endpoint, + stream=stream, + auth_header=auth_header, + url=url, + model=model, + ) + + assert result_auth_header == expected_auth_header, f"Expected auth_header {expected_auth_header}, got {result_auth_header}" + assert result_url == expected_url, f"Expected URL {expected_url}, got {result_url}" def test_check_custom_proxy_gemini_url_construction(self): """Test that Gemini URLs are constructed correctly with custom API base""" vertex_base = VertexBase() - - # Test various Gemini models with custom API base + + # Test various Gemini models with custom API base (new behavior) test_cases = [ - ("gemini-2.5-flash-lite", "generateContent", "https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent"), - ("gemini-2.5-pro", "generateContent", "https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:generateContent"), - ("gemini-1.5-flash", "streamGenerateContent", "https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-1.5-flash:streamGenerateContent"), + ("gemini-2.5-flash-lite", "generateContent", "https://proxy.example.com/v1beta/models/gemini-2.5-flash-lite:generateContent"), + ("gemini-2.5-pro", "generateContent", "https://proxy.example.com/v1beta/models/gemini-2.5-pro:generateContent"), + ("gemini-1.5-flash", "streamGenerateContent", "https://proxy.example.com/v1beta/models/gemini-1.5-flash:streamGenerateContent"), ] - + for model, endpoint, expected_url in test_cases: _, result_url = vertex_base._check_custom_proxy( - api_base="https://proxy.example.com/generativelanguage.googleapis.com/v1beta", + api_base="https://proxy.example.com", custom_llm_provider="gemini", gemini_api_key="test-api-key", endpoint=endpoint, @@ -842,16 +840,16 @@ class TestVertexBase: url=f"https://generativelanguage.googleapis.com/v1beta/models/{model}:{endpoint}", model=model, ) - + assert result_url == expected_url, f"Expected {expected_url}, got {result_url} for model {model}" def test_check_custom_proxy_streaming_parameter(self): """Test that streaming parameter correctly adds ?alt=sse to URLs""" vertex_base = VertexBase() - + # Test with streaming enabled _, result_url_streaming = vertex_base._check_custom_proxy( - api_base="https://proxy.example.com/generativelanguage.googleapis.com/v1beta", + api_base="https://proxy.example.com", custom_llm_provider="gemini", gemini_api_key="test-api-key", endpoint="generateContent", @@ -860,13 +858,13 @@ class TestVertexBase: url="https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", model="gemini-2.5-flash-lite", ) - - expected_streaming_url = "https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent?alt=sse" + + expected_streaming_url = "https://proxy.example.com/v1beta/models/gemini-2.5-flash-lite:generateContent?alt=sse" assert result_url_streaming == expected_streaming_url, f"Expected {expected_streaming_url}, got {result_url_streaming}" - + # Test with streaming disabled _, result_url_no_streaming = vertex_base._check_custom_proxy( - api_base="https://proxy.example.com/generativelanguage.googleapis.com/v1beta", + api_base="https://proxy.example.com", custom_llm_provider="gemini", gemini_api_key="test-api-key", endpoint="generateContent", @@ -875,6 +873,38 @@ class TestVertexBase: url="https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", model="gemini-2.5-flash-lite", ) - - expected_no_streaming_url = "https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent" + + expected_no_streaming_url = "https://proxy.example.com/v1beta/models/gemini-2.5-flash-lite:generateContent" assert result_url_no_streaming == expected_no_streaming_url, f"Expected {expected_no_streaming_url}, got {result_url_no_streaming}" + + def test_check_custom_proxy_cloudflare_ai_gateway(self): + """Test Cloudflare AI Gateway URL construction for both Gemini and Vertex AI""" + vertex_base = VertexBase() + + # Test Gemini with Cloudflare AI Gateway + _, gemini_url = vertex_base._check_custom_proxy( + api_base="https://gateway.ai.cloudflare.com/v1/account123/gateway456/google-ai-studio", + custom_llm_provider="gemini", + gemini_api_key="test-api-key", + endpoint="generateContent", + stream=False, + auth_header=None, + url="https://generativelanguage.googleapis.com/v1beta/models/gemini-pro:generateContent", + model="gemini-pro", + ) + expected_gemini_url = "https://gateway.ai.cloudflare.com/v1/account123/gateway456/google-ai-studio/v1beta/models/gemini-pro:generateContent" + assert gemini_url == expected_gemini_url, f"Expected {expected_gemini_url}, got {gemini_url}" + + # Test Vertex AI with Cloudflare AI Gateway + _, vertex_url = vertex_base._check_custom_proxy( + api_base="https://gateway.ai.cloudflare.com/v1/account123/gateway456/google-vertex-ai", + custom_llm_provider="vertex_ai", + gemini_api_key=None, + endpoint="streamRawPredict", + stream=False, + auth_header="Bearer token123", + url="https://us-central1-aiplatform.googleapis.com/v1/projects/my-project/locations/us-central1/publishers/anthropic/models/claude-3-sonnet@20240229:streamRawPredict", + model="claude-3-sonnet@20240229", + ) + expected_vertex_url = "https://gateway.ai.cloudflare.com/v1/account123/gateway456/google-vertex-ai/v1/projects/my-project/locations/us-central1/publishers/anthropic/models/claude-3-sonnet@20240229:streamRawPredict" + assert vertex_url == expected_vertex_url, f"Expected {expected_vertex_url}, got {vertex_url}" From a9dcb51d03dfab3ad410e8585908c19a33156010 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Mon, 29 Sep 2025 12:28:43 -0700 Subject: [PATCH 07/28] fix: add lint --- litellm/proxy/_experimental/mcp_server/server.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 6802a14fc47..2487260b615 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -425,7 +425,7 @@ if MCP_AVAILABLE: continue # Get server-specific auth header if available - server_auth_header = None + server_auth_header: Optional[Union[Dict[str, str], str]] = None if mcp_server_auth_headers and server.alias is not None: server_auth_header = mcp_server_auth_headers.get(server.alias) elif mcp_server_auth_headers and server.server_name is not None: @@ -571,16 +571,16 @@ if MCP_AVAILABLE: "litellm_logging_obj", None ) if litellm_logging_obj: - litellm_logging_obj.model_call_details[ - "mcp_tool_call_metadata" - ] = standard_logging_mcp_tool_call + litellm_logging_obj.model_call_details["mcp_tool_call_metadata"] = ( + standard_logging_mcp_tool_call + ) litellm_logging_obj.model = f"MCP: {name}" # Try managed server tool first (pass the full prefixed name) # Primary and recommended way to use MCP servers ######################################################### - mcp_server: Optional[ - MCPServer - ] = global_mcp_server_manager._get_mcp_server_from_tool_name(name) + mcp_server: Optional[MCPServer] = ( + global_mcp_server_manager._get_mcp_server_from_tool_name(name) + ) if mcp_server: standard_logging_mcp_tool_call["mcp_server_cost_info"] = ( mcp_server.mcp_info or {} From 4dc23e807995654d85b6a5d924edb4c6afdd72fa Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 29 Sep 2025 13:03:46 -0700 Subject: [PATCH 08/28] =?UTF-8?q?Revert=20"feat:=20improve=20vertex=20AI/g?= =?UTF-8?q?emini=20api=5Fbase=20handling=20for=20proxy=20services=20(?= =?UTF-8?q?=E2=80=A6"=20(#15042)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit This reverts commit ff2d19e4cab88c7ce0a417f47e5ff7be85da515a. --- litellm/llms/vertex_ai/vertex_llm_base.py | 165 ++++++++-------- .../llms/vertex_ai/test_vertex_llm_base.py | 184 ++++++++---------- 2 files changed, 156 insertions(+), 193 deletions(-) diff --git a/litellm/llms/vertex_ai/vertex_llm_base.py b/litellm/llms/vertex_ai/vertex_llm_base.py index 6769bc3fb28..6d194d41add 100644 --- a/litellm/llms/vertex_ai/vertex_llm_base.py +++ b/litellm/llms/vertex_ai/vertex_llm_base.py @@ -46,10 +46,14 @@ class VertexBase: return "global" return vertex_region or "us-central1" - def load_auth(self, credentials: Optional[VERTEX_CREDENTIALS_TYPES], project_id: Optional[str]) -> Tuple[Any, str]: + def load_auth( + self, credentials: Optional[VERTEX_CREDENTIALS_TYPES], project_id: Optional[str] + ) -> Tuple[Any, str]: if credentials is not None: if isinstance(credentials, str): - verbose_logger.debug("Vertex: Loading vertex credentials from %s", credentials) + verbose_logger.debug( + "Vertex: Loading vertex credentials from %s", credentials + ) verbose_logger.debug( "Vertex: checking if credentials is a valid path, os.path.exists(%s)=%s, current dir %s", credentials, @@ -63,18 +67,26 @@ class VertexBase: else: json_obj = json.loads(credentials) except Exception: - raise Exception("Unable to load vertex credentials from environment. Got={}".format(credentials)) + raise Exception( + "Unable to load vertex credentials from environment. Got={}".format( + credentials + ) + ) elif isinstance(credentials, dict): json_obj = credentials else: - raise ValueError("Invalid credentials type: {}".format(type(credentials))) + raise ValueError( + "Invalid credentials type: {}".format(type(credentials)) + ) # Check if the JSON object contains Workload Identity Federation configuration if "type" in json_obj and json_obj["type"] == "external_account": # If environment_id key contains "aws" value it corresponds to an AWS config file credential_source = json_obj.get("credential_source", {}) environment_id = ( - credential_source.get("environment_id", "") if isinstance(credential_source, dict) else "" + credential_source.get("environment_id", "") + if isinstance(credential_source, dict) + else "" ) if isinstance(environment_id, str) and "aws" in environment_id: creds = self._credentials_from_identity_pool_with_aws(json_obj) @@ -111,7 +123,9 @@ class VertexBase: raise ValueError("Could not resolve project_id") if not isinstance(project_id, str): - raise TypeError(f"Expected project_id to be a str but got {type(project_id)}") + raise TypeError( + f"Expected project_id to be a str but got {type(project_id)}" + ) return creds, project_id @@ -129,12 +143,16 @@ class VertexBase: def _credentials_from_authorized_user(self, json_obj, scopes): import google.oauth2.credentials - return google.oauth2.credentials.Credentials.from_authorized_user_info(json_obj, scopes=scopes) + return google.oauth2.credentials.Credentials.from_authorized_user_info( + json_obj, scopes=scopes + ) def _credentials_from_service_account(self, json_obj, scopes): import google.oauth2.service_account - return google.oauth2.service_account.Credentials.from_service_account_info(json_obj, scopes=scopes) + return google.oauth2.service_account.Credentials.from_service_account_info( + json_obj, scopes=scopes + ) def _credentials_from_default_auth(self, scopes): import google.auth as google_auth @@ -144,7 +162,9 @@ class VertexBase: def get_default_vertex_location(self) -> str: return "us-central1" - def get_api_base(self, api_base: Optional[str], vertex_location: Optional[str]) -> str: + def get_api_base( + self, api_base: Optional[str], vertex_location: Optional[str] + ) -> str: if api_base: return api_base elif vertex_location == "global": @@ -194,7 +214,9 @@ class VertexBase: stream: Optional[bool], model: str, ) -> str: - api_base = self.get_api_base(api_base=custom_api_base, vertex_location=vertex_location) + api_base = self.get_api_base( + api_base=custom_api_base, vertex_location=vertex_location + ) default_api_base = VertexBase.create_vertex_url( vertex_location=vertex_location or "us-central1", vertex_project=vertex_project or project_id, @@ -256,45 +278,6 @@ class VertexBase: """ return False - def _is_complete_gemini_url(self, api_base: str, model: str) -> bool: - import re - - # If URL already contains /models/{model_name}, consider it complete - if re.search(r"/models/" + re.escape(model), api_base): - return True - # Or check if it contains /models/ path segment (generic detection) - if "/models/" in api_base: - return True - return False - - def _is_complete_vertex_url(self, api_base: str) -> bool: - import re - - # If contains Vertex AI full path pattern, consider it complete - complete_url_patterns = [ - r"/projects/[^/]+/locations/[^/]+/publishers/", # Partner models - r"/endpoints/\d+", # Model Garden endpoints - ] - - for pattern in complete_url_patterns: - if re.search(pattern, api_base): - return True - - return False - - def _extract_gemini_path(self, url: str) -> str: - from urllib.parse import urlparse - - parsed = urlparse(url) - # Return path without query parameters (Cloudflare may need different auth) - return parsed.path - - def _extract_vertex_path(self, url: str) -> str: - from urllib.parse import urlparse - - parsed = urlparse(url) - return parsed.path - def _check_custom_proxy( self, api_base: Optional[str], @@ -316,32 +299,19 @@ class VertexBase: if custom_llm_provider == "gemini": # For Gemini (Google AI Studio), construct the full path like other providers if model is None: - raise ValueError("Model parameter is required for Gemini custom API base URLs") - - # Smart detection: is api_base a complete path or base URL? - if self._is_complete_gemini_url(api_base, model): - # Old behavior: user provided complete path, only append endpoint - url = "{}:{}".format(api_base, endpoint) - else: - # New behavior: user provided base URL, need to append full path - # Extract path from default url (/v1beta/models/{model}:{endpoint}) - path_with_endpoint = self._extract_gemini_path(url) - url = api_base.rstrip("/") + path_with_endpoint - - # Set auth_header only if gemini_api_key is provided - # Cloudflare AI Gateway can store API keys server-side - if gemini_api_key is not None: - auth_header = gemini_api_key - else: # vertex_ai - # Smart detection: is api_base a complete path or base URL? - if self._is_complete_vertex_url(api_base): - # Old behavior: user provided complete path, only append endpoint - url = "{}:{}".format(api_base, endpoint) - else: - # New behavior: user provided base URL, need to append full path - # Extract path from create_vertex_url() generated url - path_with_endpoint = self._extract_vertex_path(url) - url = api_base.rstrip("/") + path_with_endpoint + raise ValueError( + "Model parameter is required for Gemini custom API base URLs" + ) + url = "{}/models/{}:{}".format(api_base, model, endpoint) + if gemini_api_key is None: + raise ValueError( + "Missing gemini_api_key, please set `GEMINI_API_KEY`" + ) + auth_header = ( + gemini_api_key # cloudflare expects api key as bearer token + ) + else: + url = "{}:{}".format(api_base, endpoint) if stream is True: url = url + "?alt=sse" @@ -384,7 +354,9 @@ class VertexBase: ) ### SET RUNTIME ENDPOINT ### - version: Literal["v1beta1", "v1"] = "v1beta1" if should_use_v1beta1_features is True else "v1" + version: Literal["v1beta1", "v1"] = ( + "v1beta1" if should_use_v1beta1_features is True else "v1" + ) url, endpoint = _get_vertex_url( mode=mode, model=model, @@ -431,7 +403,8 @@ class VertexBase: The original error if reauthentication fails """ verbose_logger.debug( - f"Handling reauthentication for project_id: {project_id}. Clearing cache and retrying once." + f"Handling reauthentication for project_id: {project_id}. " + f"Clearing cache and retrying once." ) # Clear the cached credentials @@ -478,14 +451,20 @@ class VertexBase: """ # Convert dict credentials to string for caching - cache_credentials = json.dumps(credentials) if isinstance(credentials, dict) else credentials + cache_credentials = ( + json.dumps(credentials) if isinstance(credentials, dict) else credentials + ) credential_cache_key = (cache_credentials, project_id) _credentials: Optional[GoogleCredentialsObject] = None - verbose_logger.debug(f"Checking cached credentials for project_id: {project_id}") + verbose_logger.debug( + f"Checking cached credentials for project_id: {project_id}" + ) if credential_cache_key in self._credentials_project_mapping: - verbose_logger.debug(f"Cached credentials found for project_id: {project_id}.") + verbose_logger.debug( + f"Cached credentials found for project_id: {project_id}." + ) # Retrieve both credentials and cached project_id cached_entry = self._credentials_project_mapping[credential_cache_key] verbose_logger.debug("cached_entry: %s", cached_entry) @@ -494,7 +473,9 @@ class VertexBase: else: # Backward compatibility with old cache format _credentials = cached_entry - credential_project_id = _credentials.quota_project_id or getattr(_credentials, "project_id", None) + credential_project_id = _credentials.quota_project_id or getattr( + _credentials, "project_id", None + ) verbose_logger.debug( "Using cached credentials for project_id: %s", credential_project_id, @@ -506,7 +487,9 @@ class VertexBase: ) try: - _credentials, credential_project_id = self.load_auth(credentials=credentials, project_id=project_id) + _credentials, credential_project_id = self.load_auth( + credentials=credentials, project_id=project_id + ) except Exception as e: verbose_logger.exception( f"Failed to load vertex credentials. Check to see if credentials containing partial/invalid information. Error: {str(e)}" @@ -527,7 +510,11 @@ class VertexBase: ## VALIDATE CREDENTIALS verbose_logger.debug(f"Validating credentials for project_id: {project_id}") - if project_id is None and credential_project_id is not None and isinstance(credential_project_id, str): + if ( + project_id is None + and credential_project_id is not None + and isinstance(credential_project_id, str) + ): project_id = credential_project_id # Update cache with resolved project_id for future lookups resolved_cache_key = (cache_credentials, project_id) @@ -543,7 +530,9 @@ class VertexBase: if _credentials.expired: try: - verbose_logger.debug(f"Credentials expired, refreshing for project_id: {project_id}") + verbose_logger.debug( + f"Credentials expired, refreshing for project_id: {project_id}" + ) self.refresh_auth(_credentials) self._credentials_project_mapping[credential_cache_key] = ( _credentials, @@ -564,7 +553,9 @@ class VertexBase: ## VALIDATION STEP if _credentials.token is None or not isinstance(_credentials.token, str): raise ValueError( - "Could not resolve credentials token. Got None or non-string token - {}".format(_credentials.token) + "Could not resolve credentials token. Got None or non-string token - {}".format( + _credentials.token + ) ) if project_id is None: @@ -594,7 +585,9 @@ class VertexBase: except Exception as e: raise e - def set_headers(self, auth_header: Optional[str], extra_headers: Optional[dict]) -> dict: + def set_headers( + self, auth_header: Optional[str], extra_headers: Optional[dict] + ) -> dict: headers = { "Content-Type": "application/json", } diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py b/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py index e65db6614d6..c85f6070084 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_llm_base.py @@ -708,9 +708,9 @@ class TestVertexBase: @pytest.mark.parametrize( "api_base, custom_llm_provider, gemini_api_key, endpoint, stream, auth_header, url, model, expected_auth_header, expected_url", [ - # Test case 1: Gemini with custom API base (new behavior - appends full path) + # Test case 1: Gemini with custom API base ( - "https://proxy.example.com", + "https://proxy.example.com/generativelanguage.googleapis.com/v1beta", "gemini", "test-api-key", "generateContent", @@ -719,11 +719,11 @@ class TestVertexBase: "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", "gemini-2.5-flash-lite", "test-api-key", - "https://proxy.example.com/v1beta/models/gemini-2.5-flash-lite:generateContent" + "https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent" ), - # Test case 2: Gemini with custom API base and streaming (new behavior) + # Test case 2: Gemini with custom API base and streaming ( - "https://proxy.example.com", + "https://proxy.example.com/generativelanguage.googleapis.com/v1beta", "gemini", "test-api-key", "generateContent", @@ -732,22 +732,9 @@ class TestVertexBase: "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", "gemini-2.5-flash-lite", "test-api-key", - "https://proxy.example.com/v1beta/models/gemini-2.5-flash-lite:generateContent?alt=sse" + "https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent?alt=sse" ), - # Test case 3: Gemini with complete URL (old behavior - backward compatibility) - ( - "https://proxy.example.com/v1beta/models/gemini-2.5-flash-lite", - "gemini", - "test-api-key", - "generateContent", - False, - None, - "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", - "gemini-2.5-flash-lite", - "test-api-key", - "https://proxy.example.com/v1beta/models/gemini-2.5-flash-lite:generateContent" - ), - # Test case 4: Vertex AI with base URL (new behavior - appends full path) + # Test case 3: Non-Gemini provider with custom API base ( "https://custom-vertex-api.com", "vertex_ai", @@ -758,22 +745,9 @@ class TestVertexBase: "https://aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:generateContent", "gemini-pro", "Bearer token123", - "https://custom-vertex-api.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:generateContent" + "https://custom-vertex-api.com:generateContent" ), - # Test case 5: Vertex AI with complete URL (old behavior - backward compatibility) - ( - "https://custom-vertex-api.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro", - "vertex_ai", - None, - "generateContent", - False, - "Bearer token123", - "https://aiplatform.googleapis.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:generateContent", - "gemini-pro", - "Bearer token123", - "https://custom-vertex-api.com/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-pro:generateContent" - ), - # Test case 6: No API base provided (should return original values) + # Test case 4: No API base provided (should return original values) ( None, "gemini", @@ -786,52 +760,80 @@ class TestVertexBase: "Bearer token123", "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent" ), + # Test case 5: Gemini without API key (should raise ValueError) + ( + "https://proxy.example.com/generativelanguage.googleapis.com/v1beta", + "gemini", + None, + "generateContent", + False, + None, + "https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", + "gemini-2.5-flash-lite", + None, # This should raise an exception + None + ), ], ) def test_check_custom_proxy( - self, - api_base, - custom_llm_provider, - gemini_api_key, - endpoint, - stream, - auth_header, - url, - model, - expected_auth_header, + self, + api_base, + custom_llm_provider, + gemini_api_key, + endpoint, + stream, + auth_header, + url, + model, + expected_auth_header, expected_url ): """Test the _check_custom_proxy method for handling custom API base URLs""" vertex_base = VertexBase() - - result_auth_header, result_url = vertex_base._check_custom_proxy( - api_base=api_base, - custom_llm_provider=custom_llm_provider, - gemini_api_key=gemini_api_key, - endpoint=endpoint, - stream=stream, - auth_header=auth_header, - url=url, - model=model, - ) - - assert result_auth_header == expected_auth_header, f"Expected auth_header {expected_auth_header}, got {result_auth_header}" - assert result_url == expected_url, f"Expected URL {expected_url}, got {result_url}" + + if custom_llm_provider == "gemini" and api_base and gemini_api_key is None: + # Test case 5: Should raise ValueError for Gemini without API key + with pytest.raises(ValueError, match="Missing gemini_api_key"): + vertex_base._check_custom_proxy( + api_base=api_base, + custom_llm_provider=custom_llm_provider, + gemini_api_key=gemini_api_key, + endpoint=endpoint, + stream=stream, + auth_header=auth_header, + url=url, + model=model, + ) + else: + # Test cases 1-4: Should work correctly + result_auth_header, result_url = vertex_base._check_custom_proxy( + api_base=api_base, + custom_llm_provider=custom_llm_provider, + gemini_api_key=gemini_api_key, + endpoint=endpoint, + stream=stream, + auth_header=auth_header, + url=url, + model=model, + ) + + assert result_auth_header == expected_auth_header, f"Expected auth_header {expected_auth_header}, got {result_auth_header}" + assert result_url == expected_url, f"Expected URL {expected_url}, got {result_url}" def test_check_custom_proxy_gemini_url_construction(self): """Test that Gemini URLs are constructed correctly with custom API base""" vertex_base = VertexBase() - - # Test various Gemini models with custom API base (new behavior) + + # Test various Gemini models with custom API base test_cases = [ - ("gemini-2.5-flash-lite", "generateContent", "https://proxy.example.com/v1beta/models/gemini-2.5-flash-lite:generateContent"), - ("gemini-2.5-pro", "generateContent", "https://proxy.example.com/v1beta/models/gemini-2.5-pro:generateContent"), - ("gemini-1.5-flash", "streamGenerateContent", "https://proxy.example.com/v1beta/models/gemini-1.5-flash:streamGenerateContent"), + ("gemini-2.5-flash-lite", "generateContent", "https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent"), + ("gemini-2.5-pro", "generateContent", "https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-pro:generateContent"), + ("gemini-1.5-flash", "streamGenerateContent", "https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-1.5-flash:streamGenerateContent"), ] - + for model, endpoint, expected_url in test_cases: _, result_url = vertex_base._check_custom_proxy( - api_base="https://proxy.example.com", + api_base="https://proxy.example.com/generativelanguage.googleapis.com/v1beta", custom_llm_provider="gemini", gemini_api_key="test-api-key", endpoint=endpoint, @@ -840,16 +842,16 @@ class TestVertexBase: url=f"https://generativelanguage.googleapis.com/v1beta/models/{model}:{endpoint}", model=model, ) - + assert result_url == expected_url, f"Expected {expected_url}, got {result_url} for model {model}" def test_check_custom_proxy_streaming_parameter(self): """Test that streaming parameter correctly adds ?alt=sse to URLs""" vertex_base = VertexBase() - + # Test with streaming enabled _, result_url_streaming = vertex_base._check_custom_proxy( - api_base="https://proxy.example.com", + api_base="https://proxy.example.com/generativelanguage.googleapis.com/v1beta", custom_llm_provider="gemini", gemini_api_key="test-api-key", endpoint="generateContent", @@ -858,13 +860,13 @@ class TestVertexBase: url="https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", model="gemini-2.5-flash-lite", ) - - expected_streaming_url = "https://proxy.example.com/v1beta/models/gemini-2.5-flash-lite:generateContent?alt=sse" + + expected_streaming_url = "https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent?alt=sse" assert result_url_streaming == expected_streaming_url, f"Expected {expected_streaming_url}, got {result_url_streaming}" - + # Test with streaming disabled _, result_url_no_streaming = vertex_base._check_custom_proxy( - api_base="https://proxy.example.com", + api_base="https://proxy.example.com/generativelanguage.googleapis.com/v1beta", custom_llm_provider="gemini", gemini_api_key="test-api-key", endpoint="generateContent", @@ -873,38 +875,6 @@ class TestVertexBase: url="https://generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent", model="gemini-2.5-flash-lite", ) - - expected_no_streaming_url = "https://proxy.example.com/v1beta/models/gemini-2.5-flash-lite:generateContent" + + expected_no_streaming_url = "https://proxy.example.com/generativelanguage.googleapis.com/v1beta/models/gemini-2.5-flash-lite:generateContent" assert result_url_no_streaming == expected_no_streaming_url, f"Expected {expected_no_streaming_url}, got {result_url_no_streaming}" - - def test_check_custom_proxy_cloudflare_ai_gateway(self): - """Test Cloudflare AI Gateway URL construction for both Gemini and Vertex AI""" - vertex_base = VertexBase() - - # Test Gemini with Cloudflare AI Gateway - _, gemini_url = vertex_base._check_custom_proxy( - api_base="https://gateway.ai.cloudflare.com/v1/account123/gateway456/google-ai-studio", - custom_llm_provider="gemini", - gemini_api_key="test-api-key", - endpoint="generateContent", - stream=False, - auth_header=None, - url="https://generativelanguage.googleapis.com/v1beta/models/gemini-pro:generateContent", - model="gemini-pro", - ) - expected_gemini_url = "https://gateway.ai.cloudflare.com/v1/account123/gateway456/google-ai-studio/v1beta/models/gemini-pro:generateContent" - assert gemini_url == expected_gemini_url, f"Expected {expected_gemini_url}, got {gemini_url}" - - # Test Vertex AI with Cloudflare AI Gateway - _, vertex_url = vertex_base._check_custom_proxy( - api_base="https://gateway.ai.cloudflare.com/v1/account123/gateway456/google-vertex-ai", - custom_llm_provider="vertex_ai", - gemini_api_key=None, - endpoint="streamRawPredict", - stream=False, - auth_header="Bearer token123", - url="https://us-central1-aiplatform.googleapis.com/v1/projects/my-project/locations/us-central1/publishers/anthropic/models/claude-3-sonnet@20240229:streamRawPredict", - model="claude-3-sonnet@20240229", - ) - expected_vertex_url = "https://gateway.ai.cloudflare.com/v1/account123/gateway456/google-vertex-ai/v1/projects/my-project/locations/us-central1/publishers/anthropic/models/claude-3-sonnet@20240229:streamRawPredict" - assert vertex_url == expected_vertex_url, f"Expected {expected_vertex_url}, got {vertex_url}" From 5357b1e10239ffa2c807bb35b7d80042d8ea7cc9 Mon Sep 17 00:00:00 2001 From: Cedar Myers Date: Mon, 29 Sep 2025 16:04:33 -0400 Subject: [PATCH 09/28] fix: remove invalid vertex -latest models --- docs/my-website/docs/providers/vertex.md | 2 - ...odel_prices_and_context_window_backup.json | 90 ------------------- model_prices_and_context_window.json | 90 ------------------- 3 files changed, 182 deletions(-) diff --git a/docs/my-website/docs/providers/vertex.md b/docs/my-website/docs/providers/vertex.md index 9a969876432..943823c6386 100644 --- a/docs/my-website/docs/providers/vertex.md +++ b/docs/my-website/docs/providers/vertex.md @@ -1299,8 +1299,6 @@ litellm.vertex_location = "us-central1 # Your Location | gemini-2.5-pro | `completion('gemini-2.5-pro', messages)`, `completion('vertex_ai/gemini-2.5-pro', messages)` | | gemini-2.5-flash-preview-09-2025 | `completion('gemini-2.5-flash-preview-09-2025', messages)`, `completion('vertex_ai/gemini-2.5-flash-preview-09-2025', messages)` | | gemini-2.5-flash-lite-preview-09-2025 | `completion('gemini-2.5-flash-lite-preview-09-2025', messages)`, `completion('vertex_ai/gemini-2.5-flash-lite-preview-09-2025', messages)` | -| gemini-flash-latest | `completion('gemini-flash-latest', messages)`, `completion('vertex_ai/gemini-flash-latest', messages)` | -| gemini-flash-lite-latest | `completion('gemini-flash-lite-latest', messages)`, `completion('vertex_ai/gemini-flash-lite-latest', messages)` | ## Fine-tuned Models diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 8fc98e5d506..53bd820d3bb 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -9396,96 +9396,6 @@ "supports_vision": true, "supports_web_search": true }, - "gemini-flash-latest": { - "cache_read_input_token_cost": 7.5e-08, - "input_cost_per_audio_token": 1e-06, - "input_cost_per_token": 3e-07, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 65535, - "max_pdf_size_mb": 30, - "max_tokens": 65535, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_reasoning_token": 2.5e-06, - "output_cost_per_token": 2.5e-06, - "source": "https://developers.googleblog.com/en/continuing-to-bring-you-our-latest-models-with-an-improved-gemini-2-5-flash-and-flash-lite-release/", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/completions", - "/v1/batch" - ], - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text" - ], - "supports_audio_output": false, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_url_context": true, - "supports_vision": true, - "supports_web_search": true - }, - "gemini-flash-lite-latest": { - "cache_read_input_token_cost": 2.5e-08, - "input_cost_per_audio_token": 3e-07, - "input_cost_per_token": 1e-07, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 65535, - "max_pdf_size_mb": 30, - "max_tokens": 65535, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_reasoning_token": 4e-07, - "output_cost_per_token": 4e-07, - "source": "https://developers.googleblog.com/en/continuing-to-bring-you-our-latest-models-with-an-improved-gemini-2-5-flash-and-flash-lite-release/", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/completions", - "/v1/batch" - ], - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text" - ], - "supports_audio_output": false, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_url_context": true, - "supports_vision": true, - "supports_web_search": true - }, "gemini-2.5-flash-lite-preview-06-17": { "cache_read_input_token_cost": 2.5e-08, "input_cost_per_audio_token": 5e-07, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 8fc98e5d506..53bd820d3bb 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -9396,96 +9396,6 @@ "supports_vision": true, "supports_web_search": true }, - "gemini-flash-latest": { - "cache_read_input_token_cost": 7.5e-08, - "input_cost_per_audio_token": 1e-06, - "input_cost_per_token": 3e-07, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 65535, - "max_pdf_size_mb": 30, - "max_tokens": 65535, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_reasoning_token": 2.5e-06, - "output_cost_per_token": 2.5e-06, - "source": "https://developers.googleblog.com/en/continuing-to-bring-you-our-latest-models-with-an-improved-gemini-2-5-flash-and-flash-lite-release/", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/completions", - "/v1/batch" - ], - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text" - ], - "supports_audio_output": false, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_url_context": true, - "supports_vision": true, - "supports_web_search": true - }, - "gemini-flash-lite-latest": { - "cache_read_input_token_cost": 2.5e-08, - "input_cost_per_audio_token": 3e-07, - "input_cost_per_token": 1e-07, - "litellm_provider": "vertex_ai-language-models", - "max_audio_length_hours": 8.4, - "max_audio_per_prompt": 1, - "max_images_per_prompt": 3000, - "max_input_tokens": 1048576, - "max_output_tokens": 65535, - "max_pdf_size_mb": 30, - "max_tokens": 65535, - "max_video_length": 1, - "max_videos_per_prompt": 10, - "mode": "chat", - "output_cost_per_reasoning_token": 4e-07, - "output_cost_per_token": 4e-07, - "source": "https://developers.googleblog.com/en/continuing-to-bring-you-our-latest-models-with-an-improved-gemini-2-5-flash-and-flash-lite-release/", - "supported_endpoints": [ - "/v1/chat/completions", - "/v1/completions", - "/v1/batch" - ], - "supported_modalities": [ - "text", - "image", - "audio", - "video" - ], - "supported_output_modalities": [ - "text" - ], - "supports_audio_output": false, - "supports_function_calling": true, - "supports_parallel_function_calling": true, - "supports_pdf_input": true, - "supports_prompt_caching": true, - "supports_reasoning": true, - "supports_response_schema": true, - "supports_system_messages": true, - "supports_tool_choice": true, - "supports_url_context": true, - "supports_vision": true, - "supports_web_search": true - }, "gemini-2.5-flash-lite-preview-06-17": { "cache_read_input_token_cost": 2.5e-08, "input_cost_per_audio_token": 5e-07, From 038863a1fe8a167fc8f41b8a32b481b50005ff90 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 29 Sep 2025 13:09:00 -0700 Subject: [PATCH 10/28] [Feat] Add new claude-sonnet-4-5 model family (#15041) * add new claude-sonnet-4-5 * docs fix * fix tool_use_system_prompt_tokens * add anthropic.claude-sonnet-4-5-20250929 to bedrock converse models --- docs/my-website/docs/providers/anthropic.md | 2 + docs/my-website/docs/providers/bedrock.md | 1 + litellm/constants.py | 1 + ...odel_prices_and_context_window_backup.json | 96 +++++++++++++++++++ model_prices_and_context_window.json | 96 +++++++++++++++++++ 5 files changed, 196 insertions(+) diff --git a/docs/my-website/docs/providers/anthropic.md b/docs/my-website/docs/providers/anthropic.md index 820c2906bf0..1663d32ddfc 100644 --- a/docs/my-website/docs/providers/anthropic.md +++ b/docs/my-website/docs/providers/anthropic.md @@ -4,6 +4,7 @@ import TabItem from '@theme/TabItem'; # Anthropic LiteLLM supports all anthropic models. +- `claude-sonnet-4-5-20250929` - `claude-opus-4-1-20250805` - `claude-4` (`claude-opus-4-20250514`, `claude-sonnet-4-20250514`) - `claude-3.7` (`claude-3-7-sonnet-20250219`) @@ -268,6 +269,7 @@ print(response) | Model Name | Function Call | |------------------|--------------------------------------------| +| claude-sonnet-4-5 | `completion('claude-sonnet-4-5-20250929', messages)` | `os.environ['ANTHROPIC_API_KEY']` | | claude-opus-4 | `completion('claude-opus-4-20250514', messages)` | `os.environ['ANTHROPIC_API_KEY']` | | claude-sonnet-4 | `completion('claude-sonnet-4-20250514', messages)` | `os.environ['ANTHROPIC_API_KEY']` | | claude-3.7 | `completion('claude-3-7-sonnet-20250219', messages)` | `os.environ['ANTHROPIC_API_KEY']` | diff --git a/docs/my-website/docs/providers/bedrock.md b/docs/my-website/docs/providers/bedrock.md index fe996099145..50d32a45df3 100644 --- a/docs/my-website/docs/providers/bedrock.md +++ b/docs/my-website/docs/providers/bedrock.md @@ -1857,6 +1857,7 @@ Here's an example of using a bedrock model with LiteLLM. For a complete list, re | GPT-OSS 20B | `completion(model='bedrock/converse/openai.gpt-oss-20b-1:0', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | | GPT-OSS 120B | `completion(model='bedrock/converse/openai.gpt-oss-120b-1:0', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']`, `os.environ['AWS_REGION_NAME']` | | Deepseek R1 | `completion(model='bedrock/us.deepseek.r1-v1:0', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']` | +| Anthropic Claude Sonnet 4.5 | `completion(model='bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']` | | Anthropic Claude-V3.5 Sonnet | `completion(model='bedrock/anthropic.claude-3-5-sonnet-20240620-v1:0', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']` | | Anthropic Claude-V3 sonnet | `completion(model='bedrock/anthropic.claude-3-sonnet-20240229-v1:0', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']` | | Anthropic Claude-V3 Haiku | `completion(model='bedrock/anthropic.claude-3-haiku-20240307-v1:0', messages=messages)` | `os.environ['AWS_ACCESS_KEY_ID']`, `os.environ['AWS_SECRET_ACCESS_KEY']` | diff --git a/litellm/constants.py b/litellm/constants.py index 1ed9f237a29..b839256b78e 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -819,6 +819,7 @@ BEDROCK_CONVERSE_MODELS = [ "deepseek.v3-v1:0", "openai.gpt-oss-20b-1:0", "openai.gpt-oss-120b-1:0", + "anthropic.claude-sonnet-4-5-20250929-v1:0", "anthropic.claude-opus-4-1-20250805-v1:0", "anthropic.claude-opus-4-20250514-v1:0", "anthropic.claude-sonnet-4-20250514-v1:0", diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 8fc98e5d506..e69d850447a 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -4739,6 +4739,32 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 159 }, + "claude-sonnet-4-5-20250929": { + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "litellm_provider": "anthropic", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 200000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "tool_use_system_prompt_tokens": 346 + }, "claude-opus-4-1": { "cache_creation_input_token_cost": 1.875e-05, "cache_creation_input_token_cost_above_1hr": 3e-05, @@ -19643,6 +19669,32 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 159 }, + "us.anthropic.claude-sonnet-4-5-20250929-v1:0": { + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 200000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "tool_use_system_prompt_tokens": 346 + }, "us.anthropic.claude-opus-4-20250514-v1:0": { "cache_creation_input_token_cost": 1.875e-05, "cache_read_input_token_cost": 1.5e-06, @@ -20983,6 +21035,50 @@ "supports_tool_choice": true, "supports_vision": true }, + "vertex_ai/claude-sonnet-4-5": { + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "input_cost_per_token_batches": 1.5e-06, + "litellm_provider": "vertex_ai-anthropic_models", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 200000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "output_cost_per_token_batches": 7.5e-06, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "vertex_ai/claude-sonnet-4-5@20250929": { + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "input_cost_per_token_batches": 1.5e-06, + "litellm_provider": "vertex_ai-anthropic_models", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 200000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "output_cost_per_token_batches": 7.5e-06, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "vertex_ai/claude-opus-4@20250514": { "cache_creation_input_token_cost": 1.875e-05, "cache_read_input_token_cost": 1.5e-06, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 8fc98e5d506..e69d850447a 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -4739,6 +4739,32 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 159 }, + "claude-sonnet-4-5-20250929": { + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "litellm_provider": "anthropic", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 200000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "tool_use_system_prompt_tokens": 346 + }, "claude-opus-4-1": { "cache_creation_input_token_cost": 1.875e-05, "cache_creation_input_token_cost_above_1hr": 3e-05, @@ -19643,6 +19669,32 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 159 }, + "us.anthropic.claude-sonnet-4-5-20250929-v1:0": { + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 200000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "tool_use_system_prompt_tokens": 346 + }, "us.anthropic.claude-opus-4-20250514-v1:0": { "cache_creation_input_token_cost": 1.875e-05, "cache_read_input_token_cost": 1.5e-06, @@ -20983,6 +21035,50 @@ "supports_tool_choice": true, "supports_vision": true }, + "vertex_ai/claude-sonnet-4-5": { + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "input_cost_per_token_batches": 1.5e-06, + "litellm_provider": "vertex_ai-anthropic_models", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 200000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "output_cost_per_token_batches": 7.5e-06, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "vertex_ai/claude-sonnet-4-5@20250929": { + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "input_cost_per_token_batches": 1.5e-06, + "litellm_provider": "vertex_ai-anthropic_models", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 200000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "output_cost_per_token_batches": 7.5e-06, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "vertex_ai/claude-opus-4@20250514": { "cache_creation_input_token_cost": 1.875e-05, "cache_read_input_token_cost": 1.5e-06, From 3ab1c31e4ea04351379eecf65de0df3793c3cf90 Mon Sep 17 00:00:00 2001 From: Sameer Kankute Date: Tue, 30 Sep 2025 01:49:04 +0530 Subject: [PATCH 11/28] (Feat) Add cost tracking for Vertex AI Passthrough `/predict` endpoint (#15019) * Add cost tracking for passthrough for predict endpoint * restore file --- .../vertex_passthrough_logging_handler.py | 33 +++++--- .../test_llm_pass_through_endpoints.py | 79 +++++++++++++++++++ 2 files changed, 103 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py index 8cbda18ee3b..5b22b2746c9 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/vertex_passthrough_logging_handler.py @@ -110,7 +110,7 @@ class VertexPassthroughLoggingHandler: PassthroughCallTypes.passthrough_image_generation.value ) elif VertexPassthroughLoggingHandler._is_multimodal_embedding_response( - json_response=_json_response, + json_response=_json_response, ): # Use multimodal embedding transformation vertex_multimodal_config = VertexAIMultimodalEmbeddingConfig() @@ -137,6 +137,15 @@ class VertexPassthroughLoggingHandler: logging_obj.model = model logging_obj.model_call_details["model"] = logging_obj.model + response_cost = litellm.completion_cost( + completion_response=litellm_prediction_response, + model=model, + custom_llm_provider="vertex_ai", + ) + + kwargs["response_cost"] = response_cost + kwargs["model"] = model + logging_obj.model_call_details["response_cost"] = response_cost return { "result": litellm_prediction_response, @@ -221,7 +230,9 @@ class VertexPassthroughLoggingHandler: - Logs in litellm callbacks """ kwargs: Dict[str, Any] = {} - model = model or VertexPassthroughLoggingHandler.extract_model_from_url(url_route) + model = model or VertexPassthroughLoggingHandler.extract_model_from_url( + url_route + ) complete_streaming_response = ( VertexPassthroughLoggingHandler._build_complete_streaming_response( all_chunks=all_chunks, @@ -340,13 +351,13 @@ class VertexPassthroughLoggingHandler: """ Detect if the response is from a multimodal embedding request. - Check if the response contains multimodal embedding fields: - - Docs: https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/multimodal-embeddings-api#response-body - - + Check if the response contains multimodal embedding fields: + - Docs: https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/multimodal-embeddings-api#response-body + + Args: json_response: The JSON response from Vertex AI - + Returns: bool: True if this is a multimodal embedding response """ @@ -358,10 +369,14 @@ class VertexPassthroughLoggingHandler: # Check for multimodal embedding response fields if any( key in prediction - for key in ["textEmbedding", "imageEmbedding", "videoEmbeddings"] + for key in [ + "textEmbedding", + "imageEmbedding", + "videoEmbeddings", + ] ): return True - + return False @staticmethod diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 702ae4bd42f..239f83b21ad 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -719,6 +719,85 @@ class TestVertexAIPassThroughHandler: empty_response = {} assert VertexPassthroughLoggingHandler._is_multimodal_embedding_response(empty_response) is False + def test_vertex_passthrough_handler_predict_cost_tracking(self): + """ + Test that vertex_passthrough_handler correctly tracks costs for /predict endpoint + """ + import datetime + from unittest.mock import Mock, patch + + from litellm.litellm_core_utils.litellm_logging import ( + Logging as LiteLLMLoggingObj, + ) + from litellm.proxy.pass_through_endpoints.llm_provider_handlers.vertex_passthrough_logging_handler import ( + VertexPassthroughLoggingHandler, + ) + + # Create mock embedding response data + embedding_response_data = { + "predictions": [ + { + "embeddings": { + "values": [0.1, 0.2, 0.3, 0.4, 0.5], + "statistics": { + "token_count": 10 + } + } + } + ] + } + + # Create mock httpx.Response + mock_httpx_response = Mock() + mock_httpx_response.json.return_value = embedding_response_data + mock_httpx_response.status_code = 200 + + # Create mock logging object + mock_logging_obj = Mock(spec=LiteLLMLoggingObj) + mock_logging_obj.litellm_call_id = "test-call-id-123" + mock_logging_obj.model_call_details = {} + + # Test URL with /predict endpoint + url_route = "/v1/projects/test-project/locations/us-central1/publishers/google/models/textembedding-gecko@001:predict" + + start_time = datetime.datetime.now() + end_time = datetime.datetime.now() + + with patch("litellm.completion_cost") as mock_completion_cost: + # Mock the completion cost calculation + mock_completion_cost.return_value = 0.0001 + + # Call the handler + result = VertexPassthroughLoggingHandler.vertex_passthrough_handler( + httpx_response=mock_httpx_response, + logging_obj=mock_logging_obj, + url_route=url_route, + result="test-result", + start_time=start_time, + end_time=end_time, + cache_hit=False + ) + + # Verify cost tracking was implemented + assert result is not None + assert "result" in result + assert "kwargs" in result + + # Verify cost calculation was called + mock_completion_cost.assert_called_once() + + # Verify cost is set in kwargs + assert "response_cost" in result["kwargs"] + assert result["kwargs"]["response_cost"] == 0.0001 + + # Verify cost is set in logging object + assert "response_cost" in mock_logging_obj.model_call_details + assert mock_logging_obj.model_call_details["response_cost"] == 0.0001 + + # Verify model is set in kwargs + assert "model" in result["kwargs"] + assert result["kwargs"]["model"] == "textembedding-gecko@001" + class TestVertexAIDiscoveryPassThroughHandler: """ From 05955042d57fcc32c73223221d69c8336ecd5994 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 29 Sep 2025 15:09:40 -0700 Subject: [PATCH 12/28] Add model pricing and context window for claude-sonnet-4-5 (#15049) Co-authored-by: Cursor Agent Co-authored-by: ishaan --- model_prices_and_context_window.json | 26 ++++++++++++++++++++++++++ 1 file changed, 26 insertions(+) diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index e69d850447a..3987fe7e511 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -4739,6 +4739,32 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 159 }, + "anthropic/claude-sonnet-4-5": { + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "litellm_provider": "anthropic", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 200000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "tool_use_system_prompt_tokens": 346 + }, "claude-sonnet-4-5-20250929": { "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, From f1f58bd1d1810f267e23ed0462868126de89854e Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Tue, 30 Sep 2025 07:12:56 +0900 Subject: [PATCH 13/28] fix: resolve regression with duplicate Mcp-Protocol-Version header --- litellm/experimental_mcp_client/client.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 1176248d4f1..225349b4e8a 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -194,7 +194,7 @@ class MCPClient: def _get_auth_headers(self) -> dict: """Generate authentication headers based on auth type.""" - headers = {"MCP-Protocol-Version": "2025-06-18"} + headers = {} if self._mcp_auth_value: if isinstance(self._mcp_auth_value, str): From 619577d4e85fe442a25723aea79fecd516b420e3 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 29 Sep 2025 15:15:25 -0700 Subject: [PATCH 14/28] [Feat] Add litellm overhead metric for VertexAI (#15040) * test_litellm_overhead * vertex track overhead * fix config.yaml used for testing * test_litellm_overhead_stream * add update_response_metadata for caching handler * Revert "add update_response_metadata for caching handler" This reverts commit f2a891f2b448b878a5dbf4b5b0a6166c807b3705. --- .../vertex_and_google_ai_studio_gemini.py | 12 ++-- litellm/proxy/proxy_config.yaml | 3 + .../test_litellm_overhead.py | 63 ++++++++++++------- 3 files changed, 48 insertions(+), 30 deletions(-) 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 dc3a6cf15e5..9354ae6e67a 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 @@ -3,7 +3,6 @@ ## Initial implementation - covers gemini + image gen calls import json import time -from litellm._uuid import uuid from copy import deepcopy from functools import partial from typing import ( @@ -25,6 +24,7 @@ import litellm import litellm.litellm_core_utils import litellm.litellm_core_utils.litellm_logging from litellm import verbose_logger +from litellm._uuid import uuid from litellm.constants import ( DEFAULT_REASONING_EFFORT_DISABLE_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET, @@ -32,8 +32,8 @@ from litellm.constants import ( DEFAULT_REASONING_EFFORT_MEDIUM_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH, - DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO, DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE, + DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO, ) from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.llms.custom_httpx.http_handler import ( @@ -1596,7 +1596,7 @@ async def make_call( ) try: - response = await client.post(api_base, headers=headers, data=data, stream=True) + response = await client.post(api_base, headers=headers, data=data, stream=True, logging_obj=logging_obj) response.raise_for_status() except httpx.HTTPStatusError as e: exception_string = str(await e.response.aread()) @@ -1643,7 +1643,7 @@ def make_sync_call( if client is None: client = HTTPHandler() # Create a new client if none provided - response = client.post(api_base, headers=headers, data=data, stream=True) + response = client.post(api_base, headers=headers, data=data, stream=True, logging_obj=logging_obj) if response.status_code != 200 and response.status_code != 201: raise VertexAIError( @@ -1842,7 +1842,7 @@ class VertexLLM(VertexBase): try: response = await client.post( - api_base, headers=headers, json=cast(dict, request_body) + api_base, headers=headers, json=cast(dict, request_body), logging_obj=logging_obj ) # type: ignore response.raise_for_status() except httpx.HTTPStatusError as err: @@ -2045,7 +2045,7 @@ class VertexLLM(VertexBase): client = client try: - response = client.post(url=url, headers=headers, json=data) # type: ignore + response = client.post(url=url, headers=headers, json=data, logging_obj=logging_obj) # type: ignore response.raise_for_status() except httpx.HTTPStatusError as err: error_code = err.response.status_code diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 5445ba3e2b2..60eef9604e1 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -23,6 +23,9 @@ model_list: litellm_params: model: gemini/* api_key: os.environ/GEMINI_API_KEY + - model_name: vertex_ai/* + litellm_params: + model: vertex_ai/* guardrails: diff --git a/tests/litellm_utils_tests/test_litellm_overhead.py b/tests/litellm_utils_tests/test_litellm_overhead.py index 8d0bdf313dd..6e46e935463 100644 --- a/tests/litellm_utils_tests/test_litellm_overhead.py +++ b/tests/litellm_utils_tests/test_litellm_overhead.py @@ -20,23 +20,35 @@ import litellm "openai/gpt-4o", "openai/self_hosted", "bedrock/anthropic.claude-3-5-haiku-20241022-v1:0", + "vertex_ai/gemini-1.0-pro-vision-001", ], ) -async def test_litellm_overhead(model): +async def test_litellm_overhead_non_streaming(model): + """ + - Test we can see the litellm overhead and that it is less than 40% of the total request time + """ litellm._turn_on_debug() start_time = datetime.now() - if model == "openai/self_hosted": - response = await litellm.acompletion( - model=model, - messages=[{"role": "user", "content": "Hello, world!"}], - api_base="https://exampleopenaiendpoint-production.up.railway.app/", - ) - else: - response = await litellm.acompletion( - model=model, - messages=[{"role": "user", "content": "Hello, world!"}], - ) + kwargs ={ + "messages": [{"role": "user", "content": "Hello, world!"}], + "model": model + } + ######################################################### + # Specific cases for models + ######################################################### + if model == "vertex_ai/gemini-1.0-pro-vision-001" or model == "openai/self_hosted": + kwargs["api_base"] = "https://exampleopenaiendpoint-production.up.railway.app/" + # warmup call for auth validation on vertex_ai models + await litellm.acompletion(**kwargs) + + + response = await litellm.acompletion( + **kwargs + ) + ######################################################### + # End of specific cases for models + ######################################################### end_time = datetime.now() total_time_ms = (end_time - start_time).total_seconds() * 1000 print(response) @@ -75,19 +87,22 @@ async def test_litellm_overhead_stream(model): litellm._turn_on_debug() start_time = datetime.now() + kwargs ={ + "messages": [{"role": "user", "content": "Hello, world!"}], + "model": model, + "stream": True, + } + ######################################################### + # Specific cases for models + ######################################################### if model == "openai/self_hosted": - response = await litellm.acompletion( - model=model, - messages=[{"role": "user", "content": "Hello, world!"}], - api_base="https://exampleopenaiendpoint-production.up.railway.app/", - stream=True, - ) - else: - response = await litellm.acompletion( - model=model, - messages=[{"role": "user", "content": "Hello, world!"}], - stream=True, - ) + kwargs["api_base"] = "https://exampleopenaiendpoint-production.up.railway.app/" + # warmup call for auth validation on vertex_ai models + await litellm.acompletion(**kwargs) + + response = await litellm.acompletion( + **kwargs + ) async for chunk in response: print() From 5359a0d6a6d447078cd939d025087f7cb3e461f4 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Tue, 30 Sep 2025 07:25:32 +0900 Subject: [PATCH 15/28] fix: test --- tests/mcp_tests/test_mcp_client_unit.py | 7 +------ 1 file changed, 1 insertion(+), 6 deletions(-) diff --git a/tests/mcp_tests/test_mcp_client_unit.py b/tests/mcp_tests/test_mcp_client_unit.py index 12ed7a30245..cdeee679d18 100644 --- a/tests/mcp_tests/test_mcp_client_unit.py +++ b/tests/mcp_tests/test_mcp_client_unit.py @@ -44,7 +44,6 @@ class TestMCPClientUnitTests: headers = client._get_auth_headers() assert headers == { "Authorization": "Bearer test_token", - "MCP-Protocol-Version": "2025-06-18", } # Basic auth @@ -55,7 +54,6 @@ class TestMCPClientUnitTests: headers = client._get_auth_headers() assert headers == { "Authorization": f"Basic {expected_encoded}", - "MCP-Protocol-Version": "2025-06-18", } # API key @@ -65,7 +63,6 @@ class TestMCPClientUnitTests: headers = client._get_auth_headers() assert headers == { "X-API-Key": "api_key_123", - "MCP-Protocol-Version": "2025-06-18", } # Custom authorization header @@ -77,13 +74,12 @@ class TestMCPClientUnitTests: headers = client._get_auth_headers() assert headers == { "Authorization": "Token custom_token", - "MCP-Protocol-Version": "2025-06-18", } # No auth client = MCPClient("http://example.com") headers = client._get_auth_headers() - assert headers == {"MCP-Protocol-Version": "2025-06-18"} + assert headers == {} @pytest.mark.asyncio @patch("litellm.experimental_mcp_client.client.streamablehttp_client") @@ -112,7 +108,6 @@ class TestMCPClientUnitTests: call_args = mock_transport.call_args assert call_args[1]["headers"] == { "Authorization": "Bearer test_token", - "MCP-Protocol-Version": "2025-06-18", } # Verify session was initialized From e0172b86e2ed2388c6360042d0bdd58ba1a2ec8b Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Mon, 29 Sep 2025 15:48:32 -0700 Subject: [PATCH 16/28] test_litellm_overhead_non_streaming --- ...odel_prices_and_context_window_backup.json | 26 +++++++++++++++++++ .../test_litellm_overhead.py | 8 +++--- 2 files changed, 31 insertions(+), 3 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index e69d850447a..3987fe7e511 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -4739,6 +4739,32 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 159 }, + "anthropic/claude-sonnet-4-5": { + "cache_creation_input_token_cost": 3.75e-06, + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "litellm_provider": "anthropic", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 200000, + "mode": "chat", + "output_cost_per_token": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "tool_use_system_prompt_tokens": 346 + }, "claude-sonnet-4-5-20250929": { "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, diff --git a/tests/litellm_utils_tests/test_litellm_overhead.py b/tests/litellm_utils_tests/test_litellm_overhead.py index 6e46e935463..8b83257f9b5 100644 --- a/tests/litellm_utils_tests/test_litellm_overhead.py +++ b/tests/litellm_utils_tests/test_litellm_overhead.py @@ -20,7 +20,7 @@ import litellm "openai/gpt-4o", "openai/self_hosted", "bedrock/anthropic.claude-3-5-haiku-20241022-v1:0", - "vertex_ai/gemini-1.0-pro-vision-001", + "vertex_ai/gemini-1.5-flash", ], ) async def test_litellm_overhead_non_streaming(model): @@ -37,10 +37,12 @@ async def test_litellm_overhead_non_streaming(model): ######################################################### # Specific cases for models ######################################################### - if model == "vertex_ai/gemini-1.0-pro-vision-001" or model == "openai/self_hosted": - kwargs["api_base"] = "https://exampleopenaiendpoint-production.up.railway.app/" + if model == "vertex_ai/gemini-1.5-flash": + kwargs["api_base"] = "https://exampleopenaiendpoint-production.up.railway.app/v1/projects/pathrise-convert-1606954137718/locations/us-central1/publishers/google/models/gemini-1.0-pro-vision-001" # warmup call for auth validation on vertex_ai models await litellm.acompletion(**kwargs) + if model == "openai/self_hosted": + kwargs["api_base"] = "https://exampleopenaiendpoint-production.up.railway.app/" response = await litellm.acompletion( From d4830e34e58b492cfbb06a780c4710c2ee94d778 Mon Sep 17 00:00:00 2001 From: Alexsander Hamir Date: Mon, 29 Sep 2025 15:49:46 -0700 Subject: [PATCH 17/28] fix: remove router inefficiencies (from O(M*N) to O(1)) - 62.5% faster P99 latency (#15046) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * fix: remove redundant deep copy set_model_list already does the deep copy at the beginning of the call. * fix: remove unused model_list arguments The `model_list` parameter was being passed to classes that did not use it. * fix: reduce per-request memory and time from O(N×M) to O(N) No need to create a whole array for a simple look up. * add: missing test * fix: remove unused parameter --- .../proxy/anthropic_endpoints/endpoints.py | 2 +- .../pass_through_endpoints.py | 2 +- litellm/proxy/route_llm_request.py | 2 +- litellm/router.py | 32 ++++++++++++------- litellm/router_strategy/least_busy.py | 5 ++- litellm/router_strategy/lowest_cost.py | 3 +- litellm/router_strategy/lowest_latency.py | 3 +- litellm/router_strategy/lowest_tpm_rpm.py | 3 +- litellm/router_strategy/lowest_tpm_rpm_v2.py | 3 +- .../local_testing/test_least_busy_routing.py | 4 +-- .../local_testing/test_lowest_cost_routing.py | 6 ++-- .../test_lowest_latency_routing.py | 14 ++++---- .../local_testing/test_tpm_rpm_routing_v2.py | 9 ++---- .../test_router_index_management.py | 24 ++++++++++++++ .../test_lowest_latency_zero_tokens.py | 20 ++---------- 15 files changed, 70 insertions(+), 62 deletions(-) diff --git a/litellm/proxy/anthropic_endpoints/endpoints.py b/litellm/proxy/anthropic_endpoints/endpoints.py index 0dda5cecc83..e7ce5e888c0 100644 --- a/litellm/proxy/anthropic_endpoints/endpoints.py +++ b/litellm/proxy/anthropic_endpoints/endpoints.py @@ -120,7 +120,7 @@ async def anthropic_response( # noqa: PLR0915 ): # model in router deployments, calling a specific deployment on the router llm_coro = llm_router.aanthropic_messages(**data, specific_deployment=True) elif ( - llm_router is not None and data["model"] in llm_router.get_model_ids() + llm_router is not None and llm_router.has_model_id(data["model"]) ): # model in router model list llm_coro = llm_router.aanthropic_messages(**data) elif ( diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index c0042133b47..0eacee3b4f1 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -215,7 +215,7 @@ async def chat_completion_pass_through_endpoint( # noqa: PLR0915 llm_router.aadapter_completion(**data, specific_deployment=True) ) elif ( - llm_router is not None and data["model"] in llm_router.get_model_ids() + llm_router is not None and llm_router.has_model_id(data["model"]) ): # model in router model list llm_response = asyncio.create_task(llm_router.aadapter_completion(**data)) elif ( diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index 2a4281d6357..e1a6ca2a2be 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -130,7 +130,7 @@ async def route_request( elif ( data["model"] in router_model_names - or data["model"] in llm_router.get_model_ids() + or llm_router.has_model_id(data["model"]) ): return getattr(llm_router, f"{route_type}")(**data) diff --git a/litellm/router.py b/litellm/router.py index 3cf99a4b216..0275b636989 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -415,7 +415,6 @@ class Router: if model_list is not None: # Build model index immediately to enable O(1) lookups from the start self._build_model_id_to_deployment_index_map(model_list) - model_list = copy.deepcopy(model_list) self.set_model_list(model_list) self.healthy_deployments: List = self.model_list # type: ignore for m in model_list: @@ -700,7 +699,7 @@ class Router: or routing_strategy == RoutingStrategy.LEAST_BUSY ): self.leastbusy_logger = LeastBusyLoggingHandler( - router_cache=self.cache, model_list=self.model_list + router_cache=self.cache ) ## add callback if isinstance(litellm.input_callback, list): @@ -715,7 +714,6 @@ class Router: ): self.lowesttpm_logger = LowestTPMLoggingHandler( router_cache=self.cache, - model_list=self.model_list, routing_args=routing_strategy_args, ) if isinstance(litellm.callbacks, list): @@ -726,7 +724,6 @@ class Router: ): self.lowesttpm_logger_v2 = LowestTPMLoggingHandler_v2( router_cache=self.cache, - model_list=self.model_list, routing_args=routing_strategy_args, ) if isinstance(litellm.callbacks, list): @@ -737,7 +734,6 @@ class Router: ): self.lowestlatency_logger = LowestLatencyLoggingHandler( router_cache=self.cache, - model_list=self.model_list, routing_args=routing_strategy_args, ) if isinstance(litellm.callbacks, list): @@ -748,7 +744,6 @@ class Router: ): self.lowestcost_logger = LowestCostLoggingHandler( router_cache=self.cache, - model_list=self.model_list, routing_args={}, ) if isinstance(litellm.callbacks, list): @@ -972,7 +967,7 @@ class Router: ### DEPLOYMENT-SPECIFIC PRE-CALL CHECKS ### (e.g. update rpm pre-call. Raise error, if deployment over limit) ## only run if model group given, not model id - if model not in self.get_model_ids(): + if not self.has_model_id(model): self.routing_strategy_pre_call_checks(deployment=deployment) response = litellm.completion( @@ -5331,7 +5326,8 @@ class Router: """ # check if deployment already exists - if deployment.model_info.id in self.get_model_ids(): + _deployment_model_id = deployment.model_info.id + if _deployment_model_id and self.has_model_id(_deployment_model_id): return None # add to model list @@ -6113,7 +6109,7 @@ class Router: if 'model_name' is none, returns all. Returns list of model id's. - """ + """ ids = [] for model in self.model_list: if "model_info" in model and "id" in model["model_info"]: @@ -6126,6 +6122,19 @@ class Router: ids.append(id) return ids + def has_model_id(self, candidate_id: str) -> bool: + """ + O(1) membership check for a deployment ID without allocating large lists. + + Note: Call sites may pass a variable named `model` when it actually + contains a deployment ID. This helper expects the deployment ID string. + + Uses the existing `model_id_to_deployment_index_map` which is kept + in sync by `_build_model_id_to_deployment_index_map` and model-list + mutation helpers. + """ + return candidate_id in self.model_id_to_deployment_index_map + def map_team_model(self, team_model_name: str, team_id: str) -> Optional[str]: """ Map a team model name to a team-specific model name. @@ -6762,14 +6771,13 @@ class Router: # check if aliases set on litellm model alias map if specific_deployment is True: return model, self._get_deployment_by_litellm_model(model=model) - elif model in self.get_model_ids(): + elif self.has_model_id(model): deployment = self.get_deployment(model_id=model) if deployment is not None: deployment_model = deployment.litellm_params.model return deployment_model, deployment.model_dump(exclude_none=True) raise ValueError( - f"LiteLLM Router: Trying to call specific deployment, but Model ID :{model} does not exist in \ - Model ID List: {self.get_model_ids}" + f"LiteLLM Router: Trying to call specific deployment, but Model ID :{model} does not exist in Model ID map" ) _model_from_alias = self._get_model_from_alias(model=model) diff --git a/litellm/router_strategy/least_busy.py b/litellm/router_strategy/least_busy.py index 12f3f01c838..ae0f8433d85 100644 --- a/litellm/router_strategy/least_busy.py +++ b/litellm/router_strategy/least_busy.py @@ -18,10 +18,9 @@ class LeastBusyLoggingHandler(CustomLogger): logged_success: int = 0 logged_failure: int = 0 - def __init__(self, router_cache: DualCache, model_list: list): + def __init__(self, router_cache: DualCache): self.router_cache = router_cache - self.mapping_deployment_to_id: dict = {} - self.model_list = model_list + def log_pre_api_call(self, model, messages, kwargs): """ diff --git a/litellm/router_strategy/lowest_cost.py b/litellm/router_strategy/lowest_cost.py index bd28f6dc5a2..b0612069dfb 100644 --- a/litellm/router_strategy/lowest_cost.py +++ b/litellm/router_strategy/lowest_cost.py @@ -16,10 +16,9 @@ class LowestCostLoggingHandler(CustomLogger): logged_failure: int = 0 def __init__( - self, router_cache: DualCache, model_list: list, routing_args: dict = {} + self, router_cache: DualCache, routing_args: dict = {} ): self.router_cache = router_cache - self.model_list = model_list def log_success_event(self, kwargs, response_obj, start_time, end_time): try: diff --git a/litellm/router_strategy/lowest_latency.py b/litellm/router_strategy/lowest_latency.py index 9e7ab83bf19..7492662e3b1 100644 --- a/litellm/router_strategy/lowest_latency.py +++ b/litellm/router_strategy/lowest_latency.py @@ -32,10 +32,9 @@ class LowestLatencyLoggingHandler(CustomLogger): logged_failure: int = 0 def __init__( - self, router_cache: DualCache, model_list: list, routing_args: dict = {} + self, router_cache: DualCache, routing_args: dict = {} ): self.router_cache = router_cache - self.model_list = model_list self.routing_args = RoutingArgs(**routing_args) def log_success_event( # noqa: PLR0915 diff --git a/litellm/router_strategy/lowest_tpm_rpm.py b/litellm/router_strategy/lowest_tpm_rpm.py index 735ddb3f802..e2bb0d77c4b 100644 --- a/litellm/router_strategy/lowest_tpm_rpm.py +++ b/litellm/router_strategy/lowest_tpm_rpm.py @@ -23,10 +23,9 @@ class LowestTPMLoggingHandler(CustomLogger): default_cache_time_seconds: int = 1 * 60 * 60 # 1 hour def __init__( - self, router_cache: DualCache, model_list: list, routing_args: dict = {} + self, router_cache: DualCache, routing_args: dict = {} ): self.router_cache = router_cache - self.model_list = model_list self.routing_args = RoutingArgs(**routing_args) def log_success_event(self, kwargs, response_obj, start_time, end_time): diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index 9e6c139314f..bf3035fcc9f 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -48,10 +48,9 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): default_cache_time_seconds: int = 1 * 60 * 60 # 1 hour def __init__( - self, router_cache: DualCache, model_list: list, routing_args: dict = {} + self, router_cache: DualCache, routing_args: dict = {} ): self.router_cache = router_cache - self.model_list = model_list self.routing_args = RoutingArgs(**routing_args) BaseRoutingStrategy.__init__( self, diff --git a/tests/local_testing/test_least_busy_routing.py b/tests/local_testing/test_least_busy_routing.py index 5a2fa19562e..b30c4f3943e 100644 --- a/tests/local_testing/test_least_busy_routing.py +++ b/tests/local_testing/test_least_busy_routing.py @@ -28,7 +28,7 @@ from litellm.router_strategy.least_busy import LeastBusyLoggingHandler def test_model_added(): test_cache = DualCache() - least_busy_logger = LeastBusyLoggingHandler(router_cache=test_cache, model_list=[]) + least_busy_logger = LeastBusyLoggingHandler(router_cache=test_cache) kwargs = { "litellm_params": { "metadata": { @@ -45,7 +45,7 @@ def test_model_added(): def test_get_available_deployments(): test_cache = DualCache() - least_busy_logger = LeastBusyLoggingHandler(router_cache=test_cache, model_list=[]) + least_busy_logger = LeastBusyLoggingHandler(router_cache=test_cache) model_group = "gpt-3.5-turbo" deployment = "azure/gpt-4.1-nano" kwargs = { diff --git a/tests/local_testing/test_lowest_cost_routing.py b/tests/local_testing/test_lowest_cost_routing.py index bad8bbbb0a4..3ae123e587d 100644 --- a/tests/local_testing/test_lowest_cost_routing.py +++ b/tests/local_testing/test_lowest_cost_routing.py @@ -36,7 +36,7 @@ async def test_get_available_deployments(): }, ] lowest_cost_logger = LowestCostLoggingHandler( - router_cache=test_cache, model_list=model_list + router_cache=test_cache, ) model_group = "gpt-3.5-turbo" @@ -86,7 +86,7 @@ async def test_get_available_deployments_custom_price(): }, ] lowest_cost_logger = LowestCostLoggingHandler( - router_cache=test_cache, model_list=model_list + router_cache=test_cache, ) model_group = "gpt-3.5-turbo" @@ -187,7 +187,7 @@ async def test_get_available_endpoints_tpm_rpm_check_async(ans_rpm): }, ] lowest_cost_logger = LowestCostLoggingHandler( - router_cache=test_cache, model_list=model_list + router_cache=test_cache ) model_group = "gpt-3.5-turbo" d1 = [(lowest_cost_logger, "1234", 50, 0.01)] * non_ans_rpm diff --git a/tests/local_testing/test_lowest_latency_routing.py b/tests/local_testing/test_lowest_latency_routing.py index 2a7b0eadc42..bb2c02caca4 100644 --- a/tests/local_testing/test_lowest_latency_routing.py +++ b/tests/local_testing/test_lowest_latency_routing.py @@ -38,9 +38,8 @@ async def test_latency_memory_leak(sync_mode): - make 11th call -> no change in memory """ test_cache = DualCache() - model_list = [] lowest_latency_logger = LowestLatencyLoggingHandler( - router_cache=test_cache, model_list=model_list + router_cache=test_cache ) model_group = "gpt-3.5-turbo" deployment_id = "1234" @@ -120,9 +119,8 @@ def get_size(obj, seen=None): def test_latency_updated(): test_cache = DualCache() - model_list = [] lowest_latency_logger = LowestLatencyLoggingHandler( - router_cache=test_cache, model_list=model_list + router_cache=test_cache ) model_group = "gpt-3.5-turbo" deployment_id = "1234" @@ -165,7 +163,7 @@ def test_latency_updated_custom_ttl(): model_list = [] cache_time = 3 lowest_latency_logger = LowestLatencyLoggingHandler( - router_cache=test_cache, model_list=model_list, routing_args={"ttl": cache_time} + router_cache=test_cache, routing_args={"ttl": cache_time} ) model_group = "gpt-3.5-turbo" deployment_id = "1234" @@ -210,7 +208,7 @@ def test_get_available_deployments(): }, ] lowest_latency_logger = LowestLatencyLoggingHandler( - router_cache=test_cache, model_list=model_list + router_cache=test_cache ) model_group = "gpt-3.5-turbo" ## DEPLOYMENT 1 ## @@ -327,7 +325,7 @@ def test_get_available_endpoints_tpm_rpm_check_async(ans_rpm): }, ] lowest_latency_logger = LowestLatencyLoggingHandler( - router_cache=test_cache, model_list=model_list + router_cache=test_cache ) model_group = "gpt-3.5-turbo" d1 = [(lowest_latency_logger, "1234", 50, 0.01)] * non_ans_rpm @@ -376,7 +374,7 @@ def test_get_available_endpoints_tpm_rpm_check(ans_rpm): }, ] lowest_latency_logger = LowestLatencyLoggingHandler( - router_cache=test_cache, model_list=model_list + router_cache=test_cache ) model_group = "gpt-3.5-turbo" ## DEPLOYMENT 1 ## diff --git a/tests/local_testing/test_tpm_rpm_routing_v2.py b/tests/local_testing/test_tpm_rpm_routing_v2.py index 92d2d59785e..a418cd5b0e7 100644 --- a/tests/local_testing/test_tpm_rpm_routing_v2.py +++ b/tests/local_testing/test_tpm_rpm_routing_v2.py @@ -39,9 +39,8 @@ from create_mock_standard_logging_payload import create_standard_logging_payload def test_tpm_rpm_updated(): test_cache = DualCache() - model_list = [] lowest_tpm_logger = LowestTPMLoggingHandler( - router_cache=test_cache, model_list=model_list + router_cache=test_cache ) model_group = "gpt-3.5-turbo" deployment_id = "1234" @@ -110,7 +109,7 @@ def test_get_available_deployments(): }, ] lowest_tpm_logger = LowestTPMLoggingHandler( - router_cache=test_cache, model_list=model_list + router_cache=test_cache ) model_group = "gpt-3.5-turbo" ## DEPLOYMENT 1 ## @@ -668,12 +667,10 @@ def test_return_potential_deployments(): """ Assert deployment at limit is filtered out """ - from litellm.router_strategy.lowest_tpm_rpm_v2 import LowestTPMLoggingHandler_v2 test_cache = DualCache() - model_list = [] lowest_tpm_logger = LowestTPMLoggingHandler( - router_cache=test_cache, model_list=model_list + router_cache=test_cache ) args: Dict = { diff --git a/tests/router_unit_tests/test_router_index_management.py b/tests/router_unit_tests/test_router_index_management.py index bdd52f66ba3..ab39cc1d812 100644 --- a/tests/router_unit_tests/test_router_index_management.py +++ b/tests/router_unit_tests/test_router_index_management.py @@ -103,3 +103,27 @@ class TestRouterIndexManagement: assert router.model_id_to_deployment_index_map["id-1"] == 0 assert router.model_id_to_deployment_index_map["id-2"] == 1 assert router.model_id_to_deployment_index_map["id-3"] == 2 + + def test_has_model_id(self, router): + """Test has_model_id function for O(1) membership check""" + # Setup: Add models to router + router.model_list = [ + {"model": "test1", "model_info": {"id": "model-1"}}, + {"model": "test2", "model_info": {"id": "model-2"}}, + {"model": "test3", "model_info": {"id": "model-3"}} + ] + router.model_id_to_deployment_index_map = {"model-1": 0, "model-2": 1, "model-3": 2} + + # Test: Check existing model IDs + assert router.has_model_id("model-1") == True + assert router.has_model_id("model-2") == True + assert router.has_model_id("model-3") == True + + # Test: Check non-existing model IDs + assert router.has_model_id("non-existent") == False + assert router.has_model_id("") == False + assert router.has_model_id("model-4") == False + + # Test: Empty router + empty_router = Router(model_list=[]) + assert empty_router.has_model_id("any-id") == False diff --git a/tests/test_litellm/test_lowest_latency_zero_tokens.py b/tests/test_litellm/test_lowest_latency_zero_tokens.py index 5a5209e58c9..20ade6caf3d 100644 --- a/tests/test_litellm/test_lowest_latency_zero_tokens.py +++ b/tests/test_litellm/test_lowest_latency_zero_tokens.py @@ -23,16 +23,9 @@ def test_zero_completion_tokens_no_division_error(): (e.g., from Gemini with long contexts) caused ZeroDivisionError """ test_cache = DualCache() - model_list = [ - { - "model_name": "gemini-2.5-flash", - "litellm_params": {"model": "gemini/gemini-2.5-flash"}, - "model_info": {"id": "1234"}, - } - ] - + lowest_latency_logger = LowestLatencyLoggingHandler( - router_cache=test_cache, model_list=model_list + router_cache=test_cache ) deployment_id = "1234" @@ -98,16 +91,9 @@ def test_zero_completion_tokens_with_time_to_first_token(): Test that time_to_first_token calculation also handles zero completion tokens """ test_cache = DualCache() - model_list = [ - { - "model_name": "gemini-2.5-flash", - "litellm_params": {"model": "gemini/gemini-2.5-flash"}, - "model_info": {"id": "1234"}, - } - ] lowest_latency_logger = LowestLatencyLoggingHandler( - router_cache=test_cache, model_list=model_list + router_cache=test_cache ) deployment_id = "1234" From 67abd8880a7c0fa4a3605053d6b05e2614a2ca1f Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Mon, 29 Sep 2025 16:40:46 -0700 Subject: [PATCH 18/28] fix _is_redis_cluster --- .../hooks/parallel_request_limiter_v3.py | 75 +++++++++++++++---- 1 file changed, 61 insertions(+), 14 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index af6f77b3c3d..45681159468 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -4,6 +4,7 @@ This is a rate limiter implementation based on a similar one by Envoy proxy. This is currently in development and not yet ready for production. """ +import binascii import os from datetime import datetime from math import floor @@ -97,6 +98,9 @@ end return results """ +# Redis cluster slot count +REDIS_CLUSTER_SLOTS = 16384 +REDIS_NODE_HASHTAG_NAME="all_keys" class RateLimitDescriptorRateLimitObject(TypedDict, total=False): requests_per_unit: Optional[int] @@ -149,6 +153,20 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60)) + def _is_redis_cluster(self) -> bool: + """ + Check if the dual cache is using Redis cluster. + + Returns: + bool: True if using Redis cluster, False otherwise. + """ + from litellm.caching.redis_cluster_cache import RedisClusterCache + + return ( + self.internal_usage_cache.dual_cache.redis_cache is not None + and isinstance(self.internal_usage_cache.dual_cache.redis_cache, RedisClusterCache) + ) + async def in_memory_cache_sliding_window( self, keys: List[str], @@ -291,26 +309,55 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) return RateLimitResponse(overall_code=overall_code, statuses=statuses) + + def keyslot_for_redis_cluster(self, key: str) -> int: + """ + Compute the Redis Cluster slot for a given key. + + Simple implementation of `HASH_SLOT = CRC16(key) mod 16384` + + Read more about hash slots here: https://medium.com/@linz07m/how-hash-slots-power-data-distribution-in-redis-cluster-bc5b7e74ca7d + + Args: + key (str): The Redis key. + + Returns: + int: The slot number (0-16383). + + + """ + # Handle hash tags: use substring between { and } + start = key.find('{') + if start != -1: + end = key.find('}', start + 1) + if end != -1 and end != start + 1: + key = key[start + 1:end] + + # Compute CRC16 and mod 16384 + crc = binascii.crc_hqx(key.encode('utf-8'), 0) + return crc % REDIS_CLUSTER_SLOTS def _group_keys_by_hash_tag(self, keys: List[str]) -> Dict[str, List[str]]: """ Group keys by their Redis hash tag to ensure cluster compatibility. - Keys with the same hash tag will be processed together. + + For Redis clusters, uses slot calculation to group keys that belong to the same slot. + For regular Redis, no grouping is needed - all keys can be processed together. """ groups: Dict[str, List[str]] = {} - for key in keys: - # Extract hash tag from key like "{api_key:sk-123}:requests" - if "{" in key and "}" in key: - start = key.find("{") - end = key.find("}", start) - hash_tag = key[start : end + 1] - else: - # Fallback for keys without hash tags - hash_tag = "no_hash_tag" - - if hash_tag not in groups: - groups[hash_tag] = [] - groups[hash_tag].append(key) + + # Use slot calculation for Redis clusters only + if self._is_redis_cluster(): + for key in keys: + slot = self.keyslot_for_redis_cluster(key) + slot_key = f"slot_{slot}" + + if slot_key not in groups: + groups[slot_key] = [] + groups[slot_key].append(key) + else: + # For regular Redis, no grouping needed - process all keys together + groups[REDIS_NODE_HASHTAG_NAME] = keys return groups From 52de33787b90bf5d2be9f4321d8008cd8a61d773 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Mon, 29 Sep 2025 16:41:32 -0700 Subject: [PATCH 19/28] test_keyslot_for_redis_cluster --- .../hooks/test_parallel_request_limiter_v3.py | 292 +++++++++++------- 1 file changed, 173 insertions(+), 119 deletions(-) diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 7ebed1b5991..511eb5bbb89 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -1157,19 +1157,18 @@ async def test_async_increment_tokens_fallback_behavior(): # Redis Cluster Compatibility Tests -def test_group_keys_by_hash_tag(): +def test_group_keys_by_hash_tag_regular_redis(): """ - Test that keys are correctly grouped by Redis hash tag for cluster compatibility. + Test that keys are correctly grouped for regular Redis (non-cluster). - This ensures that keys with the same hash tag (e.g., {api_key:sk-123}) are grouped - together so they can be processed in the same Redis cluster slot. + For regular Redis, all keys should be grouped together under a single group. """ local_cache = DualCache() handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) - # Test keys with different hash tags that would cause cluster slot conflicts + # Test keys with different hash tags test_keys = [ "{api_key:sk-123}:window", "{api_key:sk-123}:requests", @@ -1181,32 +1180,77 @@ def test_group_keys_by_hash_tag(): "no_hash_tag_key" ] - # Group the keys + # Group the keys (should be single group for regular Redis) groups = handler._group_keys_by_hash_tag(test_keys) - # Verify correct grouping - expected_groups = { - "{api_key:sk-123}": [ + # Verify all keys are in single group for regular Redis + assert len(groups) == 1, f"Expected 1 group for regular Redis, got {len(groups)}" + assert "all_keys" in groups, "Expected 'all_keys' group for regular Redis" + assert set(groups["all_keys"]) == set(test_keys), "All keys should be in single group" + + +def test_group_keys_by_hash_tag_redis_cluster(): + """ + Test that keys are correctly grouped by Redis cluster slots when using Redis cluster. + + This ensures that keys are grouped by their slot number for cluster compatibility. + """ + from unittest.mock import patch + + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + # Mock _is_redis_cluster to return True + with patch.object(handler, '_is_redis_cluster', return_value=True): + # Test keys with different hash tags + test_keys = [ "{api_key:sk-123}:window", "{api_key:sk-123}:requests", - "{api_key:sk-123}:tokens" - ], - "{user:user-456}": [ "{user:user-456}:window", - "{user:user-456}:requests" - ], - "{team:team-789}": [ - "{team:team-789}:window", - "{team:team-789}:tokens" - ], - "no_hash_tag": ["no_hash_tag_key"] - } + "{user:user-456}:requests", + ] + + # Group the keys (should be grouped by slot for Redis cluster) + groups = handler._group_keys_by_hash_tag(test_keys) + + # Verify keys are grouped by slot + assert len(groups) >= 1, "Should have at least 1 slot group" + + # All group keys should start with "slot_" + for group_key in groups.keys(): + assert group_key.startswith("slot_"), f"Group key {group_key} should start with 'slot_'" + + # Verify all original keys are present across groups + all_grouped_keys = [] + for group_keys in groups.values(): + all_grouped_keys.extend(group_keys) + assert set(all_grouped_keys) == set(test_keys), "All keys should be present in groups" + + +def test_keyslot_for_redis_cluster(): + """ + Test the keyslot calculation for Redis cluster. + """ + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) - assert len(groups) == 4, f"Expected 4 groups, got {len(groups)}" + # Test basic key + slot1 = handler.keyslot_for_redis_cluster("user:1000") + assert 0 <= slot1 < 16384, "Slot should be in valid range" - for expected_tag, expected_keys in expected_groups.items(): - assert expected_tag in groups, f"Missing group {expected_tag}" - assert set(groups[expected_tag]) == set(expected_keys), f"Group {expected_tag} keys mismatch" + # Test key with hash tag + slot2 = handler.keyslot_for_redis_cluster("foo{bar}baz") + slot3 = handler.keyslot_for_redis_cluster("{bar}") + assert slot2 == slot3, "Keys with same hash tag should have same slot" + + # Test keys with same hash tag should have same slot + slot4 = handler.keyslot_for_redis_cluster("{api_key:sk-123}:requests") + slot5 = handler.keyslot_for_redis_cluster("{api_key:sk-123}:window") + assert slot4 == slot5, "Keys with same hash tag should have same slot" @pytest.mark.asyncio @@ -1217,69 +1261,76 @@ async def test_execute_redis_batch_rate_limiter_script_cluster_compatibility(): This simulates the Redis cluster error scenario and verifies fallback behavior. """ - from unittest.mock import AsyncMock + from unittest.mock import AsyncMock, patch local_cache = DualCache() handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) - # Mock script that simulates Redis cluster slot conflict - mock_script = AsyncMock() - mock_script.side_effect = [ - Exception("EVALSHA - all keys must map to the same key slot"), # First group fails - [1234, 1, 1234, 2] # Second group succeeds - ] - handler.batch_rate_limiter_script = mock_script - - # Mock in-memory fallback (returns 2 values for 2 keys: window_start, counter) - handler.in_memory_cache_sliding_window = AsyncMock(return_value=[1234, 1]) - - # Test keys from different hash tags (would fail in cluster without grouping) - test_keys = [ - "{api_key:sk-123}:window", - "{api_key:sk-123}:requests", - "{user:user-456}:window", - "{user:user-456}:requests" - ] - - # Execute the method - results = await handler._execute_redis_batch_rate_limiter_script( - keys_to_fetch=test_keys, - now_int=1234 - ) - - # Verify results: 2 from fallback + 4 from successful script = 6 total - assert len(results) == 6, f"Expected 6 results, got {len(results)}" - - # Verify script was called twice (once per hash tag group) - assert mock_script.call_count == 2 - - # Verify fallback was called for the failed group - handler.in_memory_cache_sliding_window.assert_called_once() - - # Verify the calls were made with grouped keys - call_args_list = mock_script.call_args_list - - # First call should have api_key group keys - first_call_keys = call_args_list[0][1]['keys'] - assert all(key.startswith("{api_key:sk-123}") for key in first_call_keys) - - # Second call should have user group keys - second_call_keys = call_args_list[1][1]['keys'] - assert all(key.startswith("{user:user-456}") for key in second_call_keys) + # Mock _is_redis_cluster to return True for this test + with patch.object(handler, '_is_redis_cluster', return_value=True): + # Mock script that simulates Redis cluster slot conflict + mock_script = AsyncMock() + mock_script.side_effect = [ + Exception("EVALSHA - all keys must map to the same key slot"), # First group fails + [1234, 1, 1234, 2] # Second group succeeds + ] + handler.batch_rate_limiter_script = mock_script + + # Mock in-memory fallback (returns 2 values for 2 keys: window_start, counter) + handler.in_memory_cache_sliding_window = AsyncMock(return_value=[1234, 1]) + + # Test keys from different hash tags (would fail in cluster without grouping) + test_keys = [ + "{api_key:sk-123}:window", + "{api_key:sk-123}:requests", + "{user:user-456}:window", + "{user:user-456}:requests" + ] + + # Execute the method + results = await handler._execute_redis_batch_rate_limiter_script( + keys_to_fetch=test_keys, + now_int=1234 + ) + + # Verify results: 2 from fallback + 4 from successful script = 6 total + assert len(results) == 6, f"Expected 6 results, got {len(results)}" + + # Verify script was called twice (once per slot group) + assert mock_script.call_count == 2 + + # Verify fallback was called for the failed group + handler.in_memory_cache_sliding_window.assert_called_once() + + # Verify the calls were made with grouped keys + call_args_list = mock_script.call_args_list + + # Both calls should have keys, but we can't predict exact grouping without knowing slots + # Just verify that keys were grouped and calls were made + assert len(call_args_list) == 2, "Should have made 2 script calls" + + # Verify all keys were processed + all_processed_keys = [] + for call_args in call_args_list: + all_processed_keys.extend(call_args[1]['keys']) + + # Should have processed all keys (some might be duplicated due to fallback) + unique_processed_keys = set(all_processed_keys) + assert len(unique_processed_keys) >= 2, "Should have processed at least some keys" @pytest.mark.asyncio async def test_execute_token_increment_script_cluster_compatibility(): """ Test that token increment script execution handles Redis cluster compatibility - by grouping operations by hash tag. + by grouping operations by slot. This ensures token increments work correctly in cluster environments. """ from typing import List - from unittest.mock import AsyncMock + from unittest.mock import AsyncMock, patch from litellm.types.caching import RedisPipelineIncrementOperation @@ -1288,52 +1339,55 @@ async def test_execute_token_increment_script_cluster_compatibility(): internal_usage_cache=InternalUsageCache(local_cache) ) - # Mock script - mock_script = AsyncMock() - handler.token_increment_script = mock_script - - # Create pipeline operations with different hash tags - pipeline_operations: List[RedisPipelineIncrementOperation] = [ - { - "key": "{api_key:sk-123}:tokens", - "increment_value": 100, - "ttl": 60 - }, - { - "key": "{api_key:sk-123}:max_parallel_requests", - "increment_value": -1, - "ttl": 60 - }, - { - "key": "{user:user-456}:tokens", - "increment_value": 50, - "ttl": 60 + # Mock _is_redis_cluster to return True for this test + with patch.object(handler, '_is_redis_cluster', return_value=True): + # Mock script + mock_script = AsyncMock() + handler.token_increment_script = mock_script + + # Create pipeline operations with different hash tags + pipeline_operations: List[RedisPipelineIncrementOperation] = [ + { + "key": "{api_key:sk-123}:tokens", + "increment_value": 100, + "ttl": 60 + }, + { + "key": "{api_key:sk-123}:max_parallel_requests", + "increment_value": -1, + "ttl": 60 + }, + { + "key": "{user:user-456}:tokens", + "increment_value": 50, + "ttl": 60 + } + ] + + # Execute the method + await handler._execute_token_increment_script(pipeline_operations) + + # Verify script was called (at least once, possibly more depending on slot grouping) + assert mock_script.call_count >= 1, "Script should be called at least once" + + call_args_list = mock_script.call_args_list + + # Verify all operations were processed + all_processed_keys = [] + for call_args in call_args_list: + all_processed_keys.extend(call_args[1]['keys']) + + # Should have processed all 3 keys + expected_keys = { + "{api_key:sk-123}:tokens", + "{api_key:sk-123}:max_parallel_requests", + "{user:user-456}:tokens" } - ] - - # Execute the method - await handler._execute_token_increment_script(pipeline_operations) - - # Verify script was called twice (once per hash tag group) - assert mock_script.call_count == 2 - - call_args_list = mock_script.call_args_list - - # Verify first call has api_key operations - first_call_keys = call_args_list[0][1]['keys'] - assert len(first_call_keys) == 2 - assert all(key.startswith("{api_key:sk-123}") for key in first_call_keys) - - # Verify second call has user operations - second_call_keys = call_args_list[1][1]['keys'] - assert len(second_call_keys) == 1 - assert second_call_keys[0] == "{user:user-456}:tokens" - - # Verify args are correctly mapped - first_call_args = call_args_list[0][1]['args'] - assert len(first_call_args) == 4 # 2 operations * 2 args each (increment_value, ttl) - assert first_call_args == [100, 60, -1, 60] # increment_value, ttl for each operation - - second_call_args = call_args_list[1][1]['args'] - assert len(second_call_args) == 2 # 1 operation * 2 args - assert second_call_args == [50, 60] + assert set(all_processed_keys) == expected_keys, "All operation keys should be processed" + + # Verify args structure is correct for each call + for call_args in call_args_list: + keys = call_args[1]['keys'] + args = call_args[1]['args'] + # Each key should have 2 args (increment_value, ttl) + assert len(args) == len(keys) * 2, f"Each key should have 2 args, got {len(args)} args for {len(keys)} keys" From f30386088134bf81bc9b4f7ae9b94b92d5f1318d Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Mon, 29 Sep 2025 16:42:15 -0700 Subject: [PATCH 20/28] Revert "test_keyslot_for_redis_cluster" This reverts commit 52de33787b90bf5d2be9f4321d8008cd8a61d773. --- .../hooks/test_parallel_request_limiter_v3.py | 292 +++++++----------- 1 file changed, 119 insertions(+), 173 deletions(-) diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 511eb5bbb89..7ebed1b5991 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -1157,18 +1157,19 @@ async def test_async_increment_tokens_fallback_behavior(): # Redis Cluster Compatibility Tests -def test_group_keys_by_hash_tag_regular_redis(): +def test_group_keys_by_hash_tag(): """ - Test that keys are correctly grouped for regular Redis (non-cluster). + Test that keys are correctly grouped by Redis hash tag for cluster compatibility. - For regular Redis, all keys should be grouped together under a single group. + This ensures that keys with the same hash tag (e.g., {api_key:sk-123}) are grouped + together so they can be processed in the same Redis cluster slot. """ local_cache = DualCache() handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) - # Test keys with different hash tags + # Test keys with different hash tags that would cause cluster slot conflicts test_keys = [ "{api_key:sk-123}:window", "{api_key:sk-123}:requests", @@ -1180,77 +1181,32 @@ def test_group_keys_by_hash_tag_regular_redis(): "no_hash_tag_key" ] - # Group the keys (should be single group for regular Redis) + # Group the keys groups = handler._group_keys_by_hash_tag(test_keys) - # Verify all keys are in single group for regular Redis - assert len(groups) == 1, f"Expected 1 group for regular Redis, got {len(groups)}" - assert "all_keys" in groups, "Expected 'all_keys' group for regular Redis" - assert set(groups["all_keys"]) == set(test_keys), "All keys should be in single group" - - -def test_group_keys_by_hash_tag_redis_cluster(): - """ - Test that keys are correctly grouped by Redis cluster slots when using Redis cluster. - - This ensures that keys are grouped by their slot number for cluster compatibility. - """ - from unittest.mock import patch - - local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) - - # Mock _is_redis_cluster to return True - with patch.object(handler, '_is_redis_cluster', return_value=True): - # Test keys with different hash tags - test_keys = [ + # Verify correct grouping + expected_groups = { + "{api_key:sk-123}": [ "{api_key:sk-123}:window", "{api_key:sk-123}:requests", + "{api_key:sk-123}:tokens" + ], + "{user:user-456}": [ "{user:user-456}:window", - "{user:user-456}:requests", - ] - - # Group the keys (should be grouped by slot for Redis cluster) - groups = handler._group_keys_by_hash_tag(test_keys) - - # Verify keys are grouped by slot - assert len(groups) >= 1, "Should have at least 1 slot group" - - # All group keys should start with "slot_" - for group_key in groups.keys(): - assert group_key.startswith("slot_"), f"Group key {group_key} should start with 'slot_'" - - # Verify all original keys are present across groups - all_grouped_keys = [] - for group_keys in groups.values(): - all_grouped_keys.extend(group_keys) - assert set(all_grouped_keys) == set(test_keys), "All keys should be present in groups" - - -def test_keyslot_for_redis_cluster(): - """ - Test the keyslot calculation for Redis cluster. - """ - local_cache = DualCache() - handler = _PROXY_MaxParallelRequestsHandler( - internal_usage_cache=InternalUsageCache(local_cache) - ) + "{user:user-456}:requests" + ], + "{team:team-789}": [ + "{team:team-789}:window", + "{team:team-789}:tokens" + ], + "no_hash_tag": ["no_hash_tag_key"] + } - # Test basic key - slot1 = handler.keyslot_for_redis_cluster("user:1000") - assert 0 <= slot1 < 16384, "Slot should be in valid range" + assert len(groups) == 4, f"Expected 4 groups, got {len(groups)}" - # Test key with hash tag - slot2 = handler.keyslot_for_redis_cluster("foo{bar}baz") - slot3 = handler.keyslot_for_redis_cluster("{bar}") - assert slot2 == slot3, "Keys with same hash tag should have same slot" - - # Test keys with same hash tag should have same slot - slot4 = handler.keyslot_for_redis_cluster("{api_key:sk-123}:requests") - slot5 = handler.keyslot_for_redis_cluster("{api_key:sk-123}:window") - assert slot4 == slot5, "Keys with same hash tag should have same slot" + for expected_tag, expected_keys in expected_groups.items(): + assert expected_tag in groups, f"Missing group {expected_tag}" + assert set(groups[expected_tag]) == set(expected_keys), f"Group {expected_tag} keys mismatch" @pytest.mark.asyncio @@ -1261,76 +1217,69 @@ async def test_execute_redis_batch_rate_limiter_script_cluster_compatibility(): This simulates the Redis cluster error scenario and verifies fallback behavior. """ - from unittest.mock import AsyncMock, patch + from unittest.mock import AsyncMock local_cache = DualCache() handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) - # Mock _is_redis_cluster to return True for this test - with patch.object(handler, '_is_redis_cluster', return_value=True): - # Mock script that simulates Redis cluster slot conflict - mock_script = AsyncMock() - mock_script.side_effect = [ - Exception("EVALSHA - all keys must map to the same key slot"), # First group fails - [1234, 1, 1234, 2] # Second group succeeds - ] - handler.batch_rate_limiter_script = mock_script - - # Mock in-memory fallback (returns 2 values for 2 keys: window_start, counter) - handler.in_memory_cache_sliding_window = AsyncMock(return_value=[1234, 1]) - - # Test keys from different hash tags (would fail in cluster without grouping) - test_keys = [ - "{api_key:sk-123}:window", - "{api_key:sk-123}:requests", - "{user:user-456}:window", - "{user:user-456}:requests" - ] - - # Execute the method - results = await handler._execute_redis_batch_rate_limiter_script( - keys_to_fetch=test_keys, - now_int=1234 - ) - - # Verify results: 2 from fallback + 4 from successful script = 6 total - assert len(results) == 6, f"Expected 6 results, got {len(results)}" - - # Verify script was called twice (once per slot group) - assert mock_script.call_count == 2 - - # Verify fallback was called for the failed group - handler.in_memory_cache_sliding_window.assert_called_once() - - # Verify the calls were made with grouped keys - call_args_list = mock_script.call_args_list - - # Both calls should have keys, but we can't predict exact grouping without knowing slots - # Just verify that keys were grouped and calls were made - assert len(call_args_list) == 2, "Should have made 2 script calls" - - # Verify all keys were processed - all_processed_keys = [] - for call_args in call_args_list: - all_processed_keys.extend(call_args[1]['keys']) - - # Should have processed all keys (some might be duplicated due to fallback) - unique_processed_keys = set(all_processed_keys) - assert len(unique_processed_keys) >= 2, "Should have processed at least some keys" + # Mock script that simulates Redis cluster slot conflict + mock_script = AsyncMock() + mock_script.side_effect = [ + Exception("EVALSHA - all keys must map to the same key slot"), # First group fails + [1234, 1, 1234, 2] # Second group succeeds + ] + handler.batch_rate_limiter_script = mock_script + + # Mock in-memory fallback (returns 2 values for 2 keys: window_start, counter) + handler.in_memory_cache_sliding_window = AsyncMock(return_value=[1234, 1]) + + # Test keys from different hash tags (would fail in cluster without grouping) + test_keys = [ + "{api_key:sk-123}:window", + "{api_key:sk-123}:requests", + "{user:user-456}:window", + "{user:user-456}:requests" + ] + + # Execute the method + results = await handler._execute_redis_batch_rate_limiter_script( + keys_to_fetch=test_keys, + now_int=1234 + ) + + # Verify results: 2 from fallback + 4 from successful script = 6 total + assert len(results) == 6, f"Expected 6 results, got {len(results)}" + + # Verify script was called twice (once per hash tag group) + assert mock_script.call_count == 2 + + # Verify fallback was called for the failed group + handler.in_memory_cache_sliding_window.assert_called_once() + + # Verify the calls were made with grouped keys + call_args_list = mock_script.call_args_list + + # First call should have api_key group keys + first_call_keys = call_args_list[0][1]['keys'] + assert all(key.startswith("{api_key:sk-123}") for key in first_call_keys) + + # Second call should have user group keys + second_call_keys = call_args_list[1][1]['keys'] + assert all(key.startswith("{user:user-456}") for key in second_call_keys) @pytest.mark.asyncio async def test_execute_token_increment_script_cluster_compatibility(): """ Test that token increment script execution handles Redis cluster compatibility - by grouping operations by slot. + by grouping operations by hash tag. This ensures token increments work correctly in cluster environments. """ from typing import List - from unittest.mock import AsyncMock, patch + from unittest.mock import AsyncMock from litellm.types.caching import RedisPipelineIncrementOperation @@ -1339,55 +1288,52 @@ async def test_execute_token_increment_script_cluster_compatibility(): internal_usage_cache=InternalUsageCache(local_cache) ) - # Mock _is_redis_cluster to return True for this test - with patch.object(handler, '_is_redis_cluster', return_value=True): - # Mock script - mock_script = AsyncMock() - handler.token_increment_script = mock_script - - # Create pipeline operations with different hash tags - pipeline_operations: List[RedisPipelineIncrementOperation] = [ - { - "key": "{api_key:sk-123}:tokens", - "increment_value": 100, - "ttl": 60 - }, - { - "key": "{api_key:sk-123}:max_parallel_requests", - "increment_value": -1, - "ttl": 60 - }, - { - "key": "{user:user-456}:tokens", - "increment_value": 50, - "ttl": 60 - } - ] - - # Execute the method - await handler._execute_token_increment_script(pipeline_operations) - - # Verify script was called (at least once, possibly more depending on slot grouping) - assert mock_script.call_count >= 1, "Script should be called at least once" - - call_args_list = mock_script.call_args_list - - # Verify all operations were processed - all_processed_keys = [] - for call_args in call_args_list: - all_processed_keys.extend(call_args[1]['keys']) - - # Should have processed all 3 keys - expected_keys = { - "{api_key:sk-123}:tokens", - "{api_key:sk-123}:max_parallel_requests", - "{user:user-456}:tokens" + # Mock script + mock_script = AsyncMock() + handler.token_increment_script = mock_script + + # Create pipeline operations with different hash tags + pipeline_operations: List[RedisPipelineIncrementOperation] = [ + { + "key": "{api_key:sk-123}:tokens", + "increment_value": 100, + "ttl": 60 + }, + { + "key": "{api_key:sk-123}:max_parallel_requests", + "increment_value": -1, + "ttl": 60 + }, + { + "key": "{user:user-456}:tokens", + "increment_value": 50, + "ttl": 60 } - assert set(all_processed_keys) == expected_keys, "All operation keys should be processed" - - # Verify args structure is correct for each call - for call_args in call_args_list: - keys = call_args[1]['keys'] - args = call_args[1]['args'] - # Each key should have 2 args (increment_value, ttl) - assert len(args) == len(keys) * 2, f"Each key should have 2 args, got {len(args)} args for {len(keys)} keys" + ] + + # Execute the method + await handler._execute_token_increment_script(pipeline_operations) + + # Verify script was called twice (once per hash tag group) + assert mock_script.call_count == 2 + + call_args_list = mock_script.call_args_list + + # Verify first call has api_key operations + first_call_keys = call_args_list[0][1]['keys'] + assert len(first_call_keys) == 2 + assert all(key.startswith("{api_key:sk-123}") for key in first_call_keys) + + # Verify second call has user operations + second_call_keys = call_args_list[1][1]['keys'] + assert len(second_call_keys) == 1 + assert second_call_keys[0] == "{user:user-456}:tokens" + + # Verify args are correctly mapped + first_call_args = call_args_list[0][1]['args'] + assert len(first_call_args) == 4 # 2 operations * 2 args each (increment_value, ttl) + assert first_call_args == [100, 60, -1, 60] # increment_value, ttl for each operation + + second_call_args = call_args_list[1][1]['args'] + assert len(second_call_args) == 2 # 1 operation * 2 args + assert second_call_args == [50, 60] From 55110ba6ae7e544681a31b1e13678ccdd15982ab Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Mon, 29 Sep 2025 16:42:31 -0700 Subject: [PATCH 21/28] Revert "fix _is_redis_cluster" This reverts commit 67abd8880a7c0fa4a3605053d6b05e2614a2ca1f. --- .../hooks/parallel_request_limiter_v3.py | 75 ++++--------------- 1 file changed, 14 insertions(+), 61 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index 45681159468..af6f77b3c3d 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -4,7 +4,6 @@ This is a rate limiter implementation based on a similar one by Envoy proxy. This is currently in development and not yet ready for production. """ -import binascii import os from datetime import datetime from math import floor @@ -98,9 +97,6 @@ end return results """ -# Redis cluster slot count -REDIS_CLUSTER_SLOTS = 16384 -REDIS_NODE_HASHTAG_NAME="all_keys" class RateLimitDescriptorRateLimitObject(TypedDict, total=False): requests_per_unit: Optional[int] @@ -153,20 +149,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60)) - def _is_redis_cluster(self) -> bool: - """ - Check if the dual cache is using Redis cluster. - - Returns: - bool: True if using Redis cluster, False otherwise. - """ - from litellm.caching.redis_cluster_cache import RedisClusterCache - - return ( - self.internal_usage_cache.dual_cache.redis_cache is not None - and isinstance(self.internal_usage_cache.dual_cache.redis_cache, RedisClusterCache) - ) - async def in_memory_cache_sliding_window( self, keys: List[str], @@ -309,55 +291,26 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) return RateLimitResponse(overall_code=overall_code, statuses=statuses) - - def keyslot_for_redis_cluster(self, key: str) -> int: - """ - Compute the Redis Cluster slot for a given key. - - Simple implementation of `HASH_SLOT = CRC16(key) mod 16384` - - Read more about hash slots here: https://medium.com/@linz07m/how-hash-slots-power-data-distribution-in-redis-cluster-bc5b7e74ca7d - - Args: - key (str): The Redis key. - - Returns: - int: The slot number (0-16383). - - - """ - # Handle hash tags: use substring between { and } - start = key.find('{') - if start != -1: - end = key.find('}', start + 1) - if end != -1 and end != start + 1: - key = key[start + 1:end] - - # Compute CRC16 and mod 16384 - crc = binascii.crc_hqx(key.encode('utf-8'), 0) - return crc % REDIS_CLUSTER_SLOTS def _group_keys_by_hash_tag(self, keys: List[str]) -> Dict[str, List[str]]: """ Group keys by their Redis hash tag to ensure cluster compatibility. - - For Redis clusters, uses slot calculation to group keys that belong to the same slot. - For regular Redis, no grouping is needed - all keys can be processed together. + Keys with the same hash tag will be processed together. """ groups: Dict[str, List[str]] = {} - - # Use slot calculation for Redis clusters only - if self._is_redis_cluster(): - for key in keys: - slot = self.keyslot_for_redis_cluster(key) - slot_key = f"slot_{slot}" - - if slot_key not in groups: - groups[slot_key] = [] - groups[slot_key].append(key) - else: - # For regular Redis, no grouping needed - process all keys together - groups[REDIS_NODE_HASHTAG_NAME] = keys + for key in keys: + # Extract hash tag from key like "{api_key:sk-123}:requests" + if "{" in key and "}" in key: + start = key.find("{") + end = key.find("}", start) + hash_tag = key[start : end + 1] + else: + # Fallback for keys without hash tags + hash_tag = "no_hash_tag" + + if hash_tag not in groups: + groups[hash_tag] = [] + groups[hash_tag].append(key) return groups From f6d768326166ea08cfdf124729ec0ce3edb8af54 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 29 Sep 2025 17:33:27 -0700 Subject: [PATCH 22/28] [Feat] LiteLLM Overhead metric tracking - Add support for tracking litellm overhead on cache hits (#15045) * test_litellm_overhead * vertex track overhead * fix config.yaml used for testing * test_litellm_overhead_stream * add update_response_metadata for caching handler * add CachingDetails * fix update_response_metadata import * add CachingDetails metrics * add CachingDetails * test_litellm_overhead_cache_hit * test_litellm_overhead_cache_hit * test_litellm_overhead_cache_hit --- litellm/caching/caching_handler.py | 33 +++++++++++++++++ litellm/litellm_core_utils/litellm_logging.py | 4 +++ .../llm_response_utils/response_metadata.py | 30 ++++++++++++++-- litellm/types/utils.py | 12 +++++++ litellm/utils.py | 29 +++------------ .../test_litellm_overhead.py | 36 +++++++++++++++++++ 6 files changed, 117 insertions(+), 27 deletions(-) diff --git a/litellm/caching/caching_handler.py b/litellm/caching/caching_handler.py index 9526c4a2f39..b151ebd6513 100644 --- a/litellm/caching/caching_handler.py +++ b/litellm/caching/caching_handler.py @@ -36,12 +36,16 @@ import litellm from litellm._logging import print_verbose, verbose_logger from litellm.caching import InMemoryCache from litellm.caching.caching import S3Cache +from litellm.litellm_core_utils.llm_response_utils.response_metadata import ( + update_response_metadata, +) from litellm.litellm_core_utils.logging_utils import ( _assemble_complete_response_from_streaming_chunks, ) from litellm.types.caching import CachedEmbedding from litellm.types.rerank import RerankResponse from litellm.types.utils import ( + CachingDetails, CallTypes, Embedding, EmbeddingResponse, @@ -136,6 +140,13 @@ class LLMCachingHandler: kwargs = kwargs.copy() args = args or () + ######################################################### + # Init cache timing metrics + ######################################################### + cache_check_start_time = datetime.datetime.now() + cache_check_end_time = None + ######################################################### + parent_otel_span = _get_parent_otel_span_from_kwargs(kwargs) kwargs["parent_otel_span"] = parent_otel_span @@ -157,6 +168,7 @@ class LLMCachingHandler: kwargs=kwargs, args=args, ) + cache_check_end_time = datetime.datetime.now() if cached_result is not None and not isinstance(cached_result, list): verbose_logger.debug("Cache Hit!") @@ -168,6 +180,7 @@ class LLMCachingHandler: api_base=kwargs.get("api_base", None), api_key=kwargs.get("api_key", None), ) + cache_duration_ms = (cache_check_end_time - cache_check_start_time).total_seconds() * 1000 self._update_litellm_logging_obj_environment( logging_obj=logging_obj, model=model, @@ -175,10 +188,12 @@ class LLMCachingHandler: cached_result=cached_result, is_async=True, custom_llm_provider=custom_llm_provider, + cache_duration_ms=cache_duration_ms, ) call_type = original_function.__name__ + cached_result = self._convert_cached_result_to_model_response( cached_result=cached_result, call_type=call_type, @@ -716,6 +731,18 @@ class LLMCachingHandler: and isinstance(cached_result._hidden_params, dict) ): cached_result._hidden_params["cache_hit"] = True + + ######################################################### + # Add final timing metrics to the cached result + ######################################################### + update_response_metadata( + result=cached_result, + logging_obj=logging_obj, + model=model, + kwargs=kwargs, + start_time=self.start_time, + end_time=datetime.datetime.now(), + ) return cached_result def _convert_cached_stream_response( @@ -944,6 +971,7 @@ class LLMCachingHandler: is_async: bool, is_embedding: bool = False, custom_llm_provider: Optional[str] = None, + cache_duration_ms: Optional[float] = None, ): """ Helper function to update the LiteLLMLoggingObj environment variables. @@ -995,6 +1023,11 @@ class LLMCachingHandler: custom_llm_provider=custom_llm_provider, ) + logging_obj.caching_details = CachingDetails( + cache_hit=True, + cache_duration_ms=cache_duration_ms, + ) + def convert_args_to_kwargs( original_function: Callable, diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index bbadc9c8183..24449e1bd0f 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -83,6 +83,7 @@ from litellm.types.mcp import MCPPostCallResponseObject from litellm.types.rerank import RerankResponse from litellm.types.router import CustomPricingLiteLLMParams from litellm.types.utils import ( + CachingDetails, CallTypes, CostBreakdown, CostResponseTypes, @@ -348,6 +349,9 @@ class Logging(LiteLLMLoggingBaseClass): # Initialize cost breakdown field self.cost_breakdown: Optional[CostBreakdown] = None + # Init Caching related details + self.caching_details: Optional[CachingDetails] = None + self.model_call_details: Dict[str, Any] = { "litellm_trace_id": litellm_trace_id, "litellm_call_id": litellm_call_id, diff --git a/litellm/litellm_core_utils/llm_response_utils/response_metadata.py b/litellm/litellm_core_utils/llm_response_utils/response_metadata.py index b1085c684fc..c5ef7237628 100644 --- a/litellm/litellm_core_utils/llm_response_utils/response_metadata.py +++ b/litellm/litellm_core_utils/llm_response_utils/response_metadata.py @@ -85,15 +85,37 @@ class ResponseMetadata: # Set total response time if supported if self.supports_response_time: self.result._response_ms = total_response_time_ms + + ######################################################### + # 1. Add _response_ms total duration + ######################################################### + self._update_hidden_params( + { + "_response_ms": total_response_time_ms, + } + ) - # Calculate LiteLLM overhead + ######################################################### + # 2. Add LiteLLM overhead duration + ######################################################### llm_api_duration_ms = logging_obj.model_call_details.get("llm_api_duration_ms") if llm_api_duration_ms is not None: overhead_ms = round(total_response_time_ms - llm_api_duration_ms, 4) self._update_hidden_params( { "litellm_overhead_time_ms": overhead_ms, - "_response_ms": total_response_time_ms, + } + ) + + ######################################################### + # 3. Add duration for reading from cache + # In this case overhead from litellm is the difference between the cache read duration and the total response time + ######################################################### + if logging_obj.caching_details is not None and logging_obj.caching_details.get("cache_hit") is True and (cache_duration_ms := logging_obj.caching_details.get("cache_duration_ms")) is not None: + overhead_ms = total_response_time_ms - cache_duration_ms + self._update_hidden_params( + { + "litellm_overhead_time_ms": overhead_ms, } ) @@ -113,6 +135,10 @@ def update_response_metadata( ) -> None: """ Updates response metadata including hidden params and timing metrics + Updates response metadata, adds the following: + - response._hidden_params + - response._hidden_params["litellm_overhead_time_ms"] + - response.response_time_ms """ if result is None: return diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 16a69d049c6..b0183249ba2 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2059,6 +2059,18 @@ class StandardLoggingGuardrailInformation(TypedDict, total=False): StandardLoggingPayloadStatus = Literal["success", "failure"] +class CachingDetails(TypedDict): + """ + Track all caching related metrics, fields for a given request + """ + cache_hit: Optional[bool] + """ + Whether the request hit the cache + """ + cache_duration_ms: Optional[float] + """ + Duration for reading from cache + """ class CostBreakdown(TypedDict): """ diff --git a/litellm/utils.py b/litellm/utils.py index cb340735b4e..9fdde626ae5 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7,7 +7,6 @@ # # Thank you users! We ❤️ you! - Krrish & Ishaan -from io import StringIO import ast import asyncio import base64 @@ -37,6 +36,7 @@ from dataclasses import dataclass, field from functools import lru_cache, wraps from importlib import resources from inspect import iscoroutine +from io import StringIO from os.path import abspath, dirname, join import aiohttp @@ -232,6 +232,9 @@ from typing import ( from openai import OpenAIError as OriginalError +from litellm.litellm_core_utils.llm_response_utils.response_metadata import ( + update_response_metadata, +) from litellm.litellm_core_utils.thread_pool_executor import executor from litellm.litellm_core_utils.token_counter import token_counter as token_counter_new from litellm.llms.base_llm.anthropic_messages.transformation import ( @@ -1677,30 +1680,6 @@ def _is_streaming_request( return False -def update_response_metadata( - result: Any, - logging_obj: LiteLLMLoggingObject, - model: Optional[str], - kwargs: dict, - start_time: datetime.datetime, - end_time: datetime.datetime, -) -> None: - """ - Updates response metadata, adds the following: - - response._hidden_params - - response._hidden_params["litellm_overhead_time_ms"] - - response.response_time_ms - """ - if result is None: - return - - metadata = ResponseMetadata(result) - metadata.set_hidden_params(logging_obj=logging_obj, model=model, kwargs=kwargs) - metadata.set_timing_metrics( - start_time=start_time, end_time=end_time, logging_obj=logging_obj - ) - metadata.apply() - def _select_tokenizer( model: str, custom_tokenizer: Optional[CustomHuggingfaceTokenizer] = None diff --git a/tests/litellm_utils_tests/test_litellm_overhead.py b/tests/litellm_utils_tests/test_litellm_overhead.py index 8b83257f9b5..e3472de1848 100644 --- a/tests/litellm_utils_tests/test_litellm_overhead.py +++ b/tests/litellm_utils_tests/test_litellm_overhead.py @@ -5,6 +5,7 @@ import time from datetime import datetime from unittest.mock import AsyncMock, patch, MagicMock import pytest +import asyncio sys.path.insert( 0, os.path.abspath("../..") @@ -75,6 +76,7 @@ async def test_litellm_overhead_non_streaming(model): pass + @pytest.mark.asyncio @pytest.mark.parametrize( "model", @@ -131,3 +133,37 @@ async def test_litellm_overhead_stream(model): assert overhead_percent < 40 pass + + +@pytest.mark.asyncio +async def test_litellm_overhead_cache_hit(): + """ + Test that litellm overhead is tracked on cache hits. + Makes two identical requests and checks that the second one (cache hit) has overhead in hidden params. + """ + from litellm.caching.caching import Cache + + litellm._turn_on_debug() + litellm.cache = Cache() + print("test2 for caching") + litellm.set_verbose = True + messages = [{"role": "user", "content": "Hello, world! Cache test"}] + response1 = await litellm.acompletion(model="gpt-4.1-nano", messages=messages, caching=True) + await asyncio.sleep(2) + # Wait for any pending background tasks to complete + pending_tasks = [task for task in asyncio.all_tasks() if not task.done()] + print("all pending tasks", pending_tasks) + if pending_tasks: + await asyncio.wait(pending_tasks, timeout=1.0) + + response2 = await litellm.acompletion(model="gpt-4.1-nano", messages=messages, caching=True) + print("RESPONSE 1", response1) + print("RESPONSE 2", response2) + assert response1.id == response2.id + + print("response 2 hidden params", response2._hidden_params) + + + assert "_response_ms" in response2._hidden_params + total_time_ms = response2._hidden_params["_response_ms"] + assert response2._hidden_params["litellm_overhead_time_ms"] > 0 and response2._hidden_params["litellm_overhead_time_ms"] < total_time_ms \ No newline at end of file From ebf72f5eb9f7de7fc6a4da15a80181767d4af779 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 29 Sep 2025 18:12:44 -0700 Subject: [PATCH 23/28] [Fix] Parallel Request Limiter v3 - use well known redis cluster hashing algorithm (#15052) * test_keyslot_for_redis_cluster * fix _is_redis_cluster * Update litellm/proxy/hooks/parallel_request_limiter_v3.py Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --------- Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com> --- .../hooks/parallel_request_limiter_v3.py | 75 ++++- .../hooks/test_parallel_request_limiter_v3.py | 292 +++++++++++------- 2 files changed, 234 insertions(+), 133 deletions(-) diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py index af6f77b3c3d..eda380b5165 100644 --- a/litellm/proxy/hooks/parallel_request_limiter_v3.py +++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py @@ -4,6 +4,7 @@ This is a rate limiter implementation based on a similar one by Envoy proxy. This is currently in development and not yet ready for production. """ +import binascii import os from datetime import datetime from math import floor @@ -97,6 +98,9 @@ end return results """ +# Redis cluster slot count +REDIS_CLUSTER_SLOTS = 16384 +REDIS_NODE_HASHTAG_NAME = "all_keys" class RateLimitDescriptorRateLimitObject(TypedDict, total=False): requests_per_unit: Optional[int] @@ -149,6 +153,20 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60)) + def _is_redis_cluster(self) -> bool: + """ + Check if the dual cache is using Redis cluster. + + Returns: + bool: True if using Redis cluster, False otherwise. + """ + from litellm.caching.redis_cluster_cache import RedisClusterCache + + return ( + self.internal_usage_cache.dual_cache.redis_cache is not None + and isinstance(self.internal_usage_cache.dual_cache.redis_cache, RedisClusterCache) + ) + async def in_memory_cache_sliding_window( self, keys: List[str], @@ -291,26 +309,55 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger): ) return RateLimitResponse(overall_code=overall_code, statuses=statuses) + + def keyslot_for_redis_cluster(self, key: str) -> int: + """ + Compute the Redis Cluster slot for a given key. + + Simple implementation of `HASH_SLOT = CRC16(key) mod 16384` + + Read more about hash slots here: https://medium.com/@linz07m/how-hash-slots-power-data-distribution-in-redis-cluster-bc5b7e74ca7d + + Args: + key (str): The Redis key. + + Returns: + int: The slot number (0-16383). + + + """ + # Handle hash tags: use substring between { and } + start = key.find('{') + if start != -1: + end = key.find('}', start + 1) + if end != -1 and end != start + 1: + key = key[start + 1:end] + + # Compute CRC16 and mod 16384 + crc = binascii.crc_hqx(key.encode('utf-8'), 0) + return crc % REDIS_CLUSTER_SLOTS def _group_keys_by_hash_tag(self, keys: List[str]) -> Dict[str, List[str]]: """ Group keys by their Redis hash tag to ensure cluster compatibility. - Keys with the same hash tag will be processed together. + + For Redis clusters, uses slot calculation to group keys that belong to the same slot. + For regular Redis, no grouping is needed - all keys can be processed together. """ groups: Dict[str, List[str]] = {} - for key in keys: - # Extract hash tag from key like "{api_key:sk-123}:requests" - if "{" in key and "}" in key: - start = key.find("{") - end = key.find("}", start) - hash_tag = key[start : end + 1] - else: - # Fallback for keys without hash tags - hash_tag = "no_hash_tag" - - if hash_tag not in groups: - groups[hash_tag] = [] - groups[hash_tag].append(key) + + # Use slot calculation for Redis clusters only + if self._is_redis_cluster(): + for key in keys: + slot = self.keyslot_for_redis_cluster(key) + slot_key = f"slot_{slot}" + + if slot_key not in groups: + groups[slot_key] = [] + groups[slot_key].append(key) + else: + # For regular Redis, no grouping needed - process all keys together + groups[REDIS_NODE_HASHTAG_NAME] = keys return groups diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py index 7ebed1b5991..511eb5bbb89 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py @@ -1157,19 +1157,18 @@ async def test_async_increment_tokens_fallback_behavior(): # Redis Cluster Compatibility Tests -def test_group_keys_by_hash_tag(): +def test_group_keys_by_hash_tag_regular_redis(): """ - Test that keys are correctly grouped by Redis hash tag for cluster compatibility. + Test that keys are correctly grouped for regular Redis (non-cluster). - This ensures that keys with the same hash tag (e.g., {api_key:sk-123}) are grouped - together so they can be processed in the same Redis cluster slot. + For regular Redis, all keys should be grouped together under a single group. """ local_cache = DualCache() handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) - # Test keys with different hash tags that would cause cluster slot conflicts + # Test keys with different hash tags test_keys = [ "{api_key:sk-123}:window", "{api_key:sk-123}:requests", @@ -1181,32 +1180,77 @@ def test_group_keys_by_hash_tag(): "no_hash_tag_key" ] - # Group the keys + # Group the keys (should be single group for regular Redis) groups = handler._group_keys_by_hash_tag(test_keys) - # Verify correct grouping - expected_groups = { - "{api_key:sk-123}": [ + # Verify all keys are in single group for regular Redis + assert len(groups) == 1, f"Expected 1 group for regular Redis, got {len(groups)}" + assert "all_keys" in groups, "Expected 'all_keys' group for regular Redis" + assert set(groups["all_keys"]) == set(test_keys), "All keys should be in single group" + + +def test_group_keys_by_hash_tag_redis_cluster(): + """ + Test that keys are correctly grouped by Redis cluster slots when using Redis cluster. + + This ensures that keys are grouped by their slot number for cluster compatibility. + """ + from unittest.mock import patch + + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) + + # Mock _is_redis_cluster to return True + with patch.object(handler, '_is_redis_cluster', return_value=True): + # Test keys with different hash tags + test_keys = [ "{api_key:sk-123}:window", "{api_key:sk-123}:requests", - "{api_key:sk-123}:tokens" - ], - "{user:user-456}": [ "{user:user-456}:window", - "{user:user-456}:requests" - ], - "{team:team-789}": [ - "{team:team-789}:window", - "{team:team-789}:tokens" - ], - "no_hash_tag": ["no_hash_tag_key"] - } + "{user:user-456}:requests", + ] + + # Group the keys (should be grouped by slot for Redis cluster) + groups = handler._group_keys_by_hash_tag(test_keys) + + # Verify keys are grouped by slot + assert len(groups) >= 1, "Should have at least 1 slot group" + + # All group keys should start with "slot_" + for group_key in groups.keys(): + assert group_key.startswith("slot_"), f"Group key {group_key} should start with 'slot_'" + + # Verify all original keys are present across groups + all_grouped_keys = [] + for group_keys in groups.values(): + all_grouped_keys.extend(group_keys) + assert set(all_grouped_keys) == set(test_keys), "All keys should be present in groups" + + +def test_keyslot_for_redis_cluster(): + """ + Test the keyslot calculation for Redis cluster. + """ + local_cache = DualCache() + handler = _PROXY_MaxParallelRequestsHandler( + internal_usage_cache=InternalUsageCache(local_cache) + ) - assert len(groups) == 4, f"Expected 4 groups, got {len(groups)}" + # Test basic key + slot1 = handler.keyslot_for_redis_cluster("user:1000") + assert 0 <= slot1 < 16384, "Slot should be in valid range" - for expected_tag, expected_keys in expected_groups.items(): - assert expected_tag in groups, f"Missing group {expected_tag}" - assert set(groups[expected_tag]) == set(expected_keys), f"Group {expected_tag} keys mismatch" + # Test key with hash tag + slot2 = handler.keyslot_for_redis_cluster("foo{bar}baz") + slot3 = handler.keyslot_for_redis_cluster("{bar}") + assert slot2 == slot3, "Keys with same hash tag should have same slot" + + # Test keys with same hash tag should have same slot + slot4 = handler.keyslot_for_redis_cluster("{api_key:sk-123}:requests") + slot5 = handler.keyslot_for_redis_cluster("{api_key:sk-123}:window") + assert slot4 == slot5, "Keys with same hash tag should have same slot" @pytest.mark.asyncio @@ -1217,69 +1261,76 @@ async def test_execute_redis_batch_rate_limiter_script_cluster_compatibility(): This simulates the Redis cluster error scenario and verifies fallback behavior. """ - from unittest.mock import AsyncMock + from unittest.mock import AsyncMock, patch local_cache = DualCache() handler = _PROXY_MaxParallelRequestsHandler( internal_usage_cache=InternalUsageCache(local_cache) ) - # Mock script that simulates Redis cluster slot conflict - mock_script = AsyncMock() - mock_script.side_effect = [ - Exception("EVALSHA - all keys must map to the same key slot"), # First group fails - [1234, 1, 1234, 2] # Second group succeeds - ] - handler.batch_rate_limiter_script = mock_script - - # Mock in-memory fallback (returns 2 values for 2 keys: window_start, counter) - handler.in_memory_cache_sliding_window = AsyncMock(return_value=[1234, 1]) - - # Test keys from different hash tags (would fail in cluster without grouping) - test_keys = [ - "{api_key:sk-123}:window", - "{api_key:sk-123}:requests", - "{user:user-456}:window", - "{user:user-456}:requests" - ] - - # Execute the method - results = await handler._execute_redis_batch_rate_limiter_script( - keys_to_fetch=test_keys, - now_int=1234 - ) - - # Verify results: 2 from fallback + 4 from successful script = 6 total - assert len(results) == 6, f"Expected 6 results, got {len(results)}" - - # Verify script was called twice (once per hash tag group) - assert mock_script.call_count == 2 - - # Verify fallback was called for the failed group - handler.in_memory_cache_sliding_window.assert_called_once() - - # Verify the calls were made with grouped keys - call_args_list = mock_script.call_args_list - - # First call should have api_key group keys - first_call_keys = call_args_list[0][1]['keys'] - assert all(key.startswith("{api_key:sk-123}") for key in first_call_keys) - - # Second call should have user group keys - second_call_keys = call_args_list[1][1]['keys'] - assert all(key.startswith("{user:user-456}") for key in second_call_keys) + # Mock _is_redis_cluster to return True for this test + with patch.object(handler, '_is_redis_cluster', return_value=True): + # Mock script that simulates Redis cluster slot conflict + mock_script = AsyncMock() + mock_script.side_effect = [ + Exception("EVALSHA - all keys must map to the same key slot"), # First group fails + [1234, 1, 1234, 2] # Second group succeeds + ] + handler.batch_rate_limiter_script = mock_script + + # Mock in-memory fallback (returns 2 values for 2 keys: window_start, counter) + handler.in_memory_cache_sliding_window = AsyncMock(return_value=[1234, 1]) + + # Test keys from different hash tags (would fail in cluster without grouping) + test_keys = [ + "{api_key:sk-123}:window", + "{api_key:sk-123}:requests", + "{user:user-456}:window", + "{user:user-456}:requests" + ] + + # Execute the method + results = await handler._execute_redis_batch_rate_limiter_script( + keys_to_fetch=test_keys, + now_int=1234 + ) + + # Verify results: 2 from fallback + 4 from successful script = 6 total + assert len(results) == 6, f"Expected 6 results, got {len(results)}" + + # Verify script was called twice (once per slot group) + assert mock_script.call_count == 2 + + # Verify fallback was called for the failed group + handler.in_memory_cache_sliding_window.assert_called_once() + + # Verify the calls were made with grouped keys + call_args_list = mock_script.call_args_list + + # Both calls should have keys, but we can't predict exact grouping without knowing slots + # Just verify that keys were grouped and calls were made + assert len(call_args_list) == 2, "Should have made 2 script calls" + + # Verify all keys were processed + all_processed_keys = [] + for call_args in call_args_list: + all_processed_keys.extend(call_args[1]['keys']) + + # Should have processed all keys (some might be duplicated due to fallback) + unique_processed_keys = set(all_processed_keys) + assert len(unique_processed_keys) >= 2, "Should have processed at least some keys" @pytest.mark.asyncio async def test_execute_token_increment_script_cluster_compatibility(): """ Test that token increment script execution handles Redis cluster compatibility - by grouping operations by hash tag. + by grouping operations by slot. This ensures token increments work correctly in cluster environments. """ from typing import List - from unittest.mock import AsyncMock + from unittest.mock import AsyncMock, patch from litellm.types.caching import RedisPipelineIncrementOperation @@ -1288,52 +1339,55 @@ async def test_execute_token_increment_script_cluster_compatibility(): internal_usage_cache=InternalUsageCache(local_cache) ) - # Mock script - mock_script = AsyncMock() - handler.token_increment_script = mock_script - - # Create pipeline operations with different hash tags - pipeline_operations: List[RedisPipelineIncrementOperation] = [ - { - "key": "{api_key:sk-123}:tokens", - "increment_value": 100, - "ttl": 60 - }, - { - "key": "{api_key:sk-123}:max_parallel_requests", - "increment_value": -1, - "ttl": 60 - }, - { - "key": "{user:user-456}:tokens", - "increment_value": 50, - "ttl": 60 + # Mock _is_redis_cluster to return True for this test + with patch.object(handler, '_is_redis_cluster', return_value=True): + # Mock script + mock_script = AsyncMock() + handler.token_increment_script = mock_script + + # Create pipeline operations with different hash tags + pipeline_operations: List[RedisPipelineIncrementOperation] = [ + { + "key": "{api_key:sk-123}:tokens", + "increment_value": 100, + "ttl": 60 + }, + { + "key": "{api_key:sk-123}:max_parallel_requests", + "increment_value": -1, + "ttl": 60 + }, + { + "key": "{user:user-456}:tokens", + "increment_value": 50, + "ttl": 60 + } + ] + + # Execute the method + await handler._execute_token_increment_script(pipeline_operations) + + # Verify script was called (at least once, possibly more depending on slot grouping) + assert mock_script.call_count >= 1, "Script should be called at least once" + + call_args_list = mock_script.call_args_list + + # Verify all operations were processed + all_processed_keys = [] + for call_args in call_args_list: + all_processed_keys.extend(call_args[1]['keys']) + + # Should have processed all 3 keys + expected_keys = { + "{api_key:sk-123}:tokens", + "{api_key:sk-123}:max_parallel_requests", + "{user:user-456}:tokens" } - ] - - # Execute the method - await handler._execute_token_increment_script(pipeline_operations) - - # Verify script was called twice (once per hash tag group) - assert mock_script.call_count == 2 - - call_args_list = mock_script.call_args_list - - # Verify first call has api_key operations - first_call_keys = call_args_list[0][1]['keys'] - assert len(first_call_keys) == 2 - assert all(key.startswith("{api_key:sk-123}") for key in first_call_keys) - - # Verify second call has user operations - second_call_keys = call_args_list[1][1]['keys'] - assert len(second_call_keys) == 1 - assert second_call_keys[0] == "{user:user-456}:tokens" - - # Verify args are correctly mapped - first_call_args = call_args_list[0][1]['args'] - assert len(first_call_args) == 4 # 2 operations * 2 args each (increment_value, ttl) - assert first_call_args == [100, 60, -1, 60] # increment_value, ttl for each operation - - second_call_args = call_args_list[1][1]['args'] - assert len(second_call_args) == 2 # 1 operation * 2 args - assert second_call_args == [50, 60] + assert set(all_processed_keys) == expected_keys, "All operation keys should be processed" + + # Verify args structure is correct for each call + for call_args in call_args_list: + keys = call_args[1]['keys'] + args = call_args[1]['args'] + # Each key should have 2 args (increment_value, ttl) + assert len(args) == len(keys) * 2, f"Each key should have 2 args, got {len(args)} args for {len(keys)} keys" From f22fd4cddd77bebd257ea98e3f2010120e0a3762 Mon Sep 17 00:00:00 2001 From: Copilot <198982749+Copilot@users.noreply.github.com> Date: Mon, 29 Sep 2025 18:16:52 -0700 Subject: [PATCH 24/28] Fix: Add /v1/messages/count_tokens to Anthropic routes for non-admin user access (#15034) * Initial plan * Fix: Add /v1/messages/count_tokens to Anthropic routes for user access Co-authored-by: ishaan-jaff <29436595+ishaan-jaff@users.noreply.github.com> --------- Co-authored-by: copilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com> Co-authored-by: ishaan-jaff <29436595+ishaan-jaff@users.noreply.github.com> --- litellm/proxy/_types.py | 1 + .../proxy/auth/test_route_checks.py | 19 +++++++++++++++++++ 2 files changed, 20 insertions(+) diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index c5b4ee0753e..c5370eb7d70 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -330,6 +330,7 @@ class LiteLLMRoutes(enum.Enum): anthropic_routes = [ "/v1/messages", + "/v1/messages/count_tokens", ] mcp_routes = [ diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index ac09917e4cd..539ee4a9ba8 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -228,3 +228,22 @@ def test_virtual_key_allowed_routes_with_no_member_names_only_explicit(): ) assert "Virtual key is not allowed to call this route" in str(exc_info.value) + + +def test_anthropic_count_tokens_route_is_llm_api_route(): + """Test that /v1/messages/count_tokens is recognized as an LLM API route for Anthropic""" + + # Test the core anthropic routes + assert RouteChecks.is_llm_api_route("/v1/messages") is True + assert RouteChecks.is_llm_api_route("/v1/messages/count_tokens") is True + + +def test_anthropic_count_tokens_route_accessible_to_internal_users(): + """Test that internal users can access the Anthropic count_tokens route""" + + # Test that the route is recognized as an LLM API route (which means it's accessible to internal users) + # This is the core check that was failing in the original issue + assert RouteChecks.is_llm_api_route("/v1/messages/count_tokens") is True + + # Also test that the regular messages route still works + assert RouteChecks.is_llm_api_route("/v1/messages") is True From f1578b49e23e14f0e987c56d0dbd1cf2772f1fe8 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Mon, 29 Sep 2025 18:25:54 -0700 Subject: [PATCH 25/28] vertex_httpx_mock_post --- tests/local_testing/test_amazing_vertex_completion.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/tests/local_testing/test_amazing_vertex_completion.py b/tests/local_testing/test_amazing_vertex_completion.py index d7b1f95d5f6..8262d43a0d6 100644 --- a/tests/local_testing/test_amazing_vertex_completion.py +++ b/tests/local_testing/test_amazing_vertex_completion.py @@ -977,7 +977,7 @@ def vertex_httpx_mock_reject_prompt_post(*args, **kwargs): # @pytest.mark.skip(reason="exhausted vertex quota. need to refactor to mock the call") -def vertex_httpx_mock_post(url, data=None, json=None, headers=None): +def vertex_httpx_mock_post(url, data=None, json=None, headers=None, **kwargs): mock_response = MagicMock() mock_response.status_code = 200 mock_response.headers = {"Content-Type": "application/json"} From 3e474b9e8161250970902fcd3b96d8f1e78ba7f2 Mon Sep 17 00:00:00 2001 From: Ishaan Jaffer Date: Mon, 29 Sep 2025 18:26:29 -0700 Subject: [PATCH 26/28] fix claude-sonnet-4-5 model cost map --- litellm/model_prices_and_context_window_backup.json | 2 +- model_prices_and_context_window.json | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 3987fe7e511..dc40ffa1562 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -4739,7 +4739,7 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 159 }, - "anthropic/claude-sonnet-4-5": { + "claude-sonnet-4-5": { "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 3987fe7e511..dc40ffa1562 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -4739,7 +4739,7 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 159 }, - "anthropic/claude-sonnet-4-5": { + "claude-sonnet-4-5": { "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, "input_cost_per_token": 3e-06, From 708c0bd78db66850a8a96265df5ecd6759d54cb2 Mon Sep 17 00:00:00 2001 From: Ishaan Jaff Date: Mon, 29 Sep 2025 19:47:04 -0700 Subject: [PATCH 27/28] [Feat] Return Cost for Responses API Streaming requests (#15053) * test_basic_openai_responses_api_streaming * _transform_chat_completion_usage_to_responses_usage * ResponseAPIUsage.cost * test fixes for anthropic cost with /responses * fix mypy typng --- litellm/proxy/proxy_config.yaml | 1 + .../streaming_iterator.py | 11 ++++++++++- .../transformation.py | 9 ++++++++- litellm/responses/streaming_iterator.py | 16 ++++++++++++++++ litellm/types/llms/openai.py | 3 +++ .../base_responses_api.py | 10 ++++++++++ 6 files changed, 48 insertions(+), 2 deletions(-) diff --git a/litellm/proxy/proxy_config.yaml b/litellm/proxy/proxy_config.yaml index 60eef9604e1..73177fdd482 100644 --- a/litellm/proxy/proxy_config.yaml +++ b/litellm/proxy/proxy_config.yaml @@ -42,6 +42,7 @@ guardrails: litellm_settings: callbacks: ["datadog"] + include_cost_in_streaming_usage: true datadog_params: turn_off_message_logging: true datadog_llm_observability_params: diff --git a/litellm/responses/litellm_completion_transformation/streaming_iterator.py b/litellm/responses/litellm_completion_transformation/streaming_iterator.py index 64ea93028f6..93abdee778f 100644 --- a/litellm/responses/litellm_completion_transformation/streaming_iterator.py +++ b/litellm/responses/litellm_completion_transformation/streaming_iterator.py @@ -49,6 +49,7 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): self.litellm_metadata: Optional[dict] = litellm_metadata or {} self.collected_chat_completion_chunks: List[ModelResponseStream] = [] self.finished: bool = False + self.litellm_logging_obj = litellm_custom_stream_wrapper.logging_obj async def __anext__( self, @@ -167,8 +168,16 @@ class LiteLLMCompletionStreamingIterator(ResponsesAPIStreamingIterator): def _emit_response_completed_event(self) -> Optional[ResponseCompletedEvent]: litellm_model_response: Optional[ Union[ModelResponse, TextCompletionResponse] - ] = stream_chunk_builder(chunks=self.collected_chat_completion_chunks) + ] = stream_chunk_builder(chunks=self.collected_chat_completion_chunks, logging_obj=self.litellm_logging_obj) if litellm_model_response and isinstance(litellm_model_response, ModelResponse): + # Add cost to usage object if include_cost_in_streaming_usage is True + if litellm.include_cost_in_streaming_usage and self.litellm_logging_obj is not None: + usage = getattr(litellm_model_response, "usage", None) + if usage is not None: + setattr( + usage, "cost", self.litellm_logging_obj._response_cost_calculator(result=litellm_model_response) + ) + # Transform the response responses_api_response = LiteLLMCompletionResponsesConfig.transform_chat_completion_response_to_responses_api_response( request_input=self.request_input, diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index 82d3980b370..a43e02a0f3e 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -851,8 +851,15 @@ class LiteLLMCompletionResponsesConfig: output_tokens=0, total_tokens=0, ) - return ResponseAPIUsage( + + response_usage = ResponseAPIUsage( input_tokens=usage.prompt_tokens, output_tokens=usage.completion_tokens, total_tokens=usage.total_tokens, ) + + # Preserve cost field if it exists (for streaming usage with cost calculation) + if hasattr(usage, "cost") and usage.cost is not None: + setattr(response_usage, "cost", usage.cost) + + return response_usage diff --git a/litellm/responses/streaming_iterator.py b/litellm/responses/streaming_iterator.py index e9e41789f09..eda3e6921da 100644 --- a/litellm/responses/streaming_iterator.py +++ b/litellm/responses/streaming_iterator.py @@ -5,6 +5,7 @@ from typing import Any, Dict, Optional import httpx +import litellm from litellm.constants import STREAM_SSE_DONE_STRING from litellm.litellm_core_utils.asyncify import run_async_function from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj @@ -13,6 +14,7 @@ from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfi from litellm.responses.utils import ResponsesAPIRequestUtils from litellm.types.llms.openai import ( OutputTextDeltaEvent, + ResponseAPIUsage, ResponseCompletedEvent, ResponsesAPIResponse, ResponsesAPIStreamEvents, @@ -95,6 +97,20 @@ class BaseResponsesAPIStreamingIterator: == ResponsesAPIStreamEvents.RESPONSE_COMPLETED ): self.completed_response = openai_responses_api_chunk + # Add cost to usage object if include_cost_in_streaming_usage is True + if litellm.include_cost_in_streaming_usage and self.logging_obj is not None: + response_obj: Optional[ResponsesAPIResponse] = getattr(openai_responses_api_chunk, "response", None) + if response_obj: + usage_obj: Optional[ResponseAPIUsage] = getattr(response_obj, "usage", None) + if usage_obj is not None: + try: + cost: Optional[float] = self.logging_obj._response_cost_calculator(result=response_obj) + if cost is not None: + setattr(usage_obj, "cost", cost) + except Exception: + # If cost calculation fails, continue without cost + pass + self._handle_logging_completed_response() return openai_responses_api_chunk diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 434035b809e..9f4ae03b39d 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -1033,6 +1033,9 @@ class ResponseAPIUsage(BaseLiteLLMOpenAIResponseObject): total_tokens: int """The total number of tokens used.""" + cost: Optional[float] = None + """The cost of the request.""" + model_config = {"extra": "allow"} diff --git a/tests/llm_responses_api_testing/base_responses_api.py b/tests/llm_responses_api_testing/base_responses_api.py index 855aeff246c..afb30f15f08 100644 --- a/tests/llm_responses_api_testing/base_responses_api.py +++ b/tests/llm_responses_api_testing/base_responses_api.py @@ -146,6 +146,8 @@ class BaseResponsesAPITest(ABC): @pytest.mark.flaky(retries=3, delay=2) async def test_basic_openai_responses_api_streaming(self, sync_mode): litellm._turn_on_debug() + # Enable cost calculation for streaming usage + litellm.include_cost_in_streaming_usage = True base_completion_call_args = self.get_base_completion_call_args() collected_content_string = "" response_completed_event = None @@ -208,6 +210,14 @@ class BaseResponsesAPITest(ABC): + response_completed_event.response.usage.output_tokens ) + # assert the response completed event includes cost when include_cost_in_streaming_usage is True + assert hasattr(response_completed_event.response.usage, "cost"), "Cost should be included in streaming responses API usage object" + assert response_completed_event.response.usage.cost > 0, "Cost should be greater than 0" + print(f"Cost found in streaming response: {response_completed_event.response.usage.cost}") + + # Reset the setting + litellm.include_cost_in_streaming_usage = False + @pytest.mark.parametrize("sync_mode", [False, True]) @pytest.mark.asyncio async def test_basic_openai_responses_delete_endpoint(self, sync_mode): From def4afedd7a4b4c02f2552a1fa763f8cd9a383c5 Mon Sep 17 00:00:00 2001 From: Yuta Saito Date: Tue, 30 Sep 2025 12:52:17 +0900 Subject: [PATCH 28/28] doc: add missing api_key parameter --- docs/my-website/docs/providers/bedrock.md | 1 + 1 file changed, 1 insertion(+) diff --git a/docs/my-website/docs/providers/bedrock.md b/docs/my-website/docs/providers/bedrock.md index 50d32a45df3..28cae80cc42 100644 --- a/docs/my-website/docs/providers/bedrock.md +++ b/docs/my-website/docs/providers/bedrock.md @@ -101,6 +101,7 @@ aws_profile_name: Optional[str], aws_role_name: Optional[str], aws_web_identity_token: Optional[str], aws_bedrock_runtime_endpoint: Optional[str], +api_key: Optional[str], ``` ### 2. Start the proxy