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 f5d6d867786..3a2b4860eaa 100644 --- a/ui/litellm-dashboard/src/components/chat_ui/ChatUI.test.tsx +++ b/ui/litellm-dashboard/src/components/chat_ui/ChatUI.test.tsx @@ -1,6 +1,21 @@ import { fireEvent, render, screen, waitFor } from "@testing-library/react"; -import { beforeEach, describe, expect, it } from "vitest"; +import { beforeEach, describe, expect, it, vi } from "vitest"; import ChatUI from "./ChatUI"; +import * as fetchModelsModule from "./llm_calls/fetch_models"; + +// Mock the fetchAvailableModels function +vi.mock("./llm_calls/fetch_models", () => ({ + fetchAvailableModels: vi.fn(), +})); + +// Mock other networking functions that cause errors +vi.mock("../networking", () => ({ + tagListCall: vi.fn().mockResolvedValue({ data: [] }), + vectorStoreListCall: vi.fn().mockResolvedValue({ data: [] }), + getGuardrailsList: vi.fn().mockResolvedValue({ data: [] }), + mcpToolsCall: vi.fn().mockResolvedValue({ data: [] }), + modelHubCall: vi.fn().mockResolvedValue({ data: [] }), +})); // Mock scrollIntoView which is not available in jsdom beforeEach(() => { @@ -8,7 +23,23 @@ beforeEach(() => { }); describe("ChatUI", () => { - it("should render the chat UI", () => { + beforeEach(() => { + // Reset mocks before each test + vi.clearAllMocks(); + sessionStorage.clear(); + + // Mock scrollIntoView which is not available in JSDOM + Element.prototype.scrollIntoView = vi.fn(); + + // Mock the fetchAvailableModels to return test models + vi.mocked(fetchModelsModule.fetchAvailableModels).mockResolvedValue([ + { model_group: "Model 1", mode: "chat" }, + { model_group: "Model 2", mode: "chat" }, + { model_group: "Model 3", mode: "chat" }, + ]); + }); + + it("should render the chat UI", async () => { const { getByText } = render( { disabledPersonalKeyCreation={false} />, ); - expect(getByText("Test Key")).toBeInTheDocument(); + + await waitFor(() => { + expect(getByText("Test Key")).toBeInTheDocument(); + }); + }); + + it("should render the chat UI with selected models", async () => { + const { getByText, container, getAllByText } = render( + , + ); + + await waitFor(() => { + expect(getByText("Test Key")).toBeInTheDocument(); + }); + + // Wait for the first model to be auto-selected (default behavior) + await waitFor(() => { + expect(getByText("Model 1")).toBeInTheDocument(); + }); + + // Now click the Select dropdown to open it and check other models + const selectComponent = container.querySelectorAll(".ant-select-selector")[1]; // Second select is for models + expect(selectComponent).toBeTruthy(); + + // Click to open the dropdown + fireEvent.mouseDown(selectComponent!); + + // Wait for dropdown to open and check for all models + await waitFor(() => { + const dropdown = document.querySelector(".ant-select-dropdown"); + expect(dropdown).toBeTruthy(); + + // All three models should be in the dropdown as options + // Use getAllByTitle since Ant Design Select renders options with title attributes + const model2Elements = screen.queryAllByText("Model 2"); + const model3Elements = screen.queryAllByText("Model 3"); + + expect(model2Elements.length).toBeGreaterThan(0); + expect(model3Elements.length).toBeGreaterThan(0); + }); + }); + + it("should allow users to pick multiple models and render badges for selected models", async () => { + const { getByText, container } = render( + , + ); + + await waitFor(() => { + expect(getByText("Test Key")).toBeInTheDocument(); + }); + + await waitFor(() => { + expect(getByText("Model 1")).toBeInTheDocument(); + }); + + const selectComponent = container.querySelectorAll(".ant-select-selector")[1]; + expect(selectComponent).toBeTruthy(); + + fireEvent.mouseDown(selectComponent!); + + let model2Option: HTMLElement | null = null; + await waitFor(() => { + const dropdown = document.querySelector(".ant-select-dropdown"); + expect(dropdown).toBeTruthy(); + const options = document.querySelectorAll(".ant-select-item-option-content"); + model2Option = Array.from(options).find((opt) => opt.textContent === "Model 2") as HTMLElement; + expect(model2Option).toBeTruthy(); + }); + + fireEvent.click(model2Option!); + + await waitFor(() => { + const model1Badges = screen.getAllByText("Model 1"); + const model2Badges = screen.getAllByText("Model 2"); + + expect(model1Badges.length).toBeGreaterThan(0); + expect(model2Badges.length).toBeGreaterThan(0); + }); + + fireEvent.mouseDown(selectComponent!); + + let model3Option: HTMLElement | null = null; + await waitFor(() => { + const dropdown = document.querySelector(".ant-select-dropdown"); + expect(dropdown).toBeTruthy(); + const options = document.querySelectorAll(".ant-select-item-option-content"); + model3Option = Array.from(options).find((opt) => opt.textContent === "Model 3") as HTMLElement; + expect(model3Option).toBeTruthy(); + }); + + fireEvent.click(model3Option!); + + // Verify all three models have badges rendered + await waitFor(() => { + const model1Badges = screen.getAllByText("Model 1"); + const model2Badges = screen.getAllByText("Model 2"); + const model3Badges = screen.getAllByText("Model 3"); + + expect(model1Badges.length).toBeGreaterThan(0); + expect(model2Badges.length).toBeGreaterThan(0); + expect(model3Badges.length).toBeGreaterThan(0); + }); + const tags = container.querySelectorAll(".ant-tag"); + expect(tags.length).toBe(3); + }); + + it("should manage selectedModels state and its side effects", async () => { + const { getByText, container } = render( + , + ); + + await waitFor(() => { + expect(getByText("Test Key")).toBeInTheDocument(); + }); + + await waitFor(() => { + expect(getByText("Model 1")).toBeInTheDocument(); + }); + + await waitFor(() => { + const storedModels = sessionStorage.getItem("selectedModels"); + expect(storedModels).toBeTruthy(); + const parsedModels = JSON.parse(storedModels!); + expect(parsedModels).toContain("Model 1"); + }); + + const selectComponent = container.querySelectorAll(".ant-select-selector")[1]; + + fireEvent.mouseDown(selectComponent!); + + let model2Option: HTMLElement | null = null; + await waitFor(() => { + const dropdown = document.querySelector(".ant-select-dropdown"); + expect(dropdown).toBeTruthy(); + const options = document.querySelectorAll(".ant-select-item-option-content"); + model2Option = Array.from(options).find((opt) => opt.textContent === "Model 2") as HTMLElement; + expect(model2Option).toBeTruthy(); + }); + + fireEvent.click(model2Option!); + + await waitFor(() => { + const model2Badges = screen.getAllByText("Model 2"); + expect(model2Badges.length).toBeGreaterThan(0); + }); + + await waitFor(() => { + const storedModels = sessionStorage.getItem("selectedModels"); + const parsedModels = JSON.parse(storedModels!); + expect(parsedModels).toContain("Model 1"); + expect(parsedModels).toContain("Model 2"); + expect(parsedModels.length).toBe(2); + }); + + const tags = container.querySelectorAll(".ant-tag"); + expect(tags.length).toBe(2); + + const firstTagCloseButton = tags[0].querySelector(".ant-tag-close-icon"); + expect(firstTagCloseButton).toBeTruthy(); + fireEvent.click(firstTagCloseButton!); + + await waitFor(() => { + const storedModels = sessionStorage.getItem("selectedModels"); + const parsedModels = JSON.parse(storedModels!); + expect(parsedModels).not.toContain("Model 1"); + expect(parsedModels).toContain("Model 2"); + expect(parsedModels.length).toBe(1); + }); + + await waitFor(() => { + const remainingTags = container.querySelectorAll(".ant-tag"); + expect(remainingTags.length).toBe(1); + }); }); it("should show the voice selector when the endpoint type is audio_speech", async () => { diff --git a/ui/litellm-dashboard/src/components/chat_ui/ChatUI.tsx b/ui/litellm-dashboard/src/components/chat_ui/ChatUI.tsx index 5e85733749a..f0e8bfb1aff 100644 --- a/ui/litellm-dashboard/src/components/chat_ui/ChatUI.tsx +++ b/ui/litellm-dashboard/src/components/chat_ui/ChatUI.tsx @@ -1,63 +1,62 @@ -import React, { useState, useEffect, useRef } from "react"; +import { Card, Text, TextInput, Title, Button as TremorButton } from "@tremor/react"; +import React, { useEffect, useRef, useState } from "react"; import ReactMarkdown from "react-markdown"; -import { Card, Title, Text, TextInput, Button as TremorButton } from "@tremor/react"; import { v4 as uuidv4 } from "uuid"; - -import { Select, Spin, Typography, Tooltip, Input, Upload, Modal, Button } from "antd"; -import { makeOpenAIChatCompletionRequest } from "./llm_calls/chat_completion"; -import { makeOpenAIImageGenerationRequest } from "./llm_calls/image_generation"; -import { makeOpenAIImageEditsRequest } from "./llm_calls/image_edits"; -import { makeOpenAIAudioSpeechRequest } from "./llm_calls/audio_speech"; -import { makeOpenAIAudioTranscriptionRequest } from "./llm_calls/audio_transcriptions"; -import { makeOpenAIResponsesRequest } from "./llm_calls/responses_api"; -import { makeAnthropicMessagesRequest } from "./llm_calls/anthropic_messages"; -import { fetchAvailableModels, ModelGroup } from "./llm_calls/fetch_models"; -import { fetchAvailableMCPTools } from "./llm_calls/fetch_mcp_tools"; -import type { MCPTool } from "./llm_calls/fetch_mcp_tools"; -import { EndpointType } from "./mode_endpoint_mapping"; -import { Prism as SyntaxHighlighter } from "react-syntax-highlighter"; -import { coy } from "react-syntax-highlighter/dist/esm/styles/prism"; -import EndpointSelector from "./EndpointSelector"; -import TagSelector from "../tag_management/TagSelector"; -import VectorStoreSelector from "../vector_store_management/VectorStoreSelector"; -import GuardrailSelector from "../guardrails/GuardrailSelector"; -import { generateCodeSnippet } from "./CodeSnippets"; -import { MessageType } from "./types"; -import ReasoningContent from "./ReasoningContent"; -import ResponseMetrics, { TokenUsage } from "./ResponseMetrics"; -import ResponsesImageUpload from "./ResponsesImageUpload"; -import ResponsesImageRenderer from "./ResponsesImageRenderer"; -import { createMultimodalMessage, createDisplayMessage } from "./ResponsesImageUtils"; -import ChatImageUpload from "./ChatImageUpload"; -import ChatImageRenderer from "./ChatImageRenderer"; -import { createChatMultimodalMessage, createChatDisplayMessage } from "./ChatImageUtils"; -import AudioRenderer from "./AudioRenderer"; -import SessionManagement from "./SessionManagement"; -import MCPEventsDisplay, { MCPEvent } from "./MCPEventsDisplay"; -import { SearchResultsDisplay } from "./SearchResultsDisplay"; import { ApiOutlined, - KeyOutlined, - ClearOutlined, - RobotOutlined, - UserOutlined, - DeleteOutlined, - LoadingOutlined, - TagsOutlined, - DatabaseOutlined, - InfoCircleOutlined, - SafetyOutlined, - PictureOutlined, - CodeOutlined, - SoundOutlined, - ToolOutlined, - FilePdfOutlined, ArrowUpOutlined, + ClearOutlined, + CodeOutlined, + DatabaseOutlined, + DeleteOutlined, + FilePdfOutlined, + InfoCircleOutlined, + KeyOutlined, + LoadingOutlined, + PictureOutlined, + RobotOutlined, + SafetyOutlined, + SoundOutlined, + TagsOutlined, + ToolOutlined, + UserOutlined, } from "@ant-design/icons"; +import { Button, Input, Modal, Select, Spin, Tag, Tooltip, Typography, Upload } from "antd"; +import { Prism as SyntaxHighlighter } from "react-syntax-highlighter"; +import { coy } from "react-syntax-highlighter/dist/esm/styles/prism"; +import GuardrailSelector from "../guardrails/GuardrailSelector"; import NotificationsManager from "../molecules/notifications_manager"; -import { makeOpenAIEmbeddingsRequest } from "./llm_calls/embeddings_api"; -import { truncateString } from "./chatUtils"; +import TagSelector from "../tag_management/TagSelector"; +import VectorStoreSelector from "../vector_store_management/VectorStoreSelector"; +import AudioRenderer from "./AudioRenderer"; import { OPEN_AI_VOICE_SELECT_OPTIONS, OpenAIVoice } from "./chatConstants"; +import ChatImageRenderer from "./ChatImageRenderer"; +import ChatImageUpload from "./ChatImageUpload"; +import { createChatDisplayMessage, createChatMultimodalMessage } from "./ChatImageUtils"; +import { truncateString } from "./chatUtils"; +import { generateCodeSnippet } from "./CodeSnippets"; +import EndpointSelector from "./EndpointSelector"; +import { makeAnthropicMessagesRequest } from "./llm_calls/anthropic_messages"; +import { makeOpenAIAudioSpeechRequest } from "./llm_calls/audio_speech"; +import { makeOpenAIAudioTranscriptionRequest } from "./llm_calls/audio_transcriptions"; +import { makeOpenAIChatCompletionRequest } from "./llm_calls/chat_completion"; +import { makeOpenAIEmbeddingsRequest } from "./llm_calls/embeddings_api"; +import type { MCPTool } from "./llm_calls/fetch_mcp_tools"; +import { fetchAvailableMCPTools } from "./llm_calls/fetch_mcp_tools"; +import { fetchAvailableModels, ModelGroup } from "./llm_calls/fetch_models"; +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 ReasoningContent from "./ReasoningContent"; +import ResponseMetrics, { TokenUsage } from "./ResponseMetrics"; +import ResponsesImageRenderer from "./ResponsesImageRenderer"; +import ResponsesImageUpload from "./ResponsesImageUpload"; +import { createDisplayMessage, createMultimodalMessage } from "./ResponsesImageUtils"; +import { SearchResultsDisplay } from "./SearchResultsDisplay"; +import SessionManagement from "./SessionManagement"; +import { MessageType } from "./types"; const { TextArea } = Input; const { Dragger } = Upload; @@ -107,10 +106,17 @@ const ChatUI: React.FC = ({ accessToken, token, userRole, userID, d return []; } }); - const [selectedModel, setSelectedModel] = useState( - () => sessionStorage.getItem("selectedModel") || undefined, - ); + const [selectedModels, setSelectedModels] = useState(() => { + const saved = sessionStorage.getItem("selectedModels"); + try { + return saved ? JSON.parse(saved) : []; + } catch (error) { + console.error("Error parsing selectedModels from sessionStorage", error); + return []; + } + }); const [showCustomModelInput, setShowCustomModelInput] = useState(false); + const [customModelInput, setCustomModelInput] = useState(""); const [modelInfo, setModelInfo] = useState([]); const customModelTimeout = useRef(null); const [endpointType, setEndpointType] = useState( @@ -214,7 +220,7 @@ const ChatUI: React.FC = ({ accessToken, token, userRole, userID, d selectedGuardrails, selectedMCPTools, endpointType, - selectedModel, + selectedModel: selectedModels[0], // Use first model for code snippet selectedSdk, selectedVoice, }); @@ -233,7 +239,7 @@ const ChatUI: React.FC = ({ accessToken, token, userRole, userID, d selectedGuardrails, selectedMCPTools, endpointType, - selectedModel, + selectedModels, ]); useEffect(() => { @@ -256,11 +262,7 @@ const ChatUI: React.FC = ({ accessToken, token, userRole, userID, d sessionStorage.setItem("selectedMCPTools", JSON.stringify(selectedMCPTools)); sessionStorage.setItem("selectedVoice", selectedVoice); - if (selectedModel) { - sessionStorage.setItem("selectedModel", selectedModel); - } else { - sessionStorage.removeItem("selectedModel"); - } + sessionStorage.setItem("selectedModels", JSON.stringify(selectedModels)); if (messageTraceId) { sessionStorage.setItem("messageTraceId", messageTraceId); } else { @@ -275,7 +277,7 @@ const ChatUI: React.FC = ({ accessToken, token, userRole, userID, d }, [ apiKeySource, apiKey, - selectedModel, + selectedModels, endpointType, selectedTags, selectedVectorStores, @@ -307,11 +309,11 @@ const ChatUI: React.FC = ({ accessToken, token, userRole, userID, d setModelInfo(uniqueModels); // check for selection overlap or empty model list - const hasSelection = uniqueModels.some((m) => m.model_group === selectedModel); if (!uniqueModels.length) { - setSelectedModel(undefined); - } else if (!hasSelection) { - setSelectedModel(uniqueModels[0].model_group); + setSelectedModels([]); + } else if (selectedModels.length === 0) { + // If no models selected, auto-select the first one + setSelectedModels([uniqueModels[0].model_group]); } } catch (error) { console.error("Error fetching model info:", error); @@ -724,7 +726,8 @@ const ChatUI: React.FC = ({ accessToken, token, userRole, userID, d setIsLoading(true); try { - if (selectedModel) { + if (selectedModels.length > 0) { + const selectedModel = selectedModels[0]; // For now, use first model (will be updated later for multi-model) if (endpointType === EndpointType.CHAT) { // Create chat history for API call - strip out model field and isImage field // For chat completions, we preserve the multimodal content structure @@ -932,11 +935,43 @@ const ChatUI: React.FC = ({ accessToken, token, userRole, userID, d ); } - const onModelChange = (value: string) => { - console.log(`selected ${value}`); - setSelectedModel(value); + const onModelsChange = (values: string[]) => { + console.log(`selected models:`, values); - setShowCustomModelInput(value === "custom"); + // Check if "custom" was added + if (values.includes("custom")) { + setShowCustomModelInput(true); + // Remove "custom" from the actual selection + setSelectedModels(values.filter((v) => v !== "custom")); + } else { + // Limit to 3 models maximum + if (values.length > 3) { + NotificationsManager.info("Maximum 3 models can be selected"); + return; + } + setSelectedModels(values); + } + }; + + const handleAddCustomModel = () => { + const trimmedModel = customModelInput.trim(); + if (!trimmedModel) return; + + if (selectedModels.includes(trimmedModel)) { + NotificationsManager.info("Model already selected"); + return; + } + if (selectedModels.length >= 3) { + NotificationsManager.info("Maximum 3 models can be selected"); + return; + } + setSelectedModels([...selectedModels, trimmedModel]); + setCustomModelInput(""); + setShowCustomModelInput(false); + }; + + const handleRemoveModel = (modelToRemove: string) => { + setSelectedModels(selectedModels.filter((m) => m !== modelToRemove)); }; const antIcon = ; @@ -980,40 +1015,65 @@ const ChatUI: React.FC = ({ accessToken, token, userRole, userID, d
- Select Model + Select Models