This commit is contained in:
Vedant Madane 2026-09-28 14:46:34 -04:00 • committed by GitHub
commit 086c46bff5
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 202 additions and 6 deletions

View file

@ -1685,12 +1685,16 @@ _MODEL_ROUTING_ROUTE_MARKERS: Final = (
"/evals",
"/fine_tuning",
"/videos",
# vLLM GET passthrough selects the router model via ?model=; include so
# key/team allowlists and model budgets see the same model as the body path.
"/vllm",
)
_MODEL_ROUTING_HEADER_OR_QUERY_ROUTE_MARKERS: Final = (
"/files",
"/batches",
"/skills",
"/evals",
"/vllm",
)
_MODEL_ROUTING_QUERY_TARGET_MODEL_ROUTE_MARKERS: Final = (
"/files",

View file

@ -24,6 +24,7 @@ from typing import TYPE_CHECKING, Annotated, Final, Literal, Protocol, cast
import httpx
from fastapi import APIRouter, Depends, HTTPException, Request, Response, WebSocket
from fastapi.responses import StreamingResponse
from starlette.datastructures import QueryParams
from starlette.websockets import WebSocketState
from typing_extensions import ReadOnly, TypedDict
@ -58,7 +59,7 @@ from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
from litellm.proxy._types import *
from litellm.proxy.auth.auth_checks import enforced_model_allowlists
from litellm.proxy.auth.auth_checks import can_key_call_model, enforced_model_allowlists
from litellm.proxy.auth.handle_jwt import JWTHandler
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.auth.user_api_key_auth import (
@ -140,12 +141,32 @@ def create_request_copy(request: Request):
}
def is_passthrough_request_using_router_model(request_body: dict, llm_router: litellm.Router | None) -> bool:
def _model_from_body_or_query(
request_body: Mapping[str, object],
query_params: Mapping[str, str] | QueryParams | None = None,
) -> str | None:
"""Resolve model name from JSON body, falling back to query string (vLLM GET)."""
body_model: Final = request_body.get("model")
if isinstance(body_model, str) and body_model:
return body_model
if query_params is None:
return None
query_model: Final = query_params.get("model")
if isinstance(query_model, str) and query_model:
return query_model
return None
def is_passthrough_request_using_router_model(
request_body: Mapping[str, object],
llm_router: litellm.Router | None,
query_params: Mapping[str, str] | QueryParams | None = None,
) -> bool:
"""
Returns True if the model is in the llm_router model names
"""
try:
model: Final = request_body.get("model")
model: Final = _model_from_body_or_query(request_body, query_params)
return is_known_model(model, llm_router)
except Exception:
return False
@ -169,6 +190,27 @@ def _models_served_by_group(llm_router: litellm.Router, model_group: str) -> fro
)
async def _authorize_passthrough_route_model(
model_name: str,
user_api_key_dict: UserAPIKeyAuth,
llm_router: litellm.Router | None,
) -> None:
"""Enforce key model allowlists for router-selected passthrough models.
Defense-in-depth for routes (for example vLLM GET) where the model may
arrive via query string after user_api_key_auth has already run. Uses the
public can_key_call_model helper (same path as body-selected models).
"""
if llm_router is not None and not hasattr(llm_router, "model_group_alias"):
return
await can_key_call_model(
model=model_name,
llm_model_list=None,
valid_token=user_api_key_dict,
llm_router=llm_router,
)
def is_passthrough_request_streaming(request_body: object) -> bool:
"""
Returns True if the request is streaming.
@ -492,13 +534,28 @@ async def vllm_proxy_route(
from litellm.proxy.proxy_server import llm_router
request_body: Final = await get_request_body(request)
is_router_model: Final = is_passthrough_request_using_router_model(request_body, llm_router)
body_for_router_check = request_body
if not request_body.get("model") and request.query_params.get("model"):
body_for_router_check = dict(request_body)
body_for_router_check["model"] = request.query_params.get("model")
is_router_model: Final = is_passthrough_request_using_router_model(body_for_router_check, llm_router)
is_streaming_request: Final = is_passthrough_request_streaming(request_body)
if is_router_model and llm_router:
model_name: Final = _model_from_body_or_query(request_body, query_params=request.query_params)
if not model_name:
raise HTTPException(
status_code=400,
detail="model is required in the request body or query string for vLLM router passthrough",
)
await _authorize_passthrough_route_model(
model_name=model_name,
user_api_key_dict=user_api_key_dict,
llm_router=llm_router,
)
result: Final = cast(
httpx.Response,
await llm_router.allm_passthrough_route(
model=request_body.get("model"),
model=model_name,
method=request.method,
endpoint=endpoint,
request_query_params=request.query_params,

View file

@ -940,6 +940,28 @@ def test_get_model_from_request_includes_file_endpoint_header_model():
)
def test_get_model_from_request_extracts_vllm_query_model():
"""vLLM GET passthrough selects models via ?model=; auth must see that model."""
assert (
get_model_from_request(
request_data={},
route="/vllm/v1/models",
request_query_params={"model": "restricted-vllm-model"},
)
== "restricted-vllm-model"
)
def test_get_model_from_request_vllm_prefers_body_and_includes_query():
models = get_model_from_request(
request_data={"model": "body-vllm-model"},
route="/vllm/chat/completions",
request_query_params={"model": "query-vllm-model"},
)
assert isinstance(models, list)
assert set(models) == {"body-vllm-model", "query-vllm-model"}
def test_get_model_from_request_ignores_routing_header_on_standard_llm_routes():
assert (
get_model_from_request(

View file

@ -2188,6 +2188,10 @@ class TestLLMPassthroughFactoryProxyRoute:
class TestVLLMProxyRoute:
@pytest.mark.asyncio
@patch( # test-quality-ok: patching litellm internal for unit test isolation
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._authorize_passthrough_route_model",
new_callable=AsyncMock,
)
@patch( # test-quality-ok: patching litellm internal for unit test isolation
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
return_value={"model": "router-model", "stream": False},
@ -2198,7 +2202,7 @@ class TestVLLMProxyRoute:
)
@patch("litellm.proxy.proxy_server.llm_router") # test-quality-ok: patching litellm internal for unit test isolation
async def test_vllm_proxy_route_with_router_model(
self, mock_llm_router, mock_is_router, mock_get_body
self, mock_llm_router, mock_is_router, mock_get_body, mock_authorize
):
mock_request = MagicMock(spec=Request)
mock_request.method = "POST"
@ -2218,8 +2222,100 @@ class TestVLLMProxyRoute:
)
mock_is_router.assert_called_once()
mock_authorize.assert_awaited_once_with(
model_name="router-model",
user_api_key_dict=mock_user_api_key_dict,
llm_router=mock_llm_router,
)
mock_llm_router.allm_passthrough_route.assert_awaited_once()
@pytest.mark.asyncio
@patch( # test-quality-ok: patching litellm internal for unit test isolation
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._authorize_passthrough_route_model",
new_callable=AsyncMock,
)
@patch( # test-quality-ok: patching litellm internal for unit test isolation
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
return_value={},
)
@patch( # test-quality-ok: patching litellm internal for unit test isolation
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model",
return_value=True,
)
@patch("litellm.proxy.proxy_server.llm_router") # test-quality-ok: patching litellm internal for unit test isolation
async def test_vllm_proxy_route_query_only_model_uses_router_and_auth(
self, mock_llm_router, mock_is_router, mock_get_body, mock_authorize
):
"""GET /vllm/...?model= must route and authorize like body model selection."""
mock_request = MagicMock(spec=Request)
mock_request.method = "GET"
mock_request.headers = {}
mock_request.query_params = {"model": "query-router-model"}
mock_fastapi_response = MagicMock(spec=Response)
mock_user_api_key_dict = MagicMock()
mock_llm_router.allm_passthrough_route = AsyncMock(
return_value=httpx.Response(200, json={"response": "success"})
)
await vllm_proxy_route(
endpoint="/v1/models",
request=mock_request,
fastapi_response=mock_fastapi_response,
user_api_key_dict=mock_user_api_key_dict,
)
mock_authorize.assert_awaited_once_with(
model_name="query-router-model",
user_api_key_dict=mock_user_api_key_dict,
llm_router=mock_llm_router,
)
call_kwargs = mock_llm_router.allm_passthrough_route.await_args.kwargs
assert call_kwargs["model"] == "query-router-model"
assert call_kwargs["method"] == "GET"
@pytest.mark.asyncio
@patch( # test-quality-ok: patching litellm internal for unit test isolation
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._authorize_passthrough_route_model",
new_callable=AsyncMock,
)
@patch( # test-quality-ok: patching litellm internal for unit test isolation
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
return_value={},
)
@patch( # test-quality-ok: patching litellm internal for unit test isolation
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_passthrough_request_using_router_model",
return_value=True,
)
@patch("litellm.proxy.proxy_server.llm_router") # test-quality-ok: patching litellm internal for unit test isolation
async def test_vllm_proxy_route_query_model_auth_denies_disallowed(
self, mock_llm_router, mock_is_router, mock_get_body, mock_authorize
):
from litellm.proxy._types import ProxyErrorTypes, ProxyException
mock_request = MagicMock(spec=Request)
mock_request.method = "GET"
mock_request.headers = {}
mock_request.query_params = {"model": "disallowed-model"}
mock_fastapi_response = MagicMock(spec=Response)
mock_user_api_key_dict = MagicMock()
mock_authorize.side_effect = ProxyException(
message="key not allowed to access model",
type=ProxyErrorTypes.key_model_access_denied,
param="model",
code=403,
)
mock_llm_router.allm_passthrough_route = AsyncMock()
with pytest.raises(ProxyException):
await vllm_proxy_route(
endpoint="/v1/models",
request=mock_request,
fastapi_response=mock_fastapi_response,
user_api_key_dict=mock_user_api_key_dict,
)
mock_llm_router.allm_passthrough_route.assert_not_awaited()
@pytest.mark.asyncio
@patch( # test-quality-ok: patching litellm internal for unit test isolation
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.get_request_body",
@ -2250,6 +2346,23 @@ class TestVLLMProxyRoute:
assert result == "factory_success"
mock_factory_route.assert_awaited_once()
def test_is_passthrough_request_using_router_model_reads_query_params(self):
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
is_passthrough_request_using_router_model,
)
with patch( # test-quality-ok: isolate is_known_model for query-param routing unit test
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_known_model",
return_value=True,
) as mock_known:
assert (
is_passthrough_request_using_router_model(
{}, MagicMock(), query_params={"model": "from-query"}
)
is True
)
mock_known.assert_called_once_with("from-query", mock.ANY)
class TestGigachatProxyRoute:
@pytest.mark.asyncio