mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Add additional model settings to chat models in test key (#16793)
This commit is contained in:
parent
fe05e33723
commit
d0e806d1b3
6 changed files with 375 additions and 6 deletions
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
|
|
@ -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;
|
||||
|
|
@ -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!);
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
@ -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 },
|
||||
);
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue