diff --git a/litellm/llms/vertex_ai.py b/litellm/llms/vertex_ai.py index d3bb2c78ab3..84fec734fd0 100644 --- a/litellm/llms/vertex_ai.py +++ b/litellm/llms/vertex_ai.py @@ -867,6 +867,8 @@ async def async_completion( Add support for acompletion calls for gemini-pro """ try: + import proto # type: ignore + if mode == "vision": print_verbose("\nMaking VertexAI Gemini Pro/Vision Call") print_verbose(f"\nProcessing input messages = {messages}") @@ -901,9 +903,21 @@ async def async_completion( ): function_call = response.candidates[0].content.parts[0].function_call args_dict = {} - for k, v in function_call.args.items(): - args_dict[k] = v - args_str = json.dumps(args_dict) + + # Check if it's a RepeatedComposite instance + for key, val in function_call.args.items(): + if isinstance( + val, proto.marshal.collections.repeated.RepeatedComposite + ): + # If so, convert to list + args_dict[key] = [v for v in val] + else: + args_dict[key] = val + + try: + args_str = json.dumps(args_dict) + except Exception as e: + raise VertexAIError(status_code=422, message=str(e)) message = litellm.Message( content=None, tool_calls=[ diff --git a/litellm/tests/test_amazing_vertex_completion.py b/litellm/tests/test_amazing_vertex_completion.py index a56d7fe5a9a..ce9e6286ff2 100644 --- a/litellm/tests/test_amazing_vertex_completion.py +++ b/litellm/tests/test_amazing_vertex_completion.py @@ -590,19 +590,20 @@ def test_gemini_pro_vision_base64(): pytest.fail(f"An exception occurred - {str(e)}") +@pytest.mark.parametrize("sync_mode", [True, False]) @pytest.mark.asyncio -def test_gemini_pro_function_calling(): +async def test_gemini_pro_function_calling(sync_mode): try: load_vertex_ai_credentials() - response = litellm.completion( - model="vertex_ai/gemini-pro", - messages=[ + data = { + "model": "vertex_ai/gemini-pro", + "messages": [ { "role": "user", "content": "Call the submit_cities function with San Francisco and New York", } ], - tools=[ + "tools": [ { "type": "function", "function": { @@ -618,11 +619,13 @@ def test_gemini_pro_function_calling(): }, } ], - ) + } + if sync_mode: + response = litellm.completion(**data) + else: + response = await litellm.acompletion(**data) print(f"response: {response}") - except litellm.APIError as e: - pass except litellm.RateLimitError as e: pass except Exception as e: