mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(proxy): stop treating upstream model body field as a LiteLLM model on auth-enforced pass-through routes (#33710)
* fix(proxy): stop treating upstream model body field as a LiteLLM model on auth-enforced pass-through routes An auth: true user-defined pass-through endpoint runs full virtual-key auth, and get_model_from_request unconditionally extracted the request body model field, so key/team/user/project model allowlist checks rejected requests whose model only exists upstream (key_model_access_denied), even when the key was explicitly granted the route via allowed_passthrough_routes. The pass-through route registry moves to a leaf module (route_registry.py) that the auth layer can import without re-entering the pass_through_endpoints -> user_api_key_auth -> auth_utils import cycle. get_model_from_request now returns None for routes registered as user-defined pass-through endpoints (exact and subpath), which skips model allowlist and per-model budget enforcement on those routes while key auth, allowed_passthrough_routes, and spend/budget checks stay intact. Built-in provider passthrough routes (/vertex_ai, /gemini, ...) keep model enforcement. Resolves LIT-4299 * fix(proxy): key pass-through model-access skip on the dispatched endpoint, not the request path Addresses a model-authorization bypass: the first version decided whether to skip model-allowlist extraction by matching the request path against the pass-through route registry. That ignored the HTTP method and, more importantly, whether the request was actually dispatched to a pass-through handler. A custom pass-through whose path collides with a built-in route (e.g. /v1/chat/completions, or an include_subpath prefix of one) still writes a registry entry even though FastAPI serves the built-in handler, so a normal request to that route had its model checks skipped and could reach a model outside the key/team/user/project allowlist. The skip is now keyed off the FastAPI-resolved endpoint. create_pass_through_route tags its handler with LITELLM_PASS_THROUGH_ENDPOINT_MARKER, and get_model_from_request returns None only when request.scope["endpoint"] carries that marker. Because routing runs before auth dependencies, this reflects the handler that actually serves the request: on a collision the built-in handler is dispatched and carries no marker, so model enforcement stays on. This also removes the need for the separate route_registry module, so that extraction is reverted. Regression tests cover a pass-through-dispatched request (model suppressed), a built-in-dispatched request on the same path (model still enforced), and the no-request budget path. Resolves LIT-4299
This commit is contained in:
parent
561b6796bc
commit
adb1ffb119
7 changed files with 199 additions and 0 deletions
|
|
@ -527,6 +527,7 @@ async def common_checks(
|
|||
request_headers=_safe_get_request_headers(request=request),
|
||||
request_query_params=_safe_get_request_query_params(request=request),
|
||||
llm_router=llm_router,
|
||||
request=request,
|
||||
)
|
||||
|
||||
if route in MODEL_DISCOVERY_ROUTES:
|
||||
|
|
|
|||
|
|
@ -14,6 +14,9 @@ from litellm.constants import MINIMUM_CUSTOM_KEY_LENGTH, STANDARD_CUSTOMER_ID_HE
|
|||
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
|
||||
from litellm.litellm_core_utils.url_utils import SSRFError, validate_url
|
||||
from litellm.proxy._types import *
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
||||
LITELLM_PASS_THROUGH_ENDPOINT_MARKER,
|
||||
)
|
||||
from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS
|
||||
from litellm.types.utils import CustomPricingLiteLLMParams
|
||||
|
||||
|
|
@ -1482,13 +1485,50 @@ def _format_model_candidates(
|
|||
return candidates
|
||||
|
||||
|
||||
def _request_dispatched_to_pass_through_endpoint(request: Request | None) -> bool:
|
||||
"""Whether FastAPI resolved this request to a user-defined pass-through handler.
|
||||
|
||||
Reads the marker set by ``create_pass_through_route`` off the dispatched endpoint
|
||||
(``request.scope["endpoint"]``). Because routing has already run by the time auth
|
||||
dependencies execute, this reflects the handler that actually serves the request:
|
||||
a custom path colliding with a built-in route resolves to the built-in handler,
|
||||
which carries no marker, so model-access checks are never wrongly skipped.
|
||||
"""
|
||||
if request is None:
|
||||
return False
|
||||
scope = getattr(request, "scope", None)
|
||||
if not isinstance(scope, dict):
|
||||
return False
|
||||
endpoint = scope.get("endpoint")
|
||||
# Identity check against True (not truthiness): the marker is set to the literal
|
||||
# True, and this keeps a spec'd Mock request (whose attribute access yields truthy
|
||||
# child mocks) from being misread as a pass-through dispatch.
|
||||
return getattr(endpoint, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, False) is True
|
||||
|
||||
|
||||
def get_model_from_request(
|
||||
request_data: dict,
|
||||
route: str,
|
||||
request_headers: Optional[Mapping[str, Any]] = None,
|
||||
request_query_params: Optional[Mapping[str, Any]] = None,
|
||||
llm_router: Optional[Router] = None,
|
||||
request: Request | None = None,
|
||||
) -> Optional[Union[str, List[str]]]:
|
||||
"""Resolve the model(s) a request targets, for model-access and budget checks.
|
||||
|
||||
Returns ``None`` when the request was dispatched to a user-defined pass-through
|
||||
endpoint: its body is forwarded verbatim to the configured upstream, so a
|
||||
``model`` field there names an upstream model, not a LiteLLM-managed one, and
|
||||
enforcing key/team model allowlists against it would reject valid requests. The
|
||||
check reads the FastAPI-resolved endpoint (``request.scope["endpoint"]``), not the
|
||||
request path, so a custom path that collides with a built-in route never
|
||||
suppresses model-access checks: on a collision the built-in handler is dispatched
|
||||
and does not carry the marker. Built-in provider passthrough routes
|
||||
(``/vertex_ai``, ``/gemini``, ...) are separate handlers and keep model enforcement.
|
||||
"""
|
||||
if _request_dispatched_to_pass_through_endpoint(request):
|
||||
return None
|
||||
|
||||
candidates = _extract_model_candidates_from_request(
|
||||
request_data=request_data,
|
||||
route=route,
|
||||
|
|
|
|||
|
|
@ -162,6 +162,7 @@ def _get_model_from_request_context(
|
|||
request_headers=_safe_get_request_headers(request=request),
|
||||
request_query_params=_safe_get_request_query_params(request=request),
|
||||
llm_router=llm_router,
|
||||
request=request,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -68,6 +68,7 @@ from litellm.secret_managers.main import get_secret_str
|
|||
from litellm.types.llms.custom_http import httpxSpecialProvider
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
||||
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
|
||||
LITELLM_PASS_THROUGH_ENDPOINT_MARKER,
|
||||
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
|
||||
EndpointType,
|
||||
PassthroughStandardLoggingPayload,
|
||||
|
|
@ -1771,6 +1772,7 @@ def create_pass_through_route(
|
|||
if hasattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY):
|
||||
delattr(request.state, LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY)
|
||||
|
||||
setattr(endpoint_func, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, True)
|
||||
return endpoint_func
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -11,6 +11,14 @@ LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY = "litellm_pass_through_custom_body"
|
|||
# exact byte/string body, such as AWS SigV4-signed requests.
|
||||
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY = "litellm_pass_through_raw_body"
|
||||
|
||||
# Attribute set on the FastAPI endpoint function of every user-defined pass-through
|
||||
# route. Auth reads it off the dispatched endpoint (``request.scope["endpoint"]``) to
|
||||
# decide whether a request body ``model`` names an upstream model rather than a
|
||||
# LiteLLM-managed one. Keying off the resolved endpoint (not the request path) means a
|
||||
# custom path that collides with a built-in route never suppresses model-access checks:
|
||||
# on a collision FastAPI dispatches the built-in handler, which does not carry this flag.
|
||||
LITELLM_PASS_THROUGH_ENDPOINT_MARKER = "__litellm_pass_through_endpoint__"
|
||||
|
||||
|
||||
class EndpointType(str, Enum):
|
||||
VERTEX_AI = "vertex-ai"
|
||||
|
|
|
|||
|
|
@ -2047,6 +2047,88 @@ async def test_common_checks_metadata_route_keeps_key_tags_out_of_provider_metad
|
|||
assert "metadata" not in request_body
|
||||
|
||||
|
||||
def _pass_through_request() -> "Request":
|
||||
"""A Request whose FastAPI-resolved endpoint carries the pass-through marker,
|
||||
i.e. the request was dispatched to a user-defined pass-through handler."""
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
||||
LITELLM_PASS_THROUGH_ENDPOINT_MARKER,
|
||||
)
|
||||
|
||||
def pass_through_endpoint():
|
||||
...
|
||||
|
||||
setattr(pass_through_endpoint, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, True)
|
||||
return Request(scope={"type": "http", "headers": [], "endpoint": pass_through_endpoint})
|
||||
|
||||
|
||||
def _builtin_request() -> "Request":
|
||||
"""A Request dispatched to a built-in (non-pass-through) handler, e.g. what a
|
||||
custom path colliding with a core route actually resolves to."""
|
||||
from fastapi import Request
|
||||
|
||||
def chat_completions():
|
||||
...
|
||||
|
||||
return Request(scope={"type": "http", "headers": [], "endpoint": chat_completions})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_common_checks_auth_enforced_pass_through_ignores_upstream_model():
|
||||
"""An auth-enforced (`auth: true`) user-defined pass-through endpoint must
|
||||
authenticate the key but forward the body unchanged; a body `model` naming an
|
||||
upstream-only model must not be rejected against the team/key model allowlist
|
||||
when the request was dispatched to the pass-through handler. The same body on a
|
||||
request dispatched to a built-in handler (e.g. a path collision) must still be
|
||||
enforced."""
|
||||
from litellm.proxy.auth.auth_checks import common_checks
|
||||
|
||||
team_object = LiteLLM_TeamTable(team_id="team-1", models=["gpt-4o"])
|
||||
valid_token = UserAPIKeyAuth(
|
||||
token="test-token",
|
||||
team_id="team-1",
|
||||
models=[],
|
||||
metadata={"allowed_passthrough_routes": ["/my-custom-endpoint"]},
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
|
||||
new_callable=AsyncMock,
|
||||
return_value={},
|
||||
):
|
||||
result = await common_checks(
|
||||
request_body={"model": "upstream-special-model", "prompt": "hi"},
|
||||
team_object=team_object,
|
||||
user_object=None,
|
||||
end_user_object=None,
|
||||
global_proxy_spend=None,
|
||||
general_settings={},
|
||||
route="/my-custom-endpoint",
|
||||
llm_router=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
valid_token=valid_token,
|
||||
request=_pass_through_request(),
|
||||
)
|
||||
assert result is True
|
||||
|
||||
with pytest.raises(ProxyException) as exc_info:
|
||||
await common_checks(
|
||||
request_body={"model": "upstream-special-model", "prompt": "hi"},
|
||||
team_object=team_object,
|
||||
user_object=None,
|
||||
end_user_object=None,
|
||||
global_proxy_spend=None,
|
||||
general_settings={},
|
||||
route="/v1/chat/completions",
|
||||
llm_router=None,
|
||||
proxy_logging_obj=MagicMock(),
|
||||
valid_token=valid_token,
|
||||
request=_builtin_request(),
|
||||
)
|
||||
assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_virtual_key_soft_budget_check_with_user_obj():
|
||||
"""Test _virtual_key_soft_budget_check includes user_email when user_obj is provided"""
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from typing import Optional
|
|||
from unittest.mock import MagicMock, patch
|
||||
|
||||
import pytest
|
||||
from fastapi import Request
|
||||
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_utils import (
|
||||
|
|
@ -331,6 +332,70 @@ class TestGetEndUserIdFromRequestBodyWithStandardHeaders:
|
|||
assert result == "body-user"
|
||||
|
||||
|
||||
def _request_dispatched_to(endpoint) -> Request:
|
||||
"""Build a minimal Request whose FastAPI-resolved endpoint is ``endpoint``,
|
||||
mirroring what Starlette sets in ``scope`` once routing has matched."""
|
||||
return Request(scope={"type": "http", "headers": [], "endpoint": endpoint})
|
||||
|
||||
|
||||
def _pass_through_endpoint():
|
||||
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
||||
LITELLM_PASS_THROUGH_ENDPOINT_MARKER,
|
||||
)
|
||||
|
||||
def endpoint(): # stand-in for create_pass_through_route's handler
|
||||
...
|
||||
|
||||
setattr(endpoint, LITELLM_PASS_THROUGH_ENDPOINT_MARKER, True)
|
||||
return endpoint
|
||||
|
||||
|
||||
def test_get_model_from_request_skips_pass_through_dispatched_request():
|
||||
"""When FastAPI dispatched the request to a user-defined pass-through handler,
|
||||
the body `model` names an upstream model and must not be treated as a LiteLLM
|
||||
model for allowlist/budget enforcement."""
|
||||
assert (
|
||||
get_model_from_request(
|
||||
request_data={"model": "upstream-special-model"},
|
||||
route="/my-custom-endpoint",
|
||||
request=_request_dispatched_to(_pass_through_endpoint()),
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
def test_get_model_from_request_enforces_when_builtin_handler_dispatched():
|
||||
"""A custom pass-through path that collides with a built-in route resolves to the
|
||||
built-in handler (no marker), so the body `model` must still be extracted and
|
||||
enforced. Same request path as above, but dispatched to a non-pass-through
|
||||
endpoint: the model must NOT be suppressed."""
|
||||
|
||||
def builtin_chat_completions():
|
||||
...
|
||||
|
||||
assert (
|
||||
get_model_from_request(
|
||||
request_data={"model": "gpt-4o"},
|
||||
route="/v1/chat/completions",
|
||||
request=_request_dispatched_to(builtin_chat_completions),
|
||||
)
|
||||
== "gpt-4o"
|
||||
)
|
||||
|
||||
|
||||
def test_get_model_from_request_no_request_extracts_model():
|
||||
"""Callers without a request object (e.g. budget reservation) still extract the
|
||||
model; the pass-through suppression only applies to a dispatched pass-through
|
||||
handler."""
|
||||
assert (
|
||||
get_model_from_request(
|
||||
request_data={"model": "gpt-4o"},
|
||||
route="/v1/chat/completions",
|
||||
)
|
||||
== "gpt-4o"
|
||||
)
|
||||
|
||||
|
||||
def test_get_model_from_request_supports_google_model_names_with_slashes():
|
||||
assert (
|
||||
get_model_from_request(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue