mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
fix(proxy): evict cached user on every proxy for tpm/rpm updates and edit limits in the users UI (#44130)
* test(proxy): cover cross-worker cache eviction for user tpm/rpm limit updates Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * 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> * feat(proxy): return tpm_limit and rpm_limit from /v2/user/info Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * feat(ui): edit user tpm and rpm limits from the user edit form Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): clear user tpm_limit and rpm_limit when sent as null Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): cover user tpm/rpm limit updates across proxies Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): cover user rate-limit routes and bulk updates Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(ui): reject unsafe rate limit integers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(ui): cover user rate limit seed and saved state Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): type user endpoint test doubles Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * revert: drop rate limit upper bound Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(ui): validate user rate limits with zod and clear user edit lint warnings Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(ui): rename user edit schema and name its input and output types Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): drop redundant event loop yields in rate limit eviction test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: mrinal <mrinal@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
3ca3e1c686
commit
ec0af8d5f8
10 changed files with 915 additions and 49 deletions
|
|
@ -3645,6 +3645,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
|
||||
|
|
|
|||
|
|
@ -119,7 +119,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(
|
||||
|
|
@ -1158,6 +1158,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"),
|
||||
|
|
@ -1298,7 +1300,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":
|
||||
|
|
@ -1627,7 +1629,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,
|
||||
|
|
@ -1985,7 +1987,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(
|
||||
*(
|
||||
|
|
|
|||
381
tests/integration/management/test_user_rate_limit_updates.py
Normal file
381
tests/integration/management/test_user_rate_limit_updates.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
@ -47,6 +47,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):
|
||||
|
|
@ -2250,6 +2254,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
|
||||
|
|
@ -2478,6 +2494,178 @@ 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
|
||||
|
||||
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
|
||||
|
|
@ -2888,49 +3076,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,
|
||||
|
|
@ -2943,6 +3147,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"
|
||||
|
|
@ -3273,6 +3479,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",
|
||||
|
|
|
|||
|
|
@ -457,6 +457,151 @@ 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("submits zero when the stored TPM limit changes to zero", async () => {
|
||||
const onSubmit = vi.fn();
|
||||
renderWithProviders(<UserEditView {...defaultProps} userData={userDataWithRateLimits()} onSubmit={onSubmit} />);
|
||||
|
||||
fireEvent.change(await screen.findByRole("spinbutton", { name: /tpm limit/i }), {
|
||||
target: { value: "0" },
|
||||
});
|
||||
await userEvent.click(await screen.findByRole("button", { name: /save changes/i }));
|
||||
|
||||
await waitFor(() => {
|
||||
expect(onSubmit).toHaveBeenCalled();
|
||||
});
|
||||
expect(onSubmit.mock.calls[0][0].tpm_limit).toBe(0);
|
||||
});
|
||||
|
||||
it("omits the TPM limit when the stored value is re-entered", async () => {
|
||||
const onSubmit = vi.fn();
|
||||
renderWithProviders(<UserEditView {...defaultProps} userData={userDataWithRateLimits()} onSubmit={onSubmit} />);
|
||||
|
||||
fireEvent.change(await screen.findByRole("spinbutton", { name: /tpm limit/i }), {
|
||||
target: { value: "100000" },
|
||||
});
|
||||
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");
|
||||
});
|
||||
|
||||
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();
|
||||
|
|
@ -483,7 +628,7 @@ describe("UserEditView", () => {
|
|||
"user_id",
|
||||
"user_role",
|
||||
]);
|
||||
expect(payload).toStrictEqual({
|
||||
const expectedPayload = {
|
||||
user_id: "user-123",
|
||||
user_email: "test@example.com",
|
||||
user_alias: "Test User",
|
||||
|
|
@ -494,7 +639,8 @@ describe("UserEditView", () => {
|
|||
metadata: { key1: "value1", key2: "value2" },
|
||||
mcp_servers_and_groups: { servers: [], accessGroups: [], toolsets: [] },
|
||||
mcp_tool_permissions: {},
|
||||
});
|
||||
};
|
||||
expect(payload).toStrictEqual(expectedPayload);
|
||||
expect(typeof payload.max_budget).toBe("number");
|
||||
});
|
||||
|
||||
|
|
@ -567,23 +713,29 @@ describe("UserEditView", () => {
|
|||
await waitFor(() => {
|
||||
expect(onSubmit).toHaveBeenCalled();
|
||||
});
|
||||
expect(onSubmit.mock.calls[0][0]).toMatchObject({
|
||||
const expectedPayload = {
|
||||
user_id: "user-null",
|
||||
user_email: "null@example.com",
|
||||
user_alias: null,
|
||||
user_role: null,
|
||||
budget_duration: null,
|
||||
max_budget: null,
|
||||
});
|
||||
};
|
||||
expect(onSubmit.mock.calls[0][0]).toMatchObject(expectedPayload);
|
||||
});
|
||||
|
||||
it("should keep the budget input's native step constraint armed", async () => {
|
||||
renderWithProviders(<UserEditView {...defaultProps} />);
|
||||
|
||||
const budgetInput = await screen.findByRole("spinbutton", { name: /max budget/i });
|
||||
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");
|
||||
}
|
||||
expect(budgetInput).toHaveAttribute("step", "0.01");
|
||||
expect(budgetInput).not.toHaveAttribute("min");
|
||||
expect(budgetInput.closest("form")).not.toHaveAttribute("novalidate");
|
||||
expect(form).not.toHaveAttribute("novalidate");
|
||||
});
|
||||
|
||||
it("shows the tool matrix for servers the user reaches only through an access group or toolset", async () => {
|
||||
|
|
|
|||
|
|
@ -21,6 +21,15 @@ import { Tooltip, TooltipContent, TooltipProvider, TooltipTrigger } from "@/comp
|
|||
import { useZodForm } from "@/lib/forms/useZodForm";
|
||||
import { CircleHelp } from "lucide-react";
|
||||
|
||||
const RATE_LIMIT_ERROR = "Enter a non-negative whole number, or leave empty for unlimited";
|
||||
const isBlank = (value: string | number | null | undefined): boolean =>
|
||||
value === null || value === undefined || String(value).trim() === "";
|
||||
const rateLimitField = z
|
||||
.union([z.string(), z.number()])
|
||||
.nullish()
|
||||
.transform((value) => (isBlank(value) ? null : Number(value)))
|
||||
.pipe(z.number({ error: RATE_LIMIT_ERROR }).int(RATE_LIMIT_ERROR).nonnegative(RATE_LIMIT_ERROR).nullable());
|
||||
|
||||
interface UserEditViewProps {
|
||||
userData: any;
|
||||
onCancel: () => void;
|
||||
|
|
@ -53,23 +62,23 @@ const userEditShape = {
|
|||
models: z.array(z.string()),
|
||||
budget_duration: z.string().nullish(),
|
||||
metadata: z.string().nullish(),
|
||||
tpm_limit: rateLimitField,
|
||||
rpm_limit: rateLimitField,
|
||||
mcp_servers_and_groups: MCP_SELECTION_SHAPE.optional(),
|
||||
mcp_tool_permissions: z.record(z.string(), z.array(z.string())).optional(),
|
||||
};
|
||||
|
||||
const budgetSchema = (unlimitedBudget: boolean) =>
|
||||
const userEditSchema = (unlimitedBudget: boolean) =>
|
||||
z.object({
|
||||
...userEditShape,
|
||||
max_budget: z
|
||||
.union([z.string(), z.number()])
|
||||
.nullish()
|
||||
.refine(
|
||||
(value) => unlimitedBudget || (value !== "" && value !== null && value !== undefined),
|
||||
"Please enter a budget or select Unlimited Budget",
|
||||
),
|
||||
.refine((value) => unlimitedBudget || !isBlank(value), "Please enter a budget or select Unlimited Budget"),
|
||||
});
|
||||
|
||||
type UserEditFormValues = z.infer<ReturnType<typeof budgetSchema>>;
|
||||
type UserEditFormInput = z.input<ReturnType<typeof userEditSchema>>;
|
||||
type UserEditFormValues = z.output<ReturnType<typeof userEditSchema>>;
|
||||
|
||||
const buildMcpFieldValues = (objectPermission: ObjectPermission | null | undefined) => ({
|
||||
mcp_servers_and_groups: {
|
||||
|
|
@ -88,11 +97,18 @@ const toFormValues = (
|
|||
objectPermission: ObjectPermission | null | undefined,
|
||||
isBulkEdit: boolean,
|
||||
canEditMcpPermissions: boolean,
|
||||
): UserEditFormValues => {
|
||||
): UserEditFormInput => {
|
||||
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 || [],
|
||||
|
|
@ -117,6 +133,9 @@ const parseMetadata = (metadata: string | null | undefined): ParsedMetadata => {
|
|||
}
|
||||
};
|
||||
|
||||
const changedLimit = (value: number | null, stored: number | null | undefined): number | null | undefined =>
|
||||
value === (stored ?? null) ? undefined : value;
|
||||
|
||||
const labelWithHint = (label: string, hint: string): React.ReactNode => (
|
||||
<>
|
||||
{label}
|
||||
|
|
@ -147,7 +166,7 @@ export function UserEditView({
|
|||
userData.user_id,
|
||||
() => userData.user_info?.model_max_budget ?? {},
|
||||
);
|
||||
const schema = useMemo(() => budgetSchema(unlimitedBudget), [unlimitedBudget]);
|
||||
const schema = useMemo(() => userEditSchema(unlimitedBudget), [unlimitedBudget]);
|
||||
const form = useZodForm(schema, {
|
||||
defaultValues: toFormValues(userData, objectPermission, isBulkEdit, canEditMcpPermissions),
|
||||
});
|
||||
|
|
@ -171,14 +190,20 @@ export function UserEditView({
|
|||
return;
|
||||
}
|
||||
|
||||
const { tpm_limit: tpmLimitInput, rpm_limit: rpmLimitInput, ...formValues } = values;
|
||||
const modelBudgets = modelMaxBudgetUpdate(modelMaxBudget, userData.user_info?.model_max_budget);
|
||||
onSubmit({
|
||||
...values,
|
||||
const tpmLimit = changedLimit(tpmLimitInput, userData.user_info?.tpm_limit);
|
||||
const rpmLimit = changedLimit(rpmLimitInput, userData.user_info?.rpm_limit);
|
||||
const payload = {
|
||||
...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,
|
||||
});
|
||||
};
|
||||
onSubmit(payload);
|
||||
};
|
||||
|
||||
const modelOptions = [
|
||||
|
|
@ -293,6 +318,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 && (
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1091,6 +1091,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;
|
||||
|
|
|
|||
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
4
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -49790,6 +49790,8 @@ export interface components {
|
|||
*/
|
||||
models: string[];
|
||||
object_permission?: components["schemas"]["LiteLLM_ObjectPermissionTable"] | null;
|
||||
/** Rpm Limit */
|
||||
rpm_limit?: number | null;
|
||||
/**
|
||||
* Spend
|
||||
* @default 0
|
||||
|
|
@ -49802,6 +49804,8 @@ export interface components {
|
|||
* @default []
|
||||
*/
|
||||
teams: string[];
|
||||
/** Tpm Limit */
|
||||
tpm_limit?: number | null;
|
||||
/** Updated At */
|
||||
updated_at?: string | null;
|
||||
/** User Alias */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue