mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
feat: support service_tier in gemini
This commit is contained in:
parent
885fea2c58
commit
25b266fe2c
6 changed files with 112 additions and 0 deletions
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -325,6 +325,7 @@ class RequestBody(TypedDict, total=False):
|
|||
generationConfig: GenerationConfig
|
||||
cachedContent: str
|
||||
labels: Dict[str, str]
|
||||
serviceTier: str
|
||||
|
||||
|
||||
class CachedContentRequestBody(TypedDict, total=False):
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue