diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index c0123ae45a3..8803538b0cd 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -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", diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index eaa03b67b40..362b453ecde 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -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, diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 83ac56c4c85..bba46ab296a 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -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( diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index 227921d6150..10929193908 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -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