mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
7b8a0c2f1c
commit
109fe0df45
3 changed files with 60 additions and 25 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue