fix: fix checking optional params from logging object for function call

This commit is contained in:
Krrish Dholakia 2025-05-22 14:46:42 -07:00
parent 754a94db97
commit 889f0093e0
3 changed files with 23 additions and 13 deletions

View file

@ -573,3 +573,12 @@ def get_tool_call_names(tools: List[ChatCompletionToolParam]) -> List[str]:
if tool_call_name:
tool_call_names.append(tool_call_name)
return tool_call_names
def is_function_call(optional_params: dict) -> bool:
"""
Checks if the optional params contain the function call
"""
if "functions" in optional_params and optional_params.get("functions"):
return True
return False

View file

@ -149,14 +149,14 @@ class CustomStreamWrapper:
)
def check_is_function_call(self, logging_obj) -> bool:
from litellm.litellm_core_utils.prompt_templates.common_utils import (
is_function_call,
)
if hasattr(logging_obj, "optional_params") and isinstance(
logging_obj.optional_params, dict
):
if (
"litellm_param_is_function_call" in logging_obj.optional_params
and logging_obj.optional_params["litellm_param_is_function_call"]
is True
):
if is_function_call(logging_obj.optional_params):
return True
return False

View file

@ -455,9 +455,6 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
and value
):
optional_params["tools"] = self._map_function(value=value)
optional_params["litellm_param_is_function_call"] = (
True if param == "functions" else False
)
elif param == "tool_choice" and (
isinstance(value, str) or isinstance(value, dict)
):
@ -880,8 +877,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
else:
return "stop"
def _process_candidates(self, _candidates, model_response, litellm_params):
def _process_candidates(
self, _candidates, model_response, standard_optional_params: dict
):
"""Helper method to process candidates and extract metadata"""
from litellm.litellm_core_utils.prompt_templates.common_utils import (
is_function_call,
)
grounding_metadata: List[dict] = []
safety_ratings: List = []
citation_metadata: List = []
@ -918,9 +921,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
functions, tools = self._transform_parts(
parts=candidate["content"]["parts"],
index=candidate.get("index", idx),
is_function_call=litellm_params.get(
"litellm_param_is_function_call"
),
is_function_call=is_function_call(standard_optional_params),
)
if "logprobsResult" in candidate:
@ -1019,7 +1020,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
safety_ratings,
citation_metadata,
) = self._process_candidates(
_candidates, model_response, litellm_params
_candidates, model_response, logging_obj.optional_params
)
usage = self._calculate_usage(completion_response=completion_response)