chore: add defaut/standard mapping, and some tests

This commit is contained in:
Mark McDonald 2026-03-30 17:44:15 +08:00 • committed by Sameer Kankute
parent 48af865358
commit e04af74986
No known key found for this signature in database
5 changed files with 139 additions and 6 deletions

View file

@ -306,7 +306,10 @@ LiteLLM propagates OpenAI's `service_tier` parameter to Gemini, and also extract
| `"auto"` | `"priority"` | LiteLLM maps OpenAI's `"auto"` to Gemini's `"priority"` tier, as `priority` will fall back on Gemini. |
| `"flex"` | `"flex"` | Direct mapping. |
| `"priority"` | `"priority"` | Direct mapping. |
| Any other value | Passed as-is | Unrecognized values are sent blindly so that the provider can reject the request directly if invalid. |
| `"default"` | `"standard"` | LiteLLM maps `"default"` to `"standard"`. |
| Any other value | Passed as-is (lowercased) | Values are case-insensitive and normalized to lowercase. |
On the response, LiteLLM maps `"standard"` back to `"default"` for the Gemini API.
## Text-to-Speech (TTS) Audio Output

View file

@ -765,7 +765,13 @@ def _transform_request_body( # noqa: PLR0915
data["cachedContent"] = cached_content
if service_tier := optional_params.pop("service_tier", None):
data["serviceTier"] = service_tier
if isinstance(service_tier, str):
if service_tier.lower() == "default":
data["serviceTier"] = "standard"
else:
data["serviceTier"] = service_tier.lower()
else:
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:

View file

@ -2432,7 +2432,13 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
## ADD SERVICE TIER ##
if getattr(raw_response, "headers", None):
if service_tier := raw_response.headers.get("x-gemini-service-tier"):
setattr(model_response, "service_tier", service_tier)
if isinstance(service_tier, str):
if service_tier.lower() == "standard":
setattr(model_response, "service_tier", "default")
else:
setattr(model_response, "service_tier", service_tier.lower())
else:
setattr(model_response, "service_tier", service_tier)
except Exception as e:
raise VertexAIError(
@ -3167,7 +3173,13 @@ class ModelResponseIterator:
service_tier = self.response_headers.get("x-gemini-service-tier")
if service_tier:
setattr(model_response, "service_tier", service_tier)
if isinstance(service_tier, str):
if service_tier.lower() == "standard":
setattr(model_response, "service_tier", "default")
else:
setattr(model_response, "service_tier", service_tier.lower())
else:
setattr(model_response, "service_tier", service_tier)
setattr(model_response, "usage", usage) # type: ignore

View file

@ -0,0 +1,112 @@
import pytest
from unittest.mock import MagicMock
from litellm.llms.vertex_ai.gemini.transformation import _transform_request_body
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexGeminiConfig, ModelResponseIterator
from litellm.types.utils import ModelResponse
def test_gemini_service_tier_request_mapping():
"""Test that default service_tier is mapped to standard for Gemini API, case-insensitively."""
messages = [{"role": "user", "content": "test"}]
litellm_params = {}
# Test default -> standard
optional_params = {"service_tier": "default"}
result = _transform_request_body(
messages=messages,
model="gemini-2.5-pro",
optional_params=optional_params,
custom_llm_provider="gemini",
litellm_params=litellm_params,
cached_content=None,
)
assert result["serviceTier"] == "standard"
# Test DEFAULT -> standard
optional_params = {"service_tier": "DEFAULT"}
result = _transform_request_body(
messages=messages,
model="gemini-2.5-pro",
optional_params=optional_params,
custom_llm_provider="gemini",
litellm_params=litellm_params,
cached_content=None,
)
assert result["serviceTier"] == "standard"
# Test flex -> flex
optional_params = {"service_tier": "FLEX"}
result = _transform_request_body(
messages=messages,
model="gemini-2.5-pro",
optional_params=optional_params,
custom_llm_provider="gemini",
litellm_params=litellm_params,
cached_content=None,
)
assert result["serviceTier"] == "flex"
def test_gemini_service_tier_response_mapping():
"""Test that standard service_tier is mapped back to default for Gemini API, case-insensitively."""
config = VertexGeminiConfig()
raw_response = MagicMock()
raw_response.headers = {"x-gemini-service-tier": "STANDARD"}
logging_obj = MagicMock()
logging_obj.custom_llm_provider = "gemini"
completion_response = {
"candidates": [{"content": {"parts": [{"text": "hi"}], "role": "model"}, "finishReason": "STOP"}],
"usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}
}
result = config._transform_google_generate_content_to_openai_model_response(
completion_response=completion_response,
model_response=ModelResponse(),
model="gemini-pro",
logging_obj=logging_obj,
raw_response=raw_response
)
assert result.service_tier == "default"
# Test with lowercase standard
raw_response.headers = {"x-gemini-service-tier": "standard"}
result = config._transform_google_generate_content_to_openai_model_response(
completion_response=completion_response,
model_response=ModelResponse(),
model="gemini-pro",
logging_obj=logging_obj,
raw_response=raw_response
)
assert result.service_tier == "default"
# Test with flex -> flex
raw_response.headers = {"x-gemini-service-tier": "FLEX"}
result = config._transform_google_generate_content_to_openai_model_response(
completion_response=completion_response,
model_response=ModelResponse(),
model="gemini-pro",
logging_obj=logging_obj,
raw_response=raw_response
)
assert result.service_tier == "flex"
def test_gemini_service_tier_streaming_response_mapping():
"""Test streaming response mapping."""
logging_obj = MagicMock()
logging_obj.custom_llm_provider = "gemini"
iterator = ModelResponseIterator(
streaming_response=[],
sync_stream=True,
logging_obj=logging_obj,
response_headers={"x-gemini-service-tier": "STANDARD"}
)
chunk = {
"candidates": [{"content": {"parts": [{"text": "hi"}]}}],
"usageMetadata": {"promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2}
}
result = iterator.chunk_parser(chunk)
assert result.service_tier == "default"

View file

@ -3522,7 +3522,7 @@ def test_vertex_ai_service_tier_streaming():
)
result = iterator.chunk_parser(chunk)
assert result.service_tier == "FLEX"
assert result.service_tier == "flex"
def test_vertex_ai_service_tier_non_streaming():
@ -3561,7 +3561,7 @@ def test_vertex_ai_service_tier_non_streaming():
encoding=None,
)
assert result.service_tier == "FLEX"
assert result.service_tier == "flex"
def test_vertex_ai_traffic_type_surfaced_in_responses_api():