From 0f58f8715df873f854bd8c74891f162c63ef8e64 Mon Sep 17 00:00:00 2001 From: ryan-crabbe-berri Date: Sat, 1 Aug 2026 17:35:26 -0700 Subject: [PATCH] fix(team): harden team model budget against review findings Key precedence now requires an enforceable key entry: a key-level model_max_budget entry only exempts the key from the team cap when its max_budget is finite and greater than zero, since zero or missing caps are never enforced key-side and would otherwise disable both team enforcement and team accounting for that key. Team spend counters and window start times are now keyed by the canonical budget-config entry name plus the window duration, so provider-prefixed and bare spellings of one model share a single counter and each (model, duration) window resets independently instead of sharing one per-team start time. --- .../proxy/hooks/model_max_budget_limiter.py | 92 +++++++----- .../credential_migration.py | 73 +++------- ...test_unit_test_max_model_budget_limiter.py | 132 +++++++++++++++++- 3 files changed, 205 insertions(+), 92 deletions(-) diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index 7371c053f9f..c1e90d76f2c 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -1,4 +1,5 @@ import json +import math from collections.abc import Mapping from types import MappingProxyType from typing import List, Optional @@ -158,22 +159,42 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): {_model: BudgetConfig(**_budget_info) for _model, _budget_info in model_max_budget.items()} ) + def _get_matched_budget_model_name( + self, model: str, internal_model_max_budget: Mapping[str, BudgetConfig] + ) -> str | None: + """ + Resolve the budget-config entry name `model` matches (exact, then + `{provider}/{model}`-normalized). The entry name is the canonical form for + spend counter keys so prefixed and bare spellings of one model share a counter. + """ + if model in internal_model_max_budget: + return model + stripped_model = self._get_model_without_custom_llm_provider(model) + if stripped_model in internal_model_max_budget: + return stripped_model + return None + def _key_already_covers_model( self, key_model_max_budget: Mapping[str, Mapping[str, str | float]] | None, model: str ) -> bool: """ - True iff the requesting key declares its own model_max_budget entry for `model` - (exact or `{provider}/{model}`-normalized match). Key entries take precedence - over team-level defaults, so the team budget is neither checked nor incremented - for such (key, model) pairs. + True iff the requesting key declares its own ENFORCEABLE model_max_budget entry + for `model` (exact or `{provider}/{model}`-normalized match) — a finite cap + strictly greater than zero, matching what is_key_within_model_budget actually + enforces. Only such entries take precedence over team-level defaults; a zero, + missing, or non-finite cap is never enforced key-side, so treating it as + covering would silently disable the team budget for that key. """ if not key_model_max_budget: return False + budget_config = self._get_request_model_budget_config( + model=model, internal_model_max_budget=self._coerce_budget_configs(key_model_max_budget) + ) return ( - self._get_request_model_budget_config( - model=model, internal_model_max_budget=self._coerce_budget_configs(key_model_max_budget) - ) - is not None + budget_config is not None + and budget_config.max_budget is not None + and math.isfinite(budget_config.max_budget) + and budget_config.max_budget > 0 ) async def is_team_within_model_budget( @@ -201,19 +222,20 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): verbose_proxy_logger.debug("team internal_model_max_budget %s", internal_model_max_budget) - budget_config = self._get_request_model_budget_config( + matched_model = self._get_matched_budget_model_name( model=model, internal_model_max_budget=internal_model_max_budget ) - if budget_config is None: + if matched_model is None: verbose_proxy_logger.debug(f"Model {model} not found in team_model_max_budget") return True + budget_config = internal_model_max_budget[matched_model] if not budget_config.max_budget or budget_config.max_budget <= 0: return True current_spend = await self._get_team_spend_for_model( team_id=team_id, - model=model, + matched_model=matched_model, budget_config=budget_config, ) verbose_proxy_logger.debug( @@ -233,30 +255,22 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): async def _get_team_spend_for_model( self, team_id: str, - model: str, + matched_model: str, budget_config: BudgetConfig, ) -> float | None: """ - Get the current team spend for a model. - - Lookup model in this order: - 1. model: directly look up `model` - 2. If 1, does not exist, check if passed as {custom_llm_provider}/model + Get the current team spend for a model. `matched_model` is the canonical + budget-config entry name from _get_matched_budget_model_name, so read and + write always address one counter regardless of how the request spelled the + model. """ team_model_spend_cache_key = ( - f"{TEAM_MODEL_SPEND_CACHE_KEY_PREFIX}:{team_id}:{model}:{budget_config.budget_duration}" + f"{TEAM_MODEL_SPEND_CACHE_KEY_PREFIX}:{team_id}:{matched_model}:{budget_config.budget_duration}" ) - _current_spend = await self.dual_cache.async_get_cache( + return await self.dual_cache.async_get_cache( key=team_model_spend_cache_key, ) - if _current_spend is None: - team_model_spend_cache_key = f"{TEAM_MODEL_SPEND_CACHE_KEY_PREFIX}:{team_id}:{self._get_model_without_custom_llm_provider(model)}:{budget_config.budget_duration}" - _current_spend = await self.dual_cache.async_get_cache( - key=team_model_spend_cache_key, - ) - return _current_spend - async def _track_team_spend_for_model( self, team_id: str, @@ -267,21 +281,31 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): ) -> None: """ Increment the shared team counter for `model`, unless the requesting key - declares its own model_max_budget entry for it — mirror of the precedence - rule in is_team_within_model_budget, so keys with private caps never drive - the team counter and block sibling keys. + declares its own enforceable model_max_budget entry for it — mirror of the + precedence rule in is_team_within_model_budget, so keys with private caps + never drive the team counter and block sibling keys. + + The counter and its window start time are keyed by the canonical config + entry name plus the window duration: spellings of one model share a + counter, and each (model, duration) window resets independently. """ if self._key_already_covers_model(key_model_max_budget, model): return - budget_config = self._get_request_model_budget_config( - model=model, internal_model_max_budget=self._coerce_budget_configs(team_model_max_budget) + internal_model_max_budget = self._coerce_budget_configs(team_model_max_budget) + matched_model = self._get_matched_budget_model_name( + model=model, internal_model_max_budget=internal_model_max_budget ) - if budget_config is None or not budget_config.budget_duration: + if matched_model is None: + return + budget_config = internal_model_max_budget[matched_model] + if not budget_config.budget_duration: return - team_spend_key = f"{TEAM_MODEL_SPEND_CACHE_KEY_PREFIX}:{team_id}:{model}:{budget_config.budget_duration}" - team_start_time_key = f"team_model_budget_start_time:{team_id}" + team_spend_key = ( + f"{TEAM_MODEL_SPEND_CACHE_KEY_PREFIX}:{team_id}:{matched_model}:{budget_config.budget_duration}" + ) + team_start_time_key = f"team_model_budget_start_time:{team_id}:{matched_model}:{budget_config.budget_duration}" await self._increment_spend_for_key( budget_config=budget_config, spend_key=team_spend_key, diff --git a/litellm/proxy/management_endpoints/credential_migration.py b/litellm/proxy/management_endpoints/credential_migration.py index 4d51295f8dc..6f79a39c883 100644 --- a/litellm/proxy/management_endpoints/credential_migration.py +++ b/litellm/proxy/management_endpoints/credential_migration.py @@ -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 diff --git a/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py b/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py index 7eceadbcc2d..b8117ad61c3 100644 --- a/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py +++ b/tests/proxy_unit_tests/test_unit_test_max_model_budget_limiter.py @@ -706,7 +706,7 @@ async def test_team_spend_cache_key_format(budget_limiter): budget_limiter.dual_cache, "async_get_cache", side_effect=fake_get ): await budget_limiter._get_team_spend_for_model( - team_id="team-xyz", model="gpt-4", budget_config=budget_cfg + team_id="team-xyz", matched_model="gpt-4", budget_config=budget_cfg ) assert captured["key"] == f"{TEAM_MODEL_SPEND_CACHE_KEY_PREFIX}:team-xyz:gpt-4:1d" @@ -836,3 +836,133 @@ async def test_async_log_success_event_skips_team_without_team_id(budget_limiter assert not any( k.startswith(TEAM_MODEL_SPEND_CACHE_KEY_PREFIX) for k in incremented ) + + +@pytest.mark.asyncio +async def test_zero_limit_key_entry_does_not_bypass_team_budget(budget_limiter): + """A key entry with budget_limit 0 is never enforced key-side, so it must NOT + count as covering the model; the team cap still applies (Veria finding).""" + with patch.object( + budget_limiter, "_get_team_spend_for_model", AsyncMock(return_value=150.0) + ): + with pytest.raises(litellm.BudgetExceededError): + await budget_limiter.is_team_within_model_budget( + team_id="team-1", + team_model_max_budget={ + "gpt-4": {"budget_limit": 100.0, "time_period": "1d"} + }, + model="gpt-4", + key_model_max_budget={ + "gpt-4": {"budget_limit": 0, "time_period": "1d"} + }, + ) + + +@pytest.mark.asyncio +async def test_zero_limit_key_entry_still_tracks_team_spend(budget_limiter): + """Spend from a zero-limit key entry must still drive the shared team counter, + otherwise the bypass extends to accounting as well.""" + from litellm.proxy.hooks.model_max_budget_limiter import ( + TEAM_MODEL_SPEND_CACHE_KEY_PREFIX, + ) + + incremented = {} + + async def fake_increment(budget_config, spend_key, start_time_key, response_cost): + incremented[spend_key] = response_cost + + kwargs = { + "standard_logging_object": { + "response_cost": 4.0, + "model_group": "gpt-4", + "model": "openai/gpt-4", + "metadata": {"user_api_key_hash": "hash-1"}, + }, + "litellm_params": { + "metadata": { + "user_api_key_team_id": "team-1", + "user_api_key_team_model_max_budget": { + "gpt-4": {"budget_limit": 100.0, "time_period": "1d"} + }, + "user_api_key_model_max_budget": { + "gpt-4": {"budget_limit": 0, "time_period": "1d"} + }, + } + }, + } + with patch.object( + budget_limiter, "_increment_spend_for_key", side_effect=fake_increment + ): + await budget_limiter.async_log_success_event(kwargs, None, 0, 0) + + assert any(k.startswith(TEAM_MODEL_SPEND_CACHE_KEY_PREFIX) for k in incremented) + + +@pytest.mark.asyncio +async def test_team_counter_shared_across_model_spellings(budget_limiter): + """Spend tracked under openai/gpt-4 and enforcement for gpt-4 (and vice versa) + must hit ONE canonical counter, or alternating spellings doubles the cap + (Greptile fragmentation finding). Uses the real DualCache, no spend-read mocks.""" + team_budget = {"gpt-4": {"budget_limit": 5.0, "time_period": "1d"}} + kwargs = { + "standard_logging_object": { + "response_cost": 6.0, + "model_group": "openai/gpt-4", + "model": "openai/gpt-4", + "metadata": {"user_api_key_hash": "hash-1"}, + }, + "litellm_params": { + "metadata": { + "user_api_key_team_id": "team-frag", + "user_api_key_team_model_max_budget": team_budget, + } + }, + } + await budget_limiter.async_log_success_event(kwargs, None, 0, 0) + + for request_model in ("gpt-4", "openai/gpt-4"): + with pytest.raises(litellm.BudgetExceededError): + await budget_limiter.is_team_within_model_budget( + team_id="team-frag", + team_model_max_budget=team_budget, + model=request_model, + ) + + +@pytest.mark.asyncio +async def test_team_window_start_key_is_per_model_and_duration(budget_limiter): + """Each (model, duration) window needs its own start-time key; a shared per-team + key lets a short window's reset restart a longer window early (Greptile finding).""" + captured = [] + + async def fake_increment(budget_config, spend_key, start_time_key, response_cost): + captured.append(start_time_key) + + team_budget = { + "gpt-4": {"budget_limit": 100.0, "time_period": "1d"}, + "claude-3": {"budget_limit": 500.0, "time_period": "7d"}, + } + for model_group in ("gpt-4", "claude-3"): + kwargs = { + "standard_logging_object": { + "response_cost": 1.0, + "model_group": model_group, + "model": model_group, + "metadata": {"user_api_key_hash": "hash-1"}, + }, + "litellm_params": { + "metadata": { + "user_api_key_team_id": "team-windows", + "user_api_key_team_model_max_budget": team_budget, + } + }, + } + with patch.object( + budget_limiter, "_increment_spend_for_key", side_effect=fake_increment + ): + await budget_limiter.async_log_success_event(kwargs, None, 0, 0) + + assert captured == [ + "team_model_budget_start_time:team-windows:gpt-4:1d", + "team_model_budget_start_time:team-windows:claude-3:7d", + ]