diff --git a/litellm/llms/vertex_ai/context_caching/transformation.py b/litellm/llms/vertex_ai/context_caching/transformation.py index 950edbeb478..6ba0a68f1d7 100644 --- a/litellm/llms/vertex_ai/context_caching/transformation.py +++ b/litellm/llms/vertex_ai/context_caching/transformation.py @@ -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 diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index 7945c44d44c..d3b63ec0540 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -243,6 +243,42 @@ def _camel_to_snake(camel_str: str) -> str: return re.sub(r"(? 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 diff --git a/tests/litellm/llms/vertex_ai/gemini/test_transformation.py b/tests/litellm/llms/vertex_ai/gemini/test_transformation.py index 20f48b6f393..94937032ebd 100644 --- a/tests/litellm/llms/vertex_ai/gemini/test_transformation.py +++ b/tests/litellm/llms/vertex_ai/gemini/test_transformation.py @@ -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] \ No newline at end of file + 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"]