mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
Merge 1e994f5b8d into 0f96d09588
This commit is contained in:
commit
086c46bff5
4 changed files with 202 additions and 6 deletions
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue