From c8bdc552fb0a8c2f379929ec7398241f6c2d6785 Mon Sep 17 00:00:00 2001 From: Krish Dholakia Date: Wed, 18 Jun 2025 21:30:59 -0700 Subject: [PATCH] Revert "Revert "UI - allow setting default team for new users (#11874)" (#11876)" (#11877) This reverts commit 8cfb3cfa94efdbb4e5a339135c7c1f8aa653699f. --- litellm/proxy/_new_secret_config.yaml | 4 - litellm/proxy/_types.py | 13 +- .../internal_user_endpoints.py | 181 +++++++++++++----- .../test_internal_user_endpoints.py | 125 ++++++++++++ .../src/components/SSOSettings.tsx | 167 +++++++++++++++- 5 files changed, 432 insertions(+), 58 deletions(-) diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 1f907035a3..e1cf34c924 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -111,7 +111,3 @@ litellm_settings: # supported_call_types: ["acompletion", "completion"] callbacks: ["prometheus", "langfuse"] - - - - diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 0516a0f0c8..f0c122690c 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -893,6 +893,12 @@ class LiteLLM_MCPServerTable(LiteLLMPydanticObjectBase): updated_by: Optional[str] = None +class NewUserRequestTeam(LiteLLMPydanticObjectBase): + team_id: str + max_budget_in_team: Optional[float] = None + user_role: Literal["user", "admin"] = "user" + + class NewUserRequest(GenerateRequestBase): max_budget: Optional[float] = None user_email: Optional[str] = None @@ -905,7 +911,7 @@ class NewUserRequest(GenerateRequestBase): LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, ] ] = None - teams: Optional[list] = None + teams: Optional[Union[List[str], List[NewUserRequestTeam]]] = None auto_create_key: bool = ( True # flag used for returning a key as part of the /user/new response ) @@ -3031,6 +3037,11 @@ class DefaultInternalUserParams(LiteLLMPydanticObjectBase): default=None, description="Default list of models that new users can access" ) + teams: Optional[Union[List[str], List[NewUserRequestTeam]]] = Field( + default=None, + description="Default teams for new users created", + ) + class BaseDailySpendTransaction(TypedDict): date: str diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 070fb6e200..415affb358 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -55,11 +55,13 @@ router = APIRouter() def _update_internal_new_user_params(data_json: dict, data: NewUserRequest) -> dict: if "user_id" in data_json and data_json["user_id"] is None: data_json["user_id"] = str(uuid.uuid4()) + auto_create_key = data_json.pop("auto_create_key", True) + if auto_create_key is False: - data_json[ - "table_name" - ] = "user" # only create a user, don't create key if 'auto_create_key' set to False + data_json["table_name"] = ( + "user" # only create a user, don't create key if 'auto_create_key' set to False + ) if litellm.default_internal_user_params: for key, value in litellm.default_internal_user_params.items(): @@ -91,6 +93,7 @@ def _update_internal_new_user_params(data_json: dict, data: NewUserRequest) -> d ): data_json["budget_duration"] = litellm.internal_user_budget_duration + data_json.pop("teams", None) # handled separately return data_json @@ -158,6 +161,85 @@ async def _add_user_to_organizations( await asyncio.gather(*tasks, return_exceptions=True) +async def _add_user_to_team( + user_id: str, + team_id: str, + user_api_key_dict: UserAPIKeyAuth, + user_email: Optional[str] = None, + max_budget_in_team: Optional[float] = None, + user_role: Literal["user", "admin"] = "user", +): + from litellm.proxy.management_endpoints.team_endpoints import team_member_add + + try: + await team_member_add( + data=TeamMemberAddRequest( + team_id=team_id, + member=Member( + user_id=user_id, + role=user_role, + user_email=user_email, + ), + max_budget_in_team=max_budget_in_team, + ), + user_api_key_dict=user_api_key_dict, + ) + except HTTPException as e: + if e.status_code == 400 and ( + "already exists" in str(e) or "doesn't exist" in str(e) + ): + verbose_proxy_logger.debug( + "litellm.proxy.management_endpoints.internal_user_endpoints.new_user(): User already exists in team - {}".format( + str(e) + ) + ) + else: + verbose_proxy_logger.debug( + "litellm.proxy.management_endpoints.internal_user_endpoints.new_user(): Exception occured - {}".format( + str(e) + ) + ) + except ProxyException as e: + if "already exists" in str(e) or "doesn't exist" in str(e): + verbose_proxy_logger.debug( + "litellm.proxy.management_endpoints.internal_user_endpoints.new_user(): User already exists in team - {}".format( + str(e) + ) + ) + elif ProxyErrorTypes.team_member_already_in_team in e.type: + verbose_proxy_logger.debug( + "litellm.proxy.management_endpoints.internal_user_endpoints.new_user(): User already exists in team - {}".format( + str(e) + ) + ) + else: + raise e + + +def check_if_default_team_set() -> Optional[Union[List[str], List[NewUserRequestTeam]]]: + if litellm.default_internal_user_params is None: + return None + teams = litellm.default_internal_user_params.get("teams") + if teams is not None: + if all(isinstance(team, str) for team in teams): + return teams + elif all(isinstance(team, dict) for team in teams): + return [ + NewUserRequestTeam( + team_id=team.get("team_id"), + max_budget_in_team=team.get("max_budget_in_team"), + user_role=team.get("user_role", "user"), + ) + for team in teams + ] + else: + verbose_proxy_logger.error( + "Invalid team type in default internal user params: %s", + teams, + ) + return None + + @router.post( "/user/new", tags=["Internal User management"], @@ -252,6 +334,9 @@ async def new_user( data_json = data.json() # type: ignore data_json = _update_internal_new_user_params(data_json, data) + teams = data.teams + if teams is None: + teams = check_if_default_team_set() organization_ids = cast( Optional[List[str]], data_json.pop("organizations", None) ) @@ -260,47 +345,41 @@ async def new_user( # Admin UI Logic # Add User to Team and Organization # if team_id passed add this user to the team - if data_json.get("team_id", None) is not None: - from litellm.proxy.management_endpoints.team_endpoints import ( - team_member_add, + _team_id = data_json.get("team_id", None) + if _team_id is not None: + await _add_user_to_team( + user_id=cast(str, response.get("user_id")), + team_id=_team_id, + user_api_key_dict=user_api_key_dict, + user_email=data.user_email, + max_budget_in_team=None, + user_role="user", ) + elif teams is not None: + tasks = [] + for team in teams: + max_budget_in_team: Optional[float] = None + user_role: Literal["user", "admin"] = "user" + if isinstance(team, str): + team_id = team + elif isinstance(team, NewUserRequestTeam): + team_id = team.team_id + max_budget_in_team = team.max_budget_in_team + user_role = team.user_role + else: + raise ValueError(f"Invalid team type: {type(team)}") - try: - await team_member_add( - data=TeamMemberAddRequest( - team_id=data_json.get("team_id", None), - member=Member( - user_id=data_json.get("user_id", None), - role="user", - user_email=data_json.get("user_email", None), - ), - ), - user_api_key_dict=user_api_key_dict, + tasks.append( + _add_user_to_team( + user_id=cast(str, response.get("user_id")), + team_id=team_id, + user_email=data.user_email, + user_api_key_dict=user_api_key_dict, + max_budget_in_team=max_budget_in_team, + user_role=user_role, + ) ) - except HTTPException as e: - if e.status_code == 400 and ( - "already exists" in str(e) or "doesn't exist" in str(e) - ): - verbose_proxy_logger.debug( - "litellm.proxy.management_endpoints.internal_user_endpoints.new_user(): User already exists in team - {}".format( - str(e) - ) - ) - else: - verbose_proxy_logger.debug( - "litellm.proxy.management_endpoints.internal_user_endpoints.new_user(): Exception occured - {}".format( - str(e) - ) - ) - except Exception as e: - if "already exists" in str(e) or "doesn't exist" in str(e): - verbose_proxy_logger.debug( - "litellm.proxy.management_endpoints.internal_user_endpoints.new_user(): User already exists in team - {}".format( - str(e) - ) - ) - else: - raise e + await asyncio.gather(*tasks, return_exceptions=True) user_id = cast(Optional[str], response.get("user_id", None)) @@ -676,9 +755,9 @@ def _update_internal_user_params(data_json: dict, data: UpdateUserRequest) -> di "budget_duration" not in non_default_values ): # applies internal user limits, if user role updated if is_internal_user and litellm.internal_user_budget_duration is not None: - non_default_values[ - "budget_duration" - ] = litellm.internal_user_budget_duration + non_default_values["budget_duration"] = ( + litellm.internal_user_budget_duration + ) from litellm.proxy.common_utils.timezone_utils import get_budget_reset_time non_default_values["budget_reset_at"] = get_budget_reset_time( @@ -1322,13 +1401,13 @@ async def ui_view_users( } # Query users with pagination and filters - users: Optional[ - List[BaseModel] - ] = await prisma_client.db.litellm_usertable.find_many( - where=where_conditions, - skip=skip, - take=page_size, - order={"created_at": "desc"}, + users: Optional[List[BaseModel]] = ( + await prisma_client.db.litellm_usertable.find_many( + where=where_conditions, + skip=skip, + take=page_size, + order={"created_at": "desc"}, + ) ) if not users: diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index 9ff79f4f78..2a0e4cc88d 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -349,3 +349,128 @@ async def test_user_info_url_encoding_plus_character(mocker): f"mock_prisma_client.get_data.call_args: {mock_prisma_client.get_data.call_args.kwargs}" ) assert mock_prisma_client.get_data.call_args.kwargs["user_id"] == expected_user_id + + +@pytest.mark.asyncio +async def test_new_user_default_teams_flow(mocker): + """ + Test that when teams are set via default_internal_user_params: + - Teams are NOT sent to generate_key_helper_fn + - Teams ARE sent to _add_user_to_team + """ + import litellm + from litellm.proxy._types import NewUserRequest, NewUserRequestTeam, UserAPIKeyAuth + from litellm.proxy.management_endpoints.internal_user_endpoints import new_user + + # Mock the prisma client + mock_prisma_client = mocker.MagicMock() + + # Setup the mock count response (under license limit) + async def mock_count(*args, **kwargs): + return 5 # Low user count, under limit + + mock_prisma_client.db.litellm_usertable.count = mock_count + + # Mock check_duplicate_user_email to pass + async def mock_check_duplicate_user_email(*args, **kwargs): + return None # No duplicate found + + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._check_duplicate_user_email", + mock_check_duplicate_user_email, + ) + + # Mock the license check to return False (under limit) + mock_license_check = mocker.MagicMock() + mock_license_check.is_over_limit.return_value = False + + # Mock generate_key_helper_fn + mock_generate_key_helper_fn = mocker.AsyncMock() + mock_generate_key_helper_fn.return_value = { + "user_id": "test-user-123", + "token": "sk-test-token-123", + "expires": None, + "max_budget": 100, + } + + # Mock _add_user_to_team + mock_add_user_to_team = mocker.AsyncMock() + + # Mock UserManagementEventHooks.async_user_created_hook + mock_user_created_hook = mocker.AsyncMock() + + # Setup default_internal_user_params with teams + original_default_params = getattr(litellm, "default_internal_user_params", None) + litellm.default_internal_user_params = { + "teams": [ + { + "team_id": "96fed65b-0182-4ff4-8429-2721cd7d42af", + "max_budget_in_team": 100, + "user_role": "user", + } + ] + } + + try: + # Patch all the imports + mocker.patch("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + mocker.patch("litellm.proxy.proxy_server._license_check", mock_license_check) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints.generate_key_helper_fn", + mock_generate_key_helper_fn, + ) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints._add_user_to_team", + mock_add_user_to_team, + ) + mocker.patch( + "litellm.proxy.management_endpoints.internal_user_endpoints.UserManagementEventHooks.async_user_created_hook", + mock_user_created_hook, + ) + + # Create test request data WITHOUT teams (teams should come from defaults) + user_request = NewUserRequest( + user_email="test@example.com", user_role="internal_user" + ) + + # Mock user_api_key_dict + mock_user_api_key_dict = UserAPIKeyAuth(user_id="test_admin") + + # Call new_user function + response = await new_user( + data=user_request, user_api_key_dict=mock_user_api_key_dict + ) + + # Verify generate_key_helper_fn was called WITHOUT teams + mock_generate_key_helper_fn.assert_called_once() + call_kwargs = mock_generate_key_helper_fn.call_args.kwargs + + # Teams should be removed from the data passed to generate_key_helper_fn + assert ( + "teams" not in call_kwargs + ), "Teams should not be passed to generate_key_helper_fn" + assert call_kwargs["request_type"] == "user" + assert call_kwargs["user_email"] == "test@example.com" + assert call_kwargs["user_role"] == "internal_user" + + # Verify _add_user_to_team was called with the default team + mock_add_user_to_team.assert_called_once() + team_call_kwargs = mock_add_user_to_team.call_args.kwargs + + assert team_call_kwargs["user_id"] == "test-user-123" + assert team_call_kwargs["team_id"] == "96fed65b-0182-4ff4-8429-2721cd7d42af" + assert team_call_kwargs["user_email"] == "test@example.com" + assert team_call_kwargs["max_budget_in_team"] == 100 + assert team_call_kwargs["user_role"] == "user" + + # Verify response structure + assert response.user_id == "test-user-123" + assert response.key == "sk-test-token-123" + + finally: + # Restore original default params + if original_default_params is not None: + litellm.default_internal_user_params = original_default_params + else: + if hasattr(litellm, "default_internal_user_params"): + delattr(litellm, "default_internal_user_params") diff --git a/ui/litellm-dashboard/src/components/SSOSettings.tsx b/ui/litellm-dashboard/src/components/SSOSettings.tsx index 7a4502d651..bbcb9598de 100644 --- a/ui/litellm-dashboard/src/components/SSOSettings.tsx +++ b/ui/litellm-dashboard/src/components/SSOSettings.tsx @@ -1,6 +1,7 @@ import React, { useState, useEffect } from "react"; import { Card, Title, Text, Divider, Button, TextInput } from "@tremor/react"; -import { Typography, Spin, message, Switch, Select, Form } from "antd"; +import { Typography, Spin, message, Switch, Select, Form, InputNumber } from "antd"; +import { PlusOutlined, DeleteOutlined } from "@ant-design/icons"; import { getInternalUserSettings, updateInternalUserSettings, modelAvailableCall } from "./networking"; import BudgetDurationDropdown, { getBudgetDurationLabel } from "./common_components/budget_duration_dropdown"; import { getModelDisplayName } from "./key_team_helpers/fetch_available_models_team_key"; @@ -12,6 +13,12 @@ interface SSOSettingsProps { userRole: string; } +interface TeamEntry { + team_id: string; + max_budget_in_team?: number; + user_role: "user" | "admin"; +} + const SSOSettings: React.FC = ({ accessToken, possibleUIRoles, userID, userRole }) => { const [loading, setLoading] = useState(true); const [settings, setSettings] = useState(null); @@ -80,10 +87,133 @@ const SSOSettings: React.FC = ({ accessToken, possibleUIRoles, })); }; + // Helper function to normalize teams array to consistent format + const normalizeTeams = (teams: any[]): TeamEntry[] => { + if (!teams || !Array.isArray(teams)) return []; + + return teams.map(team => { + if (typeof team === "string") { + return { + team_id: team, + user_role: "user" as const + }; + } else if (typeof team === "object" && team.team_id) { + return { + team_id: team.team_id, + max_budget_in_team: team.max_budget_in_team, + user_role: team.user_role || "user" + }; + } + return { + team_id: "", + user_role: "user" as const + }; + }); + }; + + // Teams editor component + const renderTeamsEditor = (teams: any[]) => { + const normalizedTeams = normalizeTeams(teams); + + const updateTeam = (index: number, field: keyof TeamEntry, value: any) => { + const updatedTeams = [...normalizedTeams]; + updatedTeams[index] = { + ...updatedTeams[index], + [field]: value + }; + handleTextInputChange("teams", updatedTeams); + }; + + const addTeam = () => { + const newTeam: TeamEntry = { + team_id: "", + user_role: "user" + }; + handleTextInputChange("teams", [...normalizedTeams, newTeam]); + }; + + const removeTeam = (index: number) => { + const updatedTeams = normalizedTeams.filter((_, i) => i !== index); + handleTextInputChange("teams", updatedTeams); + }; + + return ( +
+ {normalizedTeams.map((team, index) => ( +
+
+ Team {index + 1} + +
+ +
+
+ Team ID + updateTeam(index, "team_id", e.target.value)} + placeholder="Enter team ID" + /> +
+ +
+ Max Budget in Team + updateTeam(index, "max_budget_in_team", value)} + placeholder="Optional" + min={0} + step={0.01} + precision={2} + /> +
+ +
+ User Role + +
+
+
+ ))} + + +
+ ); + }; + const renderEditableField = (key: string, property: any, value: any) => { const type = property.type; - if (key === "user_role" && possibleUIRoles) { + if (key === "teams") { + return ( +
+ {renderTeamsEditor(editedValues[key] || [])} +
+ ); + } else if (key === "user_role" && possibleUIRoles) { return (