fix(auth): only skip centralized common_checks for requests custom auth actually authenticated

The skip was keyed on the global user_custom_auth being configured, so on a
deployment with custom_auth set, requests authenticated by LiteLLM's own key,
JWT or OAuth2 path also bypassed key/team/budget/guardrail authorization.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-08-02 00:33:57 +00:00
parent 23de7a15d9
commit 2c9542692f
4 changed files with 122 additions and 12 deletions

View file

@ -1005,6 +1005,26 @@ async def _resolve_jwt_to_virtual_key(
return None
CUSTOM_AUTH_REQUEST_STATE_KEY = "litellm_authenticated_via_custom_auth"
def _mark_authenticated_via_custom_auth(request: Request) -> None:
setattr(request.state, CUSTOM_AUTH_REQUEST_STATE_KEY, True)
def was_authenticated_via_custom_auth(request: Request) -> bool:
"""True only when ``user_custom_auth`` itself produced the identity for
*this* request.
Configuring ``custom_auth`` does not mean every request is authenticated by
it: the enterprise wrapper's ``auto``/``off`` modes fall back to LiteLLM's
own key auth, and a custom auth function may return an api key string that
is then resolved against the DB. Those requests carry a normal virtual key
and must still be authorized.
"""
return bool(getattr(getattr(request, "state", None), CUSTOM_AUTH_REQUEST_STATE_KEY, False))
def _ensure_parent_otel_span_on_request_state(request: Request) -> None:
"""Idempotently create the OTEL SERVER span and stash it on
``request.state.parent_otel_span``. Safe to call multiple times.
@ -1121,6 +1141,7 @@ async def _user_api_key_auth_builder(
)
if response is not None and isinstance(response, UserAPIKeyAuth):
validated = UserAPIKeyAuth.model_validate(response)
_mark_authenticated_via_custom_auth(request)
if getattr(litellm, "enable_post_custom_auth_checks", False):
validated = await _run_post_custom_auth_checks(
valid_token=validated,
@ -1136,6 +1157,7 @@ async def _user_api_key_auth_builder(
elif user_custom_auth is not None:
response = await user_custom_auth(request=request, api_key=api_key) # type: ignore
validated = UserAPIKeyAuth.model_validate(response)
_mark_authenticated_via_custom_auth(request)
if getattr(litellm, "enable_post_custom_auth_checks", False):
validated = await _run_post_custom_auth_checks(
valid_token=validated,
@ -2108,10 +2130,12 @@ async def _run_centralized_common_checks(
model-access, budgets, guardrails, org, and vector-store checks.
Invariants:
- ``user_custom_auth`` with ``custom_auth_run_common_checks`` unset
skips the gate — matches the existing custom-auth RPS guarantee.
Custom-auth deployments don't use OAuth2 / DB-fallback paths, so
the skip does not re-open any bypass.
- A request whose identity came from ``user_custom_auth`` skips the
gate when ``custom_auth_run_common_checks`` is unset — matches the
existing custom-auth RPS guarantee. Requests that merely ran on a
deployment where custom auth is *configured* but were authenticated
by LiteLLM's own key / JWT / OAuth2 path (enterprise ``auto``/``off``
modes, custom auth returning an api key string) are still gated.
- ``PROXY_ADMIN`` tokens still run through ``common_checks`` so
team-blocked / team-budget / end-user-budget / tag-budget /
vector-store / tool-allowlist enforcement applies to admin keys
@ -2126,7 +2150,6 @@ async def _run_centralized_common_checks(
prisma_client,
proxy_logging_obj,
user_api_key_cache,
user_custom_auth,
)
# Public routes (e.g. /health/liveness) are exempt from
@ -2164,7 +2187,7 @@ async def _run_centralized_common_checks(
):
return
if user_custom_auth is not None and not general_settings.get("custom_auth_run_common_checks", False):
if was_authenticated_via_custom_auth(request) and not general_settings.get("custom_auth_run_common_checks", False):
return
parent_otel_span = user_api_key_auth_obj.parent_otel_span

View file

@ -1010,12 +1010,12 @@ async def _enforce_responses_ws_first_frame_model_auth(
from litellm.proxy.auth.user_api_key_auth import (
_enforce_key_and_fallback_model_access,
_run_centralized_common_checks,
was_authenticated_via_custom_auth,
)
from litellm.proxy.proxy_server import (
general_settings,
llm_model_list,
master_key,
user_custom_auth,
)
request_data = {"model": model}
@ -1026,7 +1026,7 @@ async def _enforce_responses_ws_first_frame_model_auth(
or general_settings.get("enable_oauth2_proxy_auth", False)
):
return
if user_custom_auth is not None and not general_settings.get("custom_auth_run_common_checks", False):
if was_authenticated_via_custom_auth(request) and not general_settings.get("custom_auth_run_common_checks", False):
return
await _enforce_key_and_fallback_model_access(
valid_token=user_api_key_dict,

View file

@ -2763,8 +2763,11 @@ async def test_custom_auth_common_checks_opt_in():
import litellm.proxy.proxy_server as _proxy_server_mod
from litellm.proxy.auth.user_api_key_auth import _run_centralized_common_checks
from types import SimpleNamespace
valid_token = UserAPIKeyAuth(token="test-token", user_id="u1")
mock_request = MagicMock()
mock_request.state = SimpleNamespace(litellm_authenticated_via_custom_auth=True)
def _attrs(flag, user_custom_auth):
return {

View file

@ -3160,10 +3160,9 @@ async def test_centralized_common_checks_routes_header_tags_to_litellm_metadata(
@pytest.mark.asyncio
async def test_centralized_common_checks_skipped_for_custom_auth_without_flag():
"""Existing RPS guarantee: custom-auth deployments without
custom_auth_run_common_checks must not pay the centralized gate.
Custom-auth paths don't use OAuth2/DB-fallback so this skip does
not widen any bypass."""
"""Existing RPS guarantee: a request authenticated by custom auth on a
deployment without custom_auth_run_common_checks must not pay the
centralized gate."""
import litellm.proxy.proxy_server as _proxy_server_mod
from fastapi import Request
from starlette.datastructures import URL
@ -3171,6 +3170,7 @@ async def test_centralized_common_checks_skipped_for_custom_auth_without_flag():
token = UserAPIKeyAuth(api_key="sk-test", user_id="u1")
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
request.state.litellm_authenticated_via_custom_auth = True
attrs = _proxy_attrs_for_centralized_checks(
user_custom_auth=AsyncMock(), flag=False
@ -3208,6 +3208,7 @@ async def test_centralized_common_checks_runs_for_custom_auth_with_flag():
request._url = URL(url="/chat/completions")
attrs = _proxy_attrs_for_centralized_checks(user_custom_auth=AsyncMock(), flag=True)
request.state.litellm_authenticated_via_custom_auth = True
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
try:
for k, v in attrs.items():
@ -3228,6 +3229,89 @@ async def test_centralized_common_checks_runs_for_custom_auth_with_flag():
setattr(_proxy_server_mod, k, v)
@pytest.mark.asyncio
async def test_centralized_common_checks_run_when_custom_auth_configured_but_unused():
"""Configuring custom_auth must not disable authorization for requests
that custom auth did not authenticate (enterprise auto/off fallback,
custom auth returning an api key string, JWT/OAuth2 requests). Those
carry a normal virtual key, so common_checks must still run."""
import litellm.proxy.proxy_server as _proxy_server_mod
from fastapi import Request
from starlette.datastructures import URL
token = UserAPIKeyAuth(api_key="sk-test", user_id="u1")
request = Request(scope={"type": "http"})
request._url = URL(url="/chat/completions")
attrs = _proxy_attrs_for_centralized_checks(
user_custom_auth=AsyncMock(), flag=False
)
originals = {a: getattr(_proxy_server_mod, a, None) for a in attrs}
try:
for k, v in attrs.items():
setattr(_proxy_server_mod, k, v)
with patch(
"litellm.proxy.auth.user_api_key_auth.common_checks",
new_callable=AsyncMock,
) as mock_checks:
await _run_centralized_common_checks(
user_api_key_auth_obj=token,
request=request,
request_data={"model": "gpt-4o"},
route="/chat/completions",
)
mock_checks.assert_awaited_once()
finally:
for k, v in originals.items():
setattr(_proxy_server_mod, k, v)
@pytest.mark.asyncio
async def test_custom_auth_identity_marks_request_state():
"""The builder must record that custom auth produced the identity, so the
centralized gate can tell custom-auth requests apart from key-auth ones."""
import litellm.proxy.proxy_server as _proxy_server_mod
from fastapi import Request
from starlette.datastructures import URL
from litellm.proxy.auth.user_api_key_auth import (
_user_api_key_auth_builder,
was_authenticated_via_custom_auth,
)
request = Request(scope={"type": "http", "headers": []})
request._url = URL(url="/chat/completions")
async def _custom_auth(request, api_key):
return UserAPIKeyAuth(api_key=api_key, user_id="custom-user")
originals = {
"user_custom_auth": getattr(_proxy_server_mod, "user_custom_auth", None),
"general_settings": getattr(_proxy_server_mod, "general_settings", {}),
}
try:
_proxy_server_mod.user_custom_auth = _custom_auth
_proxy_server_mod.general_settings = {}
with patch(
"litellm.proxy.auth.user_api_key_auth.enterprise_custom_auth", None
):
result = await _user_api_key_auth_builder(
request=request,
api_key="Bearer sk-custom",
azure_api_key_header="",
anthropic_api_key_header=None,
google_ai_studio_api_key_header=None,
azure_apim_header=None,
request_data={"model": "gpt-4o"},
)
finally:
for k, v in originals.items():
setattr(_proxy_server_mod, k, v)
assert result.user_id == "custom-user"
assert was_authenticated_via_custom_auth(request) is True
@pytest.mark.asyncio
async def test_centralized_common_checks_runs_for_oauth2_fallback_token():
"""VERIA-18 regression: an OAuth2 token that would previously early-