mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
Add provider hook for runtime model metadata
This commit is contained in:
parent
6f41a84c85
commit
e87f900a78
5 changed files with 203 additions and 62 deletions
|
|
@ -5,7 +5,7 @@ Utility functions for base LLM classes.
|
|||
import copy
|
||||
import json
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any, Dict, List, Optional, Type, Union
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Type, Union
|
||||
|
||||
from openai.lib import _parsing, _pydantic
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -14,6 +14,9 @@ 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
|
||||
|
|
@ -41,6 +44,17 @@ 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,10 +1,13 @@
|
|||
from typing import List, Optional, Union
|
||||
from typing import TYPE_CHECKING, 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__(
|
||||
|
|
@ -78,12 +81,27 @@ class OllamaModelInfo(BaseLLMModelInfo):
|
|||
# env var OLLAMA_API_BASE or default
|
||||
return api_base or get_secret_str("OLLAMA_API_BASE") or "http://localhost:11434"
|
||||
|
||||
@classmethod
|
||||
def get_server_api_base(cls, api_base: Optional[str] = None) -> str:
|
||||
api_base = cls.get_api_base(api_base).rstrip("/")
|
||||
for suffix in (
|
||||
"/api/generate",
|
||||
"/api/chat",
|
||||
"/api/embed",
|
||||
"/api/embeddings",
|
||||
"/api/show",
|
||||
"/api/tags",
|
||||
):
|
||||
if api_base.endswith(suffix):
|
||||
return api_base[: -len(suffix)]
|
||||
return api_base
|
||||
|
||||
def get_models(self, api_key=None, api_base: Optional[str] = None) -> List[str]:
|
||||
"""
|
||||
List all models available on the Ollama server via /api/tags endpoint.
|
||||
"""
|
||||
|
||||
base = self.get_api_base(api_base)
|
||||
base = self.get_server_api_base(api_base)
|
||||
api_key = self.get_api_key()
|
||||
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
|
||||
|
||||
|
|
@ -126,6 +144,92 @@ class OllamaModelInfo(BaseLLMModelInfo):
|
|||
result = sorted(names)
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _strip_ollama_model_prefix(model: str) -> str:
|
||||
if model.startswith("ollama/") or model.startswith("ollama_chat/"):
|
||||
return model.split("/", 1)[1]
|
||||
return model
|
||||
|
||||
@staticmethod
|
||||
def _is_static_ollama_model(model: str) -> bool:
|
||||
from litellm import model_cost
|
||||
|
||||
stripped_model = OllamaModelInfo._strip_ollama_model_prefix(model)
|
||||
potential_model_names = {
|
||||
model,
|
||||
stripped_model,
|
||||
"ollama/" + stripped_model,
|
||||
"ollama_chat/" + stripped_model,
|
||||
}
|
||||
model_cost_keys = {key.lower() for key in model_cost}
|
||||
return any(name.lower() in model_cost_keys for name in potential_model_names)
|
||||
|
||||
@staticmethod
|
||||
def _supports_function_calling(ollama_model_info: dict) -> bool:
|
||||
_template: str = str(ollama_model_info.get("template", "") or "")
|
||||
return "tools" in _template.lower()
|
||||
|
||||
@staticmethod
|
||||
def _get_max_tokens(ollama_model_info: dict) -> Optional[int]:
|
||||
_model_info: dict = ollama_model_info.get("model_info", {})
|
||||
|
||||
for key, value in _model_info.items():
|
||||
if "context_length" in key:
|
||||
return value
|
||||
return None
|
||||
|
||||
def get_runtime_model_info(
|
||||
self, model: str, api_base: Optional[str] = None
|
||||
) -> "ModelInfoBase":
|
||||
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)
|
||||
api_key = self.get_api_key()
|
||||
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
|
||||
|
||||
try:
|
||||
response = module_level_client.post(
|
||||
url=f"{api_base}/api/show",
|
||||
json={"name": model},
|
||||
headers=headers,
|
||||
)
|
||||
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,
|
||||
)
|
||||
|
||||
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,
|
||||
)
|
||||
|
||||
def get_model_info(
|
||||
self, model: str, api_base: Optional[str] = None
|
||||
) -> Optional["ModelInfoBase"]:
|
||||
if self._is_static_ollama_model(model):
|
||||
return None
|
||||
return self.get_runtime_model_info(model=model, api_base=api_base)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
|
|
|
|||
|
|
@ -6,7 +6,7 @@ from typing import TYPE_CHECKING, Any, AsyncIterator, Iterator, List, Optional,
|
|||
from httpx._models import Headers, Response
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger, verbose_proxy_logger
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
get_str_from_messages,
|
||||
)
|
||||
|
|
@ -17,7 +17,6 @@ from litellm.litellm_core_utils.prompt_templates.factory import (
|
|||
)
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUsageBlock
|
||||
from litellm.types.utils import (
|
||||
Delta,
|
||||
|
|
@ -29,7 +28,7 @@ from litellm.types.utils import (
|
|||
StreamingChoices,
|
||||
)
|
||||
|
||||
from ..common_utils import OllamaError, _convert_image
|
||||
from ..common_utils import OllamaError, OllamaModelInfo, _convert_image
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
|
@ -231,53 +230,7 @@ class OllamaConfig(BaseConfig):
|
|||
"name": "mistral"
|
||||
}'
|
||||
"""
|
||||
if model.startswith("ollama/") or model.startswith("ollama_chat/"):
|
||||
model = model.split("/", 1)[1]
|
||||
api_base = (
|
||||
api_base or get_secret_str("OLLAMA_API_BASE") or "http://localhost:11434"
|
||||
)
|
||||
api_key = self.get_api_key()
|
||||
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
|
||||
|
||||
try:
|
||||
response = litellm.module_level_client.post(
|
||||
url=f"{api_base}/api/show",
|
||||
json={"name": model},
|
||||
headers=headers,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_logger.debug(
|
||||
"OllamaError: Could not get model info for %s from %s. Error: %s",
|
||||
model,
|
||||
api_base,
|
||||
e,
|
||||
)
|
||||
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,
|
||||
)
|
||||
|
||||
model_info = response.json()
|
||||
|
||||
_max_tokens: Optional[int] = 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 OllamaModelInfo().get_runtime_model_info(model=model, api_base=api_base)
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, Headers]
|
||||
|
|
|
|||
|
|
@ -5703,6 +5703,22 @@ def _get_model_info_helper( # noqa: PLR0915
|
|||
split_model = potential_model_names["split_model"]
|
||||
custom_llm_provider = potential_model_names["custom_llm_provider"]
|
||||
#########################
|
||||
provider_config: Optional[BaseLLMModelInfo] = None
|
||||
if custom_llm_provider and custom_llm_provider in LlmProvidersSet:
|
||||
provider_config = ProviderConfigManager.get_provider_model_info(
|
||||
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.")
|
||||
|
||||
if custom_llm_provider == "huggingface":
|
||||
max_tokens = _get_max_position_embeddings(model_name=model)
|
||||
return ModelInfoBase(
|
||||
|
|
@ -5723,14 +5739,6 @@ def _get_model_info_helper( # noqa: PLR0915
|
|||
supports_computer_use=None,
|
||||
supports_pdf_input=None,
|
||||
)
|
||||
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, api_base=api_base)
|
||||
elif custom_llm_provider == "lemonade":
|
||||
return litellm.LemonadeChatConfig().get_model_info(
|
||||
model=model, api_base=api_base
|
||||
)
|
||||
else:
|
||||
"""
|
||||
Check if: (in order of specificity)
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
import os
|
||||
import sys
|
||||
from unittest.mock import patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
|
@ -23,6 +22,7 @@ if "httpx" not in sys.modules:
|
|||
sys.modules["httpx"] = httpx_mod
|
||||
|
||||
import httpx
|
||||
import litellm
|
||||
|
||||
from litellm.llms.ollama.common_utils import OllamaModelInfo
|
||||
|
||||
|
|
@ -214,6 +214,23 @@ class TestOllamaGetModelInfo:
|
|||
|
||||
assert captured_urls[0] == "http://env-server:11434/api/show"
|
||||
|
||||
def test_get_model_info_normalizes_generate_api_base(self, monkeypatch):
|
||||
"""When completion passes the final generate URL, model info should use the server base."""
|
||||
from litellm.llms.ollama.completion.transformation import OllamaConfig
|
||||
|
||||
captured_urls = []
|
||||
|
||||
def mock_post(url, json, headers=None):
|
||||
captured_urls.append(url)
|
||||
return DummyResponse({"template": "", "model_info": {}}, status_code=200)
|
||||
|
||||
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
|
||||
|
||||
config = OllamaConfig()
|
||||
config.get_model_info("llama3", api_base="http://localhost:11434/api/generate")
|
||||
|
||||
assert captured_urls[0] == "http://localhost:11434/api/show"
|
||||
|
||||
def test_get_model_info_graceful_fallback_on_connection_error(self, monkeypatch):
|
||||
"""When the Ollama server is unreachable, should return defaults instead of raising."""
|
||||
from litellm.llms.ollama.completion.transformation import OllamaConfig
|
||||
|
|
@ -252,6 +269,51 @@ class TestOllamaGetModelInfo:
|
|||
config.get_model_info("ollama_chat/llama3", api_base="http://localhost:11434")
|
||||
assert captured_json[1]["name"] == "llama3"
|
||||
|
||||
def test_litellm_get_model_info_uses_provider_hook_for_unknown_model(
|
||||
self, monkeypatch
|
||||
):
|
||||
"""Unmapped Ollama models should use the provider-level dynamic hook."""
|
||||
captured_json = []
|
||||
|
||||
def mock_post(url, json, headers=None):
|
||||
captured_json.append(json)
|
||||
return DummyResponse(
|
||||
{
|
||||
"template": "{{ .System }} tools {{ .Prompt }}",
|
||||
"model_info": {"llama.context_length": 32768},
|
||||
},
|
||||
status_code=200,
|
||||
)
|
||||
|
||||
litellm.get_model_info.cache_clear()
|
||||
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
|
||||
try:
|
||||
model_info = litellm.get_model_info(
|
||||
"ollama/unknown-model", api_base="http://localhost:11434"
|
||||
)
|
||||
finally:
|
||||
litellm.get_model_info.cache_clear()
|
||||
|
||||
assert model_info["max_input_tokens"] == 32768
|
||||
assert model_info["supports_function_calling"] is True
|
||||
assert captured_json[0]["name"] == "unknown-model"
|
||||
|
||||
def test_litellm_get_model_info_keeps_static_map_for_known_model(self, monkeypatch):
|
||||
"""Mapped Ollama models should keep using the static model map."""
|
||||
|
||||
def mock_post(url, json, headers=None):
|
||||
raise AssertionError("Static Ollama model should not query /api/show")
|
||||
|
||||
litellm.get_model_info.cache_clear()
|
||||
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
|
||||
try:
|
||||
model_info = litellm.get_model_info("ollama/llama2")
|
||||
finally:
|
||||
litellm.get_model_info.cache_clear()
|
||||
|
||||
assert model_info["key"] == "ollama/llama2"
|
||||
assert model_info["litellm_provider"] == "ollama"
|
||||
|
||||
|
||||
class TestOllamaAuthHeaders:
|
||||
"""Tests for Ollama authentication header handling in completion calls."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue