mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
fix(proxy): enforce user budget on team keys and honor fail-closed policy in max_budget_limiter
This commit is contained in:
parent
65ca095d4d
commit
3148146b9d
3 changed files with 119 additions and 67 deletions
|
|
@ -30,9 +30,12 @@ class _PROXY_MaxBudgetLimiter(CustomLogger):
|
|||
if max_budget is None or user_id is None:
|
||||
return
|
||||
|
||||
# Personal budget applies only to non-team requests, matching
|
||||
# the explicit team-key exemption in common_checks section 4.1.
|
||||
if user_api_key_dict.team_id is not None:
|
||||
# User budgets apply to team keys too, matching common_checks; the
|
||||
# legacy team-key exemption only applies when general_settings
|
||||
# skip_user_budget_on_team_key is explicitly enabled (see #12905).
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
if user_api_key_dict.team_id is not None and general_settings.get("skip_user_budget_on_team_key") is True:
|
||||
return
|
||||
|
||||
# The reservation path admits at the strict-`<` boundary and
|
||||
|
|
@ -74,6 +77,13 @@ class _PROXY_MaxBudgetLimiter(CustomLogger):
|
|||
except HTTPException as e:
|
||||
raise e
|
||||
except Exception as e:
|
||||
# An infra error (Redis/cache/DB) left spend unverifiable. Under
|
||||
# fail_closed_budget_enforcement reject rather than silently drop
|
||||
# this enforcement layer; otherwise keep the documented fail-open.
|
||||
from litellm.proxy.proxy_server import _fail_closed_budget_enforcement
|
||||
|
||||
if _fail_closed_budget_enforcement():
|
||||
raise
|
||||
verbose_logger.exception(
|
||||
"litellm.proxy.hooks.max_budget_limiter.py::async_pre_call_hook(): Exception occured - {}".format(
|
||||
str(e)
|
||||
|
|
|
|||
|
|
@ -130,9 +130,7 @@ def classify_value(value: object, key: str = "scan") -> ValueClass:
|
|||
return "plaintext"
|
||||
if value.startswith(_V2_GCM_PREFIX):
|
||||
return "migrated"
|
||||
decrypted = decrypt_value_helper(
|
||||
value=value, key=key, exception_type="debug", return_original_value=False
|
||||
)
|
||||
decrypted = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False)
|
||||
if decrypted is None:
|
||||
# Did not decrypt under nacl and has no v2 marker: legacy plaintext.
|
||||
return "plaintext"
|
||||
|
|
@ -151,9 +149,7 @@ def reencrypt_value(value: object, key: str = "migrate") -> object:
|
|||
return value
|
||||
if value.startswith(_V2_GCM_PREFIX):
|
||||
return value # idempotent: already migrated
|
||||
decrypted = decrypt_value_helper(
|
||||
value=value, key=key, exception_type="debug", return_original_value=False
|
||||
)
|
||||
decrypted = decrypt_value_helper(value=value, key=key, exception_type="debug", return_original_value=False)
|
||||
if decrypted is None:
|
||||
# Either legacy plaintext (no ciphertext to migrate) or corrupt. Either
|
||||
# way, do not overwrite — preserve the value as stored.
|
||||
|
|
@ -161,9 +157,7 @@ def reencrypt_value(value: object, key: str = "migrate") -> object:
|
|||
return encrypt_value_helper(decrypted)
|
||||
|
||||
|
||||
def reencrypt_selective_dict(
|
||||
data: dict[str, object], sensitive_keys: list[str]
|
||||
) -> dict[str, object]:
|
||||
def reencrypt_selective_dict(data: dict[str, object], sensitive_keys: list[str]) -> dict[str, object]:
|
||||
"""Return a copy of ``data`` with only ``sensitive_keys`` re-encrypted.
|
||||
|
||||
Non-sensitive fields (e.g. ``base_url``, ``connection_id``) are left as-is.
|
||||
|
|
@ -212,9 +206,7 @@ async def _migrate_config_settings_row(
|
|||
dict with selected sensitive fields (vantage_settings / cloudzero_settings).
|
||||
"""
|
||||
report = LocationReport(location=param_name)
|
||||
record = await prisma_client.db.litellm_config.find_unique(
|
||||
where={"param_name": param_name}
|
||||
)
|
||||
record = await prisma_client.db.litellm_config.find_unique(where={"param_name": param_name})
|
||||
if record is None or record.param_value is None:
|
||||
return report
|
||||
|
||||
|
|
@ -266,9 +258,7 @@ async def _migrate_sso_config(prisma_client: object, dry_run: bool) -> LocationR
|
|||
every present string field.
|
||||
"""
|
||||
report = LocationReport(location="sso_config")
|
||||
record = await prisma_client.db.litellm_ssoconfig.find_unique(
|
||||
where={"id": "sso_config"}
|
||||
)
|
||||
record = await prisma_client.db.litellm_ssoconfig.find_unique(where={"id": "sso_config"})
|
||||
if record is None or record.sso_settings is None:
|
||||
return report
|
||||
|
||||
|
|
@ -344,9 +334,7 @@ async def _migrate_callback_vars_table(
|
|||
rows = await table.find_many()
|
||||
for row in rows or []:
|
||||
metadata = getattr(row, "metadata", None)
|
||||
if not isinstance(metadata, dict) or (
|
||||
"logging" not in metadata and "callback_settings" not in metadata
|
||||
):
|
||||
if not isinstance(metadata, dict) or ("logging" not in metadata and "callback_settings" not in metadata):
|
||||
continue
|
||||
|
||||
# Classify every callback-var value directly (strip the litellm_enc::
|
||||
|
|
@ -534,9 +522,7 @@ async def _scan_config_env_vars(prisma_client: object) -> LocationReport:
|
|||
"""Scan the ``environment_variables`` config row (``param_value`` dict)."""
|
||||
report = LocationReport(location="config_environment_variables")
|
||||
try:
|
||||
record = await prisma_client.db.litellm_config.find_unique(
|
||||
where={"param_name": "environment_variables"}
|
||||
)
|
||||
record = await prisma_client.db.litellm_config.find_unique(where={"param_name": "environment_variables"})
|
||||
except Exception as e: # pragma: no cover - defensive
|
||||
verbose_proxy_logger.debug("scan: config env vars unavailable: %s", str(e))
|
||||
return report
|
||||
|
|
@ -557,11 +543,7 @@ async def _scan_covered_tables(prisma_client: object) -> list[LocationReport]:
|
|||
"""Read-only classification of every rotation-covered table. No writes."""
|
||||
reports: list[LocationReport] = []
|
||||
for location, db_attr, json_cols, scalar_cols in _COVERED_TABLE_SPECS:
|
||||
reports.append(
|
||||
await _scan_one_table(
|
||||
prisma_client, location, db_attr, json_cols, scalar_cols
|
||||
)
|
||||
)
|
||||
reports.append(await _scan_one_table(prisma_client, location, db_attr, json_cols, scalar_cols))
|
||||
reports.append(await _scan_config_env_vars(prisma_client))
|
||||
return reports
|
||||
|
||||
|
|
@ -575,9 +557,7 @@ _VANTAGE_SENSITIVE = ["api_key", "integration_token"]
|
|||
_CLOUDZERO_SENSITIVE = ["api_key"]
|
||||
|
||||
|
||||
async def _migrate_covered_tables(
|
||||
prisma_client: object, user_api_key_dict: object
|
||||
) -> list[LocationReport]:
|
||||
async def _migrate_covered_tables(prisma_client: object, user_api_key_dict: object) -> list[LocationReport]:
|
||||
"""Re-encrypt the tables already covered by ``_rotate_master_key`` (model
|
||||
table, credentials, MCP credential/env tables, config environment_variables)
|
||||
by running that orchestrator in *same-key* mode. With the AES gate on, the
|
||||
|
|
@ -597,8 +577,7 @@ async def _migrate_covered_tables(
|
|||
current_key = _get_salt_key()
|
||||
if current_key is None:
|
||||
raise RuntimeError(
|
||||
"Cannot migrate covered tables: no salt key / master key is set. "
|
||||
"Set LITELLM_SALT_KEY before migrating."
|
||||
"Cannot migrate covered tables: no salt key / master key is set. Set LITELLM_SALT_KEY before migrating."
|
||||
)
|
||||
await _rotate_master_key(
|
||||
prisma_client=cast("PrismaClient", prisma_client),
|
||||
|
|
@ -648,19 +627,9 @@ async def migrate_encryption(
|
|||
|
||||
# Net-new walkers (items 3, 4, 11, 12, 13).
|
||||
report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run))
|
||||
report.add(
|
||||
await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run)
|
||||
)
|
||||
report.add(
|
||||
await _migrate_config_settings_row(
|
||||
prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run
|
||||
)
|
||||
)
|
||||
report.add(
|
||||
await _migrate_config_settings_row(
|
||||
prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run
|
||||
)
|
||||
)
|
||||
report.add(await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run))
|
||||
report.add(await _migrate_config_settings_row(prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run))
|
||||
report.add(await _migrate_config_settings_row(prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run))
|
||||
report.add(await _migrate_sso_config(prisma_client, dry_run))
|
||||
|
||||
return report
|
||||
|
|
@ -683,20 +652,10 @@ async def check_encryption(prisma_client: object) -> MigrationReport:
|
|||
|
||||
# Net-new walker locations, in dry-run (read-only) mode.
|
||||
report.add(await _migrate_callback_vars_table(prisma_client, "team", dry_run=True))
|
||||
report.add(await _migrate_callback_vars_table(prisma_client, "verification_token", dry_run=True))
|
||||
report.add(await _migrate_config_settings_row(prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run=True))
|
||||
report.add(
|
||||
await _migrate_callback_vars_table(
|
||||
prisma_client, "verification_token", dry_run=True
|
||||
)
|
||||
)
|
||||
report.add(
|
||||
await _migrate_config_settings_row(
|
||||
prisma_client, "vantage_settings", _VANTAGE_SENSITIVE, dry_run=True
|
||||
)
|
||||
)
|
||||
report.add(
|
||||
await _migrate_config_settings_row(
|
||||
prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run=True
|
||||
)
|
||||
await _migrate_config_settings_row(prisma_client, "cloudzero_settings", _CLOUDZERO_SENSITIVE, dry_run=True)
|
||||
)
|
||||
report.add(await _migrate_sso_config(prisma_client, dry_run=True))
|
||||
return report
|
||||
|
|
|
|||
|
|
@ -163,7 +163,10 @@ async def test_does_not_skip_when_reservation_covers_a_different_counter():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_keys_skip_personal_budget():
|
||||
async def test_team_keys_enforce_personal_budget_by_default():
|
||||
"""Regression for #33323: with the default policy, a user's personal budget
|
||||
is enforced even when the key belongs to a team. The hook must read spend
|
||||
and reject, matching common_checks._user_max_budget_check."""
|
||||
handler = _PROXY_MaxBudgetLimiter()
|
||||
user_api_key_dict = _make_user_api_key_auth(
|
||||
user_max_budget=10.0,
|
||||
|
|
@ -174,17 +177,97 @@ async def test_team_keys_skip_personal_budget():
|
|||
"litellm.proxy.proxy_server.get_current_spend",
|
||||
new=AsyncMock(return_value=999.0),
|
||||
) as mock_get_spend:
|
||||
result = await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={},
|
||||
call_type="completion",
|
||||
)
|
||||
with patch.dict("litellm.proxy.proxy_server.general_settings", {}, clear=True):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 429
|
||||
mock_get_spend.assert_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_team_keys_skip_personal_budget_when_opt_out_set():
|
||||
"""skip_user_budget_on_team_key=True restores the legacy behavior where a
|
||||
team key's request bypasses the user's personal budget entirely."""
|
||||
handler = _PROXY_MaxBudgetLimiter()
|
||||
user_api_key_dict = _make_user_api_key_auth(
|
||||
user_max_budget=10.0,
|
||||
team_id="team-1",
|
||||
)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.get_current_spend",
|
||||
new=AsyncMock(return_value=999.0),
|
||||
) as mock_get_spend:
|
||||
with patch.dict(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"skip_user_budget_on_team_key": True},
|
||||
clear=True,
|
||||
):
|
||||
result = await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert result is None
|
||||
mock_get_spend.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spend_lookup_error_fails_open_by_default():
|
||||
"""Default policy: an infra error reading spend is swallowed so the hook
|
||||
does not hard-fail requests."""
|
||||
handler = _PROXY_MaxBudgetLimiter()
|
||||
user_api_key_dict = _make_user_api_key_auth(user_max_budget=10.0)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.get_current_spend",
|
||||
new=AsyncMock(side_effect=RuntimeError("redis down")),
|
||||
):
|
||||
with patch.dict("litellm.proxy.proxy_server.general_settings", {}, clear=True):
|
||||
result = await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
assert result is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_spend_lookup_error_fails_closed_when_enabled():
|
||||
"""Regression for #33323: with fail_closed_budget_enforcement, an infra
|
||||
error reading spend must not be silently swallowed; the request is rejected
|
||||
rather than admitted on an unenforceable budget."""
|
||||
handler = _PROXY_MaxBudgetLimiter()
|
||||
user_api_key_dict = _make_user_api_key_auth(user_max_budget=10.0)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.proxy_server.get_current_spend",
|
||||
new=AsyncMock(side_effect=RuntimeError("redis down")),
|
||||
):
|
||||
with patch.dict(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{"fail_closed_budget_enforcement": True},
|
||||
clear=True,
|
||||
):
|
||||
with pytest.raises(RuntimeError):
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
cache=DualCache(),
|
||||
data={},
|
||||
call_type="completion",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_max_budget_passes():
|
||||
handler = _PROXY_MaxBudgetLimiter()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue