diff --git a/ui/litellm-dashboard/src/app/(dashboard)/admin-panel/_components/AdminPanel.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/admin-panel/_components/AdminPanel.test.tsx index 220db23338e..f11edd7023d 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/admin-panel/_components/AdminPanel.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/admin-panel/_components/AdminPanel.test.tsx @@ -1,4 +1,4 @@ -import { render, screen, waitFor } from "@testing-library/react"; +import { render, screen, waitFor, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { beforeEach, describe, expect, it, vi } from "vitest"; import AdminPanel from "./AdminPanel"; @@ -323,3 +323,73 @@ describe("AdminPanel", () => { }); }); }); + +describe("AdminPanel add allowed IP form", () => { + beforeEach(async () => { + vi.clearAllMocks(); + mockUseAuthorized.mockReturnValue({ + premiumUser: true, + accessToken: "test-token", + userId: "user-1", + }); + mockGetSSOSettings.mockResolvedValue({ values: {} }); + mockGetAllowedIPs.mockResolvedValue(["10.0.0.1"]); + mockAddAllowedIP.mockResolvedValue({}); + + const user = userEvent.setup(); + render(); + await user.click(screen.getByRole("tab", { name: /security settings/i })); + await user.click(screen.getByRole("button", { name: /allowed ips/i })); + const manageDialog = await screen.findByRole("dialog", { name: /manage allowed ip addresses/i }); + await user.click(within(manageDialog).getByRole("button", { name: /add ip address/i })); + await screen.findByPlaceholderText("Enter IP address"); + }); + + const ipField = () => screen.getByPlaceholderText("Enter IP address") as HTMLInputElement; + + const submitAddIP = async (user: ReturnType) => { + const addIpForm = ipField().form as HTMLFormElement; + await user.click(within(addIpForm).getByText("Add IP Address")); + }; + + it("sends the access token and the typed IP address", async () => { + const user = userEvent.setup(); + + await user.type(ipField(), "192.168.1.50"); + await submitAddIP(user); + + await waitFor(() => { + expect(mockAddAllowedIP).toHaveBeenCalledWith("test-token", "192.168.1.50"); + }); + expect(mockAddAllowedIP).toHaveBeenCalledTimes(1); + }); + + it("blocks the submit and shows the required message when no IP is typed", async () => { + const user = userEvent.setup(); + + await submitAddIP(user); + + expect(await screen.findByText("Please enter an IP address")).toBeInTheDocument(); + expect(mockAddAllowedIP).not.toHaveBeenCalled(); + }); + + it("submits on Enter from the IP field", async () => { + const user = userEvent.setup(); + + await user.type(ipField(), "172.16.0.9{Enter}"); + + await waitFor(() => { + expect(mockAddAllowedIP).toHaveBeenCalledWith("test-token", "172.16.0.9"); + }); + }); + + it("refreshes the allowed IP list after a successful add", async () => { + const user = userEvent.setup(); + mockGetAllowedIPs.mockResolvedValue(["10.0.0.1", "192.168.1.50"]); + + await user.type(ipField(), "192.168.1.50"); + await submitAddIP(user); + + expect(await screen.findByText("192.168.1.50")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/admin-panel/_components/AdminPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/admin-panel/_components/AdminPanel.tsx index 9c1343ccbea..3f3f7fb5db7 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/admin-panel/_components/AdminPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/admin-panel/_components/AdminPanel.tsx @@ -14,7 +14,7 @@ import { TableHeaderCell, TableRow, } from "@tremor/react"; -import { Alert, Button as Button2, Form, Input, Modal, Space, Tabs, Typography } from "antd"; +import { Alert, Modal, Space, Tabs, Typography } from "antd"; import React, { useEffect, useState } from "react"; import NewBadge from "@/components/common_components/NewBadge"; import { useBaseUrl } from "@/components/constants"; @@ -28,17 +28,51 @@ import UserBannerSettings from "@/components/Settings/AdminSettings/UserBannerSe import HashicorpVault from "@/components/Settings/AdminSettings/HashicorpVault/HashicorpVault"; import PluginSettings from "@/components/Settings/AdminSettings/PluginSettings/PluginSettings"; import SSOModals from "@/components/SSOModals"; +import { + emptySSOSettingsFormValues, + useSSOSettingsForm, + type SSOSettingsFormValues, +} from "@/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm"; import UIAccessControlForm from "@/components/UIAccessControlForm"; +import { z } from "zod/v4"; +import { FieldGroup } from "@/components/shared/form/field"; +import { FormField } from "@/components/shared/form/FormField"; +import { Button as ShadcnButton } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { useZodForm } from "@/lib/forms/useZodForm"; const { Title, Paragraph, Text } = Typography; +const allowedIPSchema = z.object({ + ip: z.string().min(1, "Please enter an IP address"), +}); + +type AllowedIPFormValues = z.infer; + +const AddAllowedIPForm = ({ onSubmit }: { onSubmit: (values: AllowedIPFormValues) => Promise }) => { + const form = useZodForm(allowedIPSchema, { defaultValues: { ip: "" } }); + + return ( +
+ + + {({ ref, ...field }) => } + +
+ Add IP Address +
+
+
+ ); +}; + interface AdminPanelProps { proxySettings?: any; } const AdminPanel: React.FC = ({ proxySettings }) => { const { premiumUser, accessToken, userId: userID } = useAuthorized(); - const [form] = Form.useForm(); + const form = useSSOSettingsForm("admin-panel"); const [isAddSSOModalVisible, setIsAddSSOModalVisible] = useState(false); const [isInstructionsModalVisible, setIsInstructionsModalVisible] = useState(false); const [isAllowedIPModalVisible, setIsAllowedIPModalVisible] = useState(false); @@ -141,7 +175,7 @@ const AdminPanel: React.FC = ({ proxySettings }) => { const handleAddSSOOk = () => { setIsAddSSOModalVisible(false); - form.resetFields(); + form.reset(emptySSOSettingsFormValues); if (accessToken && premiumUser) { checkSSOConfiguration(); } @@ -149,10 +183,10 @@ const AdminPanel: React.FC = ({ proxySettings }) => { const handleAddSSOCancel = () => { setIsAddSSOModalVisible(false); - form.resetFields(); + form.reset(emptySSOSettingsFormValues); }; - const handleShowInstructions = (formValues: Record) => { + const handleShowInstructions = (formValues: SSOSettingsFormValues) => { setIsAddSSOModalVisible(false); setIsInstructionsModalVisible(true); }; @@ -293,14 +327,7 @@ const AdminPanel: React.FC = ({ proxySettings }) => { onCancel={() => setIsAddIPModalVisible(false)} footer={null} > -
- - - - - Add IP Address - -
+ ({ + keyCreateCall: vi.fn(), +})); + +vi.mock("@/lib/toast", () => ({ + toast: { success: vi.fn(), fromError: vi.fn() }, +})); + +const ACCESS_TOKEN = "sk-access-token"; +const USER_ID = "user-1234"; + +const renderSCIM = (props?: { accessToken?: string | null; userID?: string | null }) => + renderWithProviders( + , + ); + +describe("SCIMConfig", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("sends exactly the SCIM key payload when a token name is submitted", async () => { + vi.mocked(keyCreateCall).mockResolvedValue({ key: "sk-scim-generated" }); + const user = userEvent.setup(); + renderSCIM(); + + await user.type(screen.getByLabelText("Token Name"), "My SCIM Token"); + await user.click(screen.getByRole("button", { name: /create scim token/i })); + + await waitFor(() => { + expect(keyCreateCall).toHaveBeenCalledWith(ACCESS_TOKEN, USER_ID, { + key_alias: "My SCIM Token", + team_id: null, + models: [], + allowed_routes: ["/scim/*"], + }); + }); + }); + + it("blocks the submit and shows the required message when the token name is empty", async () => { + const user = userEvent.setup(); + renderSCIM(); + + await user.click(screen.getByRole("button", { name: /create scim token/i })); + + expect(await screen.findByText("Please enter a name for your token")).toBeInTheDocument(); + expect(keyCreateCall).not.toHaveBeenCalled(); + }); + + it("submits on Enter from the token name field", async () => { + vi.mocked(keyCreateCall).mockResolvedValue({ key: "sk-scim-generated" }); + const user = userEvent.setup(); + renderSCIM(); + + await user.type(screen.getByLabelText("Token Name"), "Entered With Return{Enter}"); + + await waitFor(() => { + expect(keyCreateCall).toHaveBeenCalledWith(ACCESS_TOKEN, USER_ID, { + key_alias: "Entered With Return", + team_id: null, + models: [], + allowed_routes: ["/scim/*"], + }); + }); + }); + + it("reveals the created token and hides the creation form on success", async () => { + vi.mocked(keyCreateCall).mockResolvedValue({ key: "sk-scim-generated" }); + const user = userEvent.setup(); + renderSCIM(); + + await user.type(screen.getByLabelText("Token Name"), "My SCIM Token"); + await user.click(screen.getByRole("button", { name: /create scim token/i })); + + expect(await screen.findByText("Your SCIM Token")).toBeInTheDocument(); + expect(screen.queryByLabelText("Token Name")).not.toBeInTheDocument(); + expect(toast.success).toHaveBeenCalledWith("SCIM token created successfully"); + }); + + it("returns to the creation form when creating another token", async () => { + vi.mocked(keyCreateCall).mockResolvedValue({ key: "sk-scim-generated" }); + const user = userEvent.setup(); + renderSCIM(); + + await user.type(screen.getByLabelText("Token Name"), "My SCIM Token"); + await user.click(screen.getByRole("button", { name: /create scim token/i })); + await user.click(await screen.findByRole("button", { name: /create another token/i })); + + expect(await screen.findByLabelText("Token Name")).toBeInTheDocument(); + }); + + it("does not call the API when there is no access token", async () => { + const user = userEvent.setup(); + renderSCIM({ accessToken: null }); + + await user.type(screen.getByLabelText("Token Name"), "My SCIM Token"); + await user.click(screen.getByRole("button", { name: /create scim token/i })); + + await waitFor(() => { + expect(toast.fromError).toHaveBeenCalledWith("You need to be logged in to create a SCIM token"); + }); + expect(keyCreateCall).not.toHaveBeenCalled(); + }); + + it("surfaces a creation failure and keeps the form mounted", async () => { + vi.mocked(keyCreateCall).mockRejectedValue(new Error("boom")); + const user = userEvent.setup(); + renderSCIM(); + + await user.type(screen.getByLabelText("Token Name"), "My SCIM Token"); + await user.click(screen.getByRole("button", { name: /create scim token/i })); + + await waitFor(() => { + expect(toast.fromError).toHaveBeenCalledWith("Failed to create SCIM token: boom"); + }); + expect(screen.getByLabelText("Token Name")).toBeInTheDocument(); + }); + + it("shows the SCIM tenant URL derived from the proxy base url", () => { + renderSCIM(); + + expect(screen.getByDisplayValue("https://proxy.example.com/scim/v2")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/SCIM.tsx b/ui/litellm-dashboard/src/components/SCIM.tsx index 6b8cd5ead97..4110db4daba 100644 --- a/ui/litellm-dashboard/src/components/SCIM.tsx +++ b/ui/litellm-dashboard/src/components/SCIM.tsx @@ -1,6 +1,6 @@ import React, { useState, useEffect } from "react"; -import { Card, Title, Text, Grid, Button as TremorButton, Callout, TextInput, Divider } from "@tremor/react"; -import { Form } from "antd"; +import { Card, Title, Text, Grid, Callout, Divider } from "@tremor/react"; +import { z } from "zod/v4"; import { keyCreateCall } from "./networking"; import { CopyToClipboard } from "react-copy-to-clipboard"; import { @@ -12,6 +12,12 @@ import { } from "@ant-design/icons"; import { parseErrorMessage } from "./shared/errorUtils"; import { toast } from "@/lib/toast"; +import { FieldGroup } from "@/components/shared/form/field"; +import { FormField } from "@/components/shared/form/FormField"; +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; +import { useZodForm } from "@/lib/forms/useZodForm"; interface SCIMConfigProps { accessToken: string | null; @@ -19,8 +25,14 @@ interface SCIMConfigProps { proxySettings: any; } +const scimTokenSchema = z.object({ + key_alias: z.string().min(1, "Please enter a name for your token"), +}); + +type SCIMTokenFormValues = z.infer; + const SCIMConfig: React.FC = ({ accessToken, userID, proxySettings }) => { - const [form] = Form.useForm(); + const form = useZodForm(scimTokenSchema, { defaultValues: { key_alias: "" } }); const [isCreatingToken, setIsCreatingToken] = useState(false); const [tokenData, setTokenData] = useState(null); const [baseUrl, setBaseUrl] = useState(""); @@ -40,7 +52,7 @@ const SCIMConfig: React.FC = ({ accessToken, userID, proxySetti const scimBaseUrl = `${baseUrl}/scim/v2`; - const handleCreateSCIMToken = async (values: any) => { + const handleCreateSCIMToken = async (values: SCIMTokenFormValues) => { if (!accessToken || !userID) { toast.fromError("You need to be logged in to create a SCIM token"); return; @@ -73,7 +85,7 @@ const SCIMConfig: React.FC = ({ accessToken, userID, proxySetti
SCIM Configuration
- + System for Cross-domain Identity Management (SCIM) allows you to automatically provision and manage users and groups in LiteLLM. @@ -84,7 +96,7 @@ const SCIMConfig: React.FC = ({ accessToken, userID, proxySetti {/* Step 1: SCIM URL */}
-
+
1
@@ -92,16 +104,16 @@ const SCIMConfig: React.FC<SCIMConfigProps> = ({ accessToken, userID, proxySetti SCIM Tenant URL
- + Use this URL in your identity provider SCIM integration settings.
- + toast.success("URL copied to clipboard")}> - +
@@ -109,7 +121,7 @@ const SCIMConfig: React.FC = ({ accessToken, userID, proxySetti {/* Step 2: SCIM Token */}
-
+
2
@@ -124,50 +136,52 @@ const SCIMConfig: React.FC<SCIMConfigProps> = ({ accessToken, userID, proxySetti </Callout> {!tokenData ? ( - <div className="bg-gray-50 p-4 rounded-lg"> - <Form form={form} onFinish={handleCreateSCIMToken} layout="vertical"> - <Form.Item - name="key_alias" - label="Token Name" - rules={[{ required: true, message: "Please enter a name for your token" }]} - > - <TextInput placeholder="SCIM Access Token" /> - </Form.Item> - <Form.Item> - <TremorButton - variant="primary" - type="submit" - loading={isCreatingToken} - className="flex items-center" - > - <KeyOutlined className="h-4 w-4 mr-1" /> - Create SCIM Token - </TremorButton> - </Form.Item> - </Form> + <div className="bg-muted p-4 rounded-lg"> + <form onSubmit={form.handleSubmit(handleCreateSCIMToken)}> + <FieldGroup> + <FormField control={form.control} name="key_alias" label="Token Name"> + {({ ref, ...field }) => <Input {...field} ref={ref} placeholder="SCIM Access Token" />} + </FormField> + <div> + <Button type="submit" disabled={isCreatingToken} className="flex items-center"> + {isCreatingToken ? ( + <UiLoadingSpinner className="size-4 mr-1" /> + ) : ( + <KeyOutlined className="h-4 w-4 mr-1" /> + )} + Create SCIM Token + </Button> + </div> + </FieldGroup> + </form> </div> ) : ( - <Card className="border border-yellow-300 bg-yellow-50"> - <div className="flex items-center mb-2 text-yellow-800"> + <Card className="border border-yellow-300 bg-yellow-50 dark:border-yellow-800 dark:bg-yellow-950"> + <div className="flex items-center mb-2 text-yellow-800 dark:text-yellow-300"> <ExclamationCircleOutlined className="h-5 w-5 mr-2" /> - <Title className="text-lg text-yellow-800">Your SCIM Token + Your SCIM Token
- + Make sure to copy this token now. You will not be able to see it again.
- + toast.success("Token copied to clipboard")}> - +
- setTokenData(null)}> + )}
diff --git a/ui/litellm-dashboard/src/components/SSOModals.test.tsx b/ui/litellm-dashboard/src/components/SSOModals.test.tsx index 47b1851459b..bee5f13c626 100644 --- a/ui/litellm-dashboard/src/components/SSOModals.test.tsx +++ b/ui/litellm-dashboard/src/components/SSOModals.test.tsx @@ -1,7 +1,10 @@ import { fireEvent, render, screen, waitFor } from "@testing-library/react"; -import { Form, type FormInstance } from "antd"; +import userEvent from "@testing-library/user-event"; import { describe, expect, it, vi } from "vitest"; import SSOModals from "./SSOModals"; +import { useSSOSettingsForm } from "./Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm"; + +const user = () => userEvent.setup({ pointerEventsCheck: 0 }); // Mock the networking functions vi.mock("./networking", () => ({ @@ -20,7 +23,7 @@ import { getSSOSettings, updateSSOSettings } from "./networking"; describe("SSOModals", () => { it("should render the SSOModals component", () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("admin-panel"); return ( { it("should show validation error if proxy base url is not a valid URL", async () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("admin-panel"); return ( { render(); // Find and interact with the SSO provider select - const ssoProviderSelect = screen.getByLabelText("SSO Provider"); - fireEvent.mouseDown(ssoProviderSelect); + await user().click(screen.getByLabelText("SSO Provider")); // Wait for dropdown and select Google const googleOption = await screen.findByText("Google SSO"); - fireEvent.click(googleOption); + await user().click(googleOption); // Fill in the email field const emailInput = screen.getByLabelText("Proxy Admin Email"); @@ -94,7 +96,7 @@ describe("SSOModals", () => { it("should show validation error if proxy base url ends with trailing slash", async () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("admin-panel"); return ( { render(); // Find and interact with the SSO provider select - const ssoProviderSelect = screen.getByLabelText("SSO Provider"); - fireEvent.mouseDown(ssoProviderSelect); + await user().click(screen.getByLabelText("SSO Provider")); // Wait for dropdown and select Google const googleOption = await screen.findByText("Google SSO"); - fireEvent.click(googleOption); + await user().click(googleOption); // Fill in the email field const emailInput = screen.getByLabelText("Proxy Admin Email"); @@ -139,7 +140,7 @@ describe("SSOModals", () => { it("should allow typing https:// without interfering with slashes", async () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("admin-panel"); return ( { it("should only show URL format error for incomplete URLs, not trailing slash error", async () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("admin-panel"); return ( { render(); // Find and interact with the SSO provider select - const ssoProviderSelect = screen.getByLabelText("SSO Provider"); - fireEvent.mouseDown(ssoProviderSelect); + await user().click(screen.getByLabelText("SSO Provider")); // Wait for dropdown and select Google const googleOption = await screen.findByText("Google SSO"); - fireEvent.click(googleOption); + await user().click(googleOption); // Fill in the email field const emailInput = screen.getByLabelText("Proxy Admin Email"); @@ -258,7 +258,7 @@ describe("SSOModals", () => { (getSSOSettings as any).mockResolvedValue(mockSSOData); const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("admin-panel"); return ( { // Mock getSSOSettings to return empty data so form starts clean (getSSOSettings as any).mockResolvedValue({ values: {} }); - let formInstance: any = null; + let formInstance: ReturnType | null = null; const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("admin-panel"); formInstance = form; return ( @@ -333,16 +333,16 @@ describe("SSOModals", () => { }); // Set the provider directly using the form to trigger conditional rendering - formInstance.setFieldsValue({ sso_provider: "okta" }); + formInstance!.setValue("sso_provider", "okta"); // Wait for the "Use Role Mappings" checkbox to appear await waitFor(() => { - expect(screen.getByLabelText("Use Role Mappings")).toBeInTheDocument(); + expect(screen.getAllByLabelText("Use Role Mappings")[0]).toBeInTheDocument(); }); // Enable role mappings - const roleMappingsCheckbox = screen.getByLabelText("Use Role Mappings"); - fireEvent.click(roleMappingsCheckbox); + const roleMappingsCheckbox = screen.getAllByLabelText("Use Role Mappings")[0]; + await user().click(roleMappingsCheckbox); // Fill required fields const emailInput = screen.getByLabelText("Proxy Admin Email"); @@ -411,10 +411,10 @@ describe("SSOModals", () => { vi.mocked(updateSSOSettings).mockResolvedValue({}); vi.mocked(getSSOSettings).mockResolvedValue({ values: {} }); - let formInstance: FormInstance | null = null; + let formInstance: ReturnType | null = null; const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("admin-panel"); formInstance = form; return ( @@ -439,7 +439,7 @@ describe("SSOModals", () => { expect(getSSOSettings).toHaveBeenCalledWith("test-token"); }); - formInstance?.setFieldsValue({ sso_provider: "saml" }); + formInstance!.setValue("sso_provider", "saml"); await waitFor(() => { expect(screen.getByLabelText("IdP Metadata URL")).toBeInTheDocument(); @@ -457,7 +457,7 @@ describe("SSOModals", () => { fireEvent.change(screen.getByLabelText("SP Entity ID"), { target: { value: "https://proxy.example.com/sso/saml/metadata" }, }); - fireEvent.click(screen.getByLabelText("Allow IdP-initiated (unsolicited) responses")); + await user().click(screen.getAllByLabelText("Allow IdP-initiated (unsolicited) responses")[0]); fireEvent.click(screen.getByText("Save")); @@ -482,7 +482,7 @@ describe("SSOModals", () => { (toast.success as any).mockImplementation(() => {}); const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("admin-panel"); return ( { it("renders provider logos in the SSO provider dropdown", async () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("admin-panel"); return ( { render(); - fireEvent.mouseDown(screen.getByLabelText("SSO Provider")); + await user().click(screen.getByLabelText("SSO Provider")); await waitFor(() => { expect(screen.getAllByAltText("Google SSO logo").length).toBeGreaterThan(0); diff --git a/ui/litellm-dashboard/src/components/SSOModals.tsx b/ui/litellm-dashboard/src/components/SSOModals.tsx index 6c7c68e326f..4a156c3b30a 100644 --- a/ui/litellm-dashboard/src/components/SSOModals.tsx +++ b/ui/litellm-dashboard/src/components/SSOModals.tsx @@ -1,22 +1,34 @@ import React, { useEffect, useState } from "react"; -import { Modal, Form, Button as Button2, Select, Checkbox } from "antd"; -import { Text, TextInput } from "@tremor/react"; +import { FormProvider, useWatch, type UseFormReturn } from "react-hook-form"; +import { Modal } from "antd"; +import { Text } from "@tremor/react"; import { getSSOSettings, updateSSOSettings } from "./networking"; import { toast } from "@/lib/toast"; import { parseErrorMessage } from "./shared/errorUtils"; -import { Logo } from "@/components/molecules/logo/Logo"; -import { ssoProviderDisplayNames, ssoProviderLogoMap } from "./Settings/AdminSettings/SSOSettings/constants"; -import { renderProviderFields } from "./Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm"; +import { Button } from "@/components/ui/button"; +import { FieldGroup } from "@/components/shared/form/field"; +import { + GroupClaimField, + MappingToggleField, + ProxyAdminEmailField, + ProxyBaseUrlField, + RoleMappingTeamFields, + SSOProviderSelectField, + emptySSOSettingsFormValues, + renderProviderFields, + submitMountedSSOValues, + type SSOSettingsFormValues, +} from "./Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm"; interface SSOModalsProps { isAddSSOModalVisible: boolean; isInstructionsModalVisible: boolean; handleAddSSOOk: () => void; handleAddSSOCancel: () => void; - handleShowInstructions: (formValues: Record) => void; + handleShowInstructions: (formValues: SSOSettingsFormValues) => void; handleInstructionsOk: () => void; handleInstructionsCancel: () => void; - form: any; // Replace with proper Form type if available + form: UseFormReturn; accessToken: string | null; ssoConfigured?: boolean; // Add optional prop to indicate if SSO is configured } @@ -45,6 +57,9 @@ const SSOModals: React.FC = ({ ssoConfigured = false, // Default to false if not provided }) => { const [isClearConfirmModalVisible, setIsClearConfirmModalVisible] = useState(false); + const provider = useWatch({ control: form.control, name: "sso_provider" }); + const useRoleMappings = useWatch({ control: form.control, name: "use_role_mappings" }); + const showRoleMappingToggle = provider === "okta" || provider === "generic"; // Load existing SSO settings when modal opens useEffect(() => { @@ -79,20 +94,29 @@ const SSOModals: React.FC = ({ } // Set form values with existing data (excluding UI access control fields) - const formValues = { - sso_provider: selectedProvider, + const formValues: SSOSettingsFormValues = { + sso_provider: selectedProvider ?? "", proxy_base_url: ssoData.values.proxy_base_url, user_email: ssoData.values.user_email, - ...ssoData.values, + google_client_id: ssoData.values.google_client_id, + google_client_secret: ssoData.values.google_client_secret, + microsoft_client_id: ssoData.values.microsoft_client_id, + microsoft_client_secret: ssoData.values.microsoft_client_secret, + microsoft_tenant: ssoData.values.microsoft_tenant, + generic_client_id: ssoData.values.generic_client_id, + generic_client_secret: ssoData.values.generic_client_secret, + generic_authorization_endpoint: ssoData.values.generic_authorization_endpoint, + generic_token_endpoint: ssoData.values.generic_token_endpoint, + generic_userinfo_endpoint: ssoData.values.generic_userinfo_endpoint, + generic_scope: ssoData.values.generic_scope, + saml_idp_metadata_url: ssoData.values.saml_idp_metadata_url, + saml_idp_metadata_xml: ssoData.values.saml_idp_metadata_xml, + saml_sp_entity_id: ssoData.values.saml_sp_entity_id, ...roleMappingFields, saml_allow_unsolicited: ssoData.values.saml_allow_unsolicited === "true", }; - // Clear form first, then set values with a small delay to ensure proper initialization - form.resetFields(); - setTimeout(() => { - form.setFieldsValue(formValues); - }, 100); + form.reset({ ...emptySSOSettingsFormValues, ...formValues }); } } catch (error) { console.error("Failed to load SSO settings:", error); @@ -104,7 +128,7 @@ const SSOModals: React.FC = ({ }, [isAddSSOModalVisible, accessToken, form]); // Enhanced form submission handler - const handleFormSubmit = async (formValues: Record) => { + const handleFormSubmit = async (formValues: SSOSettingsFormValues) => { if (!accessToken) { toast.fromError("No access token available"); return; @@ -122,7 +146,7 @@ const SSOModals: React.FC = ({ ...rest } = formValues; - const payload: any = { + const payload: Record = { ...rest, }; @@ -152,7 +176,7 @@ const SSOModals: React.FC = ({ payload.role_mappings = { provider: "generic", group_claim, - default_role: defaultRoleMapping[default_role] || "internal_user", + default_role: (default_role ? defaultRoleMapping[default_role] : undefined) || "internal_user", roles: { proxy_admin: splitTeams(proxy_admin_teams), proxy_admin_viewer: splitTeams(admin_viewer_teams), @@ -206,7 +230,7 @@ const SSOModals: React.FC = ({ await updateSSOSettings(accessToken, clearSettings); // Clear the form - form.resetFields(); + form.reset(emptySSOSettingsFormValues); // Close the confirmation modal setIsClearConfirmModalVisible(false); @@ -232,186 +256,36 @@ const SSOModals: React.FC = ({ onOk={handleAddSSOOk} onCancel={handleAddSSOCancel} > -
- <> - - - - - prevValues.sso_provider !== currentValues.sso_provider} - > - {({ getFieldValue }) => { - const provider = getFieldValue("sso_provider"); - return provider ? renderProviderFields(provider) : null; - }} - - - - - - value?.trim()} - rules={[ - { required: true, message: "Please enter the proxy base url" }, - { - pattern: /^https?:\/\/.+/, - message: "URL must start with http:// or https://", - }, - { - validator: (_, value) => { - // Only check for trailing slash if the URL starts with http:// or https:// - if (value && /^https?:\/\/.+/.test(value) && value.endsWith("/")) { - return Promise.reject("URL must not end with a trailing slash"); - } - return Promise.resolve(); - }, - }, - ]} - > - - - - prevValues.sso_provider !== currentValues.sso_provider} - > - {({ getFieldValue }) => { - const provider = getFieldValue("sso_provider"); - return provider === "okta" || provider === "generic" ? ( - - - - ) : null; - }} - - - - prevValues.use_role_mappings !== currentValues.use_role_mappings - } - > - {({ getFieldValue }) => { - const useRoleMappings = getFieldValue("use_role_mappings"); - return useRoleMappings ? ( - - - - ) : null; - }} - - - - prevValues.use_role_mappings !== currentValues.use_role_mappings - } - > - {({ getFieldValue }) => { - const useRoleMappings = getFieldValue("use_role_mappings"); - return useRoleMappings ? ( - <> - - - - - - - - - - - - - - - - - - - - - ) : null; - }} - - -
+ { + event.preventDefault(); + submitMountedSSOValues(form, "admin-panel", handleFormSubmit)(); }} > - {ssoConfigured && ( - setIsClearConfirmModalVisible(true)} - style={{ - backgroundColor: "#6366f1", - borderColor: "#6366f1", - color: "white", - }} - onMouseEnter={(e) => { - e.currentTarget.style.backgroundColor = "#5558eb"; - e.currentTarget.style.borderColor = "#5558eb"; - }} - onMouseLeave={(e) => { - e.currentTarget.style.backgroundColor = "#6366f1"; - e.currentTarget.style.borderColor = "#6366f1"; - }} - > - Clear - - )} - Save -
-
+ + + {provider ? renderProviderFields(provider) : null} + + + {showRoleMappingToggle && } + {useRoleMappings && ( + <> + + + + )} + +
+ {ssoConfigured && ( + + )} + +
+ + {/* Clear Confirmation Modal */} @@ -448,7 +322,9 @@ const SSOModals: React.FC = ({ 3. Confirm your SSO is configured correctly and you can login on the new Tab 4. If Step 3 is successful, you can close this tab
- Done +
diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/HashicorpVault/EditHashicorpVaultModal.test.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/HashicorpVault/EditHashicorpVaultModal.test.tsx new file mode 100644 index 00000000000..2eb2137ab6d --- /dev/null +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/HashicorpVault/EditHashicorpVaultModal.test.tsx @@ -0,0 +1,190 @@ +import { screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; + +import { renderWithProviders } from "../../../../../tests/test-utils"; +import EditHashicorpVaultModal from "./EditHashicorpVaultModal"; +import { useHashicorpVaultConfig } from "@/app/(dashboard)/hooks/configOverrides/useHashicorpVaultConfig"; +import { useUpdateHashicorpVaultConfig } from "@/app/(dashboard)/hooks/configOverrides/useUpdateHashicorpVaultConfig"; + +vi.mock("@/app/(dashboard)/hooks/configOverrides/useHashicorpVaultConfig", () => ({ + useHashicorpVaultConfig: vi.fn(), +})); + +vi.mock("@/app/(dashboard)/hooks/configOverrides/useUpdateHashicorpVaultConfig", () => ({ + useUpdateHashicorpVaultConfig: vi.fn(), +})); + +vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ + default: () => ({ accessToken: "sk-access-token" }), +})); + +vi.mock("@/lib/toast", () => ({ + toast: { success: vi.fn(), fromError: vi.fn() }, +})); + +const ALL_FIELDS = [ + "vault_addr", + "vault_namespace", + "vault_mount_name", + "vault_path_prefix", + "vault_token", + "approle_role_id", + "approle_secret_id", + "approle_mount_path", + "client_cert", + "client_key", + "vault_cert_role", +] as const; + +const propertiesFor = (fields: readonly string[]) => + Object.fromEntries(fields.map((name) => [name, { description: `${name} description` }])); + +const mutate = vi.fn(); + +const setup = (options?: { values?: Record; fields?: readonly string[] }) => { + vi.mocked(useHashicorpVaultConfig).mockReturnValue({ + data: { + field_schema: { properties: propertiesFor(options?.fields ?? ALL_FIELDS) }, + values: options?.values ?? {}, + }, + } as unknown as ReturnType); + + vi.mocked(useUpdateHashicorpVaultConfig).mockReturnValue({ + mutate, + isPending: false, + } as unknown as ReturnType); +}; + +const renderModal = (onSuccess = vi.fn(), onCancel = vi.fn()) => + renderWithProviders(); + +const save = async (user: ReturnType) => + user.click(screen.getByRole("button", { name: "Save" })); + +describe("EditHashicorpVaultModal", () => { + beforeEach(() => { + vi.clearAllMocks(); + }); + + it("clears untouched non-sensitive fields and omits untouched sensitive fields", async () => { + setup({ + values: { + vault_addr: "https://vault.example.com", + vault_namespace: "team-ns", + vault_token: "super-secret-token", + approle_secret_id: "super-secret-id", + }, + }); + const user = userEvent.setup(); + renderModal(); + + await save(user); + + await waitFor(() => { + expect(mutate).toHaveBeenCalledTimes(1); + }); + expect(mutate.mock.calls[0][0]).toEqual({ + vault_addr: "https://vault.example.com", + vault_namespace: "team-ns", + vault_mount_name: "", + vault_path_prefix: "", + approle_role_id: "", + approle_mount_path: "", + client_cert: "", + vault_cert_role: "", + }); + }); + + it("sends a sensitive field only once it is typed into", async () => { + setup({ values: { vault_addr: "https://vault.example.com", vault_token: "super-secret-token" } }); + const user = userEvent.setup(); + renderModal(); + + await user.type(screen.getByLabelText("Token"), "rotated-token"); + await save(user); + + await waitFor(() => { + expect(mutate).toHaveBeenCalledTimes(1); + }); + expect(mutate.mock.calls[0][0]).toMatchObject({ vault_token: "rotated-token" }); + }); + + it("never seeds a stored secret into its input", () => { + setup({ values: { vault_token: "super-secret-token", approle_secret_id: "super-secret-id" } }); + renderModal(); + + expect(screen.getByLabelText("Token")).toHaveValue(""); + expect(screen.getByLabelText("Secret ID")).toHaveValue(""); + }); + + it("renders only the fields the schema declares, and sends only those", async () => { + setup({ fields: ["vault_addr", "vault_token"], values: { vault_addr: "https://vault.example.com" } }); + const user = userEvent.setup(); + renderModal(); + + expect(screen.queryByLabelText("Namespace")).not.toBeInTheDocument(); + expect(screen.queryByLabelText("Role ID")).not.toBeInTheDocument(); + + await save(user); + + await waitFor(() => { + expect(mutate).toHaveBeenCalledTimes(1); + }); + expect(mutate.mock.calls[0][0]).toEqual({ vault_addr: "https://vault.example.com" }); + }); + + it("blocks the submit when the vault address does not start with http", async () => { + setup({ values: {} }); + const user = userEvent.setup(); + renderModal(); + + await user.type(screen.getByLabelText("Vault Address"), "vault.example.com"); + await save(user); + + expect(await screen.findByText("Must start with http:// or https://")).toBeInTheDocument(); + expect(mutate).not.toHaveBeenCalled(); + }); + + it("accepts an empty vault address, because the pattern rule is not a required rule", async () => { + setup({ values: {} }); + const user = userEvent.setup(); + renderModal(); + + await save(user); + + await waitFor(() => { + expect(mutate).toHaveBeenCalledTimes(1); + }); + expect(mutate.mock.calls[0][0]).toMatchObject({ vault_addr: "" }); + }); + + it("tells the admin a stored secret is kept when the field is left blank", () => { + setup({ values: { vault_token: "super-secret-token" } }); + renderModal(); + + expect(screen.getByLabelText("Token")).toHaveAttribute( + "placeholder", + "Leave blank to keep existing (super-secret-token)", + ); + }); + + it("falls back to the schema description when no secret is stored yet", () => { + setup({ values: {} }); + renderModal(); + + expect(screen.getByLabelText("Token")).toHaveAttribute("placeholder", "vault_token description"); + }); + + it("closes without saving when cancelled", async () => { + setup({ values: {} }); + const onCancel = vi.fn(); + const user = userEvent.setup(); + renderModal(vi.fn(), onCancel); + + await user.click(screen.getByRole("button", { name: "Cancel" })); + + expect(onCancel).toHaveBeenCalledTimes(1); + expect(mutate).not.toHaveBeenCalled(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/HashicorpVault/EditHashicorpVaultModal.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/HashicorpVault/EditHashicorpVaultModal.tsx index 4904fa016a9..0658f28611a 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/HashicorpVault/EditHashicorpVaultModal.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/HashicorpVault/EditHashicorpVaultModal.tsx @@ -4,17 +4,26 @@ import { useHashicorpVaultConfig } from "@/app/(dashboard)/hooks/configOverrides import { useUpdateHashicorpVaultConfig } from "@/app/(dashboard)/hooks/configOverrides/useUpdateHashicorpVaultConfig"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; import { toast } from "@/lib/toast"; -import { Button, Divider, Form, Input, Modal, Space, Typography } from "antd"; -import React, { useEffect } from "react"; +import { Modal } from "antd"; +import React, { useMemo } from "react"; +import { z } from "zod/v4"; +import { FieldGroup } from "@/components/shared/form/field"; +import { FormField } from "@/components/shared/form/FormField"; +import { PasswordInput } from "@/components/shared/PasswordInput"; +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; +import { Separator } from "@/components/ui/separator"; +import { useZodForm } from "@/lib/forms/useZodForm"; import { SENSITIVE_FIELDS, FIELD_LABELS } from "./constants"; -interface FieldGroup { +interface VaultFieldGroup { title: string; subtitle?: string; fields: string[]; } -const FIELD_GROUPS: FieldGroup[] = [ +const FIELD_GROUPS: VaultFieldGroup[] = [ { title: "Connection", fields: ["vault_addr", "vault_namespace", "vault_mount_name", "vault_path_prefix"], @@ -36,6 +45,22 @@ const FIELD_GROUPS: FieldGroup[] = [ }, ]; +type VaultFormValues = Record; + +const buildSchema = (fields: readonly string[]): z.ZodType => + z.object( + Object.fromEntries( + fields.map((name) => [ + name, + name === "vault_addr" + ? z.string().refine((value) => value.length === 0 || /^https?:\/\/.+/.test(value), { + message: "Must start with http:// or https://", + }) + : z.string(), + ]), + ), + ) as unknown as z.ZodType; + interface EditHashicorpVaultModalProps { isVisible: boolean; onCancel: () => void; @@ -43,41 +68,40 @@ interface EditHashicorpVaultModalProps { } const EditHashicorpVaultModal: React.FC = ({ isVisible, onCancel, onSuccess }) => { - const [form] = Form.useForm(); const { accessToken } = useAuthorized(); const { data } = useHashicorpVaultConfig(); const { mutate, isPending } = useUpdateHashicorpVaultConfig(accessToken); - const schema = data?.field_schema; - const properties = schema?.properties ?? {}; - const rawValues = data?.values ?? {}; + const properties: Record = useMemo( + () => data?.field_schema?.properties ?? {}, + [data], + ); + const rawValues: Record = useMemo(() => data?.values ?? {}, [data]); - useEffect(() => { - if (isVisible && data) { - form.resetFields(); - // Only set non-sensitive fields — sensitive ones show as placeholders - const formValues: Record = {}; - for (const [key, value] of Object.entries(rawValues)) { - if (!SENSITIVE_FIELDS.has(key)) { - formValues[key] = value; - } - } - form.setFieldsValue(formValues); - } - }, [isVisible, data, form]); + const visibleFields = useMemo( + () => FIELD_GROUPS.flatMap((group) => group.fields).filter((name) => properties[name] !== undefined), + [properties], + ); - const handleSubmit = (formValues: Record) => { - const config: Record = {}; - for (const [key, value] of Object.entries(formValues)) { - if (value !== undefined && value !== null && value !== "") { - // Non-empty value → update - config[key] = value; - } else if (!SENSITIVE_FIELDS.has(key)) { - // Non-sensitive field cleared → send "" to clear it on the backend - config[key] = ""; - } - // Sensitive field left blank → omit from payload (keep existing) - } + const seededValues = useMemo( + () => + Object.fromEntries( + visibleFields.map((name) => [name, SENSITIVE_FIELDS.has(name) ? "" : ((rawValues[name] ?? "") as string)]), + ), + [visibleFields, rawValues], + ); + + const schema = useMemo(() => buildSchema(visibleFields), [visibleFields]); + const form = useZodForm(schema, { values: seededValues }); + + const handleSubmit = (formValues: VaultFormValues) => { + const config: Record = Object.fromEntries( + Object.entries(formValues).flatMap(([key, value]) => { + if (value !== undefined && value !== null && value !== "") return [[key, value]]; + if (!SENSITIVE_FIELDS.has(key)) return [[key, ""]]; + return []; + }), + ); mutate(config, { onSuccess: () => { @@ -91,7 +115,7 @@ const EditHashicorpVaultModal: React.FC = ({ isVis }; const handleCancel = () => { - form.resetFields(); + form.reset(seededValues); onCancel(); }; @@ -99,20 +123,21 @@ const EditHashicorpVaultModal: React.FC = ({ isVis const fieldSchema = properties[fieldName]; if (!fieldSchema) return null; - const rules = - fieldName === "vault_addr" - ? [{ pattern: /^https?:\/\/.+/, message: "Must start with http:// or https://" }] - : undefined; - const isSensitive = SENSITIVE_FIELDS.has(fieldName); const existingValue = rawValues[fieldName]; const hasExistingValue = isSensitive && existingValue != null && existingValue !== ""; const placeholder = hasExistingValue ? `Leave blank to keep existing (${existingValue})` : fieldSchema?.description; return ( - - {isSensitive ? : } - + + {({ ref, ...field }) => + isSensitive ? ( + + ) : ( + + ) + } + ); }; @@ -122,33 +147,28 @@ const EditHashicorpVaultModal: React.FC = ({ isVis open={isVisible} width={700} footer={ - - - - +
} onCancel={handleCancel} > -
+ {FIELD_GROUPS.map((group, index) => (
- {index > 0 && } - - {group.title} - - {group.subtitle && ( - - {group.subtitle} - - )} - {group.fields.map(renderField)} + {index > 0 && } +
{group.title}
+ {group.subtitle &&

{group.subtitle}

} + {group.fields.map(renderField)}
))} -
+ ); }; diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/AddSSOSettingsModal.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/AddSSOSettingsModal.tsx index a1f62ea4939..fcbb6300d70 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/AddSSOSettingsModal.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/AddSSOSettingsModal.tsx @@ -2,9 +2,16 @@ import { toast } from "@/lib/toast"; import { parseErrorMessage } from "@/components/shared/errorUtils"; -import { Button, Form, Modal, Space } from "antd"; +import { Modal } from "antd"; import React from "react"; -import BaseSSOSettingsForm from "./BaseSSOSettingsForm"; +import BaseSSOSettingsForm, { + emptySSOSettingsFormValues, + submitMountedSSOValues, + useSSOSettingsForm, + type SSOSettingsFormValues, +} from "./BaseSSOSettingsForm"; +import { Button } from "@/components/ui/button"; +import { UiLoadingSpinner } from "@/components/ui/ui-loading-spinner"; import { useEditSSOSettings } from "@/app/(dashboard)/hooks/sso/useEditSSOSettings"; import { processSSOSettingsPayload } from "../utils"; @@ -15,11 +22,10 @@ interface AddSSOSettingsModalProps { } const AddSSOSettingsModal: React.FC = ({ isVisible, onCancel, onSuccess }) => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("sso-settings"); const { mutateAsync, isPending } = useEditSSOSettings(); - // Enhanced form submission handler - const handleFormSubmit = async (formValues: Record) => { + const handleFormSubmit = async (formValues: SSOSettingsFormValues) => { const payload = processSSOSettingsPayload(formValues); await mutateAsync(payload, { @@ -34,7 +40,7 @@ const AddSSOSettingsModal: React.FC = ({ isVisible, on }; const handleCancel = () => { - form.resetFields(); + form.reset(emptySSOSettingsFormValues); onCancel(); }; @@ -44,14 +50,19 @@ const AddSSOSettingsModal: React.FC = ({ isVisible, on open={isVisible} width={800} footer={ - - - - +
} onCancel={handleCancel} > diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.test.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.test.tsx index 8043e41a2de..54566e1d75a 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.test.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.test.tsx @@ -1,8 +1,20 @@ -import { Form } from "antd"; import { act, fireEvent, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; import { renderWithProviders } from "../../../../../../tests/test-utils"; import { afterEach, describe, expect, it, vi } from "vitest"; -import BaseSSOSettingsForm, { renderProviderFields, ssoProviderConfigs } from "./BaseSSOSettingsForm"; +import BaseSSOSettingsForm, { + emptySSOSettingsFormValues, + renderProviderFields, + ssoProviderConfigs, + submitMountedSSOValues, + useSSOSettingsForm, +} from "./BaseSSOSettingsForm"; + +const user = () => userEvent.setup({ pointerEventsCheck: 0 }); + +const openProviderDropdown = async () => { + await user().click(screen.getByLabelText("SSO Provider")); +}; describe("BaseSSOSettingsForm", () => { afterEach(() => { @@ -11,7 +23,7 @@ describe("BaseSSOSettingsForm", () => { it("should render", () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("sso-settings"); const handleSubmit = vi.fn(); return ; @@ -26,7 +38,7 @@ describe("BaseSSOSettingsForm", () => { it("should render provider fields when provider is selected", async () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("sso-settings"); const handleSubmit = vi.fn(); return ; @@ -34,13 +46,10 @@ describe("BaseSSOSettingsForm", () => { renderWithProviders(); - const providerSelect = screen.getByLabelText("SSO Provider"); - await act(async () => { - fireEvent.mouseDown(providerSelect); - }); + await openProviderDropdown(); const googleOption = await screen.findByText(/google sso/i); - fireEvent.click(googleOption); + await user().click(googleOption); await waitFor(() => { expect(screen.getByText("Google Client ID")).toBeInTheDocument(); @@ -50,7 +59,7 @@ describe("BaseSSOSettingsForm", () => { it("should show role mappings fields for okta provider", async () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("sso-settings"); const handleSubmit = vi.fn(); return ; @@ -58,13 +67,10 @@ describe("BaseSSOSettingsForm", () => { renderWithProviders(); - const providerSelect = screen.getByLabelText("SSO Provider"); - await act(async () => { - fireEvent.mouseDown(providerSelect); - }); + await openProviderDropdown(); const oktaOption = await screen.findByText(/okta/i); - fireEvent.click(oktaOption); + await user().click(oktaOption); await waitFor(() => { expect(screen.getByText("Use Role Mappings")).toBeInTheDocument(); @@ -73,7 +79,7 @@ describe("BaseSSOSettingsForm", () => { it("should validate proxy base url format", async () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("sso-settings"); const handleSubmit = vi.fn(); return ; @@ -94,7 +100,7 @@ describe("BaseSSOSettingsForm", () => { it("should validate proxy base url trailing slash", async () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("sso-settings"); const handleSubmit = vi.fn(); return ; @@ -115,7 +121,7 @@ describe("BaseSSOSettingsForm", () => { it("should show role mappings fields when use_role_mappings is checked for generic provider", async () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("sso-settings"); const handleSubmit = vi.fn(); return ; @@ -123,22 +129,16 @@ describe("BaseSSOSettingsForm", () => { renderWithProviders(); - const providerSelect = screen.getByLabelText("SSO Provider"); - await act(async () => { - fireEvent.mouseDown(providerSelect); - }); + await openProviderDropdown(); const genericOption = await screen.findByText(/generic sso/i); - fireEvent.click(genericOption); + await user().click(genericOption); await waitFor(() => { expect(screen.getByText("Use Role Mappings")).toBeInTheDocument(); }); - const checkbox = screen.getByLabelText("Use Role Mappings"); - await act(async () => { - fireEvent.click(checkbox); - }); + await user().click(screen.getAllByLabelText("Use Role Mappings")[0]); await waitFor(() => { expect(screen.getByText("Group Claim")).toBeInTheDocument(); @@ -148,7 +148,7 @@ describe("BaseSSOSettingsForm", () => { it("should show team mappings checkbox for okta provider", async () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("sso-settings"); const handleSubmit = vi.fn(); return ; @@ -156,13 +156,10 @@ describe("BaseSSOSettingsForm", () => { renderWithProviders(); - const providerSelect = screen.getByLabelText("SSO Provider"); - await act(async () => { - fireEvent.mouseDown(providerSelect); - }); + await openProviderDropdown(); const oktaOption = await screen.findByText(/okta/i); - fireEvent.click(oktaOption); + await user().click(oktaOption); await waitFor(() => { expect(screen.getByText("Use Team Mappings")).toBeInTheDocument(); @@ -171,7 +168,7 @@ describe("BaseSSOSettingsForm", () => { it("should show team mappings checkbox for generic provider", async () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("sso-settings"); const handleSubmit = vi.fn(); return ; @@ -179,13 +176,10 @@ describe("BaseSSOSettingsForm", () => { renderWithProviders(); - const providerSelect = screen.getByLabelText("SSO Provider"); - await act(async () => { - fireEvent.mouseDown(providerSelect); - }); + await openProviderDropdown(); const genericOption = await screen.findByText(/generic sso/i); - fireEvent.click(genericOption); + await user().click(genericOption); await waitFor(() => { expect(screen.getByText("Use Team Mappings")).toBeInTheDocument(); @@ -194,7 +188,7 @@ describe("BaseSSOSettingsForm", () => { it("should show team IDs JWT field when use_team_mappings is checked for okta provider", async () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("sso-settings"); const handleSubmit = vi.fn(); return ; @@ -202,22 +196,16 @@ describe("BaseSSOSettingsForm", () => { renderWithProviders(); - const providerSelect = screen.getByLabelText("SSO Provider"); - await act(async () => { - fireEvent.mouseDown(providerSelect); - }); + await openProviderDropdown(); const oktaOption = await screen.findByText(/okta/i); - fireEvent.click(oktaOption); + await user().click(oktaOption); await waitFor(() => { expect(screen.getByText("Use Team Mappings")).toBeInTheDocument(); }); - const checkbox = screen.getByLabelText("Use Team Mappings"); - await act(async () => { - fireEvent.click(checkbox); - }); + await user().click(screen.getAllByLabelText("Use Team Mappings")[0]); await waitFor(() => { expect(screen.getByText("Team IDs JWT Field")).toBeInTheDocument(); @@ -226,7 +214,7 @@ describe("BaseSSOSettingsForm", () => { it("should not show team mappings checkbox for google provider", async () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("sso-settings"); const handleSubmit = vi.fn(); return ; @@ -234,13 +222,10 @@ describe("BaseSSOSettingsForm", () => { renderWithProviders(); - const providerSelect = screen.getByLabelText("SSO Provider"); - await act(async () => { - fireEvent.mouseDown(providerSelect); - }); + await openProviderDropdown(); const googleOption = await screen.findByText(/google sso/i); - fireEvent.click(googleOption); + await user().click(googleOption); await waitFor(() => { expect(screen.getByText("Google Client ID")).toBeInTheDocument(); @@ -297,9 +282,9 @@ describe("renderProviderFields", () => { // to the provider default. Dropping the field from ssoProviderConfigs must // fail here rather than silently in production. const handleSubmit = vi.fn(); - let form: any; + let form!: ReturnType; const TestWrapper = () => { - const [formInstance] = Form.useForm(); + const formInstance = useSSOSettingsForm("sso-settings"); form = formInstance; return ; }; @@ -308,7 +293,8 @@ describe("renderProviderFields", () => { // Mirror EditSSOSettingsModal hydrating the form from the GET response. await act(async () => { - form.setFieldsValue({ + form.reset({ + ...emptySSOSettingsFormValues, sso_provider: "generic", generic_client_id: "client-id", generic_client_secret: "client-secret", @@ -323,8 +309,8 @@ describe("renderProviderFields", () => { // The admin edits something else entirely and saves. await act(async () => { - form.setFieldsValue({ generic_token_endpoint: "https://idp.example.com/token/v2" }); - form.submit(); + form.setValue("generic_token_endpoint", "https://idp.example.com/token/v2"); + submitMountedSSOValues(form, "sso-settings", handleSubmit)(); }); await waitFor(() => { @@ -339,15 +325,13 @@ describe("renderProviderFields", () => { it("renders provider logos in the dropdown and falls back to a letter avatar on load error", async () => { const TestWrapper = () => { - const [form] = Form.useForm(); + const form = useSSOSettingsForm("sso-settings"); return ; }; renderWithProviders(); - await act(async () => { - fireEvent.mouseDown(screen.getByLabelText("SSO Provider")); - }); + await openProviderDropdown(); await waitFor(() => { expect(screen.getAllByAltText("Google SSO logo").length).toBeGreaterThan(0); @@ -372,4 +356,50 @@ describe("renderProviderFields", () => { expect(screen.getByText("O")).toBeInTheDocument(); }); }); + + it("blocks the submit and names the missing provider credential", async () => { + const handleSubmit = vi.fn(); + let form!: ReturnType; + const TestWrapper = () => { + const formInstance = useSSOSettingsForm("sso-settings"); + form = formInstance; + return ; + }; + + renderWithProviders(); + + await act(async () => { + form.reset({ + ...emptySSOSettingsFormValues, + sso_provider: "google", + google_client_secret: "a-secret", + user_email: "admin@example.com", + proxy_base_url: "https://gateway.example.com", + }); + }); + + await act(async () => { + submitMountedSSOValues(form, "sso-settings", handleSubmit)(); + }); + + expect(await screen.findByText("Please enter the google client id")).toBeInTheDocument(); + expect(handleSubmit).not.toHaveBeenCalled(); + }); + + it("shows the effective default role on an untouched form", async () => { + let form!: ReturnType; + const TestWrapper = () => { + const formInstance = useSSOSettingsForm("sso-settings"); + form = formInstance; + return ; + }; + + renderWithProviders(); + + await act(async () => { + form.reset({ ...emptySSOSettingsFormValues, sso_provider: "okta", use_role_mappings: true }); + }); + + expect(await screen.findByLabelText("Default Role")).toHaveTextContent("Internal User"); + }); }); diff --git a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx index 7deac9cbbbc..26d0785529c 100644 --- a/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx +++ b/ui/litellm-dashboard/src/components/Settings/AdminSettings/SSOSettings/Modals/BaseSSOSettingsForm.tsx @@ -1,29 +1,65 @@ "use client"; -import { TextInput } from "@tremor/react"; -import { Checkbox, Form, Input, Select } from "antd"; import React from "react"; +import { FormProvider, useFormContext, useWatch, type UseFormReturn } from "react-hook-form"; +import { z } from "zod/v4"; import { ssoProviderLogoMap, ssoProviderDisplayNames } from "../constants"; import { Logo } from "@/components/molecules/logo/Logo"; +import { FieldGroup } from "@/components/shared/form/field"; +import { FormField } from "@/components/shared/form/FormField"; +import { PasswordInput } from "@/components/shared/PasswordInput"; +import { Checkbox } from "@/components/ui/checkbox"; +import { Input } from "@/components/ui/input"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { Textarea } from "@/components/ui/textarea"; +import { useZodForm } from "@/lib/forms/useZodForm"; -export interface BaseSSOSettingsFormProps { - form: any; // Replace with proper Form type if available - onFormSubmit: (formValues: Record) => Promise; +export interface SSOSettingsFormValues { + sso_provider?: string; + google_client_id?: string; + google_client_secret?: string; + microsoft_client_id?: string; + microsoft_client_secret?: string; + microsoft_tenant?: string; + generic_client_id?: string; + generic_client_secret?: string; + generic_authorization_endpoint?: string; + generic_token_endpoint?: string; + generic_userinfo_endpoint?: string; + generic_scope?: string; + saml_idp_metadata_url?: string; + saml_idp_metadata_xml?: string; + saml_sp_entity_id?: string; + saml_allow_unsolicited?: boolean; + user_email?: string; + proxy_base_url?: string; + use_role_mappings?: boolean; + group_claim?: string; + default_role?: string; + proxy_admin_teams?: string; + admin_viewer_teams?: string; + internal_user_teams?: string; + internal_viewer_teams?: string; + use_team_mappings?: boolean; + team_ids_jwt_field?: string; +} + +export interface BaseSSOSettingsFormProps { + form: UseFormReturn; + onFormSubmit: (formValues: SSOSettingsFormValues) => Promise; } -// Define the SSO provider configuration type export interface SSOProviderConfig { envVarMap: Record; fields: Array<{ label: string; - name: string; + name: keyof SSOSettingsFormValues; placeholder?: string; required?: boolean; type?: "password" | "textarea" | "checkbox"; }>; } -// Define configurations for each SSO provider export const ssoProviderConfigs: Record = { google: { envVarMap: { @@ -128,224 +164,381 @@ export const ssoProviderConfigs: Record = { }, }; -// Helper function to render provider fields +const ROLE_MAPPING_TEAM_FIELDS = [ + "proxy_admin_teams", + "admin_viewer_teams", + "internal_user_teams", + "internal_viewer_teams", +] as const; + +const supportsMappings = (provider: string | undefined): boolean => provider === "okta" || provider === "generic"; + +const providerFieldNames = (provider: string | undefined): readonly string[] => + provider ? ssoProviderConfigs[provider]?.fields.map((field) => field.name) ?? [] : []; + +export type SSOFormVariant = "sso-settings" | "admin-panel"; + +export const mountedSSOFieldNames = (values: SSOSettingsFormValues, variant: SSOFormVariant): readonly string[] => { + const provider = values.sso_provider; + const showMappingToggles = supportsMappings(provider); + const roleFieldsVisible = + variant === "sso-settings" + ? Boolean(values.use_role_mappings) && showMappingToggles + : Boolean(values.use_role_mappings); + const teamFieldsVisible = variant === "sso-settings" && Boolean(values.use_team_mappings) && showMappingToggles; + + return [ + "sso_provider", + ...providerFieldNames(provider), + "user_email", + "proxy_base_url", + ...(showMappingToggles ? ["use_role_mappings"] : []), + ...(roleFieldsVisible ? ["group_claim", "default_role", ...ROLE_MAPPING_TEAM_FIELDS] : []), + ...(variant === "sso-settings" && showMappingToggles ? ["use_team_mappings"] : []), + ...(teamFieldsVisible ? ["team_ids_jwt_field"] : []), + ]; +}; + +export const pickMountedSSOValues = (values: SSOSettingsFormValues, variant: SSOFormVariant): SSOSettingsFormValues => + Object.fromEntries( + mountedSSOFieldNames(values, variant).map((name) => [name, values[name as keyof SSOSettingsFormValues]]), + ); + +export const submitMountedSSOValues = + ( + form: UseFormReturn, + variant: SSOFormVariant, + onFormSubmit: (formValues: SSOSettingsFormValues) => Promise | void, + ) => + () => + void form.handleSubmit((values) => onFormSubmit(pickMountedSSOValues(values, variant)))(); + +const REQUIRED_MESSAGES: Record = { + sso_provider: "Please select an SSO provider", + user_email: "Please enter the email of the proxy admin", + proxy_base_url: "Please enter the proxy base url", + group_claim: "Please enter the group claim", + team_ids_jwt_field: "Please enter the team IDs JWT field", +}; + +const isBlank = (value: unknown): boolean => value === undefined || value === null || value === ""; + +export const buildSSOSettingsSchema = (variant: SSOFormVariant) => + z.custom().superRefine((values, ctx) => { + const mounted = new Set(mountedSSOFieldNames(values, variant)); + + const requireField = (name: string) => { + if (mounted.has(name) && isBlank(values[name as keyof SSOSettingsFormValues])) { + ctx.addIssue({ code: "custom", path: [name], message: REQUIRED_MESSAGES[name] }); + } + }; + + requireField("sso_provider"); + requireField("user_email"); + requireField("group_claim"); + requireField("team_ids_jwt_field"); + + const providerConfig = values.sso_provider ? ssoProviderConfigs[values.sso_provider] : undefined; + providerConfig?.fields.forEach((field) => { + if (field.required === false) return; + if (!isBlank(values[field.name])) return; + ctx.addIssue({ + code: "custom", + path: [field.name], + message: `Please enter the ${field.label.toLowerCase()}`, + }); + }); + + const proxyBaseUrl = values.proxy_base_url; + if (isBlank(proxyBaseUrl)) { + ctx.addIssue({ code: "custom", path: ["proxy_base_url"], message: REQUIRED_MESSAGES.proxy_base_url }); + return; + } + if (!/^https?:\/\/.+/.test(proxyBaseUrl as string)) { + ctx.addIssue({ + code: "custom", + path: ["proxy_base_url"], + message: "URL must start with http:// or https://", + }); + return; + } + if ((proxyBaseUrl as string).endsWith("/")) { + ctx.addIssue({ + code: "custom", + path: ["proxy_base_url"], + message: "URL must not end with a trailing slash", + }); + } + }); + +export const emptySSOSettingsFormValues: SSOSettingsFormValues = { + sso_provider: "", + google_client_id: "", + google_client_secret: "", + microsoft_client_id: "", + microsoft_client_secret: "", + microsoft_tenant: "", + generic_client_id: "", + generic_client_secret: "", + generic_authorization_endpoint: "", + generic_token_endpoint: "", + generic_userinfo_endpoint: "", + user_email: "", + proxy_base_url: "", + default_role: "internal_user", +}; + +export const useSSOSettingsForm = ( + variant: SSOFormVariant, + values?: SSOSettingsFormValues, +): UseFormReturn => + useZodForm(buildSSOSettingsSchema(variant), { + mode: "onChange", + defaultValues: emptySSOSettingsFormValues, + ...(values ? { values } : {}), + }); + +const SSOProviderField = ({ field }: { field: SSOProviderConfig["fields"][number] }) => { + const { control } = useFormContext(); + + if (field.type === "checkbox") { + return ( + + {({ value, onChange, onBlur, id, ...rest }) => ( + + )} + + ); + } + + return ( + + {({ ref, value, ...rest }) => { + const shared = { placeholder: field.placeholder, value: (value as string) ?? "", ...rest }; + if (field.type === "textarea") return