Add Model uses endpoint info (#16664)

This commit is contained in:
yuneng-jiang 2025-11-14 16:08:27 -08:00 • committed by GitHub
parent 2bd6d0d82b
commit f1ff195bd8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 516 additions and 591 deletions

View file

@ -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

View file

@ -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)

View file

@ -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: [],
},
]),
};
});

View file

@ -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<AddModelTabProps> = ({
// Using a unique ID to force the ConnectionErrorDisplay to remount and run a fresh test
const [connectionTestId, setConnectionTestId] = useState<string>("");
// Provider metadata for driving the provider select from backend config
const [providerMetadata, setProviderMetadata] = useState<ProviderCreateInfo[] | null>(null);
const [isProviderMetadataLoading, setIsProviderMetadataLoading] = useState<boolean>(false);
const [providerMetadataError, setProviderMetadataError] = useState<string | null>(null);
useEffect(() => {
const fetchGuardrails = async () => {
try {
@ -95,6 +107,37 @@ const AddModelTab: React.FC<AddModelTabProps> = ({
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<AddModelTabProps> = ({
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<AddModelTabProps> = ({
labelAlign="left"
>
<AntdSelect
showSearch={true}
value={selectedProvider}
showSearch
loading={isProviderMetadataLoading}
placeholder={isProviderMetadataLoading ? "Loading providers..." : "Select a provider"}
optionFilterProp="data-label"
onChange={(value) => {
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]) => (
<AntdSelect.Option key={providerEnum} value={providerEnum}>
<div className="flex items-center space-x-2">
<img
src={providerLogoMap[providerDisplayName]}
alt={`${providerEnum} logo`}
className="w-5 h-5"
onError={(e) => {
// 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);
}
}}
/>
<span>{providerDisplayName}</span>
</div>
{providerMetadataError && sortedProviderMetadata.length === 0 && (
<AntdSelect.Option key="__error" value="">
{providerMetadataError}
</AntdSelect.Option>
))}
)}
{sortedProviderMetadata.map((providerInfo) => {
const displayName = providerInfo.provider_display_name;
const providerKey = providerInfo.provider;
const logoSrc = providerLogoMap[displayName] ?? "";
return (
<AntdSelect.Option key={providerKey} value={providerKey} data-label={displayName}>
<div className="flex items-center space-x-2">
{logoSrc ? (
<img
src={logoSrc}
alt={`${displayName} logo`}
className="w-5 h-5"
onError={(e) => {
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);
}
}}
/>
) : (
<div className="w-5 h-5 rounded-full bg-gray-200 flex items-center justify-center text-xs">
{displayName.charAt(0)}
</div>
)}
<span>{displayName}</span>
</div>
</AntdSelect.Option>
);
})}
</AntdSelect>
</Form.Item>
<LiteLLMModelNameField

View file

@ -0,0 +1,58 @@
import { describe, expect, it, vi } from "vitest";
import { prepareModelAddRequest } from "./handle_add_model_submit";
vi.mock("../molecules/notifications_manager", () => ({
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");
});
});

View file

@ -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<string, any>, accessToken: string, form: any) => {
try {
@ -14,8 +14,10 @@ export const prepareModelAddRequest = async (formValues: Record<string, any>, 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<string, any>, 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") {

View file

@ -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", {

View file

@ -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<string, ProviderCredentialField[]> = {};
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, ProviderCredentialField[]> = {
[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://<test>.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<ProviderSpecificFieldsProps> = ({ 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<ProviderCreateInfo[] | null>(null);
const [isLoading, setIsLoading] = React.useState<boolean>(false);
const [loadError, setLoadError] = React.useState<string | null>(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<ProviderSpecificFieldsProps> = ({ selecte
return (
<>
{isLoading && allFields.length === 0 && (
<Row>
<Col span={24}>
<Text className="mb-2">Loading provider fields...</Text>
</Col>
</Row>
)}
{loadError && allFields.length === 0 && (
<Row>
<Col span={24}>
<Text className="mb-2 text-red-500">{loadError}</Text>
</Col>
</Row>
)}
{allFields.map((field) => (
<React.Fragment key={field.key}>
<Form.Item

View file

@ -419,15 +419,20 @@ export default function ModelInfoView({
alt={`${modelData.provider} logo`}
className="w-4 h-4"
onError={(e) => {
// 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);
}
}}
/>

View file

@ -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);
}
}}
/>

View file

@ -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<ProviderCreateInfo[]> => {
/**
* 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<string, any>
) => {
export const callMCPTool = async (accessToken: string, toolName: string, toolArguments: Record<string, any>) => {
try {
// Construct base URL
let url = proxyBaseUrl ? `${proxyBaseUrl}/mcp-rest/tools/call` : `/mcp-rest/tools/call`;