mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
feat(proxy): warn at startup when custom_auth skips common_checks enforcement (#30665)
When general_settings.custom_auth is configured but custom_auth_run_common_checks is not set, project/team/org enforcement (budgets, model-level rate limits, and model-access lists) silently does nothing for custom-auth requests, since the centralized common_checks gate returns early for custom auth. Emit a startup warning pointing operators at the flag so the misconfiguration is visible instead of failing silently.
This commit is contained in:
parent
43dadc5138
commit
ba29657d09
3 changed files with 121 additions and 0 deletions
|
|
@ -2,6 +2,7 @@ import os
|
|||
import re
|
||||
import sys
|
||||
from functools import lru_cache
|
||||
from logging import Logger
|
||||
from typing import Any, Dict, FrozenSet, List, Mapping, Optional, Tuple, Union
|
||||
|
||||
from fastapi import HTTPException, Request, status
|
||||
|
|
@ -995,6 +996,45 @@ def get_project_model_tpm_limit(
|
|||
return None
|
||||
|
||||
|
||||
def custom_auth_common_checks_warning(
|
||||
*,
|
||||
custom_auth_configured: bool,
|
||||
run_common_checks: bool,
|
||||
) -> str | None:
|
||||
if not custom_auth_configured or run_common_checks:
|
||||
return None
|
||||
return (
|
||||
"custom_auth is configured but 'custom_auth_run_common_checks' is not set. "
|
||||
"Problem: budgets, model-access allowlists, and per-model rate limits configured "
|
||||
"on your DB team/project records will NOT be enforced for custom-auth requests "
|
||||
"(rate limits set directly on the returned UserAPIKeyAuth still apply). "
|
||||
"Fix: set 'general_settings.custom_auth_run_common_checks: true'. "
|
||||
"Docs: https://docs.litellm.ai/docs/proxy/custom_auth"
|
||||
)
|
||||
|
||||
|
||||
_custom_auth_common_checks_warning_emitted = False
|
||||
|
||||
|
||||
def warn_once_if_custom_auth_skips_common_checks(
|
||||
*,
|
||||
custom_auth_configured: bool,
|
||||
run_common_checks: bool,
|
||||
logger: Logger = verbose_proxy_logger,
|
||||
) -> None:
|
||||
global _custom_auth_common_checks_warning_emitted
|
||||
if _custom_auth_common_checks_warning_emitted:
|
||||
return
|
||||
message = custom_auth_common_checks_warning(
|
||||
custom_auth_configured=custom_auth_configured,
|
||||
run_common_checks=run_common_checks,
|
||||
)
|
||||
if message is None:
|
||||
return
|
||||
logger.warning(message)
|
||||
_custom_auth_common_checks_warning_emitted = True
|
||||
|
||||
|
||||
def is_pass_through_provider_route(route: str) -> bool:
|
||||
PROVIDER_SPECIFIC_PASS_THROUGH_ROUTES = [
|
||||
"vertex-ai",
|
||||
|
|
|
|||
|
|
@ -260,6 +260,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
from litellm.proxy.auth.auth_utils import (
|
||||
check_response_size_is_safe,
|
||||
is_request_body_safe,
|
||||
warn_once_if_custom_auth_skips_common_checks,
|
||||
)
|
||||
from litellm.proxy.auth.handle_jwt import JWTHandler
|
||||
from litellm.proxy.auth.litellm_license import LicenseCheck
|
||||
|
|
@ -4373,6 +4374,12 @@ class ProxyConfig:
|
|||
user_custom_auth = get_instance_fn(
|
||||
value=custom_auth, config_file_path=config_file_path
|
||||
)
|
||||
warn_once_if_custom_auth_skips_common_checks(
|
||||
custom_auth_configured=custom_auth is not None,
|
||||
run_common_checks=bool(
|
||||
general_settings.get("custom_auth_run_common_checks", False)
|
||||
),
|
||||
)
|
||||
|
||||
custom_key_generate = general_settings.get("custom_key_generate", None)
|
||||
if custom_key_generate is not None:
|
||||
|
|
|
|||
|
|
@ -13,6 +13,8 @@ from litellm.proxy.auth.auth_utils import (
|
|||
_get_customer_id_from_standard_headers,
|
||||
abbreviate_api_key,
|
||||
check_complete_credentials,
|
||||
custom_auth_common_checks_warning,
|
||||
warn_once_if_custom_auth_skips_common_checks,
|
||||
get_end_user_id_from_request_body,
|
||||
get_key_mcp_rpm_limit,
|
||||
get_key_model_rpm_limit,
|
||||
|
|
@ -25,6 +27,78 @@ from litellm.proxy.auth.auth_utils import (
|
|||
)
|
||||
|
||||
|
||||
class TestCustomAuthCommonChecksWarning:
|
||||
"""custom_auth_common_checks_warning only warns when custom auth is configured
|
||||
and the common-checks opt-in is off, since that is the only state where
|
||||
project/team enforcement silently does nothing."""
|
||||
|
||||
def test_warns_when_custom_auth_configured_and_checks_off(self):
|
||||
warning = custom_auth_common_checks_warning(
|
||||
custom_auth_configured=True,
|
||||
run_common_checks=False,
|
||||
)
|
||||
assert warning is not None
|
||||
assert "custom_auth_run_common_checks: true" in warning
|
||||
assert "https://docs.litellm.ai/docs/proxy/custom_auth" in warning
|
||||
|
||||
def test_no_warning_when_common_checks_enabled(self):
|
||||
assert (
|
||||
custom_auth_common_checks_warning(
|
||||
custom_auth_configured=True,
|
||||
run_common_checks=True,
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
def test_no_warning_when_custom_auth_not_configured(self):
|
||||
assert (
|
||||
custom_auth_common_checks_warning(
|
||||
custom_auth_configured=False,
|
||||
run_common_checks=False,
|
||||
)
|
||||
is None
|
||||
)
|
||||
assert (
|
||||
custom_auth_common_checks_warning(
|
||||
custom_auth_configured=False,
|
||||
run_common_checks=True,
|
||||
)
|
||||
is None
|
||||
)
|
||||
|
||||
|
||||
class TestWarnOnceIfCustomAuthSkipsCommonChecks:
|
||||
"""The startup warning must fire at most once per process, since load_config
|
||||
re-runs on hot-reload / config refresh and would otherwise spam the log."""
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _reset_sentinel(self, monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.auth.auth_utils._custom_auth_common_checks_warning_emitted",
|
||||
False,
|
||||
)
|
||||
|
||||
def test_warns_only_once_across_repeated_calls(self):
|
||||
logger = MagicMock()
|
||||
for _ in range(3):
|
||||
warn_once_if_custom_auth_skips_common_checks(
|
||||
custom_auth_configured=True,
|
||||
run_common_checks=False,
|
||||
logger=logger,
|
||||
)
|
||||
assert logger.warning.call_count == 1
|
||||
assert "custom_auth_run_common_checks" in logger.warning.call_args[0][0]
|
||||
|
||||
def test_does_not_warn_when_common_checks_enabled(self):
|
||||
logger = MagicMock()
|
||||
warn_once_if_custom_auth_skips_common_checks(
|
||||
custom_auth_configured=True,
|
||||
run_common_checks=True,
|
||||
logger=logger,
|
||||
)
|
||||
assert logger.warning.call_count == 0
|
||||
|
||||
|
||||
class TestGetKeyModelRpmLimit:
|
||||
"""Tests for get_key_model_rpm_limit function."""
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue