mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge ad4488b168 into dab2deb5ed
This commit is contained in:
commit
1659f3b29f
5 changed files with 600 additions and 35 deletions
|
|
@ -1,15 +1,30 @@
|
|||
from typing import Final
|
||||
from typing import Annotated, Final, Literal
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.base_utils import BaseLLMModelInfo
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelInfoBase
|
||||
from litellm.utils import _add_path_to_api_base
|
||||
|
||||
|
||||
class _VLLMModelEntry(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
|
||||
id: str
|
||||
max_model_len: Annotated[int, Field(strict=True, gt=0)] | None = None
|
||||
|
||||
|
||||
class _VLLMModelsResponse(BaseModel):
|
||||
model_config = ConfigDict(extra="ignore", frozen=True)
|
||||
|
||||
data: tuple[object, ...]
|
||||
|
||||
|
||||
class VLLMError(BaseLLMException):
|
||||
def __init__(
|
||||
self,
|
||||
|
|
@ -29,6 +44,9 @@ class VLLMError(BaseLLMException):
|
|||
|
||||
|
||||
class VLLMModelInfo(BaseLLMModelInfo):
|
||||
def __init__(self, provider: Literal["vllm", "hosted_vllm"] = "vllm") -> None:
|
||||
self._provider: Final = provider
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
|
|
@ -56,29 +74,63 @@ class VLLMModelInfo(BaseLLMModelInfo):
|
|||
def get_api_key(api_key: str | None = None) -> str | None:
|
||||
return None
|
||||
|
||||
def _get_discovery_api_base(self, api_base: str | None) -> str:
|
||||
environment_variable: Final = "HOSTED_VLLM_API_BASE" if self._provider == "hosted_vllm" else "VLLM_API_BASE"
|
||||
resolved_api_base: Final = api_base or get_secret_str(environment_variable)
|
||||
if resolved_api_base is None:
|
||||
raise ValueError(f"{environment_variable} is required to query vLLM's `/models` endpoint.")
|
||||
return resolved_api_base
|
||||
|
||||
@staticmethod
|
||||
def get_base_model(model: str) -> str | None:
|
||||
return model
|
||||
|
||||
def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]:
|
||||
api_base = VLLMModelInfo.get_api_base(api_base)
|
||||
api_key = VLLMModelInfo.get_api_key(api_key)
|
||||
endpoint: Final = "/v1/models"
|
||||
if api_base is None or api_key is None:
|
||||
raise ValueError(
|
||||
"VLLM_API_BASE or VLLM_API_KEY is not set. Please set the environment variable, to query VLLM's `/models` endpoint."
|
||||
)
|
||||
|
||||
url: Final = _add_path_to_api_base(api_base, endpoint)
|
||||
def _query_models(self, api_base: str | None, api_key: str | None) -> httpx.Response:
|
||||
resolved_api_base: Final = self._get_discovery_api_base(api_base)
|
||||
environment_variable: Final = "HOSTED_VLLM_API_KEY" if self._provider == "hosted_vllm" else "VLLM_API_KEY"
|
||||
resolved_api_key: Final = api_key or get_secret_str(environment_variable)
|
||||
headers: Final = {"Authorization": f"Bearer {resolved_api_key}"} if resolved_api_key else None
|
||||
response: Final = litellm.module_level_client.get(
|
||||
url=url,
|
||||
url=_add_path_to_api_base(resolved_api_base, "/v1/models"),
|
||||
headers=headers,
|
||||
follow_redirects=False,
|
||||
timeout=5.0,
|
||||
)
|
||||
|
||||
response.raise_for_status()
|
||||
return response
|
||||
|
||||
def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]:
|
||||
response: Final = self._query_models(api_base, api_key)
|
||||
models: Final = response.json()["data"]
|
||||
|
||||
return [model["id"] for model in models]
|
||||
|
||||
def get_model_info(
|
||||
self,
|
||||
model: str,
|
||||
api_base: str | None = None,
|
||||
api_key: str | None = None,
|
||||
) -> ModelInfoBase | None:
|
||||
response: Final = self._query_models(api_base, api_key)
|
||||
target: Final = model.removeprefix(f"{self._provider}/")
|
||||
discovered: Final = _VLLMModelsResponse.model_validate(response.json())
|
||||
for raw_entry in discovered.data:
|
||||
try:
|
||||
entry = _VLLMModelEntry.model_validate(raw_entry)
|
||||
except ValidationError:
|
||||
continue
|
||||
if entry.id == target and entry.max_model_len is not None:
|
||||
return ModelInfoBase(
|
||||
key=model,
|
||||
litellm_provider=self._provider,
|
||||
mode="chat",
|
||||
input_cost_per_token=0.0,
|
||||
output_cost_per_token=0.0,
|
||||
max_tokens=None,
|
||||
max_input_tokens=entry.max_model_len,
|
||||
max_output_tokens=None,
|
||||
)
|
||||
return None
|
||||
|
||||
def get_error_class(self, error_message: str, status_code: int, headers: dict | httpx.Headers) -> BaseLLMException:
|
||||
return VLLMError(status_code=status_code, message=error_message, headers=headers)
|
||||
|
|
|
|||
|
|
@ -9859,17 +9859,51 @@ def _pricing_override_stamps(
|
|||
|
||||
|
||||
def get_litellm_model_info(model: dict = {}):
|
||||
"""Model-info route enrichment; vLLM lookups perform I/O and must run off the event loop."""
|
||||
model_info: Final = model.get("model_info", {})
|
||||
model_to_lookup = model.get("litellm_params", {}).get("model", None)
|
||||
litellm_params: Final = model.get("litellm_params") or _EMPTY_MAPPING
|
||||
configured_model: Final = litellm_params.get("model", None)
|
||||
use_base_model: Final = (isinstance(configured_model, str) and "azure" in configured_model) or bool(
|
||||
model_info.get("base_model")
|
||||
)
|
||||
model_to_lookup: Final = model_info.get("base_model", None) if use_base_model else configured_model
|
||||
try:
|
||||
if "azure" in model_to_lookup or model_info.get("base_model"):
|
||||
model_to_lookup = model_info.get("base_model", None)
|
||||
litellm_model_info: Final = litellm.get_model_info(model_to_lookup)
|
||||
return litellm_model_info
|
||||
static_model_info = litellm.get_model_info(model_to_lookup)
|
||||
except Exception: # noqa: BLE001 # get_model_info wraps unmapped-model failures in bare Exception.
|
||||
static_model_info = _EMPTY_MAPPING
|
||||
|
||||
provider: Final = litellm_params.get("custom_llm_provider") or (
|
||||
configured_model.partition("/")[0] if isinstance(configured_model, str) else None
|
||||
)
|
||||
if provider not in ("vllm", "hosted_vllm"):
|
||||
return static_model_info
|
||||
try:
|
||||
credential_name: Final = litellm_params.get("litellm_credential_name")
|
||||
credentials: Final = (
|
||||
CredentialAccessor.get_credential_values(credential_name)
|
||||
if isinstance(credential_name, str)
|
||||
else _EMPTY_MAPPING
|
||||
)
|
||||
live_model_info: Final = litellm.get_model_info(
|
||||
configured_model,
|
||||
custom_llm_provider=provider,
|
||||
api_base=litellm_params.get("api_base") or credentials.get("api_base"),
|
||||
api_key=litellm_params.get("api_key") or credentials.get("api_key"),
|
||||
discover_model_info=True,
|
||||
)
|
||||
return {
|
||||
**live_model_info,
|
||||
**static_model_info,
|
||||
"max_input_tokens": (
|
||||
live_model_info["max_input_tokens"]
|
||||
if live_model_info["max_input_tokens"] is not None
|
||||
else static_model_info.get("max_input_tokens")
|
||||
),
|
||||
}
|
||||
except Exception:
|
||||
# this should not block returning on /model/info
|
||||
# if litellm does not have info on the model it should return {}
|
||||
return {}
|
||||
return static_model_info
|
||||
|
||||
|
||||
def on_backoff(details):
|
||||
|
|
@ -15229,12 +15263,17 @@ async def model_info_v2(
|
|||
|
||||
# Fill in model info based on config.yaml and litellm model_prices_and_context_window.json
|
||||
# This must happen before teamId filtering so that direct_access and access_via_team_ids are populated
|
||||
for i, _model in enumerate(all_models):
|
||||
all_models[i] = _enrich_model_info_with_litellm_data(
|
||||
model=_model,
|
||||
debug=debug if debug is not None else False,
|
||||
llm_router=llm_router,
|
||||
all_models = await asyncio.gather(
|
||||
*(
|
||||
asyncio.to_thread(
|
||||
_enrich_model_info_with_litellm_data,
|
||||
model=_model,
|
||||
debug=debug if debug is not None else False,
|
||||
llm_router=llm_router,
|
||||
)
|
||||
for _model in all_models
|
||||
)
|
||||
)
|
||||
|
||||
# Apply teamId filter if provided
|
||||
if teamId is not None and teamId.strip():
|
||||
|
|
@ -15946,7 +15985,9 @@ async def model_info_v1(
|
|||
status_code=400,
|
||||
detail={"error": f"Model id = {litellm_model_id} not found on litellm proxy"},
|
||||
)
|
||||
_deployment_info_dict = _get_proxy_model_info(model=deployment_info.model_dump(exclude_none=True))
|
||||
_deployment_info_dict = await asyncio.to_thread(
|
||||
_get_proxy_model_info, model=deployment_info.model_dump(exclude_none=True)
|
||||
)
|
||||
single_model_list: list[dict] = [_deployment_info_dict]
|
||||
if prisma_client is not None:
|
||||
single_model_list = await _populate_team_access_on_models(
|
||||
|
|
@ -16012,8 +16053,13 @@ async def model_info_v1(
|
|||
all_models = _filter_models_to_user_accessible(all_models)
|
||||
|
||||
all_models = [
|
||||
_translate_model_name_for_response(_enrich_model_info_with_litellm_data(model=model, llm_router=llm_router))
|
||||
for model in all_models
|
||||
_translate_model_name_for_response(model)
|
||||
for model in await asyncio.gather(
|
||||
*(
|
||||
asyncio.to_thread(_enrich_model_info_with_litellm_data, model=model, llm_router=llm_router)
|
||||
for model in all_models
|
||||
)
|
||||
)
|
||||
]
|
||||
|
||||
if teamId is not None and teamId.strip():
|
||||
|
|
|
|||
|
|
@ -5881,6 +5881,7 @@ def _get_model_info_helper(
|
|||
custom_llm_provider: str | None = None,
|
||||
api_base: str | None = None,
|
||||
api_key: str | None = None,
|
||||
discover_model_info: bool = False,
|
||||
) -> ModelInfoBase:
|
||||
"""
|
||||
Helper for 'get_model_info'. Separated out to avoid infinite loop caused by returning 'supported_openai_param's
|
||||
|
|
@ -5920,7 +5921,9 @@ def _get_model_info_helper(
|
|||
provider_config = ProviderConfigManager.get_provider_model_info(
|
||||
model=model, provider=LlmProviders(custom_llm_provider)
|
||||
)
|
||||
if provider_config is not None:
|
||||
dynamic_model_info: ModelInfoBase | None = None
|
||||
should_query_provider: Final = custom_llm_provider not in ("vllm", "hosted_vllm") or discover_model_info
|
||||
if provider_config is not None and should_query_provider:
|
||||
provider_get_model_info: Final = getattr(provider_config, "get_model_info", None)
|
||||
if callable(provider_get_model_info):
|
||||
try:
|
||||
|
|
@ -5930,14 +5933,16 @@ def _get_model_info_helper(
|
|||
api_key=api_key,
|
||||
)
|
||||
if provider_model_info is not None:
|
||||
return provider_model_info
|
||||
if custom_llm_provider not in ("vllm", "hosted_vllm"):
|
||||
return provider_model_info
|
||||
dynamic_model_info = provider_model_info
|
||||
except Exception as e:
|
||||
verbose_logger.warning(
|
||||
"Could not get dynamic model info for model=%s, provider=%s; "
|
||||
"falling back to the static cost map: %s",
|
||||
model,
|
||||
custom_llm_provider,
|
||||
e,
|
||||
type(e).__name__ if custom_llm_provider in ("vllm", "hosted_vllm") else e,
|
||||
)
|
||||
|
||||
if custom_llm_provider == "huggingface":
|
||||
|
|
@ -6058,6 +6063,8 @@ def _get_model_info_helper(
|
|||
key, _model_info = generalization
|
||||
|
||||
if _model_info is None or key is None:
|
||||
if dynamic_model_info is not None:
|
||||
return dynamic_model_info
|
||||
raise ModelNotMappedError(_model_not_mapped_message(model, custom_llm_provider))
|
||||
_input_cost_per_token: float | None = _model_info.get("input_cost_per_token")
|
||||
if _input_cost_per_token is None:
|
||||
|
|
@ -6307,6 +6314,12 @@ def _get_model_info_helper(
|
|||
for cost_key, cost_value in _model_info.items():
|
||||
if cost_key not in returned_model_info and _ABOVE_THRESHOLD_COST_KEY.search(cost_key) is not None:
|
||||
returned_model_info[cost_key] = cost_value
|
||||
if dynamic_model_info is not None and dynamic_model_info.get("max_input_tokens") is not None:
|
||||
merged_model_info: Final[ModelInfoBase] = {
|
||||
**returned_model_info,
|
||||
"max_input_tokens": dynamic_model_info["max_input_tokens"],
|
||||
}
|
||||
return merged_model_info
|
||||
return returned_model_info
|
||||
except ModelNotMappedError:
|
||||
raise
|
||||
|
|
@ -6320,6 +6333,7 @@ def _build_model_info(
|
|||
custom_llm_provider: str | None = None,
|
||||
api_base: str | None = None,
|
||||
api_key: str | None = None,
|
||||
discover_model_info: bool = False,
|
||||
) -> ModelInfo:
|
||||
supported_openai_params = litellm.get_supported_openai_params(model=model, custom_llm_provider=custom_llm_provider)
|
||||
|
||||
|
|
@ -6328,6 +6342,7 @@ def _build_model_info(
|
|||
custom_llm_provider=custom_llm_provider,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
discover_model_info=discover_model_info,
|
||||
)
|
||||
|
||||
provider_info: Final = get_provider_info(model=model, custom_llm_provider=custom_llm_provider)
|
||||
|
|
@ -6356,6 +6371,7 @@ def get_model_info(
|
|||
custom_llm_provider: str | None = None,
|
||||
api_base: str | None = None,
|
||||
api_key: str | None = None,
|
||||
discover_model_info: bool = False,
|
||||
) -> ModelInfo:
|
||||
"""
|
||||
Get a dict for the maximum tokens (context window), input_cost_per_token, output_cost_per_token for a given model.
|
||||
|
|
@ -6363,6 +6379,10 @@ def get_model_info(
|
|||
Parameters:
|
||||
- model (str): The name of the model.
|
||||
- custom_llm_provider (str | null): the provider used for the model. If provided, used to check if the litellm model info is for that provider.
|
||||
- api_base (str | null): the deployment endpoint used for provider-scoped discovery.
|
||||
- api_key (str | null): the deployment credential used for provider-scoped discovery.
|
||||
- discover_model_info (bool): opt in to a synchronous, uncached vLLM metadata lookup; defaults to False.
|
||||
Discovery overlays context only, not output limits.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary containing the following information:
|
||||
|
|
@ -6428,10 +6448,16 @@ def get_model_info(
|
|||
"supported_openai_params": ["temperature", "max_tokens", "top_p", "frequency_penalty", "presence_penalty"]
|
||||
}
|
||||
"""
|
||||
# api_key is a per-caller credential, not part of the model identity, so it is
|
||||
# kept out of the cache key; explicit keys are resolved without the cache.
|
||||
if api_key is not None:
|
||||
return _build_model_info(model, custom_llm_provider, api_base, api_key)
|
||||
# Credentials are per caller, and live discovery must observe endpoint changes,
|
||||
# so neither path can safely reuse the static metadata cache.
|
||||
if api_key is not None or discover_model_info:
|
||||
return _build_model_info(
|
||||
model,
|
||||
custom_llm_provider,
|
||||
api_base,
|
||||
api_key,
|
||||
discover_model_info,
|
||||
)
|
||||
return _cached_get_model_info(model, custom_llm_provider, api_base)
|
||||
|
||||
|
||||
|
|
@ -9242,7 +9268,7 @@ class ProviderConfigManager:
|
|||
VLLMModelInfo, # experimental approach, to reduce bloat on __init__.py
|
||||
)
|
||||
|
||||
return VLLMModelInfo()
|
||||
return VLLMModelInfo(provider="hosted_vllm" if provider == LlmProviders.HOSTED_VLLM else "vllm")
|
||||
elif LlmProviders.LEMONADE == provider:
|
||||
return litellm.LemonadeChatConfig()
|
||||
elif LlmProviders.CLARIFAI == provider:
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ Pins (PR2):
|
|||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
import copy
|
||||
from collections.abc import Callable
|
||||
from contextlib import AbstractContextManager
|
||||
|
|
@ -286,6 +287,172 @@ def test_get_proxy_model_info_surfaces_supports_parallel_function_calling(local_
|
|||
assert enriched["model_info"]["supports_parallel_function_calling"] is True
|
||||
|
||||
|
||||
def test_get_proxy_model_info_discovers_vllm_context_with_config_precedence(
|
||||
monkeypatch,
|
||||
):
|
||||
response = MagicMock()
|
||||
response.json.return_value = {"data": [{"id": "shared", "max_model_len": 262_144}]}
|
||||
request = MagicMock(return_value=response)
|
||||
monkeypatch.setattr(proxy_server.litellm.module_level_client, "get", request)
|
||||
|
||||
enriched = proxy_server._get_proxy_model_info(
|
||||
model={
|
||||
"model_name": "vllm-model",
|
||||
"litellm_params": {
|
||||
"model": "hosted_vllm/shared",
|
||||
"api_base": "https://vllm.example/v1",
|
||||
"api_key": "endpoint-secret",
|
||||
},
|
||||
"model_info": {
|
||||
"id": "vllm-deployment",
|
||||
"max_input_tokens": 200_000,
|
||||
"max_output_tokens": 32_768,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
assert enriched["model_info"]["max_input_tokens"] == 200_000
|
||||
assert enriched["model_info"]["max_output_tokens"] == 32_768
|
||||
assert enriched["model_info"]["max_tokens"] is None
|
||||
assert "api_key" not in enriched["litellm_params"]
|
||||
assert request.call_args.kwargs["headers"] == {"Authorization": "Bearer endpoint-secret"}
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"path,params",
|
||||
[
|
||||
("/v1/model/info", {"litellm_model_id": "vllm-route-deployment"}),
|
||||
("/v1/model/info", {}),
|
||||
("/v2/model/info", {}),
|
||||
],
|
||||
)
|
||||
@pytest.mark.parametrize("explicit_context", [None, 200_000])
|
||||
@pytest.mark.parametrize("named_credential", [False, True])
|
||||
def test_model_info_routes_refresh_discovered_context_below_explicit_config(
|
||||
client,
|
||||
auth_as,
|
||||
monkeypatch,
|
||||
path,
|
||||
params,
|
||||
explicit_context,
|
||||
named_credential,
|
||||
local_model_cost_map,
|
||||
mock_prisma,
|
||||
):
|
||||
response = MagicMock()
|
||||
response.json.return_value = {"data": [{"id": "shared", "max_model_len": 262_144}]}
|
||||
request = MagicMock(return_value=response)
|
||||
monkeypatch.setattr(litellm.module_level_client, "get", request)
|
||||
model_list = [
|
||||
{
|
||||
"model_name": "vllm-model",
|
||||
"litellm_params": {
|
||||
"model": "hosted_vllm/shared",
|
||||
"api_base": "https://vllm.example/v1",
|
||||
"api_key": "endpoint-secret",
|
||||
},
|
||||
"model_info": {
|
||||
"id": "vllm-route-deployment",
|
||||
"base_model": "gpt-4o",
|
||||
**({"max_input_tokens": explicit_context} if explicit_context else {}),
|
||||
},
|
||||
}
|
||||
]
|
||||
if named_credential:
|
||||
from litellm.types.utils import CredentialItem
|
||||
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"credential_list",
|
||||
[
|
||||
CredentialItem(
|
||||
credential_name="vllm-credential",
|
||||
credential_values={"api_base": "https://vllm.example/v1", "api_key": "endpoint-secret"},
|
||||
credential_info={},
|
||||
)
|
||||
],
|
||||
)
|
||||
model_list[0]["litellm_params"] = {
|
||||
"model": "hosted_vllm/shared",
|
||||
"api_base": "https://overridden.example/v1",
|
||||
"litellm_credential_name": "vllm-credential",
|
||||
}
|
||||
router = litellm.Router(model_list=model_list)
|
||||
from litellm.proxy.auth.auth_checks import model_has_no_cost_mapping
|
||||
|
||||
router.get_model_group_info(model_group="vllm-model")
|
||||
model_has_no_cost_mapping(model="vllm-model", llm_router=router)
|
||||
request.assert_not_called()
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
monkeypatch.setattr(proxy_server, "llm_model_list", model_list)
|
||||
monkeypatch.setattr(proxy_server, "user_model", None)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", mock_prisma if path == "/v2/model/info" else None)
|
||||
monkeypatch.setattr(proxy_server.proxy_config, "get_config", AsyncMock(return_value={}))
|
||||
|
||||
base = litellm.get_model_info("gpt-4o")
|
||||
for context in (262_144, 65_536, None):
|
||||
response.json.return_value = {"data": [{"id": "shared", "max_model_len": context}] if context else []}
|
||||
with auth_as():
|
||||
result = client.get(path, params=params)
|
||||
assert result.status_code == 200, result.text
|
||||
deployment = result.json()["data"][0]
|
||||
assert deployment["model_info"]["max_input_tokens"] == (explicit_context or context or base["max_input_tokens"])
|
||||
assert deployment["model_info"]["max_output_tokens"] == base["max_output_tokens"]
|
||||
assert deployment["model_info"]["input_cost_per_token"] == base["input_cost_per_token"]
|
||||
assert "api_key" not in deployment["litellm_params"]
|
||||
assert "endpoint-secret" not in result.text
|
||||
assert request.call_count == 3
|
||||
chat_api_base = model_list[0]["litellm_params"]["api_base"]
|
||||
for call in request.call_args_list:
|
||||
assert call.kwargs["url"] == f"{chat_api_base}/models"
|
||||
assert call.kwargs["headers"] == {"Authorization": "Bearer endpoint-secret"}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_model_info_discovery_runs_outside_the_event_loop(app, auth_as, configured_router, monkeypatch):
|
||||
event_loop_thread = threading.get_ident()
|
||||
lookup_threads: list[int] = []
|
||||
|
||||
def lookup(model):
|
||||
lookup_threads.append(threading.get_ident())
|
||||
return model
|
||||
|
||||
monkeypatch.setattr(proxy_server, "_get_proxy_model_info", lookup)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", None)
|
||||
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://test") as client:
|
||||
with auth_as():
|
||||
response = await client.get("/v1/model/info", params={"litellm_model_id": "abc"})
|
||||
assert response.status_code == 200
|
||||
assert len(lookup_threads) == 1
|
||||
assert lookup_threads[0] != event_loop_thread
|
||||
|
||||
|
||||
@pytest.mark.parametrize("path", ["/v1/model/info", "/v2/model/info"])
|
||||
def test_model_info_list_routes_enrich_deployments_concurrently(client, auth_as, monkeypatch, mock_prisma, path):
|
||||
model_list = [
|
||||
{"model_name": name, "litellm_params": {"model": f"hosted_vllm/{name}"}, "model_info": {"id": name}}
|
||||
for name in ("slow-a", "slow-b")
|
||||
]
|
||||
both_lookups_started = threading.Barrier(2, timeout=5)
|
||||
|
||||
def enrich(model, **kwargs):
|
||||
both_lookups_started.wait()
|
||||
return model
|
||||
|
||||
monkeypatch.setattr(proxy_server, "_enrich_model_info_with_litellm_data", enrich)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", litellm.Router(model_list=model_list))
|
||||
monkeypatch.setattr(proxy_server, "llm_model_list", model_list)
|
||||
monkeypatch.setattr(proxy_server, "user_model", None)
|
||||
monkeypatch.setattr(proxy_server, "prisma_client", mock_prisma if path == "/v2/model/info" else None)
|
||||
monkeypatch.setattr(proxy_server.proxy_config, "get_config", AsyncMock(return_value={}))
|
||||
|
||||
with auth_as():
|
||||
result = client.get(path)
|
||||
|
||||
assert result.status_code == 200, result.text
|
||||
assert [deployment["model_info"]["id"] for deployment in result.json()["data"]] == ["slow-a", "slow-b"]
|
||||
|
||||
|
||||
def _enriched_model_info(monkeypatch, litellm_params: dict, model_info: dict) -> dict:
|
||||
monkeypatch.setattr(proxy_server, "llm_router", None)
|
||||
enriched: Final = proxy_server._get_proxy_model_info(
|
||||
|
|
|
|||
|
|
@ -1,12 +1,18 @@
|
|||
import json
|
||||
import logging
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.constants import (
|
||||
DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET,
|
||||
DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import HTTPHandler
|
||||
from litellm.llms.hosted_vllm.chat.transformation import HostedVLLMChatConfig
|
||||
from litellm.llms.vllm.common_utils import VLLMModelInfo
|
||||
|
||||
|
||||
def test_hosted_vllm_chat_transformation_file_url():
|
||||
|
|
@ -366,3 +372,271 @@ def test_hosted_vllm_custom_tools_use_top_level_input_schema():
|
|||
assert tools[0]["function"]["name"] == "search"
|
||||
assert tools[0]["function"]["description"] == "Search docs"
|
||||
assert tools[0]["function"]["parameters"] == input_schema
|
||||
|
||||
|
||||
def _model_list_response(*entries: dict[str, object]) -> MagicMock:
|
||||
response = MagicMock()
|
||||
response.json.return_value = {"data": list(entries)}
|
||||
return response
|
||||
|
||||
|
||||
def test_vllm_model_info_maps_context_only_and_authenticates(monkeypatch) -> None:
|
||||
request = MagicMock(return_value=_model_list_response({"id": "Qwen/Qwen3-8B", "max_model_len": 262_144}))
|
||||
monkeypatch.setattr(litellm.module_level_client, "get", request)
|
||||
|
||||
info = VLLMModelInfo(provider="hosted_vllm").get_model_info(
|
||||
model="hosted_vllm/Qwen/Qwen3-8B",
|
||||
api_base="https://vllm.example/v1",
|
||||
api_key="secret-key",
|
||||
)
|
||||
|
||||
assert info is not None
|
||||
assert info["max_input_tokens"] == 262_144
|
||||
assert info["max_tokens"] is None
|
||||
assert info["max_output_tokens"] is None
|
||||
request.assert_called_once()
|
||||
assert request.call_args.kwargs["url"] == "https://vllm.example/v1/models"
|
||||
assert request.call_args.kwargs["headers"] == {"Authorization": "Bearer secret-key"}
|
||||
|
||||
|
||||
def test_vllm_model_info_does_not_discover_without_opt_in(monkeypatch) -> None:
|
||||
request = MagicMock()
|
||||
monkeypatch.setattr(litellm.module_level_client, "get", request)
|
||||
|
||||
with pytest.raises(Exception, match="isn't mapped yet"):
|
||||
litellm.get_model_info(
|
||||
"hosted_vllm/not-in-static-map",
|
||||
api_base="https://no-discovery.example/v1",
|
||||
)
|
||||
|
||||
request.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.parametrize("redirect", [False, True])
|
||||
def test_vllm_discovery_http_transport_bounds_and_credentials(monkeypatch, redirect) -> None:
|
||||
requests: list[httpx.Request] = []
|
||||
|
||||
def serve(request: httpx.Request) -> httpx.Response:
|
||||
requests.append(request)
|
||||
if redirect:
|
||||
return httpx.Response(302, headers={"Location": "https://other.example/models"})
|
||||
return httpx.Response(200, json={"data": [{"id": "shared", "max_model_len": 262144}]})
|
||||
|
||||
with httpx.Client(transport=httpx.MockTransport(serve), follow_redirects=True) as client:
|
||||
monkeypatch.setattr(litellm, "module_level_client", HTTPHandler(client=client))
|
||||
provider = VLLMModelInfo(provider="hosted_vllm")
|
||||
if redirect:
|
||||
with pytest.raises(httpx.HTTPStatusError):
|
||||
provider.get_model_info("hosted_vllm/shared", api_base="https://vllm.example/v1", api_key="key")
|
||||
else:
|
||||
info = provider.get_model_info("hosted_vllm/shared", api_base="https://vllm.example/v1", api_key="key")
|
||||
assert info is not None and info["max_input_tokens"] == 262144
|
||||
|
||||
assert len(requests) == 1
|
||||
assert str(requests[0].url) == "https://vllm.example/v1/models"
|
||||
assert requests[0].headers["Authorization"] == "Bearer key"
|
||||
assert set(requests[0].extensions["timeout"].values()) == {5.0}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("value", [None, 0, -1, True, "262144", 262144.0, 1.5])
|
||||
def test_vllm_model_info_ignores_non_positive_or_non_integer_context(
|
||||
monkeypatch,
|
||||
value: object,
|
||||
) -> None:
|
||||
monkeypatch.setattr(
|
||||
litellm.module_level_client,
|
||||
"get",
|
||||
MagicMock(return_value=_model_list_response({"id": "model", "max_model_len": value})),
|
||||
)
|
||||
|
||||
assert (
|
||||
VLLMModelInfo(provider="hosted_vllm").get_model_info(
|
||||
model="hosted_vllm/model",
|
||||
api_base="https://vllm.example/v1",
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_vllm_explicit_base_uses_the_same_environment_key_as_chat(monkeypatch) -> None:
|
||||
request = MagicMock(return_value=_model_list_response({"id": "model"}))
|
||||
monkeypatch.setenv("HOSTED_VLLM_API_KEY", "env-secret")
|
||||
monkeypatch.setattr(litellm.module_level_client, "get", request)
|
||||
api_base = "https://operator-supplied.example/v1"
|
||||
|
||||
VLLMModelInfo(provider="hosted_vllm").get_model_info(model="hosted_vllm/model", api_base=api_base)
|
||||
|
||||
_, chat_api_key = HostedVLLMChatConfig()._get_openai_compatible_provider_info(api_base=api_base, api_key=None)
|
||||
assert request.call_args.kwargs["headers"] == {"Authorization": f"Bearer {chat_api_key}"}
|
||||
|
||||
|
||||
def test_hosted_vllm_model_info_uses_provider_environment(monkeypatch) -> None:
|
||||
request = MagicMock(return_value=_model_list_response({"id": "env-model", "max_model_len": 32_768}))
|
||||
monkeypatch.setenv("HOSTED_VLLM_API_BASE", "https://hosted.example/v1")
|
||||
monkeypatch.setenv("HOSTED_VLLM_API_KEY", "hosted-secret")
|
||||
monkeypatch.setattr(litellm.module_level_client, "get", request)
|
||||
|
||||
info = litellm.get_model_info("hosted_vllm/env-model", discover_model_info=True)
|
||||
|
||||
assert info["litellm_provider"] == "hosted_vllm"
|
||||
assert info["max_input_tokens"] == 32_768
|
||||
request.assert_called_once()
|
||||
assert request.call_args.kwargs["url"] == "https://hosted.example/v1/models"
|
||||
assert request.call_args.kwargs["headers"] == {"Authorization": "Bearer hosted-secret"}
|
||||
|
||||
|
||||
def test_bare_vllm_model_keeps_explicit_provider_identity(monkeypatch) -> None:
|
||||
monkeypatch.setattr(
|
||||
litellm.module_level_client,
|
||||
"get",
|
||||
MagicMock(return_value=_model_list_response({"id": "shared", "max_model_len": 65_536})),
|
||||
)
|
||||
|
||||
info = litellm.get_model_info(
|
||||
"shared",
|
||||
custom_llm_provider="vllm",
|
||||
api_base="https://vllm.example/v1",
|
||||
discover_model_info=True,
|
||||
)
|
||||
|
||||
assert info["litellm_provider"] == "vllm"
|
||||
assert info["max_input_tokens"] == 65_536
|
||||
|
||||
|
||||
def test_vllm_endpoint_scoped_lookup_refreshes_and_removes_without_global_registration(
|
||||
monkeypatch,
|
||||
) -> None:
|
||||
responses = [
|
||||
_model_list_response({"id": "shared", "max_model_len": 32_768}),
|
||||
_model_list_response({"id": "shared", "max_model_len": 65_536}),
|
||||
_model_list_response(),
|
||||
]
|
||||
request = MagicMock(side_effect=responses)
|
||||
monkeypatch.setattr(litellm.module_level_client, "get", request)
|
||||
key = "hosted_vllm/shared"
|
||||
original_entry = litellm.model_cost.get(key)
|
||||
|
||||
first = litellm.get_model_info(
|
||||
key,
|
||||
api_base="https://vllm-a.example/v1",
|
||||
api_key="key-a",
|
||||
discover_model_info=True,
|
||||
)
|
||||
second = litellm.get_model_info(
|
||||
key,
|
||||
api_base="https://vllm-a.example/v1",
|
||||
api_key="key-a",
|
||||
discover_model_info=True,
|
||||
)
|
||||
with pytest.raises(Exception, match="isn't mapped yet"):
|
||||
litellm.get_model_info(
|
||||
key,
|
||||
api_base="https://vllm-a.example/v1",
|
||||
api_key="key-a",
|
||||
discover_model_info=True,
|
||||
)
|
||||
|
||||
assert first["max_input_tokens"] == 32_768
|
||||
assert second["max_input_tokens"] == 65_536
|
||||
assert litellm.model_cost.get(key) is original_entry
|
||||
|
||||
|
||||
def test_vllm_same_model_id_is_isolated_by_endpoint(monkeypatch) -> None:
|
||||
def get(url: str, headers: dict[str, str], **kwargs) -> MagicMock:
|
||||
del headers
|
||||
context = 32_768 if "vllm-a" in url else 131_072
|
||||
return _model_list_response({"id": "shared", "max_model_len": context})
|
||||
|
||||
monkeypatch.setattr(litellm.module_level_client, "get", get)
|
||||
|
||||
info_a = litellm.get_model_info(
|
||||
"hosted_vllm/shared",
|
||||
api_base="https://vllm-a.example/v1",
|
||||
api_key="key-a",
|
||||
discover_model_info=True,
|
||||
)
|
||||
info_b = litellm.get_model_info(
|
||||
"hosted_vllm/shared",
|
||||
api_base="https://vllm-b.example/v1",
|
||||
api_key="key-b",
|
||||
discover_model_info=True,
|
||||
)
|
||||
|
||||
assert info_a["max_input_tokens"] == 32_768
|
||||
assert info_b["max_input_tokens"] == 131_072
|
||||
|
||||
|
||||
def test_vllm_discovery_failure_does_not_log_the_api_key(monkeypatch, caplog) -> None:
|
||||
secret = "do-not-log-this-key"
|
||||
monkeypatch.setattr(
|
||||
litellm.module_level_client,
|
||||
"get",
|
||||
MagicMock(return_value=MagicMock(json=lambda: {"error": {"authorization": secret}})),
|
||||
)
|
||||
|
||||
with caplog.at_level(logging.WARNING), pytest.raises(Exception, match="isn't mapped yet"):
|
||||
litellm.get_model_info(
|
||||
"hosted_vllm/unmapped",
|
||||
api_base="https://vllm.example/v1",
|
||||
api_key=secret,
|
||||
discover_model_info=True,
|
||||
)
|
||||
|
||||
assert secret not in caplog.text
|
||||
|
||||
|
||||
def test_vllm_endpoint_discovery_survives_price_map_replacement(monkeypatch) -> None:
|
||||
monkeypatch.setattr(litellm, "model_cost", {})
|
||||
monkeypatch.setattr(
|
||||
litellm.module_level_client,
|
||||
"get",
|
||||
MagicMock(return_value=_model_list_response({"id": "model", "max_model_len": 98_304})),
|
||||
)
|
||||
|
||||
before = litellm.get_model_info("hosted_vllm/model", api_base="https://vllm.example/v1", discover_model_info=True)
|
||||
monkeypatch.setattr(litellm, "model_cost", {"unrelated": {"litellm_provider": "openai"}})
|
||||
info = litellm.get_model_info(
|
||||
"hosted_vllm/model",
|
||||
api_base="https://vllm.example/v1",
|
||||
api_key="key",
|
||||
discover_model_info=True,
|
||||
)
|
||||
|
||||
assert info["max_input_tokens"] == 98_304
|
||||
assert before["max_input_tokens"] == info["max_input_tokens"]
|
||||
assert "hosted_vllm/model" not in litellm.model_cost
|
||||
|
||||
|
||||
def test_vllm_discovery_preserves_static_pricing(monkeypatch) -> None:
|
||||
from litellm.utils import _invalidate_model_cost_lowercase_map
|
||||
|
||||
static_info = {
|
||||
"litellm_provider": "hosted_vllm",
|
||||
"mode": "chat",
|
||||
"max_input_tokens": 4_096,
|
||||
"max_output_tokens": 1_024,
|
||||
"supports_vision": True,
|
||||
"input_cost_per_token": 0.25,
|
||||
"output_cost_per_token": 0.5,
|
||||
}
|
||||
with monkeypatch.context() as scoped:
|
||||
scoped.setattr(litellm, "model_cost", {"hosted_vllm/priced": static_info})
|
||||
scoped.setattr(
|
||||
litellm.module_level_client,
|
||||
"get",
|
||||
MagicMock(return_value=_model_list_response({"id": "priced", "max_model_len": 131_072})),
|
||||
)
|
||||
_invalidate_model_cost_lowercase_map()
|
||||
|
||||
info = litellm.get_model_info(
|
||||
"hosted_vllm/priced",
|
||||
api_base="https://vllm.example/v1",
|
||||
discover_model_info=True,
|
||||
)
|
||||
|
||||
assert info["max_input_tokens"] == 131_072
|
||||
assert info["input_cost_per_token"] == 0.25
|
||||
assert info["output_cost_per_token"] == 0.5
|
||||
assert info["max_output_tokens"] == 1_024
|
||||
assert info["supports_vision"] is True
|
||||
_invalidate_model_cost_lowercase_map()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue