mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Address provider model info review feedback
Keep the runtime model info hook duck-typed instead of extending the base model-info class, and avoid importing ModelInfoBase from Ollama common utilities to reduce CodeQL cyclic-import noise. Co-authored-by: openhands <openhands@all-hands.dev>
This commit is contained in:
parent
e87f900a78
commit
db2ce18d5c
3 changed files with 36 additions and 52 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue