Fix self-exclusion hash mismatch and missing throughput field checks

The self-exclusion filter compared raw key strings against SHA-256
hashed tokens from the DB, so keys were never excluded and
double-counting persisted. Now hash data.key before comparison.

Also add tpm_limit_type/rpm_limit_type to _throughput_fields_changed
guard, fall back to existing_key_row.team_id for team limit checks
(matching the org pattern), and add team self-exclusion test.

Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
yuneng-jiang 2026-03-13 15:59:23 -07:00
parent 1038a119ce
commit 818c097ca9
2 changed files with 107 additions and 28 deletions

View file

@ -905,9 +905,11 @@ async def _check_team_key_limits(
keys = await prisma_client.db.litellm_verificationtoken.find_many(
where={"team_id": team_table.team_id},
)
# Exclude the key being updated to avoid double-counting its limits
# Exclude the key being updated to avoid double-counting its limits.
# key.token is the SHA-256 hash stored in DB; data.key is the raw key string.
if isinstance(data, UpdateKeyRequest):
keys = [key for key in keys if key.token != data.key]
hashed_key = hash_token(data.key)
keys = [key for key in keys if key.token != hashed_key]
check_team_key_model_specific_limits(
keys=keys,
team_table=team_table,
@ -1062,9 +1064,11 @@ async def _check_org_key_limits(
keys = await prisma_client.db.litellm_verificationtoken.find_many(
where={"organization_id": org_table.organization_id},
)
# Exclude the key being updated to avoid double-counting its limits
# Exclude the key being updated to avoid double-counting its limits.
# key.token is the SHA-256 hash stored in DB; data.key is the raw key string.
if isinstance(data, UpdateKeyRequest):
keys = [key for key in keys if key.token != data.key]
hashed_key = hash_token(data.key)
keys = [key for key in keys if key.token != hashed_key]
check_org_key_model_specific_limits(
keys=keys,
org_table=org_table,
@ -1932,11 +1936,14 @@ async def update_key_fn(
user_api_key_cache=user_api_key_cache,
)
# Only check team limits if key has a team_id
# Check team limits if key has a team_id (from request or existing key)
team_obj: Optional[LiteLLM_TeamTableCachedObj] = None
if data.team_id is not None:
_team_id_to_check = data.team_id or getattr(
existing_key_row, "team_id", None
)
if _team_id_to_check is not None:
team_obj = await get_team_object(
team_id=data.team_id,
team_id=_team_id_to_check,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
check_db_only=True,
@ -1971,6 +1978,8 @@ async def update_key_fn(
data.organization_id is not None
or data.tpm_limit is not None
or data.rpm_limit is not None
or data.tpm_limit_type is not None
or data.rpm_limit_type is not None
)
if _org_id_to_check is not None and _throughput_fields_changed:
org_table = await get_org_object(

View file

@ -1836,6 +1836,65 @@ async def test_check_team_key_limits_rpm_overallocation():
)
@pytest.mark.asyncio
async def test_check_team_key_limits_on_update_excludes_self():
"""
Test that _check_team_key_limits excludes the key being updated from the
allocated totals. Without this, the key's current limits would be
double-counted: once from find_many and once from data.tpm_limit/rpm_limit.
"""
from litellm.proxy._types import hash_token as _ht
# The key being updated is returned by find_many with its current limits.
# In the DB, token is stored as a SHA-256 hash of the raw key.
self_key = MagicMock()
self_key.token = _ht("sk-self-team-key")
self_key.tpm_limit = 6000
self_key.rpm_limit = 600
self_key.metadata = {}
# Another key in the team
other_key = MagicMock()
other_key.token = _ht("sk-other-team-key")
other_key.tpm_limit = 3000
other_key.rpm_limit = 300
other_key.metadata = {}
mock_prisma_client = AsyncMock()
mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(
return_value=[self_key, other_key]
)
team_table = LiteLLM_TeamTableCachedObj(
team_id="test-team-self",
team_alias="test-team",
tpm_limit=10000,
rpm_limit=1000,
max_budget=100.0,
spend=0.0,
models=[],
blocked=False,
members_with_roles=[],
)
# Updating the key to 7000 TPM. Other key uses 3000, so total = 10000 <= 10000.
# Without the fix, this would be 6000 (self) + 3000 (other) + 7000 = 16000 > 10000.
data = UpdateKeyRequest(
key="sk-self-team-key",
tpm_limit=7000,
rpm_limit=700,
tpm_limit_type="guaranteed_throughput",
rpm_limit_type="guaranteed_throughput",
)
# Should not raise - the key's own limits should be excluded from the sum
await _check_team_key_limits(
team_table=team_table,
data=data,
prisma_client=mock_prisma_client,
)
@pytest.mark.asyncio
async def test_check_team_key_limits_no_team_limits():
"""
@ -6819,8 +6878,10 @@ async def test_check_org_key_limits_on_update_overallocation():
Test that _check_org_key_limits raises HTTPException when updating a key
would exceed organization TPM limits.
"""
from litellm.proxy._types import hash_token as _hash_token
existing_key = MagicMock()
existing_key.token = "sk-other-key"
existing_key.token = _hash_token("sk-other-key")
existing_key.tpm_limit = 15000
existing_key.rpm_limit = 1500
existing_key.metadata = {}
@ -6869,16 +6930,19 @@ async def test_check_org_key_limits_on_update_excludes_self():
allocated totals. Without this, the key's current limits would be
double-counted: once from find_many and once from data.tpm_limit/rpm_limit.
"""
# The key being updated is returned by find_many with its current limits
from litellm.proxy._types import hash_token
# The key being updated is returned by find_many with its current limits.
# In the DB, token is stored as a SHA-256 hash of the raw key.
self_key = MagicMock()
self_key.token = "sk-test-key"
self_key.token = hash_token("sk-test-key")
self_key.tpm_limit = 10000
self_key.rpm_limit = 1000
self_key.metadata = {}
# Another key in the org
other_key = MagicMock()
other_key.token = "sk-other-key"
other_key.token = hash_token("sk-other-key")
other_key.tpm_limit = 5000
other_key.rpm_limit = 500
other_key.metadata = {}
@ -6927,34 +6991,40 @@ def test_update_key_skips_org_check_when_no_throughput_fields_changed():
when only non-throughput fields change on a key that belongs to an org.
This prevents blocking updates when the org has been deleted.
"""
def _check_throughput_changed(data: UpdateKeyRequest) -> bool:
return (
data.organization_id is not None
or data.tpm_limit is not None
or data.rpm_limit is not None
or data.tpm_limit_type is not None
or data.rpm_limit_type is not None
)
# Updating only key_alias — no throughput fields changed
data = UpdateKeyRequest(key="sk-test-key", key_alias="new-alias")
_throughput_fields_changed = (
data.organization_id is not None
or data.tpm_limit is not None
or data.rpm_limit is not None
)
assert _throughput_fields_changed is False
assert _check_throughput_changed(data) is False
# Updating tpm_limit — throughput field changed
data_with_tpm = UpdateKeyRequest(key="sk-test-key", tpm_limit=5000)
_throughput_fields_changed_tpm = (
data_with_tpm.organization_id is not None
or data_with_tpm.tpm_limit is not None
or data_with_tpm.rpm_limit is not None
)
assert _throughput_fields_changed_tpm is True
assert _check_throughput_changed(data_with_tpm) is True
# Updating organization_id — org change triggers check
data_with_org = UpdateKeyRequest(
key="sk-test-key", organization_id="new-org"
)
_throughput_fields_changed_org = (
data_with_org.organization_id is not None
or data_with_org.tpm_limit is not None
or data_with_org.rpm_limit is not None
assert _check_throughput_changed(data_with_org) is True
# Updating tpm_limit_type — limit type change triggers check
data_with_tpm_type = UpdateKeyRequest(
key="sk-test-key", tpm_limit_type="guaranteed_throughput"
)
assert _throughput_fields_changed_org is True
assert _check_throughput_changed(data_with_tpm_type) is True
# Updating rpm_limit_type — limit type change triggers check
data_with_rpm_type = UpdateKeyRequest(
key="sk-test-key", rpm_limit_type="guaranteed_throughput"
)
assert _check_throughput_changed(data_with_rpm_type) is True
def test_update_key_request_has_organization_id():