Support Lemonade runtime context metadata (#28135)

* Support Lemonade runtime context metadata

* Add provider hook for runtime model metadata

* 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>

* Fix CI after staging rebase

Relax the Ollama runtime metadata return annotation to match the provider-hook dict response and update the Google Interactions OpenAPI status expectation for the current live spec.

Co-authored-by: openhands <openhands@all-hands.dev>

* Normalize Lemonade runtime model metadata

* Avoid leaking Ollama metadata auth

* Avoid leaking Lemonade metadata auth

---------

Co-authored-by: Graham Neubig <398875+neubig@users.noreply.github.com>
Co-authored-by: openhands <openhands@all-hands.dev>
This commit is contained in:
Graham Neubig 2026-05-28 09:02:14 -07:00 • committed by GitHub
parent aee740ce25
commit dbba9f1bb3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
9 changed files with 756 additions and 75 deletions

View file

@ -868,6 +868,7 @@ openai_text_completion_compatible_providers: List = (
_openai_like_providers: List = [
"predibase",
"databricks",
"lemonade",
"watsonx",
] # private helper. similar to openai but require some custom auth / endpoint handling, so can't use the openai sdk
# well supported replicate llms

View file

@ -87,6 +87,7 @@ class ExceptionCheckers:
"is longer than the model's context length",
"input tokens exceed the configured limit",
"`inputs` tokens + `max_new_tokens` must be",
"exceeds the available context size", # llama.cpp/Lemonade
"exceeds the maximum number of tokens allowed", # Gemini
]
for substring in known_exception_substrings:
@ -891,12 +892,14 @@ def exception_type( # type: ignore # noqa: PLR0915
response=getattr(original_exception, "response", None),
litellm_debug_info=extra_information,
)
elif "model's maximum context limit" in error_str:
elif ExceptionCheckers.is_error_str_context_window_exceeded(error_str):
exception_mapping_worked = True
raise ContextWindowExceededError(
message=f"{custom_llm_provider.capitalize()}Exception: Context Window Error - {error_str}",
model=model,
llm_provider=custom_llm_provider,
response=getattr(original_exception, "response", None),
litellm_debug_info=extra_information,
)
elif "token_quota_reached" in error_str:
exception_mapping_worked = True

View file

@ -3,10 +3,12 @@ Translate from OpenAI's `/v1/chat/completions` to Lemonade's `/v1/chat/completio
"""
from typing import Any, List, Optional, Tuple, Union
from urllib.parse import quote
import httpx
import litellm
from litellm._logging import verbose_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import (
@ -18,6 +20,8 @@ from ...openai_like.chat.transformation import OpenAILikeChatConfig
class LemonadeChatConfig(OpenAILikeChatConfig):
_DEFAULT_API_KEY = "lemonade"
repeat_penalty: Optional[float] = None
functions: Optional[list] = None
logit_bias: Optional[dict] = None
@ -68,7 +72,7 @@ class LemonadeChatConfig(OpenAILikeChatConfig):
This method queries the Lemonade /models endpoint to retrieve the list of available models.
Args:
api_key: Optional API key (Lemonade doesn't require authentication)
api_key: Optional API key for authenticated Lemonade servers
api_base: Optional API base URL (defaults to LEMONADE_API_BASE env var or http://localhost:8000)
Returns:
@ -87,6 +91,7 @@ class LemonadeChatConfig(OpenAILikeChatConfig):
try:
response = litellm.module_level_client.get(
url=f"{api_base}/models",
headers=self._get_auth_headers(api_key),
)
except Exception as e:
raise ValueError(
@ -101,19 +106,131 @@ class LemonadeChatConfig(OpenAILikeChatConfig):
model_list = response.json().get("data", [])
return ["lemonade/" + model["id"] for model in model_list]
@staticmethod
def _get_positive_int(value: Any) -> Optional[int]:
if isinstance(value, bool):
return None
if isinstance(value, int) and value > 0:
return value
if isinstance(value, str):
try:
parsed = int(value)
except ValueError:
return None
if parsed > 0:
return parsed
return None
@staticmethod
def _get_provider_specific_entry(model_info: dict) -> dict:
provider_specific_entry = model_info.get("provider_specific_entry")
if not isinstance(provider_specific_entry, dict):
provider_specific_entry = {}
else:
provider_specific_entry = provider_specific_entry.copy()
for key in ("recipe_options", "context_window", "max_context_window"):
if key in model_info:
provider_specific_entry[key] = model_info[key]
return provider_specific_entry
def _get_context_window(self, model_info: dict) -> Optional[int]:
provider_specific_entry = self._get_provider_specific_entry(model_info)
recipe_options = provider_specific_entry.get("recipe_options")
if not isinstance(recipe_options, dict):
recipe_options = {}
for value in (
recipe_options.get("ctx_size"),
model_info.get("max_input_tokens"),
provider_specific_entry.get("context_window"),
provider_specific_entry.get("max_context_window"),
):
parsed = self._get_positive_int(value)
if parsed is not None:
return parsed
return None
def _get_default_model_info(self, model: str) -> dict:
return {
"key": "lemonade/" + model,
"litellm_provider": "lemonade",
"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,
}
def get_model_info(
self,
model: str,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
) -> Any:
if model.startswith("lemonade/"):
model = model.split("/", 1)[1]
api_base, api_key = self._get_openai_compatible_provider_info(
api_base=api_base, api_key=api_key
)
encoded_model = quote(model, safe="")
try:
response = litellm.module_level_client.get(
url=f"{api_base}/models/{encoded_model}",
headers=self._get_auth_headers(api_key),
)
response.raise_for_status()
model_info = response.json()
except Exception:
verbose_logger.debug("LemonadeError: Could not get model info.")
return self._get_default_model_info(model)
max_input_tokens = self._get_context_window(model_info)
max_output_tokens = self._get_positive_int(model_info.get("max_output_tokens"))
max_tokens = self._get_positive_int(model_info.get("max_tokens"))
provider_specific_entry = self._get_provider_specific_entry(model_info)
model_info_response = self._get_default_model_info(model)
model_info_response.update(
{
"max_tokens": max_tokens or max_output_tokens,
"max_input_tokens": max_input_tokens,
"max_output_tokens": max_output_tokens,
}
)
if provider_specific_entry:
model_info_response["provider_specific_entry"] = provider_specific_entry
return model_info_response
def _get_openai_compatible_provider_info(
self, api_base: Optional[str], api_key: Optional[str]
) -> Tuple[Optional[str], Optional[str]]:
# lemonade is openai compatible, we just need to set this to custom_openai and have the api_base be lemonade's endpoint
passed_api_base = api_base
api_base = (
api_base
or get_secret_str("LEMONADE_API_BASE")
or "http://localhost:8000/api/v1"
) # type: ignore
# Lemonade doesn't check the key
key = "lemonade"
key = self._DEFAULT_API_KEY
if api_key is not None or passed_api_base is None:
key = (
api_key
or litellm.lemonade_key
or get_secret_str("LEMONADE_API_KEY")
or self._DEFAULT_API_KEY
)
return api_base, key
def _get_auth_headers(self, api_key: Optional[str]) -> dict:
if api_key is None or api_key == self._DEFAULT_API_KEY:
return {}
return {"Authorization": f"Bearer {api_key}"}
def transform_response(
self,
model: str,

View file

@ -1,4 +1,4 @@
from typing import List, Optional, Union
from typing import Any, List, Optional, Union
import httpx
@ -65,7 +65,8 @@ class OllamaModelInfo(BaseLLMModelInfo):
from litellm.secret_managers.main import get_secret_str
return (
os.environ.get("OLLAMA_API_KEY")
api_key
or os.environ.get("OLLAMA_API_KEY")
or litellm.api_key
or litellm.openai_key
or get_secret_str("OLLAMA_API_KEY")
@ -78,13 +79,33 @@ 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)
api_key = self.get_api_key()
passed_api_base = api_base
base = self.get_server_api_base(api_base)
api_key = (
self.get_api_key(api_key)
if api_key is not None or passed_api_base is None
else None
)
headers = {"Authorization": f"Bearer {api_key}"} if api_key else {}
names: set[str] = set()
@ -126,6 +147,104 @@ 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,
api_key: Optional[str] = None,
) -> dict[str, Any]:
from litellm import module_level_client
model = self._strip_ollama_model_prefix(model)
passed_api_base = api_base
api_base = self.get_server_api_base(api_base)
api_key = (
self.get_api_key(api_key)
if api_key is not None or passed_api_base is None
else None
)
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 {
"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 {
"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,
api_key: Optional[str] = None,
) -> 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, api_key=api_key
)
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,19 +17,17 @@ 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,
GenericStreamingChunk,
ModelInfoBase,
ModelResponse,
ModelResponseStream,
ProviderField,
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
@ -224,59 +222,18 @@ class OllamaConfig(BaseConfig):
)
def get_model_info(
self, model: str, api_base: Optional[str] = None
) -> ModelInfoBase:
self,
model: str,
api_base: Optional[str] = None,
api_key: Optional[str] = None,
) -> Any:
"""
curl http://localhost:11434/api/show -d '{
"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, api_key=api_key
)
def get_error_class(

View file

@ -5754,6 +5754,24 @@ 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:
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)
return ModelInfoBase(
@ -5774,10 +5792,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)
else:
"""
Check if: (in order of specificity)

View file

@ -14,6 +14,7 @@ from litellm.litellm_core_utils.exception_mapping_utils import (
exception_type,
extract_and_raise_litellm_exception,
)
from litellm.llms.openai.common_utils import OpenAIError
# Test cases for is_error_str_context_window_exceeded
# Tuple format: (error_message, expected_result)
@ -41,6 +42,10 @@ context_window_test_cases = [
"`inputs` tokens + `max_new_tokens` must be <= 4096",
True,
),
(
"request (67311 tokens) exceeds the available context size (65536 tokens), try increasing it",
True,
),
# Gemini 2.5/3 format
(
"The input token count exceeds the maximum number of tokens allowed 1048576.",
@ -182,7 +187,6 @@ class TestExceptionCheckers:
]
for error_str in positive_cases:
print("testing positive case=", error_str)
result = ExceptionCheckers.is_azure_content_policy_violation_error(
error_str
)
@ -255,6 +259,33 @@ def test_gemini_context_window_error_mapping(
)
def test_lemonade_context_window_error_mapping():
"""Lemonade's llama.cpp backend should map context overflows to LiteLLM's standard error."""
model = "lemonade/Qwen3.6-35B-A3B-GGUF"
error_message = (
'{"error":{"code":"context_length_exceeded","message":"request '
"(80010 tokens) exceeds the available context size (65536 tokens), "
'try increasing it","status_code":400,"type":"invalid_request_error"}}'
)
original_exception = OpenAIError(
status_code=400,
message=error_message,
headers={},
)
with pytest.raises(litellm.ContextWindowExceededError) as excinfo:
exception_type(
model=model,
original_exception=original_exception,
custom_llm_provider="lemonade",
)
assert excinfo.value.status_code == 400
assert excinfo.value.llm_provider == "lemonade"
assert excinfo.value.model == model
# Test cases for Vertex AI RateLimitError mapping
# As per https://github.com/BerriAI/litellm/issues/16189
vertex_rate_limit_test_cases = [

View file

@ -1,17 +1,14 @@
import json
import os
import sys
import pytest
sys.path.insert(
0, os.path.abspath("../../../../..")
) # Adds the parent directory to the system path
from unittest.mock import MagicMock, patch
import litellm
from litellm.llms.lemonade.chat.transformation import LemonadeChatConfig
from litellm.types.utils import ModelResponse
import httpx
def test_lemonade_config_initialization():
@ -28,8 +25,11 @@ def test_lemonade_config_initialization():
assert config.repeat_penalty == 1.1
def test_get_openai_compatible_provider_info():
def test_get_openai_compatible_provider_info(monkeypatch):
"""Test the provider info method returns correct API base and key"""
monkeypatch.delenv("LEMONADE_API_KEY", raising=False)
monkeypatch.setattr(litellm, "lemonade_key", None)
monkeypatch.setattr(litellm, "api_key", None)
config = LemonadeChatConfig()
api_base, key = config._get_openai_compatible_provider_info(
@ -40,8 +40,11 @@ def test_get_openai_compatible_provider_info():
assert key == "lemonade"
def test_get_openai_compatible_provider_info_with_custom_base():
def test_get_openai_compatible_provider_info_with_custom_base(monkeypatch):
"""Test the provider info method with custom API base"""
monkeypatch.delenv("LEMONADE_API_KEY", raising=False)
monkeypatch.setattr(litellm, "lemonade_key", None)
monkeypatch.setattr(litellm, "api_key", None)
config = LemonadeChatConfig()
custom_api_base = "https://custom.lemonade.ai/v1"
@ -53,6 +56,280 @@ def test_get_openai_compatible_provider_info_with_custom_base():
assert key == "lemonade"
def test_get_openai_compatible_provider_info_with_api_key_env(monkeypatch):
"""Test the provider info method reads Lemonade's API key from the environment."""
monkeypatch.setenv("LEMONADE_API_KEY", "test-key")
monkeypatch.setattr(litellm, "lemonade_key", None)
monkeypatch.setattr(litellm, "api_key", None)
config = LemonadeChatConfig()
api_base, key = config._get_openai_compatible_provider_info(
api_base=None, api_key=None
)
assert api_base == "http://localhost:8000/api/v1"
assert key == "test-key"
def test_get_openai_compatible_provider_info_skips_env_key_for_custom_base(
monkeypatch,
):
"""Test that caller-supplied bases do not receive server-side Lemonade keys."""
monkeypatch.setenv("LEMONADE_API_KEY", "server-side-lemonade-key")
monkeypatch.setattr(litellm, "lemonade_key", "configured-lemonade-key")
monkeypatch.setattr(litellm, "api_key", None)
config = LemonadeChatConfig()
api_base, key = config._get_openai_compatible_provider_info(
api_base="https://attacker.example/v1", api_key=None
)
assert api_base == "https://attacker.example/v1"
assert key == "lemonade"
assert config._get_auth_headers(key) == {}
def test_get_openai_compatible_provider_info_uses_explicit_key_for_custom_base(
monkeypatch,
):
"""Test that explicitly supplied Lemonade keys are sent to supplied bases."""
monkeypatch.setenv("LEMONADE_API_KEY", "server-side-lemonade-key")
monkeypatch.setattr(litellm, "lemonade_key", "configured-lemonade-key")
monkeypatch.setattr(litellm, "api_key", None)
config = LemonadeChatConfig()
api_base, key = config._get_openai_compatible_provider_info(
api_base="https://lemonade.example/v1", api_key="explicit-lemonade-key"
)
assert api_base == "https://lemonade.example/v1"
assert key == "explicit-lemonade-key"
assert config._get_auth_headers(key) == {
"Authorization": "Bearer explicit-lemonade-key"
}
def test_get_openai_compatible_provider_info_ignores_global_api_key(monkeypatch):
"""Test that Lemonade discovery does not send unrelated global API keys."""
monkeypatch.delenv("LEMONADE_API_KEY", raising=False)
monkeypatch.setattr(litellm, "lemonade_key", None)
monkeypatch.setattr(litellm, "api_key", "global-openai-key")
config = LemonadeChatConfig()
api_base, key = config._get_openai_compatible_provider_info(
api_base="http://lemonade.test/v1", api_key=None
)
assert api_base == "http://lemonade.test/v1"
assert key == "lemonade"
assert config._get_auth_headers(key) == {}
def test_get_models_does_not_leak_lemonade_key_to_custom_base(monkeypatch):
"""Test Lemonade discovery does not send server-side keys to supplied bases."""
monkeypatch.setenv("LEMONADE_API_KEY", "server-side-lemonade-key")
monkeypatch.setattr(litellm, "lemonade_key", "configured-lemonade-key")
monkeypatch.setattr(litellm, "api_key", "global-provider-key")
config = LemonadeChatConfig()
response = MagicMock()
response.status_code = 200
response.json.return_value = {"data": []}
with patch.object(
litellm.module_level_client, "get", return_value=response
) as mock_get:
models = config.get_models(api_base="https://attacker.example/v1")
assert models == []
assert mock_get.call_args.kwargs["headers"] == {}
def test_get_model_info_uses_loaded_context_size():
"""Test that Lemonade model info prefers the effective loaded ctx_size."""
config = LemonadeChatConfig()
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"id": "Qwen3.6-35B-A3B-GGUF",
"recipe_options": {"ctx_size": 65536},
"max_context_window": 262144,
}
with patch.object(
litellm.module_level_client, "get", return_value=response
) as mock_get:
model_info = config.get_model_info(
model="lemonade/Qwen3.6-35B-A3B-GGUF",
api_base="http://lemonade.test/v1",
)
assert model_info["key"] == "lemonade/Qwen3.6-35B-A3B-GGUF"
assert model_info["litellm_provider"] == "lemonade"
assert model_info["max_input_tokens"] == 65536
assert model_info["provider_specific_entry"] == {
"recipe_options": {"ctx_size": 65536},
"max_context_window": 262144,
}
assert "supports_function_calling" not in model_info
assert "supports_response_schema" not in model_info
assert "supports_tool_choice" not in model_info
assert mock_get.call_args.kwargs["headers"] == {}
def test_get_model_info_falls_back_when_server_unavailable():
"""Test that Lemonade metadata lookup failures return safe defaults."""
config = LemonadeChatConfig()
with patch.object(
litellm.module_level_client, "get", side_effect=Exception("boom")
):
model_info = config.get_model_info(
model="lemonade/Qwen3.6-35B-A3B-GGUF",
api_base="http://lemonade.test/v1",
)
assert model_info["key"] == "lemonade/Qwen3.6-35B-A3B-GGUF"
assert model_info["litellm_provider"] == "lemonade"
assert model_info["mode"] == "chat"
assert model_info["input_cost_per_token"] == 0.0
assert model_info["output_cost_per_token"] == 0.0
assert model_info["max_tokens"] is None
assert model_info["max_input_tokens"] is None
assert model_info["max_output_tokens"] is None
assert "supports_function_calling" not in model_info
assert "supports_response_schema" not in model_info
assert "supports_tool_choice" not in model_info
def test_get_model_info_reads_context_from_provider_specific_entry():
"""Test that Lemonade model info uses provider-specific runtime metadata."""
config = LemonadeChatConfig()
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"id": "Qwen3.6-35B-A3B-GGUF",
"provider_specific_entry": {
"recipe_options": {"ctx_size": "32768"},
"max_context_window": 262144,
},
}
with patch.object(litellm.module_level_client, "get", return_value=response):
model_info = config.get_model_info(
model="lemonade/Qwen3.6-35B-A3B-GGUF",
api_base="http://lemonade.test/v1",
)
assert model_info["max_input_tokens"] == 32768
assert model_info["provider_specific_entry"] == {
"recipe_options": {"ctx_size": "32768"},
"max_context_window": 262144,
}
def test_get_model_info_sends_lemonade_api_key_for_configured_base(monkeypatch):
"""Test that Lemonade model info uses auth for configured servers."""
monkeypatch.setenv("LEMONADE_API_KEY", "test-key")
monkeypatch.setenv("LEMONADE_API_BASE", "http://lemonade.test/v1")
monkeypatch.setattr(litellm, "lemonade_key", None)
monkeypatch.setattr(litellm, "api_key", None)
config = LemonadeChatConfig()
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"id": "Qwen3.6-35B-A3B-GGUF",
"recipe_options": {"ctx_size": 65536},
}
with patch.object(
litellm.module_level_client, "get", return_value=response
) as mock_get:
config.get_model_info(
model="lemonade/Qwen3.6-35B-A3B-GGUF",
)
assert mock_get.call_args.kwargs["headers"] == {"Authorization": "Bearer test-key"}
def test_get_model_info_sends_explicit_lemonade_api_key_for_custom_base(monkeypatch):
"""Test that Lemonade model info sends explicitly supplied auth to supplied bases."""
monkeypatch.setenv("LEMONADE_API_KEY", "server-side-key")
monkeypatch.setattr(litellm, "lemonade_key", None)
monkeypatch.setattr(litellm, "api_key", None)
config = LemonadeChatConfig()
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"id": "Qwen3.6-35B-A3B-GGUF",
"recipe_options": {"ctx_size": 65536},
}
with patch.object(
litellm.module_level_client, "get", return_value=response
) as mock_get:
config.get_model_info(
model="lemonade/Qwen3.6-35B-A3B-GGUF",
api_base="http://lemonade.test/v1",
api_key="explicit-test-key",
)
assert mock_get.call_args.kwargs["headers"] == {
"Authorization": "Bearer explicit-test-key"
}
def test_litellm_get_model_info_does_not_leak_lemonade_key_to_custom_base(
monkeypatch,
):
"""Test top-level model info does not send server-side keys to supplied bases."""
monkeypatch.setenv("LEMONADE_API_KEY", "server-side-lemonade-key")
monkeypatch.setattr(litellm, "lemonade_key", "configured-lemonade-key")
monkeypatch.setattr(litellm, "api_key", "global-provider-key")
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"id": "Qwen3.6-35B-A3B-GGUF",
"max_input_tokens": 65536,
"max_context_window": 262144,
}
litellm.get_model_info.cache_clear()
with patch.object(
litellm.module_level_client, "get", return_value=response
) as mock_get:
try:
model_info = litellm.get_model_info(
model="lemonade/Qwen3.6-35B-A3B-GGUF",
api_base="https://attacker.example/v1",
)
finally:
litellm.get_model_info.cache_clear()
assert model_info["max_input_tokens"] == 65536
assert mock_get.call_args.kwargs["headers"] == {}
def test_litellm_get_model_info_uses_lemonade_api_base():
"""Test that LiteLLM model info is wired to Lemonade's model metadata API."""
response = MagicMock()
response.status_code = 200
response.json.return_value = {
"id": "Qwen3.6-35B-A3B-GGUF",
"max_input_tokens": 65536,
"max_context_window": 262144,
}
with patch.object(litellm.module_level_client, "get", return_value=response):
model_info = litellm.get_model_info(
model="lemonade/Qwen3.6-35B-A3B-GGUF",
api_base="http://lemonade.test/v1",
)
assert model_info["max_input_tokens"] == 65536
assert response.raise_for_status.called
assert response.json.called
def test_transform_response():
"""Test the response transformation adds lemonade prefix to model name"""
config = LemonadeChatConfig()

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
@ -105,6 +105,47 @@ class TestOllamaModelInfo:
"Authorization": "Bearer test_api_key"
}
def test_get_models_does_not_leak_server_key_to_provided_api_base(
self, monkeypatch
):
"""Model discovery should not send server-side keys to caller-supplied bases."""
call_headers = []
def mock_get(url, headers):
call_headers.append(headers)
return DummyResponse({"models": []}, status_code=200)
monkeypatch.setenv("OLLAMA_API_KEY", "server-side-ollama-key")
monkeypatch.setattr(litellm, "api_key", "global-provider-key")
monkeypatch.setattr(litellm, "openai_key", "global-openai-key")
monkeypatch.setattr(httpx, "get", mock_get)
info = OllamaModelInfo()
models = info.get_models(api_base="https://attacker.example")
assert models == []
assert call_headers[0] == {}
def test_get_models_uses_explicit_api_key_for_provided_api_base(self, monkeypatch):
"""Model discovery should send an explicitly supplied key to the provided base."""
call_headers = []
def mock_get(url, headers):
call_headers.append(headers)
return DummyResponse({"models": []}, status_code=200)
monkeypatch.setenv("OLLAMA_API_KEY", "server-side-ollama-key")
monkeypatch.setattr(httpx, "get", mock_get)
info = OllamaModelInfo()
models = info.get_models(
api_base="https://ollama.example",
api_key="explicit-api-key",
)
assert models == []
assert call_headers[0] == {"Authorization": "Bearer explicit-api-key"}
def test_get_models_from_list_response(self, monkeypatch):
"""
When the /api/tags endpoint returns a list of dicts,
@ -200,6 +241,83 @@ class TestOllamaGetModelInfo:
"""When no api_base is passed, should fall back to OLLAMA_API_BASE env var."""
from litellm.llms.ollama.completion.transformation import OllamaConfig
captured_urls = []
captured_headers = []
def mock_post(url, json, headers=None):
captured_urls.append(url)
captured_headers.append(headers)
return DummyResponse({"template": "", "model_info": {}}, status_code=200)
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
monkeypatch.setenv("OLLAMA_API_BASE", "http://env-server:11434")
monkeypatch.setenv("OLLAMA_API_KEY", "env-api-key")
config = OllamaConfig()
config.get_model_info("llama3")
assert captured_urls[0] == "http://env-server:11434/api/show"
assert captured_headers[0] == {"Authorization": "Bearer env-api-key"}
def test_get_model_info_uses_explicit_api_key_for_provided_api_base(
self, monkeypatch
):
"""When api_key is explicit, model info should send it to the provided api_base."""
from litellm.llms.ollama.completion.transformation import OllamaConfig
captured_headers = []
def mock_post(url, json, headers=None):
captured_headers.append(headers)
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://my-remote-server:11434",
api_key="explicit-api-key",
)
assert captured_headers[0] == {"Authorization": "Bearer explicit-api-key"}
def test_litellm_get_model_info_does_not_leak_server_key_to_provided_api_base(
self, monkeypatch
):
"""Global model info should not send server-side keys to caller-supplied bases."""
captured_headers = []
def mock_post(url, json, headers=None):
captured_headers.append(headers)
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)
monkeypatch.setenv("OLLAMA_API_KEY", "server-side-ollama-key")
monkeypatch.setattr(litellm, "api_key", "global-provider-key")
monkeypatch.setattr(litellm, "openai_key", "global-openai-key")
try:
model_info = litellm.get_model_info(
"ollama/unknown-model",
api_base="https://attacker.example",
)
finally:
litellm.get_model_info.cache_clear()
assert model_info["max_input_tokens"] == 32768
assert captured_headers[0] == {}
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):
@ -207,12 +325,11 @@ class TestOllamaGetModelInfo:
return DummyResponse({"template": "", "model_info": {}}, status_code=200)
monkeypatch.setattr("litellm.module_level_client.post", mock_post)
monkeypatch.setenv("OLLAMA_API_BASE", "http://env-server:11434")
config = OllamaConfig()
config.get_model_info("llama3")
config.get_model_info("llama3", api_base="http://localhost:11434/api/generate")
assert captured_urls[0] == "http://env-server:11434/api/show"
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."""
@ -252,6 +369,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."""