From af68763f5d5010e8308ab7f49a61bf87108a79a9 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Mon, 10 Nov 2025 18:38:01 -0800 Subject: [PATCH] Only show models based on selected endpoint (#16452) --- .../src/components/chat_ui/ChatUI.test.tsx | 62 ++++++++-- .../src/components/chat_ui/ChatUI.tsx | 114 +++++++++++------- .../chat_ui/mode_endpoint_mapping.tsx | 2 + 3 files changed, 127 insertions(+), 51 deletions(-) diff --git a/ui/litellm-dashboard/src/components/chat_ui/ChatUI.test.tsx b/ui/litellm-dashboard/src/components/chat_ui/ChatUI.test.tsx index 6670395aad5..53fa2b0916e 100644 --- a/ui/litellm-dashboard/src/components/chat_ui/ChatUI.test.tsx +++ b/ui/litellm-dashboard/src/components/chat_ui/ChatUI.test.tsx @@ -1,4 +1,4 @@ -import { fireEvent, render, screen, waitFor } from "@testing-library/react"; +import { act, fireEvent, render, screen, waitFor } from "@testing-library/react"; import { beforeEach, describe, expect, it, vi } from "vitest"; import ChatUI from "./ChatUI"; import * as fetchModelsModule from "./llm_calls/fetch_models"; @@ -118,18 +118,64 @@ describe("ChatUI", () => { expect(getByText("Test Key")).toBeInTheDocument(); }); - await waitFor(() => { - expect(getByText("Model 1")).toBeInTheDocument(); - }); + // Open the "Select Model" dropdown (AntD renders options in a portal) + const selectModelLabel = getByText("Select Model"); + const modelSelect = selectModelLabel.parentElement?.querySelector(".ant-select-selector"); + expect(modelSelect).toBeTruthy(); - const selectComponent = container.querySelectorAll(".ant-select-selector")[1]; - expect(selectComponent).toBeTruthy(); - - fireEvent.mouseDown(selectComponent!); + fireEvent.mouseDown(modelSelect!); await waitFor(() => { const model1Label = screen.getAllByText("Model 1"); expect(model1Label.length).toBeGreaterThan(0); }); }); + + it("shows only chat-compatible models when chat endpoint is selected", async () => { + vi.mocked(fetchModelsModule.fetchAvailableModels).mockResolvedValueOnce([ + { model_group: "ChatModel", mode: "chat" }, + { model_group: "SpeechModel", mode: "audio_speech" }, + { model_group: "ImageModel", mode: "image_generation" }, + { model_group: "ResponsesModel", mode: "responses" }, + ]); + + const { getByText, baseElement } = render( + , + ); + + await waitFor(() => { + expect(getByText("Test Key")).toBeInTheDocument(); + }); + + // Open endpoint selector and explicitly select /v1/chat/completions + const endpointTypeText = getByText("Endpoint Type"); + const endpointSelect = endpointTypeText.parentElement?.querySelector(".ant-select-selector"); + expect(endpointSelect).toBeTruthy(); + act(() => { + fireEvent.mouseDown(endpointSelect!); + fireEvent.click(screen.getByText("/v1/chat/completions")); + }); + + // Open model selector + const selectModelLabel = getByText("Select Model"); + const modelSelect = selectModelLabel.parentElement?.querySelector(".ant-select-selector"); + expect(modelSelect).toBeTruthy(); + act(() => { + fireEvent.mouseDown(modelSelect!); + }); + + await waitFor(() => { + // Chat-compatible: ChatModel should be visible + expect(screen.getAllByText("ChatModel").length).toBeGreaterThan(0); + expect(screen.queryByText("SpeechModel")).toBeNull(); + expect(screen.queryByText("ImageModel")).toBeNull(); + expect(screen.queryByText("ResponsesModel")).toBeNull(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/chat_ui/ChatUI.tsx b/ui/litellm-dashboard/src/components/chat_ui/ChatUI.tsx index a006d6734b2..b619c755b57 100644 --- a/ui/litellm-dashboard/src/components/chat_ui/ChatUI.tsx +++ b/ui/litellm-dashboard/src/components/chat_ui/ChatUI.tsx @@ -48,7 +48,7 @@ import { makeOpenAIImageEditsRequest } from "./llm_calls/image_edits"; import { makeOpenAIImageGenerationRequest } from "./llm_calls/image_generation"; import { makeOpenAIResponsesRequest } from "./llm_calls/responses_api"; import MCPEventsDisplay, { MCPEvent } from "./MCPEventsDisplay"; -import { EndpointType } from "./mode_endpoint_mapping"; +import { EndpointType, getEndpointType } from "./mode_endpoint_mapping"; import ReasoningContent from "./ReasoningContent"; import ResponseMetrics, { TokenUsage } from "./ResponseMetrics"; import ResponsesImageRenderer from "./ResponsesImageRenderer"; @@ -106,9 +106,7 @@ const ChatUI: React.FC = ({ accessToken, token, userRole, userID, d return []; } }); - const [selectedModel, setSelectedModel] = useState( - () => sessionStorage.getItem("selectedModel") || undefined, - ); + const [selectedModel, setSelectedModel] = useState(undefined); const [showCustomModelInput, setShowCustomModelInput] = useState(false); const [modelInfo, setModelInfo] = useState([]); const customModelTimeout = useRef(null); @@ -311,7 +309,7 @@ const ChatUI: React.FC = ({ accessToken, token, userRole, userID, d if (!uniqueModels.length) { setSelectedModel(undefined); } else if (!hasSelection) { - setSelectedModel(uniqueModels[0].model_group); + setSelectedModel(undefined); } } catch (error) { console.error("Error fetching model info:", error); @@ -978,44 +976,6 @@ const ChatUI: React.FC = ({ accessToken, token, userRole, userID, d )} -
- - Select Model - - { + if (!option.mode) { + //If no mode, show all models + return true; + } + const optionEndpoint = getEndpointType(option.mode); + // Show chat models for responses/anthropic_messages endpoints as they are compatible + if ( + endpointType === EndpointType.RESPONSES || + endpointType === EndpointType.ANTHROPIC_MESSAGES + ) { + return optionEndpoint === endpointType || optionEndpoint === EndpointType.CHAT; + } + // Show image models for image_edits endpoint as they are compatible + if (endpointType === EndpointType.IMAGE_EDITS) { + return optionEndpoint === endpointType || optionEndpoint === EndpointType.IMAGE; + } + return optionEndpoint === endpointType; + }) + .map((option) => option.model_group), + ), + ).map((model_group, index) => ({ + value: model_group, + label: model_group, + key: index, + })), + { value: "custom", label: "Enter custom model", key: "custom" }, + ]} + style={{ width: "100%" }} + showSearch={true} + className="rounded-md" + /> + {showCustomModelInput && ( + { + // Using setTimeout to create a simple debounce effect + if (customModelTimeout.current) { + clearTimeout(customModelTimeout.current); + } + + customModelTimeout.current = setTimeout(() => { + setSelectedModel(value); + }, 500); // 500ms delay after typing stops + }} + /> + )} +
+
Tags diff --git a/ui/litellm-dashboard/src/components/chat_ui/mode_endpoint_mapping.tsx b/ui/litellm-dashboard/src/components/chat_ui/mode_endpoint_mapping.tsx index 2aad64fb619..479d3cb8940 100644 --- a/ui/litellm-dashboard/src/components/chat_ui/mode_endpoint_mapping.tsx +++ b/ui/litellm-dashboard/src/components/chat_ui/mode_endpoint_mapping.tsx @@ -10,6 +10,7 @@ export enum ModelMode { RESPONSES = "responses", IMAGE_EDITS = "image_edits", ANTHROPIC_MESSAGES = "anthropic_messages", + EMBEDDING = "embedding", // add additional modes as needed } @@ -37,6 +38,7 @@ export const litellmModeMapping: Record = { [ModelMode.ANTHROPIC_MESSAGES]: EndpointType.ANTHROPIC_MESSAGES, [ModelMode.AUDIO_SPEECH]: EndpointType.SPEECH, [ModelMode.AUDIO_TRANSCRIPTION]: EndpointType.TRANSCRIPTION, + [ModelMode.EMBEDDING]: EndpointType.EMBEDDINGS, }; export const getEndpointType = (mode: string): EndpointType => {