mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
[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:
parent
78ed5126a5
commit
2b0aea87c0
3 changed files with 385 additions and 103 deletions
|
|
@ -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 () => {
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue