mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-17 23:51:30 +00:00
Merge pull request #41338 from BerriAI/litellm_fix_gemini_model_version
fix(gemini): propagate the provider's modelVersion to the response model
This commit is contained in:
commit
54e1b9e3de
2 changed files with 140 additions and 3 deletions
|
|
@ -124,6 +124,12 @@ def _unsupported_reasoning_effort(reasoning_effort: str) -> UnsupportedParamsErr
|
|||
)
|
||||
|
||||
|
||||
def _served_model_name(model_version: object) -> str | None:
|
||||
if not isinstance(model_version, str) or not model_version:
|
||||
return None
|
||||
return model_version.split("@", 1)[0]
|
||||
|
||||
|
||||
class VertexAIBaseConfig:
|
||||
def get_mapped_special_auth_params(self) -> dict:
|
||||
"""
|
||||
|
|
@ -1951,6 +1957,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
def _check_prompt_level_content_filter(
|
||||
processed_chunk: GenerateContentResponseBody,
|
||||
response_id: str | None,
|
||||
model: str | None = None,
|
||||
) -> Optional["ModelResponseStream"]:
|
||||
"""
|
||||
Check if prompt is blocked due to content filtering at the prompt level.
|
||||
|
|
@ -1990,7 +1997,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
enhancements=None,
|
||||
)
|
||||
|
||||
model_response: Final = ModelResponseStream(choices=[choice], id=response_id)
|
||||
model_response: Final = ModelResponseStream(choices=[choice], id=response_id, model=model)
|
||||
return model_response
|
||||
|
||||
return None
|
||||
|
|
@ -2434,7 +2441,8 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
|
|||
completion_response = GenerateContentResponseBody(**completion_response)
|
||||
|
||||
## GET MODEL ##
|
||||
model_response.model = model
|
||||
served: Final = _served_model_name(completion_response.get("modelVersion"))
|
||||
model_response.model = served if served is not None else model
|
||||
|
||||
## CHECK IF RESPONSE FLAGGED
|
||||
if "promptFeedback" in completion_response and "blockReason" in completion_response["promptFeedback"]:
|
||||
|
|
@ -3264,12 +3272,18 @@ class ModelResponseIterator:
|
|||
|
||||
processed_chunk: Final = GenerateContentResponseBody(**chunk)
|
||||
response_id: Final = processed_chunk.get("responseId")
|
||||
model_response = ModelResponseStream(choices=[], id=response_id)
|
||||
served: Final = _served_model_name(processed_chunk.get("modelVersion"))
|
||||
model_response = ModelResponseStream(
|
||||
choices=[],
|
||||
id=response_id,
|
||||
model=served,
|
||||
)
|
||||
|
||||
# Check if prompt is blocked due to content filtering
|
||||
blocked_response: Final = VertexGeminiConfig._check_prompt_level_content_filter(
|
||||
processed_chunk=processed_chunk,
|
||||
response_id=response_id,
|
||||
model=served,
|
||||
)
|
||||
if blocked_response is not None:
|
||||
model_response = blocked_response
|
||||
|
|
|
|||
|
|
@ -5836,3 +5836,126 @@ def test_supported_reasoning_efforts_still_map(model):
|
|||
drop_params=False,
|
||||
)
|
||||
assert "thinkingConfig" in result
|
||||
|
||||
|
||||
def _generate_content_body() -> dict:
|
||||
return {
|
||||
"candidates": [
|
||||
{
|
||||
"content": {"role": "model", "parts": [{"text": "hi"}]},
|
||||
"finishReason": "STOP",
|
||||
"index": 0,
|
||||
}
|
||||
],
|
||||
"usageMetadata": {
|
||||
"promptTokenCount": 5,
|
||||
"candidatesTokenCount": 7,
|
||||
"totalTokenCount": 12,
|
||||
},
|
||||
}
|
||||
|
||||
|
||||
def test_generate_content_transform_uses_reported_model_version():
|
||||
"""The served modelVersion must win over the requested name so downstream
|
||||
pricing sees what actually ran."""
|
||||
import httpx
|
||||
|
||||
body = {**_generate_content_body(), "modelVersion": "gemini-x-served"}
|
||||
response: Final = VertexGeminiConfig()._transform_google_generate_content_to_openai_model_response(
|
||||
completion_response=body,
|
||||
model_response=ModelResponse(),
|
||||
model="gemini-x",
|
||||
logging_obj=MagicMock(),
|
||||
raw_response=httpx.Response(200, headers={}),
|
||||
)
|
||||
|
||||
assert response.model == "gemini-x-served"
|
||||
|
||||
|
||||
def test_generate_content_transform_falls_back_to_requested_model():
|
||||
import httpx
|
||||
|
||||
response: Final = VertexGeminiConfig()._transform_google_generate_content_to_openai_model_response(
|
||||
completion_response=_generate_content_body(),
|
||||
model_response=ModelResponse(),
|
||||
model="gemini-x",
|
||||
logging_obj=MagicMock(),
|
||||
raw_response=httpx.Response(200, headers={}),
|
||||
)
|
||||
|
||||
assert response.model == "gemini-x"
|
||||
|
||||
|
||||
def test_streaming_chunk_carries_model_version():
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
ModelResponseIterator,
|
||||
)
|
||||
|
||||
chunk = {**_generate_content_body(), "modelVersion": "gemini-x-served"}
|
||||
iterator: Final = ModelResponseIterator(streaming_response=[], sync_stream=True, logging_obj=MagicMock())
|
||||
streaming_chunk: Final = iterator.chunk_parser(chunk)
|
||||
|
||||
assert streaming_chunk.model == "gemini-x-served"
|
||||
|
||||
|
||||
def test_served_model_version_reaches_assembled_stream_through_custom_stream_wrapper():
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
ModelResponseIterator,
|
||||
)
|
||||
|
||||
served_model: Final = "gemini-3.8-flash-001"
|
||||
iterator: Final = ModelResponseIterator(
|
||||
streaming_response=iter(
|
||||
[json.dumps({**_generate_content_body(), "modelVersion": served_model}) for _ in range(3)]
|
||||
),
|
||||
sync_stream=True,
|
||||
logging_obj=MagicMock(),
|
||||
)
|
||||
wrapper: Final = CustomStreamWrapper(
|
||||
completion_stream=iter(iterator),
|
||||
model="gemini/gemini-3.8-flash",
|
||||
custom_llm_provider="gemini",
|
||||
logging_obj=MagicMock(),
|
||||
)
|
||||
|
||||
chunks: Final = list(wrapper)
|
||||
|
||||
assert len(chunks) >= 3
|
||||
for chunk in chunks[:-1]:
|
||||
assert chunk._hidden_params["provider_response_model"] == served_model
|
||||
assembled: Final = litellm.stream_chunk_builder(chunks=list(chunks), messages=[{"role": "user", "content": "hi"}])
|
||||
assert assembled._hidden_params["provider_response_model"] == served_model
|
||||
|
||||
|
||||
def test_generate_content_transform_strips_version_suffix_from_model_version():
|
||||
import httpx
|
||||
|
||||
body: Final = {**_generate_content_body(), "modelVersion": "gemini-3.8-flash-001@default"}
|
||||
response: Final = VertexGeminiConfig()._transform_google_generate_content_to_openai_model_response(
|
||||
completion_response=body,
|
||||
model_response=ModelResponse(),
|
||||
model="gemini-3.8-flash",
|
||||
logging_obj=MagicMock(),
|
||||
raw_response=httpx.Response(200, headers={}),
|
||||
)
|
||||
|
||||
assert response.model == "gemini-3.8-flash-001"
|
||||
|
||||
|
||||
def test_prompt_blocked_chunk_keeps_served_model_version():
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
ModelResponseIterator,
|
||||
)
|
||||
|
||||
chunk: Final = {
|
||||
"promptFeedback": {"blockReason": "SAFETY", "blockReasonMessage": "prompt was blocked"},
|
||||
"modelVersion": "gemini-3.8-flash-001",
|
||||
"responseId": "resp-1",
|
||||
}
|
||||
iterator: Final = ModelResponseIterator(streaming_response=[], sync_stream=True, logging_obj=MagicMock())
|
||||
|
||||
streaming_chunk: Final = iterator.chunk_parser(chunk)
|
||||
|
||||
assert streaming_chunk.model == "gemini-3.8-flash-001"
|
||||
assert streaming_chunk.choices[0].finish_reason == "content_filter"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue