diff --git a/litellm/llms/base_llm/base_utils.py b/litellm/llms/base_llm/base_utils.py index 2afac26e7b6..d2d3d5c0a96 100644 --- a/litellm/llms/base_llm/base_utils.py +++ b/litellm/llms/base_llm/base_utils.py @@ -5,7 +5,7 @@ Utility functions for base LLM classes. import copy import json from abc import ABC, abstractmethod -from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type, Union +from typing import Any, Dict, List, Optional, Type, Union from openai.lib import _parsing, _pydantic from pydantic import BaseModel @@ -14,9 +14,6 @@ from litellm._logging import verbose_logger from litellm.types.llms.openai import AllMessageValues, ChatCompletionToolCallChunk from litellm.types.utils import Message, ProviderSpecificModelInfo, TokenCountResponse -if TYPE_CHECKING: - from litellm.types.utils import ModelInfoBase - class BaseTokenCounter(ABC): @abstractmethod @@ -44,17 +41,6 @@ class BaseTokenCounter(ABC): class BaseLLMModelInfo(ABC): - def get_model_info( - self, - model: str, - api_base: Optional[str] = None, - ) -> Optional["ModelInfoBase"]: - """ - Provider-specific model metadata when it cannot be represented in the - static model cost map. - """ - return None - def get_provider_info( self, model: str, diff --git a/litellm/llms/ollama/common_utils.py b/litellm/llms/ollama/common_utils.py index c4cfd514114..4787e532fae 100644 --- a/litellm/llms/ollama/common_utils.py +++ b/litellm/llms/ollama/common_utils.py @@ -1,13 +1,10 @@ -from typing import TYPE_CHECKING, List, Optional, Union +from typing import Any, List, Optional, Union import httpx from litellm import verbose_logger from litellm.llms.base_llm.chat.transformation import BaseLLMException -if TYPE_CHECKING: - from litellm.types.utils import ModelInfoBase - class OllamaError(BaseLLMException): def __init__( @@ -180,9 +177,8 @@ class OllamaModelInfo(BaseLLMModelInfo): def get_runtime_model_info( self, model: str, api_base: Optional[str] = None - ) -> "ModelInfoBase": + ) -> dict[str, Any]: from litellm import module_level_client - from litellm.types.utils import ModelInfoBase model = self._strip_ollama_model_prefix(model) api_base = self.get_server_api_base(api_base) @@ -197,35 +193,35 @@ class OllamaModelInfo(BaseLLMModelInfo): ) except Exception: verbose_logger.debug("OllamaError: Could not get model info.") - return ModelInfoBase( - key=model, - litellm_provider="ollama", - mode="chat", - input_cost_per_token=0.0, - output_cost_per_token=0.0, - max_tokens=None, - max_input_tokens=None, - max_output_tokens=None, - ) + return { + "key": model, + "litellm_provider": "ollama", + "mode": "chat", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "max_tokens": None, + "max_input_tokens": None, + "max_output_tokens": None, + } model_info = response.json() max_tokens = self._get_max_tokens(model_info) - return ModelInfoBase( - key=model, - litellm_provider="ollama", - mode="chat", - supports_function_calling=self._supports_function_calling(model_info), - input_cost_per_token=0.0, - output_cost_per_token=0.0, - max_tokens=max_tokens, - max_input_tokens=max_tokens, - max_output_tokens=max_tokens, - ) + return { + "key": model, + "litellm_provider": "ollama", + "mode": "chat", + "supports_function_calling": self._supports_function_calling(model_info), + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + "max_tokens": max_tokens, + "max_input_tokens": max_tokens, + "max_output_tokens": max_tokens, + } def get_model_info( self, model: str, api_base: Optional[str] = None - ) -> Optional["ModelInfoBase"]: + ) -> Optional[dict[str, Any]]: if self._is_static_ollama_model(model): return None return self.get_runtime_model_info(model=model, api_base=api_base) diff --git a/litellm/utils.py b/litellm/utils.py index 7b30a7a05a4..7a60f795e58 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -5709,15 +5709,17 @@ def _get_model_info_helper( # noqa: PLR0915 model=model, provider=LlmProviders(custom_llm_provider) ) if provider_config is not None: - try: - provider_model_info = provider_config.get_model_info( - model=model, - api_base=api_base, - ) - if provider_model_info is not None: - return provider_model_info - except Exception: - verbose_logger.debug("Could not get dynamic model info.") + get_model_info = getattr(provider_config, "get_model_info", None) + if callable(get_model_info): + try: + provider_model_info = get_model_info( + model=model, + api_base=api_base, + ) + if provider_model_info is not None: + return provider_model_info + except Exception: + verbose_logger.debug("Could not get dynamic model info.") if custom_llm_provider == "huggingface": max_tokens = _get_max_position_embeddings(model_name=model)