fix(hosted-vllm): harden context discovery

This commit is contained in:
Aiden McComiskey 2026-09-07 12:45:35 -04:00
parent 1599abead2
commit c6ce646e0c
8 changed files with 212 additions and 206 deletions

View file

@ -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,
):

View file

@ -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)

View file

@ -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
]

View file

@ -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:

View file

@ -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:

View file

@ -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()

View file

@ -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

View file

@ -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