mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix(proxy): broadcast user cache eviction when tpm_limit or rpm_limit changes
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
5197236444
commit
8baee043d5
2 changed files with 23 additions and 11 deletions
|
|
@ -114,7 +114,7 @@ if TYPE_CHECKING:
|
|||
router: Final = APIRouter()
|
||||
_USER_MODEL_BUDGET_ADAPTER: Final = TypeAdapter(dict[str, float | BudgetConfig])
|
||||
_USER_BUDGET_CACHE_INVALIDATION_BATCH_SIZE: Final = 50
|
||||
_USER_BUDGET_CACHE_FIELDS: Final = frozenset({"max_budget", "model_max_budget"})
|
||||
_USER_LIMIT_CACHE_FIELDS: Final = frozenset({"max_budget", "model_max_budget", "tpm_limit", "rpm_limit"})
|
||||
|
||||
|
||||
def _user_table(
|
||||
|
|
@ -1619,7 +1619,7 @@ async def _update_single_user_helper(
|
|||
|
||||
await _invalidate_user_spend_counter_if_changed(non_default_values)
|
||||
|
||||
if not _USER_BUDGET_CACHE_FIELDS.isdisjoint(non_default_values) or "metadata" in data_json:
|
||||
if not _USER_LIMIT_CACHE_FIELDS.isdisjoint(non_default_values) or "metadata" in data_json:
|
||||
await evict_and_broadcast(
|
||||
cache_keys=(non_default_values["user_id"],),
|
||||
user_api_key_cache=user_api_key_cache,
|
||||
|
|
@ -1977,7 +1977,7 @@ async def bulk_user_update(
|
|||
),
|
||||
)
|
||||
|
||||
if not _USER_BUDGET_CACHE_FIELDS.isdisjoint(non_default_values):
|
||||
if not _USER_LIMIT_CACHE_FIELDS.isdisjoint(non_default_values):
|
||||
for start in range(0, len(all_users_in_db), _USER_BUDGET_CACHE_INVALIDATION_BATCH_SIZE):
|
||||
await asyncio.gather(
|
||||
*(
|
||||
|
|
|
|||
|
|
@ -2346,8 +2346,9 @@ async def test_user_max_budget_update_evicts_cached_user_on_every_worker(mocker:
|
|||
[("tpm_limit", 100), ("rpm_limit", 1)],
|
||||
ids=["tpm_limit", "rpm_limit"],
|
||||
)
|
||||
@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
|
||||
mocker: MockerFixture, field: str, new_limit: int, all_users: bool
|
||||
) -> None:
|
||||
from redis.asyncio import Redis
|
||||
|
||||
|
|
@ -2355,7 +2356,8 @@ async def test_user_rate_limit_update_reaches_cached_user_on_every_worker(
|
|||
from litellm.proxy.auth.auth_checks import get_user_object
|
||||
from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import AuthCacheInvalidationSubscriber
|
||||
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import user_update
|
||||
from litellm.proxy.management_endpoints.internal_user_endpoints import bulk_user_update, user_update
|
||||
from litellm.types.proxy.management_endpoints.internal_user_endpoints import BulkUpdateUserRequest
|
||||
|
||||
published: Final[list[tuple[str, str]]] = [] # mutable-ok: captures messages from the async Redis publisher
|
||||
|
||||
|
|
@ -2381,6 +2383,8 @@ async def test_user_rate_limit_update_reaches_cached_user_on_every_worker(
|
|||
prisma_client: Final = mocker.MagicMock()
|
||||
prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=saved_user)
|
||||
prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock(return_value=updated_user)
|
||||
prisma_client.db.litellm_usertable.find_many = mocker.AsyncMock(return_value=[saved_user])
|
||||
prisma_client.db.litellm_usertable.update_many = mocker.AsyncMock(return_value=1)
|
||||
prisma_client.get_data = mocker.AsyncMock(return_value=saved_user)
|
||||
prisma_client.update_data = mocker.AsyncMock(return_value={"user_id": saved_user.user_id, "data": updated_user})
|
||||
mocker.patch( # test-quality-ok: substitute the database dependency
|
||||
|
|
@ -2424,10 +2428,20 @@ async def test_user_rate_limit_update_reaches_cached_user_on_every_worker(
|
|||
assert other_user_before is not None
|
||||
assert other_user_before.model_dump()[field] == old_limit
|
||||
|
||||
await user_update(
|
||||
data=UpdateUserRequest(user_id=saved_user.user_id, **{field: new_limit}),
|
||||
user_api_key_dict=UserAPIKeyAuth(user_id="admin-spruce", user_role=LitellmUserRoles.PROXY_ADMIN),
|
||||
)
|
||||
admin: Final = UserAPIKeyAuth(user_id="admin-spruce", user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
if all_users:
|
||||
await bulk_user_update(
|
||||
data=BulkUpdateUserRequest(all_users=True, user_updates={field: new_limit}),
|
||||
user_api_key_dict=admin,
|
||||
litellm_changed_by=None,
|
||||
)
|
||||
prisma_client.db.litellm_usertable.update_many.assert_awaited_once_with(where={}, data={field: new_limit})
|
||||
else:
|
||||
await user_update(
|
||||
data=UpdateUserRequest(user_id=saved_user.user_id, **{field: new_limit}),
|
||||
user_api_key_dict=admin,
|
||||
)
|
||||
assert prisma_client.update_data.call_args.kwargs["data"][field] == new_limit
|
||||
await asyncio.sleep(0)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
|
|
@ -2440,8 +2454,6 @@ async def test_user_rate_limit_update_reaches_cached_user_on_every_worker(
|
|||
{"type": "message", "data": message}
|
||||
)
|
||||
|
||||
assert prisma_client.update_data.call_args.kwargs["data"][field] == new_limit
|
||||
|
||||
handling_user_after: Final = await get_user_object(
|
||||
user_id=saved_user.user_id,
|
||||
prisma_client=prisma_client,
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue