This commit is contained in:
devin-ai-integration[bot] 2026-10-04 21:37:26 +00:00 • committed by GitHub
commit a85fbe848d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
12 changed files with 925 additions and 34 deletions

View file

@ -3611,6 +3611,8 @@ class UserInfoV2Response(LiteLLMPydanticObjectBase):
user_role: str | None = None
spend: float = 0.0
max_budget: float | None = None
tpm_limit: int | None = None
rpm_limit: int | None = None
models: list[str] = []
budget_duration: str | None = None
budget_reset_at: datetime | None = None

View file

@ -117,7 +117,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(
@ -1150,6 +1150,8 @@ async def user_info_v2(
user_role=user_data.get("user_role"),
spend=user_data.get("spend", 0.0),
max_budget=user_data.get("max_budget"),
tpm_limit=user_data.get("tpm_limit"),
rpm_limit=user_data.get("rpm_limit"),
models=user_data.get("models") or [],
budget_duration=user_data.get("budget_duration"),
budget_reset_at=user_data.get("budget_reset_at"),
@ -1290,7 +1292,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":
@ -1623,7 +1625,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,
@ -1981,7 +1983,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(
*(

View file

@ -0,0 +1,381 @@
import json
from collections.abc import Mapping
from contextlib import ExitStack
from typing import Final
from uuid import uuid4
import httpx
import pytest
from pydantic import JsonValue, TypeAdapter
from tests.integration._support.client import (
JSON_OBJECT,
Gateway,
eventually,
object_value,
string_value,
)
from tests.integration._support.database import read_rows
from tests.integration._support.wire import Reply, Request, wire_server
_HEADER_VALUE_ADAPTER: Final[TypeAdapter[str | None]] = TypeAdapter(str | None)
_UPSTREAM_REPLIES: Final[Mapping[str, Mapping[str, JsonValue]]] = {
"/v1/chat/completions": {
"id": "chatcmpl_hook_isolation",
"object": "chat.completion",
"created": 1,
"model": "gpt-5.6",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
},
"/v1/responses": {
"id": "resp_hook_isolation",
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-5.6",
"output": [
{
"type": "message",
"id": "msg_hook_isolation",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "hi", "annotations": []}],
}
],
"parallel_tool_calls": False,
"tool_choice": "auto",
"tools": [],
"usage": {"input_tokens": 1, "output_tokens": 1, "total_tokens": 2},
},
}
def _chat(proxy: Gateway, model: str, key: str) -> httpx.Response:
return proxy.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": f"user rate limit probe {uuid4().hex}"}]},
key=key,
)
def _assert_user_rate_limit_error(response: httpx.Response, user: str, limit_type: str) -> None:
request_limit: Final[str | None] = _HEADER_VALUE_ADAPTER.validate_python(
response.headers.get("x-ratelimit-user-limit-requests")
)
token_limit: Final[str | None] = _HEADER_VALUE_ADAPTER.validate_python(
response.headers.get("x-ratelimit-user-limit-tokens")
)
context: Final = (
f"Expected a user {limit_type} limit error for {user}, received HTTP {response.status_code} with "
f"user limits requests={request_limit}, tokens={token_limit}: {response.text}"
)
assert response.status_code == 429, context
body: Final = JSON_OBJECT.validate_json(response.content)
error: Final = object_value(body["error"])
message: Final = string_value(error["message"])
assert error.get("type") == "throttling_error", context
assert message.startswith(f"Rate limit exceeded for user: {user}. Limit type: {limit_type}. Current limit: 1,"), (
context
)
def _route_request(proxy: Gateway, route: str, model: str, key: str, stream: bool) -> httpx.Response:
marker: Final = f"user rpm route probe {uuid4().hex}"
if route == "/v1/messages":
return proxy.request(
"POST",
route,
{"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": marker}]},
key=key,
headers={"anthropic-version": "2023-06-01"},
)
if route == "/v1/responses":
return proxy.request(
"POST",
route,
{"model": model, "input": marker, "max_output_tokens": 16, "store": False},
key=key,
)
return proxy.request(
"POST",
route,
{"model": model, "messages": [{"role": "user", "content": marker}], "stream": stream},
key=key,
)
def _assert_route_user_requests_limit_error(response: httpx.Response, user: str, route: str) -> None:
request_limit: Final[str | None] = _HEADER_VALUE_ADAPTER.validate_python(
response.headers.get("x-ratelimit-user-limit-requests")
)
token_limit: Final[str | None] = _HEADER_VALUE_ADAPTER.validate_python(
response.headers.get("x-ratelimit-user-limit-tokens")
)
context: Final = (
f"Expected a user requests limit error for {user} on {route}, received HTTP {response.status_code} with "
f"user limits requests={request_limit}, tokens={token_limit}: {response.text}"
)
assert response.status_code == 429, context
body: Final = JSON_OBJECT.validate_json(response.content)
error: Final = object_value(body["error"])
assert route != "/v1/messages" or body.get("type") == "error", context
expected_error_type: Final = "rate_limit_error" if route == "/v1/messages" else "throttling_error"
assert error.get("type") == expected_error_type, context
message: Final = string_value(error["message"])
assert message.startswith(f"Rate limit exceeded for user: {user}. Limit type: requests. Current limit: 1,"), context
def _assert_user_rate_limit_on_every_proxy(
gateway: Gateway,
peer: Gateway,
model: str,
user: str,
key: str,
) -> None:
responses: Final = eventually(
lambda: (_chat(gateway, model, key), _chat(peer, model, key)),
lambda observed: all(response.status_code == 429 for response in observed),
seconds=10,
return_last_on_timeout=True,
)
context: Final = tuple(
(
response.status_code,
response.headers.get("x-ratelimit-user-limit-requests"),
response.headers.get("x-ratelimit-user-limit-tokens"),
response.text,
)
for response in responses
)
assert tuple(response.status_code for response in responses) == (429, 429), (
f"Expected the user RPM limit on gateway and peer for {user}, received {context!r}"
)
_assert_user_rate_limit_error(responses[0], user, "requests")
_assert_user_rate_limit_error(responses[1], user, "requests")
@pytest.mark.parametrize("field", ("tpm_limit", "rpm_limit"))
def test_user_rate_limit_lowered_on_gateway_is_enforced_by_peer(gateway: Gateway, peer: Gateway, field: str) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
user: Final = scenario.user(tpm_limit=100000, rpm_limit=1000)
key: Final = scenario.key(user_id=user, models=[model])
gateway_warm: Final = _chat(gateway, model, key)
peer_warm: Final = _chat(peer, model, key)
assert gateway_warm.status_code == 200, (
f"Gateway rejected the initial user-limited request: {gateway_warm.text}"
)
assert peer_warm.status_code == 200, f"Peer rejected the initial user-limited request: {peer_warm.text}"
assert peer_warm.headers.get("x-ratelimit-user-limit-requests") == "1000", peer_warm.headers
assert peer_warm.headers.get("x-ratelimit-user-limit-tokens") == "100000", peer_warm.headers
gateway.post("/user/update", {"user_id": user, field: 1})
expected_tpm: Final = 1 if field == "tpm_limit" else 100000
expected_rpm: Final = 1 if field == "rpm_limit" else 1000
rows: Final = read_rows(
'SELECT tpm_limit, rpm_limit FROM "LiteLLM_UserTable" WHERE user_id = %s',
(user,),
)
assert rows == [{"tpm_limit": expected_tpm, "rpm_limit": expected_rpm}], (
f"User {field} update did not persist without changing the other limit: {rows!r}"
)
limit_type: Final = "tokens" if field == "tpm_limit" else "requests"
peer_limited: Final = eventually(
lambda: _chat(peer, model, key),
lambda response: response.status_code == 429,
seconds=10,
return_last_on_timeout=True,
)
_assert_user_rate_limit_error(peer_limited, user, limit_type)
def test_user_rate_limit_explicit_null_clears_and_omitted_limit_is_untouched(gateway: Gateway, peer: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
user: Final = scenario.user(tpm_limit=100000, rpm_limit=1)
key: Final = scenario.key(user_id=user, models=[model])
gateway_warm: Final = _chat(gateway, model, key)
assert gateway_warm.status_code == 200, f"Gateway rejected the initial request under RPM 1: {gateway_warm.text}"
peer_limited: Final = _chat(peer, model, key)
_assert_user_rate_limit_error(peer_limited, user, "requests")
gateway.post("/user/update", {"user_id": user, "rpm_limit": None})
cleared_rows: Final = read_rows(
'SELECT tpm_limit, rpm_limit FROM "LiteLLM_UserTable" WHERE user_id = %s',
(user,),
)
assert cleared_rows == [{"tpm_limit": 100000, "rpm_limit": None}], (
f"Clearing RPM changed the wrong user limits: {cleared_rows!r}"
)
info: Final = gateway.get("/v2/user/info", {"user_id": user})
assert info["tpm_limit"] == 100000, f"User info omitted or changed TPM after RPM clear: {info!r}"
assert info["rpm_limit"] is None, f"User info did not report the cleared RPM limit: {info!r}"
gateway_after_clear: Final = _chat(gateway, model, key)
assert gateway_after_clear.status_code == 200, (
f"Gateway still enforced RPM after it was cleared: {gateway_after_clear.text}"
)
peer_after_clear: Final = eventually(
lambda: _chat(peer, model, key),
lambda response: response.status_code == 200,
seconds=10,
)
assert peer_after_clear.status_code == 200, (
f"Peer did not stop enforcing RPM after it was cleared: {peer_after_clear.text}"
)
gateway.post("/user/update", {"user_id": user, "tpm_limit": 50000})
omitted_rows: Final = read_rows(
'SELECT tpm_limit, rpm_limit FROM "LiteLLM_UserTable" WHERE user_id = %s',
(user,),
)
assert omitted_rows == [{"tpm_limit": 50000, "rpm_limit": None}], (
f"Omitting RPM during the TPM update changed it: {omitted_rows!r}"
)
@pytest.mark.parametrize(
("route", "stream", "upstream_target", "expected_rpm_header"),
(
pytest.param("/v1/messages", False, "/v1/responses", "1000", id="messages"),
pytest.param("/v1/responses", False, "/v1/responses", "1000", id="responses"),
pytest.param("/v1/chat/completions", True, None, None, id="streaming-chat-completions"),
),
)
def test_user_rpm_lowered_on_gateway_is_enforced_by_peer_on_llm_route(
gateway: Gateway,
peer: Gateway,
route: str,
stream: bool,
upstream_target: str | None,
expected_rpm_header: str | None,
) -> None:
with gateway.scenario() as scenario, ExitStack() as resources:
def upstream(request: Request) -> Reply:
assert request.target == upstream_target, request.target
reply: Final = _UPSTREAM_REPLIES[request.target]
return Reply(body=json.dumps(reply).encode())
provider: Final = resources.enter_context(wire_server(upstream)) if upstream_target is not None else None
model: Final = (
scenario.model()
if provider is None
else scenario.model(model="openai/gpt-5.6", api_base=provider.url + "/v1")
)
user: Final = scenario.user(tpm_limit=100000, rpm_limit=1000)
key: Final = scenario.key(user_id=user, models=[model])
gateway_warm: Final = _route_request(gateway, route, model, key, stream)
peer_warm: Final = _route_request(peer, route, model, key, stream)
assert gateway_warm.status_code == 200, f"Gateway rejected {route}: {gateway_warm.text}"
assert peer_warm.status_code == 200, f"Peer rejected {route}: {peer_warm.text}"
assert expected_rpm_header is None or (
peer_warm.headers.get("x-ratelimit-user-limit-requests") == expected_rpm_header
), f"Peer returned unexpected user RPM headers for {route}: {dict(peer_warm.headers)!r}"
targets: Final = tuple(request.target for request in provider.drain()) if provider is not None else ()
expected_targets: Final = (upstream_target, upstream_target) if upstream_target is not None else ()
assert targets == expected_targets, targets
gateway.post("/user/update", {"user_id": user, "rpm_limit": 1})
rows: Final = read_rows(
'SELECT tpm_limit, rpm_limit FROM "LiteLLM_UserTable" WHERE user_id = %s',
(user,),
)
assert rows == [{"tpm_limit": 100000, "rpm_limit": 1}], f"User RPM update changed the wrong limits: {rows!r}"
peer_limited: Final = eventually(
lambda: _route_request(peer, route, model, key, stream),
lambda response: response.status_code == 429,
seconds=10,
return_last_on_timeout=True,
)
_assert_route_user_requests_limit_error(peer_limited, user, route)
def test_internal_user_cannot_clear_own_rpm_limit(gateway: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
user: Final = scenario.user(tpm_limit=100000, rpm_limit=1, user_role="internal_user")
key: Final = scenario.key(user_id=user, models=[model])
denied: Final = gateway.request(
"POST",
"/user/update",
{"user_id": user, "rpm_limit": None},
key=key,
)
context: Final = f"Expected internal-user route denial, received HTTP {denied.status_code}: {denied.text}"
assert denied.status_code == 401, context
assert "Only proxy admin can be used to generate" in denied.text, context
assert "Route=/user/update" in denied.text, context
rows: Final = read_rows(
'SELECT tpm_limit, rpm_limit FROM "LiteLLM_UserTable" WHERE user_id = %s',
(user,),
)
assert rows == [{"tpm_limit": 100000, "rpm_limit": 1}], f"Denied self-update changed user limits: {rows!r}"
first_chat: Final = _chat(gateway, model, key)
assert first_chat.status_code == 200, f"Internal-user first chat was rejected: {first_chat.text}"
second_chat: Final = _chat(gateway, model, key)
_assert_user_rate_limit_error(second_chat, user, "requests")
def test_bulk_update_lowered_rpm_is_enforced_on_every_proxy(gateway: Gateway, peer: Gateway) -> None:
with gateway.scenario() as scenario:
model: Final = scenario.model()
first_user: Final = scenario.user(tpm_limit=100000, rpm_limit=1000)
first_key: Final = scenario.key(user_id=first_user, models=[model])
second_user: Final = scenario.user(tpm_limit=100000, rpm_limit=1000)
second_key: Final = scenario.key(user_id=second_user, models=[model])
warm_responses: Final = (
_chat(gateway, model, first_key),
_chat(peer, model, first_key),
_chat(gateway, model, second_key),
_chat(peer, model, second_key),
)
assert tuple(response.status_code for response in warm_responses) == (200, 200, 200, 200), (
f"Expected both users to warm on gateway and peer: {tuple(response.text for response in warm_responses)!r}"
)
bulk_update: Final = gateway.post(
"/user/bulk_update",
{
"users": [
{"user_id": first_user, "rpm_limit": 1},
{"user_id": second_user, "rpm_limit": 1},
]
},
)
assert (
bulk_update["total_requested"],
bulk_update["successful_updates"],
bulk_update["failed_updates"],
) == (2, 2, 0), bulk_update
results_json: Final = bulk_update.get("results")
assert isinstance(results_json, list), bulk_update
results: Final = tuple(object_value(result) for result in results_json)
assert tuple((string_value(result["user_id"]), result["success"]) for result in results) == (
(first_user, True),
(second_user, True),
), results
rows: Final = read_rows(
'SELECT user_id, tpm_limit, rpm_limit FROM "LiteLLM_UserTable" WHERE user_id IN (%s, %s) ORDER BY user_id',
(first_user, second_user),
)
expected_rows: Final = tuple(
{"user_id": user, "tpm_limit": 100000, "rpm_limit": 1} for user in sorted((first_user, second_user))
)
assert tuple(rows) == expected_rows, f"Bulk RPM update changed unexpected limits: {rows!r}"
_assert_user_rate_limit_on_every_proxy(gateway, peer, model, first_user, first_key)
_assert_user_rate_limit_on_every_proxy(gateway, peer, model, second_user, second_key)

View file

@ -5,7 +5,7 @@ import logging
from collections.abc import Mapping, Sequence
from datetime import datetime, timezone
from types import SimpleNamespace
from typing import Final
from typing import TYPE_CHECKING, Final, cast
from unittest.mock import AsyncMock, MagicMock
import httpx
@ -42,6 +42,10 @@ from tests.unit.proxy.management_endpoints.jwt_key_mapping_doubles import (
client = TestClient(app)
if TYPE_CHECKING:
from litellm.caching.redis_cache import RedisCache
from litellm.proxy.utils import PrismaClient
@pytest.mark.asyncio
async def test_ui_view_users_with_null_email(mocker, caplog):
@ -2109,6 +2113,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
@ -2337,6 +2353,180 @@ async def test_user_max_budget_update_evicts_cached_user_on_every_worker(mocker:
broadcast.assert_awaited_once_with(cache_key=saved_user.user_id)
@pytest.mark.asyncio
@pytest.mark.parametrize(
("field", "new_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 | None, all_users: bool
) -> None:
from litellm.proxy._types import LiteLLM_UserTable
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 bulk_user_update, user_update
from litellm.types.proxy.management_endpoints.internal_user_endpoints import (
BulkUpdateUserRequest,
UpdateUserRequestNoUserIDorEmail,
)
published: Final[list[tuple[str, str]]] = [] # mutable-ok: captures messages from the async Redis publisher
class _RecordingRedisClient:
async def publish(self, channel: str, message: str) -> int:
published.append((channel, message))
return 1
class _FakeRedisCache:
namespace: str | None = None
def init_pubsub_client(self) -> _RecordingRedisClient:
return _RecordingRedisClient()
class _UserTableMocks:
def __init__(
self,
find_first: AsyncMock,
find_unique: AsyncMock,
find_many: AsyncMock,
update_many: AsyncMock,
) -> None:
self.find_first = find_first
self.find_unique = find_unique
self.find_many = find_many
self.update_many = update_many
class _DatabaseMocks:
def __init__(self, litellm_usertable: _UserTableMocks) -> None:
self.litellm_usertable = litellm_usertable
class _PrismaClientMock:
def __init__(
self,
db: _DatabaseMocks,
get_data: AsyncMock,
updated_user: LiteLLM_UserTable,
) -> None:
self.db = db
self.get_data = get_data
self.updated_user = updated_user
self.update_data_payload: dict[str, object] | None = None
async def update_data(self, user_id: str, data: dict[str, object], table_name: str) -> dict[str, object]:
self.update_data_payload = data
return {"user_id": user_id, "data": self.updated_user}
saved_user: Final = LiteLLM_UserTable(user_id="user-spruce", tpm_limit=100000, rpm_limit=1000)
updated_user: Final = saved_user.model_copy(update={field: new_limit})
old_limit: Final = 100000 if field == "tpm_limit" else 1000
prisma_client: Final = _PrismaClientMock(
db=_DatabaseMocks(
litellm_usertable=_UserTableMocks(
find_first=mocker.AsyncMock(return_value=saved_user),
find_unique=mocker.AsyncMock(return_value=updated_user),
find_many=mocker.AsyncMock(return_value=[saved_user]),
update_many=mocker.AsyncMock(return_value=1),
)
),
get_data=mocker.AsyncMock(return_value=saved_user),
updated_user=updated_user,
)
prisma_client_for_auth: Final = cast("PrismaClient", prisma_client)
mocker.patch( # test-quality-ok: substitute the database dependency
"litellm.proxy.proxy_server.prisma_client", prisma_client
)
handling_worker_cache: Final = UserApiKeyCache()
other_worker_cache: Final = UserApiKeyCache()
await handling_worker_cache.async_set_cache(
key=saved_user.user_id,
value=saved_user,
model_type=LiteLLM_UserTable,
)
await other_worker_cache.async_set_cache(
key=saved_user.user_id,
value=saved_user,
model_type=LiteLLM_UserTable,
)
mocker.patch( # test-quality-ok: exercise a real isolated cache for the endpoint's worker
"litellm.proxy.proxy_server.user_api_key_cache", handling_worker_cache
)
mocker.patch( # test-quality-ok: inject an in-memory pub/sub client without live Redis
"litellm.proxy.common_utils.auth_cache_invalidation_pubsub.coordination_redis_cache",
return_value=_FakeRedisCache(),
)
handling_user_before: Final = await get_user_object(
user_id=saved_user.user_id,
prisma_client=prisma_client_for_auth,
user_api_key_cache=handling_worker_cache,
user_id_upsert=False,
)
other_user_before: Final = await get_user_object(
user_id=saved_user.user_id,
prisma_client=prisma_client_for_auth,
user_api_key_cache=other_worker_cache,
user_id_upsert=False,
)
assert handling_user_before is not None
assert handling_user_before.model_dump()[field] == old_limit
assert other_user_before is not None
assert other_user_before.model_dump()[field] == old_limit
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=UpdateUserRequestNoUserIDorEmail.model_validate({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.model_validate({"user_id": saved_user.user_id, field: new_limit}),
user_api_key_dict=admin,
)
assert prisma_client.update_data_payload is not None
assert prisma_client.update_data_payload[field] == new_limit
await asyncio.sleep(0)
await asyncio.sleep(0)
remote_subscriber: Final = AuthCacheInvalidationSubscriber(
redis_cache=cast("RedisCache", _FakeRedisCache()),
user_api_key_cache=other_worker_cache,
)
for _, message in published:
remote_subscriber._apply_message( # pyright: ignore[reportPrivateUsage] # exercising the real cross-worker message handler, not a public API
{"type": "message", "data": message}
)
handling_user_after: Final = await get_user_object(
user_id=saved_user.user_id,
prisma_client=prisma_client_for_auth,
user_api_key_cache=handling_worker_cache,
user_id_upsert=False,
)
other_user_after: Final = await get_user_object(
user_id=saved_user.user_id,
prisma_client=prisma_client_for_auth,
user_api_key_cache=other_worker_cache,
user_id_upsert=False,
)
assert handling_user_after is not None
assert handling_user_after.model_dump()[field] == new_limit
assert other_user_after is not None
assert other_user_after.model_dump()[field] == new_limit, (
"another worker still enforces the old limit; the update was never broadcast"
)
def test_generate_request_base_validator():
"""
Test that GenerateRequestBase validator converts empty string to None for max_budget
@ -2747,49 +2937,65 @@ async def test_user_update_rejects_silent_create_for_non_proxy_admin(mocker):
@pytest.mark.asyncio
async def test_user_info_v2_proxy_admin_can_query_any_user(mocker):
async def test_user_info_v2_proxy_admin_can_query_any_user(mocker: MockerFixture) -> None:
"""
Test that proxy admin can query any user via /v2/user/info.
"""
from fastapi import Request
from litellm.proxy._types import UserInfoV2Response
from litellm.proxy._types import LiteLLM_UserTable, UserInfoV2Response
from litellm.proxy.management_endpoints.internal_user_endpoints import user_info_v2
mock_prisma_client = mocker.MagicMock()
mock_user_row: Final = LiteLLM_UserTable(
user_id="target-user-123",
user_email="target@example.com",
user_alias="Target User",
user_role="internal_user",
spend=42.5,
max_budget=100.0,
tpm_limit=100000,
rpm_limit=1000,
models=["gpt-4"],
budget_duration="30d",
budget_reset_at=None,
metadata={"team": "engineering"},
created_at=datetime(2024, 1, 1, tzinfo=timezone.utc),
updated_at=datetime(2024, 6, 1, tzinfo=timezone.utc),
sso_user_id="sso-abc",
teams=["team-1", "team-2"],
)
mock_user_row = mocker.MagicMock()
mock_user_row.model_dump.return_value = {
"user_id": "target-user-123",
"user_email": "target@example.com",
"user_alias": "Target User",
"user_role": "internal_user",
"spend": 42.5,
"max_budget": 100.0,
"models": ["gpt-4"],
"budget_duration": "30d",
"budget_reset_at": None,
"metadata": {"team": "engineering"},
"created_at": datetime(2024, 1, 1, tzinfo=timezone.utc),
"updated_at": datetime(2024, 6, 1, tzinfo=timezone.utc),
"sso_user_id": "sso-abc",
"teams": ["team-1", "team-2"],
}
class _UserTable:
def __init__(self, find_unique: AsyncMock) -> None:
self.find_unique = find_unique
async def mock_find_unique(*args, **kwargs):
if kwargs.get("where", {}).get("user_id") == "target-user-123":
class _Database:
def __init__(self, litellm_usertable: _UserTable) -> None:
self.litellm_usertable = litellm_usertable
class _PrismaClient:
def __init__(self, db: _Database) -> None:
self.db = db
async def mock_find_unique(*_args: object, **kwargs: object) -> LiteLLM_UserTable | None:
where: Final = kwargs.get("where")
if isinstance(where, Mapping) and where.get("user_id") == "target-user-123":
return mock_user_row
return None
mock_prisma_client.db.litellm_usertable.find_unique = mocker.AsyncMock(side_effect=mock_find_unique)
mock_prisma_client: Final = _PrismaClient(
db=_Database(
litellm_usertable=_UserTable(find_unique=mocker.AsyncMock(side_effect=mock_find_unique))
)
)
mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
mock_request = mocker.MagicMock(spec=Request)
mock_request: Final = mocker.MagicMock(spec=Request)
admin_key = UserAPIKeyAuth(user_id="admin-user", user_role=LitellmUserRoles.PROXY_ADMIN)
admin_key: Final = UserAPIKeyAuth(user_id="admin-user", user_role=LitellmUserRoles.PROXY_ADMIN)
response = await user_info_v2(
response: Final = await user_info_v2(
request=mock_request,
user_id="target-user-123",
user_api_key_dict=admin_key,
@ -2802,6 +3008,8 @@ async def test_user_info_v2_proxy_admin_can_query_any_user(mocker):
assert response.user_role == "internal_user"
assert response.spend == 42.5
assert response.max_budget == 100.0
assert response.tpm_limit == 100000
assert response.rpm_limit == 1000
assert response.models == ["gpt-4"]
assert response.teams == ["team-1", "team-2"]
assert response.sso_user_id == "sso-abc"
@ -3132,6 +3340,8 @@ async def test_user_info_v2_response_shape(mocker):
"user_role",
"spend",
"max_budget",
"tpm_limit",
"rpm_limit",
"models",
"budget_duration",
"budget_reset_at",

View file

@ -0,0 +1,47 @@
import { describe, expect, it } from "vitest";
import { isValidRateLimitInput, rateLimitUpdate } from "./userRateLimitPayload";
describe("rateLimitUpdate", () => {
it("omits an untouched input when the stored value is null", () => {
expect(rateLimitUpdate(undefined, null)).toBeUndefined();
});
it.each([500, "500"])("omits an unchanged stored limit from input %s", (input) => {
expect(rateLimitUpdate(input, 500)).toBeUndefined();
});
it.each(["", null])("sends null when a stored limit is deliberately cleared with %s", (input) => {
expect(rateLimitUpdate(input, 500)).toBeNull();
});
it("sends a new limit entered as a string", () => {
expect(rateLimitUpdate("100", null)).toBe(100);
});
it("preserves zero as a changed limit", () => {
expect(rateLimitUpdate("0", 500)).toBe(0);
});
it("omits a whitespace-only input when no limit was stored", () => {
expect(rateLimitUpdate(" ", null)).toBeUndefined();
});
});
describe("isValidRateLimitInput", () => {
it.each([
["empty string", ""],
["null", null],
["undefined", undefined],
["whitespace", " "],
["string zero", "0"],
["number zero", 0],
["integer string", "12"],
["integer", 12],
])("accepts %s", (_label, value) => {
expect(isValidRateLimitInput(value)).toBe(true);
});
it.each(["1.5", "-1", "abc", "1e400"])("rejects %s", (value) => {
expect(isValidRateLimitInput(value)).toBe(false);
});
});

View file

@ -0,0 +1,17 @@
export const rateLimitUpdate = (
input: string | number | null | undefined,
stored: number | null | undefined,
): number | null | undefined => {
const isNullishInput = input === null || input === undefined;
const isBlankString = typeof input === "string" && input.trim() === "";
const normalized = isNullishInput || isBlankString ? null : Number(input);
return normalized === (stored ?? null) ? undefined : normalized;
};
export const isValidRateLimitInput = (value: string | number | null | undefined): boolean => {
if (value === "" || value === null || value === undefined) {
return true;
}
const number = Number(value);
return Number.isFinite(number) && Number.isInteger(number) && number >= 0;
};

View file

@ -457,6 +457,121 @@ describe("UserEditView", () => {
expect(checkbox).toBeChecked();
});
});
describe("user rate limits", () => {
const userDataWithRateLimits = () => ({
...MOCK_USER_DATA,
user_info: {
...MOCK_USER_DATA.user_info,
tpm_limit: 100000,
rpm_limit: 50,
},
});
it("seeds the TPM and RPM inputs from the selected user", async () => {
renderWithProviders(<UserEditView {...defaultProps} userData={userDataWithRateLimits()} />);
expect(await screen.findByRole("spinbutton", { name: /tpm limit/i })).toHaveValue(100000);
expect(await screen.findByRole("spinbutton", { name: /rpm limit/i })).toHaveValue(50);
});
it("keeps unset rate limits empty and omits them from an untouched save", async () => {
const onSubmit = vi.fn();
const userDataWithNullRateLimits = {
...MOCK_USER_DATA,
user_info: {
...MOCK_USER_DATA.user_info,
tpm_limit: null,
rpm_limit: null,
},
};
renderWithProviders(<UserEditView {...defaultProps} userData={userDataWithNullRateLimits} onSubmit={onSubmit} />);
expect(await screen.findByRole("spinbutton", { name: /tpm limit/i })).toHaveValue(null);
expect(await screen.findByRole("spinbutton", { name: /rpm limit/i })).toHaveValue(null);
await userEvent.click(screen.getByRole("button", { name: /save changes/i }));
await waitFor(() => {
expect(onSubmit).toHaveBeenCalled();
});
expect(onSubmit.mock.calls[0][0]).not.toHaveProperty("tpm_limit");
expect(onSubmit.mock.calls[0][0]).not.toHaveProperty("rpm_limit");
});
it("omits unchanged rate limits from the submit payload", async () => {
const onSubmit = vi.fn();
renderWithProviders(<UserEditView {...defaultProps} userData={userDataWithRateLimits()} onSubmit={onSubmit} />);
await userEvent.click(await screen.findByRole("button", { name: /save changes/i }));
await waitFor(() => {
expect(onSubmit).toHaveBeenCalled();
});
expect(onSubmit.mock.calls[0][0]).not.toHaveProperty("tpm_limit");
expect(onSubmit.mock.calls[0][0]).not.toHaveProperty("rpm_limit");
});
it("sends null only for a deliberately cleared TPM limit", async () => {
const onSubmit = vi.fn();
renderWithProviders(<UserEditView {...defaultProps} userData={userDataWithRateLimits()} onSubmit={onSubmit} />);
fireEvent.change(await screen.findByRole("spinbutton", { name: /tpm limit/i }), {
target: { value: "" },
});
await userEvent.click(screen.getByRole("button", { name: /save changes/i }));
await waitFor(() => {
expect(onSubmit).toHaveBeenCalled();
});
expect(onSubmit.mock.calls[0][0].tpm_limit).toBeNull();
expect(onSubmit.mock.calls[0][0]).not.toHaveProperty("rpm_limit");
});
it("submits a new RPM limit as a number", async () => {
const onSubmit = vi.fn();
renderWithProviders(<UserEditView {...defaultProps} userData={userDataWithRateLimits()} onSubmit={onSubmit} />);
fireEvent.change(await screen.findByRole("spinbutton", { name: /rpm limit/i }), {
target: { value: "1" },
});
await userEvent.click(screen.getByRole("button", { name: /save changes/i }));
await waitFor(() => {
expect(onSubmit).toHaveBeenCalled();
});
expect(onSubmit.mock.calls[0][0].rpm_limit).toBe(1);
expect(typeof onSubmit.mock.calls[0][0].rpm_limit).toBe("number");
});
it.each(["-1", "1.5"])("rejects an invalid TPM limit of %s", async (value) => {
const onSubmit = vi.fn();
renderWithProviders(<UserEditView {...defaultProps} userData={userDataWithRateLimits()} onSubmit={onSubmit} />);
fireEvent.change(await screen.findByRole("spinbutton", { name: /tpm limit/i }), {
target: { value },
});
const submitButton = screen.getByRole("button", { name: /save changes/i }) as HTMLButtonElement;
const form = submitButton.form;
if (!form) {
throw new Error("User edit form was not rendered");
}
fireEvent.submit(form);
expect(
await screen.findByText("Enter a non-negative whole number, or leave empty for unlimited"),
).toBeInTheDocument();
expect(onSubmit).not.toHaveBeenCalled();
});
it("hides both rate-limit inputs in bulk edit mode", async () => {
renderWithProviders(<UserEditView {...defaultProps} userData={userDataWithRateLimits()} isBulkEdit={true} />);
await screen.findByRole("button", { name: /save changes/i });
expect(screen.queryByRole("spinbutton", { name: /tpm limit/i })).not.toBeInTheDocument();
expect(screen.queryByRole("spinbutton", { name: /rpm limit/i })).not.toBeInTheDocument();
});
});
describe("submit payload parity", () => {
const submittedPayload = async (props: Partial<Parameters<typeof UserEditView>[0]> = {}) => {
const onSubmit = vi.fn();

View file

@ -5,6 +5,7 @@ import BudgetDurationDropdown from "@/components/common_components/budget_durati
import { ModelMaxBudget, ModelMaxBudgetField } from "@/components/key_team_helpers/ModelMaxBudgetEditor";
import { modelMaxBudgetUpdate } from "@/components/key_team_helpers/modelMaxBudgetPayload";
import { useSeededState } from "@/components/key_team_helpers/useSeededState";
import { isValidRateLimitInput, rateLimitUpdate } from "./userRateLimitPayload";
import { getModelDisplayName } from "@/components/key_team_helpers/fetch_available_models_team_key";
import MCPServerSelector from "@/components/mcp_server_management/MCPServerSelector";
import MCPToolPermissions from "@/components/mcp_server_management/MCPToolPermissions";
@ -67,6 +68,14 @@ const budgetSchema = (unlimitedBudget: boolean) =>
(value) => unlimitedBudget || (value !== "" && value !== null && value !== undefined),
"Please enter a budget or select Unlimited Budget",
),
tpm_limit: z
.union([z.string(), z.number()])
.nullish()
.refine(isValidRateLimitInput, "Enter a non-negative whole number, or leave empty for unlimited"),
rpm_limit: z
.union([z.string(), z.number()])
.nullish()
.refine(isValidRateLimitInput, "Enter a non-negative whole number, or leave empty for unlimited"),
});
type UserEditFormValues = z.infer<ReturnType<typeof budgetSchema>>;
@ -92,7 +101,14 @@ const toFormValues = (
const maxBudget = userData.user_info?.max_budget;
const isUnlimited = maxBudget === null || maxBudget === undefined;
return {
...(isBulkEdit ? {} : { user_id: userData.user_id, user_email: userData.user_info?.user_email }),
...(isBulkEdit
? {}
: {
user_id: userData.user_id,
user_email: userData.user_info?.user_email,
tpm_limit: userData.user_info?.tpm_limit ?? "",
rpm_limit: userData.user_info?.rpm_limit ?? "",
}),
user_alias: userData.user_info?.user_alias,
user_role: userData.user_info?.user_role,
models: userData.user_info?.models || [],
@ -171,11 +187,16 @@ export function UserEditView({
return;
}
const { tpm_limit: tpmLimitInput, rpm_limit: rpmLimitInput, ...formValues } = values;
const modelBudgets = modelMaxBudgetUpdate(modelMaxBudget, userData.user_info?.model_max_budget);
const tpmLimit = rateLimitUpdate(tpmLimitInput, isBulkEdit ? undefined : userData.user_info?.tpm_limit);
const rpmLimit = rateLimitUpdate(rpmLimitInput, isBulkEdit ? undefined : userData.user_info?.rpm_limit);
onSubmit({
...values,
...formValues,
...("metadata" in values ? { metadata: metadata.value } : {}),
...(modelBudgets !== undefined && { model_max_budget: modelBudgets }),
...(tpmLimit !== undefined && { tpm_limit: tpmLimit }),
...(rpmLimit !== undefined && { rpm_limit: rpmLimit }),
max_budget:
unlimitedBudget || values.max_budget === "" || values.max_budget === undefined ? null : values.max_budget,
});
@ -293,6 +314,56 @@ export function UserEditView({
{({ id, value, onChange }) => <BudgetDurationDropdown id={id} value={value} onChange={onChange} />}
</FormField>
{!isBulkEdit && (
<>
<FormField
control={form.control}
name="tpm_limit"
label={labelWithHint(
"TPM Limit",
"Applies across all keys owned by this user. Team and key limits still apply as ceilings.",
)}
>
{({ ref, value, onChange, ...control }) => (
<Input
{...control}
ref={ref}
type="number"
min={0}
step={1}
value={value ?? ""}
onChange={(event) => onChange(event.target.value)}
onWheel={(event) => event.currentTarget.blur()}
placeholder="Unlimited"
/>
)}
</FormField>
<FormField
control={form.control}
name="rpm_limit"
label={labelWithHint(
"RPM Limit",
"Applies across all keys owned by this user. Team and key limits still apply as ceilings.",
)}
>
{({ ref, value, onChange, ...control }) => (
<Input
{...control}
ref={ref}
type="number"
min={0}
step={1}
value={value ?? ""}
onChange={(event) => onChange(event.target.value)}
onWheel={(event) => event.currentTarget.blur()}
placeholder="Unlimited"
/>
)}
</FormField>
</>
)}
{/* Bulk edit forwards a fixed field list and has no single stored budget to
diff against, so the editor would silently discard whatever was typed. */}
{!isBulkEdit && (

View file

@ -1,4 +1,4 @@
import { render, screen, waitFor, within } from "@testing-library/react";
import { fireEvent, render, screen, waitFor, within } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { describe, expect, it, vi, beforeEach } from "vitest";
import UserInfoView from "./user_info_view";
@ -130,6 +130,42 @@ describe("UserInfoView", () => {
expect(aliases.length).toBeGreaterThan(0);
});
it("seeds the user rate limits when opening the edit form", async () => {
mockUserGetInfoV2.mockResolvedValue({
...MOCK_USER_DATA,
tpm_limit: 100000,
rpm_limit: 50,
});
render(<UserInfoView {...defaultProps} userRole="proxy_admin" initialTab={1} startInEditMode />);
expect(await screen.findByRole("spinbutton", { name: /tpm limit/i })).toHaveValue(100000);
expect(await screen.findByRole("spinbutton", { name: /rpm limit/i })).toHaveValue(50);
});
it("keeps the updated TPM and stored RPM when reopening the edit form", async () => {
mockUserGetInfoV2.mockResolvedValue({
...MOCK_USER_DATA,
tpm_limit: 100000,
rpm_limit: 50,
});
render(<UserInfoView {...defaultProps} userRole="proxy_admin" initialTab={1} startInEditMode />);
fireEvent.change(await screen.findByRole("spinbutton", { name: /tpm limit/i }), {
target: { value: "" },
});
await userEvent.click(screen.getByRole("button", { name: /save changes/i }));
await waitFor(() => {
expect(mockUserUpdateUserCall).toHaveBeenCalledTimes(1);
});
await userEvent.click(await screen.findByRole("button", { name: /edit settings/i }));
expect(await screen.findByRole("spinbutton", { name: /tpm limit/i })).toHaveValue(null);
expect(await screen.findByRole("spinbutton", { name: /rpm limit/i })).toHaveValue(50);
});
it("should render overview spend and budget with two decimal places", async () => {
mockUserGetInfoV2.mockResolvedValue({
...MOCK_USER_DATA,

View file

@ -341,6 +341,8 @@ export default function UserInfoView({
user_alias: formValues.user_alias ?? userData.user_alias,
models: formValues.models ?? userData.models,
max_budget: formValues.max_budget === undefined ? userData.max_budget : formValues.max_budget,
tpm_limit: formValues.tpm_limit === undefined ? userData.tpm_limit : formValues.tpm_limit,
rpm_limit: formValues.rpm_limit === undefined ? userData.rpm_limit : formValues.rpm_limit,
budget_duration:
formValues.budget_duration === undefined ? userData.budget_duration : formValues.budget_duration,
metadata: formValues.metadata ?? userData.metadata,
@ -401,6 +403,8 @@ export default function UserInfoView({
user_role: userData.user_role,
models: userData.models,
max_budget: userData.max_budget,
tpm_limit: userData.tpm_limit,
rpm_limit: userData.rpm_limit,
budget_duration: userData.budget_duration,
metadata: userData.metadata,
// Without these the per-model budget editor mounts empty and a save

View file

@ -1082,6 +1082,8 @@ export interface UserInfoV2Response {
user_role: string | null;
spend: number;
max_budget: number | null;
tpm_limit?: number | null;
rpm_limit?: number | null;
models: string[];
budget_duration: string | null;
budget_reset_at: string | null;

View file

@ -49799,6 +49799,8 @@ export interface components {
*/
models: string[];
object_permission?: components["schemas"]["LiteLLM_ObjectPermissionTable"] | null;
/** Rpm Limit */
rpm_limit?: number | null;
/**
* Spend
* @default 0
@ -49811,6 +49813,8 @@ export interface components {
* @default []
*/
teams: string[];
/** Tpm Limit */
tpm_limit?: number | null;
/** Updated At */
updated_at?: string | null;
/** User Alias */