From 2c9542692f792118b711e585a0d1837e54305c4a Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 2 Aug 2026 00:33:57 +0000 Subject: [PATCH] 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> --- litellm/proxy/auth/user_api_key_auth.py | 35 +++++-- .../proxy/response_api_endpoints/endpoints.py | 4 +- .../proxy/auth/test_auth_checks.py | 3 + .../proxy/auth/test_user_api_key_auth.py | 92 ++++++++++++++++++- 4 files changed, 122 insertions(+), 12 deletions(-) diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index f4d07c1a674..8a53f27f1f5 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -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 diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index 05c36406f36..0927a43d3bf 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -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, diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 5f3b0f36b95..eb7078809e8 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -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 { diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index affaaa3fbf4..6ef724ce438 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -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-