mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-08 22:21:35 +00:00
* fix(logging): redact credential query params from the uvicorn access log Raw virtual keys reached container stdout two ways: - `GET /key/info?key=sk-...`, `/global/spend/report?api_key=sk-...`, `/key/spend/report`, `/spend/logs`, `/user/daily/activity` and the Gemini passthrough routes all put the credential in the request target, and `uvicorn.access` had no redaction filter (only `uvicorn.error` did). - the key budget error interpolates `LiteLLM_VerificationToken.key_name`, a column with no enforced shape, into a message that is both logged and returned to the caller. `SecretRedactionFilter` cannot be reused on an access logger: it collapses the record into `record.msg` and clears `record.args`, and uvicorn's AccessFormatter unpacks those args at emit time, so every access line would raise TypeError. `AccessLogRedactionFilter` scrubs the positional args in place instead. An access line is the one input to the secret regex an unauthenticated caller controls end to end, so two bounds go with it. The request target is cut back to a whole query parameter under 512 characters before it is scanned, since a half parameter is too short to match its own pattern and would be logged raw, and the dropped tail is not logged at all. The connection-string pattern is bounded too, because its user half could previously re-scan the rest of the string from every `://`: a 16 KB URL of `a://` pairs took 314s and now takes 0.12s, with the caps set high enough that an RDS IAM auth token used as a DSN password still redacts. Credential query params are terminated by `&` like the existing `key=` and `sig=` patterns, so redacting one param no longer swallows the rest of the request line, and a second credential in the same query string is now redacted on its own instead of surviving once the first one stops the span. `key_name` is echoed into the budget error only when it still has the masked `sk-...abcd` shape `abbreviate_api_key` writes, so a value put there by a direct DB write or a migration falls back to the key alias. Also point the `/key/info` and spend-report examples at the sha256 hash both endpoints already accept, so callers stop putting raw keys in URLs that third-party access logs record. Resolves LIT-5909 * test(logging): assert on emitted access lines instead of filter registration The two registration tests checked that an AccessLogRedactionFilter instance sits in uvicorn.access.filters, which is the shape of the code rather than its behavior. Handing the logger a real access record and reading what a handler wrote covers the same wiring and still fails when the registration is removed. * fix(logging): redact percent-encoded credentials from access logs ?k%65y=sk%2D... is a working credential once the request parser decodes it, but the redaction patterns match literal text and never see it. Decode the request target as a detector and drop the query when decoding reveals a secret. The decoded text is never logged back, so a %0A cannot forge a following log line Also accept any four non-space characters in the masked key_name check, since abbreviate_api_key copies the last four characters of a custom key verbatim and those can be punctuation or non-ASCII * fix(auth): keep control codes out of the masked key label /key/generate accepts a custom key ending in an escape sequence, and abbreviate_api_key copies those four characters into key_name verbatim, so the over-budget message carried them to a terminal and a log viewer. Bar whitespace and C0/C1 control codes from the four, and keep everything else
7511 lines
267 KiB
Python
7511 lines
267 KiB
Python
import asyncio
|
|
import json
|
|
from types import SimpleNamespace
|
|
from typing import TYPE_CHECKING, Optional
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
|
|
|
if TYPE_CHECKING:
|
|
from litellm.router import Router
|
|
|
|
|
|
from datetime import datetime, timedelta, timezone
|
|
|
|
import httpx
|
|
import pytest
|
|
from fastapi import Request, status
|
|
|
|
import litellm
|
|
from litellm.proxy._types import (
|
|
CallInfo,
|
|
Litellm_EntityType,
|
|
LiteLLM_BudgetTable,
|
|
LiteLLM_EndUserTable,
|
|
LiteLLM_ObjectPermissionTable,
|
|
LiteLLM_TagTable,
|
|
LiteLLM_TeamTable,
|
|
LiteLLM_UserTable,
|
|
LitellmUserRoles,
|
|
ProxyErrorTypes,
|
|
ProxyException,
|
|
SSOUserDefinedValues,
|
|
UserAPIKeyAuth,
|
|
)
|
|
from litellm.proxy.auth.auth_checks import (
|
|
ExperimentalUIJWTToken,
|
|
_cache_management_object,
|
|
_can_object_call_model,
|
|
_can_object_call_vector_stores,
|
|
_check_end_user_budget,
|
|
_check_team_member_budget,
|
|
_get_fuzzy_user_object,
|
|
_get_team_db_check,
|
|
_log_budget_lookup_failure,
|
|
_tag_max_budget_check,
|
|
_team_max_budget_check,
|
|
_virtual_key_max_budget_alert_check,
|
|
_virtual_key_max_budget_check,
|
|
_virtual_key_soft_budget_check,
|
|
get_key_object,
|
|
get_user_object,
|
|
invalidate_team_member_spend_state,
|
|
vector_store_access_check,
|
|
)
|
|
from litellm.caching.in_memory_cache import InMemoryCache
|
|
from litellm.caching.redis_cache import RedisCache
|
|
from litellm.constants import (
|
|
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
|
|
END_USER_RESTRICTED_REGISTRY_MAX_SIZE,
|
|
REGISTRY_ERROR_NEGATIVE_CACHE_TTL,
|
|
TAG_REGISTRY_MAX_SIZE,
|
|
)
|
|
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
|
|
from litellm.proxy.common_utils.user_api_key_cache import (
|
|
END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL,
|
|
TAG_REGISTRY_OVERFLOW_SENTINEL,
|
|
UserApiKeyCache,
|
|
end_user_cache_key,
|
|
end_user_restricted_registry_cache_key,
|
|
tag_cache_key,
|
|
tag_registry_cache_key,
|
|
)
|
|
from litellm.utils import get_utc_datetime
|
|
|
|
|
|
def _rendered_log_message(call):
|
|
message = str(call.args[0])
|
|
values = call.args[1:]
|
|
return message % values if values else message
|
|
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def set_salt_key(monkeypatch):
|
|
"""Automatically set LITELLM_SALT_KEY for all tests"""
|
|
monkeypatch.setenv("LITELLM_SALT_KEY", "sk-1234")
|
|
|
|
|
|
@pytest.fixture(autouse=True)
|
|
def reset_constants_module():
|
|
"""Reset constants module to ensure clean state before each test"""
|
|
import importlib
|
|
|
|
from litellm import constants
|
|
from litellm.proxy.auth import auth_checks
|
|
|
|
# Reload modules before test
|
|
importlib.reload(constants)
|
|
importlib.reload(auth_checks)
|
|
|
|
yield
|
|
|
|
# Reload modules after test to clean up
|
|
importlib.reload(constants)
|
|
importlib.reload(auth_checks)
|
|
|
|
|
|
@pytest.fixture
|
|
def valid_sso_user_defined_values():
|
|
return LiteLLM_UserTable(
|
|
user_id="test_user",
|
|
user_email="test@example.com",
|
|
user_role=LitellmUserRoles.PROXY_ADMIN.value,
|
|
models=["gpt-3.5-turbo"],
|
|
max_budget=100.0,
|
|
)
|
|
|
|
|
|
@pytest.fixture
|
|
def invalid_sso_user_defined_values():
|
|
return LiteLLM_UserTable(
|
|
user_id="test_user",
|
|
user_email="test@example.com",
|
|
user_role=None, # Missing user role
|
|
models=["gpt-3.5-turbo"],
|
|
max_budget=100.0,
|
|
)
|
|
|
|
|
|
def test_get_experimental_ui_login_jwt_auth_token_valid(valid_sso_user_defined_values):
|
|
"""Test generating JWT token with valid user role"""
|
|
token = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(
|
|
valid_sso_user_defined_values
|
|
)
|
|
|
|
# Decrypt and verify token contents
|
|
decrypted_token = decrypt_value_helper(
|
|
token, key="ui_hash_key", exception_type="debug"
|
|
)
|
|
# Check that decrypted_token is not None before using json.loads
|
|
assert decrypted_token is not None
|
|
token_data = json.loads(decrypted_token)
|
|
|
|
assert token_data["user_id"] == "test_user"
|
|
assert token_data["user_role"] == LitellmUserRoles.PROXY_ADMIN.value
|
|
assert token_data["models"] == ["gpt-3.5-turbo"]
|
|
assert token_data["max_budget"] == litellm.max_ui_session_budget
|
|
|
|
# Verify expiration time is set and valid (Experimental UI uses fixed 10-min expiry)
|
|
assert "expires" in token_data
|
|
expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00"))
|
|
now = get_utc_datetime()
|
|
# Allow 2 second buffer for test execution timing
|
|
assert expires > now
|
|
assert expires <= now + timedelta(minutes=10, seconds=2)
|
|
|
|
|
|
def test_get_cli_jwt_auth_token_includes_team_alias(valid_sso_user_defined_values):
|
|
token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(
|
|
valid_sso_user_defined_values,
|
|
team_id="team-123",
|
|
team_alias="test-team",
|
|
)
|
|
|
|
decrypted_token = decrypt_value_helper(
|
|
token, key="ui_hash_key", exception_type="debug"
|
|
)
|
|
assert decrypted_token is not None
|
|
token_data = json.loads(decrypted_token)
|
|
|
|
assert token_data["team_id"] == "team-123"
|
|
assert token_data["team_alias"] == "test-team"
|
|
|
|
|
|
def test_get_cli_jwt_auth_token_carries_team_grants_not_user_allowlist(
|
|
valid_sso_user_defined_values,
|
|
):
|
|
"""A team-bound `lite login` session token must snapshot the team's grants.
|
|
|
|
Without team_models the /v1/models bail-out (`not key_models and not team_models`)
|
|
treats the session as unrestricted and lists the whole proxy; without
|
|
team_model_aliases a team alias never resolves on /chat/completions. The user's
|
|
personal allowlist must stay out of the key `models` slot, since a team-bound
|
|
credential is governed by the team grant, not by a per-user list.
|
|
"""
|
|
token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(
|
|
valid_sso_user_defined_values,
|
|
team_id="team-123",
|
|
team_alias="test-team",
|
|
team_models=("claude-sonnet-4-5", "gpt-4.1"),
|
|
team_model_aliases={"team-fast": "gpt-4.1-mini"},
|
|
)
|
|
|
|
decrypted_token = decrypt_value_helper(
|
|
token, key="ui_hash_key", exception_type="debug"
|
|
)
|
|
assert decrypted_token is not None
|
|
token_data = json.loads(decrypted_token)
|
|
|
|
assert token_data["team_id"] == "team-123"
|
|
assert token_data["team_models"] == ["claude-sonnet-4-5", "gpt-4.1"]
|
|
assert token_data["team_model_aliases"] == {"team-fast": "gpt-4.1-mini"}
|
|
assert valid_sso_user_defined_values.models == ["gpt-3.5-turbo"]
|
|
assert token_data["models"] == []
|
|
|
|
|
|
def test_get_cli_jwt_auth_token_keeps_user_allowlist_when_no_team(
|
|
valid_sso_user_defined_values,
|
|
):
|
|
"""A session token with no team bound still carries the user's own allowlist."""
|
|
token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
|
|
|
|
decrypted_token = decrypt_value_helper(
|
|
token, key="ui_hash_key", exception_type="debug"
|
|
)
|
|
assert decrypted_token is not None
|
|
token_data = json.loads(decrypted_token)
|
|
|
|
assert token_data.get("team_id") is None
|
|
assert token_data["models"] == ["gpt-3.5-turbo"]
|
|
assert token_data["team_models"] == []
|
|
|
|
|
|
def test_get_experimental_ui_login_jwt_auth_token_uses_10_min_expiry(
|
|
valid_sso_user_defined_values,
|
|
):
|
|
"""Test that Experimental UI token uses fixed 10-minute expiry (does not use LITELLM_UI_SESSION_DURATION)."""
|
|
token = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(
|
|
valid_sso_user_defined_values
|
|
)
|
|
decrypted_token = decrypt_value_helper(
|
|
token, key="ui_hash_key", exception_type="debug"
|
|
)
|
|
assert decrypted_token is not None
|
|
token_data = json.loads(decrypted_token)
|
|
expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00"))
|
|
now = get_utc_datetime()
|
|
# Should expire in ~10 minutes (allow 2 second buffer)
|
|
assert expires > now + timedelta(minutes=9)
|
|
assert expires <= now + timedelta(minutes=10, seconds=2)
|
|
|
|
|
|
def test_experimental_ui_token_ignores_litellm_ui_session_duration(
|
|
valid_sso_user_defined_values,
|
|
):
|
|
"""Regression test: LITELLM_UI_SESSION_DURATION must NOT affect Experimental UI token expiry.
|
|
Experimental UI intentionally uses fixed 10-min expiry. If this test fails, the constant
|
|
was incorrectly wired to the experimental flow."""
|
|
# Default LITELLM_UI_SESSION_DURATION is "24h" - token must still expire in ~10 min
|
|
token = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(
|
|
valid_sso_user_defined_values
|
|
)
|
|
decrypted_token = decrypt_value_helper(
|
|
token, key="ui_hash_key", exception_type="debug"
|
|
)
|
|
assert decrypted_token is not None
|
|
token_data = json.loads(decrypted_token)
|
|
expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00"))
|
|
now = get_utc_datetime()
|
|
# Must be ~10 min, NOT 24h. If LITELLM_UI_SESSION_DURATION were incorrectly used, this would fail.
|
|
assert expires <= now + timedelta(
|
|
minutes=11
|
|
), "Experimental UI must use 10-min expiry, not LITELLM_UI_SESSION_DURATION"
|
|
|
|
|
|
def test_get_experimental_ui_login_jwt_auth_token_invalid(
|
|
invalid_sso_user_defined_values,
|
|
):
|
|
"""Test generating JWT token with missing user role"""
|
|
with pytest.raises(Exception, match='User role is required for experimental UI login') as exc_info:
|
|
ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(
|
|
invalid_sso_user_defined_values
|
|
)
|
|
|
|
assert str(exc_info.value) == "User role is required for experimental UI login"
|
|
|
|
|
|
def test_get_key_object_from_ui_hash_key_valid(
|
|
valid_sso_user_defined_values, monkeypatch
|
|
):
|
|
"""Test getting key object from valid UI hash key"""
|
|
monkeypatch.setenv("EXPERIMENTAL_UI_LOGIN", "True")
|
|
# Generate a valid token
|
|
token = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(
|
|
valid_sso_user_defined_values
|
|
)
|
|
|
|
# Get key object
|
|
key_object = ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(token)
|
|
|
|
assert key_object is not None
|
|
assert key_object.user_id == "test_user"
|
|
assert key_object.user_role == LitellmUserRoles.PROXY_ADMIN
|
|
assert key_object.models == ["gpt-3.5-turbo"]
|
|
assert key_object.max_budget == litellm.max_ui_session_budget
|
|
|
|
|
|
def test_get_key_object_from_ui_hash_key_invalid():
|
|
"""Test getting key object from invalid UI hash key"""
|
|
# Test with invalid token
|
|
key_object = ExperimentalUIJWTToken.get_key_object_from_ui_hash_key("invalid_token")
|
|
assert key_object is None
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"object_type,expected_error_type",
|
|
[
|
|
("key", ProxyErrorTypes.key_model_access_denied),
|
|
("team", ProxyErrorTypes.team_model_access_denied),
|
|
("user", ProxyErrorTypes.user_model_access_denied),
|
|
("org", ProxyErrorTypes.org_model_access_denied),
|
|
("project", ProxyErrorTypes.project_model_access_denied),
|
|
],
|
|
)
|
|
def test_can_object_call_model_denials_return_forbidden(
|
|
object_type, expected_error_type
|
|
):
|
|
with pytest.raises(ProxyException) as exc_info:
|
|
_can_object_call_model(
|
|
model="restricted-model",
|
|
llm_router=None,
|
|
models=["allowed-model"],
|
|
object_type=object_type,
|
|
)
|
|
|
|
assert exc_info.value.type == expected_error_type
|
|
assert int(exc_info.value.code) == status.HTTP_403_FORBIDDEN
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_can_user_call_model_no_default_models_returns_forbidden():
|
|
from litellm.proxy._types import SpecialModelNames
|
|
from litellm.proxy.auth.auth_checks import can_user_call_model
|
|
|
|
user_object = LiteLLM_UserTable(
|
|
user_id="test-user",
|
|
models=[SpecialModelNames.no_default_models.value],
|
|
)
|
|
|
|
with pytest.raises(ProxyException) as exc_info:
|
|
await can_user_call_model(
|
|
model="restricted-model",
|
|
llm_router=None,
|
|
user_object=user_object,
|
|
)
|
|
|
|
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
|
|
assert int(exc_info.value.code) == status.HTTP_403_FORBIDDEN
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_can_key_call_model_all_team_models_uses_team_allowlist():
|
|
from litellm.proxy._types import SpecialModelNames
|
|
from litellm.proxy.auth.auth_checks import can_key_call_model
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
api_key="sk-team-key",
|
|
team_id="team-123",
|
|
models=[SpecialModelNames.all_team_models.value],
|
|
team_models=["openai/openai/gpt-5.5-batch"],
|
|
)
|
|
|
|
assert (
|
|
await can_key_call_model(
|
|
model="openai/openai/gpt-5.5-batch",
|
|
llm_model_list=None,
|
|
valid_token=valid_token,
|
|
llm_router=None,
|
|
)
|
|
is True
|
|
)
|
|
|
|
with pytest.raises(ProxyException) as exc_info:
|
|
await can_key_call_model(
|
|
model="gpt-4o",
|
|
llm_model_list=None,
|
|
valid_token=valid_token,
|
|
llm_router=None,
|
|
)
|
|
|
|
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_can_key_call_model_all_team_models_empty_team_models_is_unrestricted():
|
|
"""Team-bound key with empty team_models expands to [] -> unrestricted (same as get_key_models)."""
|
|
from litellm.proxy._types import SpecialModelNames
|
|
from litellm.proxy.auth.auth_checks import can_key_call_model
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
api_key="sk-team-key",
|
|
team_id="team-123",
|
|
models=[SpecialModelNames.all_team_models.value],
|
|
team_models=[],
|
|
)
|
|
|
|
assert (
|
|
await can_key_call_model(
|
|
model="any-model",
|
|
llm_model_list=None,
|
|
valid_token=valid_token,
|
|
llm_router=None,
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_can_key_call_model_all_team_models_no_team_id_is_unrestricted():
|
|
"""A teamless key with all-team-models inherits the full proxy model list
|
|
(empty resolved list = unrestricted access), the same as leaving the models
|
|
field empty. This test will fail if someone re-introduces a teamless denial."""
|
|
from litellm.proxy._types import SpecialModelNames
|
|
from litellm.proxy.auth.auth_checks import can_key_call_model
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
api_key="sk-orphan-key",
|
|
models=[SpecialModelNames.all_team_models.value],
|
|
team_models=[],
|
|
)
|
|
|
|
assert (
|
|
await can_key_call_model(
|
|
model="gpt-4o",
|
|
llm_model_list=None,
|
|
valid_token=valid_token,
|
|
llm_router=None,
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
def test_resolve_key_models_teamless_all_team_models_returns_empty():
|
|
"""_resolve_key_models_for_auth_check must return [] for a teamless key
|
|
with all-team-models, making it equivalent to an unscoped key (unrestricted
|
|
access). Fails if someone returns the sentinel list for teamless keys."""
|
|
from litellm.proxy._types import SpecialModelNames
|
|
from litellm.proxy.auth.auth_checks import _resolve_key_models_for_auth_check
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
api_key="sk-orphan",
|
|
models=[SpecialModelNames.all_team_models.value],
|
|
team_models=[],
|
|
)
|
|
|
|
result = _resolve_key_models_for_auth_check(valid_token)
|
|
assert result == [], "teamless all-team-models must resolve to [] (unrestricted)"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_enforce_key_access_teamless_all_team_models_passes():
|
|
"""_enforce_key_and_fallback_model_access must not deny a teamless key with
|
|
all-team-models. The inference path skips the key-level model check when
|
|
the sentinel is present, regardless of team_id. Fails if someone adds a
|
|
team_id guard to the pass branch."""
|
|
from litellm.proxy._types import SpecialModelNames
|
|
from litellm.proxy.auth.user_api_key_auth import _enforce_key_and_fallback_model_access
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
api_key="sk-orphan",
|
|
models=[SpecialModelNames.all_team_models.value],
|
|
team_models=[],
|
|
)
|
|
|
|
await _enforce_key_and_fallback_model_access(
|
|
valid_token=valid_token,
|
|
request_data={"model": "gpt-4o"},
|
|
route="/chat/completions",
|
|
request=None,
|
|
llm_model_list=None,
|
|
llm_router=None,
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_can_key_call_resolved_model_teamless_all_team_models_passes():
|
|
"""can_key_call_resolved_model must skip the key model check for a teamless
|
|
key with all-team-models. Fails if someone adds a team_id guard to the
|
|
skip_key_model_check condition."""
|
|
from unittest.mock import AsyncMock, patch
|
|
|
|
from litellm.proxy._types import SpecialModelNames
|
|
from litellm.proxy.auth.auth_checks import can_key_call_resolved_model
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
api_key="sk-orphan",
|
|
models=[SpecialModelNames.all_team_models.value],
|
|
team_models=[],
|
|
)
|
|
|
|
with patch("litellm.proxy.auth.auth_checks.can_key_call_model", new_callable=AsyncMock) as mock_call:
|
|
with patch("litellm.proxy.proxy_server.prisma_client", None):
|
|
with patch("litellm.proxy.proxy_server.proxy_logging_obj", None):
|
|
with patch("litellm.proxy.proxy_server.user_api_key_cache", None):
|
|
await can_key_call_resolved_model(
|
|
model="gpt-4o",
|
|
llm_model_list=None,
|
|
valid_token=valid_token,
|
|
llm_router=None,
|
|
)
|
|
mock_call.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_can_team_access_model_all_team_models_expands_router_models():
|
|
from litellm import Router
|
|
from litellm.proxy._types import SpecialModelNames
|
|
from litellm.proxy.auth.auth_checks import can_team_access_model
|
|
|
|
team_object = LiteLLM_TeamTable(
|
|
team_id="team-123",
|
|
models=[SpecialModelNames.all_team_models.value],
|
|
)
|
|
router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "allowed-model",
|
|
"litellm_params": {"model": "openai/gpt-4o", "api_key": "sk-test"},
|
|
}
|
|
]
|
|
)
|
|
|
|
assert (
|
|
await can_team_access_model(
|
|
model="allowed-model",
|
|
team_object=team_object,
|
|
llm_router=router,
|
|
)
|
|
is True
|
|
)
|
|
with pytest.raises(ProxyException) as exc_info:
|
|
await can_team_access_model(
|
|
model="blocked-model",
|
|
team_object=team_object,
|
|
llm_router=router,
|
|
)
|
|
|
|
assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_key_object_should_reconnect_once_on_db_connection_error():
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.get_data = AsyncMock(
|
|
side_effect=[
|
|
httpx.ConnectError("db connection reset"),
|
|
UserAPIKeyAuth(token="hashed-token-1"),
|
|
]
|
|
)
|
|
mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True)
|
|
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
key_obj = await get_key_object(
|
|
hashed_token="hashed-token-1",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
)
|
|
|
|
assert key_obj.token == "hashed-token-1"
|
|
assert mock_prisma_client.get_data.await_count == 2
|
|
mock_prisma_client.attempt_db_reconnect.assert_awaited_once_with(
|
|
reason="auth_get_key_object_lookup_failure",
|
|
timeout_seconds=2.0,
|
|
lock_timeout_seconds=0.1,
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_key_object_should_raise_if_reconnect_fails_on_db_connection_error():
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.get_data = AsyncMock(
|
|
side_effect=httpx.ConnectError("db not reachable after outage")
|
|
)
|
|
mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=False)
|
|
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
with pytest.raises(Exception, match="db not reachable after outage"):
|
|
await get_key_object(
|
|
hashed_token="hashed-token-2",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
)
|
|
|
|
mock_prisma_client.attempt_db_reconnect.assert_awaited_once_with(
|
|
reason="auth_get_key_object_lookup_failure",
|
|
timeout_seconds=2.0,
|
|
lock_timeout_seconds=0.1,
|
|
)
|
|
assert mock_prisma_client.get_data.await_count == 1
|
|
|
|
|
|
def _fake_redis_cache():
|
|
fake_redis = MagicMock()
|
|
fake_redis.async_get_cache = AsyncMock(return_value=None)
|
|
fake_redis.async_set_cache = AsyncMock()
|
|
fake_redis.async_set_cache_pipeline = AsyncMock()
|
|
fake_redis.async_delete_cache = AsyncMock()
|
|
return fake_redis
|
|
|
|
|
|
class TestAuthCacheRedisWritePolicy:
|
|
"""Redis auth-cache entries may only be written from fresh DB loads.
|
|
|
|
With ``enable_redis_auth_cache`` and multiple replicas, a pod that re-publishes
|
|
a cache-derived key object to Redis can resurrect a stale auth blob after
|
|
``/key/update`` or ``/key/delete`` already deleted it, so limit changes never
|
|
propagate fleet-wide while traffic keeps refreshing the stale entry's TTL.
|
|
"""
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_key_object_db_load_publishes_to_redis(self):
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.get_data = AsyncMock(
|
|
return_value=UserAPIKeyAuth(token="hashed-token-db")
|
|
)
|
|
|
|
fake_redis = _fake_redis_cache()
|
|
cache = UserApiKeyCache()
|
|
cache.redis_cache = fake_redis
|
|
|
|
key_obj = await get_key_object(
|
|
hashed_token="hashed-token-db",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=cache,
|
|
)
|
|
|
|
assert key_obj.token == "hashed-token-db"
|
|
fake_redis.async_set_cache.assert_awaited_once()
|
|
assert (
|
|
fake_redis.async_set_cache.await_args.kwargs.get("key")
|
|
or fake_redis.async_set_cache.await_args.args[0]
|
|
) == "hashed-token-db"
|
|
|
|
|
|
def test_get_cli_jwt_auth_token_default_expiration(valid_sso_user_defined_values):
|
|
"""Test generating CLI JWT token with default 24-hour expiration"""
|
|
token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
|
|
|
|
# Decrypt and verify token contents
|
|
decrypted_token = decrypt_value_helper(
|
|
token, key="ui_hash_key", exception_type="debug"
|
|
)
|
|
assert decrypted_token is not None
|
|
token_data = json.loads(decrypted_token)
|
|
|
|
assert token_data["user_id"] == "test_user"
|
|
assert token_data["user_role"] == LitellmUserRoles.PROXY_ADMIN.value
|
|
assert token_data["models"] == ["gpt-3.5-turbo"]
|
|
# CLI session tokens carry no per-key budget; spend is enforced via the
|
|
# shared team/user counters. The $0.25 UI session cap must not leak in.
|
|
assert token_data.get("max_budget") is None
|
|
# is_session_token=True causes key_management_endpoints to use the team
|
|
# budget as the delegation ceiling instead of treating None as unlimited.
|
|
assert token_data.get("is_session_token") is True
|
|
|
|
# Verify expiration time is set to 24 hours (default)
|
|
assert "expires" in token_data
|
|
expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00"))
|
|
assert expires > get_utc_datetime()
|
|
assert expires <= get_utc_datetime() + timedelta(hours=24, minutes=1)
|
|
assert expires >= get_utc_datetime() + timedelta(hours=23, minutes=59)
|
|
|
|
|
|
def test_get_cli_jwt_auth_token_custom_expiration(
|
|
valid_sso_user_defined_values, monkeypatch
|
|
):
|
|
"""Test generating CLI JWT token with custom expiration via environment variable"""
|
|
import importlib
|
|
|
|
from litellm import constants
|
|
from litellm.proxy.auth import auth_checks
|
|
|
|
# Set custom expiration to 48 hours
|
|
monkeypatch.setenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", "48")
|
|
|
|
# Reload the constants module to pick up the new env var
|
|
importlib.reload(constants)
|
|
# Also reload auth_checks to pick up the new constant value
|
|
importlib.reload(auth_checks)
|
|
|
|
token = auth_checks.ExperimentalUIJWTToken.get_cli_jwt_auth_token(
|
|
valid_sso_user_defined_values
|
|
)
|
|
|
|
# Decrypt and verify token contents
|
|
decrypted_token = decrypt_value_helper(
|
|
token, key="ui_hash_key", exception_type="debug"
|
|
)
|
|
assert decrypted_token is not None
|
|
token_data = json.loads(decrypted_token)
|
|
|
|
# Verify expiration time is set to 48 hours
|
|
assert "expires" in token_data
|
|
expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00"))
|
|
assert expires > get_utc_datetime() + timedelta(hours=47, minutes=59)
|
|
assert expires <= get_utc_datetime() + timedelta(hours=48, minutes=1)
|
|
|
|
|
|
def test_get_cli_jwt_auth_token_unique_per_session(valid_sso_user_defined_values):
|
|
"""Each CLI login mints a unique token id (per-session spend isolation) while
|
|
keeping a stable, user-scoped key_alias for log grouping. A regression that
|
|
pins token back to a constant would collapse both ids and fail here."""
|
|
from litellm.constants import CLI_SESSION_KEY_PREFIX
|
|
|
|
def _decode(token: str) -> dict:
|
|
decrypted = decrypt_value_helper(
|
|
token, key="ui_hash_key", exception_type="debug"
|
|
)
|
|
assert decrypted is not None
|
|
return json.loads(decrypted)
|
|
|
|
first = _decode(
|
|
ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
|
|
)
|
|
second = _decode(
|
|
ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
|
|
)
|
|
|
|
assert first["token"].startswith(f"{CLI_SESSION_KEY_PREFIX}-")
|
|
assert second["token"].startswith(f"{CLI_SESSION_KEY_PREFIX}-")
|
|
assert first["token"] != second["token"]
|
|
|
|
expected_alias = f"{CLI_SESSION_KEY_PREFIX}-test_user"
|
|
assert first["key_alias"] == second["key_alias"] == expected_alias
|
|
assert first["key_name"] == second["key_name"] == expected_alias
|
|
|
|
|
|
def test_get_cli_jwt_auth_token_applies_fallback_budget(valid_sso_user_defined_values):
|
|
token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(
|
|
valid_sso_user_defined_values, max_budget=litellm.max_ui_session_budget
|
|
)
|
|
decrypted = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug")
|
|
assert decrypted is not None
|
|
assert json.loads(decrypted).get("max_budget") == litellm.max_ui_session_budget
|
|
|
|
|
|
def test_get_cli_jwt_auth_token_no_fallback_when_budget_provided(
|
|
valid_sso_user_defined_values,
|
|
):
|
|
token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(
|
|
valid_sso_user_defined_values, max_budget=None
|
|
)
|
|
decrypted = decrypt_value_helper(token, key="ui_hash_key", exception_type="debug")
|
|
assert decrypted is not None
|
|
assert json.loads(decrypted).get("max_budget") is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_default_internal_user_params_with_get_user_object(monkeypatch):
|
|
"""Test that default_internal_user_params is used when creating a new user via get_user_object"""
|
|
# Set up default_internal_user_params
|
|
default_params = {
|
|
"models": ["gpt-4", "claude-3-opus"],
|
|
"max_budget": 200.0,
|
|
"user_role": "internal_user",
|
|
}
|
|
monkeypatch.setattr(litellm, "default_internal_user_params", default_params)
|
|
|
|
# Mock the necessary dependencies
|
|
mock_prisma_client = MagicMock()
|
|
mock_db = AsyncMock()
|
|
mock_prisma_client.db = mock_db
|
|
|
|
# Set up the user creation mock - create a complete user model that can be converted to a dict
|
|
mock_user = MagicMock()
|
|
mock_user.user_id = "new_test_user"
|
|
mock_user.models = ["gpt-4", "claude-3-opus"]
|
|
mock_user.max_budget = 200.0
|
|
mock_user.user_role = "internal_user"
|
|
mock_user.organization_memberships = []
|
|
|
|
# Make the mock model_dump or dict method return appropriate data
|
|
mock_user.dict = lambda: {
|
|
"user_id": "new_test_user",
|
|
"models": ["gpt-4", "claude-3-opus"],
|
|
"max_budget": 200.0,
|
|
"user_role": "internal_user",
|
|
"organization_memberships": [],
|
|
}
|
|
|
|
# Setup the mock returns
|
|
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
|
|
mock_prisma_client.db.litellm_usertable.create = AsyncMock(return_value=mock_user)
|
|
|
|
# Create a mock cache - use AsyncMock for async methods
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
# Call get_user_object with user_id_upsert=True to trigger user creation
|
|
try:
|
|
user_obj = await get_user_object(
|
|
user_id="new_test_user",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
user_id_upsert=True,
|
|
proxy_logging_obj=None,
|
|
)
|
|
except Exception as e:
|
|
# this fails since the mock object is a MagicMock and not a LiteLLM_UserTable
|
|
print(e)
|
|
|
|
# Verify the user was created with the default params
|
|
mock_prisma_client.db.litellm_usertable.create.assert_called_once()
|
|
creation_args = mock_prisma_client.db.litellm_usertable.create.call_args[1]["data"]
|
|
|
|
# Verify defaults were applied to the creation args
|
|
assert "models" in creation_args
|
|
assert creation_args["models"] == ["gpt-4", "claude-3-opus"]
|
|
assert creation_args["max_budget"] == 200.0
|
|
assert creation_args["user_role"] == "internal_user"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("has_budget_duration", [True, False])
|
|
async def test_get_user_object_upsert_sets_budget_reset_at(monkeypatch, has_budget_duration):
|
|
"""The JWT first-login upsert must compute budget_reset_at when
|
|
default_internal_user_params carries a budget_duration; otherwise the row
|
|
lands with budget_reset_at=NULL and shows a null reset time until the next
|
|
reset sweep heals it. Without a budget_duration, no reset time is written."""
|
|
default_params = {"max_budget": 300.0}
|
|
if has_budget_duration:
|
|
default_params["budget_duration"] = "24h"
|
|
monkeypatch.setattr(litellm, "default_internal_user_params", default_params)
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db = AsyncMock()
|
|
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
|
|
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
|
|
mock_prisma_client.db.litellm_usertable.create = AsyncMock(return_value=MagicMock(organization_memberships=[]))
|
|
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
user_id = f"jwt_upsert_reset_at_{has_budget_duration}"
|
|
try:
|
|
await get_user_object(
|
|
user_id=user_id,
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
user_id_upsert=True,
|
|
proxy_logging_obj=None,
|
|
)
|
|
except Exception as e:
|
|
print(e)
|
|
|
|
mock_prisma_client.db.litellm_usertable.create.assert_called_once()
|
|
creation_args = mock_prisma_client.db.litellm_usertable.create.call_args[1]["data"]
|
|
|
|
if has_budget_duration:
|
|
reset_at = creation_args.get("budget_reset_at")
|
|
assert isinstance(reset_at, datetime), f"expected a computed budget_reset_at, got {creation_args!r}"
|
|
assert reset_at > datetime.now(timezone.utc)
|
|
else:
|
|
assert "budget_reset_at" not in creation_args
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_user_object_wraps_db_outage_as_valueerror_preserving_context():
|
|
"""Pin get_user_object's exception contract: it catches every DB failure in a broad except and
|
|
re-raises a bare ValueError, so a real outage survives only as __context__ rather than as the
|
|
exception type. The MCP dcr_bridge admission and refresh paths depend on this to tell a transient
|
|
outage (retry, 503) from a missing user (fail closed), which is why they classify across the cause
|
|
chain instead of the top exception's type. If this wrapping ever changes, that classification must
|
|
change with it, so this test guards the contract the callers rely on."""
|
|
from unittest.mock import AsyncMock, MagicMock, patch
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db = AsyncMock()
|
|
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
|
|
side_effect=ConnectionError("can't reach database server")
|
|
)
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
with patch("litellm.proxy.auth.auth_checks._should_check_db", return_value=True):
|
|
with pytest.raises(ValueError, match="User doesn't exist in db\\.") as exc_info:
|
|
await get_user_object(
|
|
user_id="outage-contract-probe-user",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
user_id_upsert=False,
|
|
proxy_logging_obj=None,
|
|
)
|
|
|
|
assert isinstance(exc_info.value.__context__, ConnectionError)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_user_object_upsert_includes_user_email():
|
|
"""Test that user_email is included when creating a new user via get_user_object upsert"""
|
|
# Mock the necessary dependencies
|
|
mock_prisma_client = MagicMock()
|
|
mock_db = AsyncMock()
|
|
mock_prisma_client.db = mock_db
|
|
|
|
# Set up the user creation mock
|
|
mock_user = MagicMock()
|
|
mock_user.user_id = "new_test_user"
|
|
mock_user.user_email = "test@example.com"
|
|
mock_user.models = []
|
|
mock_user.max_budget = None
|
|
mock_user.user_role = None
|
|
mock_user.organization_memberships = []
|
|
|
|
mock_user.dict = lambda: {
|
|
"user_id": "new_test_user",
|
|
"user_email": "test@example.com",
|
|
"models": [],
|
|
"max_budget": None,
|
|
"user_role": None,
|
|
"organization_memberships": [],
|
|
}
|
|
|
|
# Setup the mock returns - user does not exist
|
|
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
|
|
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
|
|
mock_prisma_client.db.litellm_usertable.create = AsyncMock(return_value=mock_user)
|
|
|
|
# Create a mock cache
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
# Call get_user_object with user_id_upsert=True and user_email
|
|
try:
|
|
await get_user_object(
|
|
user_id="new_test_user",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
user_id_upsert=True,
|
|
proxy_logging_obj=None,
|
|
user_email="test@example.com",
|
|
)
|
|
except Exception as e:
|
|
# May fail since mock object is not a real LiteLLM_UserTable
|
|
print(e)
|
|
|
|
# Verify the user was created with user_email included
|
|
mock_prisma_client.db.litellm_usertable.create.assert_called_once()
|
|
creation_args = mock_prisma_client.db.litellm_usertable.create.call_args[1]["data"]
|
|
|
|
assert (
|
|
"user_email" in creation_args
|
|
), "user_email should be included when upserting a new user"
|
|
assert creation_args["user_email"] == "test@example.com"
|
|
assert creation_args["user_id"] == "new_test_user"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_user_object_backfills_null_email_from_cache_hit():
|
|
"""
|
|
Regression (LIT-4710): an existing user row with a null user_email must be
|
|
backfilled from the JWT-provided email even when served from cache, so the
|
|
JWT-to-virtual-key path (which resolves straight to the cached user) stops
|
|
logging user_api_key_user_email=null forever. Before the fix the cached row
|
|
was returned unchanged and the DB was never updated.
|
|
"""
|
|
cache = UserApiKeyCache()
|
|
existing = LiteLLM_UserTable(
|
|
user_id="jwt-user-1", user_email=None, user_role="internal_user"
|
|
)
|
|
await cache.async_set_cache(
|
|
key="jwt-user-1", value=existing, model_type=LiteLLM_UserTable
|
|
)
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=1)
|
|
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
|
|
return_value=LiteLLM_UserTable(
|
|
user_id="jwt-user-1",
|
|
user_email="jwt-user-1@example.com",
|
|
user_role="internal_user",
|
|
)
|
|
)
|
|
|
|
result = await get_user_object(
|
|
user_id="jwt-user-1",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=cache,
|
|
user_id_upsert=False,
|
|
proxy_logging_obj=None,
|
|
user_email="jwt-user-1@example.com",
|
|
)
|
|
|
|
assert result is not None
|
|
assert result.user_email == "jwt-user-1@example.com"
|
|
|
|
mock_prisma_client.db.litellm_usertable.update_many.assert_called_once()
|
|
update_kwargs = mock_prisma_client.db.litellm_usertable.update_many.call_args.kwargs
|
|
assert update_kwargs["where"] == {"user_id": "jwt-user-1", "user_email": None}
|
|
assert update_kwargs["data"]["user_email"] == "jwt-user-1@example.com"
|
|
|
|
refreshed = await cache.async_get_cache(
|
|
key="jwt-user-1", model_type=LiteLLM_UserTable
|
|
)
|
|
assert refreshed is not None
|
|
assert refreshed.user_email == "jwt-user-1@example.com"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_user_object_backfills_null_email_from_db_read():
|
|
"""
|
|
Regression (LIT-4710): a user row read from the DB with a null user_email is
|
|
backfilled from the JWT-provided email before it is cached and returned.
|
|
"""
|
|
cache = UserApiKeyCache()
|
|
db_row = LiteLLM_UserTable(
|
|
user_id="jwt-user-3", user_email=None, user_role="internal_user"
|
|
)
|
|
backfilled_row = LiteLLM_UserTable(
|
|
user_id="jwt-user-3",
|
|
user_email="jwt-user-3@example.com",
|
|
user_role="internal_user",
|
|
)
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
|
|
side_effect=[db_row, backfilled_row]
|
|
)
|
|
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
|
|
mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=1)
|
|
|
|
with patch(
|
|
"litellm.proxy.auth.auth_checks._should_check_db", return_value=True
|
|
):
|
|
result = await get_user_object(
|
|
user_id="jwt-user-3",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=cache,
|
|
user_id_upsert=False,
|
|
proxy_logging_obj=None,
|
|
user_email="jwt-user-3@example.com",
|
|
)
|
|
|
|
assert result is not None
|
|
assert result.user_email == "jwt-user-3@example.com"
|
|
mock_prisma_client.db.litellm_usertable.update_many.assert_called_once()
|
|
|
|
refreshed = await cache.async_get_cache(
|
|
key="jwt-user-3", model_type=LiteLLM_UserTable
|
|
)
|
|
assert refreshed is not None
|
|
assert refreshed.user_email == "jwt-user-3@example.com"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_user_object_does_not_overwrite_existing_email():
|
|
"""
|
|
LIT-4710 guardrail: backfill is scoped to null-to-value. An existing non-null
|
|
user_email (e.g. one an operator set intentionally) must never be overwritten
|
|
by the JWT-provided email.
|
|
"""
|
|
cache = UserApiKeyCache()
|
|
existing = LiteLLM_UserTable(
|
|
user_id="jwt-user-2",
|
|
user_email="operator-set@example.com",
|
|
user_role="internal_user",
|
|
)
|
|
await cache.async_set_cache(
|
|
key="jwt-user-2", value=existing, model_type=LiteLLM_UserTable
|
|
)
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=0)
|
|
|
|
result = await get_user_object(
|
|
user_id="jwt-user-2",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=cache,
|
|
user_id_upsert=False,
|
|
proxy_logging_obj=None,
|
|
user_email="different@example.com",
|
|
)
|
|
|
|
assert result is not None
|
|
assert result.user_email == "operator-set@example.com"
|
|
mock_prisma_client.db.litellm_usertable.update_many.assert_not_called()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_user_object_backfill_race_prefers_db_email():
|
|
"""
|
|
LIT-4710 race guard: when the null-guarded update matches 0 rows because a
|
|
concurrent writer already backfilled an email, the cache must be refreshed
|
|
with the value the DB accepted, not this request's proposed email.
|
|
"""
|
|
cache = UserApiKeyCache()
|
|
existing = LiteLLM_UserTable(
|
|
user_id="jwt-user-4", user_email=None, user_role="internal_user"
|
|
)
|
|
await cache.async_set_cache(
|
|
key="jwt-user-4", value=existing, model_type=LiteLLM_UserTable
|
|
)
|
|
|
|
winner_row = LiteLLM_UserTable(
|
|
user_id="jwt-user-4",
|
|
user_email="winner@example.com",
|
|
user_role="internal_user",
|
|
)
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=0)
|
|
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
|
|
return_value=winner_row
|
|
)
|
|
|
|
result = await get_user_object(
|
|
user_id="jwt-user-4",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=cache,
|
|
user_id_upsert=False,
|
|
proxy_logging_obj=None,
|
|
user_email="loser@example.com",
|
|
)
|
|
|
|
assert result is not None
|
|
assert result.user_email == "winner@example.com"
|
|
|
|
refreshed = await cache.async_get_cache(
|
|
key="jwt-user-4", model_type=LiteLLM_UserTable
|
|
)
|
|
assert refreshed is not None
|
|
assert refreshed.user_email == "winner@example.com"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_user_object_backfill_caches_persisted_email_not_proposed():
|
|
"""
|
|
LIT-4710 cache-coherence: even when the null-guarded update succeeds, the
|
|
cache must be refreshed from the row the DB actually holds, not this
|
|
request's proposed email. A concurrent ordinary user update (not null
|
|
guarded) can change the email in the window before the cache write, so
|
|
optimistically caching the proposed email would serve a stale value.
|
|
"""
|
|
cache = UserApiKeyCache()
|
|
existing = LiteLLM_UserTable(
|
|
user_id="jwt-user-5", user_email=None, user_role="internal_user"
|
|
)
|
|
await cache.async_set_cache(
|
|
key="jwt-user-5", value=existing, model_type=LiteLLM_UserTable
|
|
)
|
|
|
|
persisted_row = LiteLLM_UserTable(
|
|
user_id="jwt-user-5",
|
|
user_email="admin-edited@example.com",
|
|
user_role="internal_user",
|
|
)
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_usertable.update_many = AsyncMock(return_value=1)
|
|
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(
|
|
return_value=persisted_row
|
|
)
|
|
|
|
result = await get_user_object(
|
|
user_id="jwt-user-5",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=cache,
|
|
user_id_upsert=False,
|
|
proxy_logging_obj=None,
|
|
user_email="jwt-user-5@example.com",
|
|
)
|
|
|
|
assert result is not None
|
|
assert result.user_email == "admin-edited@example.com"
|
|
|
|
refreshed = await cache.async_get_cache(
|
|
key="jwt-user-5", model_type=LiteLLM_UserTable
|
|
)
|
|
assert refreshed is not None
|
|
assert refreshed.user_email == "admin-edited@example.com"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_user_object_upsert_routes_default_team_to_membership(monkeypatch):
|
|
"""Regression for LIT-4324: a configured default team (list of NewUserRequestTeam
|
|
dicts) must not be written into the Prisma create payload (teams is a String[] column
|
|
that rejects dicts). Instead it must be routed through add_new_user_to_default_team so
|
|
the JWT-provisioned user gets a real team membership."""
|
|
default_params = {
|
|
"user_role": "internal_user",
|
|
"teams": [{"team_id": "default-team", "user_role": "user"}],
|
|
}
|
|
monkeypatch.setattr(litellm, "default_internal_user_params", default_params)
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db = AsyncMock()
|
|
mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
|
|
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
|
|
|
|
mock_user = MagicMock()
|
|
mock_user.organization_memberships = []
|
|
mock_prisma_client.db.litellm_usertable.create = AsyncMock(return_value=mock_user)
|
|
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
with patch(
|
|
"litellm.proxy.management_endpoints.internal_user_endpoints.add_new_user_to_default_team",
|
|
new_callable=AsyncMock,
|
|
) as mock_add_to_team:
|
|
try:
|
|
await get_user_object(
|
|
user_id="new_jwt_user",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
user_id_upsert=True,
|
|
proxy_logging_obj=None,
|
|
)
|
|
except Exception as e:
|
|
# mock_user is a MagicMock, so the post-create LiteLLM_UserTable(**dict(...))
|
|
# conversion raises; irrelevant to what we assert.
|
|
print(e)
|
|
|
|
creation_args = mock_prisma_client.db.litellm_usertable.create.call_args[1]["data"]
|
|
assert "teams" not in creation_args, "teams must be popped before the Prisma create"
|
|
assert creation_args["user_role"] == "internal_user"
|
|
|
|
mock_add_to_team.assert_awaited_once()
|
|
passed_teams = mock_add_to_team.await_args[1]["teams"]
|
|
assert [team.team_id for team in passed_teams] == ["default-team"]
|
|
assert (
|
|
mock_add_to_team.await_args[1]["user_api_key_dict"].user_role
|
|
== LitellmUserRoles.PROXY_ADMIN
|
|
)
|
|
|
|
|
|
def test_log_budget_lookup_failure_dry_run():
|
|
"""Dry run: verify _log_budget_lookup_failure logs for schema/DB errors."""
|
|
with patch("litellm.proxy.auth.auth_checks.verbose_proxy_logger") as mock_logger:
|
|
err = Exception("column 'policies' does not exist in prisma schema")
|
|
_log_budget_lookup_failure("user", err)
|
|
mock_logger.error.assert_called_once()
|
|
call_msg = _rendered_log_message(mock_logger.error.call_args)
|
|
assert "user" in call_msg
|
|
assert "cache will not be populated" in call_msg
|
|
assert "policies" in call_msg or "prisma" in call_msg
|
|
assert "prisma db push" in call_msg
|
|
|
|
|
|
def test_log_budget_lookup_failure_skips_user_not_found():
|
|
"""Verify _log_budget_lookup_failure does NOT log for expected user-not-found."""
|
|
with patch("litellm.proxy.auth.auth_checks.verbose_proxy_logger") as mock_logger:
|
|
err = Exception() # bare Exception from get_user_object when user not found
|
|
_log_budget_lookup_failure("user", err)
|
|
mock_logger.error.assert_not_called()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@patch(
|
|
"litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock
|
|
)
|
|
async def test_get_team_db_check_calls_new_team_on_upsert(mock_new_team, monkeypatch):
|
|
"""
|
|
Test that _get_team_db_check correctly calls the `new_team` function
|
|
when a team does not exist and upsert is enabled.
|
|
"""
|
|
mock_prisma_client = MagicMock()
|
|
mock_db = AsyncMock()
|
|
mock_prisma_client.db = mock_db
|
|
mock_prisma_client.db.litellm_teamtable.find_unique.return_value = None
|
|
|
|
# Define what our mocked `new_team` function should return
|
|
team_id_to_create = "new-jwt-team"
|
|
mock_new_team.return_value = {"team_id": team_id_to_create, "max_budget": 123.45}
|
|
|
|
await _get_team_db_check(
|
|
team_id=team_id_to_create,
|
|
prisma_client=mock_prisma_client,
|
|
team_id_upsert=True,
|
|
)
|
|
|
|
# Verify that our mocked `new_team` function was called exactly once
|
|
mock_new_team.assert_called_once()
|
|
|
|
call_args = mock_new_team.call_args[1]
|
|
data_arg = call_args["data"]
|
|
|
|
# Verify that `new_team` was called with the correct team_id and that
|
|
# `max_budget` was None, as our function's job is to delegate, not to set defaults.
|
|
assert data_arg.team_id == team_id_to_create
|
|
assert data_arg.max_budget is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@patch(
|
|
"litellm.proxy.management_endpoints.team_endpoints.new_team", new_callable=AsyncMock
|
|
)
|
|
async def test_get_team_db_check_does_not_call_new_team_if_exists(
|
|
mock_new_team, monkeypatch
|
|
):
|
|
"""
|
|
Test that _get_team_db_check does NOT call the `new_team` function
|
|
if the team already exists in the database.
|
|
"""
|
|
mock_prisma_client = MagicMock()
|
|
mock_db = AsyncMock()
|
|
mock_prisma_client.db = mock_db
|
|
mock_prisma_client.db.litellm_teamtable.find_unique.return_value = MagicMock()
|
|
|
|
team_id_to_find = "existing-jwt-team"
|
|
|
|
await _get_team_db_check(
|
|
team_id=team_id_to_find,
|
|
prisma_client=mock_prisma_client,
|
|
team_id_upsert=True,
|
|
)
|
|
|
|
# Verify that `new_team` was NEVER called, because the team was found.
|
|
mock_new_team.assert_not_called()
|
|
|
|
|
|
# Vector Store Auth Check Tests
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize(
|
|
"prisma_client,vector_store_registry,expected_result",
|
|
[
|
|
(None, MagicMock(), True), # No prisma client
|
|
(MagicMock(), None, True), # No vector store registry
|
|
(MagicMock(), MagicMock(), True), # No vector stores to run
|
|
],
|
|
)
|
|
async def test_vector_store_access_check_early_returns(
|
|
prisma_client, vector_store_registry, expected_result
|
|
):
|
|
"""Test vector_store_access_check returns True for early exit conditions"""
|
|
request_body = {"messages": [{"role": "user", "content": "test"}]}
|
|
|
|
if vector_store_registry:
|
|
vector_store_registry.get_vector_store_ids_to_run.return_value = None
|
|
|
|
with (
|
|
patch("litellm.proxy.proxy_server.prisma_client", prisma_client),
|
|
patch("litellm.vector_store_registry", vector_store_registry),
|
|
):
|
|
result = await vector_store_access_check(
|
|
request_body=request_body,
|
|
team_object=None,
|
|
valid_token=None,
|
|
)
|
|
|
|
assert result == expected_result
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"object_permissions,vector_store_ids,should_raise,error_type",
|
|
[
|
|
(None, ["store-1"], False, None), # None permissions - should pass
|
|
(
|
|
{"vector_stores": []},
|
|
["store-1"],
|
|
False,
|
|
None,
|
|
), # Empty vector_stores - should pass (access to all)
|
|
(
|
|
{"vector_stores": ["store-1", "store-2"]},
|
|
["store-1"],
|
|
False,
|
|
None,
|
|
), # Has access
|
|
(
|
|
{"vector_stores": ["store-1", "store-2"]},
|
|
["store-3"],
|
|
True,
|
|
ProxyErrorTypes.key_vector_store_access_denied,
|
|
), # No access
|
|
(
|
|
{"vector_stores": ["store-1"]},
|
|
["store-1", "store-3"],
|
|
True,
|
|
ProxyErrorTypes.team_vector_store_access_denied,
|
|
), # Partial access
|
|
],
|
|
)
|
|
def test_can_object_call_vector_stores_scenarios(
|
|
object_permissions, vector_store_ids, should_raise, error_type
|
|
):
|
|
"""Test _can_object_call_vector_stores with various permission scenarios"""
|
|
# Convert dict to object if not None
|
|
if object_permissions is not None:
|
|
mock_permissions = MagicMock()
|
|
mock_permissions.vector_stores = object_permissions["vector_stores"]
|
|
object_permissions = mock_permissions
|
|
|
|
object_type = (
|
|
"key"
|
|
if error_type == ProxyErrorTypes.key_vector_store_access_denied
|
|
else "team"
|
|
)
|
|
|
|
if should_raise:
|
|
with pytest.raises(ProxyException) as exc_info:
|
|
_can_object_call_vector_stores(
|
|
object_type=object_type,
|
|
vector_store_ids_to_run=vector_store_ids,
|
|
object_permissions=object_permissions,
|
|
)
|
|
assert exc_info.value.type == error_type
|
|
else:
|
|
result = _can_object_call_vector_stores(
|
|
object_type=object_type,
|
|
vector_store_ids_to_run=vector_store_ids,
|
|
object_permissions=object_permissions,
|
|
)
|
|
assert result is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_vector_store_access_check_with_permissions():
|
|
"""Test vector_store_access_check with actual permission checking"""
|
|
request_body = {"tools": [{"type": "function", "function": {"name": "test"}}]}
|
|
|
|
# Test with valid token that has access
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
object_permission_id="perm-123",
|
|
models=["gpt-4"],
|
|
max_budget=100.0,
|
|
)
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_permissions = MagicMock()
|
|
mock_permissions.vector_stores = ["store-1", "store-2"]
|
|
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
|
return_value=mock_permissions
|
|
)
|
|
|
|
mock_vector_store_registry = MagicMock()
|
|
mock_vector_store_registry.get_vector_store_ids_to_run.return_value = ["store-1"]
|
|
|
|
with (
|
|
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
|
|
patch("litellm.vector_store_registry", mock_vector_store_registry),
|
|
):
|
|
result = await vector_store_access_check(
|
|
request_body=request_body,
|
|
team_object=None,
|
|
valid_token=valid_token,
|
|
)
|
|
|
|
assert result is True
|
|
|
|
# Test with denied access
|
|
mock_vector_store_registry.get_vector_store_ids_to_run.return_value = ["store-3"]
|
|
|
|
with (
|
|
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
|
|
patch("litellm.vector_store_registry", mock_vector_store_registry),
|
|
):
|
|
with pytest.raises(ProxyException) as exc_info:
|
|
await vector_store_access_check(
|
|
request_body=request_body,
|
|
team_object=None,
|
|
valid_token=valid_token,
|
|
)
|
|
|
|
assert exc_info.value.type == ProxyErrorTypes.key_vector_store_access_denied
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_vector_store_access_check_with_team_permissions():
|
|
"""Ensure teams restricted to specific vector stores cannot access others."""
|
|
request_body = {}
|
|
valid_token = UserAPIKeyAuth(token="team-test-token", object_permission_id=None)
|
|
|
|
team_object = MagicMock()
|
|
team_object.object_permission_id = "team-permission"
|
|
|
|
mock_prisma_client = MagicMock()
|
|
team_permissions = MagicMock()
|
|
team_permissions.vector_stores = ["team-store-allowed"]
|
|
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(
|
|
return_value=team_permissions
|
|
)
|
|
|
|
mock_vector_store_registry = MagicMock()
|
|
mock_vector_store_registry.get_vector_store_ids_to_run.return_value = [
|
|
"team-store-allowed"
|
|
]
|
|
|
|
with (
|
|
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
|
|
patch("litellm.vector_store_registry", mock_vector_store_registry),
|
|
):
|
|
result = await vector_store_access_check(
|
|
request_body=request_body,
|
|
team_object=team_object,
|
|
valid_token=valid_token,
|
|
)
|
|
|
|
assert result is True
|
|
|
|
mock_vector_store_registry.get_vector_store_ids_to_run.return_value = [
|
|
"team-store-denied"
|
|
]
|
|
|
|
with (
|
|
patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client),
|
|
patch("litellm.vector_store_registry", mock_vector_store_registry),
|
|
):
|
|
with pytest.raises(ProxyException) as exc_info:
|
|
await vector_store_access_check(
|
|
request_body=request_body,
|
|
team_object=team_object,
|
|
valid_token=valid_token,
|
|
)
|
|
|
|
assert exc_info.value.type == ProxyErrorTypes.team_vector_store_access_denied
|
|
|
|
|
|
def test_can_object_call_model_with_alias():
|
|
"""Test that can_object_call_model works with model aliases"""
|
|
from litellm import Router
|
|
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
|
|
|
model = "[ip-approved] gpt-4o"
|
|
llm_router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "gpt-3.5-turbo",
|
|
"litellm_params": {
|
|
"model": "gpt-3.5-turbo",
|
|
"api_key": "test-api-key",
|
|
},
|
|
}
|
|
],
|
|
model_group_alias={
|
|
"[ip-approved] gpt-4o": {
|
|
"model": "gpt-3.5-turbo",
|
|
"hidden": True,
|
|
},
|
|
},
|
|
)
|
|
|
|
result = _can_object_call_model(
|
|
model=model,
|
|
llm_router=llm_router,
|
|
models=["gpt-3.5-turbo"],
|
|
team_model_aliases=None,
|
|
object_type="key",
|
|
fallback_depth=0,
|
|
)
|
|
|
|
print(result)
|
|
|
|
|
|
def test_can_object_call_model_access_via_alias_only():
|
|
"""
|
|
Test that a key can access a model via alias even when it doesn't have access to the underlying model.
|
|
|
|
This tests the scenario where:
|
|
- Router has model alias: "my-fake-gpt" -> "gpt-4"
|
|
- Key has access to: ["my-fake-gpt"] (alias)
|
|
- Key does NOT have access to: ["gpt-4"] (underlying model)
|
|
- The call should succeed because access is granted via the alias
|
|
"""
|
|
from litellm import Router
|
|
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
|
|
|
model = "my-fake-gpt"
|
|
llm_router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "gpt-4",
|
|
"litellm_params": {
|
|
"model": "gpt-4",
|
|
"api_key": "test-api-key",
|
|
},
|
|
}
|
|
],
|
|
model_group_alias={
|
|
"my-fake-gpt": {
|
|
"model": "gpt-4",
|
|
"hidden": False,
|
|
},
|
|
},
|
|
)
|
|
|
|
# Key has access to the alias but NOT the underlying model
|
|
result = _can_object_call_model(
|
|
model=model,
|
|
llm_router=llm_router,
|
|
models=["my-fake-gpt"], # Only has access to alias, not "gpt-4"
|
|
team_model_aliases=None,
|
|
object_type="key",
|
|
fallback_depth=0,
|
|
)
|
|
|
|
# Should return True because access is granted via the alias
|
|
assert result is True
|
|
|
|
|
|
def test_can_object_call_model_access_via_underlying_model_only():
|
|
"""
|
|
Test that a key can access a model via underlying model even when using an alias.
|
|
|
|
This tests the scenario where:
|
|
- Router has model alias: "my-fake-gpt" -> "gpt-4"
|
|
- Key has access to: ["gpt-4"] (underlying model)
|
|
- Key does NOT have access to: ["my-fake-gpt"] (alias)
|
|
- The call should succeed because access is granted via the underlying model
|
|
"""
|
|
from litellm import Router
|
|
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
|
|
|
model = "my-fake-gpt"
|
|
llm_router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "gpt-4",
|
|
"litellm_params": {
|
|
"model": "gpt-4",
|
|
"api_key": "test-api-key",
|
|
},
|
|
}
|
|
],
|
|
model_group_alias={
|
|
"my-fake-gpt": {
|
|
"model": "gpt-4",
|
|
"hidden": False,
|
|
},
|
|
},
|
|
)
|
|
|
|
# Key has access to the underlying model but NOT the alias
|
|
result = _can_object_call_model(
|
|
model=model,
|
|
llm_router=llm_router,
|
|
models=["gpt-4"], # Only has access to underlying model, not "my-fake-gpt"
|
|
team_model_aliases=None,
|
|
object_type="key",
|
|
fallback_depth=0,
|
|
)
|
|
|
|
# Should return True because access is granted via the underlying model
|
|
assert result is True
|
|
|
|
|
|
def test_can_object_call_model_no_access_to_alias_or_underlying():
|
|
"""
|
|
Test that a key cannot access a model when it has no access to either alias or underlying model.
|
|
"""
|
|
from litellm import Router
|
|
from litellm.proxy._types import ProxyErrorTypes, ProxyException
|
|
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
|
|
|
model = "my-fake-gpt"
|
|
llm_router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "gpt-4",
|
|
"litellm_params": {
|
|
"model": "gpt-4",
|
|
"api_key": "test-api-key",
|
|
},
|
|
}
|
|
],
|
|
model_group_alias={
|
|
"my-fake-gpt": {
|
|
"model": "gpt-4",
|
|
"hidden": False,
|
|
},
|
|
},
|
|
)
|
|
|
|
# Key has access to neither the alias nor the underlying model
|
|
with pytest.raises(ProxyException) as exc_info:
|
|
_can_object_call_model(
|
|
model=model,
|
|
llm_router=llm_router,
|
|
models=["gpt-3.5-turbo"], # Has access to different model entirely
|
|
team_model_aliases=None,
|
|
object_type="key",
|
|
fallback_depth=0,
|
|
)
|
|
|
|
# Should raise ProxyException with appropriate error type
|
|
assert exc_info.value.type == ProxyErrorTypes.key_model_access_denied
|
|
assert "key not allowed to access model" in str(exc_info.value.message)
|
|
assert "my-fake-gpt" in str(exc_info.value.message)
|
|
|
|
|
|
# -- Team-member access-group resolution with team-scoped DB models -----------
|
|
|
|
|
|
def _make_team_scoped_router(team_id: str = "team-a"):
|
|
"""
|
|
Build a Router whose model_list looks like what the proxy creates for
|
|
team-scoped BYOK DB models: the internal model_name is
|
|
``<public_name>_<team_id>_<uuid>`` and the public name lives in
|
|
``model_info.team_public_model_name``. Two models belong to the
|
|
access group ``fast-models``; one (``mock-power``) does not.
|
|
"""
|
|
from litellm import Router
|
|
|
|
model_list = [
|
|
{
|
|
"model_name": f"mock-fast-1_{team_id}_aaa",
|
|
"litellm_params": {
|
|
"model": "openai/mock-fast-1",
|
|
"api_key": "fake",
|
|
},
|
|
"model_info": {
|
|
"id": f"demo-mock-fast-1-{team_id}",
|
|
"team_id": team_id,
|
|
"team_public_model_name": "mock-fast-1",
|
|
"access_groups": ["fast-models"],
|
|
},
|
|
},
|
|
{
|
|
"model_name": f"mock-fast-2_{team_id}_bbb",
|
|
"litellm_params": {
|
|
"model": "openai/mock-fast-2",
|
|
"api_key": "fake",
|
|
},
|
|
"model_info": {
|
|
"id": f"demo-mock-fast-2-{team_id}",
|
|
"team_id": team_id,
|
|
"team_public_model_name": "mock-fast-2",
|
|
"access_groups": ["fast-models"],
|
|
},
|
|
},
|
|
{
|
|
"model_name": f"mock-power_{team_id}_ccc",
|
|
"litellm_params": {
|
|
"model": "openai/mock-power",
|
|
"api_key": "fake",
|
|
},
|
|
"model_info": {
|
|
"id": f"demo-mock-power-{team_id}",
|
|
"team_id": team_id,
|
|
"team_public_model_name": "mock-power",
|
|
},
|
|
},
|
|
]
|
|
return Router(model_list=model_list)
|
|
|
|
|
|
def test_can_object_call_model_access_group_with_team_id():
|
|
"""
|
|
When team_id is passed, _can_object_call_model should resolve
|
|
model_info.access_groups for team-scoped DB models and allow
|
|
access via group name.
|
|
"""
|
|
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
|
|
|
router = _make_team_scoped_router()
|
|
|
|
result = _can_object_call_model(
|
|
model="mock-fast-1",
|
|
llm_router=router,
|
|
models=["fast-models", "mock-power"],
|
|
object_type="team",
|
|
team_id="team-a",
|
|
)
|
|
assert result is True
|
|
|
|
|
|
def test_can_object_call_model_access_group_without_team_id_fails():
|
|
"""
|
|
Without team_id the router cannot find team-scoped DB models, so
|
|
access group resolution fails and the call is denied.
|
|
This is the pre-fix behavior.
|
|
"""
|
|
from litellm.proxy._types import ProxyException
|
|
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
|
|
|
router = _make_team_scoped_router()
|
|
|
|
with pytest.raises(ProxyException):
|
|
_can_object_call_model(
|
|
model="mock-fast-1",
|
|
llm_router=router,
|
|
models=["fast-models", "mock-power"],
|
|
object_type="team",
|
|
# team_id intentionally omitted
|
|
)
|
|
|
|
|
|
def test_can_object_call_model_literal_name_with_team_id():
|
|
"""
|
|
Literal model name matching should still work when team_id is
|
|
passed — no regression from adding team_id.
|
|
"""
|
|
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
|
|
|
router = _make_team_scoped_router()
|
|
|
|
result = _can_object_call_model(
|
|
model="mock-power",
|
|
llm_router=router,
|
|
models=["fast-models", "mock-power"],
|
|
object_type="team",
|
|
team_id="team-a",
|
|
)
|
|
assert result is True
|
|
|
|
|
|
def test_can_object_call_model_denied_model_with_team_id():
|
|
"""
|
|
A model not in the allowed list (by name or access group) should
|
|
still be denied even when team_id is passed.
|
|
"""
|
|
from litellm.proxy._types import ProxyException
|
|
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
|
|
|
router = _make_team_scoped_router()
|
|
|
|
with pytest.raises(ProxyException):
|
|
_can_object_call_model(
|
|
model="mock-vision",
|
|
llm_router=router,
|
|
models=["fast-models", "mock-power"],
|
|
object_type="team",
|
|
team_id="team-a",
|
|
)
|
|
|
|
|
|
def test_can_object_call_model_second_group_member_with_team_id():
|
|
"""
|
|
Both models in the access group should be reachable, not just
|
|
the first one.
|
|
"""
|
|
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
|
|
|
router = _make_team_scoped_router()
|
|
|
|
result = _can_object_call_model(
|
|
model="mock-fast-2",
|
|
llm_router=router,
|
|
models=["fast-models"],
|
|
object_type="team",
|
|
team_id="team-a",
|
|
)
|
|
assert result is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_check_team_member_model_access_with_access_group():
|
|
"""
|
|
End-to-end test of _check_team_member_model_access: a member whose
|
|
allowed_models contains an access group name should be allowed to
|
|
call models in that group for team-scoped DB models.
|
|
"""
|
|
from litellm.proxy._types import (
|
|
LiteLLM_BudgetTable,
|
|
LiteLLM_TeamMembership,
|
|
LiteLLM_TeamTable,
|
|
UserAPIKeyAuth,
|
|
)
|
|
from litellm.proxy.auth.auth_checks import _check_team_member_model_access
|
|
|
|
router = _make_team_scoped_router()
|
|
team = LiteLLM_TeamTable(team_id="team-a")
|
|
token = UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a")
|
|
membership = LiteLLM_TeamMembership(
|
|
user_id="alice",
|
|
team_id="team-a",
|
|
litellm_budget_table=LiteLLM_BudgetTable(
|
|
allowed_models=["fast-models", "mock-power"],
|
|
),
|
|
)
|
|
|
|
with patch(
|
|
"litellm.proxy.auth.auth_checks.get_team_membership",
|
|
return_value=membership,
|
|
):
|
|
# Should not raise — mock-fast-1 is in the fast-models group
|
|
await _check_team_member_model_access(
|
|
model="mock-fast-1",
|
|
team_object=team,
|
|
valid_token=token,
|
|
llm_router=router,
|
|
prisma_client=None,
|
|
user_api_key_cache=MagicMock(),
|
|
proxy_logging_obj=MagicMock(),
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_check_team_member_model_access_denied_model():
|
|
"""
|
|
A member with per-member allowed_models should be denied access to
|
|
a model that is neither listed by name nor covered by an access group.
|
|
"""
|
|
from litellm.proxy._types import (
|
|
LiteLLM_BudgetTable,
|
|
LiteLLM_TeamMembership,
|
|
LiteLLM_TeamTable,
|
|
ProxyException,
|
|
UserAPIKeyAuth,
|
|
)
|
|
from litellm.proxy.auth.auth_checks import _check_team_member_model_access
|
|
|
|
router = _make_team_scoped_router()
|
|
team = LiteLLM_TeamTable(team_id="team-a")
|
|
token = UserAPIKeyAuth(token="sk-test", user_id="alice", team_id="team-a")
|
|
membership = LiteLLM_TeamMembership(
|
|
user_id="alice",
|
|
team_id="team-a",
|
|
litellm_budget_table=LiteLLM_BudgetTable(
|
|
allowed_models=["fast-models", "mock-power"],
|
|
),
|
|
)
|
|
|
|
with patch(
|
|
"litellm.proxy.auth.auth_checks.get_team_membership",
|
|
return_value=membership,
|
|
):
|
|
with pytest.raises(ProxyException) as exc_info:
|
|
await _check_team_member_model_access(
|
|
model="mock-vision",
|
|
team_object=team,
|
|
valid_token=token,
|
|
llm_router=router,
|
|
prisma_client=None,
|
|
user_api_key_cache=MagicMock(),
|
|
proxy_logging_obj=MagicMock(),
|
|
)
|
|
assert exc_info.value.type == ProxyErrorTypes.team_model_access_denied
|
|
assert int(exc_info.value.code) == status.HTTP_403_FORBIDDEN
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_check_team_member_model_access_no_override_inherits_team():
|
|
"""
|
|
When a member has no allowed_models (empty budget table), the function
|
|
should return without raising — the team-level check applies instead.
|
|
"""
|
|
from litellm.proxy._types import (
|
|
LiteLLM_BudgetTable,
|
|
LiteLLM_TeamMembership,
|
|
LiteLLM_TeamTable,
|
|
UserAPIKeyAuth,
|
|
)
|
|
from litellm.proxy.auth.auth_checks import _check_team_member_model_access
|
|
|
|
router = _make_team_scoped_router()
|
|
team = LiteLLM_TeamTable(team_id="team-a")
|
|
token = UserAPIKeyAuth(token="sk-test", user_id="bob", team_id="team-a")
|
|
membership = LiteLLM_TeamMembership(
|
|
user_id="bob",
|
|
team_id="team-a",
|
|
litellm_budget_table=LiteLLM_BudgetTable(),
|
|
)
|
|
|
|
with patch(
|
|
"litellm.proxy.auth.auth_checks.get_team_membership",
|
|
return_value=membership,
|
|
):
|
|
# Should return without raising — no per-member restriction
|
|
await _check_team_member_model_access(
|
|
model="mock-vision",
|
|
team_object=team,
|
|
valid_token=token,
|
|
llm_router=router,
|
|
prisma_client=None,
|
|
user_api_key_cache=MagicMock(),
|
|
proxy_logging_obj=MagicMock(),
|
|
)
|
|
|
|
|
|
# Tag Budget Enforcement Tests
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_tag_objects_batch():
|
|
"""
|
|
Test batch fetching of tags validates:
|
|
- Cached tags are fetched from cache (no DB call for them)
|
|
- Uncached tags are fetched in ONE batch DB query
|
|
- After fetching, uncached tags are cached
|
|
"""
|
|
from litellm.proxy._types import LiteLLM_TagTable
|
|
from litellm.proxy.auth.auth_checks import get_tag_objects_batch
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_cache = MagicMock()
|
|
mock_proxy_logging = MagicMock()
|
|
|
|
# Simulate 5 tags: 2 cached, 3 uncached
|
|
tag_names = ["cached-1", "uncached-1", "cached-2", "uncached-2", "uncached-3"]
|
|
|
|
# Mock cached tags — must be LiteLLM_TagTable instances: the mocked async_get_cache
|
|
# bypasses UserApiKeyCache deserialization, so returning plain dicts would flow through
|
|
# as dict (production returns models after Codec.deserialize inside the cache).
|
|
cached_tag_1 = LiteLLM_TagTable(
|
|
tag_name="cached-1",
|
|
spend=10.0,
|
|
models=[],
|
|
litellm_budget_table=None,
|
|
)
|
|
cached_tag_2 = LiteLLM_TagTable(
|
|
tag_name="cached-2",
|
|
spend=20.0,
|
|
models=[],
|
|
litellm_budget_table=None,
|
|
)
|
|
|
|
# Mock DB response for uncached tags
|
|
uncached_tag_1 = MagicMock()
|
|
uncached_tag_1.tag_name = "uncached-1"
|
|
uncached_tag_1.spend = 30.0
|
|
uncached_tag_1.models = []
|
|
uncached_tag_1.litellm_budget_table = None
|
|
uncached_tag_1.dict = MagicMock(
|
|
return_value={
|
|
"tag_name": "uncached-1",
|
|
"spend": 30.0,
|
|
"models": [],
|
|
"litellm_budget_table": None,
|
|
}
|
|
)
|
|
|
|
uncached_tag_2 = MagicMock()
|
|
uncached_tag_2.tag_name = "uncached-2"
|
|
uncached_tag_2.spend = 40.0
|
|
uncached_tag_2.models = []
|
|
uncached_tag_2.litellm_budget_table = None
|
|
uncached_tag_2.dict = MagicMock(
|
|
return_value={
|
|
"tag_name": "uncached-2",
|
|
"spend": 40.0,
|
|
"models": [],
|
|
"litellm_budget_table": None,
|
|
}
|
|
)
|
|
|
|
uncached_tag_3 = MagicMock()
|
|
uncached_tag_3.tag_name = "uncached-3"
|
|
uncached_tag_3.spend = 50.0
|
|
uncached_tag_3.models = []
|
|
uncached_tag_3.litellm_budget_table = None
|
|
uncached_tag_3.dict = MagicMock(
|
|
return_value={
|
|
"tag_name": "uncached-3",
|
|
"spend": 50.0,
|
|
"models": [],
|
|
"litellm_budget_table": None,
|
|
}
|
|
)
|
|
|
|
# Mock cache behavior - return cached tags, None for uncached
|
|
async def mock_get_cache(*args, **kwargs):
|
|
key = kwargs.get("key")
|
|
if key == "tag:cached-1":
|
|
return cached_tag_1
|
|
if key == "tag:cached-2":
|
|
return cached_tag_2
|
|
return None
|
|
|
|
mock_cache.async_get_cache = AsyncMock(side_effect=mock_get_cache)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
# Mock DB to return all uncached tags in ONE query
|
|
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(
|
|
return_value=[uncached_tag_1, uncached_tag_2, uncached_tag_3]
|
|
)
|
|
|
|
# Call batch fetch
|
|
tag_objects = await get_tag_objects_batch(
|
|
tag_names=tag_names,
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=mock_cache,
|
|
proxy_logging_obj=mock_proxy_logging,
|
|
)
|
|
|
|
# Verify results
|
|
assert len(tag_objects) == 5
|
|
assert "cached-1" in tag_objects
|
|
assert "cached-2" in tag_objects
|
|
assert "uncached-1" in tag_objects
|
|
assert "uncached-2" in tag_objects
|
|
assert "uncached-3" in tag_objects
|
|
|
|
# Verify cached tags have correct values
|
|
assert tag_objects["cached-1"].spend == 10.0
|
|
assert tag_objects["cached-2"].spend == 20.0
|
|
|
|
# Verify uncached tags have correct values
|
|
assert tag_objects["uncached-1"].spend == 30.0
|
|
assert tag_objects["uncached-2"].spend == 40.0
|
|
assert tag_objects["uncached-3"].spend == 50.0
|
|
|
|
# Verify the DB saw exactly the registry query plus ONE batch query for all 3 uncached tags
|
|
assert mock_prisma.db.litellm_tagtable.find_many.call_count == 2
|
|
registry_call, batch_call = mock_prisma.db.litellm_tagtable.find_many.call_args_list
|
|
assert "where" not in registry_call.kwargs
|
|
assert batch_call.kwargs["where"]["tag_name"]["in"] == [
|
|
"uncached-1",
|
|
"uncached-2",
|
|
"uncached-3",
|
|
]
|
|
|
|
# Verify uncached tags were cached after fetching, alongside the tag-name registry
|
|
cache_calls = mock_cache.async_set_cache.call_args_list
|
|
cached_keys = [call.kwargs["key"] for call in cache_calls]
|
|
assert sorted(cached_keys) == [
|
|
"tag:uncached-1",
|
|
"tag:uncached-2",
|
|
"tag:uncached-3",
|
|
"tag_registry",
|
|
]
|
|
# Every write is TTL-bounded; an unbounded tag entry would outlive budget updates.
|
|
assert all("ttl" in call.kwargs for call in cache_calls)
|
|
|
|
|
|
class _TtlRecordingCache(UserApiKeyCache):
|
|
"""A real cache that also records the ttl each write carried, so tests can catch unbounded entries."""
|
|
|
|
def __init__(self):
|
|
super().__init__()
|
|
self.writes = []
|
|
|
|
async def async_set_cache(self, key, value, local_only: bool = False, **kwargs):
|
|
self.writes.append((key, kwargs.get("ttl")))
|
|
return await super().async_set_cache(key, value, local_only=local_only, **kwargs)
|
|
|
|
|
|
def _tag_registry_row(tag_name: str):
|
|
"""A row as the names-only registry query sees it: only ``tag_name`` is read off it."""
|
|
return SimpleNamespace(tag_name=tag_name)
|
|
|
|
|
|
def _tag_db_row(tag_name: str, max_budget=None):
|
|
row = MagicMock()
|
|
row.tag_name = tag_name
|
|
budget = None if max_budget is None else {"max_budget": max_budget}
|
|
row.dict = MagicMock(
|
|
return_value={
|
|
"tag_name": tag_name,
|
|
"spend": 0.0,
|
|
"models": [],
|
|
"litellm_budget_table": budget,
|
|
}
|
|
)
|
|
return row
|
|
|
|
|
|
def _registry_calls(find_many):
|
|
return [call for call in find_many.call_args_list if "where" not in call.kwargs]
|
|
|
|
|
|
def _batch_calls(find_many):
|
|
return [call for call in find_many.call_args_list if "where" in call.kwargs]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_tag_objects_batch_never_queries_db_for_unregistered_tags():
|
|
"""
|
|
Regression: a request tag with no LiteLLM_TagTable row must not cost a DB read per request.
|
|
|
|
Cost-attribution tags are free-form, so most carry no tag row. Before the cached name
|
|
registry, every request carrying one ran its own Postgres find_many, forever, which is what
|
|
saturated a customer's Prisma pool.
|
|
"""
|
|
from litellm.proxy.auth.auth_checks import get_tag_objects_batch
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(
|
|
return_value=[_tag_registry_row("some-other-tag")]
|
|
)
|
|
cache = UserApiKeyCache()
|
|
|
|
first = await get_tag_objects_batch(
|
|
tag_names=["unregistered-tag"],
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert first == {}
|
|
|
|
# The only query is the names-only registry fetch; the tag itself is never looked up.
|
|
mock_prisma.db.litellm_tagtable.find_many.assert_called_once_with(
|
|
take=TAG_REGISTRY_MAX_SIZE + 1
|
|
)
|
|
|
|
second = await get_tag_objects_batch(
|
|
tag_names=["unregistered-tag"],
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert second == {}
|
|
assert mock_prisma.db.litellm_tagtable.find_many.call_count == 1
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_tag_objects_batch_fetches_only_registered_uncached_tags():
|
|
"""Cached tags skip the DB, registered ones are batch-fetched, unregistered ones are dropped."""
|
|
from litellm.proxy.auth.auth_checks import get_tag_objects_batch
|
|
|
|
cache = UserApiKeyCache()
|
|
await cache.async_set_cache(
|
|
key=tag_cache_key("cached-tag"),
|
|
value=LiteLLM_TagTable(tag_name="cached-tag", spend=7.0, models=[]),
|
|
model_type=LiteLLM_TagTable,
|
|
)
|
|
|
|
async def fake_find_many(**kwargs):
|
|
if "where" not in kwargs:
|
|
return [_tag_registry_row("cached-tag"), _tag_registry_row("registered-tag")]
|
|
requested = kwargs["where"]["tag_name"]["in"]
|
|
return [_tag_db_row(name) for name in requested if name == "registered-tag"]
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(side_effect=fake_find_many)
|
|
|
|
tag_objects = await get_tag_objects_batch(
|
|
tag_names=["cached-tag", "registered-tag", "unregistered-tag"],
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
|
|
assert sorted(tag_objects) == ["cached-tag", "registered-tag"]
|
|
assert tag_objects["cached-tag"].spend == 7.0
|
|
|
|
batch_calls = _batch_calls(mock_prisma.db.litellm_tagtable.find_many)
|
|
assert len(batch_calls) == 1
|
|
assert batch_calls[0].kwargs["where"]["tag_name"]["in"] == ["registered-tag"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_tag_objects_batch_caches_empty_registry():
|
|
"""An empty tag table is a valid registry answer and must be cached, not re-queried."""
|
|
from litellm.proxy.auth.auth_checks import get_tag_objects_batch
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(return_value=[])
|
|
cache = UserApiKeyCache()
|
|
|
|
assert (
|
|
await get_tag_objects_batch(
|
|
tag_names=["tag-a", "tag-b"],
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
== {}
|
|
)
|
|
# "No tags registered" is a cached answer, not a cache miss (which would be None).
|
|
cached_registry = await cache.async_get_cache(key=tag_registry_cache_key())
|
|
assert cached_registry is not None
|
|
assert tuple(cached_registry) == ()
|
|
|
|
assert (
|
|
await get_tag_objects_batch(
|
|
tag_names=["tag-a", "tag-b"],
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
== {}
|
|
)
|
|
assert mock_prisma.db.litellm_tagtable.find_many.call_count == 1
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_tag_objects_batch_registry_db_error_negative_caches_and_keeps_per_tag_fetch():
|
|
"""
|
|
A degraded database must not be re-asked for the registry on every request.
|
|
|
|
Without the negative cache the failing scan re-runs per request on top of the per-tag fallback
|
|
it triggers, doubling load exactly when Postgres is least able to take it. Tag budgets keep
|
|
being enforced through the per-tag path throughout, and the registry is retried once the
|
|
negative entry expires.
|
|
"""
|
|
from litellm.proxy.auth.auth_checks import get_tag_objects_batch
|
|
|
|
async def fake_find_many(**kwargs):
|
|
if "where" not in kwargs:
|
|
raise Exception("registry query failed")
|
|
requested = kwargs["where"]["tag_name"]["in"]
|
|
return [_tag_db_row(name) for name in requested]
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(side_effect=fake_find_many)
|
|
cache = _TtlRecordingCache()
|
|
|
|
first = await get_tag_objects_batch(
|
|
tag_names=["tag-a"],
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert list(first) == ["tag-a"]
|
|
assert await cache.async_get_cache(key=tag_registry_cache_key()) == TAG_REGISTRY_OVERFLOW_SENTINEL
|
|
assert (tag_registry_cache_key(), REGISTRY_ERROR_NEGATIVE_CACHE_TTL) in cache.writes
|
|
|
|
second = await get_tag_objects_batch(
|
|
tag_names=["tag-b"],
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert list(second) == ["tag-b"]
|
|
assert len(_registry_calls(mock_prisma.db.litellm_tagtable.find_many)) == 1
|
|
|
|
# The window closing (here: the entry expiring) puts the registry back in play.
|
|
await cache.async_delete_cache(key=tag_registry_cache_key())
|
|
third = await get_tag_objects_batch(
|
|
tag_names=["tag-c"],
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert list(third) == ["tag-c"]
|
|
assert len(_registry_calls(mock_prisma.db.litellm_tagtable.find_many)) == 2
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_tag_registry_load_is_single_flighted_across_concurrent_requests():
|
|
"""
|
|
A cold registry under load must run one scan, not one per in-flight request.
|
|
|
|
The registry query is an unindexed table scan; a TTL expiry on a busy worker would otherwise
|
|
fan out into as many identical scans as there are concurrent requests.
|
|
"""
|
|
from litellm.proxy.auth.auth_checks import get_tag_objects_batch
|
|
|
|
async def fake_find_many(**kwargs):
|
|
if "where" not in kwargs:
|
|
await asyncio.sleep(0)
|
|
return [_tag_registry_row("registered-tag")]
|
|
return [_tag_db_row(name) for name in kwargs["where"]["tag_name"]["in"]]
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(side_effect=fake_find_many)
|
|
cache = UserApiKeyCache()
|
|
|
|
results = await asyncio.gather(
|
|
*(
|
|
get_tag_objects_batch(
|
|
tag_names=["registered-tag"],
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
for _ in range(8)
|
|
)
|
|
)
|
|
|
|
assert all(list(result) == ["registered-tag"] for result in results)
|
|
assert len(_registry_calls(mock_prisma.db.litellm_tagtable.find_many)) == 1
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_tag_objects_batch_oversized_registry_falls_back_and_stops_refetching():
|
|
"""Past the cap the registry is unusable: keep the old per-tag path, but stop rebuilding it."""
|
|
from litellm.proxy.auth.auth_checks import get_tag_objects_batch
|
|
|
|
oversized = [
|
|
_tag_registry_row(f"tag-{index}") for index in range(TAG_REGISTRY_MAX_SIZE + 1)
|
|
]
|
|
|
|
async def fake_find_many(**kwargs):
|
|
if "where" not in kwargs:
|
|
return oversized
|
|
return [_tag_db_row(name) for name in kwargs["where"]["tag_name"]["in"]]
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(side_effect=fake_find_many)
|
|
cache = UserApiKeyCache()
|
|
|
|
first = await get_tag_objects_batch(
|
|
tag_names=["tag-a"],
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert list(first) == ["tag-a"]
|
|
assert (
|
|
await cache.async_get_cache(key=tag_registry_cache_key())
|
|
== TAG_REGISTRY_OVERFLOW_SENTINEL
|
|
)
|
|
|
|
second = await get_tag_objects_batch(
|
|
tag_names=["tag-b"],
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert list(second) == ["tag-b"]
|
|
|
|
find_many = mock_prisma.db.litellm_tagtable.find_many
|
|
assert len(_registry_calls(find_many)) == 1
|
|
assert [call.kwargs["where"]["tag_name"]["in"] for call in _batch_calls(find_many)] == [
|
|
["tag-a"],
|
|
["tag-b"],
|
|
]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_tag_max_budget_check_still_enforces_registered_tag_over_budget():
|
|
"""The registry filter must not swallow a real tag: an over-budget tag still raises."""
|
|
from litellm.proxy.utils import ProxyLogging
|
|
|
|
async def fake_find_many(**kwargs):
|
|
if "where" not in kwargs:
|
|
return [_tag_registry_row("paid-tag")]
|
|
return [
|
|
_tag_db_row(name, max_budget=1.0)
|
|
for name in kwargs["where"]["tag_name"]["in"]
|
|
]
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_tagtable.find_many = AsyncMock(side_effect=fake_find_many)
|
|
|
|
async def mock_get_current_spend(
|
|
counter_key, fallback_spend, max_budget=None, **kwargs
|
|
):
|
|
if counter_key == "spend:tag:paid-tag":
|
|
return 1.5
|
|
return fallback_spend
|
|
|
|
with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend):
|
|
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
|
await _tag_max_budget_check(
|
|
request_body={"metadata": {"tags": ["paid-tag", "unregistered-tag"]}},
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=UserApiKeyCache(),
|
|
proxy_logging_obj=ProxyLogging(user_api_key_cache=None),
|
|
valid_token=UserAPIKeyAuth(token="test-token"),
|
|
)
|
|
assert exc_info.value.current_cost == 1.5
|
|
assert exc_info.value.entity_id == "paid-tag"
|
|
|
|
# The unregistered tag alongside it never reached the DB.
|
|
batch_calls = _batch_calls(mock_prisma.db.litellm_tagtable.find_many)
|
|
assert [call.kwargs["where"]["tag_name"]["in"] for call in batch_calls] == [["paid-tag"]]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_team_object_raises_404_when_not_found():
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
from fastapi import HTTPException
|
|
|
|
from litellm.proxy.auth.auth_checks import get_team_object
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_db = AsyncMock()
|
|
mock_prisma_client.db = mock_db
|
|
mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=None)
|
|
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
|
|
with pytest.raises(HTTPException) as exc_info:
|
|
await get_team_object(
|
|
team_id="nonexistent-team",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
check_cache_only=False,
|
|
check_db_only=True,
|
|
)
|
|
|
|
assert exc_info.value.status_code == 404
|
|
assert "Team doesn't exist in db" in str(exc_info.value.detail)
|
|
|
|
|
|
def _mock_prisma_for_team_lookup(find_unique):
|
|
from unittest.mock import MagicMock
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_teamtable.find_unique = find_unique
|
|
return mock_prisma_client
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_team_object_distinguishes_absent_team_from_unreadable_row():
|
|
"""A deleted team and a database that would not answer both surface as a 404,
|
|
which leaves callers unable to tell a definitive answer from a degraded read.
|
|
Only the row being positively absent raises the subclass; anything else keeps
|
|
the plain 404 so every existing caller is unaffected."""
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
from fastapi import HTTPException
|
|
|
|
from litellm.proxy.auth.auth_checks import TeamNotFoundError, get_team_object
|
|
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
|
|
# The database answered, and the row is not there.
|
|
with pytest.raises(TeamNotFoundError) as absent_info:
|
|
await get_team_object(
|
|
team_id="absent-team-lit5522",
|
|
prisma_client=_mock_prisma_for_team_lookup(AsyncMock(return_value=None)),
|
|
user_api_key_cache=mock_cache,
|
|
check_db_only=True,
|
|
)
|
|
assert absent_info.value.status_code == 404
|
|
assert "Team doesn't exist in db" in str(absent_info.value.detail)
|
|
|
|
# The database did not answer. Same status and detail, but not the subclass,
|
|
# so a caller keying on it does not read this as proof the team is gone.
|
|
with pytest.raises(HTTPException) as unreadable_info:
|
|
await get_team_object(
|
|
team_id="unreadable-team-lit5522",
|
|
prisma_client=_mock_prisma_for_team_lookup(AsyncMock(side_effect=ConnectionError("db unreachable"))),
|
|
user_api_key_cache=mock_cache,
|
|
check_db_only=True,
|
|
)
|
|
assert unreadable_info.value.status_code == 404
|
|
assert not isinstance(unreadable_info.value, TeamNotFoundError)
|
|
|
|
|
|
# Reject Client-Side Metadata Tags Tests
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_reject_clientside_metadata_tags_enabled_with_tags():
|
|
"""Test that common_checks rejects request when reject_clientside_metadata_tags is True and metadata.tags is present."""
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
request_body = {
|
|
"model": "gpt-3.5-turbo",
|
|
"messages": [{"role": "user", "content": "test"}],
|
|
"metadata": {"tags": ["custom-tag"]},
|
|
}
|
|
|
|
general_settings = {"reject_clientside_metadata_tags": True}
|
|
|
|
# Create a mock request object
|
|
mock_request = MagicMock(spec=Request)
|
|
|
|
# Create a valid token for the test
|
|
valid_token = UserAPIKeyAuth(token="test-token", models=["gpt-3.5-turbo"])
|
|
|
|
with pytest.raises(ProxyException) as exc_info:
|
|
await common_checks(
|
|
request_body=request_body,
|
|
team_object=None,
|
|
user_object=None,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings=general_settings,
|
|
route="/chat/completions",
|
|
llm_router=None,
|
|
proxy_logging_obj=MagicMock(),
|
|
valid_token=valid_token,
|
|
request=mock_request,
|
|
)
|
|
|
|
assert exc_info.value.type == ProxyErrorTypes.bad_request_error
|
|
assert "metadata.tags" in exc_info.value.message
|
|
assert exc_info.value.param == "metadata.tags"
|
|
assert exc_info.value.code == "400"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_reject_clientside_metadata_tags_enabled_without_tags():
|
|
"""Test that common_checks allows request when reject_clientside_metadata_tags is True but no metadata.tags is present."""
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
request_body = {
|
|
"model": "gpt-3.5-turbo",
|
|
"messages": [{"role": "user", "content": "test"}],
|
|
"metadata": {"custom_field": "value"}, # No tags field
|
|
}
|
|
|
|
general_settings = {"reject_clientside_metadata_tags": True}
|
|
|
|
# Create a mock request object
|
|
mock_request = MagicMock(spec=Request)
|
|
|
|
# Create a valid token for the test
|
|
valid_token = UserAPIKeyAuth(token="test-token", models=["gpt-3.5-turbo"])
|
|
|
|
# Should not raise an exception
|
|
result = await common_checks(
|
|
request_body=request_body,
|
|
team_object=None,
|
|
user_object=None,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings=general_settings,
|
|
route="/chat/completions",
|
|
llm_router=None,
|
|
proxy_logging_obj=MagicMock(),
|
|
valid_token=valid_token,
|
|
request=mock_request,
|
|
)
|
|
|
|
assert result is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_reject_clientside_metadata_tags_disabled_with_tags():
|
|
"""Test that common_checks allows request with metadata.tags when reject_clientside_metadata_tags is False."""
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
request_body = {
|
|
"model": "gpt-3.5-turbo",
|
|
"messages": [{"role": "user", "content": "test"}],
|
|
"metadata": {"tags": ["custom-tag"]},
|
|
}
|
|
|
|
general_settings = {"reject_clientside_metadata_tags": False}
|
|
|
|
# Create a mock request object
|
|
mock_request = MagicMock(spec=Request)
|
|
|
|
# Create a valid token for the test
|
|
valid_token = UserAPIKeyAuth(token="test-token", models=["gpt-3.5-turbo"])
|
|
|
|
# Should not raise an exception
|
|
result = await common_checks(
|
|
request_body=request_body,
|
|
team_object=None,
|
|
user_object=None,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings=general_settings,
|
|
route="/chat/completions",
|
|
llm_router=None,
|
|
proxy_logging_obj=MagicMock(),
|
|
valid_token=valid_token,
|
|
request=mock_request,
|
|
)
|
|
|
|
assert result is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_reject_clientside_metadata_tags_not_set_with_tags():
|
|
"""Test that common_checks allows request with metadata.tags when reject_clientside_metadata_tags is not set."""
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
request_body = {
|
|
"model": "gpt-3.5-turbo",
|
|
"messages": [{"role": "user", "content": "test"}],
|
|
"metadata": {"tags": ["custom-tag"]},
|
|
}
|
|
|
|
general_settings = {} # No reject_clientside_metadata_tags setting
|
|
|
|
# Create a mock request object
|
|
mock_request = MagicMock(spec=Request)
|
|
|
|
# Create a valid token for the test
|
|
valid_token = UserAPIKeyAuth(token="test-token", models=["gpt-3.5-turbo"])
|
|
|
|
# Should not raise an exception
|
|
result = await common_checks(
|
|
request_body=request_body,
|
|
team_object=None,
|
|
user_object=None,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings=general_settings,
|
|
route="/chat/completions",
|
|
llm_router=None,
|
|
proxy_logging_obj=MagicMock(),
|
|
valid_token=valid_token,
|
|
request=mock_request,
|
|
)
|
|
|
|
assert result is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_reject_clientside_metadata_tags_non_llm_route():
|
|
"""Test that reject_clientside_metadata_tags check only applies to LLM API routes."""
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
request_body = {
|
|
"metadata": {"tags": ["custom-tag"]},
|
|
}
|
|
|
|
general_settings = {"reject_clientside_metadata_tags": True}
|
|
|
|
# Create a mock request object
|
|
mock_request = MagicMock(spec=Request)
|
|
|
|
# Create a valid token for the test
|
|
valid_token = UserAPIKeyAuth(token="test-token", models=["gpt-3.5-turbo"])
|
|
|
|
# Create an admin user object for the management route
|
|
admin_user = LiteLLM_UserTable(
|
|
user_id="admin-user",
|
|
user_role=LitellmUserRoles.PROXY_ADMIN,
|
|
)
|
|
|
|
# Should not raise an exception for non-LLM route
|
|
result = await common_checks(
|
|
request_body=request_body,
|
|
team_object=None,
|
|
user_object=admin_user,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings=general_settings,
|
|
route="/key/generate", # Management route, not LLM route
|
|
llm_router=None,
|
|
proxy_logging_obj=MagicMock(),
|
|
valid_token=valid_token,
|
|
request=mock_request,
|
|
)
|
|
|
|
assert result is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_reject_clientside_metadata_tags_allows_key_tags_without_client_tags():
|
|
"""Key metadata.tags are injected after the reject check; requests without
|
|
client metadata.tags must not be blocked when reject_clientside_metadata_tags is on.
|
|
"""
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
request_body = {
|
|
"model": "gpt-3.5-turbo",
|
|
"messages": [{"role": "user", "content": "test"}],
|
|
}
|
|
|
|
general_settings = {"reject_clientside_metadata_tags": True}
|
|
mock_request = MagicMock(spec=Request)
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
models=["gpt-3.5-turbo"],
|
|
metadata={"tags": ["engineering"]},
|
|
)
|
|
|
|
with patch(
|
|
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
|
|
new_callable=AsyncMock,
|
|
return_value={},
|
|
):
|
|
result = await common_checks(
|
|
request_body=request_body,
|
|
team_object=None,
|
|
user_object=None,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings=general_settings,
|
|
route="/chat/completions",
|
|
llm_router=None,
|
|
proxy_logging_obj=MagicMock(),
|
|
valid_token=valid_token,
|
|
request=mock_request,
|
|
)
|
|
|
|
assert result is True
|
|
assert request_body["metadata"]["tags"] == ["engineering"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize(
|
|
"route",
|
|
[
|
|
"/bedrock/model/us.anthropic.claude-sonnet-4-6/invoke",
|
|
"/v1/messages",
|
|
],
|
|
)
|
|
async def test_common_checks_metadata_route_keeps_key_tags_out_of_provider_metadata(
|
|
route,
|
|
):
|
|
"""GH#30629: on routes that track tags in litellm_metadata (bedrock, /v1/messages,
|
|
responses, ...) key-level tags must land in litellm_metadata, never in the
|
|
provider-facing metadata field (Bedrock rejects non-user_id metadata with HTTP 400).
|
|
The auth-time pre-seed keys off LITELLM_METADATA_ROUTES, so hardcoding a single route
|
|
or dropping the pre-seed makes apply_key_tags_pre_auth fall back to metadata; this
|
|
guards that regression.
|
|
"""
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
request_body = {"messages": [{"role": "user", "content": "test"}]}
|
|
|
|
mock_request = MagicMock(spec=Request)
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
metadata={"tags": ["engineering"]},
|
|
)
|
|
|
|
with patch(
|
|
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
|
|
new_callable=AsyncMock,
|
|
return_value={},
|
|
):
|
|
result = await common_checks(
|
|
request_body=request_body,
|
|
team_object=None,
|
|
user_object=None,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings={},
|
|
route=route,
|
|
llm_router=None,
|
|
proxy_logging_obj=MagicMock(),
|
|
valid_token=valid_token,
|
|
request=mock_request,
|
|
)
|
|
|
|
assert result is True
|
|
assert request_body["litellm_metadata"]["tags"] == ["engineering"]
|
|
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 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."""
|
|
|
|
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"""
|
|
alert_triggered = False
|
|
captured_call_info = None
|
|
|
|
class MockProxyLogging:
|
|
async def budget_alerts(self, type, user_info):
|
|
nonlocal alert_triggered, captured_call_info
|
|
alert_triggered = True
|
|
captured_call_info = user_info
|
|
assert type == "soft_budget"
|
|
assert isinstance(user_info, CallInfo)
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
spend=100.0,
|
|
soft_budget=50.0,
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
team_alias="test-team-alias",
|
|
org_id="test-org",
|
|
key_alias="test-key",
|
|
max_budget=200.0,
|
|
)
|
|
|
|
user_obj = LiteLLM_UserTable(
|
|
user_id="test-user",
|
|
user_email="test@example.com",
|
|
max_budget=None,
|
|
)
|
|
|
|
proxy_logging_obj = MockProxyLogging()
|
|
|
|
await _virtual_key_soft_budget_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
user_obj=user_obj,
|
|
)
|
|
|
|
await asyncio.sleep(0.1)
|
|
|
|
assert alert_triggered is True
|
|
assert captured_call_info is not None
|
|
assert captured_call_info.user_email == "test@example.com"
|
|
assert captured_call_info.token == "test-token"
|
|
assert captured_call_info.spend == 100.0
|
|
assert captured_call_info.soft_budget == 50.0
|
|
assert captured_call_info.max_budget == 200.0
|
|
assert captured_call_info.user_id == "test-user"
|
|
assert captured_call_info.team_id == "test-team"
|
|
assert captured_call_info.team_alias == "test-team-alias"
|
|
assert captured_call_info.organization_id == "test-org"
|
|
assert captured_call_info.key_alias == "test-key"
|
|
assert captured_call_info.event_group == Litellm_EntityType.KEY
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_virtual_key_soft_budget_check_without_user_obj():
|
|
"""Test _virtual_key_soft_budget_check sets user_email to None when user_obj is not provided"""
|
|
alert_triggered = False
|
|
captured_call_info = None
|
|
|
|
class MockProxyLogging:
|
|
async def budget_alerts(self, type, user_info):
|
|
nonlocal alert_triggered, captured_call_info
|
|
alert_triggered = True
|
|
captured_call_info = user_info
|
|
assert type == "soft_budget"
|
|
assert isinstance(user_info, CallInfo)
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
spend=100.0,
|
|
soft_budget=50.0,
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
key_alias="test-key",
|
|
)
|
|
|
|
proxy_logging_obj = MockProxyLogging()
|
|
|
|
await _virtual_key_soft_budget_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
user_obj=None,
|
|
)
|
|
|
|
await asyncio.sleep(0.1)
|
|
|
|
assert alert_triggered is True
|
|
assert captured_call_info is not None
|
|
assert captured_call_info.user_email is None
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"spend, soft_budget, expect_alert",
|
|
[
|
|
(100.0, 50.0, True), # Over soft budget
|
|
(50.0, 50.0, True), # At soft budget
|
|
(25.0, 50.0, False), # Under soft budget
|
|
(100.0, None, False), # No soft budget set
|
|
],
|
|
)
|
|
@pytest.mark.asyncio
|
|
async def test_virtual_key_soft_budget_check_scenarios(
|
|
spend, soft_budget, expect_alert
|
|
):
|
|
"""Test _virtual_key_soft_budget_check with various spend and soft_budget scenarios"""
|
|
alert_triggered = False
|
|
|
|
class MockProxyLogging:
|
|
async def budget_alerts(self, type, user_info):
|
|
nonlocal alert_triggered
|
|
alert_triggered = True
|
|
assert type == "soft_budget"
|
|
assert isinstance(user_info, CallInfo)
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
spend=spend,
|
|
soft_budget=soft_budget,
|
|
user_id="test-user",
|
|
key_alias="test-key",
|
|
)
|
|
|
|
proxy_logging_obj = MockProxyLogging()
|
|
|
|
await _virtual_key_soft_budget_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
user_obj=None,
|
|
)
|
|
|
|
await asyncio.sleep(0.1)
|
|
|
|
assert (
|
|
alert_triggered == expect_alert
|
|
), f"Expected alert_triggered to be {expect_alert} for spend={spend}, soft_budget={soft_budget}"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_virtual_key_max_budget_alert_check_with_user_obj():
|
|
"""Test _virtual_key_max_budget_alert_check includes user_email when user_obj is provided"""
|
|
alert_triggered = False
|
|
captured_call_info = None
|
|
|
|
class MockProxyLogging:
|
|
async def budget_alerts(self, type, user_info):
|
|
nonlocal alert_triggered, captured_call_info
|
|
alert_triggered = True
|
|
captured_call_info = user_info
|
|
assert type == "max_budget_alert"
|
|
assert isinstance(user_info, CallInfo)
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
spend=90.0,
|
|
max_budget=100.0,
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
team_alias="test-team-alias",
|
|
org_id="test-org",
|
|
key_alias="test-key",
|
|
soft_budget=50.0,
|
|
)
|
|
|
|
user_obj = LiteLLM_UserTable(
|
|
user_id="test-user",
|
|
user_email="test@example.com",
|
|
max_budget=None,
|
|
)
|
|
|
|
proxy_logging_obj = MockProxyLogging()
|
|
|
|
await _virtual_key_max_budget_alert_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
user_obj=user_obj,
|
|
)
|
|
|
|
await asyncio.sleep(0.1)
|
|
|
|
assert alert_triggered is True
|
|
assert captured_call_info is not None
|
|
assert captured_call_info.user_email == "test@example.com"
|
|
assert captured_call_info.token == "test-token"
|
|
assert captured_call_info.spend == 90.0
|
|
assert captured_call_info.max_budget == 100.0
|
|
assert captured_call_info.soft_budget == 50.0
|
|
assert captured_call_info.user_id == "test-user"
|
|
assert captured_call_info.team_id == "test-team"
|
|
assert captured_call_info.team_alias == "test-team-alias"
|
|
assert captured_call_info.organization_id == "test-org"
|
|
assert captured_call_info.key_alias == "test-key"
|
|
assert captured_call_info.event_group == Litellm_EntityType.KEY
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_virtual_key_max_budget_alert_check_without_user_obj():
|
|
"""Test _virtual_key_max_budget_alert_check sets user_email to None when user_obj is not provided"""
|
|
alert_triggered = False
|
|
captured_call_info = None
|
|
|
|
class MockProxyLogging:
|
|
async def budget_alerts(self, type, user_info):
|
|
nonlocal alert_triggered, captured_call_info
|
|
alert_triggered = True
|
|
captured_call_info = user_info
|
|
assert type == "max_budget_alert"
|
|
assert isinstance(user_info, CallInfo)
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
spend=90.0,
|
|
max_budget=100.0,
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
key_alias="test-key",
|
|
)
|
|
|
|
proxy_logging_obj = MockProxyLogging()
|
|
|
|
await _virtual_key_max_budget_alert_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
user_obj=None,
|
|
)
|
|
|
|
await asyncio.sleep(0.1)
|
|
|
|
assert alert_triggered is True
|
|
assert captured_call_info is not None
|
|
assert captured_call_info.user_email is None
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"spend, max_budget, expect_alert",
|
|
[
|
|
(80.0, 100.0, True), # At 80% threshold (alert threshold)
|
|
(90.0, 100.0, True), # Above threshold, below max_budget
|
|
(79.0, 100.0, False), # Below threshold
|
|
(100.0, 100.0, False), # At max_budget (not below, so no alert)
|
|
(110.0, 100.0, False), # Above max_budget (already exceeded)
|
|
(100.0, None, False), # No max_budget set
|
|
(0.0, 100.0, False), # Spend is 0
|
|
],
|
|
)
|
|
@pytest.mark.asyncio
|
|
async def test_virtual_key_max_budget_alert_check_scenarios(
|
|
spend, max_budget, expect_alert
|
|
):
|
|
"""Test _virtual_key_max_budget_alert_check with various spend and max_budget scenarios"""
|
|
alert_triggered = False
|
|
|
|
class MockProxyLogging:
|
|
async def budget_alerts(self, type, user_info):
|
|
nonlocal alert_triggered
|
|
alert_triggered = True
|
|
assert type == "max_budget_alert"
|
|
assert isinstance(user_info, CallInfo)
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
spend=spend,
|
|
max_budget=max_budget,
|
|
user_id="test-user",
|
|
key_alias="test-key",
|
|
)
|
|
|
|
proxy_logging_obj = MockProxyLogging()
|
|
|
|
await _virtual_key_max_budget_alert_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
user_obj=None,
|
|
)
|
|
|
|
await asyncio.sleep(0.1)
|
|
|
|
assert (
|
|
alert_triggered == expect_alert
|
|
), f"Expected alert_triggered to be {expect_alert} for spend={spend}, max_budget={max_budget}"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_virtual_key_max_budget_alert_check_with_multi_threshold_map():
|
|
"""Test that max_budget_alert_emails map from metadata is attached to CallInfo on the new path"""
|
|
alert_triggered = False
|
|
captured_call_info = None
|
|
|
|
class MockProxyLogging:
|
|
async def budget_alerts(self, type, user_info):
|
|
nonlocal alert_triggered, captured_call_info
|
|
alert_triggered = True
|
|
captured_call_info = user_info
|
|
|
|
alert_config = {
|
|
"50": ["finance@co.com"],
|
|
"75": ["finance@co.com", "bu_lead@co.com"],
|
|
}
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
spend=60.0,
|
|
max_budget=100.0,
|
|
user_id="test-user",
|
|
key_alias="test-key",
|
|
metadata={"max_budget_alert_emails": alert_config},
|
|
)
|
|
user_obj = LiteLLM_UserTable(
|
|
user_id="test-user",
|
|
user_email="owner@co.com",
|
|
max_budget=None,
|
|
)
|
|
|
|
await _virtual_key_max_budget_alert_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=MockProxyLogging(),
|
|
user_obj=user_obj,
|
|
)
|
|
await asyncio.sleep(0.1)
|
|
|
|
assert alert_triggered is True
|
|
assert captured_call_info is not None
|
|
assert captured_call_info.max_budget_alert_emails == alert_config
|
|
assert captured_call_info.user_email == "owner@co.com"
|
|
assert captured_call_info.event_group == Litellm_EntityType.KEY
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_virtual_key_max_budget_alert_check_old_path_no_map():
|
|
"""Test that old single-threshold path is used when no max_budget_alert_emails in metadata"""
|
|
alert_triggered = False
|
|
captured_call_info = None
|
|
|
|
class MockProxyLogging:
|
|
async def budget_alerts(self, type, user_info):
|
|
nonlocal alert_triggered, captured_call_info
|
|
alert_triggered = True
|
|
captured_call_info = user_info
|
|
|
|
# spend=90 is above 80% of 100 → old path should fire
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
spend=90.0,
|
|
max_budget=100.0,
|
|
user_id="test-user",
|
|
key_alias="test-key",
|
|
metadata={},
|
|
)
|
|
|
|
await _virtual_key_max_budget_alert_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=MockProxyLogging(),
|
|
user_obj=None,
|
|
)
|
|
await asyncio.sleep(0.1)
|
|
|
|
assert alert_triggered is True
|
|
assert captured_call_info is not None
|
|
assert captured_call_info.max_budget_alert_emails is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_virtual_key_max_budget_alert_check_old_path_below_threshold_no_alert():
|
|
"""Test that old path does NOT fire when spend is below 80% and no map is set"""
|
|
alert_triggered = False
|
|
|
|
class MockProxyLogging:
|
|
async def budget_alerts(self, type, user_info):
|
|
nonlocal alert_triggered
|
|
alert_triggered = True
|
|
|
|
# spend=50 is below 80% of 100 → should NOT fire
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
spend=50.0,
|
|
max_budget=100.0,
|
|
user_id="test-user",
|
|
key_alias="test-key",
|
|
metadata={},
|
|
)
|
|
|
|
await _virtual_key_max_budget_alert_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=MockProxyLogging(),
|
|
user_obj=None,
|
|
)
|
|
await asyncio.sleep(0.1)
|
|
|
|
assert alert_triggered is False
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_virtual_key_max_budget_alert_check_global_fallback():
|
|
"""Test that litellm.default_key_max_budget_alert_emails is used when key metadata has no map"""
|
|
alert_triggered = False
|
|
captured_call_info = None
|
|
|
|
class MockProxyLogging:
|
|
async def budget_alerts(self, type, user_info):
|
|
nonlocal alert_triggered, captured_call_info
|
|
alert_triggered = True
|
|
captured_call_info = user_info
|
|
|
|
global_config = {
|
|
"50": ["global-finance@co.com"],
|
|
"75": ["global-finance@co.com", "global-lead@co.com"],
|
|
}
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
spend=60.0,
|
|
max_budget=100.0,
|
|
user_id="test-user",
|
|
key_alias="test-key",
|
|
metadata={}, # no per-key config
|
|
)
|
|
|
|
import litellm
|
|
|
|
original = litellm.default_key_max_budget_alert_emails
|
|
try:
|
|
litellm.default_key_max_budget_alert_emails = global_config
|
|
await _virtual_key_max_budget_alert_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=MockProxyLogging(),
|
|
user_obj=None,
|
|
)
|
|
await asyncio.sleep(0.1)
|
|
|
|
assert alert_triggered is True
|
|
assert captured_call_info.max_budget_alert_emails == global_config
|
|
finally:
|
|
litellm.default_key_max_budget_alert_emails = original
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_virtual_key_max_budget_alert_check_per_key_merges_with_global():
|
|
"""Test that per-key and global configs are additively merged"""
|
|
captured_call_info = None
|
|
|
|
class MockProxyLogging:
|
|
async def budget_alerts(self, type, user_info):
|
|
nonlocal captured_call_info
|
|
captured_call_info = user_info
|
|
|
|
per_key_config = {"50": ["per-key@co.com"]}
|
|
global_config = {"75": ["global@co.com"]}
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
spend=60.0,
|
|
max_budget=100.0,
|
|
user_id="test-user",
|
|
key_alias="test-key",
|
|
metadata={"max_budget_alert_emails": per_key_config},
|
|
)
|
|
|
|
import litellm
|
|
|
|
original = litellm.default_key_max_budget_alert_emails
|
|
try:
|
|
litellm.default_key_max_budget_alert_emails = global_config
|
|
await _virtual_key_max_budget_alert_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=MockProxyLogging(),
|
|
user_obj=None,
|
|
)
|
|
await asyncio.sleep(0.1)
|
|
|
|
# Additive merge: both thresholds present, recipients merged per threshold
|
|
assert captured_call_info.max_budget_alert_emails == {
|
|
"50": ["per-key@co.com"],
|
|
"75": ["global@co.com"],
|
|
}
|
|
finally:
|
|
litellm.default_key_max_budget_alert_emails = original
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_fuzzy_user_object_case_insensitive_email():
|
|
"""Test that _get_fuzzy_user_object uses case-insensitive email lookup"""
|
|
# Setup mock Prisma client
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db = MagicMock()
|
|
mock_prisma.db.litellm_usertable = MagicMock()
|
|
|
|
# Mock user data with mixed case email
|
|
test_user = LiteLLM_UserTable(
|
|
user_id="test_123",
|
|
sso_user_id=None,
|
|
user_email="Test@Example.com", # Mixed case in DB
|
|
organization_memberships=[],
|
|
max_budget=None,
|
|
)
|
|
|
|
# Test: SSO ID not found, find by email with different casing
|
|
mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
|
|
mock_prisma.db.litellm_usertable.find_first = AsyncMock(return_value=test_user)
|
|
|
|
# Search with lowercase email (different from DB)
|
|
result = await _get_fuzzy_user_object(
|
|
prisma_client=mock_prisma,
|
|
sso_user_id=None,
|
|
user_email="test@example.com", # Lowercase search
|
|
)
|
|
|
|
# Verify user was found despite case difference
|
|
assert result == test_user
|
|
|
|
# Verify the query used case-insensitive mode
|
|
mock_prisma.db.litellm_usertable.find_first.assert_called_once()
|
|
call_args = mock_prisma.db.litellm_usertable.find_first.call_args
|
|
assert call_args.kwargs["where"]["user_email"]["equals"] == "test@example.com"
|
|
assert call_args.kwargs["where"]["user_email"]["mode"] == "insensitive"
|
|
assert call_args.kwargs["include"] == {"organization_memberships": True}
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_custom_auth_common_checks_opt_in():
|
|
"""
|
|
Test that common_checks only runs for a custom-auth deployment when
|
|
custom_auth_run_common_checks is explicitly set to True in general_settings.
|
|
|
|
After the centralization refactor, common_checks runs in the
|
|
``user_api_key_auth`` wrapper via ``_run_centralized_common_checks``
|
|
(not inside ``_run_post_custom_auth_checks``). The opt-in flag now
|
|
gates the centralized gate for custom-auth deployments, preserving
|
|
the pre-existing RPS guarantee for custom-auth hot paths.
|
|
"""
|
|
import litellm.proxy.proxy_server as _proxy_server_mod
|
|
from litellm.proxy.auth.user_api_key_auth import _run_centralized_common_checks
|
|
|
|
valid_token = UserAPIKeyAuth(token="test-token", user_id="u1")
|
|
mock_request = MagicMock()
|
|
|
|
def _attrs(flag, user_custom_auth):
|
|
return {
|
|
"prisma_client": None,
|
|
"user_api_key_cache": MagicMock(),
|
|
"proxy_logging_obj": MagicMock(),
|
|
"general_settings": (
|
|
{"custom_auth_run_common_checks": True} if flag else {}
|
|
),
|
|
"llm_router": None,
|
|
"user_custom_auth": user_custom_auth,
|
|
"litellm_proxy_admin_name": "admin",
|
|
"master_key": "sk-test-master",
|
|
}
|
|
|
|
# Default (no flag) with custom auth configured — centralized gate
|
|
# SHOULD skip to preserve custom-auth RPS.
|
|
attrs = _attrs(flag=False, user_custom_auth=AsyncMock())
|
|
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_common:
|
|
await _run_centralized_common_checks(
|
|
user_api_key_auth_obj=valid_token,
|
|
request=mock_request,
|
|
request_data={},
|
|
route="/chat/completions",
|
|
)
|
|
mock_common.assert_not_called()
|
|
finally:
|
|
for k, v in originals.items():
|
|
setattr(_proxy_server_mod, k, v)
|
|
|
|
# With flag=True and custom auth configured — common_checks SHOULD run.
|
|
attrs = _attrs(flag=True, user_custom_auth=AsyncMock())
|
|
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_common:
|
|
await _run_centralized_common_checks(
|
|
user_api_key_auth_obj=valid_token,
|
|
request=mock_request,
|
|
request_data={},
|
|
route="/chat/completions",
|
|
)
|
|
mock_common.assert_called_once()
|
|
finally:
|
|
for k, v in originals.items():
|
|
setattr(_proxy_server_mod, k, v)
|
|
|
|
|
|
# =====================================================================
|
|
# Spend counter budget check tests (v2 — Redis-backed spend counters)
|
|
# =====================================================================
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_virtual_key_budget_check_reads_from_spend_counter():
|
|
"""Budget check should use get_current_spend when counter exists,
|
|
even if cached object shows lower spend."""
|
|
from litellm.proxy.utils import ProxyLogging
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-hashed-token",
|
|
spend=0.0, # stale — counter has 1.5
|
|
max_budget=1.0,
|
|
user_id="test-user",
|
|
)
|
|
|
|
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
|
|
proxy_logging_obj.budget_alerts = AsyncMock()
|
|
|
|
async def mock_get_current_spend(
|
|
counter_key, fallback_spend, max_budget=None, **kwargs
|
|
):
|
|
if counter_key == "spend:key:test-hashed-token":
|
|
return 1.5
|
|
return fallback_spend
|
|
|
|
with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend):
|
|
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
|
await _virtual_key_max_budget_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
assert exc_info.value.current_cost == 1.5
|
|
assert exc_info.value.max_budget == 1.0
|
|
assert exc_info.value.entity_type == "key"
|
|
assert exc_info.value.entity_id == "test-hashed-token"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_virtual_key_budget_check_fallback_no_counter():
|
|
"""When counter doesn't exist, budget check should fall back
|
|
to cached object's spend via fallback_spend."""
|
|
from litellm.proxy.utils import ProxyLogging
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-hashed-token",
|
|
spend=15.0,
|
|
max_budget=10.0,
|
|
user_id="test-user",
|
|
)
|
|
|
|
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
|
|
proxy_logging_obj.budget_alerts = AsyncMock()
|
|
|
|
# get_current_spend returns fallback_spend when no counter exists
|
|
async def mock_get_current_spend(
|
|
counter_key, fallback_spend, max_budget=None, **kwargs
|
|
):
|
|
return fallback_spend
|
|
|
|
with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend):
|
|
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
|
await _virtual_key_max_budget_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
assert exc_info.value.current_cost == 15.0
|
|
|
|
|
|
# =====================================================================
|
|
# Throttle-on-budget-exceeded tests (LIT-3894): an over-budget key that
|
|
# opted in is throttled to a global % of its TPM/RPM instead of blocked.
|
|
# =====================================================================
|
|
|
|
|
|
def _over_budget_token(**overrides) -> UserAPIKeyAuth:
|
|
base = dict(
|
|
token="throttle-token",
|
|
spend=20.0,
|
|
max_budget=10.0,
|
|
user_id="test-user",
|
|
)
|
|
base.update(overrides)
|
|
return UserAPIKeyAuth(**base)
|
|
|
|
|
|
def _patched_spend(value: float):
|
|
async def mock_get_current_spend(
|
|
counter_key, fallback_spend, max_budget=None, **kwargs
|
|
):
|
|
return value
|
|
|
|
return patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend)
|
|
|
|
|
|
def _budget_logging_obj():
|
|
from litellm.proxy.utils import ProxyLogging
|
|
|
|
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
|
|
proxy_logging_obj.budget_alerts = AsyncMock()
|
|
return proxy_logging_obj
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"limit, pct, expected",
|
|
[
|
|
(1000, 0.1, 100),
|
|
(100, 0.1, 10),
|
|
(1, 0.1, 1), # floor would be 0; trickle of 1 keeps the key alive
|
|
(None, 0.1, None),
|
|
(50, 0.5, 25),
|
|
(1000, None, 1000), # no percentage -> limit unchanged
|
|
],
|
|
)
|
|
def test_throttled_limit(limit, pct, expected):
|
|
from litellm.proxy.auth.budget_throttle import throttled_limit
|
|
|
|
assert throttled_limit(limit, pct) == expected
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_budget_exceeded_throttles_instead_of_blocking(monkeypatch):
|
|
monkeypatch.setattr(litellm, "budget_exceeded_throttle_percentage", 0.1)
|
|
valid_token = _over_budget_token(
|
|
tpm_limit=1000,
|
|
rpm_limit=100,
|
|
metadata={"throttle_on_budget_exceeded": True},
|
|
)
|
|
|
|
with _patched_spend(20.0):
|
|
await _virtual_key_max_budget_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=_budget_logging_obj(),
|
|
)
|
|
|
|
# persistent limits are untouched (so the throttle never compounds); the
|
|
# request-scoped percentage is what the rate limiter scales by
|
|
assert valid_token.budget_throttle_pct == 0.1
|
|
assert valid_token.tpm_limit == 1000
|
|
assert valid_token.rpm_limit == 100
|
|
# the request-scoped decision must not leak into serialized responses
|
|
assert "budget_throttle_pct" not in valid_token.model_dump()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_budget_throttle_decision_cleared_before_caching():
|
|
"""The request-scoped throttle decision must not persist into the key cache,
|
|
otherwise it would re-apply (and compound) on every subsequent request."""
|
|
from litellm.proxy.auth.auth_checks import _copy_user_api_key_auth_for_cache
|
|
|
|
valid_token = _over_budget_token(
|
|
tpm_limit=1000, rpm_limit=100, metadata={"throttle_on_budget_exceeded": True}
|
|
)
|
|
valid_token.budget_throttle_pct = 0.1
|
|
|
|
cached = _copy_user_api_key_auth_for_cache(user_api_key_obj=valid_token)
|
|
|
|
assert cached.budget_throttle_pct is None
|
|
assert cached.tpm_limit == 1000
|
|
assert cached.rpm_limit == 100
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_budget_exceeded_throttle_no_configured_limits(monkeypatch):
|
|
monkeypatch.setattr(litellm, "budget_exceeded_throttle_percentage", 0.1)
|
|
valid_token = _over_budget_token(metadata={"throttle_on_budget_exceeded": True})
|
|
assert valid_token.tpm_limit is None
|
|
assert valid_token.rpm_limit is None
|
|
|
|
with _patched_spend(20.0):
|
|
with pytest.raises(litellm.BudgetExceededError):
|
|
await _virtual_key_max_budget_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=_budget_logging_obj(),
|
|
)
|
|
|
|
assert valid_token.budget_throttle_pct is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_budget_exceeded_not_opted_in_still_blocks(monkeypatch):
|
|
monkeypatch.setattr(litellm, "budget_exceeded_throttle_percentage", 0.1)
|
|
valid_token = _over_budget_token(tpm_limit=1000, rpm_limit=100)
|
|
|
|
with _patched_spend(20.0):
|
|
with pytest.raises(litellm.BudgetExceededError):
|
|
await _virtual_key_max_budget_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=_budget_logging_obj(),
|
|
)
|
|
|
|
assert valid_token.budget_throttle_pct is None
|
|
|
|
|
|
@pytest.mark.parametrize("pct", [None, 0, 1.5, -0.1, True])
|
|
@pytest.mark.asyncio
|
|
async def test_budget_exceeded_invalid_percentage_blocks(monkeypatch, pct):
|
|
monkeypatch.setattr(litellm, "budget_exceeded_throttle_percentage", pct)
|
|
valid_token = _over_budget_token(
|
|
tpm_limit=1000,
|
|
rpm_limit=100,
|
|
metadata={"throttle_on_budget_exceeded": True},
|
|
)
|
|
|
|
with _patched_spend(20.0):
|
|
with pytest.raises(litellm.BudgetExceededError):
|
|
await _virtual_key_max_budget_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=_budget_logging_obj(),
|
|
)
|
|
|
|
assert valid_token.budget_throttle_pct is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_under_budget_does_not_throttle(monkeypatch):
|
|
monkeypatch.setattr(litellm, "budget_exceeded_throttle_percentage", 0.1)
|
|
valid_token = _over_budget_token(
|
|
max_budget=100.0,
|
|
tpm_limit=1000,
|
|
rpm_limit=100,
|
|
metadata={"throttle_on_budget_exceeded": True},
|
|
)
|
|
|
|
with _patched_spend(5.0):
|
|
await _virtual_key_max_budget_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=_budget_logging_obj(),
|
|
)
|
|
|
|
assert valid_token.budget_throttle_pct is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_team_budget_check_reads_from_spend_counter():
|
|
"""Team budget check should use get_current_spend when counter exists."""
|
|
from litellm.proxy.utils import ProxyLogging
|
|
|
|
team_object = LiteLLM_TeamTable(
|
|
team_id="test-team",
|
|
spend=0.0, # stale
|
|
max_budget=1.0,
|
|
)
|
|
valid_token = UserAPIKeyAuth(token="test-token", team_id="test-team")
|
|
|
|
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
|
|
proxy_logging_obj.budget_alerts = AsyncMock()
|
|
|
|
async def mock_get_current_spend(
|
|
counter_key, fallback_spend, max_budget=None, **kwargs
|
|
):
|
|
if counter_key == "spend:team:test-team":
|
|
return 1.5
|
|
return fallback_spend
|
|
|
|
with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend):
|
|
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
|
await _team_max_budget_check(
|
|
team_object=team_object,
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
assert exc_info.value.current_cost == 1.5
|
|
assert exc_info.value.entity_type == "team"
|
|
assert exc_info.value.entity_id == "test-team"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_end_user_budget_check_reads_from_spend_counter():
|
|
"""End-user budget check should use get_current_spend when counter exists."""
|
|
end_user_object = LiteLLM_EndUserTable(
|
|
user_id="customer-1",
|
|
blocked=False,
|
|
spend=0.0,
|
|
litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0),
|
|
)
|
|
|
|
async def mock_get_current_spend(
|
|
counter_key, fallback_spend, max_budget=None, **kwargs
|
|
):
|
|
if counter_key == "spend:end_user:customer-1":
|
|
return 1.5
|
|
return fallback_spend
|
|
|
|
with patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend):
|
|
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
|
await _check_end_user_budget(
|
|
end_user_obj=end_user_object,
|
|
route="/chat/completions",
|
|
)
|
|
assert exc_info.value.current_cost == 1.5
|
|
assert exc_info.value.max_budget == 1.0
|
|
assert exc_info.value.entity_type == "end_user"
|
|
assert exc_info.value.entity_id == "customer-1"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_tag_budget_check_reads_from_spend_counter():
|
|
"""Tag budget check should use get_current_spend when counter exists."""
|
|
from litellm.proxy.utils import ProxyLogging
|
|
|
|
tag_object = LiteLLM_TagTable(
|
|
tag_name="paid-tag",
|
|
spend=0.0,
|
|
litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0),
|
|
)
|
|
|
|
async def mock_get_current_spend(
|
|
counter_key, fallback_spend, max_budget=None, **kwargs
|
|
):
|
|
if counter_key == "spend:tag:paid-tag":
|
|
return 1.5
|
|
return fallback_spend
|
|
|
|
with (
|
|
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
|
|
patch(
|
|
"litellm.proxy.auth.auth_checks.get_tag_objects_batch",
|
|
new_callable=AsyncMock,
|
|
return_value={"paid-tag": tag_object},
|
|
),
|
|
):
|
|
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
|
await _tag_max_budget_check(
|
|
request_body={"metadata": {"tags": ["paid-tag"]}},
|
|
prisma_client=MagicMock(),
|
|
user_api_key_cache=MagicMock(),
|
|
proxy_logging_obj=ProxyLogging(user_api_key_cache=None),
|
|
valid_token=UserAPIKeyAuth(token="test-token"),
|
|
)
|
|
assert exc_info.value.current_cost == 1.5
|
|
assert exc_info.value.max_budget == 1.0
|
|
assert exc_info.value.entity_type == "tag"
|
|
assert exc_info.value.entity_id == "paid-tag"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_team_member_budget_check_reads_from_spend_counter():
|
|
"""Team member budget check should use get_current_spend when counter exists."""
|
|
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TeamMembership
|
|
from litellm.proxy.utils import ProxyLogging
|
|
|
|
team_object = LiteLLM_TeamTable(team_id="test-team")
|
|
user_object = LiteLLM_UserTable(user_id="test-user")
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
)
|
|
|
|
team_membership = LiteLLM_TeamMembership(
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
spend=0.0, # stale
|
|
litellm_budget_table=LiteLLM_BudgetTable(max_budget=1.0),
|
|
)
|
|
|
|
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
|
|
|
|
async def mock_get_current_spend(
|
|
counter_key, fallback_spend, max_budget=None, **kwargs
|
|
):
|
|
if counter_key == "spend:team_member:test-user:test-team":
|
|
return 1.5
|
|
return fallback_spend
|
|
|
|
with (
|
|
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
|
|
patch(
|
|
"litellm.proxy.auth.auth_checks.get_team_membership",
|
|
new_callable=AsyncMock,
|
|
return_value=team_membership,
|
|
),
|
|
):
|
|
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
|
await _check_team_member_budget(
|
|
team_object=team_object,
|
|
user_object=user_object,
|
|
valid_token=valid_token,
|
|
prisma_client=MagicMock(),
|
|
user_api_key_cache=MagicMock(),
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
assert exc_info.value.current_cost == 1.5
|
|
assert exc_info.value.entity_type == "team_member"
|
|
assert exc_info.value.entity_id == "test-user:test-team"
|
|
|
|
|
|
class TestGuardrailModificationCheck:
|
|
"""Defense-in-depth: `_guardrail_modification_check` must 403 when the
|
|
caller's metadata attempts to modify any guardrail-related key and the
|
|
team lacks the `modify_guardrails` permission. Checks both the
|
|
historically-covered `guardrails` list and the bypass toggles that
|
|
`_get_admin_metadata` silently ignores at read time.
|
|
"""
|
|
|
|
def _call(self, request_body):
|
|
from litellm.proxy.auth.auth_checks import _guardrail_modification_check
|
|
|
|
team_object = MagicMock()
|
|
team_object.metadata = {} # no permission
|
|
return _guardrail_modification_check(
|
|
request_body=request_body, team_object=team_object
|
|
)
|
|
|
|
def test_noop_when_no_guardrail_keys_present(self):
|
|
# no-op — should return silently
|
|
self._call({"metadata": {"unrelated": "value"}})
|
|
|
|
def test_rejects_guardrails_list(self):
|
|
from fastapi import HTTPException
|
|
|
|
with patch(
|
|
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
|
return_value=False,
|
|
):
|
|
with pytest.raises(HTTPException) as exc:
|
|
self._call({"metadata": {"guardrails": ["custom"]}})
|
|
assert exc.value.status_code == 403
|
|
|
|
def test_rejects_disable_global_guardrails_plural(self):
|
|
from fastapi import HTTPException
|
|
|
|
with patch(
|
|
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
|
return_value=False,
|
|
):
|
|
with pytest.raises(HTTPException) as exc:
|
|
self._call({"metadata": {"disable_global_guardrails": True}})
|
|
assert exc.value.status_code == 403
|
|
|
|
def test_rejects_disable_global_guardrail_singular(self):
|
|
"""VERIA-28's originally-reported singular-key typo variant."""
|
|
from fastapi import HTTPException
|
|
|
|
with patch(
|
|
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
|
return_value=False,
|
|
):
|
|
with pytest.raises(HTTPException) as exc:
|
|
self._call({"metadata": {"disable_global_guardrail": True}})
|
|
assert exc.value.status_code == 403
|
|
|
|
def test_rejects_opted_out_global_guardrails(self):
|
|
from fastapi import HTTPException
|
|
|
|
with patch(
|
|
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
|
return_value=False,
|
|
):
|
|
with pytest.raises(HTTPException) as exc:
|
|
self._call(
|
|
{"metadata": {"opted_out_global_guardrails": ["some_guardrail"]}}
|
|
)
|
|
assert exc.value.status_code == 403
|
|
|
|
@pytest.mark.parametrize(
|
|
"key",
|
|
[
|
|
"guardrails",
|
|
"disable_global_guardrails",
|
|
"disable_global_guardrail",
|
|
"opted_out_global_guardrails",
|
|
],
|
|
)
|
|
@pytest.mark.parametrize("empty_value", [{}, [], "", 0, False])
|
|
def test_rejects_empty_value_modification(self, key, empty_value):
|
|
"""Regression: an explicitly-supplied empty/falsy value still expresses
|
|
intent to modify and must trigger the permission check. Truthiness-based
|
|
gating let callers bypass the check by sending e.g.
|
|
``metadata={"guardrails": {}}``, which downstream evaluation interpreted
|
|
as "disable all guardrails" while the auth layer treated it as no-op.
|
|
"""
|
|
from fastapi import HTTPException
|
|
|
|
with patch(
|
|
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
|
return_value=False,
|
|
):
|
|
with pytest.raises(HTTPException) as exc:
|
|
self._call({"metadata": {key: empty_value}})
|
|
assert exc.value.status_code == 403
|
|
|
|
def test_rejects_injection_via_litellm_metadata_key(self):
|
|
"""Caller can populate the OTHER metadata key; that must also 403."""
|
|
from fastapi import HTTPException
|
|
|
|
with patch(
|
|
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
|
return_value=False,
|
|
):
|
|
with pytest.raises(HTTPException) as exc:
|
|
self._call({"litellm_metadata": {"disable_global_guardrails": True}})
|
|
assert exc.value.status_code == 403
|
|
|
|
def test_rejects_root_level_injection(self):
|
|
"""Top-level injection (`request_body["disable_global_guardrails"]`)
|
|
was VERIA-28's easiest variant to hit — keep it rejected."""
|
|
from fastapi import HTTPException
|
|
|
|
with patch(
|
|
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
|
return_value=False,
|
|
):
|
|
with pytest.raises(HTTPException) as exc:
|
|
self._call({"disable_global_guardrails": True})
|
|
assert exc.value.status_code == 403
|
|
|
|
def test_allows_when_team_has_permission(self):
|
|
with patch(
|
|
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
|
return_value=True,
|
|
):
|
|
# no-op, should not raise
|
|
self._call({"metadata": {"disable_global_guardrails": True}})
|
|
|
|
def test_rejects_string_encoded_metadata_bypass(self):
|
|
"""Regression: attacker sends metadata as JSON string to bypass the
|
|
isinstance(dict) guard. The check must coerce the string to dict
|
|
and evaluate guardrail modification keys inside it."""
|
|
import json as _json
|
|
|
|
from fastapi import HTTPException
|
|
|
|
attacker_payload = {"disable_global_guardrails": True}
|
|
with patch(
|
|
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
|
return_value=False,
|
|
):
|
|
with pytest.raises(HTTPException) as exc:
|
|
self._call({"metadata": _json.dumps(attacker_payload)})
|
|
assert exc.value.status_code == 403
|
|
|
|
def test_rejects_string_encoded_litellm_metadata_bypass(self):
|
|
"""Same bypass via the litellm_metadata key."""
|
|
import json as _json
|
|
|
|
from fastapi import HTTPException
|
|
|
|
attacker_payload = {"guardrails": ["evaded"]}
|
|
with patch(
|
|
"litellm.proxy.guardrails.guardrail_helpers.can_modify_guardrails",
|
|
return_value=False,
|
|
):
|
|
with pytest.raises(HTTPException) as exc:
|
|
self._call({"litellm_metadata": _json.dumps(attacker_payload)})
|
|
assert exc.value.status_code == 403
|
|
|
|
def test_noop_when_string_is_not_json_object(self):
|
|
"""Unparseable strings should not trigger a 403 — they have no keys."""
|
|
self._call({"metadata": "not-json"})
|
|
self._call({"metadata": '"just a string"'})
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_team_member_budget_check_falls_back_to_team_default_budget_id():
|
|
"""When a member's TeamMembership has no linked budget row, the check
|
|
should fall back to team.metadata["team_member_budget_id"] and still
|
|
enforce the cap. Pre-fix, this path silently skipped enforcement."""
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.proxy._types import LiteLLM_TeamMembership
|
|
from litellm.proxy.utils import ProxyLogging
|
|
|
|
team_object = LiteLLM_TeamTable(
|
|
team_id="test-team",
|
|
metadata={"team_member_budget_id": "budget-default"},
|
|
)
|
|
user_object = LiteLLM_UserTable(user_id="test-user")
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
)
|
|
|
|
# Membership row without an attached budget.
|
|
team_membership = LiteLLM_TeamMembership(
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
spend=0.0,
|
|
budget_id=None,
|
|
litellm_budget_table=None,
|
|
)
|
|
|
|
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
|
|
|
|
fake_budget_row = MagicMock()
|
|
fake_budget_row.max_budget = 50.0
|
|
fake_budget_row.dict = MagicMock(
|
|
return_value={"budget_id": "budget-default", "max_budget": 50.0}
|
|
)
|
|
|
|
prisma_client = MagicMock()
|
|
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(
|
|
return_value=fake_budget_row
|
|
)
|
|
|
|
async def mock_get_current_spend(
|
|
counter_key, fallback_spend, max_budget=None, **kwargs
|
|
):
|
|
if counter_key == "spend:team_member:test-user:test-team":
|
|
return 70.0
|
|
return fallback_spend
|
|
|
|
user_api_key_cache = DualCache()
|
|
|
|
with (
|
|
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
|
|
patch(
|
|
"litellm.proxy.auth.auth_checks.get_team_membership",
|
|
new_callable=AsyncMock,
|
|
return_value=team_membership,
|
|
),
|
|
):
|
|
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
|
await _check_team_member_budget(
|
|
team_object=team_object,
|
|
user_object=user_object,
|
|
valid_token=valid_token,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
assert exc_info.value.current_cost == 70.0
|
|
assert exc_info.value.max_budget == 50.0
|
|
|
|
# First call did perform the fallback DB lookup.
|
|
prisma_client.db.litellm_budgettable.find_unique.assert_awaited_once()
|
|
|
|
# Second call hits the cached budget row, no additional prisma read.
|
|
prisma_client.db.litellm_budgettable.find_unique.reset_mock()
|
|
with (
|
|
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
|
|
patch(
|
|
"litellm.proxy.auth.auth_checks.get_team_membership",
|
|
new_callable=AsyncMock,
|
|
return_value=team_membership,
|
|
),
|
|
):
|
|
with pytest.raises(litellm.BudgetExceededError) as second_exc_info:
|
|
await _check_team_member_budget(
|
|
team_object=team_object,
|
|
user_object=user_object,
|
|
valid_token=valid_token,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
# The cached $50 cap is still being applied (not a coincidental skip)
|
|
assert second_exc_info.value.current_cost == 70.0
|
|
assert second_exc_info.value.max_budget == 50.0
|
|
prisma_client.db.litellm_budgettable.find_unique.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_team_member_budget_check_per_member_override_wins_over_team_default():
|
|
"""If a member has a per-member budget AND the team carries a
|
|
team_member_budget_id default, the per-member value wins and the
|
|
fallback prisma lookup is never performed."""
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TeamMembership
|
|
from litellm.proxy.utils import ProxyLogging
|
|
|
|
team_object = LiteLLM_TeamTable(
|
|
team_id="test-team",
|
|
metadata={"team_member_budget_id": "budget-default"},
|
|
)
|
|
user_object = LiteLLM_UserTable(user_id="test-user")
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
)
|
|
|
|
team_membership = LiteLLM_TeamMembership(
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
spend=0.0,
|
|
budget_id="budget-override",
|
|
litellm_budget_table=LiteLLM_BudgetTable(max_budget=200.0),
|
|
)
|
|
|
|
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
|
|
|
|
# Team-default row resolves to $50. If the fallback fired (it must
|
|
# not here), spend $70 would exceed that $50 cap and raise.
|
|
fake_budget_row = MagicMock()
|
|
fake_budget_row.max_budget = 50.0
|
|
|
|
prisma_client = MagicMock()
|
|
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(
|
|
return_value=fake_budget_row
|
|
)
|
|
|
|
mocked_spend = 70.0
|
|
|
|
async def mock_get_current_spend(
|
|
counter_key, fallback_spend, max_budget=None, **kwargs
|
|
):
|
|
if counter_key == "spend:team_member:test-user:test-team":
|
|
return mocked_spend
|
|
return fallback_spend
|
|
|
|
# 1. spend ($70) < per-member cap ($200) → no raise, no fallback lookup.
|
|
with (
|
|
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
|
|
patch(
|
|
"litellm.proxy.auth.auth_checks.get_team_membership",
|
|
new_callable=AsyncMock,
|
|
return_value=team_membership,
|
|
),
|
|
):
|
|
await _check_team_member_budget(
|
|
team_object=team_object,
|
|
user_object=user_object,
|
|
valid_token=valid_token,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=DualCache(),
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
|
|
prisma_client.db.litellm_budgettable.find_unique.assert_not_awaited()
|
|
|
|
# 2. Now push spend above the per-member cap ($200). Must raise with
|
|
# max_budget=200 to prove the per-member cap is the value being
|
|
# enforced (not just that enforcement silently skipped).
|
|
mocked_spend = 250.0
|
|
with (
|
|
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
|
|
patch(
|
|
"litellm.proxy.auth.auth_checks.get_team_membership",
|
|
new_callable=AsyncMock,
|
|
return_value=team_membership,
|
|
),
|
|
):
|
|
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
|
await _check_team_member_budget(
|
|
team_object=team_object,
|
|
user_object=user_object,
|
|
valid_token=valid_token,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=DualCache(),
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
assert exc_info.value.current_cost == 250.0
|
|
assert exc_info.value.max_budget == 200.0
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_team_member_budget_check_null_clone_falls_back_to_team_default():
|
|
"""Per-member NULL max_budget falls through to the team default cap."""
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TeamMembership
|
|
from litellm.proxy.utils import ProxyLogging
|
|
|
|
team_object = LiteLLM_TeamTable(
|
|
team_id="test-team",
|
|
metadata={"team_member_budget_id": "budget-default"},
|
|
)
|
|
user_object = LiteLLM_UserTable(user_id="test-user")
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
)
|
|
|
|
# Per-member row exists with NULL max_budget (the cloned-from-incomplete-default case).
|
|
team_membership = LiteLLM_TeamMembership(
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
spend=0.0,
|
|
budget_id="budget-clone",
|
|
litellm_budget_table=LiteLLM_BudgetTable(max_budget=None),
|
|
)
|
|
|
|
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
|
|
|
|
fake_default_row = MagicMock()
|
|
fake_default_row.max_budget = 65.0
|
|
fake_default_row.dict = MagicMock(
|
|
return_value={"budget_id": "budget-default", "max_budget": 65.0}
|
|
)
|
|
|
|
prisma_client = MagicMock()
|
|
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(
|
|
return_value=fake_default_row
|
|
)
|
|
|
|
async def mock_get_current_spend(
|
|
counter_key, fallback_spend, max_budget=None, **kwargs
|
|
):
|
|
if counter_key == "spend:team_member:test-user:test-team":
|
|
return 500.0
|
|
return fallback_spend
|
|
|
|
with (
|
|
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
|
|
patch(
|
|
"litellm.proxy.auth.auth_checks.get_team_membership",
|
|
new_callable=AsyncMock,
|
|
return_value=team_membership,
|
|
),
|
|
):
|
|
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
|
await _check_team_member_budget(
|
|
team_object=team_object,
|
|
user_object=user_object,
|
|
valid_token=valid_token,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=DualCache(),
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
|
|
assert exc_info.value.current_cost == 500.0
|
|
assert exc_info.value.max_budget == 65.0
|
|
prisma_client.db.litellm_budgettable.find_unique.assert_awaited_once()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_team_member_budget_check_null_clone_with_null_default_skips_enforcement():
|
|
"""When per-member and team default are both NULL, enforcement still skips."""
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TeamMembership
|
|
from litellm.proxy.utils import ProxyLogging
|
|
|
|
team_object = LiteLLM_TeamTable(
|
|
team_id="test-team",
|
|
metadata={"team_member_budget_id": "budget-default"},
|
|
)
|
|
user_object = LiteLLM_UserTable(user_id="test-user")
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
)
|
|
|
|
team_membership = LiteLLM_TeamMembership(
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
spend=0.0,
|
|
budget_id="budget-clone",
|
|
litellm_budget_table=LiteLLM_BudgetTable(max_budget=None),
|
|
)
|
|
|
|
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
|
|
|
|
fake_default_row = MagicMock()
|
|
fake_default_row.max_budget = None
|
|
fake_default_row.dict = MagicMock(
|
|
return_value={"budget_id": "budget-default", "max_budget": None}
|
|
)
|
|
|
|
prisma_client = MagicMock()
|
|
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(
|
|
return_value=fake_default_row
|
|
)
|
|
|
|
async def mock_get_current_spend(
|
|
counter_key, fallback_spend, max_budget=None, **kwargs
|
|
):
|
|
if counter_key == "spend:team_member:test-user:test-team":
|
|
return 1000.0
|
|
return fallback_spend
|
|
|
|
with (
|
|
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
|
|
patch(
|
|
"litellm.proxy.auth.auth_checks.get_team_membership",
|
|
new_callable=AsyncMock,
|
|
return_value=team_membership,
|
|
),
|
|
):
|
|
# No raise: both rows are NULL, so enforcement is correctly skipped.
|
|
await _check_team_member_budget(
|
|
team_object=team_object,
|
|
user_object=user_object,
|
|
valid_token=valid_token,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=DualCache(),
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_team_member_budget_check_zero_team_default_treated_as_no_cap():
|
|
"""A team default budget with max_budget=0.0 (likely a stale/accidental
|
|
write) must not block every member. The fallback path treats 0 as
|
|
"no cap"; per-member rows still respect 0 as an explicit disable."""
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.proxy._types import LiteLLM_TeamMembership
|
|
from litellm.proxy.utils import ProxyLogging
|
|
|
|
team_object = LiteLLM_TeamTable(
|
|
team_id="test-team",
|
|
metadata={"team_member_budget_id": "budget-default"},
|
|
)
|
|
user_object = LiteLLM_UserTable(user_id="test-user")
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
)
|
|
|
|
# No per-member row -> falls through to team default.
|
|
team_membership = LiteLLM_TeamMembership(
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
spend=0.0,
|
|
budget_id=None,
|
|
litellm_budget_table=None,
|
|
)
|
|
|
|
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
|
|
|
|
# Team default budget row with max_budget=0.0 (the regression trigger).
|
|
fake_default_row = MagicMock()
|
|
fake_default_row.max_budget = 0.0
|
|
fake_default_row.dict = MagicMock(
|
|
return_value={"budget_id": "budget-default", "max_budget": 0.0}
|
|
)
|
|
|
|
prisma_client = MagicMock()
|
|
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(
|
|
return_value=fake_default_row
|
|
)
|
|
|
|
async def mock_get_current_spend(
|
|
counter_key, fallback_spend, max_budget=None, **kwargs
|
|
):
|
|
if counter_key == "spend:team_member:test-user:test-team":
|
|
return 0.0
|
|
return fallback_spend
|
|
|
|
with (
|
|
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
|
|
patch(
|
|
"litellm.proxy.auth.auth_checks.get_team_membership",
|
|
new_callable=AsyncMock,
|
|
return_value=team_membership,
|
|
),
|
|
):
|
|
# No raise: 0.0 cap is treated as "no cap configured".
|
|
await _check_team_member_budget(
|
|
team_object=team_object,
|
|
user_object=user_object,
|
|
valid_token=valid_token,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=DualCache(),
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_team_member_budget_check_zero_per_member_row_still_blocks():
|
|
"""A per-member row with max_budget=0.0 is treated as an explicit admin
|
|
disable - enforcement still blocks. Only the team-default fallback
|
|
path treats 0 as no cap."""
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.proxy._types import LiteLLM_BudgetTable, LiteLLM_TeamMembership
|
|
from litellm.proxy.utils import ProxyLogging
|
|
|
|
team_object = LiteLLM_TeamTable(
|
|
team_id="test-team",
|
|
metadata={"team_member_budget_id": "budget-default"},
|
|
)
|
|
user_object = LiteLLM_UserTable(user_id="test-user")
|
|
valid_token = UserAPIKeyAuth(
|
|
token="test-token",
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
)
|
|
|
|
# Per-member row with max_budget=0.0 - admin intent: disable this user.
|
|
team_membership = LiteLLM_TeamMembership(
|
|
user_id="test-user",
|
|
team_id="test-team",
|
|
spend=0.0,
|
|
budget_id="budget-disable",
|
|
litellm_budget_table=LiteLLM_BudgetTable(max_budget=0.0),
|
|
)
|
|
|
|
proxy_logging_obj = ProxyLogging(user_api_key_cache=None)
|
|
|
|
prisma_client = MagicMock()
|
|
prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=None)
|
|
|
|
async def mock_get_current_spend(
|
|
counter_key, fallback_spend, max_budget=None, **kwargs
|
|
):
|
|
if counter_key == "spend:team_member:test-user:test-team":
|
|
return 0.0
|
|
return fallback_spend
|
|
|
|
with (
|
|
patch("litellm.proxy.proxy_server.get_current_spend", mock_get_current_spend),
|
|
patch(
|
|
"litellm.proxy.auth.auth_checks.get_team_membership",
|
|
new_callable=AsyncMock,
|
|
return_value=team_membership,
|
|
),
|
|
):
|
|
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
|
await _check_team_member_budget(
|
|
team_object=team_object,
|
|
user_object=user_object,
|
|
valid_token=valid_token,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=DualCache(),
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
assert exc_info.value.max_budget == 0.0
|
|
|
|
|
|
# --- resolve_and_validate_end_user_id ---------------------------------------
|
|
|
|
|
|
@pytest.fixture
|
|
def _validate_flag_on(monkeypatch):
|
|
"""Enable opt-in DB validation for the duration of a test."""
|
|
import litellm
|
|
|
|
monkeypatch.setattr(litellm, "validate_end_user_id_in_db", True)
|
|
monkeypatch.setattr(litellm, "max_end_user_budget_id", None)
|
|
|
|
|
|
def _validation_cache():
|
|
cache = MagicMock()
|
|
cache.async_get_cache = AsyncMock(return_value=None)
|
|
cache.async_set_cache = AsyncMock()
|
|
return cache
|
|
|
|
|
|
def _patch_validation_helpers(monkeypatch, *, end_user=None, user=None, fuzzy=None):
|
|
"""Stub out the DB helpers resolve_and_validate_end_user_id delegates to."""
|
|
from litellm.proxy.auth import auth_checks
|
|
|
|
monkeypatch.setattr(
|
|
auth_checks, "get_end_user_object", AsyncMock(return_value=end_user)
|
|
)
|
|
monkeypatch.setattr(auth_checks, "get_user_object", AsyncMock(return_value=user))
|
|
monkeypatch.setattr(
|
|
auth_checks, "_get_fuzzy_user_object", AsyncMock(return_value=fuzzy)
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_end_user_returns_none_for_none_input(
|
|
_validate_flag_on, monkeypatch
|
|
):
|
|
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
|
|
|
_patch_validation_helpers(monkeypatch)
|
|
cache = _validation_cache()
|
|
assert (
|
|
await resolve_and_validate_end_user_id(
|
|
raw_end_user_id=None,
|
|
prisma_client=MagicMock(),
|
|
user_api_key_cache=cache,
|
|
)
|
|
is None
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_end_user_passes_through_when_flag_disabled(monkeypatch):
|
|
"""Default behaviour: flag is off, arbitrary ids pass through untouched."""
|
|
import litellm
|
|
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
|
|
|
monkeypatch.setattr(litellm, "validate_end_user_id_in_db", False)
|
|
_patch_validation_helpers(monkeypatch)
|
|
cache = _validation_cache()
|
|
|
|
result = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id="codex-session-abc",
|
|
prisma_client=MagicMock(),
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert result == "codex-session-abc"
|
|
cache.async_set_cache.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_end_user_passes_through_when_no_prisma_client(
|
|
_validate_flag_on, monkeypatch
|
|
):
|
|
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
|
|
|
_patch_validation_helpers(monkeypatch)
|
|
cache = _validation_cache()
|
|
|
|
result = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id="alice@example.com",
|
|
prisma_client=None,
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert result == "alice@example.com"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_end_user_matches_end_user_table(_validate_flag_on, monkeypatch):
|
|
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
|
|
|
_patch_validation_helpers(monkeypatch, end_user=MagicMock())
|
|
cache = _validation_cache()
|
|
|
|
result = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id="customer-123",
|
|
prisma_client=MagicMock(),
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert result == "customer-123"
|
|
cache.async_set_cache.assert_awaited_once()
|
|
kwargs = cache.async_set_cache.await_args.kwargs
|
|
assert kwargs["key"] == "end_user_validation:customer-123"
|
|
assert kwargs["value"] == "valid"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_end_user_matches_user_table_by_user_id(
|
|
_validate_flag_on, monkeypatch
|
|
):
|
|
from litellm.proxy.auth import auth_checks
|
|
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
|
|
|
_patch_validation_helpers(monkeypatch, user=MagicMock())
|
|
cache = _validation_cache()
|
|
|
|
result = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id="user-xyz",
|
|
prisma_client=MagicMock(),
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert result == "user-xyz"
|
|
# email fallback should not run for a non-email input
|
|
auth_checks._get_fuzzy_user_object.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_end_user_matches_user_table_by_email(
|
|
_validate_flag_on, monkeypatch
|
|
):
|
|
"""Email-shaped ids route through get_user_object with user_email set.
|
|
|
|
The fuzzy lookup must happen inside get_user_object so it shares the
|
|
_should_check_db throttle and user_api_key_cache — no direct raw
|
|
Prisma calls on the auth path.
|
|
"""
|
|
from litellm.proxy.auth import auth_checks
|
|
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
|
|
|
_patch_validation_helpers(monkeypatch, user=MagicMock())
|
|
cache = _validation_cache()
|
|
|
|
result = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id="Alice@Example.com",
|
|
prisma_client=MagicMock(),
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert result == "Alice@Example.com"
|
|
auth_checks.get_user_object.assert_awaited_once()
|
|
user_kwargs = auth_checks.get_user_object.await_args.kwargs
|
|
assert user_kwargs["user_id"] == "Alice@Example.com"
|
|
assert user_kwargs["user_email"] == "Alice@Example.com"
|
|
# email branch must not bypass the cached helper with a raw fuzzy call
|
|
auth_checks._get_fuzzy_user_object.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_end_user_non_email_id_does_not_pass_user_email(
|
|
_validate_flag_on, monkeypatch
|
|
):
|
|
"""Non-email ids skip the email fuzzy path to avoid a pointless DB hit."""
|
|
from litellm.proxy.auth import auth_checks
|
|
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
|
|
|
_patch_validation_helpers(monkeypatch, user=MagicMock())
|
|
cache = _validation_cache()
|
|
|
|
await resolve_and_validate_end_user_id(
|
|
raw_end_user_id="user-xyz",
|
|
prisma_client=MagicMock(),
|
|
user_api_key_cache=cache,
|
|
)
|
|
auth_checks.get_user_object.assert_awaited_once()
|
|
user_kwargs = auth_checks.get_user_object.await_args.kwargs
|
|
assert user_kwargs["user_email"] is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_end_user_drops_codex_opaque_identifier(
|
|
_validate_flag_on, monkeypatch
|
|
):
|
|
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
|
|
|
_patch_validation_helpers(monkeypatch) # all helpers return None
|
|
cache = _validation_cache()
|
|
|
|
codex_id = (
|
|
"user_8a4a360c36621665b341e06fb76041d9b6def732bb183eea148d4abc9d97c1de"
|
|
"_account__session_a2bce4a5-8887-44ef-b491-fbf0a55c6569"
|
|
)
|
|
result = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id=codex_id,
|
|
prisma_client=MagicMock(),
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert result is None
|
|
cache.async_set_cache.assert_awaited_once()
|
|
kwargs = cache.async_set_cache.await_args.kwargs
|
|
assert kwargs["value"] == "invalid"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_end_user_preserves_id_when_default_budget_configured(
|
|
_validate_flag_on, monkeypatch
|
|
):
|
|
"""Don't drop unregistered ids when litellm.max_end_user_budget_id is set.
|
|
|
|
The default end-user budget is applied downstream when the id is present
|
|
but not found in the db — dropping the id here would bypass those limits.
|
|
"""
|
|
import litellm
|
|
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
|
|
|
monkeypatch.setattr(litellm, "max_end_user_budget_id", "default-budget")
|
|
_patch_validation_helpers(monkeypatch)
|
|
cache = _validation_cache()
|
|
|
|
result = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id="new-customer",
|
|
prisma_client=MagicMock(),
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert result == "new-customer"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_end_user_drops_unknown_email(_validate_flag_on, monkeypatch):
|
|
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
|
|
|
_patch_validation_helpers(monkeypatch)
|
|
cache = _validation_cache()
|
|
|
|
result = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id="stranger@example.com",
|
|
prisma_client=MagicMock(),
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert result is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_end_user_uses_cached_valid_result(
|
|
_validate_flag_on, monkeypatch
|
|
):
|
|
from litellm.proxy.auth import auth_checks
|
|
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
|
|
|
_patch_validation_helpers(monkeypatch)
|
|
cache = _validation_cache()
|
|
cache.async_get_cache = AsyncMock(return_value="valid")
|
|
|
|
result = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id="alice@example.com",
|
|
prisma_client=MagicMock(),
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert result == "alice@example.com"
|
|
auth_checks.get_end_user_object.assert_not_awaited()
|
|
auth_checks.get_user_object.assert_not_awaited()
|
|
auth_checks._get_fuzzy_user_object.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_end_user_uses_cached_invalid_result(
|
|
_validate_flag_on, monkeypatch
|
|
):
|
|
from litellm.proxy.auth import auth_checks
|
|
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
|
|
|
_patch_validation_helpers(monkeypatch, end_user=MagicMock())
|
|
cache = _validation_cache()
|
|
cache.async_get_cache = AsyncMock(return_value="invalid")
|
|
|
|
result = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id="bogus",
|
|
prisma_client=MagicMock(),
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert result is None
|
|
# Despite a matching row configured, helpers aren't called — cache wins.
|
|
auth_checks.get_end_user_object.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_end_user_swallows_db_errors_and_returns_none(
|
|
_validate_flag_on, monkeypatch
|
|
):
|
|
from litellm.proxy.auth import auth_checks
|
|
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
|
|
|
monkeypatch.setattr(
|
|
auth_checks,
|
|
"get_end_user_object",
|
|
AsyncMock(side_effect=Exception("db down")),
|
|
)
|
|
monkeypatch.setattr(
|
|
auth_checks,
|
|
"get_user_object",
|
|
AsyncMock(side_effect=Exception("db down")),
|
|
)
|
|
cache = _validation_cache()
|
|
|
|
result = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id="alice@example.com",
|
|
prisma_client=MagicMock(),
|
|
user_api_key_cache=cache,
|
|
)
|
|
# DB errors shouldn't raise through the auth path — treat as unknown.
|
|
assert result is None
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_resolve_end_user(_validate_flag_on, monkeypatch):
|
|
"""Verify that resolve_and_validate_end_user_id does NOT raise BudgetExceededError.
|
|
|
|
Note: As of the refactor that moved _check_end_user_budget out of
|
|
get_end_user_object, budget enforcement now happens in common_checks().
|
|
|
|
The end-user validation path should return the user ID regardless of budget status.
|
|
Budget enforcement for end users happens later in common_checks() via
|
|
_check_end_user_budget(), which respects skip_budget_checks for zero-cost models.
|
|
|
|
This test verifies that even when get_end_user_object returns a user with a budget,
|
|
resolve_and_validate_end_user_id does not block the request - budget enforcement
|
|
is deferred to common_checks() where skip_budget_checks logic can be applied.
|
|
"""
|
|
from litellm.proxy.auth import auth_checks
|
|
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
|
|
|
# Mock get_end_user_object to return a user with budget info
|
|
# (simulating a user who may have exceeded their budget)
|
|
mock_end_user = MagicMock()
|
|
mock_end_user.user_id = "customer-over-budget"
|
|
monkeypatch.setattr(
|
|
auth_checks,
|
|
"get_end_user_object",
|
|
AsyncMock(return_value=mock_end_user),
|
|
)
|
|
cache = _validation_cache()
|
|
|
|
# resolve_and_validate_end_user_id should return the user ID without raising
|
|
# BudgetExceededError - budget enforcement happens in common_checks()
|
|
result = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id="customer-over-budget",
|
|
prisma_client=MagicMock(),
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert result == "customer-over-budget"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_cache_team_object_writes_team_id_and_invalidates_team_alias():
|
|
"""
|
|
Regression pin for LIT-3244 patch/1.86.0 follow-up.
|
|
|
|
`_cache_team_object` is the canonical "refresh this team" primitive.
|
|
Two cache keys are in play:
|
|
- "team_id:<id>" — used by `get_team_object(team_id=...)`,
|
|
i.e. API-key auth and JWT-with-team_id_jwt_field
|
|
- "team_alias:<alias>" — used by `get_team_object_by_alias(team_alias=...)`,
|
|
i.e. JWT-with-team_alias_jwt_field
|
|
|
|
Invariants this test pins:
|
|
1. Writes the team_id-keyed entry with the refreshed object (team_id
|
|
is the table PK — guaranteed unique, safe to write).
|
|
2. DELETES (does NOT write) the team_alias-keyed entry. `team_alias`
|
|
has no UNIQUE constraint in schema.prisma, so writing it from
|
|
this generic refresh path would let a team admin who renames
|
|
their team to collide with another team's alias silently
|
|
overwrite the cached team for JWT-by-alias auth (veria-ai
|
|
review on #28739). Deleting forces the next JWT-by-alias
|
|
reader through `get_team_object_by_alias`, which enforces
|
|
len(teams)==1 before populating the cache.
|
|
3. When team_alias is None, NO alias-key operation happens (no
|
|
delete of an empty-keyed entry, no spurious write).
|
|
4. DELETES the team_id-keyed entry from the internal usage cache
|
|
BEFORE the fresh write (LIT-4391). `_get_team_object_from_cache` no
|
|
longer reads the internal usage cache (LIT-5944), but the delete
|
|
protects mixed-version rolling deploys where older workers still do.
|
|
"""
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
|
|
from litellm.proxy.auth.auth_checks import _cache_team_object
|
|
|
|
base_team_row = {
|
|
"team_id": "team-1234",
|
|
"team_alias": "H-Capacity",
|
|
"models": ["openai/*", "bedrock-claude-sonnet-4"],
|
|
}
|
|
|
|
# ===== team_alias is set =====
|
|
team_table = LiteLLM_TeamTableCachedObj(**base_team_row)
|
|
cache = MagicMock()
|
|
cache.async_set_cache = AsyncMock()
|
|
cache.delete_cache = MagicMock()
|
|
logging_obj = MagicMock()
|
|
logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
|
|
|
|
await _cache_team_object(
|
|
team_id="team-1234",
|
|
team_table=team_table,
|
|
user_api_key_cache=cache,
|
|
proxy_logging_obj=logging_obj,
|
|
)
|
|
|
|
# (1) team_id-keyed write fires with the refreshed object
|
|
written_keys = [
|
|
(c.kwargs.get("key") or c.args[0])
|
|
for c in cache.async_set_cache.await_args_list
|
|
]
|
|
assert written_keys == ["team_id:team-1234"], (
|
|
"Only the team_id-keyed write should fire; the alias key must be "
|
|
"deleted, NOT written. "
|
|
f"Got writes: {written_keys}"
|
|
)
|
|
written_value = (
|
|
cache.async_set_cache.await_args.kwargs.get("value")
|
|
or cache.async_set_cache.await_args.args[1]
|
|
)
|
|
assert written_value is team_table
|
|
|
|
# (2) team_alias-keyed entry is deleted in BOTH the in-memory cache
|
|
# and the Redis dual cache (mirrors _delete_cache_key_object pattern).
|
|
cache.delete_cache.assert_called_once_with(key="team_alias:H-Capacity")
|
|
|
|
# (4) internal usage cache: team_id entry deleted BEFORE the fresh
|
|
# write, alias entry deleted as before.
|
|
internal_deleted_keys = [
|
|
(c.kwargs.get("key") or c.args[0])
|
|
for c in logging_obj.internal_usage_cache.dual_cache.async_delete_cache.await_args_list
|
|
]
|
|
assert internal_deleted_keys == ["team_id:team-1234", "team_alias:H-Capacity"]
|
|
|
|
# ===== team_alias is None: no alias-key operation =====
|
|
aliasless = LiteLLM_TeamTableCachedObj(**{**base_team_row, "team_alias": None})
|
|
cache2 = MagicMock()
|
|
cache2.async_set_cache = AsyncMock()
|
|
cache2.delete_cache = MagicMock()
|
|
logging_obj2 = MagicMock()
|
|
logging_obj2.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
|
|
|
|
await _cache_team_object(
|
|
team_id="team-no-alias",
|
|
team_table=aliasless,
|
|
user_api_key_cache=cache2,
|
|
proxy_logging_obj=logging_obj2,
|
|
)
|
|
|
|
cache2.delete_cache.assert_not_called()
|
|
logging_obj2.internal_usage_cache.dual_cache.async_delete_cache.assert_awaited_once_with(
|
|
key="team_id:team-no-alias"
|
|
)
|
|
written_keys_aliasless = [
|
|
(c.kwargs.get("key") or c.args[0])
|
|
for c in cache2.async_set_cache.await_args_list
|
|
]
|
|
assert written_keys_aliasless == ["team_id:team-no-alias"]
|
|
|
|
|
|
class _SharedFakeRedis(RedisCache):
|
|
"""Dict-backed stand-in for the single Redis that both
|
|
``user_api_key_cache`` (enable_redis_auth_cache) and
|
|
``proxy_logging_obj.internal_usage_cache.dual_cache`` share in the
|
|
LIT-4391 deployment topology. Only the methods DualCache calls are
|
|
implemented; ``super().__init__`` is skipped intentionally."""
|
|
|
|
def __init__(self):
|
|
self._store: dict = {}
|
|
|
|
async def async_set_cache(self, key, value, **kwargs):
|
|
self._store[key] = json.dumps(value)
|
|
|
|
async def async_get_cache(self, key, **kwargs):
|
|
raw = self._store.get(key)
|
|
return json.loads(raw) if raw is not None else None
|
|
|
|
async def async_delete_cache(self, key):
|
|
self._store.pop(key, None)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_team_update_not_shadowed_by_internal_usage_cache_lit_4391():
|
|
"""
|
|
Regression test for LIT-4391: keys with models=["all-team-models"] kept
|
|
getting 403 team_model_access_denied for models added via /team/update.
|
|
|
|
`_get_team_object_from_cache` used to consult the internal usage cache
|
|
BEFORE `user_api_key_cache` (removed in LIT-5944; this test now also
|
|
guards against reintroducing that read).
|
|
When both share one Redis (enable_redis_auth_cache),
|
|
any team read backfills the internal cache's in-memory tier with the team
|
|
object. `_cache_team_object` (the /team/update refresh) only wrote
|
|
`user_api_key_cache`, so that backfilled copy kept shadowing the update
|
|
until its TTL expired — and the auth-time write-back then pushed the stale
|
|
copy back into the shared Redis, making the staleness self-sustaining.
|
|
|
|
Pins:
|
|
1. After `_cache_team_object` writes an updated team, `get_team_object`
|
|
returns the UPDATED model list even though the internal usage cache's
|
|
in-memory tier was backfilled with the pre-update team.
|
|
2. The shared Redis still holds the updated team afterwards — the
|
|
internal-cache invalidation must happen BEFORE the fresh write, or it
|
|
would wipe the value it just wrote.
|
|
"""
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
|
|
from litellm.proxy.auth.auth_checks import _cache_team_object, get_team_object
|
|
|
|
team_id = "team-lit-4391"
|
|
shared_redis = _SharedFakeRedis()
|
|
user_api_key_cache = UserApiKeyCache(redis_cache=shared_redis)
|
|
proxy_logging_obj = MagicMock()
|
|
proxy_logging_obj.internal_usage_cache.dual_cache = DualCache(
|
|
redis_cache=shared_redis,
|
|
default_in_memory_ttl=300,
|
|
)
|
|
prisma_client = MagicMock()
|
|
|
|
await _cache_team_object(
|
|
team_id=team_id,
|
|
team_table=LiteLLM_TeamTableCachedObj(team_id=team_id, models=["model-a"]),
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
|
|
primed = await get_team_object(
|
|
team_id=team_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
assert primed is not None and primed.models == ["model-a"]
|
|
|
|
await _cache_team_object(
|
|
team_id=team_id,
|
|
team_table=LiteLLM_TeamTableCachedObj(
|
|
team_id=team_id, models=["model-a", "model-b"]
|
|
),
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
|
|
refreshed = await get_team_object(
|
|
team_id=team_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
assert refreshed is not None and refreshed.models == ["model-a", "model-b"], (
|
|
"get_team_object served a stale team allowlist after _cache_team_object "
|
|
f"refreshed it. Got models={refreshed.models if refreshed else None}"
|
|
)
|
|
|
|
redis_copy = await shared_redis.async_get_cache(f"team_id:{team_id}")
|
|
assert redis_copy is not None and redis_copy["models"] == ["model-a", "model-b"], (
|
|
"The shared Redis lost the refreshed team object — the internal-cache "
|
|
"invalidation must run BEFORE the fresh write, not after. "
|
|
f"Got: {redis_copy}"
|
|
)
|
|
|
|
|
|
class _CountingFakeRedis(_SharedFakeRedis):
|
|
"""Counts per-key Redis round-trips so tests can pin the number of
|
|
network operations a code path issues."""
|
|
|
|
def __init__(self):
|
|
super().__init__()
|
|
self.get_calls: int = 0
|
|
|
|
async def async_get_cache(self, key, **kwargs):
|
|
self.get_calls += 1
|
|
return await super().async_get_cache(key, **kwargs)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_warm_team_object_reads_issue_no_redis_ops_lit_5944():
|
|
"""
|
|
Regression test for LIT-5944: project/team-scoped virtual-key requests
|
|
paid ~4 awaited Redis GETs per request just to re-read the team object.
|
|
|
|
`_get_team_object_from_cache` used to consult
|
|
`proxy_logging_obj.internal_usage_cache.dual_cache` (in-memory TTL 1s,
|
|
Redis-backed) BEFORE `user_api_key_cache`. Nothing writes team objects
|
|
into that internal cache — `_cache_team_object` only DELETES the key
|
|
there — so when `user_api_key_cache` has no Redis tier the shared Redis
|
|
key stays absent forever and every team lookup in the auth hot path
|
|
(4 call sites per chat-completion request) became a guaranteed-miss
|
|
Redis round-trip, saturating the event loop at high TPS.
|
|
|
|
Pins: once `_cache_team_object` has cached a team, repeated
|
|
`get_team_object` reads are served from `user_api_key_cache`'s in-memory
|
|
tier and issue ZERO Redis operations.
|
|
"""
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
|
|
from litellm.proxy.auth.auth_checks import _cache_team_object, get_team_object
|
|
|
|
team_id = "team-lit-5944"
|
|
counting_redis = _CountingFakeRedis()
|
|
user_api_key_cache = UserApiKeyCache()
|
|
proxy_logging_obj = MagicMock()
|
|
proxy_logging_obj.internal_usage_cache.dual_cache = DualCache(
|
|
redis_cache=counting_redis,
|
|
default_in_memory_ttl=1,
|
|
)
|
|
prisma_client = MagicMock()
|
|
|
|
await _cache_team_object(
|
|
team_id=team_id,
|
|
team_table=LiteLLM_TeamTableCachedObj(team_id=team_id, models=["model-a"]),
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
|
|
for _ in range(4):
|
|
team_obj = await get_team_object(
|
|
team_id=team_id,
|
|
prisma_client=prisma_client,
|
|
user_api_key_cache=user_api_key_cache,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
assert team_obj is not None and team_obj.models == ["model-a"]
|
|
|
|
assert counting_redis.get_calls == 0, (
|
|
"Warm team-object reads must be served from user_api_key_cache's "
|
|
"in-memory tier without any Redis round-trips. "
|
|
f"Got {counting_redis.get_calls} Redis GETs for 4 get_team_object calls."
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_cache_team_object_tolerates_cache_invalidation_failures():
|
|
"""
|
|
Greptile review on the LIT-4391 fix: `_cache_team_object` runs after a
|
|
successful DB fetch (inside `get_team_object`) and after every team
|
|
mutation's DB write. A cache-backend error during the best-effort
|
|
invalidations must NOT fail those operations — otherwise a Redis blip
|
|
turns a healthy team lookup into a 404 and a committed /team/update into
|
|
a 500. The authoritative team_id-keyed write must still happen.
|
|
"""
|
|
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
|
|
from litellm.proxy.auth.auth_checks import _cache_team_object
|
|
|
|
cache = MagicMock()
|
|
cache.async_set_cache = AsyncMock()
|
|
cache.delete_cache = MagicMock(side_effect=Exception("redis down"))
|
|
logging_obj = MagicMock()
|
|
logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock(
|
|
side_effect=Exception("redis down")
|
|
)
|
|
|
|
await _cache_team_object(
|
|
team_id="team-cache-outage",
|
|
team_table=LiteLLM_TeamTableCachedObj(
|
|
team_id="team-cache-outage",
|
|
team_alias="cache-outage-alias",
|
|
models=["model-a"],
|
|
),
|
|
user_api_key_cache=cache,
|
|
proxy_logging_obj=logging_obj,
|
|
)
|
|
|
|
written_keys = [
|
|
(c.kwargs.get("key") or c.args[0])
|
|
for c in cache.async_set_cache.await_args_list
|
|
]
|
|
assert written_keys == ["team_id:team-cache-outage"]
|
|
|
|
|
|
MODEL_DISCOVERY_ROUTES = [
|
|
"/v1/models",
|
|
"/models",
|
|
"/model/info",
|
|
"/v1/model/info",
|
|
"/v2/model/info",
|
|
"/model_group/info",
|
|
]
|
|
|
|
|
|
@pytest.mark.parametrize("route", MODEL_DISCOVERY_ROUTES)
|
|
@pytest.mark.asyncio
|
|
async def test_model_discovery_route_bypasses_team_budget(route):
|
|
"""Regression for #27923: an exhausted team budget must not block model-discovery routes,
|
|
otherwise OpenAI-compatible clients calling GET /v1/models at startup break."""
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
team_object = LiteLLM_TeamTable(team_id="test-team", spend=150.0, max_budget=100.0)
|
|
|
|
result = await common_checks(
|
|
request_body={},
|
|
team_object=team_object,
|
|
user_object=None,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings={},
|
|
route=route,
|
|
llm_router=None,
|
|
proxy_logging_obj=AsyncMock(),
|
|
valid_token=UserAPIKeyAuth(token="test-token", team_id="test-team"),
|
|
request=MagicMock(),
|
|
)
|
|
|
|
assert result is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_model_discovery_route_bypasses_user_budget():
|
|
"""Regression for #27923: an exhausted user budget must not block model discovery."""
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
user_object = LiteLLM_UserTable(user_id="test-user", spend=100.0, max_budget=50.0)
|
|
|
|
result = await common_checks(
|
|
request_body={},
|
|
team_object=None,
|
|
user_object=user_object,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings={},
|
|
route="/v1/models",
|
|
llm_router=None,
|
|
proxy_logging_obj=AsyncMock(),
|
|
valid_token=UserAPIKeyAuth(token="test-token", user_id="test-user"),
|
|
request=MagicMock(),
|
|
)
|
|
|
|
assert result is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_side_effectful_info_route_still_enforces_budget():
|
|
"""#27923 keeps the bypass narrow: /health/services can fire Slack/email/webhook test
|
|
messages, so an exhausted budget must still block it. Widening the exemption back to
|
|
is_info_route() would regress this."""
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
team_object = LiteLLM_TeamTable(team_id="test-team", spend=150.0, max_budget=100.0)
|
|
|
|
with pytest.raises(litellm.BudgetExceededError):
|
|
await common_checks(
|
|
request_body={},
|
|
team_object=team_object,
|
|
user_object=None,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings={},
|
|
route="/health/services",
|
|
llm_router=None,
|
|
proxy_logging_obj=AsyncMock(),
|
|
valid_token=UserAPIKeyAuth(token="test-token", team_id="test-team"),
|
|
request=MagicMock(),
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_inference_route_still_enforces_team_budget():
|
|
"""Control for #27923: inference routes stay fully budget-enforced."""
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
team_object = LiteLLM_TeamTable(team_id="test-team", spend=150.0, max_budget=100.0)
|
|
|
|
with pytest.raises(litellm.BudgetExceededError):
|
|
await common_checks(
|
|
request_body={},
|
|
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=AsyncMock(),
|
|
valid_token=UserAPIKeyAuth(token="test-token", team_id="test-team"),
|
|
request=MagicMock(),
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_virtual_key_max_budget_error_names_the_key():
|
|
"""BudgetExceededError for a virtual key must name the key (alias + masked key)
|
|
so operators don't have to reverse-map a spend figure back to a key."""
|
|
valid_token = UserAPIKeyAuth(
|
|
token="hashed-token",
|
|
key_alias="payments-prod",
|
|
key_name="sk-...um_g",
|
|
max_budget=10.0,
|
|
spend=0.0,
|
|
)
|
|
proxy_logging_obj = MagicMock()
|
|
proxy_logging_obj.budget_alerts = AsyncMock()
|
|
|
|
with patch(
|
|
"litellm.proxy.proxy_server.get_current_spend",
|
|
new=AsyncMock(return_value=25.0),
|
|
):
|
|
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
|
await _virtual_key_max_budget_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
|
|
message = str(exc_info.value)
|
|
assert "payments-prod" in message
|
|
assert "sk-...um_g" in message
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_virtual_key_max_budget_not_exceeded_does_not_raise():
|
|
"""Spend below the configured budget must not raise."""
|
|
valid_token = UserAPIKeyAuth(
|
|
token="hashed-token",
|
|
key_alias="payments-prod",
|
|
max_budget=10.0,
|
|
spend=0.0,
|
|
)
|
|
proxy_logging_obj = MagicMock()
|
|
proxy_logging_obj.budget_alerts = AsyncMock()
|
|
|
|
with patch(
|
|
"litellm.proxy.proxy_server.get_current_spend",
|
|
new=AsyncMock(return_value=1.0),
|
|
):
|
|
await _virtual_key_max_budget_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
)
|
|
|
|
|
|
class _TTLCapturingInMemoryCache(InMemoryCache):
|
|
"""Records the ``ttl`` DualCache forwards into the in-memory layer."""
|
|
|
|
def __init__(self) -> None:
|
|
super().__init__()
|
|
self.last_ttl = None
|
|
|
|
def set_cache(self, key, value, **kwargs): # type: ignore[override]
|
|
self.last_ttl = kwargs.get("ttl")
|
|
super().set_cache(key, value, **kwargs)
|
|
|
|
|
|
class TestManagementObjectTTLHonored:
|
|
"""
|
|
Regression for LIT-3338. ``_cache_management_object`` is the central writer on
|
|
the reported ``get_key_object -> _cache_key_object -> _cache_management_object``
|
|
path. It must cache for the configured ``user_api_key_cache_ttl`` (propagated to
|
|
``default_in_memory_ttl``) rather than the hardcoded 60s management default.
|
|
"""
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_uses_configured_user_api_key_cache_ttl(self):
|
|
mem = _TTLCapturingInMemoryCache()
|
|
cache = UserApiKeyCache(in_memory_cache=mem, default_in_memory_ttl=300)
|
|
|
|
await _cache_management_object(
|
|
key="team_id:lit-3338",
|
|
value=UserAPIKeyAuth(token="hash-lit-3338"),
|
|
user_api_key_cache=cache,
|
|
proxy_logging_obj=None,
|
|
model_type=UserAPIKeyAuth,
|
|
)
|
|
|
|
assert mem.last_ttl == 300
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_falls_back_to_management_default_when_unconfigured(self):
|
|
mem = _TTLCapturingInMemoryCache()
|
|
cache = UserApiKeyCache(in_memory_cache=mem)
|
|
assert cache.default_in_memory_ttl is None
|
|
|
|
await _cache_management_object(
|
|
key="team_id:lit-3338-default",
|
|
value=UserAPIKeyAuth(token="hash-default"),
|
|
user_api_key_cache=cache,
|
|
proxy_logging_obj=None,
|
|
model_type=UserAPIKeyAuth,
|
|
)
|
|
|
|
assert mem.last_ttl == DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL
|
|
|
|
|
|
class _BudgetSpendConcurrencyProbe:
|
|
"""Stand-in for get_current_spend that pins how many scope checks are in flight.
|
|
|
|
Each call registers itself, records the peak simultaneous count, and blocks on
|
|
``release`` until the test lets it proceed. ``all_arrived`` only fires once
|
|
``expected`` distinct scope reads are suspended here at the same time, which can
|
|
happen only if common_checks gathers the per-scope reads instead of awaiting
|
|
them one after another.
|
|
"""
|
|
|
|
def __init__(self, expected: int):
|
|
self.expected = expected
|
|
self.in_flight = 0
|
|
self.max_in_flight = 0
|
|
self.all_arrived = asyncio.Event()
|
|
self.release = asyncio.Event()
|
|
|
|
async def __call__(self, *args, **kwargs) -> float:
|
|
self.in_flight += 1
|
|
self.max_in_flight = max(self.max_in_flight, self.in_flight)
|
|
if self.in_flight >= self.expected:
|
|
self.all_arrived.set()
|
|
try:
|
|
await self.release.wait()
|
|
finally:
|
|
self.in_flight -= 1
|
|
return 0.0
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_common_checks_budget_reads_run_concurrently():
|
|
"""Independent per-scope budget reads in common_checks must run concurrently.
|
|
|
|
team max, team window, key window, and end-user each read a distinct spend
|
|
counter with no cross-scope dependency. With the gather they are all suspended
|
|
in get_current_spend simultaneously; reverting to sequential awaits leaves only
|
|
one in flight at a time, so ``all_arrived`` never fires and this test times out.
|
|
"""
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
team = LiteLLM_TeamTable(
|
|
team_id="t1",
|
|
spend=0.0,
|
|
max_budget=100.0,
|
|
budget_limits=[{"budget_duration": "1d", "max_budget": 100.0}],
|
|
)
|
|
token = UserAPIKeyAuth(
|
|
token="k1",
|
|
budget_limits=[{"budget_duration": "1d", "max_budget": 100.0}],
|
|
)
|
|
end_user = LiteLLM_EndUserTable(
|
|
user_id="eu1",
|
|
blocked=False,
|
|
spend=0.0,
|
|
litellm_budget_table=LiteLLM_BudgetTable(max_budget=100.0),
|
|
)
|
|
|
|
probe = _BudgetSpendConcurrencyProbe(expected=4)
|
|
|
|
with patch("litellm.proxy.proxy_server.prisma_client", None), patch(
|
|
"litellm.proxy.proxy_server.get_current_spend", probe
|
|
):
|
|
task = asyncio.create_task(
|
|
common_checks(
|
|
request_body={"messages": [{"role": "user", "content": "hi"}]},
|
|
team_object=team,
|
|
user_object=None,
|
|
end_user_object=end_user,
|
|
global_proxy_spend=None,
|
|
general_settings={},
|
|
route="/chat/completions",
|
|
llm_router=None,
|
|
proxy_logging_obj=MagicMock(),
|
|
valid_token=token,
|
|
request=MagicMock(spec=Request),
|
|
)
|
|
)
|
|
try:
|
|
await asyncio.wait_for(probe.all_arrived.wait(), timeout=3.0)
|
|
assert probe.max_in_flight == 4
|
|
finally:
|
|
probe.release.set()
|
|
assert await task is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_common_checks_budget_gather_raises_highest_priority_scope():
|
|
"""A gathered scope that is over budget must still raise BudgetExceededError.
|
|
|
|
When more than one scope is over budget the error from the highest-priority
|
|
scope (team, matching the previous sequential order) propagates; when only a
|
|
lower-priority scope (end-user) is over budget its error still surfaces. This
|
|
fails if any scope is dropped from the gather or if errors are swallowed.
|
|
"""
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
async def _spend_by_counter(counter_key, fallback_spend, max_budget=None, **kwargs):
|
|
if counter_key == "spend:team:t1":
|
|
return _spend_by_counter.team
|
|
if counter_key == "spend:end_user:eu1":
|
|
return _spend_by_counter.end_user
|
|
return 0.0
|
|
|
|
team = LiteLLM_TeamTable(team_id="t1", spend=0.0, max_budget=100.0)
|
|
end_user = LiteLLM_EndUserTable(
|
|
user_id="eu1",
|
|
blocked=False,
|
|
spend=0.0,
|
|
litellm_budget_table=LiteLLM_BudgetTable(max_budget=100.0),
|
|
)
|
|
|
|
async def _run():
|
|
return await common_checks(
|
|
request_body={"messages": [{"role": "user", "content": "hi"}]},
|
|
team_object=team,
|
|
user_object=None,
|
|
end_user_object=end_user,
|
|
global_proxy_spend=None,
|
|
general_settings={},
|
|
route="/chat/completions",
|
|
llm_router=None,
|
|
proxy_logging_obj=MagicMock(),
|
|
valid_token=None,
|
|
request=MagicMock(spec=Request),
|
|
)
|
|
|
|
with patch("litellm.proxy.proxy_server.prisma_client", None), patch(
|
|
"litellm.proxy.proxy_server.get_current_spend", _spend_by_counter
|
|
):
|
|
# Both team and end-user over budget: team wins on priority.
|
|
_spend_by_counter.team = 999.0
|
|
_spend_by_counter.end_user = 999.0
|
|
with pytest.raises(litellm.BudgetExceededError) as both_over:
|
|
await _run()
|
|
assert "Team=t1" in str(both_over.value)
|
|
|
|
# Only the lower-priority end-user scope over budget: its error still raises.
|
|
_spend_by_counter.team = 0.0
|
|
_spend_by_counter.end_user = 999.0
|
|
with pytest.raises(litellm.BudgetExceededError) as end_user_over:
|
|
await _run()
|
|
assert "End User=eu1" in str(end_user_over.value)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_common_checks_personal_user_budget_blocks_in_gather():
|
|
"""The personal-key user budget scope is enforced inside the gather.
|
|
|
|
For a personal key (no team) whose user is over budget, the gathered user
|
|
check must raise BudgetExceededError. This guards the relocated personal
|
|
user-budget read and fails if that scope is dropped from the gather.
|
|
"""
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
user = LiteLLM_UserTable(user_id="u1", spend=0.0, max_budget=100.0)
|
|
token = UserAPIKeyAuth(token="k1", user_id="u1")
|
|
|
|
async def _spend_by_counter(counter_key, fallback_spend, max_budget=None, **kwargs):
|
|
return 999.0 if counter_key == "spend:user:u1" else 0.0
|
|
|
|
with patch("litellm.proxy.proxy_server.prisma_client", None), patch(
|
|
"litellm.proxy.proxy_server.get_current_spend", _spend_by_counter
|
|
):
|
|
with pytest.raises(litellm.BudgetExceededError) as over:
|
|
await common_checks(
|
|
request_body={"messages": [{"role": "user", "content": "hi"}]},
|
|
team_object=None,
|
|
user_object=user,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings={},
|
|
route="/chat/completions",
|
|
llm_router=None,
|
|
proxy_logging_obj=MagicMock(),
|
|
valid_token=token,
|
|
request=MagicMock(spec=Request),
|
|
)
|
|
assert "User=u1" in str(over.value)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_common_checks_personal_user_budget_skipped_for_team_key():
|
|
"""A user's personal max_budget does not apply to a team-scoped key.
|
|
|
|
Team keys are governed by the team (and team-member) budgets only; the key
|
|
owner's personal budget is deliberately out of scope. This asserts the read
|
|
path lets a team key through even when the user is far over their personal
|
|
budget, and fails if personal enforcement is reintroduced for team keys.
|
|
"""
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
user = LiteLLM_UserTable(user_id="u1", spend=0.0, max_budget=100.0)
|
|
team = LiteLLM_TeamTable(team_id="t1", spend=0.0, max_budget=1000.0)
|
|
token = UserAPIKeyAuth(token="k1", user_id="u1", team_id="t1")
|
|
|
|
async def _spend_by_counter(counter_key, fallback_spend, max_budget=None, **kwargs):
|
|
return 999.0 if counter_key == "spend:user:u1" else 0.0
|
|
|
|
async def _no_membership(*args, **kwargs):
|
|
return None
|
|
|
|
with patch("litellm.proxy.proxy_server.prisma_client", None), patch(
|
|
"litellm.proxy.proxy_server.get_current_spend", _spend_by_counter
|
|
), patch("litellm.proxy.auth.auth_checks.get_team_membership", _no_membership):
|
|
result = await common_checks(
|
|
request_body={"messages": [{"role": "user", "content": "hi"}]},
|
|
team_object=team,
|
|
user_object=user,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings={},
|
|
route="/chat/completions",
|
|
llm_router=None,
|
|
proxy_logging_obj=MagicMock(),
|
|
valid_token=token,
|
|
request=MagicMock(spec=Request),
|
|
)
|
|
assert result is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_common_checks_personal_user_budget_enforced_on_team_key_when_flag_enabled():
|
|
"""general_settings.apply_user_budget_to_team_keys opts a deployment into
|
|
charging the key owner's personal budget on team-scoped keys too.
|
|
|
|
Same fixture as the default-off test above, so a regression that ignores the
|
|
flag lets this call through instead of raising.
|
|
"""
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
user = LiteLLM_UserTable(user_id="u1", spend=0.0, max_budget=100.0)
|
|
team = LiteLLM_TeamTable(team_id="t1", spend=0.0, max_budget=1000.0)
|
|
token = UserAPIKeyAuth(token="k1", user_id="u1", team_id="t1")
|
|
|
|
async def _spend_by_counter(counter_key, fallback_spend, max_budget=None, **kwargs):
|
|
return 999.0 if counter_key == "spend:user:u1" else 0.0
|
|
|
|
async def _no_membership(*args, **kwargs):
|
|
return None
|
|
|
|
with patch("litellm.proxy.proxy_server.prisma_client", None), patch(
|
|
"litellm.proxy.proxy_server.get_current_spend", _spend_by_counter
|
|
), patch("litellm.proxy.auth.auth_checks.get_team_membership", _no_membership):
|
|
with pytest.raises(litellm.BudgetExceededError) as exc_info:
|
|
await common_checks(
|
|
request_body={"messages": [{"role": "user", "content": "hi"}]},
|
|
team_object=team,
|
|
user_object=user,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings={"apply_user_budget_to_team_keys": True},
|
|
route="/chat/completions",
|
|
llm_router=None,
|
|
proxy_logging_obj=MagicMock(),
|
|
valid_token=token,
|
|
request=MagicMock(spec=Request),
|
|
)
|
|
assert "ExceededBudget: User=u1" in str(exc_info.value)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_common_checks_personal_user_budget_still_enforced_on_personal_key_with_flag_enabled():
|
|
"""The flag only widens enforcement to team keys; personal keys keep blocking."""
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
user = LiteLLM_UserTable(user_id="u1", spend=0.0, max_budget=100.0)
|
|
token = UserAPIKeyAuth(token="k1", user_id="u1")
|
|
|
|
async def _spend_by_counter(counter_key, fallback_spend, max_budget=None, **kwargs):
|
|
return 999.0 if counter_key == "spend:user:u1" else 0.0
|
|
|
|
with patch("litellm.proxy.proxy_server.prisma_client", None), patch(
|
|
"litellm.proxy.proxy_server.get_current_spend", _spend_by_counter
|
|
):
|
|
with pytest.raises(litellm.BudgetExceededError):
|
|
await common_checks(
|
|
request_body={"messages": [{"role": "user", "content": "hi"}]},
|
|
team_object=None,
|
|
user_object=user,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings={"apply_user_budget_to_team_keys": True},
|
|
route="/chat/completions",
|
|
llm_router=None,
|
|
proxy_logging_obj=MagicMock(),
|
|
valid_token=token,
|
|
request=MagicMock(spec=Request),
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"scope, route, expect_blocked",
|
|
[
|
|
("user", "/chat/completions", True),
|
|
("user", "/key/list", False),
|
|
("team", "/chat/completions", True),
|
|
("team", "/key/list", False),
|
|
("org", "/chat/completions", True),
|
|
("org", "/key/list", False),
|
|
],
|
|
)
|
|
@pytest.mark.asyncio
|
|
async def test_budget_checks_only_run_on_llm_api_routes(scope, route, expect_blocked):
|
|
"""Budgets cap spend, so they must only gate routes that can spend.
|
|
|
|
Enforcing them on management routes locked an over-budget caller out of the
|
|
Admin UI, which authenticates with a normal virtual key, leaving no way to
|
|
reach the page that raises the limit.
|
|
"""
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
over_budget_counter = {"user": "spend:user:u1", "team": "spend:team:t1", "org": "spend:org:o1"}[scope]
|
|
|
|
async def _spend_by_counter(counter_key, fallback_spend, max_budget=None, **kwargs):
|
|
return 999.0 if counter_key == over_budget_counter else 0.0
|
|
|
|
async def _no_membership(*a, **kw):
|
|
return None
|
|
|
|
org_table = MagicMock()
|
|
org_table.spend = 999.0
|
|
org_table.litellm_budget_table = MagicMock()
|
|
org_table.litellm_budget_table.max_budget = 10.0
|
|
|
|
async def _get_org(*a, **kw):
|
|
return org_table
|
|
|
|
user = LiteLLM_UserTable(user_id="u1", spend=0.0, max_budget=10.0 if scope == "user" else None)
|
|
team = LiteLLM_TeamTable(team_id="t1", max_budget=10.0) if scope == "team" else None
|
|
token = UserAPIKeyAuth(
|
|
token="k1",
|
|
user_id="u1",
|
|
team_id="t1" if scope == "team" else None,
|
|
org_id="o1" if scope == "org" else None,
|
|
)
|
|
proxy_logging_obj = MagicMock()
|
|
proxy_logging_obj.budget_alerts = AsyncMock()
|
|
|
|
async def _run():
|
|
return await common_checks(
|
|
request_body={"messages": [{"role": "user", "content": "hi"}]},
|
|
team_object=team,
|
|
user_object=user,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings={},
|
|
route=route,
|
|
llm_router=None,
|
|
proxy_logging_obj=proxy_logging_obj,
|
|
valid_token=token,
|
|
request=MagicMock(spec=Request),
|
|
)
|
|
|
|
with patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), patch(
|
|
"litellm.proxy.proxy_server.get_current_spend", _spend_by_counter
|
|
), patch("litellm.proxy.auth.auth_checks.get_team_membership", _no_membership), patch(
|
|
"litellm.proxy.auth.auth_checks.get_org_object", _get_org
|
|
):
|
|
if expect_blocked:
|
|
with pytest.raises(litellm.BudgetExceededError):
|
|
await _run()
|
|
else:
|
|
assert await _run() is True
|
|
|
|
|
|
@pytest.mark.parametrize("route", ["/health", "/health/services", "/health/test_connection"])
|
|
@pytest.mark.asyncio
|
|
async def test_spend_capable_non_llm_routes_still_enforce_budget(route):
|
|
"""These routes are not LLM API routes but still reach a provider or an
|
|
external service: /health and /health/test_connection run litellm.ahealth_check
|
|
against real deployments, and /health/services fires Slack/email/webhook sends.
|
|
Exempting them with the other management routes would let an exhausted budget
|
|
keep spending.
|
|
"""
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
team = LiteLLM_TeamTable(team_id="t1", spend=150.0, max_budget=100.0)
|
|
|
|
with pytest.raises(litellm.BudgetExceededError):
|
|
await common_checks(
|
|
request_body={},
|
|
team_object=team,
|
|
user_object=None,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings={},
|
|
route=route,
|
|
llm_router=None,
|
|
proxy_logging_obj=AsyncMock(),
|
|
valid_token=UserAPIKeyAuth(token="k1", team_id="t1"),
|
|
request=MagicMock(spec=Request),
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_default_end_user_budget_db_fetch_returns_validated_budget(monkeypatch):
|
|
from litellm.proxy.auth.auth_checks import get_default_end_user_budget
|
|
|
|
monkeypatch.setattr(litellm, "max_end_user_budget_id", "budget-default-1")
|
|
|
|
budget_row = MagicMock()
|
|
budget_row.dict = lambda: {"budget_id": "budget-default-1", "max_budget": 12.5, "tpm_limit": 100}
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row)
|
|
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
result = await get_default_end_user_budget(
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
)
|
|
|
|
assert isinstance(result, LiteLLM_BudgetTable)
|
|
assert result.max_budget == 12.5
|
|
assert result.tpm_limit == 100
|
|
mock_cache.async_set_cache.assert_awaited_once()
|
|
assert mock_cache.async_set_cache.call_args.kwargs["value"] is result
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_team_member_default_budget_caches_json_safe_payload():
|
|
"""The Redis layer json.dumps() the cached value, so datetime columns on the budget row
|
|
must be dumped to ISO strings before the write, and the read side must give back a model.
|
|
"""
|
|
from litellm.proxy.auth.auth_checks import get_team_member_default_budget
|
|
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
|
|
|
budget_row = MagicMock()
|
|
budget_row.dict = lambda: {
|
|
"budget_id": "tm-budget-1",
|
|
"max_budget": 25.0,
|
|
"created_at": datetime(2026, 1, 1, tzinfo=timezone.utc),
|
|
"updated_at": datetime(2026, 1, 2, tzinfo=timezone.utc),
|
|
}
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row)
|
|
|
|
class _JsonOnlyRedis:
|
|
"""Stands in for RedisCache, which serializes with a bare json.dumps()."""
|
|
|
|
def __init__(self):
|
|
self.writes = []
|
|
|
|
async def async_set_cache(self, key, value, **kwargs):
|
|
self.writes.append((key, json.dumps(value)))
|
|
|
|
async def async_get_cache(self, key, **kwargs):
|
|
return None
|
|
|
|
redis_cache = _JsonOnlyRedis()
|
|
cache = UserApiKeyCache(redis_cache=redis_cache)
|
|
|
|
budget = await get_team_member_default_budget(
|
|
budget_id="tm-budget-1",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=cache,
|
|
)
|
|
|
|
assert isinstance(budget, LiteLLM_BudgetTable)
|
|
assert budget.max_budget == 25.0
|
|
assert len(redis_cache.writes) == 1
|
|
written_key, written_payload = redis_cache.writes[0]
|
|
assert written_key == "team_member_default_budget:tm-budget-1"
|
|
assert json.loads(written_payload)["max_budget"] == 25.0
|
|
|
|
cached = await get_team_member_default_budget(
|
|
budget_id="tm-budget-1",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=cache,
|
|
)
|
|
|
|
assert isinstance(cached, LiteLLM_BudgetTable)
|
|
assert cached.max_budget == 25.0
|
|
mock_prisma_client.db.litellm_budgettable.find_unique.assert_awaited_once()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_end_user_object_db_fetch_returns_validated_end_user():
|
|
from litellm.proxy.auth.auth_checks import get_end_user_object
|
|
|
|
end_user_row = MagicMock()
|
|
end_user_row.dict = lambda: {"user_id": "eu-1", "blocked": False, "spend": 3.0}
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_endusertable.find_unique = AsyncMock(return_value=end_user_row)
|
|
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
result = await get_end_user_object(
|
|
end_user_id="eu-1",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
)
|
|
|
|
assert isinstance(result, LiteLLM_EndUserTable)
|
|
assert result.user_id == "eu-1"
|
|
assert result.blocked is False
|
|
assert result.spend == 3.0
|
|
|
|
|
|
def _end_user_registry_row(user_id: str):
|
|
"""A row as the restricted-id registry query sees it: only ``user_id`` is read off it."""
|
|
return SimpleNamespace(user_id=user_id)
|
|
|
|
|
|
def _end_user_db_row(user_id: str, **fields):
|
|
row = MagicMock()
|
|
row.user_id = user_id
|
|
row.dict = lambda: {"user_id": user_id, "blocked": False, "spend": 0.0, **fields}
|
|
return row
|
|
|
|
|
|
_RESTRICTED_END_USER_WHERE = {
|
|
"OR": [
|
|
{"blocked": True},
|
|
{"budget_id": {"not": None}},
|
|
{"allowed_model_region": {"not": None}},
|
|
{"default_model": {"not": None}},
|
|
{"object_permission_id": {"not": None}},
|
|
]
|
|
}
|
|
|
|
|
|
@pytest.fixture
|
|
def end_user_registry_skip_enabled(monkeypatch):
|
|
"""Both bypass gates off: the default deployment, and the only state the registry skip runs in."""
|
|
monkeypatch.setattr(litellm, "max_end_user_budget_id", None)
|
|
monkeypatch.setattr(litellm, "validate_end_user_id_in_db", False)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_end_user_object_never_queries_db_for_unrestricted_end_users(
|
|
end_user_registry_skip_enabled,
|
|
):
|
|
"""
|
|
Regression: an end user carrying no restriction must not cost a DB read per request.
|
|
|
|
Spend tracking auto-creates a row for every distinct caller-supplied ``user`` id with every
|
|
restriction field null, so a high-cardinality deployment misses the per-pod cache on virtually
|
|
every request. Before the cached registry each miss ran its own Postgres find_unique, twice per
|
|
request, and under Prisma pool contention those queued for minutes inside user_api_key_auth.
|
|
"""
|
|
from litellm.proxy.auth.auth_checks import get_end_user_object
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[_end_user_registry_row("eu-blocked")])
|
|
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-anon-1"))
|
|
cache = UserApiKeyCache()
|
|
|
|
assert (
|
|
await get_end_user_object(
|
|
end_user_id="eu-anon-1",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
is None
|
|
)
|
|
mock_prisma.db.litellm_endusertable.find_unique.assert_not_awaited()
|
|
|
|
mock_prisma.db.litellm_endusertable.find_many.assert_awaited_once()
|
|
registry_call = mock_prisma.db.litellm_endusertable.find_many.call_args
|
|
assert registry_call.kwargs["take"] == END_USER_RESTRICTED_REGISTRY_MAX_SIZE + 1
|
|
# Every field the callers of get_end_user_object consume has to be in this predicate, or an id
|
|
# the registry calls unrestricted would silently lose a restriction that is actually enforced.
|
|
assert registry_call.kwargs["where"] == _RESTRICTED_END_USER_WHERE
|
|
|
|
mock_prisma.db.litellm_endusertable.find_many.reset_mock()
|
|
assert (
|
|
await get_end_user_object(
|
|
end_user_id="eu-anon-2",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
is None
|
|
)
|
|
# A second, different unknown id inside the TTL costs nothing: no rebuild, no row fetch.
|
|
mock_prisma.db.litellm_endusertable.find_many.assert_not_awaited()
|
|
mock_prisma.db.litellm_endusertable.find_unique.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_end_user_object_still_fetches_restricted_end_user(end_user_registry_skip_enabled):
|
|
"""An id in the registry keeps today's path: fetched, TTL-bounded in cache, then served cached."""
|
|
from litellm.proxy.auth.auth_checks import get_end_user_object
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[_end_user_registry_row("eu-blocked")])
|
|
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(
|
|
return_value=_end_user_db_row("eu-blocked", blocked=True)
|
|
)
|
|
cache = _TtlRecordingCache()
|
|
|
|
blocked = await get_end_user_object(
|
|
end_user_id="eu-blocked",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert isinstance(blocked, LiteLLM_EndUserTable)
|
|
assert blocked.blocked is True
|
|
mock_prisma.db.litellm_endusertable.find_unique.assert_awaited_once()
|
|
# Without a ttl the Redis entry never expires, so a later unblock would never be picked up.
|
|
assert (end_user_cache_key("eu-blocked"), DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL) in cache.writes
|
|
|
|
mock_prisma.db.litellm_endusertable.find_unique.reset_mock()
|
|
again = await get_end_user_object(
|
|
end_user_id="eu-blocked",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert again is not None and again.blocked is True
|
|
mock_prisma.db.litellm_endusertable.find_unique.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_end_user_object_caches_empty_restricted_registry(end_user_registry_skip_enabled):
|
|
"""No restricted end users at all is a valid answer and must be cached, not re-queried."""
|
|
from litellm.proxy.auth.auth_checks import get_end_user_object
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[])
|
|
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-anon-1"))
|
|
cache = UserApiKeyCache()
|
|
|
|
assert (
|
|
await get_end_user_object(
|
|
end_user_id="eu-anon-1",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
is None
|
|
)
|
|
# "Nobody is restricted" is a cached answer, not a cache miss (which would read back as None).
|
|
cached_registry = await cache.async_get_cache(key=end_user_restricted_registry_cache_key())
|
|
assert cached_registry is not None
|
|
assert tuple(cached_registry) == ()
|
|
|
|
assert (
|
|
await get_end_user_object(
|
|
end_user_id="eu-anon-2",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
is None
|
|
)
|
|
mock_prisma.db.litellm_endusertable.find_many.assert_awaited_once()
|
|
mock_prisma.db.litellm_endusertable.find_unique.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_end_user_object_registry_db_error_negative_caches_and_keeps_per_id_fetch(
|
|
end_user_registry_skip_enabled,
|
|
):
|
|
"""
|
|
A degraded database must not be re-asked for the registry on every request.
|
|
|
|
Restrictions keep being enforced through the per-id fetch, exactly as before the registry
|
|
existed, but the failing scan is suppressed for the negative-cache window instead of running
|
|
again on every request on top of that fetch. It is retried once the window closes.
|
|
"""
|
|
from litellm.proxy.auth.auth_checks import get_end_user_object
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(side_effect=Exception("registry query failed"))
|
|
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(
|
|
side_effect=lambda **kwargs: _end_user_db_row(kwargs["where"]["user_id"], blocked=True)
|
|
)
|
|
cache = _TtlRecordingCache()
|
|
|
|
first = await get_end_user_object(
|
|
end_user_id="eu-blocked-1",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert first is not None and first.blocked is True
|
|
assert (
|
|
await cache.async_get_cache(key=end_user_restricted_registry_cache_key())
|
|
== END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL
|
|
)
|
|
assert (end_user_restricted_registry_cache_key(), REGISTRY_ERROR_NEGATIVE_CACHE_TTL) in cache.writes
|
|
|
|
second = await get_end_user_object(
|
|
end_user_id="eu-blocked-2",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert second is not None and second.blocked is True
|
|
mock_prisma.db.litellm_endusertable.find_many.assert_awaited_once()
|
|
|
|
# The window closing (here: the entry expiring) puts the registry back in play.
|
|
await cache.async_delete_cache(key=end_user_restricted_registry_cache_key())
|
|
third = await get_end_user_object(
|
|
end_user_id="eu-blocked-3",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert third is not None and third.blocked is True
|
|
assert mock_prisma.db.litellm_endusertable.find_many.await_count == 2
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_registry_db_error_is_logged_at_warning(end_user_registry_skip_enabled):
|
|
"""
|
|
A registry that stops loading is a silent enforcement degradation, so seeing it must not
|
|
require debug logging: per-id lookups still enforce restrictions, but an operator has no other
|
|
signal that the database is failing the scan and that every request is paying for it.
|
|
"""
|
|
from litellm.proxy.auth.auth_checks import get_end_user_object
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(side_effect=Exception("registry query failed"))
|
|
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-1", blocked=True))
|
|
|
|
with patch("litellm.proxy.auth.auth_checks.verbose_proxy_logger") as mock_logger:
|
|
await get_end_user_object(
|
|
end_user_id="eu-1",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=UserApiKeyCache(),
|
|
)
|
|
|
|
warnings = [_rendered_log_message(call) for call in mock_logger.warning.call_args_list]
|
|
assert any(
|
|
end_user_restricted_registry_cache_key() in message and "registry query failed" in message
|
|
for message in warnings
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_end_user_registry_load_is_single_flighted_across_concurrent_requests(
|
|
end_user_registry_skip_enabled,
|
|
):
|
|
"""
|
|
A cold registry under load must run one scan, not one per in-flight request.
|
|
|
|
The registry query is an unindexed scan over the end-user table, which for the deployments this
|
|
exists for holds hundreds of thousands of rows; a TTL expiry on a busy worker would otherwise
|
|
fan it out across every concurrent request.
|
|
"""
|
|
from litellm.proxy.auth.auth_checks import get_end_user_object
|
|
|
|
async def fake_find_many(**kwargs):
|
|
await asyncio.sleep(0)
|
|
return [_end_user_registry_row("eu-blocked")]
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(side_effect=fake_find_many)
|
|
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-anon-1"))
|
|
cache = UserApiKeyCache()
|
|
|
|
results = await asyncio.gather(
|
|
*(
|
|
get_end_user_object(
|
|
end_user_id="eu-anon-1",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
for _ in range(8)
|
|
)
|
|
)
|
|
|
|
assert all(result is None for result in results)
|
|
mock_prisma.db.litellm_endusertable.find_many.assert_awaited_once()
|
|
mock_prisma.db.litellm_endusertable.find_unique.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_end_user_object_oversized_registry_falls_back_and_stops_refetching(
|
|
end_user_registry_skip_enabled,
|
|
):
|
|
"""Past the cap the registry is unusable: keep the per-id path, but stop rebuilding the set."""
|
|
from litellm.proxy.auth.auth_checks import get_end_user_object
|
|
|
|
oversized = [_end_user_registry_row(f"eu-{index}") for index in range(END_USER_RESTRICTED_REGISTRY_MAX_SIZE + 1)]
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=oversized)
|
|
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(
|
|
side_effect=lambda **kwargs: _end_user_db_row(kwargs["where"]["user_id"], blocked=True)
|
|
)
|
|
cache = UserApiKeyCache()
|
|
|
|
first = await get_end_user_object(
|
|
end_user_id="eu-anon-1",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert first is not None and first.blocked is True
|
|
assert (
|
|
await cache.async_get_cache(key=end_user_restricted_registry_cache_key())
|
|
== END_USER_RESTRICTED_REGISTRY_OVERFLOW_SENTINEL
|
|
)
|
|
|
|
second = await get_end_user_object(
|
|
end_user_id="eu-anon-2",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
)
|
|
assert second is not None and second.blocked is True
|
|
mock_prisma.db.litellm_endusertable.find_many.assert_awaited_once()
|
|
assert mock_prisma.db.litellm_endusertable.find_unique.await_count == 2
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_end_user_object_default_budget_gate_keeps_fetching_unrestricted_end_users(monkeypatch):
|
|
"""
|
|
With ``max_end_user_budget_id`` set, an existing unrestricted row is not equivalent to a missing
|
|
one: the default budget is grafted onto whatever row exists and is then enforced, so the skip
|
|
has to stay off entirely.
|
|
"""
|
|
from litellm.proxy.auth.auth_checks import get_end_user_object
|
|
|
|
monkeypatch.setattr(litellm, "max_end_user_budget_id", "default-eu-budget")
|
|
monkeypatch.setattr(litellm, "validate_end_user_id_in_db", False)
|
|
|
|
budget_row = MagicMock()
|
|
budget_row.dict = lambda: {"budget_id": "default-eu-budget", "max_budget": 25.0}
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[])
|
|
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-anon-1"))
|
|
mock_prisma.db.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row)
|
|
|
|
result = await get_end_user_object(
|
|
end_user_id="eu-anon-1",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=UserApiKeyCache(),
|
|
)
|
|
|
|
assert result is not None
|
|
assert result.litellm_budget_table is not None
|
|
assert result.litellm_budget_table.max_budget == 25.0
|
|
mock_prisma.db.litellm_endusertable.find_many.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_end_user_object_token_budget_gate_keeps_fetching_unrestricted_end_users(
|
|
end_user_registry_skip_enabled,
|
|
):
|
|
"""
|
|
A token-supplied end-user budget is enforced against the row's recorded spend, so the row has
|
|
to be loaded even though nothing on it is restricted.
|
|
|
|
A ``user_custom_auth`` callable can set ``end_user_max_budget`` on the returned token for an
|
|
end user whose row carries no budget of its own, which keeps it out of the registry.
|
|
"""
|
|
from litellm.proxy.auth.auth_checks import get_end_user_object
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[])
|
|
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(
|
|
return_value=_end_user_db_row("eu-anon-1", spend=100.0)
|
|
)
|
|
cache = UserApiKeyCache()
|
|
|
|
result = await get_end_user_object(
|
|
end_user_id="eu-anon-1",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=cache,
|
|
token_end_user_max_budget=50.0,
|
|
)
|
|
|
|
assert result is not None
|
|
assert result.spend == 100.0
|
|
mock_prisma.db.litellm_endusertable.find_unique.assert_awaited_once()
|
|
mock_prisma.db.litellm_endusertable.find_many.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_end_user_id_validation_gate_still_resolves_unrestricted_end_users(monkeypatch):
|
|
"""
|
|
With ``validate_end_user_id_in_db`` on, existence itself is the answer, so the skip stays off.
|
|
|
|
Skipping here would turn every unrestricted customer into an unknown id and drop it from the
|
|
request, which for a deployment with no default budget means the id silently stops being tracked.
|
|
"""
|
|
from litellm.proxy.auth.auth_checks import resolve_and_validate_end_user_id
|
|
|
|
monkeypatch.setattr(litellm, "max_end_user_budget_id", None)
|
|
monkeypatch.setattr(litellm, "validate_end_user_id_in_db", True)
|
|
|
|
mock_prisma = MagicMock()
|
|
mock_prisma.db.litellm_endusertable.find_many = AsyncMock(return_value=[])
|
|
mock_prisma.db.litellm_endusertable.find_unique = AsyncMock(return_value=_end_user_db_row("eu-known-1"))
|
|
mock_prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None)
|
|
mock_prisma.db.litellm_usertable.find_first = AsyncMock(return_value=None)
|
|
|
|
resolved = await resolve_and_validate_end_user_id(
|
|
raw_end_user_id="eu-known-1",
|
|
prisma_client=mock_prisma,
|
|
user_api_key_cache=UserApiKeyCache(),
|
|
)
|
|
|
|
assert resolved == "eu-known-1"
|
|
mock_prisma.db.litellm_endusertable.find_many.assert_not_awaited()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_team_membership_db_fetch_returns_validated_membership():
|
|
from litellm.proxy._types import LiteLLM_TeamMembership
|
|
from litellm.proxy.auth.auth_checks import get_team_membership
|
|
|
|
membership_row = MagicMock()
|
|
membership_row.dict = lambda: {"user_id": "u-1", "team_id": "t-1", "spend": 1.5}
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_teammembership.find_unique = AsyncMock(return_value=membership_row)
|
|
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
result = await get_team_membership(
|
|
user_id="u-1",
|
|
team_id="t-1",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
)
|
|
|
|
assert isinstance(result, LiteLLM_TeamMembership)
|
|
assert result.user_id == "u-1"
|
|
assert result.team_id == "t-1"
|
|
assert result.spend == 1.5
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_access_object_db_fetch_returns_validated_access_group():
|
|
from litellm.proxy._types import LiteLLM_AccessGroupTable
|
|
from litellm.proxy.auth.auth_checks import get_access_object
|
|
|
|
access_row = MagicMock()
|
|
access_row.dict = lambda: {
|
|
"access_group_id": "ag-1",
|
|
"access_group_name": "group one",
|
|
"access_model_names": ["gpt-4"],
|
|
}
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_accessgrouptable.find_unique = AsyncMock(return_value=access_row)
|
|
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
result = await get_access_object(
|
|
access_group_id="ag-1",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
proxy_logging_obj=None,
|
|
)
|
|
|
|
assert isinstance(result, LiteLLM_AccessGroupTable)
|
|
assert result.access_group_id == "ag-1"
|
|
assert result.access_model_names == ["gpt-4"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_team_object_by_alias_db_fetch_returns_cached_obj():
|
|
from litellm.proxy._types import LiteLLM_TeamTableCachedObj
|
|
from litellm.proxy.auth.auth_checks import get_team_object_by_alias
|
|
|
|
team_row = MagicMock()
|
|
team_row.model_dump = lambda: {"team_id": "t-9", "team_alias": "alias-9", "models": ["gpt-4"]}
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_teamtable.find_many = AsyncMock(return_value=[team_row])
|
|
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
result = await get_team_object_by_alias(
|
|
team_alias="alias-9",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
)
|
|
|
|
assert isinstance(result, LiteLLM_TeamTableCachedObj)
|
|
assert result.team_id == "t-9"
|
|
assert result.team_alias == "alias-9"
|
|
assert result.models == ["gpt-4"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_org_object_by_alias_db_fetch_returns_validated_org():
|
|
from litellm.proxy._types import LiteLLM_OrganizationTable
|
|
from litellm.proxy.auth.auth_checks import get_org_object_by_alias
|
|
|
|
org_row = MagicMock()
|
|
org_row.model_dump = lambda: {
|
|
"organization_id": "org-1",
|
|
"organization_alias": "org-alias",
|
|
"budget_id": "b-1",
|
|
"created_by": "admin",
|
|
"updated_by": "admin",
|
|
"models": [],
|
|
}
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_organizationtable.find_many = AsyncMock(return_value=[org_row])
|
|
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
result = await get_org_object_by_alias(
|
|
org_alias="org-alias",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
)
|
|
|
|
assert isinstance(result, LiteLLM_OrganizationTable)
|
|
assert result.organization_id == "org-1"
|
|
assert result.budget_id == "b-1"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_object_permission_db_fetch_returns_validated_permission():
|
|
from litellm.proxy.auth.auth_checks import get_object_permission
|
|
|
|
perm_row = MagicMock()
|
|
perm_row.dict = lambda: {"object_permission_id": "op-1", "vector_stores": ["vs-1"]}
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock(return_value=perm_row)
|
|
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
result = await get_object_permission(
|
|
object_permission_id="op-1",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
)
|
|
|
|
assert isinstance(result, LiteLLM_ObjectPermissionTable)
|
|
assert result.object_permission_id == "op-1"
|
|
assert result.vector_stores == ["vs-1"]
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_managed_vector_store_rows_by_uuids_db_fetch_validates_rows():
|
|
from litellm.proxy._types import LiteLLM_ManagedVectorStoresTable
|
|
from litellm.proxy.auth.auth_checks import get_managed_vector_store_rows_by_uuids
|
|
|
|
vs_row = MagicMock()
|
|
vs_row.model_dump = lambda: {"vector_store_id": "vs-7", "custom_llm_provider": "openai"}
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_managedvectorstorestable.find_many = AsyncMock(return_value=[vs_row])
|
|
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
result = await get_managed_vector_store_rows_by_uuids(
|
|
uuids=["vs-7"],
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
)
|
|
|
|
assert len(result) == 1
|
|
assert isinstance(result[0], LiteLLM_ManagedVectorStoresTable)
|
|
assert result[0].vector_store_id == "vs-7"
|
|
assert result[0].custom_llm_provider == "openai"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_get_project_object_db_fetch_returns_cached_obj():
|
|
from litellm.proxy._types import LiteLLM_ProjectTableCachedObj
|
|
from litellm.proxy.auth.auth_checks import get_project_object
|
|
|
|
project_row = MagicMock()
|
|
project_row.model_dump = lambda: {"project_id": "p-1", "project_alias": "proj", "team_id": "t-1"}
|
|
|
|
mock_prisma_client = MagicMock()
|
|
mock_prisma_client.db.litellm_projecttable.find_unique = AsyncMock(return_value=project_row)
|
|
|
|
mock_cache = MagicMock()
|
|
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
|
mock_cache.async_set_cache = AsyncMock()
|
|
|
|
result = await get_project_object(
|
|
project_id="p-1",
|
|
prisma_client=mock_prisma_client,
|
|
user_api_key_cache=mock_cache,
|
|
)
|
|
|
|
assert isinstance(result, LiteLLM_ProjectTableCachedObj)
|
|
assert result.project_id == "p-1"
|
|
assert result.project_alias == "proj"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_project_allowlist_enforced_when_key_models_empty():
|
|
"""
|
|
LIT-3803: a project-bound key with models=[] has no key-level restriction,
|
|
but the project allowlist must still 403 team models outside it.
|
|
"""
|
|
from litellm.proxy._types import (
|
|
LiteLLM_ProjectTableCachedObj,
|
|
ProxyErrorTypes,
|
|
ProxyException,
|
|
)
|
|
from litellm.proxy.auth.auth_checks import _run_project_checks, can_key_call_model
|
|
|
|
valid_token = UserAPIKeyAuth(
|
|
api_key="hashed-key",
|
|
project_id="p-1",
|
|
team_id="t-1",
|
|
models=[],
|
|
)
|
|
project = LiteLLM_ProjectTableCachedObj(
|
|
project_id="p-1",
|
|
team_id="t-1",
|
|
models=["gemini-2.5-flash-image", "gemini-3.1-flash-lite-preview"],
|
|
)
|
|
|
|
assert (
|
|
await can_key_call_model(
|
|
model="gemini-2.5-flash",
|
|
llm_model_list=None,
|
|
valid_token=valid_token,
|
|
llm_router=None,
|
|
)
|
|
is True
|
|
)
|
|
|
|
await _run_project_checks(
|
|
project_object=project,
|
|
_model="gemini-2.5-flash-image",
|
|
llm_router=None,
|
|
skip_budget_checks=True,
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=MagicMock(),
|
|
)
|
|
|
|
with pytest.raises(ProxyException) as exc_info:
|
|
await _run_project_checks(
|
|
project_object=project,
|
|
_model="gemini-2.5-flash",
|
|
llm_router=None,
|
|
skip_budget_checks=True,
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=MagicMock(),
|
|
)
|
|
assert exc_info.value.type == ProxyErrorTypes.project_model_access_denied
|
|
assert exc_info.value.code == "403"
|
|
|
|
|
|
def test_is_user_proxy_admin_rejects_view_only_admin():
|
|
"""This predicate skips `non_proxy_admin_allowed_routes_check` entirely, so an
|
|
Admin Viewer answering True here would gain every write route. Read parity for
|
|
that role belongs in the route checks, never here."""
|
|
from litellm.proxy.auth.auth_checks import _is_user_proxy_admin
|
|
|
|
viewer = LiteLLM_UserTable(
|
|
user_id="viewer_user",
|
|
user_email="viewer@example.com",
|
|
user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY.value,
|
|
)
|
|
admin = LiteLLM_UserTable(
|
|
user_id="admin_user",
|
|
user_email="admin@example.com",
|
|
user_role=LitellmUserRoles.PROXY_ADMIN.value,
|
|
)
|
|
|
|
assert _is_user_proxy_admin(user_obj=viewer) is False
|
|
assert _is_user_proxy_admin(user_obj=admin) is True
|
|
assert _is_user_proxy_admin(user_obj=None) is False
|
|
|
|
|
|
def _make_wildcard_access_group_router():
|
|
"""
|
|
`openai/*` tagged into an access group, plus an untagged `azure/*`, mirroring a
|
|
proxy that fronts a whole provider behind one wildcard deployment.
|
|
"""
|
|
from litellm import Router
|
|
|
|
return Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "openai/*",
|
|
"litellm_params": {"model": "openai/*", "api_key": "fake"},
|
|
"model_info": {
|
|
"id": "wildcard-openai",
|
|
"access_groups": ["default-models"],
|
|
},
|
|
},
|
|
{
|
|
"model_name": "azure/*",
|
|
"litellm_params": {"model": "azure/*", "api_key": "fake"},
|
|
"model_info": {"id": "wildcard-azure"},
|
|
},
|
|
]
|
|
)
|
|
|
|
|
|
def test_can_object_call_model_access_group_wildcard_accepts_bare_model_name():
|
|
"""
|
|
Regression: a key holding only the access group name was denied for `gpt-4o`
|
|
while `openai/gpt-4o` was allowed, because group membership resolved through the
|
|
pattern router's raw regex and skipped the `{provider}/{model}` retry that both
|
|
routing and the direct-wildcard grant already perform.
|
|
"""
|
|
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
|
|
|
router = _make_wildcard_access_group_router()
|
|
|
|
assert (
|
|
_can_object_call_model(
|
|
model="gpt-4o",
|
|
llm_router=router,
|
|
models=["default-models"],
|
|
object_type="key",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
def test_can_object_call_model_access_group_wildcard_accepts_prefixed_model_name():
|
|
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
|
|
|
router = _make_wildcard_access_group_router()
|
|
|
|
assert (
|
|
_can_object_call_model(
|
|
model="openai/gpt-4o",
|
|
llm_router=router,
|
|
models=["default-models"],
|
|
object_type="key",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"model",
|
|
[
|
|
"totally-made-up-model-zzz", # no provider can be inferred
|
|
"azure/some-deployment", # wildcard exists but carries no access group
|
|
],
|
|
)
|
|
def test_can_object_call_model_access_group_wildcard_does_not_over_grant(model):
|
|
"""The bare-name retry must not turn an access group into a blanket grant."""
|
|
from litellm.proxy._types import ProxyException
|
|
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
|
|
|
router = _make_wildcard_access_group_router()
|
|
|
|
with pytest.raises(ProxyException):
|
|
_can_object_call_model(
|
|
model=model,
|
|
llm_router=router,
|
|
models=["default-models"],
|
|
object_type="key",
|
|
)
|
|
|
|
|
|
def test_can_object_call_model_access_group_rejects_unconsumed_namespace():
|
|
"""
|
|
`bedrockz/...` infers provider `bedrock` from a fragment of the name, so
|
|
re-prefixing would smuggle an unrecognized namespace through a `bedrock/*` group.
|
|
"""
|
|
from litellm import Router
|
|
from litellm.proxy._types import ProxyException
|
|
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
|
|
|
router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "bedrock/*",
|
|
"litellm_params": {"model": "bedrock/*"},
|
|
"model_info": {
|
|
"id": "wildcard-bedrock",
|
|
"access_groups": ["bedrock-models"],
|
|
},
|
|
}
|
|
]
|
|
)
|
|
|
|
assert (
|
|
_can_object_call_model(
|
|
model="anthropic.claude-3-5-sonnet-20240620-v1:0",
|
|
llm_router=router,
|
|
models=["bedrock-models"],
|
|
object_type="key",
|
|
)
|
|
is True
|
|
)
|
|
|
|
with pytest.raises(ProxyException):
|
|
_can_object_call_model(
|
|
model="bedrockz/anthropic.claude-3-5-sonnet-20240620-v1:0",
|
|
llm_router=router,
|
|
models=["bedrock-models"],
|
|
object_type="key",
|
|
)
|
|
|
|
|
|
def test_can_object_call_model_team_scoped_wildcard_accepts_bare_model_name():
|
|
"""
|
|
Same regression as the proxy-wide wildcard, but for a team-scoped deployment
|
|
whose public name is a wildcard: those live in a separate per-team pattern
|
|
index that needed the same `{provider}/{model}` retry.
|
|
"""
|
|
from litellm import Router
|
|
from litellm.proxy.auth.auth_checks import _can_object_call_model
|
|
|
|
router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "openai/*_team-a_abc",
|
|
"litellm_params": {"model": "openai/*", "api_key": "fake"},
|
|
"model_info": {
|
|
"id": "team-byok-wildcard",
|
|
"team_id": "team-a",
|
|
"team_public_model_name": "openai/*",
|
|
"access_groups": ["team-models"],
|
|
},
|
|
}
|
|
]
|
|
)
|
|
|
|
for model in ("gpt-4o", "openai/gpt-4o"):
|
|
assert (
|
|
_can_object_call_model(
|
|
model=model,
|
|
llm_router=router,
|
|
models=["team-models"],
|
|
object_type="team",
|
|
team_id="team-a",
|
|
)
|
|
is True
|
|
)
|
|
|
|
|
|
UNPRICED_UNDERLYING_MODEL = "openai/unpriced-model-lit4984-xyz"
|
|
|
|
|
|
def _router_with_priced_and_unpriced_models() -> "Router":
|
|
from litellm.router import Router
|
|
|
|
return Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "priced-group",
|
|
"litellm_params": {"model": "gpt-3.5-turbo", "api_key": "sk-test"},
|
|
},
|
|
{
|
|
"model_name": "unpriced-group",
|
|
"litellm_params": {"model": UNPRICED_UNDERLYING_MODEL, "api_key": "sk-test"},
|
|
},
|
|
]
|
|
)
|
|
|
|
|
|
def test_model_has_no_cost_mapping_priced_model_is_false():
|
|
from litellm.proxy.auth.auth_checks import model_has_no_cost_mapping
|
|
|
|
router = _router_with_priced_and_unpriced_models()
|
|
|
|
assert model_has_no_cost_mapping(model="priced-group", llm_router=router) is False
|
|
|
|
|
|
def test_model_has_no_cost_mapping_unpriced_model_is_true():
|
|
from litellm.proxy.auth.auth_checks import model_has_no_cost_mapping
|
|
|
|
router = _router_with_priced_and_unpriced_models()
|
|
|
|
assert model_has_no_cost_mapping(model="unpriced-group", llm_router=router) is True
|
|
|
|
|
|
def test_model_has_no_cost_mapping_no_model_or_router_is_false():
|
|
from litellm.proxy.auth.auth_checks import model_has_no_cost_mapping
|
|
|
|
router = _router_with_priced_and_unpriced_models()
|
|
|
|
assert model_has_no_cost_mapping(model=None, llm_router=router) is False
|
|
assert model_has_no_cost_mapping(model="unpriced-group", llm_router=None) is False
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"underlying_model",
|
|
[
|
|
"azure/speech/azure-tts",
|
|
"mistral/mistral-ocr-latest",
|
|
"vertex_ai/imagen-3.0-generate-001",
|
|
"dashscope/qwen-flash",
|
|
],
|
|
)
|
|
def test_model_has_no_cost_mapping_non_token_priced_model_is_false(underlying_model):
|
|
from litellm.proxy.auth.auth_checks import model_has_no_cost_mapping
|
|
from litellm.router import Router
|
|
|
|
router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "non-token-priced-group",
|
|
"litellm_params": {"model": underlying_model, "api_key": "sk-test"},
|
|
}
|
|
]
|
|
)
|
|
|
|
assert model_has_no_cost_mapping(model="non-token-priced-group", llm_router=router) is False
|
|
|
|
|
|
def test_model_has_no_cost_mapping_non_token_price_from_litellm_params_is_false():
|
|
from litellm.proxy.auth.auth_checks import model_has_no_cost_mapping
|
|
from litellm.router import Router
|
|
|
|
router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "custom-tts",
|
|
"litellm_params": {
|
|
"model": f"{UNPRICED_UNDERLYING_MODEL}-per-second",
|
|
"api_key": "sk-test",
|
|
"input_cost_per_second": 0.0001,
|
|
},
|
|
}
|
|
]
|
|
)
|
|
|
|
assert model_has_no_cost_mapping(model="custom-tts", llm_router=router) is False
|
|
|
|
|
|
@pytest.mark.parametrize("cost_field", ["input_cost_per_second", "input_cost_per_token"])
|
|
def test_model_has_no_cost_mapping_explicit_zero_price_is_false(cost_field):
|
|
from litellm.proxy.auth.auth_checks import model_has_no_cost_mapping
|
|
from litellm.router import Router
|
|
|
|
router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "free-group",
|
|
"litellm_params": {
|
|
"model": f"{UNPRICED_UNDERLYING_MODEL}-{cost_field}",
|
|
"api_key": "sk-test",
|
|
cost_field: 0,
|
|
},
|
|
}
|
|
]
|
|
)
|
|
|
|
assert model_has_no_cost_mapping(model="free-group", llm_router=router) is False
|
|
|
|
|
|
def test_model_has_no_cost_mapping_tiered_pricing_only_is_false():
|
|
from litellm.proxy.auth.auth_checks import model_has_no_cost_mapping
|
|
from litellm.router import Router
|
|
|
|
router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "tiered-group",
|
|
"litellm_params": {
|
|
"model": f"{UNPRICED_UNDERLYING_MODEL}-tiered",
|
|
"api_key": "sk-test",
|
|
"tiered_pricing": [
|
|
{"range": [0, 128000], "input_cost_per_token": 2e-7, "output_cost_per_token": 6e-7},
|
|
{"range": [128000, 256000], "input_cost_per_token": 4e-7, "output_cost_per_token": 12e-7},
|
|
],
|
|
},
|
|
}
|
|
]
|
|
)
|
|
|
|
assert model_has_no_cost_mapping(model="tiered-group", llm_router=router) is False
|
|
|
|
|
|
async def _run_common_checks(
|
|
model: Optional[str], llm_router: Optional["Router"], route: str = "/chat/completions"
|
|
) -> bool:
|
|
from fastapi import Request
|
|
|
|
from litellm.proxy.auth.auth_checks import common_checks
|
|
|
|
return await common_checks(
|
|
request_body={"model": model, "messages": [{"role": "user", "content": "hi"}]},
|
|
team_object=None,
|
|
user_object=None,
|
|
end_user_object=None,
|
|
global_proxy_spend=None,
|
|
general_settings={},
|
|
route=route,
|
|
llm_router=llm_router,
|
|
proxy_logging_obj=MagicMock(),
|
|
valid_token=UserAPIKeyAuth(token="test-token"),
|
|
request=MagicMock(spec=Request),
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_common_checks_blocks_unpriced_model_when_enabled(monkeypatch):
|
|
monkeypatch.setattr(litellm, "block_requests_for_models_without_pricing", True)
|
|
router = _router_with_priced_and_unpriced_models()
|
|
|
|
with pytest.raises(ProxyException) as exc_info:
|
|
await _run_common_checks(model="unpriced-group", llm_router=router)
|
|
|
|
assert exc_info.value.code == "403"
|
|
assert exc_info.value.type == ProxyErrorTypes.model_cost_map_missing
|
|
assert exc_info.value.param == "model"
|
|
assert "unpriced-group" in exc_info.value.message
|
|
assert "pricing" in exc_info.value.message.lower()
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_common_checks_allows_unpriced_model_when_disabled(monkeypatch):
|
|
monkeypatch.setattr(litellm, "block_requests_for_models_without_pricing", False)
|
|
router = _router_with_priced_and_unpriced_models()
|
|
|
|
result = await _run_common_checks(model="unpriced-group", llm_router=router)
|
|
|
|
assert result is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_common_checks_allows_priced_model_when_enabled(monkeypatch):
|
|
monkeypatch.setattr(litellm, "block_requests_for_models_without_pricing", True)
|
|
router = _router_with_priced_and_unpriced_models()
|
|
|
|
result = await _run_common_checks(model="priced-group", llm_router=router)
|
|
|
|
assert result is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_common_checks_ignores_non_llm_route_when_enabled(monkeypatch):
|
|
monkeypatch.setattr(litellm, "block_requests_for_models_without_pricing", True)
|
|
router = _router_with_priced_and_unpriced_models()
|
|
|
|
result = await _run_common_checks(
|
|
model="unpriced-group", llm_router=router, route="/model/new"
|
|
)
|
|
|
|
assert result is True
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_common_checks_blocks_alias_resolving_to_unpriced_model(monkeypatch):
|
|
from litellm.router import Router
|
|
|
|
monkeypatch.setattr(litellm, "block_requests_for_models_without_pricing", True)
|
|
router = Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "billed-underlying-group",
|
|
"litellm_params": {"model": UNPRICED_UNDERLYING_MODEL, "api_key": "sk-test"},
|
|
}
|
|
],
|
|
model_group_alias={"public-alias": "billed-underlying-group"},
|
|
)
|
|
|
|
with pytest.raises(ProxyException) as exc_info:
|
|
await _run_common_checks(model="public-alias", llm_router=router)
|
|
|
|
assert exc_info.value.code == "403"
|
|
assert exc_info.value.type == ProxyErrorTypes.model_cost_map_missing
|
|
assert "public-alias" in exc_info.value.message
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_common_checks_blocks_comma_separated_request_carrying_an_unpriced_model(monkeypatch):
|
|
monkeypatch.setattr(litellm, "block_requests_for_models_without_pricing", True)
|
|
router = _router_with_priced_and_unpriced_models()
|
|
|
|
with pytest.raises(ProxyException) as exc_info:
|
|
await _run_common_checks(model="priced-group,unpriced-group", llm_router=router)
|
|
|
|
assert exc_info.value.code == "403"
|
|
assert exc_info.value.type == ProxyErrorTypes.model_cost_map_missing
|
|
assert "'unpriced-group'" in exc_info.value.message
|
|
assert "'priced-group'" not in exc_info.value.message
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_common_checks_allows_comma_separated_request_when_every_model_is_priced(monkeypatch):
|
|
monkeypatch.setattr(litellm, "block_requests_for_models_without_pricing", True)
|
|
router = _router_with_priced_and_unpriced_models()
|
|
|
|
result = await _run_common_checks(model="priced-group,priced-group", llm_router=router)
|
|
|
|
assert result is True
|
|
|
|
|
|
def _router_with_a_group_priced_through_model_info() -> "Router":
|
|
from litellm.router import Router
|
|
|
|
return Router(
|
|
model_list=[
|
|
{
|
|
"model_name": "model-info-priced-group",
|
|
"litellm_params": {"model": f"{UNPRICED_UNDERLYING_MODEL}-model-info", "api_key": "sk-test"},
|
|
"model_info": {"input_cost_per_token": 0, "output_cost_per_token": 0},
|
|
}
|
|
],
|
|
model_group_alias={"model-info-priced-alias": "model-info-priced-group"},
|
|
)
|
|
|
|
|
|
def test_model_has_no_cost_mapping_group_priced_through_model_info_is_false():
|
|
from litellm.proxy.auth.auth_checks import model_has_no_cost_mapping
|
|
|
|
router = _router_with_a_group_priced_through_model_info()
|
|
|
|
assert model_has_no_cost_mapping(model="model-info-priced-group", llm_router=router) is False
|
|
|
|
|
|
def test_model_has_no_cost_mapping_alias_to_a_group_priced_through_model_info_is_false():
|
|
from litellm.proxy.auth.auth_checks import model_has_no_cost_mapping
|
|
|
|
router = _router_with_a_group_priced_through_model_info()
|
|
|
|
assert model_has_no_cost_mapping(model="model-info-priced-alias", llm_router=router) is False
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"user_route, expected",
|
|
[
|
|
("/internal-models/v1/chat/completions", True),
|
|
("/internal-models/newly-registered-model/predict", True),
|
|
("/internal-models-other/v1/chat/completions", False),
|
|
("/anthropic/v1/messages", False),
|
|
],
|
|
)
|
|
def test_team_allowed_routes_wildcard_prefix_matches_unregistered_passthrough_routes(user_route, expected):
|
|
"""A `/prefix/*` entry in `team_allowed_routes` must cover every route under that prefix, so
|
|
passthrough endpoints registered after the proxy config was written are reachable without an
|
|
exact-route config change."""
|
|
from litellm.proxy._types import LiteLLM_JWTAuth
|
|
from litellm.proxy.auth.auth_checks import allowed_routes_check
|
|
|
|
assert (
|
|
allowed_routes_check(
|
|
user_role=LitellmUserRoles.TEAM,
|
|
user_route=user_route,
|
|
litellm_proxy_roles=LiteLLM_JWTAuth(team_allowed_routes=["/internal-models/*"]),
|
|
)
|
|
is expected
|
|
)
|
|
|
|
|
|
def test_team_allowed_routes_exact_route_does_not_become_a_prefix_grant():
|
|
from litellm.proxy._types import LiteLLM_JWTAuth
|
|
from litellm.proxy.auth.auth_checks import allowed_routes_check
|
|
|
|
roles = LiteLLM_JWTAuth(team_allowed_routes=["/internal-models/model-a"])
|
|
|
|
assert (
|
|
allowed_routes_check(user_role=LitellmUserRoles.TEAM, user_route="/internal-models/model-a", litellm_proxy_roles=roles)
|
|
is True
|
|
)
|
|
assert (
|
|
allowed_routes_check(user_role=LitellmUserRoles.TEAM, user_route="/internal-models/model-b", litellm_proxy_roles=roles)
|
|
is False
|
|
)
|
|
|
|
|
|
def test_admin_allowed_routes_wildcard_prefix_is_honored():
|
|
from litellm.proxy._types import LiteLLM_JWTAuth
|
|
from litellm.proxy.auth.auth_checks import allowed_routes_check
|
|
|
|
roles = LiteLLM_JWTAuth(admin_allowed_routes=["/internal-models/*"])
|
|
|
|
assert (
|
|
allowed_routes_check(
|
|
user_role=LitellmUserRoles.PROXY_ADMIN, user_route="/internal-models/anything", litellm_proxy_roles=roles
|
|
)
|
|
is True
|
|
)
|
|
assert (
|
|
allowed_routes_check(
|
|
user_role=LitellmUserRoles.PROXY_ADMIN, user_route="/other/anything", litellm_proxy_roles=roles
|
|
)
|
|
is False
|
|
)
|
|
|
|
|
|
def test_team_allowed_routes_named_route_group_still_resolves():
|
|
from litellm.proxy._types import LiteLLM_JWTAuth
|
|
from litellm.proxy.auth.auth_checks import allowed_routes_check
|
|
|
|
roles = LiteLLM_JWTAuth(team_allowed_routes=["openai_routes"])
|
|
|
|
assert (
|
|
allowed_routes_check(
|
|
user_role=LitellmUserRoles.TEAM, user_route="/v1/chat/completions", litellm_proxy_roles=roles
|
|
)
|
|
is True
|
|
)
|
|
assert (
|
|
allowed_routes_check(user_role=LitellmUserRoles.TEAM, user_route="/key/generate", litellm_proxy_roles=roles)
|
|
is False
|
|
)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_invalidate_team_member_spend_state_sets_the_spend_counter_and_clears_both_membership_cache_keys():
|
|
"""A team-member budget reset (new_spend passed) must SET the spend counter to the reset
|
|
value, clear its DB-floor marker, AND invalidate both independently-keyed membership caches
|
|
(user_api_key_auth.py's admission check writes one key format, budget_reservation.py and
|
|
auth_checks.py's own get_team_membership() write the other) or a stale read keeps 429ing
|
|
after the reset. Asserted against real cache reads, not mock call args, so a change that
|
|
keeps the call but drops its effect still fails."""
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
|
|
|
real_cache = UserApiKeyCache()
|
|
await real_cache.async_set_cache(key="team-1_user-1", value="stale-membership")
|
|
await real_cache.async_set_cache(key="team_membership:user-1:team-1", value="stale-membership")
|
|
|
|
real_spend_counter_cache = DualCache()
|
|
real_spend_counter_cache.in_memory_cache.set_cache(key="spend:team_member:user-1:team-1", value=999.0)
|
|
real_spend_counter_cache.in_memory_cache.set_cache(
|
|
key="spend_db_floor:spend:team_member:user-1:team-1", value=999.0
|
|
)
|
|
|
|
with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
|
|
"litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache
|
|
):
|
|
await invalidate_team_member_spend_state(
|
|
user_id="user-1",
|
|
team_id="team-1",
|
|
user_api_key_cache=real_cache,
|
|
new_spend=0.0,
|
|
)
|
|
|
|
assert await real_cache.async_get_cache(key="team-1_user-1") is None
|
|
assert await real_cache.async_get_cache(key="team_membership:user-1:team-1") is None
|
|
assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-1:team-1") == 0.0
|
|
assert (
|
|
real_spend_counter_cache.in_memory_cache.get_cache(key="spend_db_floor:spend:team_member:user-1:team-1")
|
|
== 0.0
|
|
), "the DB-floor marker kept the pre-reset value; a stale-floor read can raise the counter right back up"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_invalidate_team_member_spend_state_leaves_the_live_spend_counter_alone_without_new_spend():
|
|
"""team_member_update only changes the budget cap, not the tracked spend, so it calls
|
|
invalidate_team_member_spend_state with no new_spend. Deleting the live spend counter in that
|
|
case would force the next read to reseed from the DB's own spend column, which lags the live
|
|
counter via periodic batch writes, briefly UNDER-enforcing the raised cap against a spend
|
|
value lower than what was actually tracked (regression: PR #37971 Bugbot finding)."""
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
|
|
|
real_cache = UserApiKeyCache()
|
|
await real_cache.async_set_cache(key="team-1_user-1", value="stale-membership")
|
|
|
|
real_spend_counter_cache = DualCache()
|
|
real_spend_counter_cache.in_memory_cache.set_cache(key="spend:team_member:user-1:team-1", value=999.0)
|
|
|
|
with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
|
|
"litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache
|
|
):
|
|
await invalidate_team_member_spend_state(
|
|
user_id="user-1",
|
|
team_id="team-1",
|
|
user_api_key_cache=real_cache,
|
|
)
|
|
|
|
assert await real_cache.async_get_cache(key="team-1_user-1") is None
|
|
assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-1:team-1") == 999.0
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_invalidate_team_member_spend_state_sets_new_spend_instead_of_deleting():
|
|
"""/key/{key}/reset_spend SETs its counter to the reset value rather than deleting it, so a
|
|
worker's next read reflects it directly instead of falling back through a DB reseed. A reset
|
|
caller passing new_spend must match that precedent, not merely delete the counter."""
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
|
|
|
real_cache = UserApiKeyCache()
|
|
real_spend_counter_cache = DualCache()
|
|
fake_redis_cache = MagicMock()
|
|
fake_redis_cache.async_set_cache = AsyncMock()
|
|
real_spend_counter_cache.redis_cache = fake_redis_cache
|
|
|
|
with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
|
|
"litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache
|
|
):
|
|
await invalidate_team_member_spend_state(
|
|
user_id="user-1",
|
|
team_id="team-1",
|
|
user_api_key_cache=real_cache,
|
|
new_spend=2.5,
|
|
)
|
|
|
|
assert real_spend_counter_cache.in_memory_cache.get_cache(key="spend:team_member:user-1:team-1") == 2.5
|
|
fake_redis_cache.async_set_cache.assert_awaited_once_with(key="spend:team_member:user-1:team-1", value=2.5, ttl=60)
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_invalidate_team_member_spend_state_deletes_redis_counter_when_set_fails(): # test-quality-ok: only observable effect is the fallback call on the same fake client
|
|
"""Redis reads take priority over the local in-memory copy (get_current_spend reads Redis
|
|
first), so a failed Redis SET would otherwise leave the OLD pre-reset value authoritative
|
|
for every worker even though the reset reported success. On a failed SET, the stale Redis
|
|
entry must be deleted instead, so the next read clean-misses and reseeds from the DB."""
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
|
|
|
real_cache = UserApiKeyCache()
|
|
real_spend_counter_cache = DualCache()
|
|
fake_redis_cache = MagicMock()
|
|
fake_redis_cache.async_set_cache = AsyncMock(side_effect=ConnectionError("redis down"))
|
|
fake_redis_cache.async_delete_cache = AsyncMock()
|
|
real_spend_counter_cache.redis_cache = fake_redis_cache
|
|
|
|
with patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
|
|
"litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache
|
|
):
|
|
await invalidate_team_member_spend_state(
|
|
user_id="user-1",
|
|
team_id="team-1",
|
|
user_api_key_cache=real_cache,
|
|
new_spend=2.5,
|
|
)
|
|
|
|
fake_redis_cache.async_delete_cache.assert_awaited_once_with(key="spend:team_member:user-1:team-1")
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_invalidate_team_member_spend_state_raises_503_when_both_redis_writes_fail():
|
|
"""If the Redis SET fails AND the fallback DELETE fails, the stale pre-reset counter is still
|
|
authoritative in Redis for every worker. Reporting success would silently keep 429ing the
|
|
member, so the reset must surface a 503 instead (regression: PR #37971 Greptile finding)."""
|
|
from fastapi import HTTPException
|
|
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
|
|
|
real_cache = UserApiKeyCache()
|
|
real_spend_counter_cache = DualCache()
|
|
fake_redis_cache = MagicMock()
|
|
fake_redis_cache.async_set_cache = AsyncMock(side_effect=ConnectionError("redis down"))
|
|
fake_redis_cache.async_delete_cache = AsyncMock(side_effect=ConnectionError("redis still down"))
|
|
real_spend_counter_cache.redis_cache = fake_redis_cache
|
|
|
|
with (
|
|
patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
|
|
"litellm.proxy.proxy_server.spend_counter_cache", real_spend_counter_cache
|
|
),
|
|
pytest.raises(HTTPException) as exc_info,
|
|
):
|
|
await invalidate_team_member_spend_state(
|
|
user_id="user-1",
|
|
team_id="team-1",
|
|
user_api_key_cache=real_cache,
|
|
new_spend=2.5,
|
|
)
|
|
|
|
assert exc_info.value.status_code == 503
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_invalidate_team_member_spend_state_broadcasts_the_spend_counter_to_remote_workers():
|
|
"""The test above only proves the handling worker's own spend counter is
|
|
cleared. A remote worker's spend counter is a separate DualCache instance;
|
|
if the reset never reaches it, that worker keeps enforcing the pre-reset
|
|
spend the moment its own Redis read for the counter fails and it falls
|
|
back to its own (now-stale) in-memory copy. Drives the actual message
|
|
published onto the invalidation channel through a second, independent
|
|
AuthCacheInvalidationSubscriber standing in for that remote worker, rather
|
|
than asserting on the publish call args."""
|
|
from redis.asyncio import Redis
|
|
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.caching.in_memory_cache import InMemoryCache
|
|
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import AuthCacheInvalidationSubscriber
|
|
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
|
|
|
published: list[tuple[str, str]] = []
|
|
|
|
class _RecordingRedisClient(Redis):
|
|
def __init__(self) -> None:
|
|
pass
|
|
|
|
async def publish(self, channel: str, message: str) -> int:
|
|
published.append((channel, message))
|
|
return 1
|
|
|
|
class _FakeRedisCache:
|
|
def __init__(self) -> None:
|
|
self.namespace = None
|
|
|
|
def init_async_client(self) -> object:
|
|
return _RecordingRedisClient()
|
|
|
|
local_spend_counter_cache = DualCache()
|
|
|
|
remote_user_api_key_cache = UserApiKeyCache()
|
|
remote_spend_counter_in_memory_cache = InMemoryCache()
|
|
remote_spend_counter_in_memory_cache.set_cache("spend:team_member:user-1:team-1", 999.0)
|
|
remote_spend_counter_in_memory_cache.set_cache("spend_db_floor:spend:team_member:user-1:team-1", 999.0)
|
|
|
|
with (
|
|
patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
|
|
"litellm.proxy.proxy_server.spend_counter_cache", local_spend_counter_cache
|
|
),
|
|
patch( # test-quality-ok: injects a fake pub/sub-capable redis cache; no live redis in this unit test
|
|
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache",
|
|
return_value=_FakeRedisCache(),
|
|
),
|
|
):
|
|
await invalidate_team_member_spend_state(
|
|
user_id="user-1",
|
|
team_id="team-1",
|
|
user_api_key_cache=UserApiKeyCache(),
|
|
new_spend=0.0,
|
|
)
|
|
|
|
def _published_message_for(cache_key: str) -> str:
|
|
matches = [message for _, message in published if json.loads(message)["cache_key"] == cache_key]
|
|
assert matches, f"{cache_key} never reached the cross-worker invalidation channel"
|
|
return matches[-1]
|
|
|
|
remote_subscriber = AuthCacheInvalidationSubscriber(
|
|
redis_cache=_FakeRedisCache(),
|
|
user_api_key_cache=remote_user_api_key_cache,
|
|
additional_in_memory_caches=(remote_spend_counter_in_memory_cache,),
|
|
)
|
|
for cache_key in ("spend:team_member:user-1:team-1", "spend_db_floor:spend:team_member:user-1:team-1"):
|
|
remote_subscriber._apply_message( # pyright: ignore[reportPrivateUsage] # exercising the real cross-worker message handler, not a public API
|
|
{"type": "message", "data": _published_message_for(cache_key)}
|
|
)
|
|
|
|
assert remote_spend_counter_in_memory_cache.get_cache("spend:team_member:user-1:team-1") == 0.0
|
|
assert (
|
|
remote_spend_counter_in_memory_cache.get_cache("spend_db_floor:spend:team_member:user-1:team-1") == 0.0
|
|
), "the DB-floor marker was not broadcast; a remote worker can re-raise the counter off its stale floor"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_invalidate_team_member_spend_state_self_delivered_broadcast_does_not_erase_the_reset():
|
|
"""The handling worker subscribes to the same invalidation channel it publishes on, so it
|
|
receives its own reset message. A delete-style broadcast would erase the post-reset counter
|
|
and floor marker the handler just wrote, reopening the stale-floor race the reset closed
|
|
(regression: PR #37971 Greptile finding). The broadcast carries the reset value as a SET, so
|
|
applying the self-delivered message must leave both keys at the post-reset value."""
|
|
from redis.asyncio import Redis
|
|
|
|
from litellm.caching.dual_cache import DualCache
|
|
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import AuthCacheInvalidationSubscriber
|
|
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
|
|
|
published: list[tuple[str, str]] = []
|
|
|
|
class _RecordingRedisClient(Redis):
|
|
def __init__(self) -> None:
|
|
pass
|
|
|
|
async def publish(self, channel: str, message: str) -> int:
|
|
published.append((channel, message))
|
|
return 1
|
|
|
|
class _FakeRedisCache:
|
|
def __init__(self) -> None:
|
|
self.namespace = None
|
|
|
|
def init_async_client(self) -> object:
|
|
return _RecordingRedisClient()
|
|
|
|
local_spend_counter_cache = DualCache()
|
|
local_user_api_key_cache = UserApiKeyCache()
|
|
|
|
with (
|
|
patch( # test-quality-ok: injects a real DualCache for the module global, not a behavior mock
|
|
"litellm.proxy.proxy_server.spend_counter_cache", local_spend_counter_cache
|
|
),
|
|
patch( # test-quality-ok: injects a fake pub/sub-capable redis cache; no live redis in this unit test
|
|
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache",
|
|
return_value=_FakeRedisCache(),
|
|
),
|
|
):
|
|
await invalidate_team_member_spend_state(
|
|
user_id="user-1",
|
|
team_id="team-1",
|
|
user_api_key_cache=local_user_api_key_cache,
|
|
new_spend=0.0,
|
|
)
|
|
|
|
own_subscriber = AuthCacheInvalidationSubscriber(
|
|
redis_cache=_FakeRedisCache(),
|
|
user_api_key_cache=local_user_api_key_cache,
|
|
additional_in_memory_caches=(local_spend_counter_cache.in_memory_cache,),
|
|
)
|
|
for _, message in published:
|
|
own_subscriber._apply_message( # pyright: ignore[reportPrivateUsage] # exercising the real cross-worker message handler, not a public API
|
|
{"type": "message", "data": message}
|
|
)
|
|
|
|
assert local_spend_counter_cache.in_memory_cache.get_cache("spend:team_member:user-1:team-1") == 0.0, (
|
|
"the handler's self-delivered broadcast erased the post-reset spend counter"
|
|
)
|
|
assert (
|
|
local_spend_counter_cache.in_memory_cache.get_cache("spend_db_floor:spend:team_member:user-1:team-1") == 0.0
|
|
), "the handler's self-delivered broadcast erased the post-reset floor marker, reopening the stale-floor race"
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
async def test_delete_cache_key_object_is_best_effort_when_the_cache_backend_fails(caplog):
|
|
"""
|
|
LIT-5898: `_delete_cache_key_object` must not propagate a cache-backend error.
|
|
|
|
Every caller runs it after its own write has committed, so a raise here turned a persisted
|
|
`/key/update` into `400 Authentication Error` (and `/key/block`, `/key/regenerate` into 500s)
|
|
for operators whose Redis ACL denies `DEL` on LiteLLM's unprefixed token-hash keys. The
|
|
in-memory entry is already dropped by then, so raising never made the cache less stale.
|
|
"""
|
|
import logging
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
from litellm.proxy.auth.auth_checks import _delete_cache_key_object
|
|
|
|
hashed_token = "a" * 64
|
|
caplog.set_level(logging.WARNING, logger="LiteLLM Proxy")
|
|
|
|
failing_cache = MagicMock()
|
|
failing_cache.delete_cache = MagicMock()
|
|
failing_logging_obj = MagicMock()
|
|
failing_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock(
|
|
side_effect=Exception("No permissions to access a key")
|
|
)
|
|
|
|
await _delete_cache_key_object(
|
|
hashed_token=hashed_token,
|
|
user_api_key_cache=failing_cache,
|
|
proxy_logging_obj=failing_logging_obj,
|
|
)
|
|
|
|
failing_cache.delete_cache.assert_called_once_with(key=hashed_token)
|
|
failing_logging_obj.internal_usage_cache.dual_cache.async_delete_cache.assert_awaited_once_with(key=hashed_token)
|
|
assert any("Failed to invalidate cached key entry" in record.getMessage() for record in caplog.records), (
|
|
"a swallowed cache-eviction failure must still be logged, or a stale auth entry goes unnoticed"
|
|
)
|
|
|
|
caplog.clear()
|
|
healthy_cache = MagicMock()
|
|
healthy_cache.delete_cache = MagicMock()
|
|
healthy_logging_obj = MagicMock()
|
|
healthy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache = AsyncMock()
|
|
|
|
await _delete_cache_key_object(
|
|
hashed_token=hashed_token,
|
|
user_api_key_cache=healthy_cache,
|
|
proxy_logging_obj=healthy_logging_obj,
|
|
)
|
|
|
|
healthy_cache.delete_cache.assert_called_once_with(key=hashed_token)
|
|
healthy_logging_obj.internal_usage_cache.dual_cache.async_delete_cache.assert_awaited_once_with(key=hashed_token)
|
|
assert caplog.records == [], "a healthy eviction must stay silent, and must still reach both caches"
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# Budget-exceeded error text must not carry a raw virtual key (LIT-5909)
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class _BudgetAlertRecorder:
|
|
async def budget_alerts(self, type, user_info):
|
|
return None
|
|
|
|
|
|
async def _run_key_budget_check(key_name: str) -> str:
|
|
"""Drive the over-budget key path and return the raised message."""
|
|
valid_token = UserAPIKeyAuth(
|
|
token="hashed-token",
|
|
key_name=key_name,
|
|
key_alias="prod-key",
|
|
spend=10.0,
|
|
max_budget=1.0,
|
|
)
|
|
with pytest.raises(litellm.BudgetExceededError, match="Budget has been exceeded") as exc_info:
|
|
await _virtual_key_max_budget_check(
|
|
valid_token=valid_token,
|
|
proxy_logging_obj=_BudgetAlertRecorder(),
|
|
)
|
|
await asyncio.sleep(0)
|
|
return exc_info.value.message
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize(
|
|
"key_name",
|
|
[
|
|
"sk-mx5ous1o9Iezz5fj3pkLuA",
|
|
"my-company-key-2026",
|
|
"sk-...5LuA-but-longer",
|
|
# /key/generate takes a custom key ending in an escape sequence, and this
|
|
# message reaches a terminal and a log viewer
|
|
"sk-...\x1b[2J",
|
|
"sk-...a\x9bm",
|
|
],
|
|
)
|
|
async def test_key_budget_error_does_not_carry_a_raw_key_name(key_name):
|
|
"""key_name is written masked, but the column has no enforced shape (a direct DB
|
|
write bypasses abbreviate_api_key) and this message is returned to the caller."""
|
|
message = await _run_key_budget_check(key_name)
|
|
assert key_name not in message
|
|
assert "Key=prod-key Current cost" in message
|
|
|
|
|
|
@pytest.mark.asyncio
|
|
@pytest.mark.parametrize("key_name", ["sk-...5LuA", "sk-...", "sk-...ke.!", "sk-...café"])
|
|
async def test_key_budget_error_keeps_the_masked_key_name(key_name):
|
|
"""The masked form is the whole point of naming the key, so it must survive.
|
|
|
|
abbreviate_api_key takes the last four characters of the key verbatim, and a
|
|
custom key may end in punctuation or a non-ASCII character, so those masked
|
|
names are just as valid as the alphanumeric ones."""
|
|
message = await _run_key_budget_check(key_name)
|
|
assert f"Key=prod-key ({key_name}) Current cost" in message
|