From c6ce646e0c42141d0909075ba53bdf83df5a9417 Mon Sep 17 00:00:00 2001 From: Aiden McComiskey Date: Mon, 7 Sep 2026 12:45:35 -0400 Subject: [PATCH] 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