mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
chore: add defaut/standard mapping, and some tests
This commit is contained in:
parent
48af865358
commit
e04af74986
5 changed files with 139 additions and 6 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
112
tests/test_gemini_service_tier.py
Normal file
112
tests/test_gemini_service_tier.py
Normal 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"
|
||||
|
|
@ -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():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue