diff --git a/litellm/utils.py b/litellm/utils.py index 9b8d11cfd5d..8217a1860ad 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -523,9 +523,9 @@ def function_setup( # noqa: PLR0915 function_id: Optional[str] = kwargs["id"] if "id" in kwargs else None ## DYNAMIC CALLBACKS ## - dynamic_callbacks: Optional[List[Union[str, Callable, CustomLogger]]] = ( - kwargs.pop("callbacks", None) - ) + dynamic_callbacks: Optional[ + List[Union[str, Callable, CustomLogger]] + ] = kwargs.pop("callbacks", None) all_callbacks = get_dynamic_callbacks(dynamic_callbacks=dynamic_callbacks) if len(all_callbacks) > 0: @@ -1209,9 +1209,9 @@ def client(original_function): # noqa: PLR0915 exception=e, retry_policy=kwargs.get("retry_policy"), ) - kwargs["retry_policy"] = ( - reset_retry_policy() - ) # prevent infinite loops + kwargs[ + "retry_policy" + ] = reset_retry_policy() # prevent infinite loops litellm.num_retries = ( None # set retries to None to prevent infinite loops ) @@ -2754,16 +2754,16 @@ def get_optional_params( # noqa: PLR0915 True # so that main.py adds the function call to the prompt ) if "tools" in non_default_params: - optional_params["functions_unsupported_model"] = ( - non_default_params.pop("tools") - ) + optional_params[ + "functions_unsupported_model" + ] = non_default_params.pop("tools") non_default_params.pop( "tool_choice", None ) # causes ollama requests to hang elif "functions" in non_default_params: - optional_params["functions_unsupported_model"] = ( - non_default_params.pop("functions") - ) + optional_params[ + "functions_unsupported_model" + ] = non_default_params.pop("functions") elif ( litellm.add_function_to_prompt ): # if user opts to add it to prompt instead @@ -2786,10 +2786,10 @@ def get_optional_params( # noqa: PLR0915 if "response_format" in non_default_params: if provider_config is not None: - non_default_params["response_format"] = ( - provider_config.get_json_schema_from_pydantic_object( - response_format=non_default_params["response_format"] - ) + non_default_params[ + "response_format" + ] = provider_config.get_json_schema_from_pydantic_object( + response_format=non_default_params["response_format"] ) else: non_default_params["response_format"] = type_to_response_format_param( @@ -3805,9 +3805,9 @@ def _count_characters(text: str) -> int: def get_response_string(response_obj: Union[ModelResponse, ModelResponseStream]) -> str: - _choices: Union[List[Union[Choices, StreamingChoices]], List[StreamingChoices]] = ( - response_obj.choices - ) + _choices: Union[ + List[Union[Choices, StreamingChoices]], List[StreamingChoices] + ] = response_obj.choices response_str = "" for choice in _choices: @@ -4127,6 +4127,18 @@ def get_provider_info( return model_info +def _is_potential_model_name_in_model_cost( + potential_model_names: PotentialModelNamesAndCustomLLMProvider, +) -> bool: + """ + Check if the potential model name is in the model cost. + """ + return any( + potential_model_name in litellm.model_cost + for potential_model_name in potential_model_names.values() + ) + + def _get_model_info_helper( # noqa: PLR0915 model: str, custom_llm_provider: Optional[str] = None ) -> ModelInfoBase: @@ -4182,7 +4194,9 @@ def _get_model_info_helper( # noqa: PLR0915 supports_prompt_caching=None, supports_pdf_input=None, ) - elif custom_llm_provider == "ollama" or custom_llm_provider == "ollama_chat": + elif ( + custom_llm_provider == "ollama" or custom_llm_provider == "ollama_chat" + ) and not _is_potential_model_name_in_model_cost(potential_model_names): return litellm.OllamaConfig().get_model_info(model) else: """