diff --git a/ui/litellm-dashboard/src/app/onboarding/OnboardingForm.test.tsx b/ui/litellm-dashboard/src/app/onboarding/OnboardingForm.test.tsx index 965b0db47c3..d40fc82be07 100644 --- a/ui/litellm-dashboard/src/app/onboarding/OnboardingForm.test.tsx +++ b/ui/litellm-dashboard/src/app/onboarding/OnboardingForm.test.tsx @@ -5,17 +5,14 @@ import { OnboardingForm } from "./OnboardingForm"; const mockUseOnboardingCredentials = vi.fn(); const mockClaimToken = vi.fn(); +const mockJwtDecode = vi.fn(); vi.mock("next/navigation", () => ({ useSearchParams: () => new URLSearchParams("invitation_id=inv-123"), })); vi.mock("jwt-decode", () => ({ - jwtDecode: vi.fn(() => ({ - user_email: "alice@example.com", - user_id: "user-1", - key: "access-tok", - })), + jwtDecode: (...args: unknown[]) => mockJwtDecode(...args), })); vi.mock("@/app/(dashboard)/hooks/onboarding/useOnboarding", () => ({ @@ -23,8 +20,14 @@ vi.mock("@/app/(dashboard)/hooks/onboarding/useOnboarding", () => ({ useClaimOnboardingToken: () => ({ mutate: mockClaimToken, isPending: false }), })); +let currentAuthHeaderName = "Authorization"; + vi.mock("@/components/networking", () => ({ getProxyBaseUrl: vi.fn(() => ""), + setGlobalLitellmHeaderName: (headerName: string) => { + currentAuthHeaderName = headerName; + }, + getGlobalLitellmHeaderName: () => currentAuthHeaderName, })); vi.mock("./OnboardingLoadingView", () => ({ @@ -60,6 +63,12 @@ vi.mock("./OnboardingFormBody", () => ({ describe("OnboardingForm", () => { beforeEach(() => { vi.clearAllMocks(); + currentAuthHeaderName = "Authorization"; + mockJwtDecode.mockReturnValue({ + user_email: "alice@example.com", + user_id: "user-1", + key: "access-tok", + }); document.cookie = "token=; expires=Thu, 01 Jan 1970 00:00:00 GMT; path=/"; document.cookie = "token=; expires=Thu, 01 Jan 1970 00:00:00 GMT; path=/ui"; sessionStorage.clear(); @@ -166,6 +175,62 @@ describe("OnboardingForm", () => { cookieSpy.mockRestore(); }); + it("should send the claim request with the proxy's custom auth header name", async () => { + mockJwtDecode.mockReturnValue({ + user_email: "alice@example.com", + user_id: "user-1", + key: "access-tok", + auth_header_name: "x-litellm-api-key", + }); + mockUseOnboardingCredentials.mockReturnValue({ + data: { token: "fake-jwt-token" }, + isLoading: false, + isError: false, + }); + + // The onboarding page has no session yet, so the only source for a custom + // `litellm_key_header_name` is the onboarding token's own claims; without + // propagating it the claim lands in `Authorization` and the proxy answers + // "Missing onboarding session for invitation link." + const { getGlobalLitellmHeaderName } = await import("@/components/networking"); + const headerNameAtClaimTime: string[] = []; + mockClaimToken.mockImplementation((_params, options) => { + headerNameAtClaimTime.push(getGlobalLitellmHeaderName()); + options.onSuccess({ token: "NEW_USER_TOKEN" }); + }); + + render(); + + await act(async () => { + fireEvent.click(screen.getByRole("button", { name: "Submit" })); + }); + + expect(headerNameAtClaimTime).toEqual(["x-litellm-api-key"]); + }); + + it("should keep the default auth header name when the onboarding token has no custom header claim", async () => { + mockUseOnboardingCredentials.mockReturnValue({ + data: { token: "fake-jwt-token" }, + isLoading: false, + isError: false, + }); + + const { getGlobalLitellmHeaderName } = await import("@/components/networking"); + const headerNameAtClaimTime: string[] = []; + mockClaimToken.mockImplementation((_params, options) => { + headerNameAtClaimTime.push(getGlobalLitellmHeaderName()); + options.onSuccess({ token: "NEW_USER_TOKEN" }); + }); + + render(); + + await act(async () => { + fireEvent.click(screen.getByRole("button", { name: "Submit" })); + }); + + expect(headerNameAtClaimTime).toEqual(["Authorization"]); + }); + it("should show claim error when claim response is missing final token", async () => { mockUseOnboardingCredentials.mockReturnValue({ data: { token: "fake-jwt-token" }, diff --git a/ui/litellm-dashboard/src/app/onboarding/OnboardingForm.tsx b/ui/litellm-dashboard/src/app/onboarding/OnboardingForm.tsx index 23a8bc6725a..54d2711ce5e 100644 --- a/ui/litellm-dashboard/src/app/onboarding/OnboardingForm.tsx +++ b/ui/litellm-dashboard/src/app/onboarding/OnboardingForm.tsx @@ -3,7 +3,7 @@ import React from "react"; import { useSearchParams } from "next/navigation"; import { jwtDecode } from "jwt-decode"; import { useOnboardingCredentials, useClaimOnboardingToken } from "@/app/(dashboard)/hooks/onboarding/useOnboarding"; -import { getProxyBaseUrl } from "@/components/networking"; +import { getProxyBaseUrl, setGlobalLitellmHeaderName } from "@/components/networking"; import { clearTokenCookies, storeLoginToken } from "@/utils/cookieUtils"; import { OnboardingLoadingView } from "./OnboardingLoadingView"; import { OnboardingErrorView } from "./OnboardingErrorView"; @@ -30,12 +30,17 @@ export function OnboardingForm({ variant }: OnboardingFormProps) { const userEmail: string = decoded?.user_email ?? ""; const userId: string | null = decoded?.user_id ?? null; const accessToken: string | null = decoded?.key ?? null; + const authHeaderName: string | null = decoded?.auth_header_name ?? null; const handleSubmit = (formValues: { password: string }) => { if (!accessToken || !userId || !inviteId) return; setClaimError(null); + if (authHeaderName) { + setGlobalLitellmHeaderName(authHeaderName); + } + claimToken( { accessToken, inviteId, userId, password: formValues.password }, {