fix(hosted-vllm): match chat credentials and enrich model info concurrently

This commit is contained in:
Aiden McComiskey 2026-09-24 18:02:00 +01:00
parent 3682b06887
commit a4f47b3c4c
5 changed files with 56 additions and 23 deletions

View file

@ -88,9 +88,7 @@ class VLLMModelInfo(BaseLLMModelInfo):
def _query_models(self, api_base: str | None, api_key: str | None) -> httpx.Response:
resolved_api_base: Final = self._get_discovery_api_base(api_base)
environment_variable: Final = "HOSTED_VLLM_API_KEY" if self._provider == "hosted_vllm" else "VLLM_API_KEY"
resolved_api_key: Final = (
api_key if api_key is not None or api_base is not None else get_secret_str(environment_variable)
)
resolved_api_key: Final = api_key or get_secret_str(environment_variable)
headers: Final = {"Authorization": f"Bearer {resolved_api_key}"} if resolved_api_key else {}
response: Final = litellm.module_level_client.get(
url=_add_path_to_api_base(resolved_api_base, "/v1/models"),

View file

@ -9775,8 +9775,8 @@ def get_litellm_model_info(model: dict = {}):
live_model_info: Final = litellm.get_model_info(
configured_model,
custom_llm_provider=provider,
api_base=credentials.get("api_base", litellm_params.get("api_base")),
api_key=credentials.get("api_key", litellm_params.get("api_key")),
api_base=litellm_params.get("api_base") or credentials.get("api_base"),
api_key=litellm_params.get("api_key") or credentials.get("api_key"),
discover_model_info=True,
)
return {
@ -15052,13 +15052,19 @@ async def model_info_v2(
# Fill in model info based on config.yaml and litellm model_prices_and_context_window.json
# This must happen before teamId filtering so that direct_access and access_via_team_ids are populated
for i, _model in enumerate(all_models):
all_models[i] = await asyncio.to_thread(
_enrich_model_info_with_litellm_data,
model=_model,
debug=debug if debug is not None else False,
llm_router=llm_router,
all_models = list(
await asyncio.gather(
*(
asyncio.to_thread(
_enrich_model_info_with_litellm_data,
model=_model,
debug=debug if debug is not None else False,
llm_router=llm_router,
)
for _model in all_models
)
)
)
# Apply teamId filter if provided
if teamId is not None and teamId.strip():
@ -15838,10 +15844,13 @@ async def model_info_v1(
all_models = _filter_models_to_user_accessible(all_models)
all_models = [
_translate_model_name_for_response(
await asyncio.to_thread(_enrich_model_info_with_litellm_data, model=model, llm_router=llm_router)
_translate_model_name_for_response(model)
for model in await asyncio.gather(
*(
asyncio.to_thread(_enrich_model_info_with_litellm_data, model=model, llm_router=llm_router)
for model in all_models
)
)
for model in all_models
]
if teamId is not None and teamId.strip():

View file

@ -6387,7 +6387,7 @@ def get_model_info(
- api_base (str | null): the deployment endpoint used for provider-scoped discovery.
- api_key (str | null): the deployment credential used for provider-scoped discovery.
- discover_model_info (bool): opt in to a synchronous, uncached vLLM metadata lookup; defaults to False.
Explicit api_base never inherits an ambient API key. Discovery overlays context only, not output limits.
Discovery overlays context only, not output limits.
Returns:
dict: A dictionary containing the following information:

View file

@ -402,8 +402,9 @@ def test_model_info_routes_refresh_discovered_context_below_explicit_config(
assert "api_key" not in deployment["litellm_params"]
assert "endpoint-secret" not in result.text
assert request.call_count == 3
chat_api_base = model_list[0]["litellm_params"]["api_base"]
for call in request.call_args_list:
assert call.kwargs["url"] == "https://vllm.example/v1/models"
assert call.kwargs["url"] == f"{chat_api_base}/models"
assert call.kwargs["headers"] == {"Authorization": "Bearer endpoint-secret"}
@ -426,6 +427,32 @@ async def test_model_info_discovery_runs_outside_the_event_loop(app, auth_as, co
assert lookup_threads[0] != event_loop_thread
@pytest.mark.parametrize("path", ["/v1/model/info", "/v2/model/info"])
def test_model_info_list_routes_enrich_deployments_concurrently(client, auth_as, monkeypatch, mock_prisma, path):
model_list = [
{"model_name": name, "litellm_params": {"model": f"hosted_vllm/{name}"}, "model_info": {"id": name}}
for name in ("slow-a", "slow-b")
]
both_lookups_started = threading.Barrier(2, timeout=5)
def enrich(model, **kwargs):
both_lookups_started.wait()
return model
monkeypatch.setattr(proxy_server, "_enrich_model_info_with_litellm_data", enrich)
monkeypatch.setattr(proxy_server, "llm_router", litellm.Router(model_list=model_list))
monkeypatch.setattr(proxy_server, "llm_model_list", model_list)
monkeypatch.setattr(proxy_server, "user_model", None)
monkeypatch.setattr(proxy_server, "prisma_client", mock_prisma if path == "/v2/model/info" else None)
monkeypatch.setattr(proxy_server.proxy_config, "get_config", AsyncMock(return_value={}))
with auth_as():
result = client.get(path)
assert result.status_code == 200, result.text
assert [deployment["model_info"]["id"] for deployment in result.json()["data"]] == ["slow-a", "slow-b"]
def _enriched_model_info(monkeypatch, litellm_params: dict, model_info: dict) -> dict:
monkeypatch.setattr(proxy_server, "llm_router", None)
enriched: Final = proxy_server._get_proxy_model_info(

View file

@ -458,17 +458,16 @@ def test_vllm_model_info_ignores_non_positive_or_non_integer_context(
)
def test_vllm_explicit_base_never_receives_an_ambient_key(monkeypatch) -> None:
def test_vllm_explicit_base_uses_the_same_environment_key_as_chat(monkeypatch) -> None:
request = MagicMock(return_value=_model_list_response({"id": "model"}))
monkeypatch.setenv("HOSTED_VLLM_API_KEY", "ambient-secret")
monkeypatch.setenv("HOSTED_VLLM_API_KEY", "env-secret")
monkeypatch.setattr(litellm.module_level_client, "get", request)
api_base = "https://operator-supplied.example/v1"
VLLMModelInfo(provider="hosted_vllm").get_model_info(
model="hosted_vllm/model",
api_base="https://operator-supplied.example/v1",
)
VLLMModelInfo(provider="hosted_vllm").get_model_info(model="hosted_vllm/model", api_base=api_base)
assert dict(request.call_args.kwargs["headers"]) == {}
_, chat_api_key = HostedVLLMChatConfig()._get_openai_compatible_provider_info(api_base=api_base, api_key=None)
assert request.call_args.kwargs["headers"] == {"Authorization": f"Bearer {chat_api_key}"}
def test_hosted_vllm_model_info_uses_provider_environment(monkeypatch) -> None: