mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(hosted-vllm): match chat credentials and enrich model info concurrently
This commit is contained in:
parent
3682b06887
commit
a4f47b3c4c
5 changed files with 56 additions and 23 deletions
|
|
@ -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"),
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue