[Feature] UI - Initial changes for supporting prompts to multiple models (#16223)

* Initial changes for supporting prompts to multiple models

* Additional tests to prevent regressions

* Resolve merge conflicts

* Fixing E2E build issues
This commit is contained in:
yuneng-jiang 2025-11-04 16:02:19 -08:00 • committed by GitHub
parent 78ed5126a5
commit 2b0aea87c0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 385 additions and 103 deletions

View file

@ -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(
<ChatUI
accessToken="1234567890"
@ -18,7 +49,197 @@ describe("ChatUI", () => {
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(
<ChatUI
accessToken="1234567890"
token="1234567890"
userRole="user"
userID="1234567890"
disabledPersonalKeyCreation={false}
/>,
);
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(
<ChatUI
accessToken="1234567890"
token="1234567890"
userRole="user"
userID="1234567890"
disabledPersonalKeyCreation={false}
/>,
);
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(
<ChatUI
accessToken="1234567890"
token="1234567890"
userRole="user"
userID="1234567890"
disabledPersonalKeyCreation={false}
/>,
);
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 () => {

View file

@ -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<ChatUIProps> = ({ accessToken, token, userRole, userID, d
return [];
}
});
const [selectedModel, setSelectedModel] = useState<string | undefined>(
() => sessionStorage.getItem("selectedModel") || undefined,
);
const [selectedModels, setSelectedModels] = useState<string[]>(() => {
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<boolean>(false);
const [customModelInput, setCustomModelInput] = useState<string>("");
const [modelInfo, setModelInfo] = useState<ModelGroup[]>([]);
const customModelTimeout = useRef<NodeJS.Timeout | null>(null);
const [endpointType, setEndpointType] = useState<string>(
@ -214,7 +220,7 @@ const ChatUI: React.FC<ChatUIProps> = ({ 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<ChatUIProps> = ({ accessToken, token, userRole, userID, d
selectedGuardrails,
selectedMCPTools,
endpointType,
selectedModel,
selectedModels,
]);
useEffect(() => {
@ -256,11 +262,7 @@ const ChatUI: React.FC<ChatUIProps> = ({ 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<ChatUIProps> = ({ accessToken, token, userRole, userID, d
}, [
apiKeySource,
apiKey,
selectedModel,
selectedModels,
endpointType,
selectedTags,
selectedVectorStores,
@ -307,11 +309,11 @@ const ChatUI: React.FC<ChatUIProps> = ({ 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<ChatUIProps> = ({ 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<ChatUIProps> = ({ 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 = <LoadingOutlined style={{ fontSize: 24 }} spin />;
@ -980,40 +1015,65 @@ const ChatUI: React.FC<ChatUIProps> = ({ accessToken, token, userRole, userID, d
<div>
<Text className="font-medium block mb-2 text-gray-700 flex items-center">
<RobotOutlined className="mr-2" /> Select Model
<RobotOutlined className="mr-2" /> Select Models
</Text>
<Select
value={selectedModel}
placeholder="Select a Model"
onChange={onModelChange}
mode="multiple"
value={selectedModels}
placeholder="Select one or more models"
onChange={onModelsChange}
options={[
...Array.from(new Set(modelInfo.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" },
{ value: "custom", label: "+ Add custom model", key: "custom" },
]}
style={{ width: "100%" }}
showSearch={true}
className="rounded-md"
maxTagCount={0}
maxTagPlaceholder={`${selectedModels.length} model${selectedModels.length !== 1 ? "s" : ""} selected`}
/>
{showCustomModelInput && (
<TextInput
className="mt-2"
placeholder="Enter custom model name"
onValueChange={(value) => {
// Using setTimeout to create a simple debounce effect
if (customModelTimeout.current) {
clearTimeout(customModelTimeout.current);
placeholder="Custom Model Name (Enter to add)"
value={customModelInput}
onValueChange={setCustomModelInput}
onKeyDown={(e) => {
if (e.key === "Enter") {
handleAddCustomModel();
}
customModelTimeout.current = setTimeout(() => {
setSelectedModel(value);
}, 500); // 500ms delay after typing stops
}}
/>
)}
{/* Selected Models Tags */}
{selectedModels.length > 0 && (
<div className="mt-2 overflow-hidden">
<div className="flex flex-wrap gap-1.5">
{selectedModels.map((model) => (
<Tag
key={model}
closable
onClose={(e) => {
e.preventDefault();
handleRemoveModel(model);
}}
color="blue"
className="px-2 py-0.5 text-xs max-w-[200px]"
>
<span className="truncate inline-block max-w-[160px] align-middle" title={model}>
{model}
</span>
</Tag>
))}
</div>
</div>
)}
</div>
<div>
@ -1165,8 +1225,8 @@ const ChatUI: React.FC<ChatUIProps> = ({ accessToken, token, userRole, userID, d
{/* Main Chat Area */}
<div className="w-3/4 flex flex-col bg-white">
<div className="p-4 border-b border-gray-200 flex justify-between items-center">
<Title className="text-xl font-semibold mb-0">Test Key</Title>
<div className="p-4 border-b border-gray-200 flex items-center justify-between">
<Title className="text-xl font-semibold">Test Key</Title>
<div className="flex gap-2">
<TremorButton
onClick={clearChatHistory}

View file

@ -129,10 +129,7 @@ const GuardrailDetails = ({ entry, index, total }: GuardrailDetailsProps) => {
<h5 className="font-medium mb-2">Masked Entity Summary</h5>
<div className="flex flex-wrap gap-2">
{Object.entries(maskedEntityCount).map(([entityType, count]) => (
<span
key={entityType}
className="px-3 py-1.5 bg-blue-50 text-blue-700 rounded-md text-xs font-medium"
>
<span key={entityType} className="px-3 py-1.5 bg-blue-50 text-blue-700 rounded-md text-xs font-medium">
{entityType}: {count}
</span>
))}
@ -159,28 +156,32 @@ const GuardrailViewer = ({ data }: GuardrailViewerProps) => {
const guardrailEntries = Array.isArray(data)
? data.filter((entry): entry is GuardrailInformation => Boolean(entry))
: data
? [data]
: [];
if (guardrailEntries.length === 0) {
return null;
}
? [data]
: [];
const [sectionExpanded, setSectionExpanded] = useState(true);
const primaryName = guardrailEntries.length === 1 ? guardrailEntries[0].guardrail_name : `${guardrailEntries.length} guardrails`;
const primaryName =
guardrailEntries.length === 1 ? guardrailEntries[0].guardrail_name : `${guardrailEntries.length} guardrails`;
const statuses = Array.from(new Set(guardrailEntries.map((entry) => entry.guardrail_status)));
const allSucceeded = statuses.every((status) => (status ?? "").toLowerCase() === "success");
const aggregatedStatus = allSucceeded ? "success" : "failure";
const totalMaskedEntities = guardrailEntries.reduce((sum, entry) => {
return (
sum +
Object.values(entry.masked_entity_count || {}).reduce((acc, count) => acc + (typeof count === "number" ? count : 0), 0)
Object.values(entry.masked_entity_count || {}).reduce(
(acc, count) => acc + (typeof count === "number" ? count : 0),
0,
)
);
}, 0);
const tooltipTitle = allSucceeded ? null : "Guardrail failed to run.";
if (guardrailEntries.length === 0) {
return null;
}
return (
<div className="bg-white rounded-lg shadow mb-6">
<div