mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(vertex_ai): fix multimodal field naming mismatch and preserve tool schema casing(fixes #24399)
This commit is contained in:
parent
c89496f378
commit
7b8a0c2f1c
3 changed files with 110 additions and 2 deletions
|
|
@ -169,6 +169,8 @@ def transform_openai_messages_to_gemini_context_caching(
|
|||
model=model, custom_llm_provider=custom_llm_provider
|
||||
)
|
||||
|
||||
is_vertex_ai = custom_llm_provider != "gemini"
|
||||
|
||||
transformed_system_messages, new_messages = _transform_system_message(
|
||||
supports_system_message=supports_system_message, messages=messages
|
||||
)
|
||||
|
|
@ -179,7 +181,7 @@ def transform_openai_messages_to_gemini_context_caching(
|
|||
|
||||
model_name = "models/{}".format(model)
|
||||
|
||||
if custom_llm_provider == "vertex_ai" or custom_llm_provider == "vertex_ai_beta":
|
||||
if is_vertex_ai:
|
||||
model_name = f"projects/{vertex_project}/locations/{vertex_location}/publishers/google/{model_name}"
|
||||
|
||||
data = CachedContentRequestBody(
|
||||
|
|
@ -195,4 +197,9 @@ def transform_openai_messages_to_gemini_context_caching(
|
|||
if transformed_system_messages is not None:
|
||||
data["system_instruction"] = transformed_system_messages
|
||||
|
||||
if is_vertex_ai:
|
||||
from ..gemini.transformation import _transform_part_to_httpx_format
|
||||
|
||||
return _transform_part_to_httpx_format(data) # type: ignore
|
||||
|
||||
return data
|
||||
|
|
|
|||
|
|
@ -243,6 +243,42 @@ def _camel_to_snake(camel_str: str) -> str:
|
|||
return re.sub(r"(?<!^)(?=[A-Z])", "_", camel_str).lower()
|
||||
|
||||
|
||||
def _transform_part_to_httpx_format(part: dict) -> dict:
|
||||
"""
|
||||
Recursively transform a Gemini part (PartType) to HttpxPartType (camelCase)
|
||||
Required for Vertex AI REST API.
|
||||
Google AI Studio REST API uses snake_case, so this is only for Vertex AI.
|
||||
"""
|
||||
new_part = {}
|
||||
for k, v in part.items():
|
||||
# Handle exceptions for keys that should not be camelCased or have special mapping
|
||||
# These are user-defined keys that should be preserved.
|
||||
if k in ["args", "response", "properties", "labels"] and isinstance(v, dict):
|
||||
camel_k = _snake_to_camel(k)
|
||||
if k == "properties":
|
||||
# 'properties' values are Schema objects, so they SHOULD be transformed recursively
|
||||
new_part[camel_k] = {
|
||||
pk: _transform_part_to_httpx_format(pv) if isinstance(pv, dict) else pv
|
||||
for pk, pv in v.items()
|
||||
}
|
||||
else:
|
||||
# 'args', 'response', 'labels' keys are user-defined and should NOT be transformed
|
||||
new_part[camel_k] = v
|
||||
continue
|
||||
|
||||
camel_k = _snake_to_camel(k)
|
||||
if isinstance(v, dict):
|
||||
new_part[camel_k] = _transform_part_to_httpx_format(v)
|
||||
elif isinstance(v, list):
|
||||
new_part[camel_k] = [
|
||||
_transform_part_to_httpx_format(i) if isinstance(i, dict) else i
|
||||
for i in v
|
||||
]
|
||||
else:
|
||||
new_part[camel_k] = v
|
||||
return new_part
|
||||
|
||||
|
||||
def _get_equivalent_key(key: str, available_keys: set) -> Optional[str]:
|
||||
"""
|
||||
Get the equivalent key from available keys, checking both camelCase and snake_case variants
|
||||
|
|
@ -770,6 +806,8 @@ def _transform_request_body( # noqa: PLR0915
|
|||
except Exception as e:
|
||||
raise e
|
||||
|
||||
if custom_llm_provider != LlmProviders.GEMINI:
|
||||
return _transform_part_to_httpx_format(data) # type: ignore
|
||||
return data
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -287,4 +287,67 @@ def test_map_function_enterprise_web_search_snake_case():
|
|||
result = config._map_function(tools, optional_params)
|
||||
|
||||
assert len(result) == 1
|
||||
assert "enterpriseWebSearch" in result[0]
|
||||
assert "enterpriseWebSearch" in result[0]
|
||||
@pytest.mark.asyncio
|
||||
async def test_vertex_transformation_field_casing():
|
||||
"""
|
||||
Tests that structural fields are camelCased for Vertex AI,
|
||||
while user-defined logic (args, properties, labels) preserves snake_case.
|
||||
"""
|
||||
from litellm.llms.vertex_ai.gemini.transformation import _transform_part_to_httpx_format
|
||||
|
||||
# 1. Test Multimodal field naming
|
||||
part = {
|
||||
"inline_data": {
|
||||
"mime_type": "image/png",
|
||||
"data": "base64data"
|
||||
}
|
||||
}
|
||||
transformed = _transform_part_to_httpx_format(part)
|
||||
assert "inlineData" in transformed
|
||||
assert transformed["inlineData"]["mimeType"] == "image/png"
|
||||
|
||||
# 2. Test functionCall with args preservation
|
||||
part = {
|
||||
"function_call": {
|
||||
"name": "my_func",
|
||||
"args": {
|
||||
"security_risk": "high",
|
||||
"other_param": 123
|
||||
}
|
||||
}
|
||||
}
|
||||
transformed = _transform_part_to_httpx_format(part)
|
||||
assert "functionCall" in transformed
|
||||
assert "args" in transformed["functionCall"]
|
||||
# Keys in args should NOT be camelCased
|
||||
assert "security_risk" in transformed["functionCall"]["args"]
|
||||
assert "other_param" in transformed["functionCall"]["args"]
|
||||
|
||||
# 3. Test tools with functionDeclarations and parameters schema preservation
|
||||
part = {
|
||||
"tools": [
|
||||
{
|
||||
"function_declarations": [
|
||||
{
|
||||
"name": "my_func",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"security_risk": {"type": "string", "mime_type": "text/plain"}
|
||||
},
|
||||
"required": ["security_risk"]
|
||||
}
|
||||
}
|
||||
]
|
||||
}
|
||||
]
|
||||
}
|
||||
transformed = _transform_part_to_httpx_format(part)
|
||||
schema = transformed["tools"][0]["functionDeclarations"][0]["parameters"]
|
||||
# Property name should stay snake_case
|
||||
assert "security_risk" in schema["properties"]
|
||||
# Internal schema keyword 'mime_type' should become 'mimeType'
|
||||
assert "mimeType" in schema["properties"]["security_risk"]
|
||||
# 'required' list strings should stay snake_case
|
||||
assert "security_risk" in schema["required"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue