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.
This commit is contained in:
ryan-crabbe-berri 2026-08-01 17:35:26 -07:00
parent 7a9a1a0d45
commit 0f58f8715d
3 changed files with 205 additions and 92 deletions

View file

@ -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,

View file

@ -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

View file

@ -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",
]