Add additional model settings to chat models in test key (#16793)

This commit is contained in:
yuneng-jiang 2025-11-18 19:49:16 -08:00 • committed by GitHub
parent fe05e33723
commit d0e806d1b3
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 375 additions and 6 deletions

View file

@ -0,0 +1,50 @@
import { render, screen, waitFor } from "@testing-library/react";
import userEvent from "@testing-library/user-event";
import { describe, expect, it, vi } from "vitest";
import AdditionalModelSettings from "./AdditionalModelSettings";
describe("AdditionalModelSettings", () => {
it("should render correctly", () => {
const { container } = render(<AdditionalModelSettings />);
expect(container).toBeTruthy();
expect(screen.getByText("Use Advanced Parameters")).toBeInTheDocument();
expect(screen.getByText("Temperature")).toBeInTheDocument();
expect(screen.getByText("Max Tokens")).toBeInTheDocument();
});
it("should enable the sliders when the checkbox is checked", async () => {
const user = userEvent.setup();
const mockOnTemperatureChange = vi.fn();
const mockOnMaxTokensChange = vi.fn();
render(
<AdditionalModelSettings
onTemperatureChange={mockOnTemperatureChange}
onMaxTokensChange={mockOnMaxTokensChange}
/>,
);
const checkbox = screen.getByRole("checkbox", { name: /Use Advanced Parameters/i });
expect(checkbox).toBeInTheDocument();
expect(checkbox).not.toBeChecked();
await user.click(checkbox);
await waitFor(() => {
expect(screen.getByRole("checkbox", { name: /Use Advanced Parameters/i })).toBeChecked();
});
await waitFor(() => {
const sliders = screen.getAllByRole("slider");
expect(sliders.length).toBeGreaterThan(0);
expect(sliders[0]).not.toBeDisabled();
});
const temperatureSlider = screen.getAllByRole("slider")[0];
const maxTokensSlider = screen.getAllByRole("slider")[1];
expect(temperatureSlider).not.toBeDisabled();
expect(maxTokensSlider).not.toBeDisabled();
});
});

View file

@ -0,0 +1,137 @@
import { InfoCircleOutlined } from "@ant-design/icons";
import { Text } from "@tremor/react";
import { Checkbox, InputNumber, Slider, Tooltip } from "antd";
import React, { useEffect, useState } from "react";
interface AdditionalModelSettingsProps {
temperature?: number;
maxTokens?: number;
useAdvancedParams?: boolean;
onTemperatureChange?: (value: number) => void;
onMaxTokensChange?: (value: number) => void;
onUseAdvancedParamsChange?: (value: boolean) => void;
}
const AdditionalModelSettings: React.FC<AdditionalModelSettingsProps> = ({
temperature = 1.0,
maxTokens = 2048,
useAdvancedParams: externalUseAdvancedParams,
onTemperatureChange,
onMaxTokensChange,
onUseAdvancedParamsChange,
}) => {
const [internalUseAdvancedParams, setInternalUseAdvancedParams] = useState(false);
const useAdvancedParams =
externalUseAdvancedParams !== undefined ? externalUseAdvancedParams : internalUseAdvancedParams;
const [localTemperature, setLocalTemperature] = useState(temperature);
const [localMaxTokens, setLocalMaxTokens] = useState(maxTokens);
// Sync local state with props when they change
useEffect(() => {
setLocalTemperature(temperature);
}, [temperature]);
useEffect(() => {
setLocalMaxTokens(maxTokens);
}, [maxTokens]);
const handleTemperatureChange = (value: number | null) => {
const newValue = value ?? 1.0;
setLocalTemperature(newValue);
onTemperatureChange?.(newValue);
};
const handleMaxTokensChange = (value: number | null) => {
const newValue = value ?? 1000;
setLocalMaxTokens(newValue);
onMaxTokensChange?.(newValue);
};
const disabledOpacity = useAdvancedParams ? 1 : 0.4;
const disabledTextColor = useAdvancedParams ? "text-gray-700" : "text-gray-400";
const handleUseAdvancedParamsChange = (checked: boolean) => {
if (onUseAdvancedParamsChange) {
onUseAdvancedParamsChange(checked);
} else {
setInternalUseAdvancedParams(checked);
}
};
return (
<div className="space-y-4 p-4 w-80">
<Checkbox checked={useAdvancedParams} onChange={(e) => handleUseAdvancedParamsChange(e.target.checked)}>
<span className="font-medium">Use Advanced Parameters</span>
</Checkbox>
<div className="space-y-4 transition-opacity duration-200" style={{ opacity: disabledOpacity }}>
<div>
<div className="flex items-center justify-between mb-2">
<div className="flex items-center gap-1">
<Text className={`text-sm ${disabledTextColor}`}>Temperature</Text>
<Tooltip title="Controls randomness. Lower values make output more deterministic, higher values more creative.">
<InfoCircleOutlined className={`text-xs ${disabledTextColor} cursor-help`} />
</Tooltip>
</div>
<InputNumber
min={0}
max={2}
step={0.1}
value={localTemperature}
onChange={handleTemperatureChange}
disabled={!useAdvancedParams}
precision={1}
className="w-20"
/>
</div>
<Slider
min={0}
max={2}
step={0.1}
value={localTemperature}
onChange={handleTemperatureChange}
disabled={!useAdvancedParams}
marks={{
0: "0",
1: "1.0",
2: "2.0",
}}
/>
</div>
<div>
<div className="flex items-center justify-between mb-2">
<div className="flex items-center gap-1">
<Text className={`text-sm ${disabledTextColor}`}>Max Tokens</Text>
<Tooltip title="Maximum number of tokens to generate in the response.">
<InfoCircleOutlined className={`text-xs ${disabledTextColor} cursor-help`} />
</Tooltip>
</div>
<InputNumber
min={1}
max={32768}
step={1}
value={localMaxTokens}
onChange={handleMaxTokensChange}
disabled={!useAdvancedParams}
/>
</div>
<Slider
min={1}
max={32768}
step={1}
value={localMaxTokens}
onChange={handleMaxTokensChange}
disabled={!useAdvancedParams}
marks={{
1: "1",
32768: "32768",
}}
/>
</div>
</div>
</div>
);
};
export default AdditionalModelSettings;

View file

@ -120,7 +120,9 @@ describe("ChatUI", () => {
// Open the "Select Model" dropdown (AntD renders options in a portal)
const selectModelLabel = getByText("Select Model");
const modelSelect = selectModelLabel.parentElement?.querySelector(".ant-select-selector");
// The Select component is a sibling of the Text component, so we need to find it in the parent container
const modelSelectContainer = selectModelLabel.closest("div");
const modelSelect = modelSelectContainer?.querySelector(".ant-select-selector");
expect(modelSelect).toBeTruthy();
fireEvent.mouseDown(modelSelect!);
@ -164,7 +166,9 @@ describe("ChatUI", () => {
// Open model selector
const selectModelLabel = getByText("Select Model");
const modelSelect = selectModelLabel.parentElement?.querySelector(".ant-select-selector");
// The Select component is a sibling of the Text component, so we need to find it in the parent container
const modelSelectContainer = selectModelLabel.closest("div");
const modelSelect = modelSelectContainer?.querySelector(".ant-select-selector");
expect(modelSelect).toBeTruthy();
act(() => {
fireEvent.mouseDown(modelSelect!);

View file

@ -12,28 +12,30 @@ import {
PictureOutlined,
RobotOutlined,
SafetyOutlined,
SettingOutlined,
SoundOutlined,
TagsOutlined,
ToolOutlined,
UserOutlined,
} from "@ant-design/icons";
import { Card, Text, TextInput, Title, Button as TremorButton } from "@tremor/react";
import { Button, Input, Modal, Select, Spin, Tooltip, Typography, Upload } from "antd";
import { Button, Input, Modal, Popover, Select, Spin, Tooltip, Typography, Upload } from "antd";
import React, { useEffect, useRef, useState } from "react";
import ReactMarkdown from "react-markdown";
import { Prism as SyntaxHighlighter } from "react-syntax-highlighter";
import { coy } from "react-syntax-highlighter/dist/esm/styles/prism";
import { v4 as uuidv4 } from "uuid";
import { truncateString } from "../../utils/textUtils";
import GuardrailSelector from "../guardrails/GuardrailSelector";
import NotificationsManager from "../molecules/notifications_manager";
import TagSelector from "../tag_management/TagSelector";
import VectorStoreSelector from "../vector_store_management/VectorStoreSelector";
import AdditionalModelSettings from "./AdditionalModelSettings";
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 "../../utils/textUtils";
import { generateCodeSnippet } from "./CodeSnippets";
import EndpointSelector from "./EndpointSelector";
import { makeAnthropicMessagesRequest } from "./llm_calls/anthropic_messages";
@ -184,6 +186,9 @@ const ChatUI: React.FC<ChatUIProps> = ({
const [generatedCode, setGeneratedCode] = useState("");
const [selectedSdk, setSelectedSdk] = useState<"openai" | "azure">("openai");
const [mcpEvents, setMCPEvents] = useState<MCPEvent[]>([]);
const [temperature, setTemperature] = useState<number>(1.0);
const [maxTokens, setMaxTokens] = useState<number>(2048);
const [useAdvancedParams, setUseAdvancedParams] = useState<boolean>(false);
const chatEndRef = useRef<HTMLDivElement>(null);
@ -765,6 +770,8 @@ const ChatUI: React.FC<ChatUIProps> = ({
selectedMCPTools, // Pass the selected tools array
updateChatImageUI, // Pass the image callback
updateSearchResults, // Pass the search results callback
useAdvancedParams ? temperature : undefined, // Pass temperature if enabled
useAdvancedParams ? maxTokens : undefined, // Pass max_tokens if enabled
);
} else if (endpointType === EndpointType.IMAGE) {
// For image generation
@ -950,6 +957,19 @@ const ChatUI: React.FC<ChatUIProps> = ({
setShowCustomModelInput(value === "custom");
};
// Check if the selected model is a chat model
const isChatModel = () => {
if (!selectedModel || selectedModel === "custom") {
return false;
}
const model = modelInfo.find((m) => m.model_group === selectedModel);
if (!model) {
return false;
}
// Check if mode is explicitly "chat" or undefined (which defaults to chat per backend)
return !model.mode || model.mode === "chat";
};
const antIcon = <LoadingOutlined style={{ fontSize: 24 }} spin />;
return (
@ -1037,8 +1057,44 @@ const ChatUI: React.FC<ChatUIProps> = ({
</div>
<div>
<Text className="font-medium block mb-2 text-gray-700 flex items-center">
<RobotOutlined className="mr-2" /> Select Model
<Text className="font-medium block mb-2 text-gray-700 flex items-center justify-between">
<span className="flex items-center">
<RobotOutlined className="mr-2" /> Select Model
</span>
{isChatModel() ? (
<Popover
content={
<AdditionalModelSettings
temperature={temperature}
maxTokens={maxTokens}
useAdvancedParams={useAdvancedParams}
onTemperatureChange={setTemperature}
onMaxTokensChange={setMaxTokens}
onUseAdvancedParamsChange={setUseAdvancedParams}
/>
}
title="Model Settings"
trigger="click"
placement="right"
>
<Button
type="text"
size="small"
icon={<SettingOutlined />}
className="text-gray-500 hover:text-gray-700"
/>
</Popover>
) : (
<Tooltip title="Advanced parameters are only supported for chat models currently">
<Button
type="text"
size="small"
icon={<SettingOutlined />}
className="text-gray-300 cursor-not-allowed"
disabled
/>
</Tooltip>
)}
</Text>
<Select
value={selectedModel}

View file

@ -0,0 +1,118 @@
import { afterEach, beforeEach, describe, expect, it, vi } from "vitest";
import { makeOpenAIChatCompletionRequest } from "./chat_completion";
vi.mock("@/components/networking", () => ({
getProxyBaseUrl: vi.fn(() => "https://example.com"),
}));
// Mock the OpenAI client
const mockCreate = vi.fn();
const mockChatCompletions = {
create: mockCreate,
};
const mockChat = {
completions: mockChatCompletions,
};
const mockClient = {
chat: mockChat,
};
vi.mock("openai", () => ({
default: {
OpenAI: vi.fn(() => mockClient),
},
}));
describe("chat_completion", () => {
const mockUpdateUI = vi.fn();
const mockChatHistory = [{ role: "user", content: "Hello" }];
beforeEach(() => {
// Create a mock async iterator for streaming response
const mockChunks = [
{
choices: [{ delta: { content: "Hello" }, index: 0 }],
model: "gpt-4",
},
{
choices: [{ delta: { content: " there" }, index: 0 }],
model: "gpt-4",
},
{
choices: [{ delta: {}, index: 0 }],
model: "gpt-4",
usage: {
completion_tokens: 2,
prompt_tokens: 5,
total_tokens: 7,
},
},
];
async function* mockStream() {
for (const chunk of mockChunks) {
yield chunk;
}
}
mockCreate.mockResolvedValue(mockStream());
});
afterEach(() => {
vi.clearAllMocks();
});
it("should make a basic chat completion request", async () => {
await makeOpenAIChatCompletionRequest(mockChatHistory, mockUpdateUI, "gpt-4", "test-token");
expect(mockCreate).toHaveBeenCalledTimes(1);
expect(mockCreate).toHaveBeenCalledWith(
{
model: "gpt-4",
stream: true,
stream_options: {
include_usage: true,
},
messages: mockChatHistory,
},
{ signal: undefined },
);
expect(mockUpdateUI).toHaveBeenCalledWith("Hello", "gpt-4");
expect(mockUpdateUI).toHaveBeenCalledWith(" there", "gpt-4");
});
it("should include temperature and max_tokens when provided", async () => {
await makeOpenAIChatCompletionRequest(
mockChatHistory,
mockUpdateUI,
"gpt-4",
"test-token",
undefined, // tags
undefined, // signal
undefined, // onReasoningContent
undefined, // onTimingData
undefined, // onUsageData
undefined, // traceId
undefined, // vector_store_ids
undefined, // guardrails
undefined, // selectedMCPTools
undefined, // onImageGenerated
undefined, // onSearchResults
0.7, // temperature
100, // max_tokens
);
expect(mockCreate).toHaveBeenCalledTimes(1);
const callArgs = mockCreate.mock.calls[0][0];
expect(callArgs).toMatchObject({
model: "gpt-4",
stream: true,
stream_options: {
include_usage: true,
},
messages: mockChatHistory,
temperature: 0.7,
max_tokens: 100,
});
});
});

View file

@ -20,6 +20,8 @@ export async function makeOpenAIChatCompletionRequest(
selectedMCPTools?: string[],
onImageGenerated?: (imageUrl: string, model?: string) => void,
onSearchResults?: (searchResults: VectorStoreSearchResponse[]) => void,
temperature?: number,
max_tokens?: number,
) {
// base url should be the current base_url
const isLocal = process.env.NODE_ENV === "development";
@ -80,6 +82,8 @@ export async function makeOpenAIChatCompletionRequest(
...(vector_store_ids ? { vector_store_ids } : {}),
...(guardrails ? { guardrails } : {}),
...(tools ? { tools, tool_choice: "auto" } : {}),
...(temperature !== undefined ? { temperature } : {}),
...(max_tokens !== undefined ? { max_tokens } : {}),
},
{ signal },
);