diff --git a/litellm/llms/vertex_ai.py b/litellm/llms/vertex_ai.py index b52e8689f10..4452ae43a0d 100644 --- a/litellm/llms/vertex_ai.py +++ b/litellm/llms/vertex_ai.py @@ -3,7 +3,7 @@ import json from enum import Enum import requests # type: ignore import time -from typing import Callable, Optional, Union, List, Literal +from typing import Callable, Optional, Union, List, Literal, Tuple from litellm.utils import ModelResponse, Usage, CustomStreamWrapper, map_finish_reason import litellm, uuid import httpx, inspect # type: ignore @@ -336,14 +336,39 @@ def _process_gemini_image(image_url: str) -> PartType: raise e -def _gemini_convert_messages_with_history(messages: list) -> List[ContentType]: +def _extract_system_prompt_from_messages(messages: list) -> Optional[ContentType]: + # Separate system prompt from rest of message + system_prompt_indices = [] + _parts: List[PartType] = [] + for idx, message in enumerate(messages): + if message["role"] == "system": + _part = PartType(text=message["content"]) + _parts.append(_part) + system_prompt_indices.append(idx) + if len(system_prompt_indices) > 0: + for idx in reversed(system_prompt_indices): + messages.pop(idx) + if len(_parts) > 0: + return ContentType(parts=_parts) + return None + + +def _gemini_convert_messages_with_history( + messages: list, +) -> Tuple[Optional[ContentType], List[ContentType]]: """ Converts given messages from OpenAI format to Gemini format - Parts must be iterable - Roles must alternate b/w 'user' and 'model' (same as anthropic -> merge consecutive roles) - Please ensure that function response turn comes immediately after a function call turn + + Returns: + - Tuple[Optional[system_instructions], messages] """ + + system_instructions = _extract_system_prompt_from_messages(messages=messages) + user_message_types = {"user", "system"} contents: List[ContentType] = [] @@ -404,7 +429,7 @@ def _gemini_convert_messages_with_history(messages: list) -> List[ContentType]: ) ) - return contents + return system_instructions, contents def _gemini_vision_convert_messages(messages: list): @@ -699,7 +724,9 @@ def completion( print_verbose("\nMaking VertexAI Gemini Pro / Pro Vision Call") print_verbose(f"\nProcessing input messages = {messages}") tools = optional_params.pop("tools", None) - content = _gemini_convert_messages_with_history(messages=messages) + system_instruction, content = _gemini_convert_messages_with_history( + messages=messages + ) stream = optional_params.pop("stream", False) if stream == True: request_str += f"response = llm_model.generate_content({content}, generation_config=GenerationConfig(**{optional_params}), safety_settings={safety_settings}, stream={stream})\n" @@ -736,6 +763,7 @@ def completion( ## LLM Call response = llm_model.generate_content( contents=content, + system_instruction=system_instruction, generation_config=optional_params, safety_settings=safety_settings, tools=tools, diff --git a/litellm/tests/test_amazing_vertex_completion.py b/litellm/tests/test_amazing_vertex_completion.py index d7b2dc2d403..4df39fec087 100644 --- a/litellm/tests/test_amazing_vertex_completion.py +++ b/litellm/tests/test_amazing_vertex_completion.py @@ -676,7 +676,7 @@ async def test_gemini_pro_function_calling(sync_mode): # gemini_pro_function_calling() -@pytest.mark.parametrize("sync_mode", [False, True]) +@pytest.mark.parametrize("sync_mode", [True]) @pytest.mark.asyncio async def test_gemini_pro_function_calling_streaming(sync_mode): load_vertex_ai_credentials()