From 55a17730fbd06d499fd900c311144edfe69bab73 Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Sat, 19 Apr 2025 18:03:05 -0700 Subject: [PATCH] fix(transformation.py): pass back in gemini thinking content to api (#10173) Ensures thinking content always returned --- .../llms/vertex_ai/gemini/transformation.py | 6 +++ litellm/types/llms/openai.py | 37 ++++++++++--------- litellm/types/llms/vertex_ai.py | 1 + tests/llm_translation/test_gemini.py | 28 ++++++++++++++ 4 files changed, 54 insertions(+), 18 deletions(-) diff --git a/litellm/llms/vertex_ai/gemini/transformation.py b/litellm/llms/vertex_ai/gemini/transformation.py index f4a05d9104a..e50954b8f96 100644 --- a/litellm/llms/vertex_ai/gemini/transformation.py +++ b/litellm/llms/vertex_ai/gemini/transformation.py @@ -216,6 +216,11 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915 msg_dict = messages[msg_i] # type: ignore assistant_msg = ChatCompletionAssistantMessage(**msg_dict) # type: ignore _message_content = assistant_msg.get("content", None) + reasoning_content = assistant_msg.get("reasoning_content", None) + if reasoning_content is not None: + assistant_content.append( + PartType(thought=True, text=reasoning_content) + ) if _message_content is not None and isinstance(_message_content, list): _parts = [] for element in _message_content: @@ -223,6 +228,7 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915 if element["type"] == "text": _part = PartType(text=element["text"]) _parts.append(_part) + assistant_content.extend(_parts) elif ( _message_content is not None diff --git a/litellm/types/llms/openai.py b/litellm/types/llms/openai.py index 42c86b8b19e..24aebf12afd 100644 --- a/litellm/types/llms/openai.py +++ b/litellm/types/llms/openai.py @@ -651,6 +651,7 @@ class OpenAIChatCompletionAssistantMessage(TypedDict, total=False): name: Optional[str] tool_calls: Optional[List[ChatCompletionAssistantToolCall]] function_call: Optional[ChatCompletionToolCallFunctionChunk] + reasoning_content: Optional[str] class ChatCompletionAssistantMessage(OpenAIChatCompletionAssistantMessage, total=False): @@ -823,12 +824,12 @@ class OpenAIChatCompletionChunk(ChatCompletionChunk): class Hyperparameters(BaseModel): batch_size: Optional[Union[str, int]] = None # "Number of examples in each batch." - learning_rate_multiplier: Optional[Union[str, float]] = ( - None # Scaling factor for the learning rate - ) - n_epochs: Optional[Union[str, int]] = ( - None # "The number of epochs to train the model for" - ) + learning_rate_multiplier: Optional[ + Union[str, float] + ] = None # Scaling factor for the learning rate + n_epochs: Optional[ + Union[str, int] + ] = None # "The number of epochs to train the model for" class FineTuningJobCreate(BaseModel): @@ -855,18 +856,18 @@ class FineTuningJobCreate(BaseModel): model: str # "The name of the model to fine-tune." training_file: str # "The ID of an uploaded file that contains training data." - hyperparameters: Optional[Hyperparameters] = ( - None # "The hyperparameters used for the fine-tuning job." - ) - suffix: Optional[str] = ( - None # "A string of up to 18 characters that will be added to your fine-tuned model name." - ) - validation_file: Optional[str] = ( - None # "The ID of an uploaded file that contains validation data." - ) - integrations: Optional[List[str]] = ( - None # "A list of integrations to enable for your fine-tuning job." - ) + hyperparameters: Optional[ + Hyperparameters + ] = None # "The hyperparameters used for the fine-tuning job." + suffix: Optional[ + str + ] = None # "A string of up to 18 characters that will be added to your fine-tuned model name." + validation_file: Optional[ + str + ] = None # "The ID of an uploaded file that contains validation data." + integrations: Optional[ + List[str] + ] = None # "A list of integrations to enable for your fine-tuning job." seed: Optional[int] = None # "The seed controls the reproducibility of the job." diff --git a/litellm/types/llms/vertex_ai.py b/litellm/types/llms/vertex_ai.py index 6f3707c657c..6ac7ae62751 100644 --- a/litellm/types/llms/vertex_ai.py +++ b/litellm/types/llms/vertex_ai.py @@ -39,6 +39,7 @@ class PartType(TypedDict, total=False): file_data: FileDataType function_call: FunctionCall function_response: FunctionResponse + thought: bool class HttpxFunctionCall(TypedDict): diff --git a/tests/llm_translation/test_gemini.py b/tests/llm_translation/test_gemini.py index c8cc0366d4b..35aa22722e6 100644 --- a/tests/llm_translation/test_gemini.py +++ b/tests/llm_translation/test_gemini.py @@ -89,3 +89,31 @@ def test_gemini_image_generation(): +def test_gemini_thinking(): + litellm._turn_on_debug() + from litellm.types.utils import Message, CallTypes + from litellm.utils import return_raw_request + import json + + messages = [ + {"role": "user", "content": "Explain the concept of Occam's Razor and provide a simple, everyday example"} + ] + reasoning_content = "I'm thinking about Occam's Razor." + assistant_message = Message(content='Okay, let\'s break down Occam\'s Razor.', reasoning_content=reasoning_content, role='assistant', tool_calls=None, function_call=None, provider_specific_fields=None) + + messages.append(assistant_message) + + raw_request = return_raw_request( + endpoint=CallTypes.completion, + kwargs={ + "model": "gemini/gemini-2.5-flash-preview-04-17", + "messages": messages, + } + ) + assert reasoning_content in json.dumps(raw_request) + response = completion( + model="gemini/gemini-2.5-flash-preview-04-17", + messages=messages, # make sure call works + ) + print(response.choices[0].message) + assert response.choices[0].message.content is not None \ No newline at end of file