mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(ui): route credential-only WIF variants on Add Model through an LLM Credential
Federation auth variants for OpenAI and Anthropic submit server-owned identity parameters that /model/new and /health/test_connection reject inline with a 401, so the Add Model form was a dead end for them. The provider fields payload now flags those variants credential_only, and Add Model swaps their inline fields for a "Create credential" step that opens the credential dialog preset to that provider and variant, saves the credential, and attaches it to the model. The credential dialog also keeps the typed values when saving fails instead of wiping them.
This commit is contained in:
parent
82afbbdba1
commit
64a2eb98ab
14 changed files with 438 additions and 46 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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(<AddModelForm {...props} />);
|
||||
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();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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<AddModelFormProps> = ({
|
|||
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<string | null>(null);
|
||||
|
||||
const handleCreateCredential = async (values: Record<string, unknown>): Promise<boolean> => {
|
||||
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<AddModelFormProps> = ({
|
|||
<span className="px-4 text-muted-foreground text-sm">OR</span>
|
||||
<div className="grow border-t border-border"></div>
|
||||
</div>
|
||||
<ProviderSpecificFields selectedProvider={selectedProvider} />
|
||||
<ProviderSpecificFields
|
||||
selectedProvider={selectedProvider}
|
||||
onCreateCredential={isAdmin ? setCredentialDraftVariantId : undefined}
|
||||
/>
|
||||
</>
|
||||
)}
|
||||
<div className="flex items-center my-4">
|
||||
|
|
@ -448,6 +481,18 @@ const AddModelForm: React.FC<AddModelFormProps> = ({
|
|||
</CardContent>
|
||||
</Card>
|
||||
|
||||
{credentialDraftVariantId !== null && (
|
||||
<CredentialModal
|
||||
open
|
||||
mode="add"
|
||||
initialProvider={selectedProvider}
|
||||
initialVariantId={credentialDraftVariantId}
|
||||
onCancel={() => setCredentialDraftVariantId(null)}
|
||||
onSubmit={handleCreateCredential}
|
||||
testConnection={(request) => discoverProviderModelsCall(accessToken, request)}
|
||||
loadJwks={(credentialName) => getCredentialJwksCall(accessToken, credentialName)}
|
||||
/>
|
||||
)}
|
||||
{/* Test Connection Results Modal */}
|
||||
<Dialog
|
||||
open={isResultModalVisible}
|
||||
|
|
|
|||
|
|
@ -26,18 +26,20 @@ const anthropicVariants: ProviderCredentialVariants = {
|
|||
field("anthropic_keycloak_client_secret_ref", true),
|
||||
],
|
||||
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: "WIF (token)",
|
||||
field_keys: ["anthropic_federation_rule_id", "anthropic_organization_id", "anthropic_identity_token"],
|
||||
fixed_values: {},
|
||||
credential_only: true,
|
||||
},
|
||||
{
|
||||
id: "wif_token_file",
|
||||
label: "WIF (token file)",
|
||||
field_keys: ["anthropic_federation_rule_id", "anthropic_organization_id", "anthropic_identity_token_file"],
|
||||
fixed_values: {},
|
||||
credential_only: true,
|
||||
},
|
||||
{
|
||||
id: "wif_internal_issuer",
|
||||
|
|
@ -50,6 +52,7 @@ const anthropicVariants: ProviderCredentialVariants = {
|
|||
],
|
||||
optional_field_keys: ["anthropic_federation_rule_id"],
|
||||
fixed_values: { anthropic_identity_source: "internal_issuer" },
|
||||
credential_only: true,
|
||||
},
|
||||
{
|
||||
id: "wif_keycloak",
|
||||
|
|
@ -62,6 +65,7 @@ const anthropicVariants: ProviderCredentialVariants = {
|
|||
"anthropic_keycloak_client_secret_ref",
|
||||
],
|
||||
fixed_values: { anthropic_identity_source: "keycloak" },
|
||||
credential_only: true,
|
||||
},
|
||||
],
|
||||
};
|
||||
|
|
@ -93,7 +97,15 @@ describe("resolveVariantFieldDefs", () => {
|
|||
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"]);
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
<QueryClientProvider client={createQueryClient()}>
|
||||
<MountedFormHost>
|
||||
<ProviderSpecificFields selectedProvider={Providers.Anthropic} initialVariantId="wif_token" />
|
||||
</MountedFormHost>
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
|
||||
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(
|
||||
<QueryClientProvider client={createQueryClient()}>
|
||||
<MountedFormHost>
|
||||
<ProviderSpecificFields selectedProvider={Providers.Anthropic} onCreateCredential={onCreateCredential} />
|
||||
<IdentitySourceProbe />
|
||||
</MountedFormHost>
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
|
||||
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(
|
||||
<QueryClientProvider client={createQueryClient()}>
|
||||
<MountedFormHost>
|
||||
<ProviderSpecificFields selectedProvider={Providers.Anthropic} onCreateCredential={vi.fn()} />
|
||||
</MountedFormHost>
|
||||
</QueryClientProvider>,
|
||||
);
|
||||
|
||||
expect(await screen.findByLabelText("API Key")).toBeInTheDocument();
|
||||
expect(screen.queryByRole("button", { name: "Create credential" })).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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 }) => (
|
||||
<Alert className="mb-4">
|
||||
<AlertDescription className="flex flex-col items-start gap-3">
|
||||
<span>
|
||||
{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.
|
||||
</span>
|
||||
<Button type="button" variant="outline" onClick={() => onCreateCredential(variant.id)}>
|
||||
Create credential
|
||||
</Button>
|
||||
</AlertDescription>
|
||||
</Alert>
|
||||
);
|
||||
|
||||
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<ProviderSpecificFieldsProps> = ({ selectedProvider }) => {
|
||||
const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({
|
||||
selectedProvider,
|
||||
initialVariantId,
|
||||
onCreateCredential,
|
||||
}) => {
|
||||
const selectedProviderEnum = Providers[selectedProvider as keyof typeof Providers] as Providers;
|
||||
const form = useFormContext<MountedFormValues>();
|
||||
const credentialsFileRef = React.useRef<HTMLInputElement>(null);
|
||||
|
|
@ -235,7 +264,7 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({ 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<string | undefined>(undefined);
|
||||
const [userChosenVariantId, setUserChosenVariantId] = React.useState<string | undefined>(initialVariantId);
|
||||
const validUserChoice = variants?.variants.some((variant) => variant.id === userChosenVariantId)
|
||||
? userChosenVariantId
|
||||
: undefined;
|
||||
|
|
@ -400,6 +429,7 @@ const ProviderSpecificFields: React.FC<ProviderSpecificFieldsProps> = ({ 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<ProviderSpecificFieldsProps> = ({ 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. */}
|
||||
<React.Fragment key={variants ? activeVariantId : "static"}>
|
||||
{currentFields.map(renderFieldEntry)}
|
||||
{activeVariant &&
|
||||
Object.entries(activeVariant.fixed_values).map(([key, value]) => (
|
||||
<FixedValueField key={key} name={key} value={value} />
|
||||
))}
|
||||
</React.Fragment>
|
||||
{credentialOnlyVariant && onCreateCredential ? (
|
||||
<CredentialOnlyNotice variant={credentialOnlyVariant} onCreateCredential={onCreateCredential} />
|
||||
) : (
|
||||
<React.Fragment key={variants ? activeVariantId : "static"}>
|
||||
{currentFields.map(renderFieldEntry)}
|
||||
{activeVariant &&
|
||||
Object.entries(activeVariant.fixed_values).map(([key, value]) => (
|
||||
<FixedValueField key={key} name={key} value={value} />
|
||||
))}
|
||||
</React.Fragment>
|
||||
)}
|
||||
</>
|
||||
);
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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<boolean>;
|
||||
mode: "add" | "edit";
|
||||
existingCredential?: CredentialItem | null;
|
||||
initialProvider?: Providers;
|
||||
initialVariantId?: string;
|
||||
testConnection: (request: ProviderModelDiscoveryRequest) => Promise<ProviderModelDiscoveryResponse>;
|
||||
loadJwks: (credentialName: string) => Promise<AnthropicJwks>;
|
||||
}
|
||||
|
|
@ -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<Providers>(
|
||||
(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<ConnectionTest>({ 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<MountedFormValues>({ 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({
|
|||
)}
|
||||
</MountedFormField>
|
||||
|
||||
<ProviderSpecificFields selectedProvider={selectedProvider} />
|
||||
<ProviderSpecificFields selectedProvider={selectedProvider} initialVariantId={initialVariantId} />
|
||||
|
||||
{jwksCredentialName && <JwksExport credentialName={jwksCredentialName} loadJwks={loadJwks} />}
|
||||
|
||||
|
|
|
|||
|
|
@ -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<string, unknown>, credentialValues: Record<string, unknown>) => ({
|
||||
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<string, unknown>): Record<string, unknown> =>
|
||||
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<string, unknown>, credentialValuesToDelete: string[] = []) => {
|
||||
const handleUpdateCredential = async (
|
||||
values: Record<string, unknown>,
|
||||
credentialValuesToDelete: string[] = [],
|
||||
): Promise<boolean> => {
|
||||
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<string, unknown>) => {
|
||||
const handleAddCredential = async (values: Record<string, unknown>): Promise<boolean> => {
|
||||
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;
|
||||
}
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -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<string, unknown>;
|
||||
readonly credential_info: { readonly custom_llm_provider: string };
|
||||
};
|
||||
|
||||
export const buildCredential = (
|
||||
values: Record<string, unknown>,
|
||||
credentialValues: Record<string, unknown>,
|
||||
): 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<string, unknown>): Record<string, unknown> =>
|
||||
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;
|
||||
|
|
|
|||
|
|
@ -318,6 +318,7 @@ export interface ProviderCredentialVariant {
|
|||
field_keys: string[];
|
||||
optional_field_keys?: string[];
|
||||
fixed_values: Record<string, string>;
|
||||
credential_only: boolean;
|
||||
}
|
||||
|
||||
export interface ProviderCredentialVariants {
|
||||
|
|
|
|||
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
5
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -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 */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue