mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
23de7a15d9
commit
2c9542692f
4 changed files with 122 additions and 12 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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-
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue