From 44d6abb36ce30e6e906f1a777578943ac39236db Mon Sep 17 00:00:00 2001 From: Shine1i Date: Wed, 25 Feb 2026 10:55:53 +0100 Subject: [PATCH] feat: chat model switching toast and add image detection logic --- studio/frontend/src/app/provider.tsx | 2 +- .../frontend/src/features/chat/chat-page.tsx | 104 ++++++++++++------ .../src/features/chat/chat-settings-sheet.tsx | 2 +- .../chat/hooks/use-chat-model-runtime.ts | 17 ++- .../src/features/chat/types/runtime.ts | 2 +- 5 files changed, 86 insertions(+), 41 deletions(-) diff --git a/studio/frontend/src/app/provider.tsx b/studio/frontend/src/app/provider.tsx index 5c2b93bcc9..da88e0cfa3 100644 --- a/studio/frontend/src/app/provider.tsx +++ b/studio/frontend/src/app/provider.tsx @@ -10,7 +10,7 @@ export function AppProvider({ children }: AppProviderProps) { return ( {children} - + ); } diff --git a/studio/frontend/src/features/chat/chat-page.tsx b/studio/frontend/src/features/chat/chat-page.tsx index 4557639bcb..1ef3c6bec5 100644 --- a/studio/frontend/src/features/chat/chat-page.tsx +++ b/studio/frontend/src/features/chat/chat-page.tsx @@ -33,6 +33,7 @@ import { useRef, useState, } from "react"; +import { toast } from "sonner"; import { GuidedTour, useGuidedTourController } from "@/features/tour"; import { ChatSettingsPanel } from "./chat-settings-sheet"; import { db } from "./db"; @@ -50,7 +51,7 @@ import { SharedComposer, } from "./shared-composer"; import { ThreadSidebar } from "./thread-sidebar"; -import type { ChatView } from "./types"; +import type { ChatView, MessageRecord } from "./types"; import { buildChatTourSteps } from "./tour"; type LoraCandidate = { @@ -90,6 +91,40 @@ function pickBestLoraForBase( return partial ?? sorted[0]; } +function messageHasImage(message: MessageRecord): boolean { + const contentParts = Array.isArray(message.content) ? message.content : []; + if (contentParts.some((part) => part.type === "image")) { + return true; + } + const attachments = Array.isArray(message.attachments) ? message.attachments : []; + for (const attachment of attachments) { + const parts = Array.isArray(attachment.content) ? attachment.content : []; + for (const part of parts as Array<{ type?: string }>) { + if (part?.type === "image") { + return true; + } + } + } + return false; +} + +async function resolveActiveSingleThreadId(view: ChatView): Promise { + if (view.mode !== "single") { + return undefined; + } + if (view.threadId) { + return view.threadId; + } + + // New-thread flow keeps threadId undefined in local view state. + // Fall back to most recent regular base thread. + const candidates = await db.threads.where("modelType").equals("base").toArray(); + const latest = candidates + .filter((thread) => !thread.archived && !thread.pairId) + .sort((a, b) => b.createdAt - a.createdAt)[0]; + return latest?.id; +} + const SingleContent = memo(function SingleContent({ threadId, newThreadNonce, @@ -304,15 +339,42 @@ export function ChatPage(): ReactElement { const currentCheckpoint = useChatRuntimeStore.getState().params.checkpoint; if (!value || value === currentCheckpoint) return; - setView({ mode: "single", newThreadNonce: crypto.randomUUID() }); void (async () => { - if (currentCheckpoint) { - await ejectModel(); + let switchNote: string | undefined; + const activeThreadId = await resolveActiveSingleThreadId(view); + if (activeThreadId) { + const thread = await db.threads.get(activeThreadId); + if (thread?.modelId && thread.modelId !== value) { + const messages = await db.messages + .where("threadId") + .equals(activeThreadId) + .toArray(); + const hasImage = messages.some(messageHasImage); + const targetModel = modelsFromStore.find((model) => model.id === value); + const nonVisionWithImages = hasImage && targetModel?.isVision === false; + + switchNote = nonVisionWithImages + ? "Full chat history will be sent to the new model. This chat has images; text-only models may fail." + : hasImage + ? "Full chat history will be sent to the new model. This chat includes images." + : "Full chat history will be sent to the new model."; + } } - await selectModel({ id: value, isLora: meta?.isLora }); + + if (switchNote) { + toast.warning("Model changed for this chat", { + description: switchNote, + duration: 6000, + }); + } + + await selectModel({ + id: value, + isLora: meta?.isLora, + }); })(); }, - [selectModel, ejectModel], + [modelsFromStore, selectModel, view], ); const handleEject = useCallback(() => { void ejectModel(); @@ -361,36 +423,8 @@ export function ChatPage(): ReactElement { const handleThreadSelect = useCallback( (nextView: ChatView) => { setView(nextView); - - const threadId = - nextView.mode === "single" ? nextView.threadId : undefined; - const pairId = - nextView.mode === "compare" ? nextView.pairId : undefined; - - void (async () => { - let thread: import("./types").ThreadRecord | undefined; - if (threadId) { - thread = await db.threads.get(threadId); - } else if (pairId) { - thread = await db.threads - .where("pairId") - .equals(pairId) - .first(); - } - const threadModelId = thread?.modelId; - if (!threadModelId) return; - - const currentCheckpoint = - useChatRuntimeStore.getState().params.checkpoint; - if (threadModelId === currentCheckpoint) return; - - if (currentCheckpoint) { - await ejectModel(); - } - await selectModel({ id: threadModelId }); - })(); }, - [ejectModel, selectModel], + [], ); const models = useMemo( diff --git a/studio/frontend/src/features/chat/chat-settings-sheet.tsx b/studio/frontend/src/features/chat/chat-settings-sheet.tsx index 26b6efdc09..2d7ae781fa 100644 --- a/studio/frontend/src/features/chat/chat-settings-sheet.tsx +++ b/studio/frontend/src/features/chat/chat-settings-sheet.tsx @@ -320,7 +320,7 @@ export function ChatSettingsPanel({ label="Max Tokens" value={params.maxTokens} min={64} - max={4096} + max={4092} step={64} onChange={set("maxTokens")} /> diff --git a/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts b/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts index 1c8922e929..aafbf00310 100644 --- a/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts +++ b/studio/frontend/src/features/chat/hooks/use-chat-model-runtime.ts @@ -20,6 +20,7 @@ const DEFAULT_MODEL_MAX_SEQ_LENGTH = 2048; type SelectedModelInput = { id: string; isLora?: boolean; + loadingDescription?: string; }; const LORA_SUFFIX_RE = /_(\d{9,})$/; @@ -159,11 +160,22 @@ export function useChatModelRuntime() { const explicitIsLora = typeof selection === "string" ? undefined : selection.isLora; + const extraLoadingDescription = + typeof selection === "string" ? undefined : selection.loadingDescription; const model = models.find((entry) => entry.id === modelId); const lora = loras.find((entry) => entry.id === modelId); const isLora = explicitIsLora ?? model?.isLora ?? (lora ? true : false); const displayName = model?.name || lora?.name || modelId; + const currentCheckpoint = + useChatRuntimeStore.getState().params.checkpoint; + const loadingDescription = [ + currentCheckpoint ? "Unloading previous model first." : null, + extraLoadingDescription ?? null, + "This may include downloading. Large models can take a while.", + ] + .filter(Boolean) + .join(" "); setModelsError(null); setLoadingModel({ id: modelId, displayName }); @@ -197,8 +209,7 @@ export function useChatModelRuntime() { success: `${displayName} loaded`, error: (err) => err instanceof Error ? err.message : "Failed to load model", - description: - "This may include downloading. Large models can take a while.", + description: loadingDescription, }); } catch (error) { setLoadingModel(null); @@ -224,7 +235,7 @@ export function useChatModelRuntime() { await toast.promise(performUnload(), { loading: "Unloading model", - success: "Model unloaded", + success: { message: "Model unloaded", duration: 1200 }, error: (err) => err instanceof Error ? err.message : "Failed to unload model", description: "Releases VRAM and resets inference state.", diff --git a/studio/frontend/src/features/chat/types/runtime.ts b/studio/frontend/src/features/chat/types/runtime.ts index 87eb4c565c..48f147d9bc 100644 --- a/studio/frontend/src/features/chat/types/runtime.ts +++ b/studio/frontend/src/features/chat/types/runtime.ts @@ -15,7 +15,7 @@ export const DEFAULT_INFERENCE_PARAMS: InferenceParams = { topK: 50, minP: 0.01, repetitionPenalty: 1.1, - maxTokens: 512, + maxTokens: 4092, systemPrompt: "", checkpoint: "", };