diff --git a/ui/litellm-dashboard/src/components/SSOModals.test.tsx b/ui/litellm-dashboard/src/components/SSOModals.test.tsx index e9d2389b69e..0b7cbe3f442 100644 --- a/ui/litellm-dashboard/src/components/SSOModals.test.tsx +++ b/ui/litellm-dashboard/src/components/SSOModals.test.tsx @@ -1,25 +1,22 @@ -import { render, fireEvent, waitFor } from "@testing-library/react"; -import { describe, expect, it, beforeAll } from "vitest"; +import { render, fireEvent, waitFor, screen, act } from "@testing-library/react"; +import { describe, expect, it, beforeAll, vi } from "vitest"; import { Form } from "antd"; import SSOModals from "./SSOModals"; import React from "react"; -// Mock window.matchMedia for Ant Design components -beforeAll(() => { - Object.defineProperty(window, "matchMedia", { - writable: true, - value: (query: string) => ({ - matches: false, - media: query, - onchange: null, - addListener: () => {}, // deprecated - removeListener: () => {}, // deprecated - addEventListener: () => {}, - removeEventListener: () => {}, - dispatchEvent: () => true, - }), - }); -}); +// Mock the networking functions +vi.mock("./networking", () => ({ + getSSOSettings: vi.fn(), + updateSSOSettings: vi.fn(), +})); + +// Mock parseErrorMessage +vi.mock("./shared/errorUtils", () => ({ + parseErrorMessage: vi.fn((error) => error?.message || "An error occurred"), +})); + +import { getSSOSettings, updateSSOSettings } from "./networking"; +import NotificationsManager from "./molecules/notifications_manager"; describe("SSOModals", () => { it("should render the SSOModals component", () => { @@ -42,11 +39,11 @@ describe("SSOModals", () => { ); }; - const { getByText } = render(); - expect(getByText("Add SSO")).toBeInTheDocument(); + render(); + expect(screen.getByText("Add SSO")).toBeInTheDocument(); }); - it("should have a validation error if the proxy base url is not a valid URL", async () => { + it("should show validation error if proxy base url is not a valid URL", async () => { const TestWrapper = () => { const [form] = Form.useForm(); return ( @@ -65,42 +62,40 @@ describe("SSOModals", () => { ); }; - const { getByLabelText, getByText, container } = render(); + render(); // Find and interact with the SSO provider select - const ssoProviderSelect = container.querySelector("#sso_provider"); - if (ssoProviderSelect) { - fireEvent.mouseDown(ssoProviderSelect); - // Wait for dropdown and select Google - await waitFor(() => { - const googleOption = getByText("Google SSO"); - fireEvent.click(googleOption); - }); - } + const ssoProviderSelect = screen.getByLabelText("SSO Provider"); + fireEvent.mouseDown(ssoProviderSelect); + // Wait for dropdown and select Google + await waitFor(() => { + const googleOption = screen.getByText("Google SSO"); + fireEvent.click(googleOption); + }); // Fill in the email field - const emailInput = getByLabelText("Proxy Admin Email"); + const emailInput = screen.getByLabelText("Proxy Admin Email"); fireEvent.change(emailInput, { target: { value: "test@example.com" } }); // Fill in an invalid URL - const urlInput = getByLabelText("Proxy Base URL"); + const urlInput = screen.getByLabelText("Proxy Base URL"); fireEvent.change(urlInput, { target: { value: "invalid-url" } }); // Submit the form - const saveButton = getByText("Save"); + const saveButton = screen.getByText("Save"); fireEvent.click(saveButton); // Check for validation error await waitFor( () => { - expect(getByText("URL must start with http:// or https://")).toBeInTheDocument(); + expect(screen.getByText("URL must start with http:// or https://")).toBeInTheDocument(); }, // The validation is based on a Promise, so we need to wait for it to resolve { timeout: 5000 }, ); }); - it("should show validation error if the proxy base url ends with a trailing slash", async () => { + it("should show validation error if proxy base url ends with trailing slash", async () => { const TestWrapper = () => { const [form] = Form.useForm(); return ( @@ -119,33 +114,31 @@ describe("SSOModals", () => { ); }; - const { getByLabelText, getByText, findByText, container } = render(); + render(); // Find and interact with the SSO provider select - const ssoProviderSelect = container.querySelector("#sso_provider"); - if (ssoProviderSelect) { - fireEvent.mouseDown(ssoProviderSelect); - // Wait for dropdown and select Google - await waitFor(() => { - const googleOption = getByText("Google SSO"); - fireEvent.click(googleOption); - }); - } + const ssoProviderSelect = screen.getByLabelText("SSO Provider"); + fireEvent.mouseDown(ssoProviderSelect); + // Wait for dropdown and select Google + await waitFor(() => { + const googleOption = screen.getByText("Google SSO"); + fireEvent.click(googleOption); + }); // Fill in the email field - const emailInput = getByLabelText("Proxy Admin Email"); + const emailInput = screen.getByLabelText("Proxy Admin Email"); fireEvent.change(emailInput, { target: { value: "test@example.com" } }); // Fill in a URL with trailing slash - const urlInput = getByLabelText("Proxy Base URL") as HTMLInputElement; + const urlInput = screen.getByLabelText("Proxy Base URL") as HTMLInputElement; fireEvent.change(urlInput, { target: { value: "https://example.com/" } }); // Submit the form - const saveButton = getByText("Save"); + const saveButton = screen.getByText("Save"); fireEvent.click(saveButton); // Check for validation error using findByText for async rendering - const errorMessage = await findByText("URL must not end with a trailing slash", {}, { timeout: 5000 }); + const errorMessage = await screen.findByText("URL must not end with a trailing slash", {}, { timeout: 5000 }); expect(errorMessage).toBeInTheDocument(); }); @@ -168,9 +161,9 @@ describe("SSOModals", () => { ); }; - const { getByLabelText } = render(); + render(); - const urlInput = getByLabelText("Proxy Base URL") as HTMLInputElement; + const urlInput = screen.getByLabelText("Proxy Base URL") as HTMLInputElement; // Simulate user typing "https://" fireEvent.change(urlInput, { target: { value: "h" } }); @@ -218,36 +211,266 @@ describe("SSOModals", () => { ); }; - const { getByLabelText, getByText, queryByText, container, findByText } = render(); + render(); // Find and interact with the SSO provider select - const ssoProviderSelect = container.querySelector("#sso_provider"); - if (ssoProviderSelect) { - fireEvent.mouseDown(ssoProviderSelect); - // Wait for dropdown and select Google - await waitFor(() => { - const googleOption = getByText("Google SSO"); - fireEvent.click(googleOption); - }); - } + const ssoProviderSelect = screen.getByLabelText("SSO Provider"); + fireEvent.mouseDown(ssoProviderSelect); + // Wait for dropdown and select Google + await waitFor(() => { + const googleOption = screen.getByText("Google SSO"); + fireEvent.click(googleOption); + }); // Fill in the email field - const emailInput = getByLabelText("Proxy Admin Email"); + const emailInput = screen.getByLabelText("Proxy Admin Email"); fireEvent.change(emailInput, { target: { value: "test@example.com" } }); // Fill in an incomplete URL like "http:" - const urlInput = getByLabelText("Proxy Base URL"); + const urlInput = screen.getByLabelText("Proxy Base URL"); fireEvent.change(urlInput, { target: { value: "http:" } }); // Submit the form - const saveButton = getByText("Save"); + const saveButton = screen.getByText("Save"); fireEvent.click(saveButton); // Check that only the URL format error appears (use findByText for async rendering) - const errorMessage = await findByText("URL must start with http:// or https://", {}, { timeout: 3000 }); + const errorMessage = await screen.findByText("URL must start with http:// or https://", {}, { timeout: 3000 }); expect(errorMessage).toBeInTheDocument(); // Verify the trailing slash error does NOT appear - expect(queryByText("URL must not end with a trailing slash")).not.toBeInTheDocument(); + expect(screen.queryByText("URL must not end with a trailing slash")).not.toBeInTheDocument(); + }); + + it("should load existing SSO settings when modal opens", async () => { + const mockSSOData = { + values: { + google_client_id: "test-client-id", + google_client_secret: "test-client-secret", + proxy_base_url: "https://example.com", + user_email: "admin@example.com", + role_mappings: { + group_claim: "groups", + default_role: "internal_user", + roles: { + proxy_admin: ["admin-group"], + proxy_admin_viewer: ["viewer-group"], + internal_user: ["user-group"], + internal_user_viewer: ["readonly-group"], + }, + }, + }, + }; + + (getSSOSettings as any).mockResolvedValue(mockSSOData); + + const TestWrapper = () => { + const [form] = Form.useForm(); + + return ( + {}} + handleAddSSOCancel={() => {}} + handleShowInstructions={() => {}} + handleInstructionsOk={() => {}} + handleInstructionsCancel={() => {}} + form={form} + accessToken="test-token" + ssoConfigured={false} + /> + ); + }; + + render(); + + // Wait for the useEffect to load data and populate form + await waitFor(() => { + expect(getSSOSettings).toHaveBeenCalledWith("test-token"); + }); + + // Check that form fields are populated with loaded data + await waitFor(() => { + const emailInput = screen.getByLabelText("Proxy Admin Email") as HTMLInputElement; + expect(emailInput.value).toBe("admin@example.com"); + }); + + const urlInput = screen.getByLabelText("Proxy Base URL") as HTMLInputElement; + expect(urlInput.value).toBe("https://example.com"); + + // Check that role mappings are populated + const groupClaimInput = screen.getByLabelText("Group Claim") as HTMLInputElement; + expect(groupClaimInput.value).toBe("groups"); + }); + + it("should submit form with role mappings enabled", async () => { + const mockHandleShowInstructions = vi.fn(); + (updateSSOSettings as any).mockResolvedValue({}); + // Mock getSSOSettings to return empty data so form starts clean + (getSSOSettings as any).mockResolvedValue({ values: {} }); + + let formInstance: any = null; + + const TestWrapper = () => { + const [form] = Form.useForm(); + formInstance = form; + + return ( + {}} + handleAddSSOCancel={() => {}} + handleShowInstructions={mockHandleShowInstructions} + handleInstructionsOk={() => {}} + handleInstructionsCancel={() => {}} + form={form} + accessToken="test-token" + ssoConfigured={false} + /> + ); + }; + + render(); + + // Wait for any initial loading to complete + await waitFor(() => { + expect(getSSOSettings).toHaveBeenCalledWith("test-token"); + }); + + // Set the provider directly using the form to trigger conditional rendering + formInstance.setFieldsValue({ sso_provider: "okta" }); + + // Wait for the "Use Role Mappings" checkbox to appear + await waitFor(() => { + expect(screen.getByLabelText("Use Role Mappings")).toBeInTheDocument(); + }); + + // Enable role mappings + const roleMappingsCheckbox = screen.getByLabelText("Use Role Mappings"); + fireEvent.click(roleMappingsCheckbox); + + // Fill required fields + const emailInput = screen.getByLabelText("Proxy Admin Email"); + fireEvent.change(emailInput, { target: { value: "admin@example.com" } }); + + const urlInput = screen.getByLabelText("Proxy Base URL"); + fireEvent.change(urlInput, { target: { value: "https://example.com" } }); + + // Fill Okta specific fields + const clientIdInput = screen.getByLabelText("Generic Client ID"); + fireEvent.change(clientIdInput, { target: { value: "test-client-id" } }); + + const clientSecretInput = screen.getByLabelText("Generic Client Secret"); + fireEvent.change(clientSecretInput, { target: { value: "test-client-secret" } }); + + const authEndpointInput = screen.getByLabelText("Authorization Endpoint"); + fireEvent.change(authEndpointInput, { target: { value: "https://example.okta.com/authorize" } }); + + const tokenEndpointInput = screen.getByLabelText("Token Endpoint"); + fireEvent.change(tokenEndpointInput, { target: { value: "https://example.okta.com/token" } }); + + const userinfoEndpointInput = screen.getByLabelText("Userinfo Endpoint"); + fireEvent.change(userinfoEndpointInput, { target: { value: "https://example.okta.com/userinfo" } }); + + // Fill role mapping fields + const groupClaimInput = screen.getByLabelText("Group Claim"); + fireEvent.change(groupClaimInput, { target: { value: "groups" } }); + + const proxyAdminTeamsInput = screen.getByLabelText("Proxy Admin Teams"); + fireEvent.change(proxyAdminTeamsInput, { target: { value: "admin-group, super-admin" } }); + + // Submit the form + const saveButton = screen.getByText("Save"); + fireEvent.click(saveButton); + + // Verify the API was called with correct payload including role mappings + await waitFor(() => { + expect(updateSSOSettings).toHaveBeenCalledWith("test-token", { + sso_provider: "okta", + user_email: "admin@example.com", + proxy_base_url: "https://example.com", + generic_client_id: "test-client-id", + generic_client_secret: "test-client-secret", + generic_authorization_endpoint: "https://example.okta.com/authorize", + generic_token_endpoint: "https://example.okta.com/token", + generic_userinfo_endpoint: "https://example.okta.com/userinfo", + role_mappings: { + provider: "generic", + group_claim: "groups", + default_role: "internal_user", + roles: { + proxy_admin: ["admin-group", "super-admin"], + proxy_admin_viewer: [], + internal_user: [], + internal_user_viewer: [], + }, + }, + }); + }); + + expect(mockHandleShowInstructions).toHaveBeenCalled(); + }); + + it("should show Clear button and clear SSO settings when configured", async () => { + const mockHandleAddSSOOk = vi.fn(); + (updateSSOSettings as any).mockResolvedValue({}); + (NotificationsManager.success as any).mockImplementation(() => {}); + + const TestWrapper = () => { + const [form] = Form.useForm(); + + return ( + {}} + handleShowInstructions={() => {}} + handleInstructionsOk={() => {}} + handleInstructionsCancel={() => {}} + form={form} + accessToken="test-token" + ssoConfigured={true} + /> + ); + }; + + render(); + + // Check that Clear button is visible when SSO is configured + const clearButton = screen.getByText("Clear"); + expect(clearButton).toBeInTheDocument(); + + // Click Clear button to open confirmation modal + fireEvent.click(clearButton); + + // Confirm the clear action in the modal + const confirmButton = screen.getByText("Yes, Clear"); + fireEvent.click(confirmButton); + + // Verify the clear API was called with null values + await waitFor(() => { + expect(updateSSOSettings).toHaveBeenCalledWith("test-token", { + google_client_id: null, + google_client_secret: null, + microsoft_client_id: null, + microsoft_client_secret: null, + microsoft_tenant: null, + generic_client_id: null, + generic_client_secret: null, + generic_authorization_endpoint: null, + generic_token_endpoint: null, + generic_userinfo_endpoint: null, + proxy_base_url: null, + user_email: null, + sso_provider: null, + role_mappings: null, + }); + }); + + expect(NotificationsManager.success).toHaveBeenCalledWith("SSO settings cleared successfully"); + expect(mockHandleAddSSOOk).toHaveBeenCalled(); }); }); diff --git a/ui/litellm-dashboard/src/components/SSOModals.tsx b/ui/litellm-dashboard/src/components/SSOModals.tsx index 26e33ace2d7..6cb57f41736 100644 --- a/ui/litellm-dashboard/src/components/SSOModals.tsx +++ b/ui/litellm-dashboard/src/components/SSOModals.tsx @@ -1,5 +1,5 @@ import React, { useEffect, useState } from "react"; -import { Modal, Form, Input, Button as Button2, Select } from "antd"; +import { Modal, Form, Input, Button as Button2, Select, Checkbox } from "antd"; import { Text, TextInput } from "@tremor/react"; import { getSSOSettings, updateSSOSettings } from "./networking"; import NotificationsManager from "./molecules/notifications_manager"; @@ -144,12 +144,35 @@ const SSOModals: React.FC = ({ } } + // Extract role mappings if they exist + let roleMappingFields = {}; + if (ssoData.values.role_mappings) { + const roleMappings = ssoData.values.role_mappings; + + // Helper function to join arrays into comma-separated strings + const joinTeams = (teams: string[] | undefined): string => { + if (!teams || teams.length === 0) return ""; + return teams.join(", "); + }; + + roleMappingFields = { + use_role_mappings: true, + group_claim: roleMappings.group_claim, + default_role: roleMappings.default_role || "internal_user", + proxy_admin_teams: joinTeams(roleMappings.roles?.proxy_admin), + admin_viewer_teams: joinTeams(roleMappings.roles?.proxy_admin_viewer), + internal_user_teams: joinTeams(roleMappings.roles?.internal_user), + internal_viewer_teams: joinTeams(roleMappings.roles?.internal_user_viewer), + }; + } + // Set form values with existing data (excluding UI access control fields) const formValues = { sso_provider: selectedProvider, proxy_base_url: ssoData.values.proxy_base_url, user_email: ssoData.values.user_email, ...ssoData.values, + ...roleMappingFields, }; console.log("Setting form values:", formValues); // Debug log @@ -178,8 +201,55 @@ const SSOModals: React.FC = ({ } try { + const { + proxy_admin_teams, + admin_viewer_teams, + internal_user_teams, + internal_viewer_teams, + default_role, + group_claim, + use_role_mappings, + ...rest + } = formValues; + + const payload: any = { + ...rest, + }; + + // Add role mappings if use_role_mappings is checked + if (use_role_mappings) { + // Helper function to split comma-separated string into array + const splitTeams = (teams: string | undefined): string[] => { + if (!teams || teams.trim() === "") return []; + return teams + .split(",") + .map((team) => team.trim()) + .filter((team) => team.length > 0); + }; + + // Map default role display values to backend values + const defaultRoleMapping: Record = { + internal_user_viewer: "internal_user_viewer", + internal_user: "internal_user", + proxy_admin_viewer: "proxy_admin_viewer", + proxy_admin: "proxy_admin", + }; + + payload.role_mappings = { + provider: "generic", + group_claim, + default_role: defaultRoleMapping[default_role] || "internal_user", + roles: { + proxy_admin: splitTeams(proxy_admin_teams), + proxy_admin_viewer: splitTeams(admin_viewer_teams), + internal_user: splitTeams(internal_user_teams), + internal_user_viewer: splitTeams(internal_viewer_teams), + }, + }; + } + // Save SSO settings using the new API - await updateSSOSettings(accessToken, formValues); + await updateSSOSettings(accessToken, payload); // Continue with the original flow (show instructions) handleShowInstructions(formValues); @@ -211,6 +281,7 @@ const SSOModals: React.FC = ({ proxy_base_url: null, user_email: null, sso_provider: null, + role_mappings: null, }; await updateSSOSettings(accessToken, clearSettings); @@ -334,6 +405,79 @@ const SSOModals: React.FC = ({ > + + 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; + }} +