mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
Only show models based on selected endpoint (#16452)
This commit is contained in:
parent
b8eda0ef55
commit
af68763f5d
3 changed files with 127 additions and 51 deletions
|
|
@ -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();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 => {
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue