diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 0db1118a7b4..71f961b94c0 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -1082,6 +1082,25 @@ def _rewrite_google_maps_response_format(data: RequestBody) -> None: _rewrite_mime_type_to_response_format(generation_config) +_SERVICE_TIER_ENUM_PREFIX = "service_tier_" + + +def _transform_service_tier( + service_tier: str, + custom_llm_provider: Literal["vertex_ai", "vertex_ai_beta", "gemini"], +) -> str: + """ + Google AI Studio accepts the lowercase enum names ('standard', 'flex', 'priority'), while + Vertex AI validates against google.cloud.aiplatform.v1.ServiceTier, whose names are prefixed + and uppercased ('SERVICE_TIER_FLEX'). + """ + normalized = service_tier.lower().removeprefix(_SERVICE_TIER_ENUM_PREFIX) + tier = "standard" if normalized == "default" else normalized + if custom_llm_provider == LlmProviders.GEMINI: + return tier + return f"{_SERVICE_TIER_ENUM_PREFIX}{tier}".upper() + + def _transform_request_body( messages: List[AllMessageValues], model: str, @@ -1177,10 +1196,9 @@ def _transform_request_body( if service_tier := optional_params.pop("service_tier", None): if isinstance(service_tier, str): - if service_tier.lower() == "default": - data["serviceTier"] = "standard" - else: - data["serviceTier"] = service_tier.lower() + data["serviceTier"] = _transform_service_tier( + service_tier=service_tier, custom_llm_provider=custom_llm_provider + ) else: data["serviceTier"] = service_tier diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py index d99c190c6e5..c717ec09eb2 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_ai_gemini_transformation.py @@ -1,3 +1,5 @@ +import pytest + from litellm.litellm_core_utils.prompt_templates.factory import ( convert_to_gemini_tool_call_result, ) @@ -197,23 +199,54 @@ def test_vertex_ai_includes_labels(): assert result["labels"] == {"project": "test", "team": "ai"} -def test_service_tier_forwarded_to_vertex_ai(): - """Test that service_tier in optional_params is mapped to serviceTier in request body.""" - messages = [{"role": "user", "content": "test"}] - optional_params = {"service_tier": "flex"} - litellm_params = {} - +@pytest.mark.parametrize( + "service_tier, expected", + [ + ("flex", "SERVICE_TIER_FLEX"), + ("FLEX", "SERVICE_TIER_FLEX"), + ("priority", "SERVICE_TIER_PRIORITY"), + ("standard", "SERVICE_TIER_STANDARD"), + ("default", "SERVICE_TIER_STANDARD"), + ("SERVICE_TIER_FLEX", "SERVICE_TIER_FLEX"), + ], +) +def test_service_tier_forwarded_to_vertex_ai_as_proto_enum_name(service_tier, expected): + """ + Vertex AI validates serviceTier against google.cloud.aiplatform.v1.ServiceTier, so anything + other than the prefixed, uppercased enum name is rejected with a 400 + """ result = _transform_request_body( - messages=messages, + messages=[{"role": "user", "content": "test"}], model="gemini-2.5-pro", - optional_params=optional_params, + optional_params={"service_tier": service_tier}, custom_llm_provider="vertex_ai", - litellm_params=litellm_params, + litellm_params={}, cached_content=None, ) - assert "serviceTier" in result - assert result["serviceTier"] == "flex" + assert result["serviceTier"] == expected + + +@pytest.mark.parametrize( + "service_tier, expected", + [ + ("flex", "flex"), + ("PRIORITY", "priority"), + ("default", "standard"), + ], +) +def test_service_tier_forwarded_to_google_ai_studio_as_lowercase(service_tier, expected): + """Google AI Studio's generativelanguage API only accepts the lowercase enum names""" + result = _transform_request_body( + messages=[{"role": "user", "content": "test"}], + model="gemini-2.5-pro", + optional_params={"service_tier": service_tier}, + custom_llm_provider="gemini", + litellm_params={}, + cached_content=None, + ) + + assert result["serviceTier"] == expected def test_extra_body_cache_not_forwarded_to_vertex_ai():