feat: support service_tier in gemini

This commit is contained in:
Mark McDonald 2026-03-24 13:11:43 +08:00 • committed by Sameer Kankute
parent 885fea2c58
commit 25b266fe2c
No known key found for this signature in database
6 changed files with 112 additions and 0 deletions

View file

@ -91,6 +91,7 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
"modalities",
"parallel_tool_calls",
"web_search_options",
"service_tier",
]
if supports_reasoning(model, custom_llm_provider="gemini"):
supported_params.append("reasoning_effort")

View file

@ -763,6 +763,10 @@ def _transform_request_body( # noqa: PLR0915
data["generationConfig"] = generation_config
if cached_content is not None:
data["cachedContent"] = cached_content
if service_tier := optional_params.pop("service_tier", None):
data["serviceTier"] = service_tier
# Only add labels for Vertex AI endpoints (not Google GenAI/AI Studio) and only if non-empty
if labels and custom_llm_provider != LlmProviders.GEMINI:
data["labels"] = labels

View file

@ -318,6 +318,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
"parallel_tool_calls",
"web_search_options",
"include_server_side_tool_invocations",
"service_tier",
]
# Add penalty parameters only for non-preview models
@ -1121,6 +1122,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
optional_params = self._add_tools_to_optional_params(
optional_params, [_tools]
)
elif param == "service_tier" and isinstance(value, str):
optional_params["service_tier"] = value
elif param == "include_server_side_tool_invocations" and value is True:
optional_params["include_server_side_tool_invocations"] = True
if litellm.vertex_ai_safety_settings is not None:
@ -2415,6 +2418,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
"provider_specific_fields", {}
)["traffic_type"] = traffic_type
## ADD SERVICE TIER ##
if "serviceTier" in completion_response:
setattr(model_response, "service_tier", completion_response["serviceTier"])
except Exception as e:
raise VertexAIError(
message="Received={}, Error converting to valid response block={}. File an issue if litellm error - https://github.com/BerriAI/litellm/issues".format(
@ -3139,6 +3146,9 @@ class ModelResponseIterator:
"provider_specific_fields", {}
)["traffic_type"] = traffic_type
if "serviceTier" in processed_chunk:
setattr(model_response, "service_tier", processed_chunk["serviceTier"])
setattr(model_response, "usage", usage) # type: ignore
model_response._hidden_params["is_finished"] = False

View file

@ -325,6 +325,7 @@ class RequestBody(TypedDict, total=False):
generationConfig: GenerationConfig
cachedContent: str
labels: Dict[str, str]
serviceTier: str
class CachedContentRequestBody(TypedDict, total=False):

View file

@ -127,6 +127,24 @@ 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 = {}
result = _transform_request_body(
messages=messages,
model="gemini-2.5-pro",
optional_params=optional_params,
custom_llm_provider="vertex_ai",
litellm_params=litellm_params,
cached_content=None,
)
assert "serviceTier" in result
assert result["serviceTier"] == "flex"
def test_extra_body_cache_not_forwarded_to_vertex_ai():
"""

View file

@ -3504,6 +3504,64 @@ def test_vertex_ai_traffic_type_preserved_in_hidden_params_non_streaming():
assert result._hidden_params["provider_specific_fields"]["traffic_type"] == "PROVISIONED_THROUGHPUT"
def test_vertex_ai_service_tier_streaming():
"""Test serviceTier is preserved in model_response for streaming."""
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
ModelResponseIterator,
)
chunk = {
"candidates": [{"content": {"parts": [{"text": "Hello"}]}}],
"serviceTier": "FLEX",
}
iterator = ModelResponseIterator(
streaming_response=[], sync_stream=True, logging_obj=MagicMock()
)
result = iterator.chunk_parser(chunk)
assert result.service_tier == "FLEX"
def test_vertex_ai_service_tier_non_streaming():
"""Test serviceTier is preserved in model_response for non-streaming."""
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
completion_response = {
"candidates": [
{
"content": {"parts": [{"text": "Hello"}], "role": "model"},
"finishReason": "STOP",
}
],
"usageMetadata": {
"promptTokenCount": 50,
"candidatesTokenCount": 100,
"totalTokenCount": 150,
},
"serviceTier": "FLEX",
}
raw_response = MagicMock()
raw_response.json.return_value = completion_response
result = VertexGeminiConfig().transform_response(
model="gemini-pro",
raw_response=raw_response,
model_response=ModelResponse(),
logging_obj=MagicMock(),
request_data={},
messages=[],
optional_params={},
litellm_params={},
encoding=None,
)
assert result.service_tier == "FLEX"
def test_vertex_ai_traffic_type_surfaced_in_responses_api():
"""Test trafficType is surfaced as provider_specific_fields in ResponsesAPIResponse."""
from litellm.responses.litellm_completion_transformation.transformation import (
@ -3609,6 +3667,26 @@ def test_vertex_ai_web_search_options_in_map_openai_params():
assert "web_search_options" not in optional_params, "web_search_options should be removed after transformation"
def test_vertex_ai_service_tier_in_map_openai_params():
"""Test that service_tier is correctly mapped to optional_params."""
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
v = VertexGeminiConfig()
optional_params = {}
non_default_params = {"service_tier": "FLEX"}
result = v.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model="gemini-pro",
drop_params=True,
)
assert result["service_tier"] == "FLEX"
def test_vertex_ai_usage_metadata_with_video_tokens_in_prompt():
"""Test promptTokensDetails with VIDEO modality for video inputs.