diff --git a/backend/src/controllers/v2/usersController.ts b/backend/src/controllers/v2/usersController.ts index 66a48b550..48acbf490 100644 --- a/backend/src/controllers/v2/usersController.ts +++ b/backend/src/controllers/v2/usersController.ts @@ -148,6 +148,43 @@ export const updateAuthProvider = async (req: Request, res: Response) => { }); } +/** + * Update auth provider of the current user to [authProvider] + * @param req + * @param res + * @returns + */ + export const updateAuthProviders = async (req: Request, res: Response) => { + const { + authProviders + } = req.body; + + if ( + req.user?.authProvider === AuthProvider.OKTA_SAML + || req.user?.authProvider === AuthProvider.AZURE_SAML + || req.user?.authProvider === AuthProvider.JUMPCLOUD_SAML + ) { + return res.status(400).send({ + message: "Failed to update user authentication method because SAML SSO is enforced" + }); + } + + const user = await User.findByIdAndUpdate( + req.user._id.toString(), + { + authProviders + }, + { + new: true + } + ); + + return res.status(200).send({ + user + }); +} + + /** * Return organizations that the current user is part of. * @param req diff --git a/backend/src/models/user.ts b/backend/src/models/user.ts index fb14e8b8a..77b8f522f 100644 --- a/backend/src/models/user.ts +++ b/backend/src/models/user.ts @@ -13,6 +13,7 @@ export interface IUser extends Document { _id: Types.ObjectId; authId?: string; authProvider?: AuthProvider; + authProviders?: AuthProvider[]; email: string; firstName?: string; lastName?: string; diff --git a/backend/src/routes/v2/users.ts b/backend/src/routes/v2/users.ts index 334ef523b..6042e6e3f 100644 --- a/backend/src/routes/v2/users.ts +++ b/backend/src/routes/v2/users.ts @@ -57,6 +57,24 @@ router.patch( usersController.updateAuthProvider ); +router.put( + "/me/auth-providers", + requireAuth({ + acceptedAuthModes: [AUTH_MODE_JWT, AUTH_MODE_API_KEY], + }), + body("authProviders").exists().isArray({ + min: 1, + }).custom((authProviders: AuthProvider[]) => { + return authProviders.every(provider => [ + AuthProvider.EMAIL, + AuthProvider.GOOGLE, + AuthProvider.GITHUB + ].includes(provider)) + }), + validateRequest, + usersController.updateAuthProviders, +); + router.get( "/me/organizations", requireAuth({ diff --git a/frontend/src/components/v2/Select/Select.tsx b/frontend/src/components/v2/Select/Select.tsx index cdf790634..ccac29a2a 100644 --- a/frontend/src/components/v2/Select/Select.tsx +++ b/frontend/src/components/v2/Select/Select.tsx @@ -16,6 +16,7 @@ type Props = { position?: "item-aligned" | "popper"; isDisabled?: boolean; icon?: IconProp; + isMulti?: boolean; }; export type SelectProps = Omit & Props; diff --git a/frontend/src/hooks/api/users/index.tsx b/frontend/src/hooks/api/users/index.tsx index a209367e2..8aec39c4e 100644 --- a/frontend/src/hooks/api/users/index.tsx +++ b/frontend/src/hooks/api/users/index.tsx @@ -15,5 +15,6 @@ export { useRegisterUserAction, useRevokeMySessions, useUpdateOrgUserRole, - useUpdateUserAuthProvider + useUpdateUserAuthProvider, + useUpdateUserAuthProviders, } from "./queries"; diff --git a/frontend/src/hooks/api/users/queries.tsx b/frontend/src/hooks/api/users/queries.tsx index 93d8d41de..0def469e9 100644 --- a/frontend/src/hooks/api/users/queries.tsx +++ b/frontend/src/hooks/api/users/queries.tsx @@ -80,6 +80,28 @@ export const useUpdateUserAuthProvider = () => { }); }; + +export const useUpdateUserAuthProviders = () => { + const queryClient = useQueryClient(); + + return useMutation({ + mutationFn: async ({ + authProviders + }: { + authProviders: string[]; + }) => { + const { data: { user } } = await apiRequest.put("/api/v2/users/me/auth-providers", { + authProviders + }); + + return user; + }, + onSuccess: () => { + queryClient.invalidateQueries(userKeys.getUser); + } + }); +}; + export const useGetUserAction = (action: string) => useQuery({ queryKey: userKeys.userAction, diff --git a/frontend/src/hooks/api/users/types.ts b/frontend/src/hooks/api/users/types.ts index 0c312b2b6..51bfa3452 100644 --- a/frontend/src/hooks/api/users/types.ts +++ b/frontend/src/hooks/api/users/types.ts @@ -13,6 +13,7 @@ export type User = { firstName?: string; lastName?: string; authProvider?: AuthProvider; + authProviders?: AuthProvider[]; encryptionVersion?: number; protectedKey?: string; protectedKeyIV?: string; diff --git a/frontend/src/views/Settings/PersonalSettingsPage/AuthMethodSection/AuthMethodSection.tsx b/frontend/src/views/Settings/PersonalSettingsPage/AuthMethodSection/AuthMethodSection.tsx index b8246180c..6d874de0e 100644 --- a/frontend/src/views/Settings/PersonalSettingsPage/AuthMethodSection/AuthMethodSection.tsx +++ b/frontend/src/views/Settings/PersonalSettingsPage/AuthMethodSection/AuthMethodSection.tsx @@ -1,17 +1,16 @@ import { useEffect } from "react"; -import { Controller, useForm } from "react-hook-form"; +import { useForm } from "react-hook-form"; import { yupResolver } from "@hookform/resolvers/yup"; import * as yup from "yup"; import { useNotificationContext } from "@app/components/context/Notifications/NotificationProvider"; import { Button, - FormControl, - Select, - SelectItem} from "@app/components/v2"; + Checkbox +} from "@app/components/v2"; import { useUser } from "@app/context"; import { - useUpdateUserAuthProvider + useUpdateUserAuthProviders } from "@app/hooks/api"; const authMethods = [ @@ -24,7 +23,7 @@ const authMethods = [ ]; const schema = yup.object({ - authMethod: yup.string().required("Auth method is required") + authMethods: yup.array().required("Auth method is required") }); export type FormData = yup.InferType; @@ -32,35 +31,38 @@ export type FormData = yup.InferType; export const AuthMethodSection = () => { const { createNotification } = useNotificationContext(); const { user } = useUser(); - const { mutateAsync, isLoading } = useUpdateUserAuthProvider(); + const { mutateAsync, isLoading } = useUpdateUserAuthProviders(); const { reset, - control, - handleSubmit + handleSubmit, + setValue, + watch, } = useForm({ defaultValues: { - authMethod: user?.authProvider ?? "email" + authMethods: [user?.authProvider ?? "email"] }, resolver: yupResolver(schema) }); + const selectedAuthMethods = watch("authMethods"); + useEffect(() => { if (user) { reset({ - authMethod: user?.authProvider ?? "email" + authMethods: [user?.authProvider ?? "email"] }); } }, [user]); const onFormSubmit = async ({ - authMethod + authMethods }: FormData) => { try { if ( - authMethod === "okta-saml" - || authMethod === "azure-saml" - || authMethod === "jumpcloud-saml" + authMethods.includes("okta-saml") + || authMethods.includes("azure-saml") + || authMethods.includes("jumpcloud-saml") ) { createNotification({ text: "SAML authentication can only be configured in your organization settings", @@ -71,7 +73,7 @@ export const AuthMethodSection = () => { } await mutateAsync({ - authProvider: authMethod + authProviders: authMethods }); createNotification({ @@ -96,36 +98,29 @@ export const AuthMethodSection = () => { Authentication Method
- ( - - - - )} - /> + { + authMethods.map(authMethod => ( + { + if (checked) { + setValue("authMethods", [ + ...selectedAuthMethods, + authMethod.value + ]) + } else { + setValue("authMethods", selectedAuthMethods.filter(auth => auth !== authMethod.value)) + } + }}> + {authMethod.label} + + )) + }
+ ); -} \ No newline at end of file +}