Fix return finish_reason = "tool_calls" for gemini tool calling (#10485)

* fix(vertex_and_google_ai_studio.py): fix finish reason to be 'tool_calls' when tool call returned

Vertex returns 'Stop', openai format is 'tool calls'

* test(base_llm_unit_tests.py): bump test to assert tool calls in finish reason
This commit is contained in:
Krish Dholakia 2025-05-01 22:02:56 -07:00 • committed by GitHub
parent 38d1691e20
commit a4c96d5224
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 47 additions and 19 deletions

View file

@ -44,6 +44,7 @@ from litellm.types.llms.openai import (
ChatCompletionToolCallFunctionChunk,
ChatCompletionToolParamFunctionChunk,
ChatCompletionUsageBlock,
OpenAIChatCompletionFinishReason,
)
from litellm.types.llms.vertex_ai import (
VERTEX_CREDENTIALS_TYPES,
@ -810,6 +811,22 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
return usage
def _check_finish_reason(
self,
chat_completion_message: ChatCompletionResponseMessage,
finish_reason: Optional[str],
) -> OpenAIChatCompletionFinishReason:
if chat_completion_message.get("function_call"):
return "function_call"
elif chat_completion_message.get("tool_calls"):
return "tool_calls"
elif finish_reason and (
finish_reason == "SAFETY" or finish_reason == "RECITATION"
): # vertex ai
return "content_filter"
else:
return "stop"
def _process_candidates(self, _candidates, model_response, litellm_params):
"""Helper method to process candidates and extract metadata"""
grounding_metadata: List[dict] = []
@ -865,7 +882,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
chat_completion_message["function_call"] = functions
choice = litellm.Choices(
finish_reason=candidate.get("finishReason", "stop"),
finish_reason=self._check_finish_reason(
chat_completion_message, candidate.get("finishReason")
),
index=candidate.get("index", idx),
message=chat_completion_message, # type: ignore
logprobs=chat_completion_logprobs,

View file

@ -57,6 +57,9 @@ model_list:
litellm_params:
model: text-embedding-ada-002
api_key: os.environ/OPENAI_API_KEY
- model_name: gemini/gemini-2.0-flash
litellm_params:
model: gemini/gemini-2.0-flash
litellm_settings:
num_retries: 0

View file

@ -824,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):
@ -856,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."
@ -1293,3 +1293,8 @@ class OpenAIModerationResponse(BaseLiteLLMOpenAIResponseObject):
# Define private attributes using PrivateAttr
_hidden_params: dict = PrivateAttr(default_factory=dict)
OpenAIChatCompletionFinishReason = Literal[
"stop", "content_filter", "function_call", "tool_calls", "length"
]

View file

@ -917,6 +917,7 @@ class BaseLLMChatTest(ABC):
assert isinstance(
response.choices[0].message.tool_calls[0].function.arguments, str
)
assert response.choices[0].finish_reason == "tool_calls"
messages.append(
response.choices[0].message.model_dump()
) # Add assistant tool invokes