Add provider hook for runtime model metadata

This commit is contained in:
Graham Neubig 2026-05-17 16:31:47 -04:00 • committed by openhands
parent 6f41a84c85
commit e87f900a78
5 changed files with 203 additions and 62 deletions

View file

@ -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,

View file

@ -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,

View file

@ -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]

View file

@ -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)

View file

@ -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."""