fix(vertex_ai): fix multimodal field naming mismatch and preserve tool schema casing(fixes #24399)

This commit is contained in:
prophet_system_team 2026-03-23 15:02:06 +05:30
parent c89496f378
commit 7b8a0c2f1c
3 changed files with 110 additions and 2 deletions

View file

@ -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

View file

@ -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

View file

@ -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"]