diff --git a/litellm/llms/base_llm/base_utils.py b/litellm/llms/base_llm/base_utils.py index d2d3d5c0a96..2afac26e7b6 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 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, diff --git a/litellm/llms/ollama/common_utils.py b/litellm/llms/ollama/common_utils.py index 8ca8b7d383a..c4cfd514114 100644 --- a/litellm/llms/ollama/common_utils.py +++ b/litellm/llms/ollama/common_utils.py @@ -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, diff --git a/litellm/llms/ollama/completion/transformation.py b/litellm/llms/ollama/completion/transformation.py index 32981776753..bb88c60ae64 100644 --- a/litellm/llms/ollama/completion/transformation.py +++ b/litellm/llms/ollama/completion/transformation.py @@ -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] diff --git a/litellm/utils.py b/litellm/utils.py index 1bfda06270a..7b30a7a05a4 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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) diff --git a/tests/test_litellm/llms/ollama/test_ollama_model_info.py b/tests/test_litellm/llms/ollama/test_ollama_model_info.py index 448a26bafe1..7d92b6a19c4 100644 --- a/tests/test_litellm/llms/ollama/test_ollama_model_info.py +++ b/tests/test_litellm/llms/ollama/test_ollama_model_info.py @@ -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."""