Merge remote-tracking branch 'berri/litellm_playground_shadcn' into litellm_playground_shadcn

This commit is contained in:
mubashir1osmani 2026-08-13 12:44:01 -07:00
commit 8fd97803ce
8 changed files with 521 additions and 169 deletions

View file

@ -0,0 +1,122 @@
import { fireEvent, render, screen } from "@testing-library/react";
import { describe, expect, it, vi } from "vitest";
import { ChatComposer } from "./ChatComposer";
const renderComposer = (props: Partial<React.ComponentProps<typeof ChatComposer>> = {}) =>
render(
<ChatComposer
value=""
onChange={vi.fn()}
onSubmit={vi.fn()}
placeholder="Send a message"
{...props}
/>,
);
const addonOf = (container: HTMLElement) =>
container.querySelector<HTMLElement>("[data-slot=input-group-addon]") as HTMLElement;
describe("ChatComposer", () => {
it("should submit on Enter", () => {
const onSubmit = vi.fn();
renderComposer({ onSubmit });
fireEvent.keyDown(screen.getByTestId("chat-composer-input"), { key: "Enter" });
expect(onSubmit).toHaveBeenCalledTimes(1);
});
it("should not submit on Shift+Enter", () => {
const onSubmit = vi.fn();
renderComposer({ onSubmit });
fireEvent.keyDown(screen.getByTestId("chat-composer-input"), { key: "Enter", shiftKey: true });
expect(onSubmit).not.toHaveBeenCalled();
});
it("should not submit while an IME composition is active", () => {
const onSubmit = vi.fn();
renderComposer({ onSubmit });
fireEvent.keyDown(screen.getByTestId("chat-composer-input"), { key: "Enter", isComposing: true });
expect(onSubmit).not.toHaveBeenCalled();
});
it("should not submit on Enter or click when submitDisabled", () => {
const onSubmit = vi.fn();
renderComposer({ onSubmit, submitDisabled: true });
fireEvent.keyDown(screen.getByTestId("chat-composer-input"), { key: "Enter" });
fireEvent.click(screen.getByTestId("chat-send-button"));
expect(onSubmit).not.toHaveBeenCalled();
});
it("should submit when the send button is clicked", () => {
const onSubmit = vi.fn();
renderComposer({ onSubmit });
fireEvent.click(screen.getByTestId("chat-send-button"));
expect(onSubmit).toHaveBeenCalledTimes(1);
});
it("should swap send for a stop button that cancels while loading", () => {
const onSubmit = vi.fn();
const onCancel = vi.fn();
renderComposer({ onSubmit, onCancel, isLoading: true });
expect(screen.queryByTestId("chat-send-button")).not.toBeInTheDocument();
fireEvent.click(screen.getByTestId("chat-stop-button"));
expect(onCancel).toHaveBeenCalledTimes(1);
expect(onSubmit).not.toHaveBeenCalled();
});
it("should render suggestions only when asked and report the chosen one", () => {
const onSuggestionSelect = vi.fn();
const { rerender } = renderComposer({ suggestions: ["Summarize this"], onSuggestionSelect });
expect(screen.queryByTestId("chat-suggested-actions")).not.toBeInTheDocument();
rerender(
<ChatComposer
value=""
onChange={vi.fn()}
onSubmit={vi.fn()}
placeholder="Send a message"
suggestions={["Summarize this"]}
showSuggestions
onSuggestionSelect={onSuggestionSelect}
/>,
);
fireEvent.click(screen.getByText("Summarize this"));
expect(onSuggestionSelect).toHaveBeenCalledWith("Summarize this");
});
it("should not nest a form inside the composer when body renders one", () => {
const { container } = renderComposer({
body: (
<form data-testid="body-form">
<input aria-label="tool argument" />
</form>
),
});
expect(container.querySelectorAll("form")).toHaveLength(1);
expect(screen.getByTestId("body-form")).toBeInTheDocument();
});
it("should focus the message box, not a tool input, when the toolbar gap is clicked", () => {
const { container } = renderComposer({
tools: <input type="file" aria-label="Attach file" />,
});
fireEvent.click(addonOf(container));
expect(document.activeElement).toBe(screen.getByTestId("chat-composer-input"));
});
});

View file

@ -0,0 +1,173 @@
import React from "react";
import { ArrowUp, Code2, Square } from "lucide-react";
import { Button } from "@/components/ui/button";
import { InputGroup, InputGroupAddon, InputGroupButton, InputGroupTextarea } from "@/components/ui/input-group";
import { Tooltip, TooltipContent, TooltipTrigger } from "@/components/ui/tooltip";
import { cn } from "@/lib/cva.config";
interface ChatComposerProps {
value: string;
onChange: (value: string) => void;
onSubmit: () => void;
onCancel?: () => void;
placeholder: string;
disabled?: boolean;
isLoading?: boolean;
submitDisabled?: boolean;
tools?: React.ReactNode;
body?: React.ReactNode;
suggestions?: string[];
showSuggestions?: boolean;
onSuggestionSelect?: (suggestion: string) => void;
className?: string;
}
export function ChatComposer({
value,
onChange,
onSubmit,
onCancel,
placeholder,
disabled = false,
isLoading = false,
submitDisabled = false,
tools,
body,
suggestions = [],
showSuggestions = false,
onSuggestionSelect,
className,
}: ChatComposerProps) {
const submitIfAllowed = () => {
if (!submitDisabled && !isLoading) {
onSubmit();
}
};
const handleKeyDown = (event: React.KeyboardEvent<HTMLTextAreaElement>) => {
if (event.key === "Enter" && !event.shiftKey && !event.nativeEvent.isComposing) {
event.preventDefault();
submitIfAllowed();
}
};
return (
<div className={cn("relative flex w-full flex-col gap-3", className)}>
{showSuggestions && suggestions.length > 0 && (
<div
className="flex w-full gap-2 overflow-x-auto pb-1 sm:grid sm:grid-cols-2 sm:overflow-visible"
data-testid="chat-suggested-actions"
>
{suggestions.map((suggestion) => (
<button
key={suggestion}
type="button"
className="min-w-[200px] shrink-0 rounded-xl border border-border/50 bg-card/30 px-4 py-3 text-left text-[12px] leading-relaxed text-muted-foreground transition-all duration-200 hover:-translate-y-0.5 hover:bg-card/60 hover:text-foreground sm:min-w-0 sm:whitespace-normal sm:p-4 sm:text-[13px]"
onClick={() => onSuggestionSelect?.(suggestion)}
>
{suggestion}
</button>
))}
</div>
)}
<div className="w-full">
<InputGroup
className={cn(
"h-auto min-h-[7.5rem] flex-col overflow-hidden rounded-2xl border border-border bg-card",
"shadow-[0_1px_2px_rgba(0,0,0,0.06),0_8px_24px_rgba(0,0,0,0.08)] ring-1 ring-black/5",
"transition-[box-shadow,border-color,ring] duration-200",
"has-[[data-slot=input-group-control]:focus-visible]:border-ring",
"has-[[data-slot=input-group-control]:focus-visible]:shadow-[0_2px_8px_rgba(0,0,0,0.08),0_12px_32px_rgba(0,0,0,0.12)]",
"has-[[data-slot=input-group-control]:focus-visible]:ring-2 has-[[data-slot=input-group-control]:focus-visible]:ring-ring/40",
)}
>
{body ? (
<div className="max-h-48 min-h-24 w-full overflow-y-auto px-3 pt-3">{body}</div>
) : (
<InputGroupTextarea
data-testid="chat-composer-input"
value={value}
disabled={disabled}
placeholder={placeholder}
rows={1}
className="min-h-24 max-h-48 resize-none overflow-y-auto border-0 bg-transparent px-4 pt-3.5 pb-1.5 text-[13px] leading-relaxed shadow-none placeholder:text-muted-foreground/50 focus-visible:ring-0 [field-sizing:content]"
onChange={(event) => onChange(event.target.value)}
onKeyDown={handleKeyDown}
/>
)}
<InputGroupAddon align="block-end" className="justify-between gap-2 px-3 pb-3 pt-1">
<div className="flex min-w-0 items-center gap-1">{tools}</div>
{isLoading && onCancel ? (
<InputGroupButton
type="button"
size="icon-sm"
aria-label="Stop request"
data-testid="chat-stop-button"
className="size-8 rounded-xl bg-foreground text-background hover:bg-foreground/90"
onClick={onCancel}
>
<Square className="size-3.5 fill-current" />
</InputGroupButton>
) : (
<InputGroupButton
type="button"
size="icon-sm"
aria-label="Send message"
data-testid="chat-send-button"
disabled={submitDisabled || isLoading}
onClick={submitIfAllowed}
className={cn(
"size-8 rounded-xl transition-all duration-200",
!submitDisabled && !isLoading
? "bg-foreground text-background hover:opacity-90 active:scale-95"
: "cursor-not-allowed bg-muted text-muted-foreground/40",
)}
>
<ArrowUp className="size-4" />
</InputGroupButton>
)}
</InputGroupAddon>
</InputGroup>
</div>
</div>
);
}
interface CodeInterpreterToggleProps {
enabled: boolean;
onToggle: () => void;
}
export function CodeInterpreterToggle({ enabled, onToggle }: CodeInterpreterToggleProps) {
return (
<Tooltip>
<TooltipTrigger
render={
<Button
type="button"
variant="ghost"
size="icon-sm"
className={cn(
"size-8 rounded-lg border border-border/40",
enabled
? "border-blue-200 bg-blue-50 text-blue-600 hover:bg-blue-100"
: "text-muted-foreground hover:text-foreground",
)}
aria-label={enabled ? "Code Interpreter enabled (click to disable)" : "Enable Code Interpreter"}
onClick={onToggle}
/>
}
>
<Code2 className="size-4" />
</TooltipTrigger>
<TooltipContent>
{enabled ? "Code Interpreter enabled (click to disable)" : "Enable Code Interpreter"}
</TooltipContent>
</Tooltip>
);
}
export default ChatComposer;

View file

@ -118,6 +118,74 @@ describe("ChatUI", () => {
});
});
it("shows only endpoint-compatible models when chat endpoint is selected", async () => {
(fetchModelsModule.fetchAvailableModels as ReturnType<typeof vi.fn>).mockResolvedValueOnce([
{ model_group: "ChatModel", mode: "chat" },
{ model_group: "SpeechModel", mode: "audio_speech" },
{ model_group: "ImageModel", mode: "image_generation" },
{ model_group: "ResponsesModel", mode: "responses" },
{ model_group: "RealtimeModel", mode: "realtime" },
{ model_group: "NoModeModel" },
]);
render(
<ChatUI
accessToken="1234567890"
token="1234567890"
userRole="user"
userID="1234567890"
disabledPersonalKeyCreation={false}
/>,
);
await waitFor(() => {
expect(screen.getByText("Test Key")).toBeInTheDocument();
});
await selectComboboxOption("Select an endpoint", "/v1/chat/completions");
await openComboboxByPlaceholder("Select a Model");
await waitFor(() => {
expect(screen.getAllByText("ChatModel").length).toBeGreaterThan(0);
expect(screen.getAllByText("NoModeModel").length).toBeGreaterThan(0);
expect(screen.queryByText("SpeechModel")).toBeNull();
expect(screen.queryByText("ImageModel")).toBeNull();
expect(screen.queryByText("ResponsesModel")).toBeNull();
expect(screen.queryByText("RealtimeModel")).toBeNull();
});
});
it("shows only realtime models when realtime endpoint is selected", async () => {
(fetchModelsModule.fetchAvailableModels as ReturnType<typeof vi.fn>).mockResolvedValueOnce([
{ model_group: "ChatModel", mode: "chat" },
{ model_group: "RealtimeModel", mode: "realtime" },
{ model_group: "NoModeModel" },
]);
render(
<ChatUI
accessToken="1234567890"
token="1234567890"
userRole="user"
userID="1234567890"
disabledPersonalKeyCreation={false}
/>,
);
await waitFor(() => {
expect(screen.getByText("Test Key")).toBeInTheDocument();
});
await selectComboboxOption("Select an endpoint", "/v1/realtime");
await openComboboxByPlaceholder("Select a Model");
await waitFor(() => {
expect(screen.getAllByText("RealtimeModel").length).toBeGreaterThan(0);
expect(screen.getAllByText("NoModeModel").length).toBeGreaterThan(0);
expect(screen.queryByText("ChatModel")).toBeNull();
});
});
it("should show 'Enter custom model' option in model selector", async () => {
render(
<ChatUI
@ -427,8 +495,9 @@ describe("ChatUI", () => {
expect(screen.getByPlaceholderText("Select an endpoint")).toHaveValue("/v1/responses");
});
it("should still switch endpoint when the picked model cannot be served by it", async () => {
it("should not offer a model the selected endpoint cannot serve", async () => {
(fetchModelsModule.fetchAvailableModels as ReturnType<typeof vi.fn>).mockResolvedValueOnce([
{ model_group: "ChatModel", mode: "chat" },
{ model_group: "SpeechModel", mode: "audio_speech" },
]);
@ -447,9 +516,12 @@ describe("ChatUI", () => {
});
await selectComboboxOption("Select an endpoint", "/v1/responses");
await selectComboboxOption("Select a Model", "SpeechModel");
await openComboboxByPlaceholder("Select a Model");
expect(screen.getByPlaceholderText("Select an endpoint")).toHaveValue("/v1/audio/speech");
await waitFor(() => {
expect(screen.getAllByText("ChatModel").length).toBeGreaterThan(0);
});
expect(screen.queryByText("SpeechModel")).toBeNull();
});
it("should attach an audio file dropped on the transcription upload area", async () => {

View file

@ -1,7 +1,6 @@
"use client";
import {
ArrowUp,
Bot,
Code2,
Database,
@ -48,11 +47,13 @@ import { makeOpenAIResponsesRequest } from "@/components/llm_calls/responses_api
import { makeInteractionsRequest } from "../../llm_calls/interactions_api";
import AdditionalModelSettings from "./AdditionalModelSettings";
import { OPEN_AI_VOICE_SELECT_OPTIONS, OpenAIVoice } from "./chatConstants";
import ChatComposer, { CodeInterpreterToggle } from "./ChatComposer";
import ChatImageUpload from "./ChatImageUpload";
import { createChatDisplayMessage, createChatMultimodalMessage } from "./ChatImageUtils";
import CodeInterpreterTool from "./CodeInterpreterTool";
import { generateCodeSnippet } from "@/components/chat_ui/CodeSnippets";
import EndpointSelector from "./EndpointSelector";
import { filterModelsForEndpoint, isModelCompatibleWithEndpoint } from "./EndpointUtils";
import FilePreviewCard from "./FilePreviewCard";
import ChatMessageBubble from "./ChatMessageBubble";
import MCPEventsDisplay from "@/components/chat_ui/MCPEventsDisplay";
@ -72,7 +73,6 @@ import { Dialog, DialogContent, DialogFooter, DialogHeader, DialogTitle } from "
import { Input } from "@/components/ui/input";
import { Popover, PopoverContent, PopoverTrigger } from "@/components/ui/popover";
import { Select as ShadcnSelect, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select";
import { Textarea } from "@/components/ui/textarea";
import { Tooltip, TooltipContent, TooltipTrigger } from "@/components/ui/tooltip";
import { useDebouncedCallback } from "@tanstack/react-pacer/debouncer";
import {
@ -507,14 +507,6 @@ const ChatUI: React.FC<ChatUIProps> = ({
}
}, [chatHistory]);
const handleKeyDown = (event: React.KeyboardEvent<HTMLTextAreaElement>) => {
if (event.key === "Enter" && !event.shiftKey) {
event.preventDefault(); // Prevent default to avoid newline
handleSendMessage();
}
// If Shift+Enter is pressed, the default behavior (inserting a newline) will occur
};
const handleCancelRequest = () => {
if (abortControllerRef.current) {
abortControllerRef.current.abort();
@ -1150,27 +1142,12 @@ const ChatUI: React.FC<ChatUIProps> = ({
NotificationsManager.success("Chat history cleared.");
};
const currentEndpointServes = (mode: string): boolean => {
const modelEndpoint = getEndpointType(mode);
if (
endpointType === EndpointType.RESPONSES ||
endpointType === EndpointType.ANTHROPIC_MESSAGES ||
endpointType === EndpointType.INTERACTIONS
) {
return modelEndpoint === endpointType || modelEndpoint === EndpointType.CHAT;
}
if (endpointType === EndpointType.IMAGE_EDITS) {
return modelEndpoint === endpointType || modelEndpoint === EndpointType.IMAGE;
}
return modelEndpoint === endpointType;
};
const onModelChange = (value: string) => {
setSelectedModel(value);
setShowCustomModelInput(value === "custom");
const model = modelInfo.find((option) => option.model_group === value);
if (model?.mode && !currentEndpointServes(model.mode)) {
if (model?.mode && !isModelCompatibleWithEndpoint(model, endpointType as EndpointType)) {
setEndpointType(getEndpointType(model.mode));
}
};
@ -1189,11 +1166,17 @@ const ChatUI: React.FC<ChatUIProps> = ({
};
const supportsStreamingToggle = endpointType === EndpointType.CHAT || endpointType === EndpointType.RESPONSES;
const modelsForEndpoint = useMemo(
() => filterModelsForEndpoint(modelInfo, endpointType as EndpointType),
[modelInfo, endpointType],
);
let modelEmptyText = "No models available for this key";
if (modelLoadError) {
modelEmptyText = "Unable to load models for this key";
} else if (apiKeySource === "custom" && !apiKey.trim()) {
modelEmptyText = "Enter a Virtual Key to load models";
} else if (modelInfo.length > 0 && modelsForEndpoint.length === 0) {
modelEmptyText = "No models available for this endpoint";
}
const inputPlaceholder =
@ -1423,7 +1406,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
onValueChange={onModelChange}
options={[
{ value: "custom", label: "Enter custom model" },
...modelInfo.map((model) => ({
...modelsForEndpoint.map((model) => ({
value: model.model_group,
label: model.model_group,
sublabel: model.mode ? `Mode: ${model.mode}` : undefined,
@ -2005,27 +1988,24 @@ const ChatUI: React.FC<ChatUIProps> = ({
</div>
)}
{chatHistory.length === 0 && !isLoading && endpointType !== EndpointType.MCP && (
<div className="mb-3 flex items-center gap-2 overflow-x-auto">
{(endpointType === EndpointType.A2A_AGENTS
<ChatComposer
value={inputMessage}
onChange={setInputMessage}
onSubmit={handleSendMessage}
onCancel={handleCancelRequest}
placeholder={inputPlaceholder}
disabled={isLoading}
isLoading={isLoading}
submitDisabled={sendDisabled}
showSuggestions={chatHistory.length === 0 && !isLoading && endpointType !== EndpointType.MCP}
suggestions={
endpointType === EndpointType.A2A_AGENTS
? ["What can you help me with?", "Tell me about yourself", "What tasks can you perform?"]
: ["Write me a poem", "Explain quantum computing", "Draft a polite email requesting a meeting"]
).map((prompt) => (
<button
key={prompt}
type="button"
className="shrink-0 cursor-pointer rounded-full border border-gray-200 px-3 py-1 text-xs font-medium text-gray-600 transition-colors hover:border-blue-300 hover:bg-blue-50 hover:text-blue-600"
onClick={() => setInputMessage(prompt)} // lgtm[js/xss-through-dom]
>
{prompt}
</button>
))}
</div>
)}
<div className="flex items-center gap-2">
<div className="flex min-h-[44px] flex-1 items-center rounded-xl border border-gray-300 bg-white px-3 py-1">
<div className="mr-2 flex shrink-0 items-center gap-1">
}
onSuggestionSelect={setInputMessage}
tools={
<>
{endpointType === EndpointType.RESPONSES && !responsesUploadedImage && (
<ResponsesImageUpload
responsesUploadedImage={responsesUploadedImage}
@ -2043,47 +2023,24 @@ const ChatUI: React.FC<ChatUIProps> = ({
/>
)}
{endpointType === EndpointType.RESPONSES && (
<Tooltip>
<TooltipTrigger
render={
<button
type="button"
className={`rounded-md p-1.5 transition-colors ${
codeInterpreter.enabled
? "bg-blue-100 text-blue-600"
: "text-gray-400 hover:bg-gray-100 hover:text-gray-600"
}`}
aria-label={
codeInterpreter.enabled
? "Code Interpreter enabled (click to disable)"
: "Enable Code Interpreter"
}
onClick={() => {
codeInterpreter.toggle();
if (!codeInterpreter.enabled) {
NotificationsManager.success("Code Interpreter enabled!");
}
}}
/>
<CodeInterpreterToggle
enabled={codeInterpreter.enabled}
onToggle={() => {
codeInterpreter.toggle();
if (!codeInterpreter.enabled) {
NotificationsManager.success("Code Interpreter enabled!");
}
>
<Code2 className="size-4" />
</TooltipTrigger>
<TooltipContent>
{codeInterpreter.enabled
? "Code Interpreter enabled (click to disable)"
: "Enable Code Interpreter"}
</TooltipContent>
</Tooltip>
}}
/>
)}
</div>
{endpointType === EndpointType.MCP &&
</>
}
body={
endpointType === EndpointType.MCP &&
selectedMCPServers.length === 1 &&
selectedMCPServers[0] !== "__all__" &&
selectedMCPDirectTool ? (
<div className="max-h-48 min-h-[44px] flex-1 overflow-y-auto rounded-lg border border-gray-200 bg-gray-50/50 p-2">
{(() => {
selectedMCPDirectTool
? (() => {
const rawSel = selectedMCPServers[0];
let toolPool: MCPTool[] = [];
if (rawSel.startsWith("toolset:")) {
@ -2106,45 +2063,10 @@ const ChatUI: React.FC<ChatUIProps> = ({
Loading tool schema...
</div>
);
})()}
</div>
) : (
<Textarea
value={inputMessage}
onChange={(event) => setInputMessage(event.target.value)}
onKeyDown={handleKeyDown}
placeholder={inputPlaceholder}
disabled={isLoading}
rows={1}
className="min-h-0 flex-1 resize-none border-0 bg-transparent px-0 py-1 text-sm shadow-none focus-visible:ring-0"
/>
)}
<Button
type="button"
size="icon-sm"
onClick={handleSendMessage}
disabled={sendDisabled}
className="ml-2 size-8 shrink-0 rounded-full bg-blue-600 text-white hover:bg-blue-700 disabled:bg-gray-300 disabled:text-gray-500"
aria-label="Send message"
>
<ArrowUp className="size-3.5" />
</Button>
</div>
{isLoading && (
<Button
type="button"
variant="outline"
size="sm"
className="border-red-200 bg-red-50 text-red-600 hover:bg-red-100"
onClick={handleCancelRequest}
>
<Trash2 className="size-3.5" />
Cancel
</Button>
)}
</div>
})()
: undefined
}
/>
</div>
</>
)}

View file

@ -1,37 +1,16 @@
import { beforeEach, describe, expect, it, vi } from "vitest";
import type { ModelGroup } from "@/components/llm_calls/fetch_models";
import { determineEndpointType } from "./EndpointUtils";
import { determineEndpointType, filterModelsForEndpoint, isModelCompatibleWithEndpoint } from "./EndpointUtils";
import { EndpointType } from "@/components/chat_ui/mode_endpoint_mapping";
// Mock the getEndpointType function
vi.mock("@/components/chat_ui/mode_endpoint_mapping", () => ({
EndpointType: {
IMAGE: "image",
VIDEO: "video",
CHAT: "chat",
RESPONSES: "responses",
IMAGE_EDITS: "image_edits",
ANTHROPIC_MESSAGES: "anthropic_messages",
EMBEDDINGS: "embeddings",
SPEECH: "speech",
TRANSCRIPTION: "transcription",
A2A_AGENTS: "a2a_agents",
},
getEndpointType: vi.fn(),
ModelMode: {
AUDIO_SPEECH: "audio_speech",
AUDIO_TRANSCRIPTION: "audio_transcription",
IMAGE_GENERATION: "image_generation",
VIDEO_GENERATION: "video_generation",
CHAT: "chat",
RESPONSES: "responses",
IMAGE_EDITS: "image_edits",
ANTHROPIC_MESSAGES: "anthropic_messages",
EMBEDDING: "embedding",
},
}));
vi.mock("@/components/chat_ui/mode_endpoint_mapping", async (importOriginal) => {
const actual = await importOriginal<typeof import("@/components/chat_ui/mode_endpoint_mapping")>();
return {
...actual,
getEndpointType: vi.fn(actual.getEndpointType),
};
});
// Import the mocked function
import { getEndpointType } from "@/components/chat_ui/mode_endpoint_mapping";
describe("determineEndpointType", () => {
@ -210,10 +189,72 @@ describe("determineEndpointType", () => {
vi.mocked(getEndpointType).mockReturnValue(EndpointType.CHAT);
// Test with different case - should not match
const result = determineEndpointType("gpt-3.5-turbo", mockModelInfo);
expect(getEndpointType).not.toHaveBeenCalled();
expect(result).toBe(EndpointType.CHAT);
});
});
describe("isModelCompatibleWithEndpoint / filterModelsForEndpoint", () => {
beforeEach(async () => {
const actual =
await vi.importActual<typeof import("@/components/chat_ui/mode_endpoint_mapping")>(
"@/components/chat_ui/mode_endpoint_mapping",
);
vi.mocked(getEndpointType).mockImplementation(actual.getEndpointType);
});
it("keeps models with no mode for every endpoint", () => {
const model: ModelGroup = { model_group: "custom-proxy-model" };
expect(isModelCompatibleWithEndpoint(model, EndpointType.CHAT)).toBe(true);
expect(isModelCompatibleWithEndpoint(model, EndpointType.REALTIME)).toBe(true);
expect(isModelCompatibleWithEndpoint(model, EndpointType.SPEECH)).toBe(true);
});
it("keeps chat models for responses, anthropic messages, and interactions", () => {
const chatModel: ModelGroup = { model_group: "gpt-4o", mode: "chat" };
expect(isModelCompatibleWithEndpoint(chatModel, EndpointType.RESPONSES)).toBe(true);
expect(isModelCompatibleWithEndpoint(chatModel, EndpointType.ANTHROPIC_MESSAGES)).toBe(true);
expect(isModelCompatibleWithEndpoint(chatModel, EndpointType.INTERACTIONS)).toBe(true);
expect(isModelCompatibleWithEndpoint(chatModel, EndpointType.SPEECH)).toBe(false);
});
it("keeps image models for image_edits", () => {
const imageModel: ModelGroup = { model_group: "dall-e-3", mode: "image_generation" };
expect(isModelCompatibleWithEndpoint(imageModel, EndpointType.IMAGE_EDITS)).toBe(true);
expect(isModelCompatibleWithEndpoint(imageModel, EndpointType.IMAGE)).toBe(true);
expect(isModelCompatibleWithEndpoint(imageModel, EndpointType.CHAT)).toBe(false);
});
it("keeps only realtime models for the realtime endpoint", () => {
const models: ModelGroup[] = [
{ model_group: "gpt-4o", mode: "chat" },
{ model_group: "gpt-realtime", mode: "realtime" },
{ model_group: "no-mode" },
];
expect(filterModelsForEndpoint(models, EndpointType.REALTIME).map((m) => m.model_group)).toEqual([
"gpt-realtime",
"no-mode",
]);
});
it("excludes unknown modes from conversational endpoints", () => {
const batchModel: ModelGroup = { model_group: "batch-job", mode: "batch" };
const rerankModel: ModelGroup = { model_group: "reranker", mode: "rerank" };
expect(isModelCompatibleWithEndpoint(batchModel, EndpointType.CHAT)).toBe(false);
expect(isModelCompatibleWithEndpoint(rerankModel, EndpointType.RESPONSES)).toBe(false);
expect(isModelCompatibleWithEndpoint(batchModel, EndpointType.REALTIME)).toBe(false);
});
it("keeps image-edit models for the image-edits endpoint using the mode the backend sends", () => {
const imageEditModel: ModelGroup = { model_group: "gpt-image-1", mode: "image_edit" };
const imageModel: ModelGroup = { model_group: "dall-e-3", mode: "image_generation" };
expect(isModelCompatibleWithEndpoint(imageEditModel, EndpointType.IMAGE_EDITS)).toBe(true);
expect(filterModelsForEndpoint([imageEditModel, imageModel], EndpointType.IMAGE_EDITS).map((m) => m.model_group)).toEqual(
["gpt-image-1", "dall-e-3"],
);
});
});

View file

@ -1,22 +1,43 @@
import { ModelGroup } from "@/components/llm_calls/fetch_models";
import { EndpointType, getEndpointType } from "@/components/chat_ui/mode_endpoint_mapping";
import { EndpointType, getEndpointType, ModelMode } from "@/components/chat_ui/mode_endpoint_mapping";
const KNOWN_MODEL_MODES = new Set<string>(Object.values(ModelMode));
/**
* Determines the appropriate endpoint type based on the selected model
*
* @param selectedModel - The model identifier string
* @param modelInfo - Array of model information
* @returns The appropriate endpoint type
*/
export const determineEndpointType = (selectedModel: string, modelInfo: ModelGroup[]): EndpointType => {
// Find the model information for the selected model
const selectedModelInfo = modelInfo.find((option) => option.model_group === selectedModel);
// If model info is found and it has a mode, determine the endpoint type
if (selectedModelInfo?.mode) {
return getEndpointType(selectedModelInfo.mode);
}
// Default to chat endpoint if no match is found
return EndpointType.CHAT;
};
export const isModelCompatibleWithEndpoint = (model: ModelGroup, endpointType: EndpointType): boolean => {
if (!model.mode) {
return true;
}
if (!KNOWN_MODEL_MODES.has(model.mode)) {
return false;
}
const optionEndpoint = getEndpointType(model.mode);
if (
endpointType === EndpointType.RESPONSES ||
endpointType === EndpointType.ANTHROPIC_MESSAGES ||
endpointType === EndpointType.INTERACTIONS
) {
return optionEndpoint === endpointType || optionEndpoint === EndpointType.CHAT;
}
if (endpointType === EndpointType.IMAGE_EDITS) {
return optionEndpoint === endpointType || optionEndpoint === EndpointType.IMAGE;
}
return optionEndpoint === endpointType;
};
export const filterModelsForEndpoint = (models: ModelGroup[], endpointType: EndpointType): ModelGroup[] =>
models.filter((model) => isModelCompatibleWithEndpoint(model, endpointType));

View file

@ -8,10 +8,10 @@ export enum ModelMode {
VIDEO_GENERATION = "video_generation",
CHAT = "chat",
RESPONSES = "responses",
IMAGE_EDITS = "image_edits",
IMAGE_EDITS = "image_edit",
ANTHROPIC_MESSAGES = "anthropic_messages",
EMBEDDING = "embedding",
// add additional modes as needed
REALTIME = "realtime",
}
// Define an enum for the endpoint types your UI calls
@ -42,6 +42,7 @@ export const litellmModeMapping: Record<ModelMode, EndpointType> = {
[ModelMode.AUDIO_SPEECH]: EndpointType.SPEECH,
[ModelMode.AUDIO_TRANSCRIPTION]: EndpointType.TRANSCRIPTION,
[ModelMode.EMBEDDING]: EndpointType.EMBEDDINGS,
[ModelMode.REALTIME]: EndpointType.REALTIME,
};
export const getEndpointType = (mode: string): EndpointType => {

View file

@ -53,7 +53,7 @@ function InputGroupAddon({
if ((e.target as HTMLElement).closest("button")) {
return;
}
e.currentTarget.parentElement?.querySelector("input")?.focus();
e.currentTarget.parentElement?.querySelector<HTMLElement>("[data-slot=input-group-control]")?.focus();
}}
{...props}
/>