From 1599abead24116ae5a8ab232b3485e04d0901a10 Mon Sep 17 00:00:00 2001 From: Aiden McComiskey Date: Fri, 4 Sep 2026 16:56:53 -0400 Subject: [PATCH] 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