fix: guard empty-dict team limits and malformed int in deployment default limits

- Change `if team_limit:` to `if team_limit is not None:` in both
  get_key_model_rpm_limit and get_key_model_tpm_limit so that an
  explicitly-empty team rate-limit map ({}) is returned as-is instead
  of silently falling through to deployment defaults (P1 fix).
- Replace the bare `int()` list comprehension in _get_deployment_default_limit
  with a loop that catches ValueError/TypeError so malformed config strings
  do not raise an unhandled exception during request handling (P2 fix).
- Add corresponding unit tests for both edge cases.

Co-Authored-By: Claude (claude-sonnet-4-6) <noreply@anthropic.com>
This commit is contained in:
Ephrim Stanley 2026-03-19 07:40:47 -04:00
parent e562c1d064
commit ae0769b1df
2 changed files with 64 additions and 7 deletions

View file

@ -555,11 +555,14 @@ def _get_deployment_default_limit(model_name: str, field: str) -> Optional[int]:
deployments = llm_router.get_model_list(model_name=model_name)
if not deployments:
return None
limits = [
int(deployment.get("litellm_params", {}).get(field))
for deployment in deployments
if deployment.get("litellm_params", {}).get(field) is not None
]
limits = []
for deployment in deployments:
raw = deployment.get("litellm_params", {}).get(field)
if raw is not None:
try:
limits.append(int(raw))
except (ValueError, TypeError):
pass
return min(limits) if limits else None
@ -602,7 +605,7 @@ def get_key_model_rpm_limit(
# 3. Fallback to team metadata
if user_api_key_dict.team_metadata:
team_limit = user_api_key_dict.team_metadata.get("model_rpm_limit")
if team_limit:
if team_limit is not None:
return team_limit
# 4. Fallback to deployment default_api_key_rpm_limit
@ -645,7 +648,7 @@ def get_key_model_tpm_limit(
# 3. Fallback to team metadata
if user_api_key_dict.team_metadata:
team_limit = user_api_key_dict.team_metadata.get("model_tpm_limit")
if team_limit:
if team_limit is not None:
return team_limit
# 4. Fallback to deployment default_api_key_tpm_limit

View file

@ -71,6 +71,19 @@ class TestGetKeyModelRpmLimit:
assert result is None
def test_team_metadata_empty_rpm_dict_falls_through_to_deployment_default(self):
"""Explicitly empty team model_rpm_limit ({}) should be returned as-is, not fallen through."""
# An empty dict is a valid team limit map (no per-model limits configured).
# It should be returned directly rather than falling through to deployment defaults,
# so a team with an empty map is treated as unconstrained at the team level.
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-123",
team_metadata={"model_rpm_limit": {}},
)
result = get_key_model_rpm_limit(user_api_key_dict)
assert result == {}
class TestGetKeyModelTpmLimit:
"""Tests for get_key_model_tpm_limit function."""
@ -137,6 +150,33 @@ class TestGetKeyModelTpmLimit:
assert result == {"gpt-4": 10000}
def test_team_metadata_empty_tpm_dict_falls_through_to_deployment_default(self):
"""Explicitly empty team model_tpm_limit ({}) should be returned as-is, not fallen through."""
# An empty dict is a valid team limit map (no per-model limits configured).
# It should be returned directly rather than falling through to deployment defaults,
# so a team with an empty map is treated as unconstrained at the team level.
user_api_key_dict = UserAPIKeyAuth(
api_key="sk-123",
team_metadata={"model_tpm_limit": {}},
)
result = get_key_model_tpm_limit(user_api_key_dict)
assert result == {}
def test_skips_deployments_with_malformed_limit_value(self):
"""Deployments with non-integer-parseable limit values are skipped without raising."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
{"model_name": "model1", "litellm_params": {"default_api_key_tpm_limit": "not-a-number"}},
_make_deployment_dict("model1", tpm=500),
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1")
# The malformed deployment is skipped; the valid one provides 500
assert result == {"model1": 500}
class TestGetCustomerIdFromStandardHeaders:
"""Tests for _get_customer_id_from_standard_headers helper function."""
@ -414,6 +454,20 @@ class TestDeploymentDefaultRpmLimit:
assert result == {"model1": 75}
def test_skips_deployments_with_malformed_limit_value(self):
"""Deployments with non-integer-parseable limit values are skipped without raising."""
user_api_key_dict = UserAPIKeyAuth(api_key="sk-123")
mock_router = MagicMock()
mock_router.get_model_list.return_value = [
{"model_name": "model1", "litellm_params": {"default_api_key_rpm_limit": "not-a-number"}},
_make_deployment_dict("model1", rpm=100),
]
with patch(_ROUTER_PATCH, mock_router):
result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1")
# The malformed deployment is skipped; the valid one provides 100
assert result == {"model1": 100}
class TestDeploymentDefaultTpmLimit:
"""Tests for deployment default_api_key_tpm_limit fallback in get_key_model_tpm_limit."""