mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-14 23:21:35 +00:00
fix(vertex_ai): send serviceTier as the ServiceTier proto enum name
This commit is contained in:
parent
daf22ec871
commit
19a5a3f2dc
2 changed files with 66 additions and 15 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue