diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index f6c088af311..48b89ec3f1b 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -1289,7 +1289,7 @@ def _update_internal_user_params(data_json: dict, data: UpdateUserRequest | Upda fields_set: Final = data.fields_set() if hasattr(data, "fields_set") else set() for k, v in data_json.items(): - if k in ("max_budget", "budget_duration"): + if k in ("max_budget", "budget_duration", "tpm_limit", "rpm_limit"): if k in fields_set: non_default_values[k] = v elif k == "model_max_budget": diff --git a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py index 1d6d8207865..728e156c3e6 100644 --- a/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_internal_user_endpoints.py @@ -2112,6 +2112,18 @@ def test_update_internal_user_params_ignores_other_nones(): assert non_default_values["max_budget"] == 100.0 +@pytest.mark.parametrize("field", ["tpm_limit", "rpm_limit"], ids=["tpm_limit", "rpm_limit"]) +def test_update_internal_user_params_explicit_null_clears_rate_limit_but_omitted_is_untouched( + field: str, +) -> None: + data: Final = UpdateUserRequest(user_id="limit-clear", **{field: None}) + result: Final = _update_internal_user_params(data_json=data.model_dump(exclude_unset=True), data=data) + other_field: Final = "rpm_limit" if field == "tpm_limit" else "tpm_limit" + + assert result[field] is None + assert other_field not in result + + def test_update_internal_user_params_keeps_original_max_budget_when_not_provided(): """ Test that _update_internal_user_params does not include max_budget @@ -2343,12 +2355,12 @@ async def test_user_max_budget_update_evicts_cached_user_on_every_worker(mocker: @pytest.mark.asyncio @pytest.mark.parametrize( ("field", "new_limit"), - [("tpm_limit", 100), ("rpm_limit", 1)], - ids=["tpm_limit", "rpm_limit"], + [("tpm_limit", 100), ("rpm_limit", 1), ("tpm_limit", None), ("rpm_limit", None)], + ids=["tpm_limit", "rpm_limit", "tpm_limit-cleared", "rpm_limit-cleared"], ) @pytest.mark.parametrize("all_users", [False, True], ids=["single-user", "bulk-all-users"]) async def test_user_rate_limit_update_reaches_cached_user_on_every_worker( - mocker: MockerFixture, field: str, new_limit: int, all_users: bool + mocker: MockerFixture, field: str, new_limit: int | None, all_users: bool ) -> None: from redis.asyncio import Redis