refactor(vertex_ai): improve field transformation safety and test quality

- Switch from denylist to allowlist for Vertex AI provider checks.
- Make camelCase transformation context-aware to ensure selective key preservation (args, response, properties) only in correct API structures.
- Refactor transformation unit tests to sync and address PEP 8 style issues.
- Add test cases for 'response' and 'labels' field preservation.
This commit is contained in:
prophet_system_team 2026-03-23 20:14:35 +05:30
parent 7b8a0c2f1c
commit 109fe0df45
3 changed files with 60 additions and 25 deletions

View file

@ -169,7 +169,7 @@ def transform_openai_messages_to_gemini_context_caching(
model=model, custom_llm_provider=custom_llm_provider
)
is_vertex_ai = custom_llm_provider != "gemini"
is_vertex_ai = custom_llm_provider in ["vertex_ai", "vertex_ai_beta"]
transformed_system_messages, new_messages = _transform_system_message(
supports_system_message=supports_system_message, messages=messages

View file

@ -243,7 +243,7 @@ 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:
def _transform_part_to_httpx_format(part: dict, parent_key: Optional[str] = None) -> dict:
"""
Recursively transform a Gemini part (PartType) to HttpxPartType (camelCase)
Required for Vertex AI REST API.
@ -251,14 +251,35 @@ def _transform_part_to_httpx_format(part: dict) -> dict:
"""
new_part = {}
for k, v in part.items():
camel_k = _snake_to_camel(k)
# 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)
should_preserve_keys = False
if isinstance(v, dict):
if k == "args" and parent_key in ["functionCall", "function_call"]:
should_preserve_keys = True
elif k == "response" and parent_key in [
"functionResponse",
"function_response",
"toolResponse",
"tool_response",
]:
should_preserve_keys = True
elif k == "properties":
# 'properties' keys in a JSON Schema are user-defined
should_preserve_keys = True
elif k == "labels":
# 'labels' is a top-level map for Vertex AI billing/tracking, keys are user-defined
should_preserve_keys = True
if should_preserve_keys:
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
pk: _transform_part_to_httpx_format(pv, parent_key="properties")
if isinstance(pv, dict)
else pv
for pk, pv in v.items()
}
else:
@ -266,12 +287,13 @@ def _transform_part_to_httpx_format(part: dict) -> dict:
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)
new_part[camel_k] = _transform_part_to_httpx_format(v, parent_key=camel_k)
elif isinstance(v, list):
new_part[camel_k] = [
_transform_part_to_httpx_format(i) if isinstance(i, dict) else i
_transform_part_to_httpx_format(i, parent_key=camel_k)
if isinstance(i, dict)
else i
for i in v
]
else:
@ -806,7 +828,7 @@ def _transform_request_body( # noqa: PLR0915
except Exception as e:
raise e
if custom_llm_provider != LlmProviders.GEMINI:
if custom_llm_provider in ["vertex_ai", "vertex_ai_beta"]:
return _transform_part_to_httpx_format(data) # type: ignore
return data

View file

@ -288,8 +288,8 @@ def test_map_function_enterprise_web_search_snake_case():
assert len(result) == 1
assert "enterpriseWebSearch" in result[0]
@pytest.mark.asyncio
async def test_vertex_transformation_field_casing():
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.
@ -297,12 +297,7 @@ async def test_vertex_transformation_field_casing():
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"
}
}
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"
@ -311,13 +306,10 @@ async def test_vertex_transformation_field_casing():
part = {
"function_call": {
"name": "my_func",
"args": {
"security_risk": "high",
"other_param": 123
}
"args": {"security_risk": "high", "other_param": 123},
}
}
transformed = _transform_part_to_httpx_format(part)
transformed = _transform_part_to_httpx_format(part, parent_key="functionCall")
assert "functionCall" in transformed
assert "args" in transformed["functionCall"]
# Keys in args should NOT be camelCased
@ -334,10 +326,13 @@ async def test_vertex_transformation_field_casing():
"parameters": {
"type": "object",
"properties": {
"security_risk": {"type": "string", "mime_type": "text/plain"}
"security_risk": {
"type": "string",
"mime_type": "text/plain",
}
},
"required": ["security_risk"]
}
"required": ["security_risk"],
},
}
]
}
@ -351,3 +346,21 @@ async def test_vertex_transformation_field_casing():
assert "mimeType" in schema["properties"]["security_risk"]
# 'required' list strings should stay snake_case
assert "security_risk" in schema["required"]
# 4. Test response field preservation inside functionResponse
part = {
"function_response": {
"name": "my_func",
"response": {"output_field": "value"},
}
}
transformed = _transform_part_to_httpx_format(part)
assert "functionResponse" in transformed
assert "response" in transformed["functionResponse"]
assert "output_field" in transformed["functionResponse"]["response"]
# 5. Test response field preservation inside labels
part = {"labels": {"response": "user_value"}}
transformed = _transform_part_to_httpx_format(part)
assert "labels" in transformed
assert "response" in transformed["labels"]