fix(hosted-vllm): discover context limits

This commit is contained in:
Aiden McComiskey 2026-09-04 16:56:53 -04:00
parent 300d335255
commit 1599abead2
8 changed files with 544 additions and 25 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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