mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-01 02:02:20 +00:00
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:
parent
aee740ce25
commit
dbba9f1bb3
9 changed files with 756 additions and 75 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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 = [
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue