mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
fix(ui): add nvidia riva to the model provider list
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
09889e1986
commit
6346497498
4 changed files with 113 additions and 42 deletions
|
|
@ -2062,6 +2062,44 @@
|
|||
],
|
||||
"default_model_placeholder": "gpt-3.5-turbo"
|
||||
},
|
||||
{
|
||||
"provider": "NVIDIA_RIVA",
|
||||
"provider_display_name": "Nvidia Riva",
|
||||
"litellm_provider": "nvidia_riva",
|
||||
"credential_fields": [
|
||||
{
|
||||
"key": "api_base",
|
||||
"label": "API Base",
|
||||
"placeholder": "grpc.nvcf.nvidia.com:443",
|
||||
"tooltip": "host:port of the Riva gRPC endpoint. Use grpc.nvcf.nvidia.com:443 for NVCF-hosted Riva, or your own host (e.g. localhost:50051) when self-hosting. Riva has no public default, so this is required.",
|
||||
"required": true,
|
||||
"field_type": "text",
|
||||
"options": null,
|
||||
"default_value": null
|
||||
},
|
||||
{
|
||||
"key": "api_key",
|
||||
"label": "API Key",
|
||||
"placeholder": "nvapi-...",
|
||||
"tooltip": "Sent as gRPC authorization metadata. Required for NVCF-hosted Riva, optional for self-hosted deployments without auth.",
|
||||
"required": false,
|
||||
"field_type": "password",
|
||||
"options": null,
|
||||
"default_value": null
|
||||
},
|
||||
{
|
||||
"key": "nvcf_function_id",
|
||||
"label": "NVCF Function ID",
|
||||
"placeholder": "1598d209-5e27-4d3c-8079-4751568b1081",
|
||||
"tooltip": "NVCF function id of the hosted Riva model. Setting it turns on TLS and the function-id gRPC metadata. Leave empty for self-hosted Riva.",
|
||||
"required": false,
|
||||
"field_type": "text",
|
||||
"options": null,
|
||||
"default_value": null
|
||||
}
|
||||
],
|
||||
"default_model_placeholder": "nvidia_riva/nvidia/parakeet-ctc-1_1b-asr"
|
||||
},
|
||||
{
|
||||
"provider": "Ollama",
|
||||
"provider_display_name": "Ollama",
|
||||
|
|
|
|||
|
|
@ -243,6 +243,39 @@ def test_bedrock_mantle_provider_fields():
|
|||
assert fields_by_key["api_base"]["field_type"] == "text"
|
||||
|
||||
|
||||
def test_nvidia_riva_provider_fields():
|
||||
"""The Add Model provider dropdown is populated from /public/providers/fields, so a
|
||||
missing entry meant Riva could not be added through the UI. Riva is gRPC only with no
|
||||
public default endpoint, hence the required api_base, and NVCF hosted Riva cannot be
|
||||
called without nvcf_function_id.
|
||||
"""
|
||||
app_instance = FastAPI()
|
||||
app_instance.include_router(router)
|
||||
test_client = TestClient(app_instance)
|
||||
|
||||
response = test_client.get("/public/providers/fields")
|
||||
assert response.status_code == 200
|
||||
providers = response.json()
|
||||
|
||||
riva = next((p for p in providers if p["provider"] == "NVIDIA_RIVA"), None)
|
||||
assert riva is not None, "NVIDIA Riva provider entry not found"
|
||||
|
||||
assert riva["provider_display_name"] == "Nvidia Riva"
|
||||
assert riva["litellm_provider"] == LlmProviders.NVIDIA_RIVA.value
|
||||
assert riva["default_model_placeholder"].startswith("nvidia_riva/")
|
||||
|
||||
fields_by_key = {f["key"]: f for f in riva["credential_fields"]}
|
||||
|
||||
assert fields_by_key["api_base"]["required"] is True
|
||||
assert fields_by_key["api_base"]["field_type"] == "text"
|
||||
|
||||
assert fields_by_key["api_key"]["required"] is False
|
||||
assert fields_by_key["api_key"]["field_type"] == "password"
|
||||
|
||||
assert "nvcf_function_id" in fields_by_key
|
||||
assert fields_by_key["nvcf_function_id"]["required"] is False
|
||||
|
||||
|
||||
def test_google_ai_studio_provider_fields_expose_api_base():
|
||||
"""The Google AI Studio (gemini) credential form must let admins set a custom
|
||||
api_base so they can point at a Gemini-compatible gateway (e.g. a self-hosted
|
||||
|
|
|
|||
|
|
@ -94,6 +94,17 @@ describe("provider_info_helpers", () => {
|
|||
expect(result.displayName).toBe(Providers.ZAI);
|
||||
});
|
||||
|
||||
it("should resolve the nvidia_riva provider value to the Nvidia Riva display name and logo", () => {
|
||||
// The backend registers nvidia_riva and it has a docs page, but the UI
|
||||
// registry had no entry, so it could not be picked in Add Model and the
|
||||
// slug rendered raw with no logo.
|
||||
const result = getProviderLogoAndName("nvidia_riva");
|
||||
expect(result.displayName).toBe(Providers.NVIDIA_RIVA);
|
||||
expect(provider_map.NVIDIA_RIVA).toBe("nvidia_riva");
|
||||
expect(result.logo).toBe(providerLogoMap[Providers.NVIDIA_RIVA]);
|
||||
expect(result.logo).toBeTruthy();
|
||||
});
|
||||
|
||||
it("should return provider value as display name when no mapping exists", () => {
|
||||
const unknownProvider = "unknown_provider";
|
||||
const result = getProviderLogoAndName(unknownProvider);
|
||||
|
|
@ -225,6 +236,10 @@ describe("provider_info_helpers", () => {
|
|||
expect(getPlaceholder(Providers.ZAI)).toBe("zai/glm-4.5");
|
||||
});
|
||||
|
||||
it("should return the riva asr placeholder for NVIDIA_RIVA provider", () => {
|
||||
expect(getPlaceholder(Providers.NVIDIA_RIVA)).toBe("nvidia_riva/nvidia/parakeet-ctc-1_1b-asr");
|
||||
});
|
||||
|
||||
it("should return default gpt-3.5-turbo placeholder for unknown provider", () => {
|
||||
expect(getPlaceholder("UnknownProvider" as any)).toBe("gpt-3.5-turbo");
|
||||
});
|
||||
|
|
|
|||
|
|
@ -130,6 +130,7 @@ export enum Providers {
|
|||
NOVITA = "Novita",
|
||||
NSCALE = "Nscale",
|
||||
NVIDIA_NIM = "Nvidia Nim",
|
||||
NVIDIA_RIVA = "Nvidia Riva",
|
||||
Ollama = "Ollama",
|
||||
OLLAMA_CHAT = "Ollama Chat",
|
||||
OOBABOOGA = "Oobabooga",
|
||||
|
|
@ -238,6 +239,7 @@ export const provider_map: Record<string, string> = {
|
|||
NOVITA: "novita",
|
||||
NSCALE: "nscale",
|
||||
NVIDIA_NIM: "nvidia_nim",
|
||||
NVIDIA_RIVA: "nvidia_riva",
|
||||
Ollama: "ollama",
|
||||
OLLAMA_CHAT: "ollama_chat",
|
||||
OOBABOOGA: "oobabooga",
|
||||
|
|
@ -334,6 +336,7 @@ export const providerLogoMap: Partial<Record<Providers, string>> = {
|
|||
[Providers.NEBIUS]: nebiusLogo.src,
|
||||
[Providers.NOVITA]: novitaLogo.src,
|
||||
[Providers.NVIDIA_NIM]: nvidiaNimLogo.src,
|
||||
[Providers.NVIDIA_RIVA]: nvidiaNimLogo.src,
|
||||
[Providers.Ollama]: ollamaLogo.src,
|
||||
[Providers.OLLAMA_CHAT]: ollamaLogo.src,
|
||||
[Providers.OOBABOOGA]: openaiSmallLogo.src,
|
||||
|
|
@ -400,50 +403,32 @@ export const getProviderLogoAndName = (providerValue: string): { logo: string; d
|
|||
return { logo, displayName };
|
||||
};
|
||||
|
||||
export const getPlaceholder = (selectedProvider: string): string => {
|
||||
if (selectedProvider === Providers.AIML) {
|
||||
return "aiml/flux-pro/v1.1";
|
||||
} else if (selectedProvider === Providers.Vertex_AI) {
|
||||
return "gemini-pro";
|
||||
} else if (selectedProvider == Providers.Anthropic) {
|
||||
return "claude-3-opus";
|
||||
} else if (selectedProvider == Providers.Bedrock) {
|
||||
return "claude-3-opus";
|
||||
} else if (selectedProvider == Providers.SageMaker) {
|
||||
return "sagemaker/jumpstart-dft-meta-textgeneration-llama-2-7b";
|
||||
} else if (selectedProvider == Providers.Google_AI_Studio) {
|
||||
return "gemini-pro";
|
||||
} else if (selectedProvider == Providers.Azure_AI_Studio) {
|
||||
return "azure_ai/command-r-plus";
|
||||
} else if (selectedProvider == Providers.Azure) {
|
||||
return "my-deployment";
|
||||
} else if (selectedProvider == Providers.Oracle) {
|
||||
return "oci/xai.grok-4";
|
||||
} else if (selectedProvider == Providers.Snowflake) {
|
||||
return "snowflake/mistral-7b";
|
||||
} else if (selectedProvider == Providers.Voyage) {
|
||||
return "voyage/";
|
||||
} else if (selectedProvider == Providers.JinaAI) {
|
||||
return "jina_ai/";
|
||||
} else if (selectedProvider == Providers.VolcEngine) {
|
||||
return "volcengine/<any-model-on-volcengine>";
|
||||
} else if (selectedProvider == Providers.DeepInfra) {
|
||||
return "deepinfra/<any-model-on-deepinfra>";
|
||||
} else if (selectedProvider == Providers.FalAI) {
|
||||
return "fal_ai/fal-ai/flux-pro/v1.1-ultra";
|
||||
} else if (selectedProvider == Providers.RunwayML) {
|
||||
return "runwayml/gen4_turbo";
|
||||
} else if (selectedProvider === Providers.WATSONX) {
|
||||
return "watsonx/ibm/granite-3-3-8b-instruct";
|
||||
} else if (selectedProvider === Providers.Cursor) {
|
||||
return "cursor/claude-4-sonnet";
|
||||
} else if (selectedProvider === Providers.ZAI) {
|
||||
return "zai/glm-4.5";
|
||||
} else {
|
||||
return "gpt-3.5-turbo";
|
||||
}
|
||||
const providerPlaceholderMap: Partial<Record<Providers, string>> = {
|
||||
[Providers.AIML]: "aiml/flux-pro/v1.1",
|
||||
[Providers.Anthropic]: "claude-3-opus",
|
||||
[Providers.Azure]: "my-deployment",
|
||||
[Providers.Azure_AI_Studio]: "azure_ai/command-r-plus",
|
||||
[Providers.Bedrock]: "claude-3-opus",
|
||||
[Providers.Cursor]: "cursor/claude-4-sonnet",
|
||||
[Providers.DeepInfra]: "deepinfra/<any-model-on-deepinfra>",
|
||||
[Providers.FalAI]: "fal_ai/fal-ai/flux-pro/v1.1-ultra",
|
||||
[Providers.Google_AI_Studio]: "gemini-pro",
|
||||
[Providers.JinaAI]: "jina_ai/",
|
||||
[Providers.NVIDIA_RIVA]: "nvidia_riva/nvidia/parakeet-ctc-1_1b-asr",
|
||||
[Providers.Oracle]: "oci/xai.grok-4",
|
||||
[Providers.RunwayML]: "runwayml/gen4_turbo",
|
||||
[Providers.SageMaker]: "sagemaker/jumpstart-dft-meta-textgeneration-llama-2-7b",
|
||||
[Providers.Snowflake]: "snowflake/mistral-7b",
|
||||
[Providers.Vertex_AI]: "gemini-pro",
|
||||
[Providers.VolcEngine]: "volcengine/<any-model-on-volcengine>",
|
||||
[Providers.Voyage]: "voyage/",
|
||||
[Providers.WATSONX]: "watsonx/ibm/granite-3-3-8b-instruct",
|
||||
[Providers.ZAI]: "zai/glm-4.5",
|
||||
};
|
||||
|
||||
export const getPlaceholder = (selectedProvider: string): string =>
|
||||
providerPlaceholderMap[selectedProvider as Providers] ?? "gpt-3.5-turbo";
|
||||
|
||||
export const getProviderModels = (provider: Providers, modelMap: any): Array<string> => {
|
||||
let providerKey = provider;
|
||||
let custom_llm_provider = provider_map[providerKey];
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue