diff --git a/litellm/proxy/public_endpoints/provider_create_metadata.py b/litellm/proxy/public_endpoints/provider_create_metadata.py index f2ef03e2a0c..bfb2fb2fe0f 100644 --- a/litellm/proxy/public_endpoints/provider_create_metadata.py +++ b/litellm/proxy/public_endpoints/provider_create_metadata.py @@ -6,6 +6,7 @@ from litellm.types.proxy.public_endpoints.public_endpoints import ( ProviderCreateInfo, ProviderCredentialField, ) +from litellm.types.utils import LlmProviders DEFAULT_MODEL_PLACEHOLDER = "gpt-3.5-turbo" @@ -739,6 +740,30 @@ def get_provider_create_metadata() -> List[ProviderCreateInfo]: ) ) + # Ensure we have metadata entries for all providers defined in LlmProviders. + # If a provider enum value is not already present in the litellm_provider + # field of any entry, create a default entry for it using the fallback + # credential fields (api_key + api_base) and a generated display name. + existing_litellm_providers = {p.litellm_provider for p in providers} + + for provider_enum in LlmProviders: + litellm_provider_value = provider_enum.value + if litellm_provider_value in existing_litellm_providers: + continue + + normalized_fields = [_normalize_field(field) for field in _FALLBACK_FIELDS] + provider_display_name = provider_enum.value.replace("_", " ").title() + + providers.append( + ProviderCreateInfo( + provider=provider_enum.name, + provider_display_name=provider_display_name, + litellm_provider=litellm_provider_value, + default_model_placeholder=DEFAULT_MODEL_PLACEHOLDER, + credential_fields=normalized_fields, + ) + ) + providers.sort(key=lambda item: item.provider_display_name.lower()) return providers 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 f433c64c58d..8456cf55389 100644 --- a/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py +++ b/tests/test_litellm/proxy/public_endpoints/test_public_endpoints.py @@ -45,3 +45,22 @@ def test_get_provider_fields_returns_metadata(): credential_keys = {field["key"] for field in openai_fields["credential_fields"]} assert {"api_base", "api_key"}.issubset(credential_keys) + # Every provider exposed by `/public/providers` (i.e. every LlmProviders value) + # should have a corresponding entry in `/public/providers/fields`. + expected_litellm_providers = {provider.value for provider in LlmProviders} + actual_litellm_providers = {item["litellm_provider"] for item in payload} + assert expected_litellm_providers.issubset(actual_litellm_providers) + + # Sanity check for runwayml specifically – it should be present and use the + # default API base + API key credential fields at minimum. + runway_entries = [ + item for item in payload if item["litellm_provider"] == "runwayml" + ] + assert ( + len(runway_entries) >= 1 + ), "Expected runwayml provider metadata in /public/providers/fields" + runway_credential_keys = { + field["key"] for field in runway_entries[0]["credential_fields"] + } + assert {"api_base", "api_key"}.issubset(runway_credential_keys) + diff --git a/ui/litellm-dashboard/src/components/add_model/add_model_tab.test.tsx b/ui/litellm-dashboard/src/components/add_model/add_model_tab.test.tsx index a74c6d283ca..a3937c2f2ff 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_model_tab.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_model_tab.test.tsx @@ -19,6 +19,15 @@ vi.mock("../networking", async () => { modelAvailableCall: vi.fn().mockResolvedValue({ data: [{ id: "model-group-1" }, { id: "model-group-2" }], }), + getProviderCreateMetadata: vi.fn().mockResolvedValue([ + { + provider: "OpenAI", + provider_display_name: "OpenAI", + litellm_provider: "openai", + default_model_placeholder: "gpt-3.5-turbo", + credential_fields: [], + }, + ]), }; }); diff --git a/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx index 678ba741ca8..4efcd7be907 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_model_tab.tsx @@ -1,4 +1,4 @@ -import React, { useEffect, useState } from "react"; +import React, { useEffect, useMemo, useState } from "react"; import { Card, Form, Button, Tooltip, Typography, Select as AntdSelect, Modal } from "antd"; import type { FormInstance } from "antd"; import type { UploadProps } from "antd/es/upload"; @@ -9,7 +9,14 @@ import ProviderSpecificFields from "./provider_specific_fields"; import AdvancedSettings from "./advanced_settings"; import { Providers, providerLogoMap } from "../provider_info_helpers"; import type { Team } from "../key_team_helpers/key_list"; -import { CredentialItem, getGuardrailsList, modelAvailableCall, tagListCall } from "../networking"; +import { + type CredentialItem, + type ProviderCreateInfo, + getGuardrailsList, + getProviderCreateMetadata, + modelAvailableCall, + tagListCall, +} from "../networking"; import ConnectionErrorDisplay from "./model_connection_test"; import { TEST_MODES } from "./add_model_modes"; import { Row, Col } from "antd"; @@ -68,6 +75,11 @@ const AddModelTab: React.FC = ({ // Using a unique ID to force the ConnectionErrorDisplay to remount and run a fresh test const [connectionTestId, setConnectionTestId] = useState(""); + // Provider metadata for driving the provider select from backend config + const [providerMetadata, setProviderMetadata] = useState(null); + const [isProviderMetadataLoading, setIsProviderMetadataLoading] = useState(false); + const [providerMetadataError, setProviderMetadataError] = useState(null); + useEffect(() => { const fetchGuardrails = async () => { try { @@ -95,6 +107,37 @@ const AddModelTab: React.FC = ({ fetchTags(); }, [accessToken]); + useEffect(() => { + let isMounted = true; + + const fetchProviderMetadata = async () => { + setIsProviderMetadataLoading(true); + setProviderMetadataError(null); + try { + const metadata = await getProviderCreateMetadata(); + if (!isMounted) { + return; + } + setProviderMetadata(metadata); + } catch (error) { + console.error("Failed to fetch provider metadata:", error); + if (isMounted) { + setProviderMetadataError("Failed to load providers"); + } + } finally { + if (isMounted) { + setIsProviderMetadataLoading(false); + } + } + }; + + fetchProviderMetadata(); + + return () => { + isMounted = false; + }; + }, []); + // Test connection when button is clicked const handleTestConnection = async () => { setIsTestingConnection(true); @@ -118,6 +161,13 @@ const AddModelTab: React.FC = ({ fetchModelAccessGroups(); }, [accessToken]); + const sortedProviderMetadata: ProviderCreateInfo[] = useMemo(() => { + if (!providerMetadata) { + return []; + } + return [...providerMetadata].sort((a, b) => a.provider_display_name.localeCompare(b.provider_display_name)); + }, [providerMetadata]); + const isAdmin = all_admin_roles.includes(userRole); const handleAutoRouterOk = () => { @@ -166,41 +216,68 @@ const AddModelTab: React.FC = ({ labelAlign="left" > { - setSelectedProvider(value); - setProviderModelsFn(value); + setSelectedProvider(value as Providers); + setProviderModelsFn(value as Providers); + form.setFieldsValue({ + custom_llm_provider: value, + }); form.setFieldsValue({ model: [], model_name: undefined, }); }} > - {Object.entries(Providers).map(([providerEnum, providerDisplayName]) => ( - -
- {`${providerEnum} { - // Create a div with provider initial as fallback - const target = e.target as HTMLImageElement; - const parent = target.parentElement; - if (parent) { - const fallbackDiv = document.createElement("div"); - fallbackDiv.className = - "w-5 h-5 rounded-full bg-gray-200 flex items-center justify-center text-xs"; - fallbackDiv.textContent = providerDisplayName.charAt(0); - parent.replaceChild(fallbackDiv, target); - } - }} - /> - {providerDisplayName} -
+ {providerMetadataError && sortedProviderMetadata.length === 0 && ( + + {providerMetadataError} - ))} + )} + {sortedProviderMetadata.map((providerInfo) => { + const displayName = providerInfo.provider_display_name; + const providerKey = providerInfo.provider; + const logoSrc = providerLogoMap[displayName] ?? ""; + + return ( + +
+ {logoSrc ? ( + {`${displayName} { + const target = e.currentTarget as HTMLImageElement; + const parent = target.parentElement; + if (!parent || !parent.contains(target)) { + return; + } + + try { + const fallbackDiv = document.createElement("div"); + fallbackDiv.className = + "w-5 h-5 rounded-full bg-gray-200 flex items-center justify-center text-xs"; + fallbackDiv.textContent = displayName.charAt(0); + parent.replaceChild(fallbackDiv, target); + } catch (error) { + console.error("Failed to replace provider logo fallback:", error); + } + }} + /> + ) : ( +
+ {displayName.charAt(0)} +
+ )} + {displayName} +
+
+ ); + })}
({ + default: { + fromBackend: vi.fn(), + }, +})); + +describe("prepareModelAddRequest", () => { + it("returns deployment data for the most basic form", async () => { + const formValues = { + model_mappings: [ + { + public_name: "Public Model", + litellm_model: "litellm/public", + }, + ], + model_name: "custom-model-name", + base_model: "gpt-4", + team_id: "team-123", + model_access_group: ["group-1"], + input_cost_per_token: "2000000", + output_cost_per_token: "1000000", + }; + + const deployments = await prepareModelAddRequest({ ...formValues }, "token", null); + + expect(deployments).toHaveLength(1); + const [deployment] = deployments!; + expect(deployment.modelName).toBe("Public Model"); + expect(deployment.litellmParamsObj.model).toBe("custom-model-name"); + expect(deployment.litellmParamsObj.input_cost_per_token).toBe(2); + expect(deployment.litellmParamsObj.output_cost_per_token).toBe(1); + expect(deployment.modelInfoObj.base_model).toBe("gpt-4"); + expect(deployment.modelInfoObj.access_groups).toEqual(["group-1"]); + expect(deployment.modelInfoObj.team_id).toBe("team-123"); + }); + + it("uses a lowercase fallback for unrecognized custom providers", async () => { + const fallbackValues = { + model_mappings: [ + { + public_name: "Petals Model", + litellm_model: "petals/model", + }, + ], + model_name: "petals/model", + custom_llm_provider: "Petals", + }; + + const deployments = await prepareModelAddRequest({ ...fallbackValues }, "token", null); + + expect(deployments).toHaveLength(1); + const [deployment] = deployments!; + expect(deployment.litellmParamsObj.custom_llm_provider).toBe("petals"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/add_model/handle_add_model_submit.tsx b/ui/litellm-dashboard/src/components/add_model/handle_add_model_submit.tsx index e41b114c777..8fa5ffd56a2 100644 --- a/ui/litellm-dashboard/src/components/add_model/handle_add_model_submit.tsx +++ b/ui/litellm-dashboard/src/components/add_model/handle_add_model_submit.tsx @@ -1,6 +1,6 @@ -import { provider_map, Providers } from "../provider_info_helpers"; -import { modelCreateCall, Model } from "../networking"; import NotificationManager from "../molecules/notifications_manager"; +import { Model, modelCreateCall } from "../networking"; +import { provider_map } from "../provider_info_helpers"; export const prepareModelAddRequest = async (formValues: Record, accessToken: string, form: any) => { try { @@ -14,8 +14,10 @@ export const prepareModelAddRequest = async (formValues: Record, ac // Handle wildcard case if (formValues["model"] && formValues["model"].includes("all-wildcard")) { - const customProvider: Providers = formValues["custom_llm_provider"]; - const litellm_custom_provider = provider_map[customProvider as keyof typeof Providers]; + const customProviderKey = formValues["custom_llm_provider"] as string; + const mappedProvider = + provider_map[customProviderKey as keyof typeof provider_map] ?? customProviderKey.toLowerCase(); + const litellm_custom_provider = mappedProvider; const wildcardModel = litellm_custom_provider + "/*"; formValues["model_name"] = wildcardModel; modelMappings.push({ @@ -59,7 +61,8 @@ export const prepareModelAddRequest = async (formValues: Record, ac litellmParamsObj["model"] = value; } else if (key == "custom_llm_provider") { console.log("custom_llm_provider:", value); - const mappingResult = provider_map[value]; // Get the corresponding value from the mapping + const providerKey = value as string; + const mappingResult = provider_map[providerKey as keyof typeof provider_map] ?? providerKey.toLowerCase(); litellmParamsObj["custom_llm_provider"] = mappingResult; console.log("custom_llm_provider mappingResult:", mappingResult); } else if (key == "model") { 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 454caf9e601..ab9b3e92e98 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 @@ -1,9 +1,102 @@ import { render, waitFor } from "@testing-library/react"; -import { describe, it, expect, beforeAll } from "vitest"; +import { describe, it, expect, beforeAll, vi } from "vitest"; import { Form } from "antd"; import { Providers } from "../provider_info_helpers"; import ProviderSpecificFields from "./provider_specific_fields"; +vi.mock("../networking", async () => { + const actual = await vi.importActual("../networking"); + return { + ...actual, + getProviderCreateMetadata: vi.fn().mockResolvedValue([ + { + provider: "OpenAI", + provider_display_name: Providers.OpenAI, + litellm_provider: "openai", + default_model_placeholder: "gpt-3.5-turbo", + credential_fields: [ + { + key: "api_base", + label: "API Base", + field_type: "text", + placeholder: "https://api.openai.com/v1", + tooltip: + "Common endpoints: https://api.openai.com/v1, https://eu.api.openai.com, https://us.api.openai.com", + default_value: "https://api.openai.com/v1", + }, + { + key: "organization", + label: "OpenAI Organization ID", + placeholder: "[OPTIONAL] my-unique-org", + }, + { + key: "api_key", + label: "OpenAI API Key", + field_type: "password", + required: true, + }, + ], + }, + { + provider: "Hosted_Vllm", + provider_display_name: Providers.Hosted_Vllm, + litellm_provider: "hosted_vllm", + default_model_placeholder: "vllm/any-model", + credential_fields: [ + { + key: "api_base", + label: "API Base", + placeholder: "https://...", + }, + { + key: "api_key", + label: "vLLM API Key", + field_type: "password", + }, + ], + }, + { + provider: "Azure", + provider_display_name: Providers.Azure, + litellm_provider: "azure", + default_model_placeholder: "azure/my-deployment", + credential_fields: [ + { + key: "api_base", + label: "API Base", + placeholder: "https://...", + required: true, + }, + { + key: "api_version", + label: "API Version", + placeholder: "2023-07-01-preview", + tooltip: + "By default litellm will use the latest version. If you want to use a different version, you can specify it here", + }, + { + key: "base_model", + label: "Base Model", + placeholder: "azure/gpt-3.5-turbo", + }, + { + key: "api_key", + label: "Azure API Key", + field_type: "password", + placeholder: "Enter your Azure API Key", + }, + { + key: "azure_ad_token", + label: "Azure AD Token", + field_type: "password", + placeholder: "Enter your Azure AD Token", + }, + ], + }, + ]), + }; +}); + // Mock window.matchMedia for Ant Design components beforeAll(() => { Object.defineProperty(window, "matchMedia", { 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 c34014436bb..c1b17ff4419 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 @@ -4,7 +4,12 @@ import { TextInput, Text } from "@tremor/react"; import { Row, Col, Typography, Button as Button2, Upload, UploadProps } from "antd"; import { UploadOutlined } from "@ant-design/icons"; import { provider_map, Providers } from "../provider_info_helpers"; -import { CredentialItem } from "../networking"; +import { + CredentialItem, + ProviderCreateInfo, + ProviderCredentialFieldMetadata, + getProviderCreateMetadata, +} from "../networking"; const { Link } = Typography; interface ProviderSpecificFieldsProps { @@ -28,6 +33,33 @@ export interface CredentialValues { value: string; } +const mapFieldMetadataToUiField = (field: ProviderCredentialFieldMetadata): ProviderCredentialField => { + const type: ProviderCredentialField["type"] = + field.field_type === "password" + ? "password" + : field.field_type === "select" + ? "select" + : field.field_type === "upload" + ? "upload" + : "text"; + + return { + key: field.key, + label: field.label, + placeholder: field.placeholder ?? undefined, + tooltip: field.tooltip ?? undefined, + required: field.required ?? false, + type, + options: field.options ?? undefined, + defaultValue: field.default_value ?? undefined, + }; +}; + +// In-memory cache of provider credential fields keyed by provider display name. +// This lets us reuse the data across multiple mounts and also supports +// non-React helpers like createCredentialFromModel. +const providerFieldsByDisplayName: Record = {}; + export const createCredentialFromModel = (provider: string, modelData: any): CredentialItem => { console.log("provider", provider); console.log("modelData", modelData); @@ -35,8 +67,8 @@ export const createCredentialFromModel = (provider: string, modelData: any): Cre if (!enumKey) { throw new Error(`Provider ${provider} not found in provider_map`); } - const providerEnum = Providers[enumKey as keyof typeof Providers]; - const providerFields = PROVIDER_CREDENTIAL_FIELDS[providerEnum] || []; + const providerDisplayName = Providers[enumKey as keyof typeof Providers]; + const providerFields = providerFieldsByDisplayName[providerDisplayName] || []; const credentialValues: object = {}; console.log("providerFields", providerFields); @@ -63,554 +95,103 @@ export const createCredentialFromModel = (provider: string, modelData: any): Cre return credential; }; -const PROVIDER_CREDENTIAL_FIELDS: Record = { - [Providers.OpenAI]: [ - { - key: "api_base", - label: "API Base", - type: "select", - placeholder: "Select an endpoint", - tooltip: "Select from common OpenAI endpoints", - defaultValue: "https://api.openai.com/v1", - options: [ - "https://api.openai.com/v1", - "https://us.api.openai.com/v1", - "https://eu.api.openai.com/v1", - ], - }, - { - key: "organization", - label: "OpenAI Organization ID", - placeholder: "[OPTIONAL] my-unique-org", - }, - { - key: "api_key", - label: "OpenAI API Key", - type: "password", - required: true, - }, - ], - [Providers.OpenAI_Text]: [ - { - key: "api_base", - label: "API Base", - type: "select", - placeholder: "Select an endpoint", - tooltip: "Select from common OpenAI endpoints", - defaultValue: "https://api.openai.com/v1", - options: [ - "https://api.openai.com/v1", - "https://us.api.openai.com/v1", - "https://eu.api.openai.com/v1", - ], - }, - { - key: "organization", - label: "OpenAI Organization ID", - placeholder: "[OPTIONAL] my-unique-org", - }, - { - key: "api_key", - label: "OpenAI API Key", - type: "password", - required: true, - }, - ], - [Providers.Vertex_AI]: [ - { - key: "vertex_project", - label: "Vertex Project", - placeholder: "adroit-cadet-1234..", - required: true, - }, - { - key: "vertex_location", - label: "Vertex Location", - placeholder: "us-east-1", - required: true, - }, - { - key: "vertex_credentials", - label: "Vertex Credentials", - required: true, - type: "upload", - }, - ], - [Providers.AssemblyAI]: [ - { - key: "api_base", - label: "API Base", - type: "select", - required: true, - options: ["https://api.assemblyai.com", "https://api.eu.assemblyai.com"], - }, - { - key: "api_key", - label: "AssemblyAI API Key", - type: "password", - required: true, - }, - ], - [Providers.Azure]: [ - { - key: "api_base", - label: "API Base", - placeholder: "https://...", - required: true, - }, - { - key: "api_version", - label: "API Version", - placeholder: "2023-07-01-preview", - tooltip: - "By default litellm will use the latest version. If you want to use a different version, you can specify it here", - }, - { - key: "base_model", - label: "Base Model", - placeholder: "azure/gpt-3.5-turbo", - }, - { - key: "api_key", - label: "Azure API Key", - type: "password", - placeholder: "Enter your Azure API Key", - required: false, - }, - { - key: "azure_ad_token", - label: "Azure AD Token", - type: "password", - placeholder: "Enter your Azure AD Token", - required: false, - }, - ], - [Providers.Azure_AI_Studio]: [ - { - key: "api_base", - label: "API Base", - placeholder: "https://.openai.azure.com/openai/deployments/gpt-4o/chat/completions?api-version=2024-10-21", - tooltip: - "Enter your full Target URI from Azure Foundry here. Example: https://litellm8397336933.openai.azure.com/openai/deployments/gpt-4o/chat/completions?api-version=2024-10-21", - required: true, - }, - { - key: "api_key", - label: "Azure API Key", - type: "password", - required: true, - }, - ], - [Providers.OpenAI_Compatible]: [ - { - key: "api_base", - label: "API Base", - placeholder: "https://...", - required: true, - }, - { - key: "api_key", - label: "OpenAI API Key", - type: "password", - required: true, - }, - ], - [Providers.Dashscope]: [ - { - key: "api_key", - label: "Dashscope API Key", - type: "password", - required: true, - }, - { - key: "api_base", - label: "API Base", - placeholder: "https://dashscope-intl.aliyuncs.com/compatible-mode/v1", - defaultValue: "https://dashscope-intl.aliyuncs.com/compatible-mode/v1", - required: true, - tooltip: - "The base URL for your Dashscope server. Defaults to https://dashscope-intl.aliyuncs.com/compatible-mode/v1 if not specified.", - }, - ], - [Providers.OpenAI_Text_Compatible]: [ - { - key: "api_base", - label: "API Base", - placeholder: "https://...", - required: true, - }, - { - key: "api_key", - label: "OpenAI API Key", - type: "password", - required: true, - }, - ], - [Providers.Bedrock]: [ - { - key: "aws_access_key_id", - label: "AWS Access Key ID", - type: "password", - required: false, - tooltip: "You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`).", - }, - { - key: "aws_secret_access_key", - label: "AWS Secret Access Key", - type: "password", - required: false, - tooltip: "You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`).", - }, - { - key: "aws_session_token", - label: "AWS Session Token", - type: "password", - required: false, - tooltip: - "Temporary credentials session token. You can provide the raw token or the environment variable (e.g. `os.environ/MY_SESSION_TOKEN`).", - }, - { - key: "aws_region_name", - label: "AWS Region Name", - placeholder: "us-east-1", - required: false, - tooltip: "You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`).", - }, - { - key: "aws_session_name", - label: "AWS Session Name", - placeholder: "my-session", - required: false, - tooltip: - "Name for the AWS session. You can provide the raw value or the environment variable (e.g. `os.environ/MY_SESSION_NAME`).", - }, - { - key: "aws_profile_name", - label: "AWS Profile Name", - placeholder: "default", - required: false, - tooltip: - "AWS profile name to use for authentication. You can provide the raw value or the environment variable (e.g. `os.environ/MY_PROFILE_NAME`).", - }, - { - key: "aws_role_name", - label: "AWS Role Name", - placeholder: "MyRole", - required: false, - tooltip: - "AWS IAM role name to assume. You can provide the raw value or the environment variable (e.g. `os.environ/MY_ROLE_NAME`).", - }, - { - key: "aws_web_identity_token", - label: "AWS Web Identity Token", - type: "password", - required: false, - tooltip: - "Web identity token for OIDC authentication. You can provide the raw token or the environment variable (e.g. `os.environ/MY_WEB_IDENTITY_TOKEN`).", - }, - { - key: "aws_bedrock_runtime_endpoint", - label: "AWS Bedrock Runtime Endpoint", - placeholder: "https://bedrock-runtime.us-east-1.amazonaws.com", - required: false, - tooltip: - "Custom Bedrock runtime endpoint URL. You can provide the raw value or the environment variable (e.g. `os.environ/MY_BEDROCK_ENDPOINT`).", - }, - ], - [Providers.SageMaker]: [ - { - key: "aws_access_key_id", - label: "AWS Access Key ID", - type: "password", - required: false, - tooltip: "You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`).", - }, - { - key: "aws_secret_access_key", - label: "AWS Secret Access Key", - type: "password", - required: false, - tooltip: "You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`).", - }, - { - key: "aws_region_name", - label: "AWS Region Name", - placeholder: "us-east-1", - required: false, - tooltip: "You can provide the raw key or the environment variable (e.g. `os.environ/MY_SECRET_KEY`).", - }, - ], - [Providers.Ollama]: [ - { - key: "api_base", - label: "API Base", - placeholder: "http://localhost:11434", - defaultValue: "http://localhost:11434", - required: false, - tooltip: "The base URL for your Ollama server. Defaults to http://localhost:11434 if not specified.", - }, - ], - [Providers.Anthropic]: [ - { - key: "api_key", - label: "API Key", - placeholder: "sk-", - type: "password", - required: true, - }, - ], - [Providers.Deepgram]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.ElevenLabs]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.Google_AI_Studio]: [ - { - key: "api_key", - label: "API Key", - placeholder: "aig-", - type: "password", - required: true, - }, - ], - [Providers.Groq]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.MistralAI]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.Deepseek]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.Cohere]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.Databricks]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.xAI]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.AIML]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.Cerebras]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.Sambanova]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.Perplexity]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.TogetherAI]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.Openrouter]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.FireworksAI]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.GradientAI]: [ - { - key: "api_base", - label: "GradientAI Endpoint", - placeholder: "https://...", - required: false, - }, - { - key: "api_key", - label: "GradientAI API Key", - type: "password", - required: true, - }, - ], - [Providers.Triton]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: false, - }, - { - key: "api_base", - label: "API Base", - placeholder: "http://localhost:8000/generate", - required: false, - }, - ], - [Providers.Hosted_Vllm]: [ - { - key: "api_base", - label: "API Base", - placeholder: "https://...", - required: true, - }, - { - key: "api_key", - label: "vLLM API Key", - type: "password", - required: false, - }, - ], - [Providers.Voyage]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.JinaAI]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.VolcEngine]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.DeepInfra]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.Oracle]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], - [Providers.Snowflake]: [ - { - key: "api_key", - label: "Snowflake API Key / JWT Key for Authentication", - type: "password", - required: true, - }, - { - key: "api_base", - label: "Snowflake API Endpoint", - placeholder: "https://1234567890.snowflakecomputing.com/api/v2/cortex/inference:complete", - tooltip: - "Enter the full endpoint with path here. Example: https://1234567890.snowflakecomputing.com/api/v2/cortex/inference:complete", - required: true, - }, - ], - [Providers.Infinity]: [ - { - key: "api_base", - label: "API Base", - placeholder: "http://localhost:7997", - }, - ], - [Providers.FalAI]: [ - { - key: "api_key", - label: "API Key", - type: "password", - required: true, - }, - ], -}; - const ProviderSpecificFields: React.FC = ({ selectedProvider, uploadProps }) => { const selectedProviderEnum = Providers[selectedProvider as keyof typeof Providers] as Providers; const form = Form.useFormInstance(); // Get form instance from context - // Simply use the fields as defined in PROVIDER_CREDENTIAL_FIELDS + const [providerMetadata, setProviderMetadata] = React.useState(null); + const [isLoading, setIsLoading] = React.useState(false); + const [loadError, setLoadError] = React.useState(null); + + React.useEffect(() => { + const hasCachedFields = Object.keys(providerFieldsByDisplayName).length > 0; + if (hasCachedFields) { + // We already have fields cached globally; no need to refetch. + // This is important so we can reuse credential field definitions + // across mounts and in non-React helpers. + return; + } + + let isMounted = true; + + const fetchProviderFields = async () => { + setIsLoading(true); + setLoadError(null); + try { + const metadata = await getProviderCreateMetadata(); + if (!isMounted) { + return; + } + setProviderMetadata(metadata); + + // Populate cache keyed by provider display name and identifiers + metadata.forEach((providerInfo) => { + const displayName = providerInfo.provider_display_name; + const mappedFields = providerInfo.credential_fields.map(mapFieldMetadataToUiField); + + // Primary key: human-readable display name + providerFieldsByDisplayName[displayName] = mappedFields; + + // Also cache by backend identifiers so lookups by provider slug work + if (providerInfo.provider) { + providerFieldsByDisplayName[providerInfo.provider] = mappedFields; + } + if (providerInfo.litellm_provider) { + providerFieldsByDisplayName[providerInfo.litellm_provider] = mappedFields; + } + }); + } catch (error) { + console.error("Failed to load provider credential fields:", error); + if (isMounted) { + setLoadError("Failed to load provider credential fields"); + } + } finally { + if (isMounted) { + setIsLoading(false); + } + } + }; + + fetchProviderFields(); + + return () => { + isMounted = false; + }; + }, []); + const allFields = React.useMemo(() => { - return PROVIDER_CREDENTIAL_FIELDS[selectedProviderEnum] || []; - }, [selectedProviderEnum]); + // First try to resolve from the in-memory cache. We support both the + // enum/display-name form and the raw provider slug (e.g. "petals"). + const cachedFields = + providerFieldsByDisplayName[selectedProviderEnum] ?? providerFieldsByDisplayName[selectedProvider]; + if (cachedFields) { + return cachedFields; + } + + if (!providerMetadata) { + return []; + } + + const providerInfo = providerMetadata.find( + (p) => + p.provider_display_name === selectedProviderEnum || + p.provider === selectedProvider || + p.litellm_provider === selectedProvider, + ); + if (!providerInfo) { + return []; + } + + const mapped = providerInfo.credential_fields.map(mapFieldMetadataToUiField); + providerFieldsByDisplayName[providerInfo.provider_display_name] = mapped; + if (providerInfo.provider) { + providerFieldsByDisplayName[providerInfo.provider] = mapped; + } + if (providerInfo.litellm_provider) { + providerFieldsByDisplayName[providerInfo.litellm_provider] = mapped; + } + return mapped; + }, [selectedProviderEnum, selectedProvider, providerMetadata]); const handleUpload = { name: "file", @@ -643,6 +224,20 @@ const ProviderSpecificFields: React.FC = ({ selecte return ( <> + {isLoading && allFields.length === 0 && ( + + + Loading provider fields... + + + )} + {loadError && allFields.length === 0 && ( + + + {loadError} + + + )} {allFields.map((field) => ( { - // Create a div with provider initial as fallback - const target = e.target as HTMLImageElement; + const target = e.currentTarget as HTMLImageElement; const parent = target.parentElement; - if (parent) { + if (!parent || !parent.contains(target)) { + return; + } + + try { const fallbackDiv = document.createElement("div"); fallbackDiv.className = "w-4 h-4 rounded-full bg-gray-200 flex items-center justify-center text-xs"; fallbackDiv.textContent = modelData.provider?.charAt(0) || "-"; parent.replaceChild(fallbackDiv, target); + } catch (error) { + console.error("Failed to replace provider logo fallback:", error); } }} /> diff --git a/ui/litellm-dashboard/src/components/molecules/models/columns.tsx b/ui/litellm-dashboard/src/components/molecules/models/columns.tsx index d3e01a6632a..782eb3394a0 100644 --- a/ui/litellm-dashboard/src/components/molecules/models/columns.tsx +++ b/ui/litellm-dashboard/src/components/molecules/models/columns.tsx @@ -67,14 +67,20 @@ export const columns = ( alt={`${model.provider} logo`} className="w-4 h-4" onError={(e) => { - const target = e.target as HTMLImageElement; + const target = e.currentTarget as HTMLImageElement; const parent = target.parentElement; - if (parent) { + if (!parent || !parent.contains(target)) { + return; + } + + try { const fallbackDiv = document.createElement("div"); fallbackDiv.className = "w-4 h-4 rounded-full bg-gray-200 flex items-center justify-center text-xs"; fallbackDiv.textContent = model.provider?.charAt(0) || "-"; parent.replaceChild(fallbackDiv, target); + } catch (error) { + console.error("Failed to replace provider logo fallback:", error); } }} /> diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index a8cb05310c8..28670344837 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -145,6 +145,25 @@ export interface CredentialItem { }; } +export interface ProviderCredentialFieldMetadata { + key: string; + label: string; + placeholder?: string | null; + tooltip?: string | null; + required?: boolean; + field_type?: "text" | "password" | "select" | "upload"; + options?: string[] | null; + default_value?: string | null; +} + +export interface ProviderCreateInfo { + provider: string; + provider_display_name: string; + litellm_provider: string; + default_model_placeholder?: string | null; + credential_fields: ProviderCredentialFieldMetadata[]; +} + export interface PublicModelHubInfo { docs_title: string; custom_docs_description: string | null; @@ -182,6 +201,26 @@ const handleError = async (errorData: string) => { } }; +export const getProviderCreateMetadata = async (): Promise => { + /** + * Fetch provider credential field metadata from the proxy's public endpoint. + * This is used by the UI to dynamically render provider-specific credential fields. + */ + const url = defaultProxyBaseUrl ? `${defaultProxyBaseUrl}/public/providers/fields` : `/public/providers/fields`; + const response = await fetch(url, { + method: "GET", + }); + + if (!response.ok) { + const errorText = await response.text(); + console.error("Failed to fetch provider create metadata:", response.status, errorText); + throw new Error("Failed to load provider configuration"); + } + + const jsonData: ProviderCreateInfo[] = await response.json(); + return jsonData; +}; + // Global variable for the header name let globalLitellmHeaderName: string = "Authorization"; const MCP_AUTH_HEADER: string = "x-mcp-auth"; @@ -5794,11 +5833,7 @@ export const listMCPTools = async (accessToken: string, serverId: string) => { } }; -export const callMCPTool = async ( - accessToken: string, - toolName: string, - toolArguments: Record -) => { +export const callMCPTool = async (accessToken: string, toolName: string, toolArguments: Record) => { try { // Construct base URL let url = proxyBaseUrl ? `${proxyBaseUrl}/mcp-rest/tools/call` : `/mcp-rest/tools/call`;