From b7e31638fdbdade189e931fccd5c9188f12f6976 Mon Sep 17 00:00:00 2001 From: Krrish Dholakia Date: Wed, 7 Aug 2024 18:50:40 -0700 Subject: [PATCH] fix(internal_user_endpoints.py): respect 'max_user_budget' for new internal user's --- .../internal_user_endpoints.py | 4 +++ litellm/tests/test_proxy_server.py | 36 +++++++++++++++++++ 2 files changed, 40 insertions(+) diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index b132761ae5d..de109db46ba 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -87,6 +87,10 @@ async def new_user( "user" # only create a user, don't create key if 'auto_create_key' set to False ) + if "max_budget" in data_json and data_json["max_budget"] is None: + if litellm.max_user_budget is not None: + data_json["max_budget"] = litellm.max_user_budget + response = await generate_key_helper_fn(request_type="user", **data_json) # Admin UI Logic diff --git a/litellm/tests/test_proxy_server.py b/litellm/tests/test_proxy_server.py index b0a972bddc8..35198af807e 100644 --- a/litellm/tests/test_proxy_server.py +++ b/litellm/tests/test_proxy_server.py @@ -800,3 +800,39 @@ async def test_get_team_redis(client_no_auth): pass mock_client.assert_called_once() + + +import random +import uuid +from unittest.mock import AsyncMock, MagicMock, patch + +from litellm.proxy._types import LitellmUserRoles, NewUserRequest, UserAPIKeyAuth +from litellm.proxy.management_endpoints.internal_user_endpoints import new_user +from litellm.tests.test_key_generate_prisma import prisma_client + + +@pytest.mark.asyncio +async def test_create_user_default_budget(prisma_client): + + setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client) + setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") + setattr(litellm, "max_user_budget", 10) + await litellm.proxy.proxy_server.prisma_client.connect() + user = f"ishaan {uuid.uuid4().hex}" + request = NewUserRequest(user_id=user) # create a key with no budget + with patch.object( + litellm.proxy.proxy_server.prisma_client, "insert_data", new=AsyncMock() + ) as mock_client: + await new_user( + request, + ) + + mock_client.assert_called() + + print(f"mock_client.call_args: {mock_client.call_args}") + print("mock_client.call_args.kwargs: {}".format(mock_client.call_args.kwargs)) + + assert ( + mock_client.call_args.kwargs["data"]["max_budget"] + == litellm.max_user_budget + )