diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py index 6ba0a68f1d7..c315643ac72 100644 --- a/litellm/llms/vertex_ai/context_caching/transformation.py +++ b/litellm/llms/vertex_ai/context_caching/transformation.py @@ -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 diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index d3b63ec0540..c0d024da59a 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -243,7 +243,7 @@ def _camel_to_snake(camel_str: str) -> str: return re.sub(r"(? 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 diff --git a/tests/litellm/llms/vertex_ai/gemini/test_transformation.py b/tests/litellm/llms/vertex_ai/gemini/test_transformation.py index 94937032ebd..818947325f1 100644 --- a/tests/litellm/llms/vertex_ai/gemini/test_transformation.py +++ b/tests/litellm/llms/vertex_ai/gemini/test_transformation.py @@ -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"]