diff --git a/litellm/llms/vllm/common_utils.py b/litellm/llms/vllm/common_utils.py index 2da2269e1f9..576931582fb 100644 --- a/litellm/llms/vllm/common_utils.py +++ b/litellm/llms/vllm/common_utils.py @@ -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) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7cd116c23d7..4b86981723b 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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(): diff --git a/litellm/utils.py b/litellm/utils.py index 13a46840431..31b8959b5d0 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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: diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py b/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py index 5175d92084c..b69249a82a6 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_model_info.py @@ -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( diff --git a/tests/unit/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py b/tests/unit/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py index 1cc6a1457fc..5a8f066f61f 100644 --- a/tests/unit/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py +++ b/tests/unit/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py @@ -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()