From 3322b282f85f06a4952f6784c06f68ff5646a877 Mon Sep 17 00:00:00 2001 From: Matthias Dittrich Date: Wed, 21 May 2025 17:47:01 +0200 Subject: [PATCH] Ollama wildcard support (#10982) * Add Ollama wildcard support * Add Ollama-chatas well. * Fix missing methods. * Improve logs a bit. * Add tests * Add tests --- litellm/llms/ollama/common_utils.py | 86 +++++++++++++- litellm/utils.py | 8 +- .../llms/ollama/test_ollama_model_info.py | 106 ++++++++++++++++++ 3 files changed, 196 insertions(+), 4 deletions(-) create mode 100644 tests/litellm/llms/ollama/test_ollama_model_info.py diff --git a/litellm/llms/ollama/common_utils.py b/litellm/llms/ollama/common_utils.py index 5cf213950c1..b2e781e5665 100644 --- a/litellm/llms/ollama/common_utils.py +++ b/litellm/llms/ollama/common_utils.py @@ -1,10 +1,14 @@ from typing import Union +from litellm import verbose_logger -import httpx +# dynamic import to allow usage even if httpx is not installed in dev env +try: + import httpx +except ImportError: + httpx = None # type: ignore from litellm.llms.base_llm.chat.transformation import BaseLLMException - class OllamaError(BaseLLMException): def __init__( self, status_code: int, message: str, headers: Union[dict, httpx.Headers] @@ -43,3 +47,81 @@ def _convert_image(image): image_data.convert("RGB").save(jpeg_image, "JPEG") jpeg_image.seek(0) return base64.b64encode(jpeg_image.getvalue()).decode("utf-8") + + +from litellm.llms.base_llm.base_utils import BaseLLMModelInfo + +class OllamaModelInfo(BaseLLMModelInfo): + """ + Dynamic model listing for Ollama server. + Fetches /api/models and /api/tags, then for each tag also /api/models?tag=... + Returns the union of all model names. + """ + @staticmethod + def get_api_key(api_key=None) -> None: + return None # Ollama does not use an API key by default + + @staticmethod + def get_api_base(api_base: str | None = None) -> str: + from litellm.secret_managers.main import get_secret_str + # env var OLLAMA_API_BASE or default + return api_base or get_secret_str("OLLAMA_API_BASE") or "http://localhost:11434" + + def get_models(self, api_key=None, api_base: str | None = None) -> list[str]: + """ + List all models available on the Ollama server via /api/tags endpoint. + """ + import httpx + base = self.get_api_base(api_base) + names: set[str] = set() + try: + resp = httpx.get(f"{base}/api/tags") + resp.raise_for_status() + data = resp.json() + # Expecting a dict with a 'models' list + models_list = [] + if isinstance(data, dict) and 'models' in data and isinstance(data['models'], list): + models_list = data['models'] + elif isinstance(data, list): + models_list = data + # Extract model names + for entry in models_list: + if not isinstance(entry, dict): + continue + nm = entry.get('name') or entry.get('model') + if isinstance(nm, str): + names.add(nm) + except Exception as e: + verbose_logger.warning(f"Error retrieving ollama tag endpoint: {e}") + # If tags endpoint fails, fall back to static list + try: + from litellm import models_by_provider + static = models_by_provider.get("ollama", []) or [] + return [f"ollama/{m}" for m in static] + except Exception as e1: + verbose_logger.warning(f"Error retrieving static ollama models as fallback: {e1}") + return [] + # assemble full model names + result = sorted(names) + return result + def validate_environment( + self, + headers: dict, + model: str, + messages: list, + optional_params: dict, + litellm_params: dict, + api_key=None, + api_base=None, + ) -> dict: + """ + No-op environment validation for Ollama. + """ + return {} + + @staticmethod + def get_base_model(model: str) -> str: + """ + Return the base model name for Ollama (no-op). + """ + return model diff --git a/litellm/utils.py b/litellm/utils.py index 5d6a64cd345..0dd4c8cd1e3 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -5851,7 +5851,7 @@ def _get_valid_models_from_provider_api( _model_cache.set_cached_model_info(custom_llm_provider, litellm_params, models) return models except Exception as e: - verbose_logger.debug(f"Error getting valid models: {e}") + verbose_logger.warning(f"Error getting valid models: {e}") return [] @@ -5916,7 +5916,7 @@ def get_valid_models( return valid_models except Exception as e: - verbose_logger.debug(f"Error getting valid models: {e}") + verbose_logger.warning(f"Error getting valid models: {e}") return [] # NON-Blocking @@ -6599,6 +6599,10 @@ class ProviderConfigManager: return litellm.AnthropicModelInfo() elif LlmProviders.XAI == provider: return litellm.XAIModelInfo() + elif LlmProviders.OLLAMA == provider or LlmProviders.OLLAMA_CHAT == provider: + # Dynamic model listing for Ollama server + from litellm.llms.ollama.common_utils import OllamaModelInfo + return OllamaModelInfo() elif LlmProviders.VLLM == provider: from litellm.llms.vllm.common_utils import ( VLLMModelInfo, # experimental approach, to reduce bloat on __init__.py diff --git a/tests/litellm/llms/ollama/test_ollama_model_info.py b/tests/litellm/llms/ollama/test_ollama_model_info.py new file mode 100644 index 00000000000..f5fea572ac9 --- /dev/null +++ b/tests/litellm/llms/ollama/test_ollama_model_info.py @@ -0,0 +1,106 @@ +import os +import sys +import json +import uuid +import pytest +from unittest.mock import MagicMock, patch + + +sys.path.insert( + 0, os.path.abspath("../../../../..") +) # Adds the parent directory to the system path + +""" +Unit tests for OllamaModelInfo.get_models functionality. +""" +# Ensure a dummy httpx module is available for import in tests +import sys, types +# Provide a dummy httpx module for import in get_models +if 'httpx' not in sys.modules: + # Create a minimal module with HTTPStatusError + httpx_mod = types.ModuleType('httpx') + httpx_mod.HTTPStatusError = Exception + sys.modules['httpx'] = httpx_mod + +import httpx + +from litellm.llms.ollama.common_utils import OllamaModelInfo + + +class DummyResponse: + """ + A dummy response object to simulate httpx responses. + """ + def __init__(self, json_data, status_code=200): + self._json = json_data + self.status_code = status_code + + def raise_for_status(self): + if self.status_code >= 400: + # Simulate an HTTP status error + raise httpx.HTTPStatusError("Error status code", request=None, response=None) + + def json(self): + return self._json + + +class TestOllamaModelInfo: + def test_get_models_from_dict_response(self, monkeypatch): + """ + When the /api/tags endpoint returns a dict with a 'models' list, + get_models should extract and return sorted unique model names. + """ + calls = [] + sample = {'models': [ + {'name': 'zeta'}, + {'model': 'alpha'}, + {'name': 123}, # non-str should be ignored + 'invalid', # non-dict should be ignored + ]} + + def mock_get(url): + calls.append(url) + return DummyResponse(sample, status_code=200) + + monkeypatch.setattr(httpx, 'get', mock_get) + info = OllamaModelInfo() + models = info.get_models() + # Only 'alpha' and 'zeta' should be returned, sorted alphabetically + assert models == ['alpha', 'zeta'] + # Ensure correct endpoint was called + assert calls and calls[0].endswith('/api/tags') + + + def test_get_models_from_list_response(self, monkeypatch): + """ + When the /api/tags endpoint returns a list of dicts, + get_models should extract and return sorted unique model names. + """ + sample = [ + {'name': 'm1'}, + {'model': 'm2'}, + {}, # no name/model key should be ignored + ] + + def mock_get(url): + return DummyResponse(sample, status_code=200) + + monkeypatch.setattr(httpx, 'get', mock_get) + info = OllamaModelInfo() + models = info.get_models() + assert models == ['m1', 'm2'] + + + def test_get_models_fallback_on_error(self, monkeypatch): + """ + If the httpx.get call raises an exception, get_models should + fall back to the static models_by_provider list prefixed by 'ollama/'. + """ + def mock_get(url): + raise Exception("connection failure") + + monkeypatch.setattr(httpx, 'get', mock_get) + info = OllamaModelInfo() + models = info.get_models() + # Default static ollama_models is ['llama2'], so expect ['ollama/llama2'] + assert models == ['ollama/llama2'] \ No newline at end of file