From 1599abead24116ae5a8ab232b3485e04d0901a10 Mon Sep 17 00:00:00 2001 From: Aiden McComiskey Date: Fri, 4 Sep 2026 16:56:53 -0400 Subject: [PATCH 1/4] fix(hosted-vllm): discover context limits --- litellm/llms/custom_httpx/http_handler.py | 4 +- litellm/llms/vllm/common_utils.py | 116 ++++++++- litellm/proxy/proxy_server.py | 17 +- litellm/router.py | 32 ++- litellm/utils.py | 39 ++- .../test_hosted_vllm_chat_transformation.py | 241 ++++++++++++++++++ .../proxy_server/test_routes_model_info.py | 71 ++++++ .../test_router_model_cost_isolation.py | 49 +++- 8 files changed, 544 insertions(+), 25 deletions(-) diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index e1f0fc9e7d3..129453fe827 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -651,7 +651,7 @@ class AsyncHTTPHandler: self, url: str, params: dict | None = None, - headers: dict | None = None, + headers: dict | httpx.Headers | None = None, follow_redirects: bool | None = None, timeout: float | httpx.Timeout | None = None, ): @@ -1289,7 +1289,7 @@ class HTTPHandler: self, url: str, params: dict | None = None, - headers: dict | None = None, + headers: dict | httpx.Headers | None = None, follow_redirects: bool | None = None, timeout: float | httpx.Timeout | None = None, ): diff --git a/litellm/llms/vllm/common_utils.py b/litellm/llms/vllm/common_utils.py index 2da2269e1f9..40f8f72a2f7 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 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: str = "vllm") -> None: + self._provider: Final = provider + def validate_environment( self, headers: dict, @@ -56,22 +74,34 @@ 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 + + def _get_discovery_api_key(self, api_key: str | None) -> str | None: + environment_variable: Final = "HOSTED_VLLM_API_KEY" if self._provider == "hosted_vllm" else "VLLM_API_KEY" + return api_key or get_secret_str(environment_variable) + @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) + passed_api_base: Final = api_base + resolved_api_base: Final = self._get_discovery_api_base(api_base) + resolved_api_key: Final = self._get_discovery_api_key(api_key) if passed_api_base is None or api_key else None 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) + url: Final = _add_path_to_api_base(resolved_api_base, endpoint) + headers: Final = ( + httpx.Headers((("Authorization", f"Bearer {resolved_api_key}"),)) if resolved_api_key else httpx.Headers() + ) response: Final = litellm.module_level_client.get( url=url, + headers=headers, ) response.raise_for_status() @@ -80,5 +110,75 @@ class VLLMModelInfo(BaseLLMModelInfo): return [model["id"] for model in models] + @staticmethod + def _strip_provider_prefix(model: str) -> str: + for prefix in ("hosted_vllm/", "vllm/"): + if model.startswith(prefix): + return model[len(prefix) :] + return model + + def _model_info_from_entry( + self, + raw_entry: object, + *, + target: str, + model: str, + ) -> ModelInfoBase | None: + try: + entry: Final = _VLLMModelEntry.model_validate(raw_entry) + except ValidationError: + return None + if entry.id != target or entry.max_model_len is None: + return None + + model_info: Final[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 model_info + + def get_model_info( + self, + model: str, + api_base: str | None = None, + api_key: str | None = None, + ) -> ModelInfoBase | None: + passed_api_base: Final = api_base + resolved_api_base: Final = self._get_discovery_api_base(api_base) + resolved_api_key: Final = self._get_discovery_api_key(api_key) if passed_api_base is None or api_key else None + + headers: Final = ( + httpx.Headers((("Authorization", f"Bearer {resolved_api_key}"),)) if resolved_api_key else httpx.Headers() + ) + response: Final = litellm.module_level_client.get( + url=_add_path_to_api_base(resolved_api_base, "/v1/models"), + headers=headers, + ) + response.raise_for_status() + + target: Final = self._strip_provider_prefix(model) + discovered: Final = _VLLMModelsResponse.model_validate(response.json()) + return next( + ( + model_info + for raw_entry in discovered.data + if ( + model_info := self._model_info_from_entry( + raw_entry, + target=target, + model=model, + ) + ) + is not None + ), + 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 5a39b8c610a..3a273520624 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -8917,11 +8917,20 @@ def select_data_generator( def get_litellm_model_info(model: dict = {}): 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) + litellm_model_info: Final = litellm.get_model_info( + model_to_lookup, + custom_llm_provider=litellm_params.get("custom_llm_provider"), + api_base=litellm_params.get("api_base"), + api_key=litellm_params.get("api_key"), + discover_model_info=True, + ) return litellm_model_info except Exception: # this should not block returning on /model/info diff --git a/litellm/router.py b/litellm/router.py index 6da201725b6..c7277581698 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -10321,7 +10321,12 @@ class Router: model_name: Final = model_info["model_name"] return self.get_model_list(model_name=model_name) - def get_deployment_model_info(self, model_id: str, model_name: str) -> ModelInfo | None: + def get_deployment_model_info( + self, + model_id: str, + model_name: str, + litellm_params: LiteLLM_Params | None = None, + ) -> ModelInfo | None: """ For a given model id, return the model info @@ -10340,8 +10345,25 @@ class Router: except Exception: pass + deployment: Final = self.get_deployment(model_id=model_id) if litellm_params is None else None + resolved_litellm_params: Final = ( + litellm_params + if litellm_params is not None + else deployment.litellm_params + if deployment is not None + else None + ) + try: - litellm_model_name_model_info = litellm.get_model_info(model=model_name) + litellm_model_name_model_info = litellm.get_model_info( + model=model_name, + custom_llm_provider=( + resolved_litellm_params.custom_llm_provider if resolved_litellm_params is not None else None + ), + api_base=resolved_litellm_params.api_base if resolved_litellm_params is not None else None, + api_key=resolved_litellm_params.api_key if resolved_litellm_params is not None else None, + discover_model_info=True, + ) except Exception: pass @@ -10457,7 +10479,11 @@ class Router: try: model_id = model_info_dict.get("id", None) if model_id is not None: - model_info = self.get_deployment_model_info(model_id=model_id, model_name=litellm_params.model) + model_info = self.get_deployment_model_info( + model_id=model_id, + model_name=litellm_params.model, + litellm_params=litellm_params, + ) else: model_info = None except Exception: diff --git a/litellm/utils.py b/litellm/utils.py index 8b1b32ea328..bb1d22cdfb6 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -5575,6 +5575,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 @@ -5614,7 +5615,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: @@ -5624,7 +5627,9 @@ 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; " @@ -5743,6 +5748,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 ValueError( "This model isn't mapped yet. Add it here - https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json" ) @@ -5950,6 +5957,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 Exception as e: verbose_logger.debug("Error getting model info: %s", e) @@ -5963,6 +5976,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) @@ -5971,6 +5985,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) @@ -5999,6 +6014,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. @@ -6006,6 +6022,9 @@ 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): query supported provider metadata endpoints before falling back to the static map. Returns: dict: A dictionary containing the following information: @@ -6071,10 +6090,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) @@ -8787,7 +8812,7 @@ class ProviderConfigManager: VLLMModelInfo, # experimental approach, to reduce bloat on __init__.py ) - return VLLMModelInfo() + return VLLMModelInfo(provider=provider.value) elif LlmProviders.LEMONADE == provider: return litellm.LemonadeChatConfig() elif LlmProviders.CLARIFAI == provider: diff --git a/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py b/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py index 82b05601a85..26317420b90 100644 --- a/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py +++ b/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py @@ -1,12 +1,16 @@ import json +import logging from unittest.mock import MagicMock, patch +import pytest +import litellm from litellm.constants import ( DEFAULT_REASONING_EFFORT_HIGH_THINKING_BUDGET, DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET, ) 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(): @@ -433,3 +437,240 @@ 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 dict(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("value", [None, 0, -1, True, "262144"]) +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_never_receives_an_ambient_key(monkeypatch) -> None: + request = MagicMock(return_value=_model_list_response({"id": "model"})) + monkeypatch.setenv("HOSTED_VLLM_API_KEY", "ambient-secret") + monkeypatch.setattr(litellm.module_level_client, "get", request) + + VLLMModelInfo(provider="hosted_vllm").get_model_info( + model="hosted_vllm/model", + api_base="https://operator-supplied.example/v1", + ) + + assert dict(request.call_args.kwargs["headers"]) == {} + + +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 dict(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]) -> 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(side_effect=RuntimeError("discovery failed")), + ) + + 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})), + ) + + 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 "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, + "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 + _invalidate_model_cost_lowercase_map() 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 cb38e7edbe2..5f5b10b1fbc 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 @@ -13,6 +13,7 @@ from unittest.mock import MagicMock import pytest +import litellm from litellm.proxy import proxy_server from .conftest import normalize # type: ignore[import-not-found] @@ -142,6 +143,76 @@ 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 dict(request.call_args.kwargs["headers"]) == {"authorization": "Bearer endpoint-secret"} + + +def test_v1_model_info_route_surfaces_discovered_vllm_context( + client, + auth_as, + monkeypatch, +): + 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"}, + } + ] + router = litellm.Router(model_list=model_list) + 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", None) + + with auth_as(): + result = client.get( + "/v1/model/info", + params={"litellm_model_id": "vllm-route-deployment"}, + ) + + assert result.status_code == 200 + deployment = result.json()["data"][0] + assert deployment["model_info"]["max_input_tokens"] == 262_144 + assert deployment["model_info"]["max_output_tokens"] is None + assert "api_key" not in deployment["litellm_params"] + + def test_v1_model_info_star_wildcard_filter_keeps_provider_expansion(monkeypatch): from litellm.proxy._types import SpecialModelNames, UserAPIKeyAuth from litellm.proxy.auth import model_checks diff --git a/tests/test_litellm/test_router_model_cost_isolation.py b/tests/test_litellm/test_router_model_cost_isolation.py index 30b265905f3..c510ed2436c 100644 --- a/tests/test_litellm/test_router_model_cost_isolation.py +++ b/tests/test_litellm/test_router_model_cost_isolation.py @@ -11,7 +11,7 @@ import copy import logging import os import re -from unittest.mock import patch +from unittest.mock import MagicMock, patch import pytest @@ -129,6 +129,53 @@ def test_should_not_pollute_shared_key_with_zero_cost_pricing(): ) +def test_vllm_deployments_with_same_model_use_their_own_endpoint_metadata(monkeypatch): + def get(url, headers): + expected_key = "key-a" if "vllm-a" in url else "key-b" + assert dict(headers) == {"authorization": f"Bearer {expected_key}"} + response = MagicMock() + context = 32_768 if "vllm-a" in url else 131_072 + response.json.return_value = {"data": [{"id": "shared", "max_model_len": context}]} + return response + + monkeypatch.setattr(litellm.module_level_client, "get", get) + router = Router( + model_list=[ + { + "model_name": "vllm-group", + "litellm_params": { + "model": "hosted_vllm/shared", + "api_base": "https://vllm-a.example/v1", + "api_key": "key-a", + }, + "model_info": {"id": "vllm-deployment-a"}, + }, + { + "model_name": "vllm-group", + "litellm_params": { + "model": "hosted_vllm/shared", + "api_base": "https://vllm-b.example/v1", + "api_key": "key-b", + }, + "model_info": {"id": "vllm-deployment-b"}, + }, + ] + ) + + info_a = router.get_deployment_model_info( + model_id="vllm-deployment-a", + model_name="hosted_vllm/shared", + ) + info_b = router.get_deployment_model_info( + model_id="vllm-deployment-b", + model_name="hosted_vllm/shared", + ) + + assert info_a is not None and info_a["max_input_tokens"] == 32_768 + assert info_b is not None and info_b["max_input_tokens"] == 131_072 + assert info_a["max_output_tokens"] is None + assert info_b["max_output_tokens"] is None + def test_should_not_pollute_shared_key_with_custom_nonzero_pricing(): """ A deployment with custom (non-zero) pricing should not overwrite From c6ce646e0c42141d0909075ba53bdf83df5a9417 Mon Sep 17 00:00:00 2001 From: Aiden McComiskey Date: Mon, 7 Sep 2026 12:45:35 -0400 Subject: [PATCH 2/4] fix(hosted-vllm): harden context discovery --- litellm/llms/custom_httpx/http_handler.py | 4 +- litellm/llms/vllm/common_utils.py | 108 +++++----------- litellm/proxy/proxy_server.py | 50 ++++++-- litellm/router.py | 32 +---- litellm/utils.py | 7 +- .../test_hosted_vllm_chat_transformation.py | 52 ++++++-- .../proxy_server/test_routes_model_info.py | 116 ++++++++++++++---- .../test_router_model_cost_isolation.py | 49 +------- 8 files changed, 212 insertions(+), 206 deletions(-) diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 129453fe827..e1f0fc9e7d3 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -651,7 +651,7 @@ class AsyncHTTPHandler: self, url: str, params: dict | None = None, - headers: dict | httpx.Headers | None = None, + headers: dict | None = None, follow_redirects: bool | None = None, timeout: float | httpx.Timeout | None = None, ): @@ -1289,7 +1289,7 @@ class HTTPHandler: self, url: str, params: dict | None = None, - headers: dict | httpx.Headers | None = None, + headers: dict | None = None, follow_redirects: bool | None = None, timeout: float | httpx.Timeout | None = None, ): diff --git a/litellm/llms/vllm/common_utils.py b/litellm/llms/vllm/common_utils.py index 40f8f72a2f7..3f1923050ab 100644 --- a/litellm/llms/vllm/common_utils.py +++ b/litellm/llms/vllm/common_utils.py @@ -1,4 +1,4 @@ -from typing import Annotated, Final +from typing import Annotated, Final, Literal import httpx from pydantic import BaseModel, ConfigDict, Field, ValidationError @@ -44,7 +44,7 @@ class VLLMError(BaseLLMException): class VLLMModelInfo(BaseLLMModelInfo): - def __init__(self, provider: str = "vllm") -> None: + def __init__(self, provider: Literal["vllm", "hosted_vllm"] = "vllm") -> None: self._provider: Final = provider def validate_environment( @@ -81,104 +81,58 @@ class VLLMModelInfo(BaseLLMModelInfo): raise ValueError(f"{environment_variable} is required to query vLLM's `/models` endpoint.") return resolved_api_base - def _get_discovery_api_key(self, api_key: str | None) -> str | None: - environment_variable: Final = "HOSTED_VLLM_API_KEY" if self._provider == "hosted_vllm" else "VLLM_API_KEY" - return api_key or get_secret_str(environment_variable) - @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]: - passed_api_base: Final = api_base + 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) - resolved_api_key: Final = self._get_discovery_api_key(api_key) if passed_api_base is None or api_key else None - endpoint: Final = "/v1/models" - - url: Final = _add_path_to_api_base(resolved_api_base, endpoint) - headers: Final = ( - httpx.Headers((("Authorization", f"Bearer {resolved_api_key}"),)) if resolved_api_key else httpx.Headers() + environment_variable: Final = "HOSTED_VLLM_API_KEY" if self._provider == "hosted_vllm" else "VLLM_API_KEY" + resolved_api_key: Final = ( + api_key if api_key is not None or api_base is not None else get_secret_str(environment_variable) ) + headers: Final = {"Authorization": f"Bearer {resolved_api_key}"} if resolved_api_key else {} 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] - @staticmethod - def _strip_provider_prefix(model: str) -> str: - for prefix in ("hosted_vllm/", "vllm/"): - if model.startswith(prefix): - return model[len(prefix) :] - return model - - def _model_info_from_entry( - self, - raw_entry: object, - *, - target: str, - model: str, - ) -> ModelInfoBase | None: - try: - entry: Final = _VLLMModelEntry.model_validate(raw_entry) - except ValidationError: - return None - if entry.id != target or entry.max_model_len is None: - return None - - model_info: Final[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 model_info - def get_model_info( self, model: str, api_base: str | None = None, api_key: str | None = None, ) -> ModelInfoBase | None: - passed_api_base: Final = api_base - resolved_api_base: Final = self._get_discovery_api_base(api_base) - resolved_api_key: Final = self._get_discovery_api_key(api_key) if passed_api_base is None or api_key else None - - headers: Final = ( - httpx.Headers((("Authorization", f"Bearer {resolved_api_key}"),)) if resolved_api_key else httpx.Headers() - ) - response: Final = litellm.module_level_client.get( - url=_add_path_to_api_base(resolved_api_base, "/v1/models"), - headers=headers, - ) - response.raise_for_status() - - target: Final = self._strip_provider_prefix(model) + response: Final = self._query_models(api_base, api_key) + target: Final = model.removeprefix(f"{self._provider}/") discovered: Final = _VLLMModelsResponse.model_validate(response.json()) - return next( - ( - model_info - for raw_entry in discovered.data - if ( - model_info := self._model_info_from_entry( - raw_entry, - target=target, - model=model, - ) + 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, ) - is not None - ), - 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 3a273520624..9cada3b2b36 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -8916,6 +8916,7 @@ def select_data_generator( 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", {}) litellm_params: Final = model.get("litellm_params") or _EMPTY_MAPPING configured_model: Final = litellm_params.get("model", None) @@ -8924,18 +8925,42 @@ def get_litellm_model_info(model: dict = {}): ) model_to_lookup: Final = model_info.get("base_model", None) if use_base_model else configured_model try: - litellm_model_info: Final = litellm.get_model_info( - model_to_lookup, - custom_llm_provider=litellm_params.get("custom_llm_provider"), - api_base=litellm_params.get("api_base"), - api_key=litellm_params.get("api_key"), + 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=credentials.get("api_base", litellm_params.get("api_base")), + api_key=credentials.get("api_key", litellm_params.get("api_key")), discover_model_info=True, ) - return litellm_model_info + 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 {} + ), + } 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): @@ -13917,7 +13942,8 @@ 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( + all_models[i] = await asyncio.to_thread( + _enrich_model_info_with_litellm_data, model=_model, debug=debug if debug is not None else False, llm_router=llm_router, @@ -14660,7 +14686,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( @@ -14726,7 +14754,9 @@ 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)) + _translate_model_name_for_response( + await asyncio.to_thread(_enrich_model_info_with_litellm_data, model=model, llm_router=llm_router) + ) for model in all_models ] diff --git a/litellm/router.py b/litellm/router.py index c7277581698..6da201725b6 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -10321,12 +10321,7 @@ class Router: model_name: Final = model_info["model_name"] return self.get_model_list(model_name=model_name) - def get_deployment_model_info( - self, - model_id: str, - model_name: str, - litellm_params: LiteLLM_Params | None = None, - ) -> ModelInfo | None: + def get_deployment_model_info(self, model_id: str, model_name: str) -> ModelInfo | None: """ For a given model id, return the model info @@ -10345,25 +10340,8 @@ class Router: except Exception: pass - deployment: Final = self.get_deployment(model_id=model_id) if litellm_params is None else None - resolved_litellm_params: Final = ( - litellm_params - if litellm_params is not None - else deployment.litellm_params - if deployment is not None - else None - ) - try: - litellm_model_name_model_info = litellm.get_model_info( - model=model_name, - custom_llm_provider=( - resolved_litellm_params.custom_llm_provider if resolved_litellm_params is not None else None - ), - api_base=resolved_litellm_params.api_base if resolved_litellm_params is not None else None, - api_key=resolved_litellm_params.api_key if resolved_litellm_params is not None else None, - discover_model_info=True, - ) + litellm_model_name_model_info = litellm.get_model_info(model=model_name) except Exception: pass @@ -10479,11 +10457,7 @@ class Router: try: model_id = model_info_dict.get("id", None) if model_id is not None: - model_info = self.get_deployment_model_info( - model_id=model_id, - model_name=litellm_params.model, - litellm_params=litellm_params, - ) + model_info = self.get_deployment_model_info(model_id=model_id, model_name=litellm_params.model) else: model_info = None except Exception: diff --git a/litellm/utils.py b/litellm/utils.py index bb1d22cdfb6..f479de65de7 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -5636,7 +5636,7 @@ def _get_model_info_helper( "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": @@ -6024,7 +6024,8 @@ def get_model_info( - 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): query supported provider metadata endpoints before falling back to the static map. + - discover_model_info (bool): opt in to a synchronous, uncached vLLM metadata lookup; defaults to False. + Explicit api_base never inherits an ambient API key. Discovery overlays context only, not output limits. Returns: dict: A dictionary containing the following information: @@ -8812,7 +8813,7 @@ class ProviderConfigManager: VLLMModelInfo, # experimental approach, to reduce bloat on __init__.py ) - return VLLMModelInfo(provider=provider.value) + 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/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py b/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py index 26317420b90..fc635e81957 100644 --- a/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py +++ b/tests/test_litellm/llms/hosted_vllm/chat/test_hosted_vllm_chat_transformation.py @@ -2,6 +2,7 @@ import json import logging from unittest.mock import MagicMock, patch +import httpx import pytest import litellm @@ -9,6 +10,7 @@ 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 @@ -110,9 +112,7 @@ def test_hosted_vllm_chat_transformation_with_audio_url(): def test_hosted_vllm_supports_reasoning_effort(): config = HostedVLLMChatConfig() - supported_params = config.get_supported_openai_params( - model="hosted_vllm/gpt-oss-120b" - ) + supported_params = config.get_supported_openai_params(model="hosted_vllm/gpt-oss-120b") assert "reasoning_effort" in supported_params optional_params = config.map_openai_params( non_default_params={"reasoning_effort": "high"}, @@ -133,9 +133,7 @@ def test_hosted_vllm_supports_thinking(): Related issue: https://github.com/BerriAI/litellm/issues/19761 """ config = HostedVLLMChatConfig() - supported_params = config.get_supported_openai_params( - model="hosted_vllm/GLM-4.6-FP8" - ) + supported_params = config.get_supported_openai_params(model="hosted_vllm/GLM-4.6-FP8") assert "thinking" in supported_params # Test thinking below the low threshold -> "minimal" @@ -461,7 +459,7 @@ def test_vllm_model_info_maps_context_only_and_authenticates(monkeypatch) -> Non assert info["max_output_tokens"] is None request.assert_called_once() assert request.call_args.kwargs["url"] == "https://vllm.example/v1/models" - assert dict(request.call_args.kwargs["headers"]) == {"authorization": "Bearer secret-key"} + assert request.call_args.kwargs["headers"] == {"Authorization": "Bearer secret-key"} def test_vllm_model_info_does_not_discover_without_opt_in(monkeypatch) -> None: @@ -477,7 +475,33 @@ def test_vllm_model_info_does_not_discover_without_opt_in(monkeypatch) -> None: request.assert_not_called() -@pytest.mark.parametrize("value", [None, 0, -1, True, "262144"]) +@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, @@ -522,7 +546,7 @@ def test_hosted_vllm_model_info_uses_provider_environment(monkeypatch) -> None: assert info["max_input_tokens"] == 32_768 request.assert_called_once() assert request.call_args.kwargs["url"] == "https://hosted.example/v1/models" - assert dict(request.call_args.kwargs["headers"]) == {"authorization": "Bearer hosted-secret"} + assert request.call_args.kwargs["headers"] == {"Authorization": "Bearer hosted-secret"} def test_bare_vllm_model_keeps_explicit_provider_identity(monkeypatch) -> None: @@ -582,7 +606,7 @@ def test_vllm_endpoint_scoped_lookup_refreshes_and_removes_without_global_regist def test_vllm_same_model_id_is_isolated_by_endpoint(monkeypatch) -> None: - def get(url: str, headers: dict[str, str]) -> MagicMock: + 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}) @@ -611,7 +635,7 @@ def test_vllm_discovery_failure_does_not_log_the_api_key(monkeypatch, caplog) -> monkeypatch.setattr( litellm.module_level_client, "get", - MagicMock(side_effect=RuntimeError("discovery failed")), + MagicMock(return_value=MagicMock(json=lambda: {"error": {"authorization": secret}})), ) with caplog.at_level(logging.WARNING), pytest.raises(Exception, match="isn't mapped yet"): @@ -633,6 +657,7 @@ def test_vllm_endpoint_discovery_survives_price_map_replacement(monkeypatch) -> 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", @@ -642,6 +667,7 @@ def test_vllm_endpoint_discovery_survives_price_map_replacement(monkeypatch) -> ) 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 @@ -652,6 +678,8 @@ def test_vllm_discovery_preserves_static_pricing(monkeypatch) -> None: "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, } @@ -673,4 +701,6 @@ def test_vllm_discovery_preserves_static_pricing(monkeypatch) -> None: 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() 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 5f5b10b1fbc..938cb0702e3 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,8 +9,10 @@ Pins (PR2): from __future__ import annotations -from unittest.mock import MagicMock +import threading +from unittest.mock import AsyncMock, MagicMock +import httpx import pytest import litellm @@ -129,7 +131,6 @@ def test_v1_model_info_no_model_list_error(client, auth_as, null_router, path): assert "LLM Model List not loaded" in response.text - def test_get_proxy_model_info_surfaces_supports_parallel_function_calling(local_model_cost_map): """``GET /v1/model/info`` enriches each deployment through ``_get_proxy_model_info``; a registry entry declaring parallel function calling must land in ``model_info`` instead of null.""" @@ -171,13 +172,29 @@ def test_get_proxy_model_info_discovers_vllm_context_with_config_precedence( 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 dict(request.call_args.kwargs["headers"]) == {"authorization": "Bearer endpoint-secret"} + assert request.call_args.kwargs["headers"] == {"Authorization": "Bearer endpoint-secret"} -def test_v1_model_info_route_surfaces_discovered_vllm_context( +@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}]} @@ -191,26 +208,79 @@ def test_v1_model_info_route_surfaces_discovered_vllm_context( "api_base": "https://vllm.example/v1", "api_key": "endpoint-secret", }, - "model_info": {"id": "vllm-route-deployment"}, + "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 + for call in request.call_args_list: + assert call.kwargs["url"] == "https://vllm.example/v1/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) - - with auth_as(): - result = client.get( - "/v1/model/info", - params={"litellm_model_id": "vllm-route-deployment"}, - ) - - assert result.status_code == 200 - deployment = result.json()["data"][0] - assert deployment["model_info"]["max_input_tokens"] == 262_144 - assert deployment["model_info"]["max_output_tokens"] is None - assert "api_key" not in deployment["litellm_params"] + 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 def test_v1_model_info_star_wildcard_filter_keeps_provider_expansion(monkeypatch): @@ -232,9 +302,7 @@ def test_v1_model_info_star_wildcard_filter_keeps_provider_expansion(monkeypatch router.get_model_list = MagicMock(return_value=[deployment]) monkeypatch.setattr(model_checks, "get_provider_models", fake_get_provider_models) - expanded_deployments = proxy_server.expand_wildcard_deployments_for_model_info( - [deployment] - ) + expanded_deployments = proxy_server.expand_wildcard_deployments_for_model_info([deployment]) allowed_model_names = proxy_server._get_v1_model_info_allowed_model_names( user_api_key_dict=UserAPIKeyAuth( api_key="sk-test", @@ -470,14 +538,10 @@ def test_v2_model_info_exclude_auto_routers_shrinks_total_count(client, auth_as, assert len(payload["data"]) == payload["total_count"] -def test_v2_model_info_exclude_auto_routers_paginates_over_the_filtered_set( - client, auth_as, mixed_auto_router_router -): +def test_v2_model_info_exclude_auto_routers_paginates_over_the_filtered_set(client, auth_as, mixed_auto_router_router): """Page size applies to the filtered list, so no page silently comes back short.""" with auth_as(): - response = client.get( - "/v2/model/info", params={"exclude_auto_routers": "true", "page": 1, "size": 1} - ) + response = client.get("/v2/model/info", params={"exclude_auto_routers": "true", "page": 1, "size": 1}) payload = response.json() assert payload["total_count"] == 2 assert payload["total_pages"] == 2 diff --git a/tests/test_litellm/test_router_model_cost_isolation.py b/tests/test_litellm/test_router_model_cost_isolation.py index c510ed2436c..30b265905f3 100644 --- a/tests/test_litellm/test_router_model_cost_isolation.py +++ b/tests/test_litellm/test_router_model_cost_isolation.py @@ -11,7 +11,7 @@ import copy import logging import os import re -from unittest.mock import MagicMock, patch +from unittest.mock import patch import pytest @@ -129,53 +129,6 @@ def test_should_not_pollute_shared_key_with_zero_cost_pricing(): ) -def test_vllm_deployments_with_same_model_use_their_own_endpoint_metadata(monkeypatch): - def get(url, headers): - expected_key = "key-a" if "vllm-a" in url else "key-b" - assert dict(headers) == {"authorization": f"Bearer {expected_key}"} - response = MagicMock() - context = 32_768 if "vllm-a" in url else 131_072 - response.json.return_value = {"data": [{"id": "shared", "max_model_len": context}]} - return response - - monkeypatch.setattr(litellm.module_level_client, "get", get) - router = Router( - model_list=[ - { - "model_name": "vllm-group", - "litellm_params": { - "model": "hosted_vllm/shared", - "api_base": "https://vllm-a.example/v1", - "api_key": "key-a", - }, - "model_info": {"id": "vllm-deployment-a"}, - }, - { - "model_name": "vllm-group", - "litellm_params": { - "model": "hosted_vllm/shared", - "api_base": "https://vllm-b.example/v1", - "api_key": "key-b", - }, - "model_info": {"id": "vllm-deployment-b"}, - }, - ] - ) - - info_a = router.get_deployment_model_info( - model_id="vllm-deployment-a", - model_name="hosted_vllm/shared", - ) - info_b = router.get_deployment_model_info( - model_id="vllm-deployment-b", - model_name="hosted_vllm/shared", - ) - - assert info_a is not None and info_a["max_input_tokens"] == 32_768 - assert info_b is not None and info_b["max_input_tokens"] == 131_072 - assert info_a["max_output_tokens"] is None - assert info_b["max_output_tokens"] is None - def test_should_not_pollute_shared_key_with_custom_nonzero_pricing(): """ A deployment with custom (non-zero) pricing should not overwrite From a4f47b3c4c774339fd0c7254dd811cc60b4d3e29 Mon Sep 17 00:00:00 2001 From: Aiden McComiskey Date: Thu, 24 Sep 2026 18:02:00 +0100 Subject: [PATCH 3/4] fix(hosted-vllm): match chat credentials and enrich model info concurrently --- litellm/llms/vllm/common_utils.py | 4 +-- litellm/proxy/proxy_server.py | 31 ++++++++++++------- litellm/utils.py | 2 +- .../proxy_server/test_routes_model_info.py | 29 ++++++++++++++++- .../test_hosted_vllm_chat_transformation.py | 13 ++++---- 5 files changed, 56 insertions(+), 23 deletions(-) diff --git a/litellm/llms/vllm/common_utils.py b/litellm/llms/vllm/common_utils.py index 3f1923050ab..cdfef9a42bc 100644 --- a/litellm/llms/vllm/common_utils.py +++ b/litellm/llms/vllm/common_utils.py @@ -88,9 +88,7 @@ class VLLMModelInfo(BaseLLMModelInfo): 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 if api_key is not None or api_base is not None else get_secret_str(environment_variable) - ) + 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 {} response: Final = litellm.module_level_client.get( url=_add_path_to_api_base(resolved_api_base, "/v1/models"), diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 6ebab691c02..a4453a5f480 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -9775,8 +9775,8 @@ def get_litellm_model_info(model: dict = {}): live_model_info: Final = litellm.get_model_info( configured_model, custom_llm_provider=provider, - api_base=credentials.get("api_base", litellm_params.get("api_base")), - api_key=credentials.get("api_key", litellm_params.get("api_key")), + 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 { @@ -15052,13 +15052,19 @@ 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] = await asyncio.to_thread( - _enrich_model_info_with_litellm_data, - model=_model, - debug=debug if debug is not None else False, - llm_router=llm_router, + all_models = list( + 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(): @@ -15838,10 +15844,13 @@ async def model_info_v1( all_models = _filter_models_to_user_accessible(all_models) all_models = [ - _translate_model_name_for_response( - await asyncio.to_thread(_enrich_model_info_with_litellm_data, model=model, llm_router=llm_router) + _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 + ) ) - for model in all_models ] if teamId is not None and teamId.strip(): diff --git a/litellm/utils.py b/litellm/utils.py index 2e230a7b30d..d8fcf01e533 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -6387,7 +6387,7 @@ def get_model_info( - 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. - Explicit api_base never inherits an ambient API key. Discovery overlays context only, not output limits. + Discovery overlays context only, not output limits. Returns: dict: A dictionary containing the following information: 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 36c38225d9e..4c49da078a7 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 @@ -402,8 +402,9 @@ def test_model_info_routes_refresh_discovered_context_below_explicit_config( 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"] == "https://vllm.example/v1/models" + assert call.kwargs["url"] == f"{chat_api_base}/models" assert call.kwargs["headers"] == {"Authorization": "Bearer endpoint-secret"} @@ -426,6 +427,32 @@ async def test_model_info_discovery_runs_outside_the_event_loop(app, auth_as, co 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 774b41ed163..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 @@ -458,17 +458,16 @@ def test_vllm_model_info_ignores_non_positive_or_non_integer_context( ) -def test_vllm_explicit_base_never_receives_an_ambient_key(monkeypatch) -> 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", "ambient-secret") + 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="https://operator-supplied.example/v1", - ) + VLLMModelInfo(provider="hosted_vllm").get_model_info(model="hosted_vllm/model", api_base=api_base) - assert dict(request.call_args.kwargs["headers"]) == {} + _, 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: From ad4488b168b0d0a4d6c16554588957fe24c076eb Mon Sep 17 00:00:00 2001 From: Aiden McComiskey Date: Fri, 25 Sep 2026 15:43:14 +0100 Subject: [PATCH 4/4] chore: drop mutable-collection constructions flagged by LIT002 gate --- litellm/llms/vllm/common_utils.py | 2 +- litellm/proxy/proxy_server.py | 24 +++++++++++------------- 2 files changed, 12 insertions(+), 14 deletions(-) diff --git a/litellm/llms/vllm/common_utils.py b/litellm/llms/vllm/common_utils.py index cdfef9a42bc..576931582fb 100644 --- a/litellm/llms/vllm/common_utils.py +++ b/litellm/llms/vllm/common_utils.py @@ -89,7 +89,7 @@ class VLLMModelInfo(BaseLLMModelInfo): 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 {} + headers: Final = {"Authorization": f"Bearer {resolved_api_key}"} if resolved_api_key else None response: Final = litellm.module_level_client.get( url=_add_path_to_api_base(resolved_api_base, "/v1/models"), headers=headers, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index a4453a5f480..55c00b57720 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -9782,10 +9782,10 @@ def get_litellm_model_info(model: dict = {}): return { **live_model_info, **static_model_info, - **( - {"max_input_tokens": live_model_info["max_input_tokens"]} + "max_input_tokens": ( + live_model_info["max_input_tokens"] if live_model_info["max_input_tokens"] is not None - else {} + else static_model_info.get("max_input_tokens") ), } except Exception: @@ -15052,17 +15052,15 @@ 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 - all_models = list( - 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 + 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 ) )