diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 3fa0ef51f5b..ec788471a43 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -1744,3 +1744,12 @@ class ProxyErrorTypes(str, enum.Enum): auth_error = "auth_error" internal_server_error = "internal_server_error" bad_request_error = "bad_request_error" + + +class SSOUserDefinedValues(TypedDict): + models: List[str] + user_id: str + user_email: Optional[str] + user_role: Optional[str] + max_budget: Optional[float] + budget_duration: Optional[str] diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 173b205fc30..22faec3be61 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -87,12 +87,16 @@ async def new_user( "user" # only create a user, don't create key if 'auto_create_key' set to False ) + is_internal_user = False + if data.user_role == LitellmUserRoles.INTERNAL_USER: + is_internal_user = True + if "max_budget" in data_json and data_json["max_budget"] is None: - if litellm.max_internal_user_budget is not None: + if is_internal_user and litellm.max_internal_user_budget is not None: data_json["max_budget"] = litellm.max_internal_user_budget if "budget_duration" in data_json and data_json["budget_duration"] is None: - if litellm.internal_user_budget_duration is not None: + if is_internal_user and litellm.internal_user_budget_duration is not None: data_json["budget_duration"] = litellm.internal_user_budget_duration response = await generate_key_helper_fn(request_type="user", **data_json) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index c9ea1f7a4be..419ac8f7d20 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -283,7 +283,6 @@ except Exception as e: pass server_root_path = os.getenv("SERVER_ROOT_PATH", "") -print("server root path: ", server_root_path) # noqa _license_check = LicenseCheck() premium_user: bool = _license_check.is_premium() ui_link = f"{server_root_path}/ui/" @@ -8629,8 +8628,13 @@ async def auth_callback(request: Request): _last_name = getattr(result, "last_name", "") or "" user_id = _first_name + _last_name + if user_email is not None and (user_id is None or len(user_id) == 0): + user_id = user_email + user_info = None user_id_models: List = [] + max_internal_user_budget = litellm.max_internal_user_budget + internal_user_budget_duration = litellm.internal_user_budget_duration # User might not be already created on first generation of key # But if it is, we want their models preferences @@ -8642,10 +8646,13 @@ async def auth_callback(request: Request): "spend": 0, "team_id": "litellm-dashboard", } - user_defined_values = { + user_defined_values: SSOUserDefinedValues = { "models": user_id_models, "user_id": user_id, "user_email": user_email, + "max_budget": max_internal_user_budget, + "user_role": None, + "budget_duration": internal_user_budget_duration, } _user_id_from_sso = user_id try: @@ -8657,10 +8664,16 @@ async def auth_callback(request: Request): ) if user_info is not None: user_defined_values = { - "models": getattr(user_info, "models", []), + "models": getattr(user_info, "models", user_id_models), "user_id": getattr(user_info, "user_id", user_id), "user_email": getattr(user_info, "user_id", user_email), "user_role": getattr(user_info, "user_role", None), + "max_budget": getattr( + user_info, "max_budget", max_internal_user_budget + ), + "budget_duration": getattr( + user_info, "budget_duration", internal_user_budget_duration + ), } user_role = getattr(user_info, "user_role", None) @@ -8674,6 +8687,12 @@ async def auth_callback(request: Request): "user_id": user_id, "user_email": getattr(user_info, "user_id", user_email), "user_role": getattr(user_info, "user_role", None), + "max_budget": getattr( + user_info, "max_budget", max_internal_user_budget + ), + "budget_duration": getattr( + user_info, "budget_duration", internal_user_budget_duration + ), } user_role = getattr(user_info, "user_role", None) @@ -8690,10 +8709,38 @@ async def auth_callback(request: Request): "user_email": litellm.default_user_params.get( "user_email", user_email ), + "user_role": litellm.default_user_params.get("user_role", None), + "max_budget": litellm.default_user_params.get( + "max_budget", max_internal_user_budget + ), + "budget_duration": litellm.default_user_params.get( + "budget_duration", internal_user_budget_duration + ), } + except Exception as e: pass + is_internal_user = False + if ( + user_defined_values["user_role"] is not None + and user_defined_values["user_role"] == LitellmUserRoles.INTERNAL_USER.value + ): + is_internal_user = True + if ( + is_internal_user is True + and user_defined_values["max_budget"] is None + and litellm.max_internal_user_budget is not None + ): + user_defined_values["max_budget"] = litellm.max_internal_user_budget + + if ( + is_internal_user is True + and user_defined_values["budget_duration"] is None + and litellm.internal_user_budget_duration is not None + ): + user_defined_values["budget_duration"] = litellm.internal_user_budget_duration + verbose_proxy_logger.info( f"user_defined_values for creating ui key: {user_defined_values}" ) diff --git a/litellm/tests/test_proxy_server.py b/litellm/tests/test_proxy_server.py index d0f17b16415..c910e786c88 100644 --- a/litellm/tests/test_proxy_server.py +++ b/litellm/tests/test_proxy_server.py @@ -811,15 +811,22 @@ from litellm.proxy.management_endpoints.internal_user_endpoints import new_user from litellm.tests.test_key_generate_prisma import prisma_client +@pytest.mark.parametrize( + "user_role", + [LitellmUserRoles.INTERNAL_USER.value, LitellmUserRoles.PROXY_ADMIN.value], +) @pytest.mark.asyncio -async def test_create_user_default_budget(prisma_client): +async def test_create_user_default_budget(prisma_client, user_role): setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client) setattr(litellm.proxy.proxy_server, "master_key", "sk-1234") setattr(litellm, "max_internal_user_budget", 10) + setattr(litellm, "internal_user_budget_duration", "5m") 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 + request = NewUserRequest( + user_id=user, user_role=user_role + ) # create a key with no budget with patch.object( litellm.proxy.proxy_server.prisma_client, "insert_data", new=AsyncMock() ) as mock_client: @@ -832,7 +839,16 @@ async def test_create_user_default_budget(prisma_client): 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_internal_user_budget - ) + if user_role == LitellmUserRoles.INTERNAL_USER.value: + assert ( + mock_client.call_args.kwargs["data"]["max_budget"] + == litellm.max_internal_user_budget + ) + assert ( + mock_client.call_args.kwargs["data"]["budget_duration"] + == litellm.internal_user_budget_duration + ) + + else: + assert mock_client.call_args.kwargs["data"]["max_budget"] is None + assert mock_client.call_args.kwargs["data"]["budget_duration"] is None diff --git a/ui/litellm-dashboard/src/components/user_dashboard.tsx b/ui/litellm-dashboard/src/components/user_dashboard.tsx index 61ec058475a..3d1d4ea6009 100644 --- a/ui/litellm-dashboard/src/components/user_dashboard.tsx +++ b/ui/litellm-dashboard/src/components/user_dashboard.tsx @@ -315,6 +315,7 @@ const UserDashboard: React.FC = ({ = ({ userID, userRole, accessToken, userSpend, selectedTeam }) => { +const ViewUserSpend: React.FC = ({ userID, userRole, accessToken, userSpend, userMaxBudget, selectedTeam }) => { console.log(`userSpend: ${userSpend}`) let [spend, setSpend] = useState(userSpend !== null ? userSpend : 0.0); const [maxBudget, setMaxBudget] = useState(selectedTeam ? selectedTeam.max_budget : null); - const team_budget = selectedTeam ? selectedTeam.max_budget : "unlimited"; - console.log(`maxBudget: ${maxBudget}, selectedTeam.max_budget: ${team_budget}, selectedTeam: ${JSON.stringify(selectedTeam)}`) + useEffect(() => { + console.log(`Updating team max budget for default team - userMaxBudget - ${userMaxBudget}, selectedTeam.team_alias - ${selectedTeam.team_alias}`) + if (selectedTeam) { + if (selectedTeam.team_alias === "Default Team") { + setMaxBudget(userMaxBudget); + } else { + setMaxBudget(selectedTeam.max_budget); + } + } + }, [selectedTeam, userMaxBudget]); + console.log(`maxBudget: ${maxBudget}, selectedTeam.max_budget: ${selectedTeam.max_budget}, selectedTeam: ${JSON.stringify(selectedTeam)}`) const [userModels, setUserModels] = useState([]); useEffect(() => { const fetchData = async () => { if (!accessToken || !userID || !userRole) { return; } - // if (userRole === "Admin" && userSpend == null) { - // try { - // const globalSpend = await getTotalSpendCall(accessToken); - // if (globalSpend) { - // if (globalSpend.spend) { - // setSpend(globalSpend.spend); - // } else { - // setSpend(0.0); - // } - // if (globalSpend.max_budget) { - // setMaxBudget(globalSpend.max_budget); - // } else { - // setMaxBudget(null); - // } - // } - // } catch (error) { - // console.error("Error fetching global spend data:", error); - // } - // } }; const fetchUserModels = async () => { try { @@ -104,11 +95,6 @@ const ViewUserSpend: React.FC = ({ userID, userRole, accessT setSpend(userSpend) } }, [userSpend]) - useEffect(() => { - if (selectedTeam && selectedTeam.max_budget !== maxBudget) { - setMaxBudget(selectedTeam.max_budget); - } - }, [selectedTeam, maxBudget]); // logic to decide what models to display let modelsToDisplay = [];