mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
38d1691e20
commit
a4c96d5224
4 changed files with 47 additions and 19 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue