diff --git a/litellm/types/proxy/public_endpoints/public_endpoints.py b/litellm/types/proxy/public_endpoints/public_endpoints.py index ea6c69213c7..d0a226c2af4 100644 --- a/litellm/types/proxy/public_endpoints/public_endpoints.py +++ b/litellm/types/proxy/public_endpoints/public_endpoints.py @@ -1,9 +1,11 @@ from collections.abc import Mapping, Sequence from typing import Any, Final, Literal -from pydantic import BaseModel, ConfigDict, Field, model_validator +from pydantic import BaseModel, ConfigDict, Field, computed_field, model_validator from typing_extensions import Self +from litellm.types.router import server_owned_wif_fields_named + class PublicModelHubInfo(BaseModel): docs_title: str @@ -34,7 +36,10 @@ class ProviderCredentialVariant(BaseModel): ``anthropic_identity_source: keycloak``) and are submitted without a form field for them. ``optional_field_keys`` relaxes a globally-required field for this variant alone, for a value only obtainable after the credential exists (the federation rule id an operator can only read - off the Anthropic Console once the generated JWKS is registered).""" + off the Anthropic Console once the generated JWKS is registered). + ``credential_only`` is derived, never declared: a variant that submits a server-owned workload + identity federation parameter, as a field or a fixed value, can only be saved as a named LLM + Credential, since ``/model/new`` rejects those parameters inline.""" id: str label: str @@ -42,6 +47,11 @@ class ProviderCredentialVariant(BaseModel): optional_field_keys: tuple[str, ...] = () fixed_values: Mapping[str, str] = Field(default_factory=dict) + @computed_field + @property + def credential_only(self) -> bool: + return bool(server_owned_wif_fields_named((*self.field_keys, *self.fixed_values))) + def _validate_variant(variant: "ProviderCredentialVariant", defined_keys: frozenset[str]) -> None: """Each variant may only reference declared fields, may only relax fields it actually mounts, and a diff --git a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py index d878ce9f3e3..c81d7904f8b 100644 --- a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py @@ -983,8 +983,8 @@ def test_public_mcp_hub_returns_only_whitelisted_servers(): litellm.public_mcp_servers, mirroring /public/model_hub and /public/agent_hub. Servers with available_on_public_internet=True that are not on the whitelist must not leak.""" - from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer app = FastAPI() app.include_router(router) @@ -1045,8 +1045,8 @@ def test_public_mcp_hub_returns_empty_when_whitelist_unset(): def test_public_mcp_hub_does_not_expose_upstream_url(): """Regression: /public/mcp_hub is unauthenticated, so the gateway-internal upstream url must never appear in its response even when the server has one.""" - from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.proxy._types import MCPTransport + from litellm.types.mcp_server.mcp_server_manager import MCPServer app = FastAPI() app.include_router(router) @@ -1559,3 +1559,46 @@ async def test_fetch_remote_autorouter_presets_parses_and_rejects_empty(monkeypa response.json = MagicMock(return_value={}) with pytest.raises(ValueError, match="empty"): await _fetch_remote_autorouter_presets("https://example.test/presets.json") + + +def test_credential_only_marks_variants_carrying_server_owned_wif_params(): + """A variant is credential-only exactly when a field or a fixed value it submits is a + server-owned workload identity federation parameter, which /model/new refuses inline.""" + variants = ProviderCredentialVariants( + selector_label="Authentication method", + default_variant="api_key", + field_definitions=( + ProviderCredentialField(key="api_base", label="API Base"), + ProviderCredentialField(key="api_key", label="API Key"), + ProviderCredentialField(key="openai_identity_provider_id", label="Identity Provider ID"), + ), + variants=( + ProviderCredentialVariant(id="api_key", label="API Key", field_keys=("api_base", "api_key")), + ProviderCredentialVariant( + id="wif_token_file", label="Federation", field_keys=("openai_identity_provider_id",) + ), + ProviderCredentialVariant( + id="wif_keycloak", + label="Federation (Keycloak)", + field_keys=("api_base",), + fixed_values={"anthropic_identity_source": "keycloak"}, + ), + ), + ) + expected = {"api_key": False, "wif_token_file": True, "wif_keycloak": True} + assert {variant.id: variant.credential_only for variant in variants.variants} == expected + assert {v["id"]: v["credential_only"] for v in variants.model_dump()["variants"]} == expected + + +def test_provider_fields_flag_every_federation_variant_credential_only(): + """The dashboard reads credential_only off /public/providers/fields to swap a federation + variant's inline fields for the create-credential step, so every federation variant must + carry the flag and the API-key variant must not.""" + app_instance = FastAPI() + app_instance.include_router(router) + providers = TestClient(app_instance).get("/public/providers/fields").json() + for provider_name in ("OpenAI", "Anthropic"): + provider = next(p for p in providers if p["provider"] == provider_name) + flags = {v["id"]: v["credential_only"] for v in provider["credential_variants"]["variants"]} + assert flags.pop("api_key") is False, provider_name + assert flags and all(flags.values()), (provider_name, flags) diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/credentials/useCredentials.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/credentials/useCredentials.ts index e3266de4fbc..bdbf4445514 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/hooks/credentials/useCredentials.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/credentials/useCredentials.ts @@ -3,7 +3,7 @@ import { useQuery } from "@tanstack/react-query"; import { createQueryKeys } from "../common/queryKeysFactory"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; -const credentialsKeys = createQueryKeys("credentials"); +export const credentialsKeys = createQueryKeys("credentials"); export const useCredentials = () => { const { accessToken } = useAuthorized(); diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx index 3b879da2d68..171b6b6d4cd 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.test.tsx @@ -1,8 +1,8 @@ -import { renderHook, screen, waitFor, renderWithProviders } from "../../../tests/test-utils"; +import { renderHook, screen, waitFor, within, renderWithProviders } from "../../../tests/test-utils"; import userEvent, { PointerEventsCheckLevel } from "@testing-library/user-event"; import { describe, expect, it, vi } from "vitest"; import type { Team } from "../key_team_helpers/key_list"; -import type { CredentialItem } from "../networking"; +import { credentialCreateCall, type CredentialItem } from "../networking"; import { Providers } from "../provider_info_helpers"; import { projectMountedValues, useMountRegistry, type MountedFormValues } from "../common_components/MountedFormField"; import { useForm } from "react-hook-form"; @@ -34,6 +34,9 @@ vi.mock("../networking", async () => { ], }), testConnectionRequest: vi.fn().mockResolvedValue({ status: "success" }), + credentialCreateCall: vi.fn().mockResolvedValue({ success: true }), + discoverProviderModelsCall: vi.fn().mockResolvedValue({ models: [] }), + getCredentialJwksCall: vi.fn().mockResolvedValue({ keys: [] }), getProviderCreateMetadata: vi.fn().mockResolvedValue([ { provider: "OpenAI", @@ -55,6 +58,31 @@ vi.mock("@/app/(dashboard)/hooks/providers/useProviderFields", () => ({ litellm_provider: "openai", default_model_placeholder: "gpt-3.5-turbo", credential_fields: [], + credential_variants: { + selector_label: "Authentication method", + default_variant: "api_key", + field_definitions: [ + { key: "api_key", label: "OpenAI API Key", field_type: "password" }, + { key: "openai_identity_provider_id", label: "Identity Provider ID", field_type: "text", required: true }, + { key: "openai_service_account_id", label: "Service Account ID", field_type: "text", required: true }, + { + key: "openai_identity_token_file", + label: "Identity Token File Path", + field_type: "text", + required: true, + }, + ], + variants: [ + { id: "api_key", label: "API Key", field_keys: ["api_key"], fixed_values: {}, credential_only: false }, + { + id: "wif_token_file", + label: "Workload Identity Federation (token file)", + field_keys: ["openai_identity_provider_id", "openai_service_account_id", "openai_identity_token_file"], + fixed_values: {}, + credential_only: true, + }, + ], + }, }, ], isLoading: false, @@ -66,6 +94,8 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({ default: vi.fn(), })); +vi.mock("@/lib/toast", () => ({ toast: { success: vi.fn(), error: vi.fn(), info: vi.fn(), warning: vi.fn() } })); + vi.mock("@/app/(dashboard)/hooks/teams/useTeams", () => ({ useInfiniteTeams: () => ({ data: { @@ -431,4 +461,66 @@ describe("AddModelForm", () => { expect(values).not.toHaveProperty("cache_control_injection_points"); }); }); + + describe("credential-only auth variants", () => { + const WIF_KEYS = ["openai_identity_provider_id", "openai_service_account_id", "openai_identity_token_file"]; + + const pickWifVariant = async () => { + const mockUseAuthorized = vi.mocked(await import("@/app/(dashboard)/hooks/useAuthorized")); + mockUseAuthorized.default.mockReturnValue(mockAuthorizedUser("proxy_admin", "user-1", true)); + const props = createTestProps(); + renderWithProviders(); + const user = userEvent.setup({ pointerEventsCheck: PointerEventsCheckLevel.Never }); + await screen.findByText("Provider"); + await user.click(await screen.findByRole("combobox", { name: "Authentication method" })); + await user.click(await screen.findByRole("option", { name: "Workload Identity Federation (token file)" })); + return { user, props }; + }; + + it("never mounts a credential-only variant's fields on the deployment", async () => { + const { props } = await pickWifVariant(); + + expect(await screen.findByRole("button", { name: "Create credential" })).toBeInTheDocument(); + expect(screen.queryByLabelText("Identity Provider ID")).not.toBeInTheDocument(); + expect(screen.queryByLabelText("OpenAI API Key")).not.toBeInTheDocument(); + const values = props.mountedValues(); + for (const key of [...WIF_KEYS, "api_key"]) { + expect(values).not.toHaveProperty(key); + } + }); + + it("saves the federation settings as a credential and attaches it to the model", async () => { + const { user, props } = await pickWifVariant(); + + await user.click(await screen.findByRole("button", { name: "Create credential" })); + const dialog = await screen.findByRole("dialog"); + expect(within(dialog).getByText("Add New Credential")).toBeInTheDocument(); + await user.type(within(dialog).getByLabelText("Credential Name:"), "openai-wif"); + await user.type(await within(dialog).findByLabelText("Identity Provider ID"), "idp_123"); + await user.type(within(dialog).getByLabelText("Service Account ID"), "user-456"); + await user.type(within(dialog).getByLabelText("Identity Token File Path"), "/var/run/token"); + await user.click(within(dialog).getByRole("button", { name: "Add Credential" })); + + await waitFor(() => + expect(vi.mocked(credentialCreateCall)).toHaveBeenCalledWith("test-access-token", { + credential_name: "openai-wif", + credential_values: { + openai_identity_provider_id: "idp_123", + openai_service_account_id: "user-456", + openai_identity_token_file: "/var/run/token", + }, + credential_info: { custom_llm_provider: "OpenAI" }, + }), + ); + await waitFor(() => expect(props.form.getValues("litellm_credential_name")).toBe("openai-wif")); + await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument()); + + const values = props.mountedValues(); + expect(values.litellm_credential_name).toBe("openai-wif"); + for (const key of WIF_KEYS) { + expect(values).not.toHaveProperty(key); + } + expect(screen.queryByRole("combobox", { name: "Authentication method" })).not.toBeInTheDocument(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx index 92951b68fd5..76f9b2f3aa6 100644 --- a/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx +++ b/ui/litellm-dashboard/src/components/add_model/AddModelForm.tsx @@ -24,7 +24,16 @@ import { type MountedFormValues, } from "../common_components/MountedFormField"; import type { Team } from "../key_team_helpers/key_list"; -import { type CredentialItem, type ProviderCreateInfo, modelAvailableCall } from "../networking"; +import { + type CredentialItem, + type ProviderCreateInfo, + credentialCreateCall, + discoverProviderModelsCall, + getCredentialJwksCall, + modelAvailableCall, +} from "../networking"; +import CredentialModal from "../model_add/CredentialModal"; +import { buildCredential, withoutRestrictedFields } from "../model_add/credential_form_helpers"; import { Providers } from "../provider_info_helpers"; import { ProviderLogo } from "../molecules/models/ProviderLogo"; import AccessGroupTagsCombobox from "./AccessGroupTagsCombobox"; @@ -35,6 +44,10 @@ import ConnectionErrorDisplay from "./model_connection_test"; import ProviderSpecificFields from "./provider_specific_fields"; import { TEST_MODES } from "./add_model_modes"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; +import { credentialsKeys } from "@/app/(dashboard)/hooks/credentials/useCredentials"; +import { toast } from "@/lib/toast"; +import { extractProxyErrorMessage } from "@/lib/http/client"; +import { useQueryClient } from "@tanstack/react-query"; import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "@/components/ui/dialog"; interface AddModelFormProps { @@ -92,6 +105,23 @@ const AddModelForm: React.FC = ({ const guardrailsList = guardrailsData?.guardrails.map((g) => g.guardrail_name); const { data: tagsList } = useTags(); const selectedCredentialName = useWatch({ control: form.control, name: "litellm_credential_name" }); + const queryClient = useQueryClient(); + const [credentialDraftVariantId, setCredentialDraftVariantId] = useState(null); + + const handleCreateCredential = async (values: Record): Promise => { + const credential = buildCredential(values, withoutRestrictedFields(values)); + try { + await credentialCreateCall(accessToken, credential); + } catch (error) { + toast.error(extractProxyErrorMessage(error)); + return false; + } + toast.success("Credential added successfully"); + setCredentialDraftVariantId(null); + await queryClient.invalidateQueries({ queryKey: credentialsKeys.all }); + form.setValue("litellm_credential_name", credential.credential_name, { shouldDirty: true }); + return true; + }; const handleTestConnection = async () => { setIsTestingConnection(true); @@ -317,7 +347,10 @@ const AddModelForm: React.FC = ({ OR
- + )}
@@ -448,6 +481,18 @@ const AddModelForm: React.FC = ({ + {credentialDraftVariantId !== null && ( + setCredentialDraftVariantId(null)} + onSubmit={handleCreateCredential} + testConnection={(request) => discoverProviderModelsCall(accessToken, request)} + loadJwks={(credentialName) => getCredentialJwksCall(accessToken, credentialName)} + /> + )} {/* Test Connection Results Modal */} { it("drops a field_key that has no matching field_definitions entry, rather than crashing", () => { const variants: ProviderCredentialVariants = { ...anthropicVariants, - variants: [{ id: "broken", label: "Broken", field_keys: ["api_key", "missing_field"], fixed_values: {} }], + variants: [ + { + id: "broken", + label: "Broken", + field_keys: ["api_key", "missing_field"], + fixed_values: {}, + credential_only: false, + }, + ], }; expect(resolveVariantFieldDefs(variants, "broken").map((f) => f.key)).toEqual(["api_key"]); }); diff --git a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.test.tsx b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.test.tsx index eada05d4f04..412e5e8afe0 100644 --- a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.test.tsx @@ -148,12 +148,19 @@ vi.mock("../networking", async () => { }, ], variants: [ - { id: "api_key", label: "API Key", field_keys: ["api_base", "api_key"], fixed_values: {} }, + { + id: "api_key", + label: "API Key", + field_keys: ["api_base", "api_key"], + fixed_values: {}, + credential_only: false, + }, { id: "wif_token", label: "Workload Identity Federation (external token)", field_keys: ["anthropic_federation_rule_id", "anthropic_organization_id", "anthropic_identity_token"], fixed_values: {}, + credential_only: true, }, { id: "wif_internal_issuer", @@ -165,6 +172,7 @@ vi.mock("../networking", async () => { "anthropic_issuer_signing_key_ref", ], fixed_values: { anthropic_identity_source: "internal_issuer" }, + credential_only: true, }, ], }, @@ -586,5 +594,57 @@ describe("ProviderSpecificFields", () => { expect(await screen.findByLabelText("Identity Token Reference")).toBeInTheDocument(); expect(screen.queryByLabelText("API Key")).not.toBeInTheDocument(); }); + + it("opens on the variant the caller hands in", async () => { + render( + + + + + , + ); + + expect(await screen.findByLabelText("Federation Rule ID")).toBeInTheDocument(); + expect(screen.queryByLabelText("API Key")).not.toBeInTheDocument(); + }); + + it("swaps a credential-only variant's fields for the create-credential step on a deployment form", async () => { + const onCreateCredential = vi.fn(); + render( + + + + + + , + ); + + const user = userEvent.setup(); + await screen.findByLabelText("API Key"); + await user.click(await screen.findByRole("combobox", { name: "Authentication method" })); + await user.click(await screen.findByRole("option", { name: "Workload Identity Federation (LiteLLM-signed)" })); + + const createButton = await screen.findByRole("button", { name: "Create credential" }); + expect(screen.queryByLabelText("Issuer URL")).not.toBeInTheDocument(); + expect(screen.queryByLabelText("Federation Rule ID")).not.toBeInTheDocument(); + expect(screen.queryByLabelText("API Key")).not.toBeInTheDocument(); + expect(screen.getByTestId("identity-source")).toBeEmptyDOMElement(); + + await user.click(createButton); + expect(onCreateCredential).toHaveBeenCalledWith("wif_internal_issuer"); + }); + + it("keeps the api_key variant inline on a deployment form", async () => { + render( + + + + + , + ); + + expect(await screen.findByLabelText("API Key")).toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Create credential" })).not.toBeInTheDocument(); + }); }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx index 3fc92fa5249..38acbfd9422 100644 --- a/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx +++ b/ui/litellm-dashboard/src/components/add_model/provider_specific_fields.tsx @@ -14,16 +14,41 @@ import { type MountedFieldControlProps, type MountedFormValues, } from "../common_components/MountedFormField"; -import { CredentialItem, ProviderCredentialFieldMetadata, ProviderCredentialVariants } from "../networking"; +import { + CredentialItem, + ProviderCredentialFieldMetadata, + ProviderCredentialVariant, + ProviderCredentialVariants, +} from "../networking"; import { provider_map, Providers } from "../provider_info_helpers"; import { labelWithHint } from "@/components/shared/form/LabelWithHint"; import { Field, FieldLabel } from "@/components/ui/field"; +import { Alert, AlertDescription } from "@/components/ui/alert"; import { getVariant, inferActiveVariant, resolveVariantFieldDefs } from "./provider_credential_variants"; interface ProviderSpecificFieldsProps { selectedProvider: Providers; + initialVariantId?: string; + onCreateCredential?: (variantId: string) => void; } +const CredentialOnlyNotice: React.FC<{ + variant: ProviderCredentialVariant; + onCreateCredential: (variantId: string) => void; +}> = ({ variant, onCreateCredential }) => ( + + + + {variant.label} uses server-owned identity settings, which are stored as an LLM Credential rather than entered + on the model. Create the credential and this model will use it. + + + + +); + const readTextFile = (file: File, onLoaded: (contents: string) => void) => { const reader = new FileReader(); reader.onload = (event) => { @@ -141,7 +166,11 @@ const FixedValueField: React.FC<{ name: string; value: string }> = ({ name, valu return null; }; -const ProviderSpecificFields: React.FC = ({ selectedProvider }) => { +const ProviderSpecificFields: React.FC = ({ + selectedProvider, + initialVariantId, + onCreateCredential, +}) => { const selectedProviderEnum = Providers[selectedProvider as keyof typeof Providers] as Providers; const form = useFormContext(); const credentialsFileRef = React.useRef(null); @@ -235,7 +264,7 @@ const ProviderSpecificFields: React.FC = ({ selecte // still-loading `variants` or a provider switch resolve correctly with no effect needed to // "catch up" afterward). A choice left over from a different provider's variant list is // treated as unset rather than resolving to a variant id that provider doesn't have. - const [userChosenVariantId, setUserChosenVariantId] = React.useState(undefined); + const [userChosenVariantId, setUserChosenVariantId] = React.useState(initialVariantId); const validUserChoice = variants?.variants.some((variant) => variant.id === userChosenVariantId) ? userChosenVariantId : undefined; @@ -400,6 +429,7 @@ const ProviderSpecificFields: React.FC = ({ selecte ); const activeVariant = variants ? getVariant(variants, activeVariantId) : undefined; + const credentialOnlyVariant = onCreateCredential && activeVariant?.credential_only ? activeVariant : undefined; return ( <> @@ -433,13 +463,17 @@ const ProviderSpecificFields: React.FC = ({ selecte {/* Keyed by the active variant so switching variants fully unmounts the previous variant's fields (deregistering them from submission) and re-applies fixed_values fresh, rather than leaving a stale value behind under a reused field name. */} - - {currentFields.map(renderFieldEntry)} - {activeVariant && - Object.entries(activeVariant.fixed_values).map(([key, value]) => ( - - ))} - + {credentialOnlyVariant && onCreateCredential ? ( + + ) : ( + + {currentFields.map(renderFieldEntry)} + {activeVariant && + Object.entries(activeVariant.fixed_values).map(([key, value]) => ( + + ))} + + )} ); }; diff --git a/ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx b/ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx index e29a9721abd..41570d6a7c9 100644 --- a/ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx +++ b/ui/litellm-dashboard/src/components/model_add/CredentialModal.test.tsx @@ -63,12 +63,13 @@ vi.mock("../networking", async () => { }, ], variants: [ - { id: "api_key", label: "API Key", field_keys: ["api_key"], fixed_values: {} }, + { id: "api_key", label: "API Key", field_keys: ["api_key"], fixed_values: {}, credential_only: false }, { id: "wif_token", label: "Workload Identity Federation", field_keys: ["anthropic_federation_rule_id", "anthropic_organization_id", "anthropic_identity_token"], fixed_values: {}, + credential_only: true, }, { id: "wif_internal_issuer", @@ -81,6 +82,7 @@ vi.mock("../networking", async () => { "anthropic_federation_rule_id", ], fixed_values: { anthropic_identity_source: "internal_issuer" }, + credential_only: true, }, ], }, @@ -345,4 +347,68 @@ describe("CredentialModal", () => { expect(deletedKeys).toEqual([]); }); }); + + const presetWifProps = { mode: "add" as const, initialProvider: Providers.Anthropic, initialVariantId: "wif_token" }; + + describe("after submit", () => { + it("keeps the typed values when saving fails", async () => { + const onSubmit = vi.fn().mockResolvedValue(false); + renderModal({ ...presetWifProps, onSubmit }); + const user = userEvent.setup(); + + await user.type(screen.getByLabelText("Credential Name:"), "anthropic-wif"); + await user.type(await screen.findByLabelText("Federation Rule ID"), "rule-1"); + await user.type(screen.getByLabelText("Organization ID"), "org-1"); + await user.type(screen.getByLabelText("Identity Token Reference"), "oidc/env/TOKEN"); + await user.click(screen.getByRole("button", { name: "Add Credential" })); + + await waitFor(() => expect(onSubmit).toHaveBeenCalled()); + expect(screen.getByLabelText("Credential Name:")).toHaveValue("anthropic-wif"); + expect(screen.getByLabelText("Federation Rule ID")).toHaveValue("rule-1"); + expect(screen.getByLabelText("Identity Token Reference")).toHaveValue("oidc/env/TOKEN"); + }); + + it("clears the form once saving succeeds", async () => { + const onSubmit = vi.fn().mockResolvedValue(true); + renderModal({ ...presetWifProps, onSubmit }); + const user = userEvent.setup(); + + await user.type(screen.getByLabelText("Credential Name:"), "anthropic-wif"); + await user.type(await screen.findByLabelText("Federation Rule ID"), "rule-1"); + await user.type(screen.getByLabelText("Organization ID"), "org-1"); + await user.type(screen.getByLabelText("Identity Token Reference"), "oidc/env/TOKEN"); + await user.click(screen.getByRole("button", { name: "Add Credential" })); + + await waitFor(() => expect(onSubmit).toHaveBeenCalled()); + await waitFor(() => expect(screen.getByLabelText("Credential Name:")).toHaveValue("")); + }); + }); + + describe("preset from the Add Model form", () => { + it("opens on the handed-over provider and variant and submits that provider", async () => { + const onSubmit = vi.fn(); + renderModal({ ...presetWifProps, onSubmit }); + const user = userEvent.setup(); + + expect(await screen.findByLabelText("Federation Rule ID")).toBeInTheDocument(); + expect(screen.queryByLabelText("Anthropic API Key")).not.toBeInTheDocument(); + expect(screen.queryByLabelText("OpenAI API Key")).not.toBeInTheDocument(); + + await user.type(screen.getByLabelText("Credential Name:"), "anthropic-wif"); + await user.type(screen.getByLabelText("Federation Rule ID"), "rule-1"); + await user.type(screen.getByLabelText("Organization ID"), "org-1"); + await user.type(screen.getByLabelText("Identity Token Reference"), "oidc/env/TOKEN"); + await user.click(screen.getByRole("button", { name: "Add Credential" })); + + await waitFor(() => expect(onSubmit).toHaveBeenCalled()); + const expectedCredential = { + credential_name: "anthropic-wif", + custom_llm_provider: "Anthropic", + anthropic_federation_rule_id: "rule-1", + anthropic_organization_id: "org-1", + anthropic_identity_token: "oidc/env/TOKEN", + }; + expect(onSubmit.mock.calls[0][0]).toEqual(expectedCredential); + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx b/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx index de36d77a817..12203c1d44a 100644 --- a/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx +++ b/ui/litellm-dashboard/src/components/model_add/CredentialModal.tsx @@ -32,6 +32,7 @@ import { Providers } from "../provider_info_helpers"; import { computeCredentialValuesToDelete, planCredentialTest, + providerEnumKey, resetCredentialFormOnProviderChange, summarizeDiscoveredModels, } from "./credential_form_helpers"; @@ -51,9 +52,11 @@ type ConnectionTest = interface CredentialModalProps { open: boolean; onCancel: () => void; - onSubmit: (values: any, credentialValuesToDelete: string[]) => void; + onSubmit: (values: any, credentialValuesToDelete: string[]) => Promise; mode: "add" | "edit"; existingCredential?: CredentialItem | null; + initialProvider?: Providers; + initialVariantId?: string; testConnection: (request: ProviderModelDiscoveryRequest) => Promise; loadJwks: (credentialName: string) => Promise; } @@ -93,15 +96,18 @@ export default function CredentialModal({ onSubmit, mode, existingCredential = null, + initialProvider, + initialVariantId, testConnection, loadJwks, }: CredentialModalProps) { const isEdit = mode === "edit"; const [selectedProvider, setSelectedProvider] = useState( - (existingCredential?.credential_info.custom_llm_provider as Providers) ?? Providers.OpenAI, + (existingCredential?.credential_info.custom_llm_provider as Providers) ?? initialProvider ?? Providers.OpenAI, ); const [connectionTest, setConnectionTest] = useState({ kind: "idle" }); + const presetValues = initialProvider ? { custom_llm_provider: providerEnumKey(initialProvider) } : undefined; const initialValues = existingCredential ? { credential_name: existingCredential.credential_name, @@ -110,7 +116,7 @@ export default function CredentialModal({ Object.entries(existingCredential.credential_values || {}).map(([key, value]) => [key, value ?? null]), ), } - : undefined; + : presetValues; const form = useForm({ mode: "onChange", defaultValues: initialValues }); const registry = useMountRegistry(); @@ -149,8 +155,10 @@ export default function CredentialModal({ const credentialValuesToDelete = isEdit ? computeCredentialValuesToDelete(existingCredential?.credential_values ?? {}, values) : []; - onSubmit(filteredValues, credentialValuesToDelete); - form.reset(); + const saved = await onSubmit(filteredValues, credentialValuesToDelete); + if (saved) { + form.reset(); + } }; const handleTestConnection = async () => { @@ -225,7 +233,7 @@ export default function CredentialModal({ )} - + {jwksCredentialName && } diff --git a/ui/litellm-dashboard/src/components/model_add/CredentialsPanel.tsx b/ui/litellm-dashboard/src/components/model_add/CredentialsPanel.tsx index 18a5b45f042..bcf98bab976 100644 --- a/ui/litellm-dashboard/src/components/model_add/CredentialsPanel.tsx +++ b/ui/litellm-dashboard/src/components/model_add/CredentialsPanel.tsx @@ -22,19 +22,7 @@ import DeleteResourceModal from "../common_components/DeleteResourceModal"; import { toast } from "@/lib/toast"; import CredentialModal from "./CredentialModal"; import CredentialsTable from "./CredentialsTable"; - -const restrictedFields = ["credential_name", "custom_llm_provider"]; - -const buildCredential = (values: Record, credentialValues: Record) => ({ - credential_name: values.credential_name as string, - credential_values: credentialValues, - credential_info: { - custom_llm_provider: values.custom_llm_provider as string, - }, -}); - -const withoutRestrictedFields = (values: Record): Record => - Object.fromEntries(Object.entries(values).filter(([key]) => !restrictedFields.includes(key))); +import { buildCredential, withoutRestrictedFields } from "./credential_form_helpers"; export default function CredentialsPanel() { const { accessToken, userRole } = useAuthorized(); @@ -54,9 +42,12 @@ export default function CredentialsPanel() { const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); const [isCredentialDeleting, setIsCredentialDeleting] = useState(false); - const handleUpdateCredential = async (values: Record, credentialValuesToDelete: string[] = []) => { + const handleUpdateCredential = async ( + values: Record, + credentialValuesToDelete: string[] = [], + ): Promise => { if (!accessToken) { - return; + return false; } try { const newCredential = { @@ -67,14 +58,16 @@ export default function CredentialsPanel() { toast.success("Credential updated successfully"); setIsUpdateModalOpen(false); await refetchCredentials(); + return true; } catch (error) { toast.error("Failed to update credential"); + return false; } }; - const handleAddCredential = async (values: Record) => { + const handleAddCredential = async (values: Record): Promise => { if (!accessToken) { - return; + return false; } try { const newCredential = buildCredential(values, withoutRestrictedFields(values)); @@ -82,8 +75,10 @@ export default function CredentialsPanel() { toast.success("Credential added successfully"); setIsAddModalOpen(false); await refetchCredentials(); + return true; } catch (error) { toast.error("Failed to add credential"); + return false; } }; diff --git a/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.ts b/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.ts index ffd5aaf22d5..c168fb1dcf1 100644 --- a/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.ts +++ b/ui/litellm-dashboard/src/components/model_add/credential_form_helpers.ts @@ -126,3 +126,24 @@ export function summarizeDiscoveredModels(models: readonly string[]): string { const rest = models.length > 3 ? ` and ${models.length - 3} more` : ""; return `Connection succeeded. ${models.length} model${models.length === 1 ? "" : "s"} available: ${shown}${rest}.`; } + +export type CredentialPayload = { + readonly credential_name: string; + readonly credential_values: Record; + readonly credential_info: { readonly custom_llm_provider: string }; +}; + +export const buildCredential = ( + values: Record, + credentialValues: Record, +): CredentialPayload => ({ + credential_name: values.credential_name as string, + credential_values: credentialValues, + credential_info: { custom_llm_provider: values.custom_llm_provider as string }, +}); + +export const withoutRestrictedFields = (values: Record): Record => + Object.fromEntries(Object.entries(values).filter(([key]) => !FORM_META_KEYS.has(key))); + +export const providerEnumKey = (provider: Providers): string => + Object.entries(Providers).find(([, displayName]) => displayName === provider)?.[0] ?? provider; diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 467b59ef42d..9e2a67abb68 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -318,6 +318,7 @@ export interface ProviderCredentialVariant { field_keys: string[]; optional_field_keys?: string[]; fixed_values: Record; + credential_only: boolean; } export interface ProviderCredentialVariants { diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index e03b94d97d7..5319ecbf714 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -34301,8 +34301,13 @@ export interface components { * ``optional_field_keys`` relaxes a globally-required field for this variant alone, for a value * only obtainable after the credential exists (the federation rule id an operator can only read * off the Anthropic Console once the generated JWKS is registered). + * ``credential_only`` is derived, never declared: a variant that submits a server-owned workload + * identity federation parameter, as a field or a fixed value, can only be saved as a named LLM + * Credential, since ``/model/new`` rejects those parameters inline. */ ProviderCredentialVariant: { + /** Credential Only */ + readonly credential_only: boolean; /** Field Keys */ field_keys: string[]; /** Fixed Values */