Only show models based on selected endpoint (#16452)

This commit is contained in:
yuneng-jiang 2025-11-10 18:38:01 -08:00 • committed by GitHub
parent b8eda0ef55
commit af68763f5d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 127 additions and 51 deletions

View file

@ -1,4 +1,4 @@
import { fireEvent, render, screen, waitFor } from "@testing-library/react";
import { act, fireEvent, render, screen, waitFor } from "@testing-library/react";
import { beforeEach, describe, expect, it, vi } from "vitest";
import ChatUI from "./ChatUI";
import * as fetchModelsModule from "./llm_calls/fetch_models";
@ -118,18 +118,64 @@ describe("ChatUI", () => {
expect(getByText("Test Key")).toBeInTheDocument();
});
await waitFor(() => {
expect(getByText("Model 1")).toBeInTheDocument();
});
// Open the "Select Model" dropdown (AntD renders options in a portal)
const selectModelLabel = getByText("Select Model");
const modelSelect = selectModelLabel.parentElement?.querySelector(".ant-select-selector");
expect(modelSelect).toBeTruthy();
const selectComponent = container.querySelectorAll(".ant-select-selector")[1];
expect(selectComponent).toBeTruthy();
fireEvent.mouseDown(selectComponent!);
fireEvent.mouseDown(modelSelect!);
await waitFor(() => {
const model1Label = screen.getAllByText("Model 1");
expect(model1Label.length).toBeGreaterThan(0);
});
});
it("shows only chat-compatible models when chat endpoint is selected", async () => {
vi.mocked(fetchModelsModule.fetchAvailableModels).mockResolvedValueOnce([
{ model_group: "ChatModel", mode: "chat" },
{ model_group: "SpeechModel", mode: "audio_speech" },
{ model_group: "ImageModel", mode: "image_generation" },
{ model_group: "ResponsesModel", mode: "responses" },
]);
const { getByText, baseElement } = render(
<ChatUI
accessToken="1234567890"
token="1234567890"
userRole="user"
userID="1234567890"
disabledPersonalKeyCreation={false}
/>,
);
await waitFor(() => {
expect(getByText("Test Key")).toBeInTheDocument();
});
// Open endpoint selector and explicitly select /v1/chat/completions
const endpointTypeText = getByText("Endpoint Type");
const endpointSelect = endpointTypeText.parentElement?.querySelector(".ant-select-selector");
expect(endpointSelect).toBeTruthy();
act(() => {
fireEvent.mouseDown(endpointSelect!);
fireEvent.click(screen.getByText("/v1/chat/completions"));
});
// Open model selector
const selectModelLabel = getByText("Select Model");
const modelSelect = selectModelLabel.parentElement?.querySelector(".ant-select-selector");
expect(modelSelect).toBeTruthy();
act(() => {
fireEvent.mouseDown(modelSelect!);
});
await waitFor(() => {
// Chat-compatible: ChatModel should be visible
expect(screen.getAllByText("ChatModel").length).toBeGreaterThan(0);
expect(screen.queryByText("SpeechModel")).toBeNull();
expect(screen.queryByText("ImageModel")).toBeNull();
expect(screen.queryByText("ResponsesModel")).toBeNull();
});
});
});

View file

@ -48,7 +48,7 @@ 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 { EndpointType, getEndpointType } from "./mode_endpoint_mapping";
import ReasoningContent from "./ReasoningContent";
import ResponseMetrics, { TokenUsage } from "./ResponseMetrics";
import ResponsesImageRenderer from "./ResponsesImageRenderer";
@ -106,9 +106,7 @@ const ChatUI: React.FC<ChatUIProps> = ({ accessToken, token, userRole, userID, d
return [];
}
});
const [selectedModel, setSelectedModel] = useState<string | undefined>(
() => sessionStorage.getItem("selectedModel") || undefined,
);
const [selectedModel, setSelectedModel] = useState<string | undefined>(undefined);
const [showCustomModelInput, setShowCustomModelInput] = useState<boolean>(false);
const [modelInfo, setModelInfo] = useState<ModelGroup[]>([]);
const customModelTimeout = useRef<NodeJS.Timeout | null>(null);
@ -311,7 +309,7 @@ const ChatUI: React.FC<ChatUIProps> = ({ accessToken, token, userRole, userID, d
if (!uniqueModels.length) {
setSelectedModel(undefined);
} else if (!hasSelection) {
setSelectedModel(uniqueModels[0].model_group);
setSelectedModel(undefined);
}
} catch (error) {
console.error("Error fetching model info:", error);
@ -978,44 +976,6 @@ const ChatUI: React.FC<ChatUIProps> = ({ accessToken, token, userRole, userID, d
)}
</div>
<div>
<Text className="font-medium block mb-2 text-gray-700 flex items-center">
<RobotOutlined className="mr-2" /> Select Model
</Text>
<Select
value={selectedModel}
placeholder="Select a Model"
onChange={onModelChange}
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" },
]}
style={{ width: "100%" }}
showSearch={true}
className="rounded-md"
/>
{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);
}
customModelTimeout.current = setTimeout(() => {
setSelectedModel(value);
}, 500); // 500ms delay after typing stops
}}
/>
)}
</div>
<div>
<Text className="font-medium block mb-2 text-gray-700 flex items-center">
<ApiOutlined className="mr-2" /> Endpoint Type
@ -1024,6 +984,12 @@ const ChatUI: React.FC<ChatUIProps> = ({ accessToken, token, userRole, userID, d
endpointType={endpointType}
onEndpointChange={(value) => {
setEndpointType(value);
// Clear model selection when switching endpoint type
setSelectedModel(undefined);
setShowCustomModelInput(false);
try {
sessionStorage.removeItem("selectedModel");
} catch {}
}}
className="mb-4"
/>
@ -1057,6 +1023,68 @@ const ChatUI: React.FC<ChatUIProps> = ({ accessToken, token, userRole, userID, d
/>
</div>
<div>
<Text className="font-medium block mb-2 text-gray-700 flex items-center">
<RobotOutlined className="mr-2" /> Select Model
</Text>
<Select
value={selectedModel}
placeholder="Select a Model"
onChange={onModelChange}
options={[
...Array.from(
new Set(
modelInfo
.filter((option) => {
if (!option.mode) {
//If no mode, show all models
return true;
}
const optionEndpoint = getEndpointType(option.mode);
// Show chat models for responses/anthropic_messages endpoints as they are compatible
if (
endpointType === EndpointType.RESPONSES ||
endpointType === EndpointType.ANTHROPIC_MESSAGES
) {
return optionEndpoint === endpointType || optionEndpoint === EndpointType.CHAT;
}
// Show image models for image_edits endpoint as they are compatible
if (endpointType === EndpointType.IMAGE_EDITS) {
return optionEndpoint === endpointType || optionEndpoint === EndpointType.IMAGE;
}
return optionEndpoint === endpointType;
})
.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" },
]}
style={{ width: "100%" }}
showSearch={true}
className="rounded-md"
/>
{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);
}
customModelTimeout.current = setTimeout(() => {
setSelectedModel(value);
}, 500); // 500ms delay after typing stops
}}
/>
)}
</div>
<div>
<Text className="font-medium block mb-2 text-gray-700 flex items-center">
<TagsOutlined className="mr-2" /> Tags

View file

@ -10,6 +10,7 @@ export enum ModelMode {
RESPONSES = "responses",
IMAGE_EDITS = "image_edits",
ANTHROPIC_MESSAGES = "anthropic_messages",
EMBEDDING = "embedding",
// add additional modes as needed
}
@ -37,6 +38,7 @@ export const litellmModeMapping: Record<ModelMode, EndpointType> = {
[ModelMode.ANTHROPIC_MESSAGES]: EndpointType.ANTHROPIC_MESSAGES,
[ModelMode.AUDIO_SPEECH]: EndpointType.SPEECH,
[ModelMode.AUDIO_TRANSCRIPTION]: EndpointType.TRANSCRIPTION,
[ModelMode.EMBEDDING]: EndpointType.EMBEDDINGS,
};
export const getEndpointType = (mode: string): EndpointType => {