mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
e562c1d064
commit
ae0769b1df
2 changed files with 64 additions and 7 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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."""
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue