From 1eb07f6ad2f6cb68b9ef1984ae4f55c17daa06e9 Mon Sep 17 00:00:00 2001 From: Shine1i Date: Sun, 15 Feb 2026 18:59:03 +0100 Subject: [PATCH] refactor: streamline chat runtime logic and remove warming indicator - Replaced `setThreadWarming` logic with streamlined token settlement functions (`settleFirstTokenOk` and `settleFirstTokenErr`) for improved readability and reliability. - Simplified model loading/unloading functions with reusable `performLoad` and `performUnload` patterns. - Removed `warmingByThreadId` from runtime store and associated code for reduced complexity. - Enhanced title generation flow by consolidating logic for persisting and streaming titles. --- .../src/features/chat/api/chat-adapter.ts | 68 +++++++-------- .../chat/hooks/use-chat-model-runtime.ts | 77 +++++++++-------- .../src/features/chat/runtime-provider.tsx | 85 ++++++++----------- .../chat/stores/chat-runtime-store.ts | 13 --- 4 files changed, 108 insertions(+), 135 deletions(-) diff --git a/studio/frontend/src/features/chat/api/chat-adapter.ts b/studio/frontend/src/features/chat/api/chat-adapter.ts index 89cbfcaa47..45e172bcd4 100644 --- a/studio/frontend/src/features/chat/api/chat-adapter.ts +++ b/studio/frontend/src/features/chat/api/chat-adapter.ts @@ -103,8 +103,8 @@ async function resolveUseAdapter( export function createOpenAIStreamAdapter(): ChatModelAdapter { return { async *run({ messages, abortSignal, unstable_threadId }) { - const state = useChatRuntimeStore.getState(); - const { params } = state; + const runtime = useChatRuntimeStore.getState(); + const { params } = runtime; if (!params.checkpoint) { toast.error("No model loaded", { @@ -130,19 +130,33 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter { const threadKey = unstable_threadId || "__default"; let waitingFirstChunk = true; - let hasResolvedFirstToken = false; - let resolveFirstToken: (() => void) | undefined; - let rejectFirstToken: ((err: unknown) => void) | undefined; + let firstTokenSettled = false; + let resolveFirstToken: (() => void) | null = null; + let rejectFirstToken: ((err: unknown) => void) | null = null; const firstTokenPromise = new Promise((resolve, reject) => { resolveFirstToken = resolve; rejectFirstToken = reject; }); // Avoid unhandled rejections if toast.promise never attached. void firstTokenPromise.catch(() => {}); + + function settleFirstTokenOk(): void { + if (firstTokenSettled) return; + firstTokenSettled = true; + resolveFirstToken?.(); + } + + function settleFirstTokenErr(err: unknown): void { + if (firstTokenSettled) return; + firstTokenSettled = true; + rejectFirstToken?.(err); + } + let warmupToastShown = false; const warmupDelayMs = 450; const warmupTimer = setTimeout(() => { - if (!waitingFirstChunk || abortSignal.aborted) return; + if (!waitingFirstChunk) return; + if (abortSignal.aborted) return; warmupToastShown = true; toast.promise(firstTokenPromise, { loading: "Warming up model", @@ -153,8 +167,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter { duration: 900, }); }, warmupDelayMs); - useChatRuntimeStore.getState().setThreadWarming(threadKey, true); - useChatRuntimeStore.getState().setThreadRunning(threadKey, true); + runtime.setThreadRunning(threadKey, true); let cumulativeText = ""; let reasoningStartAt: number | null = null; let reasoningDuration = 0; @@ -183,11 +196,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter { } if (waitingFirstChunk) { waitingFirstChunk = false; - useChatRuntimeStore.getState().setThreadWarming(threadKey, false); - if (!hasResolvedFirstToken) { - hasResolvedFirstToken = true; - resolveFirstToken?.(); - } + settleFirstTokenOk(); } cumulativeText += delta; @@ -207,17 +216,9 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter { }; } } - if (!hasResolvedFirstToken) { - hasResolvedFirstToken = true; - resolveFirstToken?.(); - } + settleFirstTokenOk(); } catch (err) { - if (!hasResolvedFirstToken) { - hasResolvedFirstToken = true; - rejectFirstToken?.( - err instanceof Error ? err : new Error("Generation failed"), - ); - } + settleFirstTokenErr(err instanceof Error ? err : new Error("Generation failed")); const isEarly = waitingFirstChunk; if (!abortSignal.aborted && !(warmupToastShown && isEarly)) { toast.error("Generation failed", { @@ -228,20 +229,17 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter { } finally { clearTimeout(warmupTimer); if (waitingFirstChunk) { - useChatRuntimeStore.getState().setThreadWarming(threadKey, false); - if (warmupToastShown && !hasResolvedFirstToken) { - hasResolvedFirstToken = true; - rejectFirstToken?.( - abortSignal.aborted - ? new Error("Cancelled") - : new Error("No tokens received"), - ); - } else if (!hasResolvedFirstToken) { - hasResolvedFirstToken = true; - resolveFirstToken?.(); + if (warmupToastShown && !firstTokenSettled) { + if (abortSignal.aborted) { + settleFirstTokenErr(new Error("Cancelled")); + } else { + settleFirstTokenErr(new Error("No tokens received")); + } + } else { + settleFirstTokenOk(); } } - useChatRuntimeStore.getState().setThreadRunning(threadKey, false); + runtime.setThreadRunning(threadKey, false); } }, }; 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 febf6da367..e8612eb625 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 @@ -128,31 +128,35 @@ export function useChatModelRuntime() { setModelsError(null); try { - await toast.promise( - (async () => { - if (params.checkpoint) { - await unloadModel({ model_path: params.checkpoint }); - } + async function performLoad(): Promise { + if (params.checkpoint) { + await unloadModel({ model_path: params.checkpoint }); + } - await loadModel({ - model_path: modelId, - hf_token: null, - max_seq_length: DEFAULT_MODEL_MAX_SEQ_LENGTH, - load_in_4bit: true, - is_lora: isLora, - }); + await loadModel({ + model_path: modelId, + hf_token: null, + max_seq_length: DEFAULT_MODEL_MAX_SEQ_LENGTH, + load_in_4bit: true, + is_lora: isLora, + }); - setCheckpoint(modelId); - await refresh(); - })(), - { - loading: `Loading ${displayName}`, - success: `${displayName} loaded`, - error: (err) => - err instanceof Error ? err.message : "Failed to load model", - description: isLora ? "Fine-tuned (LoRA) selected." : "Base model selected.", - }, - ); + setCheckpoint(modelId); + await refresh(); + } + + let description = "Base model selected."; + if (isLora) { + description = "Fine-tuned (LoRA) selected."; + } + + await toast.promise(performLoad(), { + loading: `Loading ${displayName}`, + success: `${displayName} loaded`, + error: (err) => + err instanceof Error ? err.message : "Failed to load model", + description, + }); } catch (error) { const message = error instanceof Error ? error.message : "Failed to load model"; @@ -168,20 +172,19 @@ export function useChatModelRuntime() { } setModelsError(null); try { - await toast.promise( - (async () => { - await unloadModel({ model_path: params.checkpoint }); - clearCheckpoint(); - await refresh(); - })(), - { - loading: "Unloading model", - success: "Model unloaded", - error: (err) => - err instanceof Error ? err.message : "Failed to unload model", - description: "Releases VRAM and resets inference state.", - }, - ); + async function performUnload(): Promise { + await unloadModel({ model_path: params.checkpoint }); + clearCheckpoint(); + await refresh(); + } + + await toast.promise(performUnload(), { + loading: "Unloading model", + success: "Model unloaded", + error: (err) => + err instanceof Error ? err.message : "Failed to unload model", + description: "Releases VRAM and resets inference state.", + }); } catch (error) { const message = error instanceof Error ? error.message : "Failed to unload model"; diff --git a/studio/frontend/src/features/chat/runtime-provider.tsx b/studio/frontend/src/features/chat/runtime-provider.tsx index 090a2cb162..28bcf83421 100644 --- a/studio/frontend/src/features/chat/runtime-provider.tsx +++ b/studio/frontend/src/features/chat/runtime-provider.tsx @@ -331,47 +331,47 @@ function createDexieAdapter( async generateTitle(remoteId: string, messages: readonly ThreadMessage[]) { const autoTitle = useChatRuntimeStore.getState().autoTitle; const thread = await db.threads.get(remoteId); - if (!thread) { - return createAssistantStream((c) => { - c.appendText("New Chat"); - c.close(); - }); - } + const defaultTitle = "New Chat"; - // Only generate once per thread/pair. - if (thread.title && thread.title !== "New Chat") { - return createAssistantStream((c) => { - c.appendText(thread.title); - c.close(); - }); - } - - const firstUser = messages.find((m) => m.role === "user"); - const userText = extractTextParts(firstUser) || "New Chat"; - - if (!autoTitle) { - const title = fallbackTitleFromUserText(userText); - await db.threads.update(remoteId, { title }); - if (pairId) { - const paired = await db.threads - .where("pairId") - .equals(pairId) - .filter((t) => t.id !== remoteId) - .first(); - if (paired) await db.threads.update(paired.id, { title }); - } + function streamTitle(title: string) { return createAssistantStream((c) => { c.appendText(title); c.close(); }); } + async function persistTitle(title: string): Promise { + await db.threads.update(remoteId, { title }); + if (!pairId) return; + const paired = await db.threads + .where("pairId") + .equals(pairId) + .filter((t) => t.id !== remoteId) + .first(); + if (paired) await db.threads.update(paired.id, { title }); + } + + if (!thread) { + return streamTitle(defaultTitle); + } + + // Only generate once per thread/pair. + if (thread.title && thread.title !== "New Chat") { + return streamTitle(thread.title); + } + + const firstUser = messages.find((m) => m.role === "user"); + const userText = extractTextParts(firstUser) || defaultTitle; + + if (!autoTitle) { + const title = fallbackTitleFromUserText(userText); + await persistTitle(title); + return streamTitle(title); + } + const key = pairId ? `pair:${pairId}` : `thread:${remoteId}`; if (inflightTitleByKey.has(key)) { - return createAssistantStream((c) => { - c.appendText(thread.title || "New Chat"); - c.close(); - }); + return streamTitle(thread.title || defaultTitle); } // Compare: wait until both threads done. @@ -388,10 +388,7 @@ function createDexieAdapter( setTimeout(() => { void createDexieAdapter(modelType, pairId).generateTitle(remoteId, messages); }, 600); - return createAssistantStream((c) => { - c.appendText(thread.title || "New Chat"); - c.close(); - }); + return streamTitle(thread.title || defaultTitle); } } } @@ -404,20 +401,8 @@ function createDexieAdapter( })) || fallbackTitleFromUserText(userText); - await db.threads.update(remoteId, { title }); - if (pairId) { - const paired = await db.threads - .where("pairId") - .equals(pairId) - .filter((t) => t.id !== remoteId) - .first(); - if (paired) await db.threads.update(paired.id, { title }); - } - - return createAssistantStream((c) => { - c.appendText(title); - c.close(); - }); + await persistTitle(title); + return streamTitle(title); } finally { inflightTitleByKey.delete(key); } diff --git a/studio/frontend/src/features/chat/stores/chat-runtime-store.ts b/studio/frontend/src/features/chat/stores/chat-runtime-store.ts index 10d5957906..03825b5861 100644 --- a/studio/frontend/src/features/chat/stores/chat-runtime-store.ts +++ b/studio/frontend/src/features/chat/stores/chat-runtime-store.ts @@ -36,14 +36,12 @@ type ChatRuntimeStore = { params: InferenceParams; models: ChatModelSummary[]; loras: ChatLoraSummary[]; - warmingByThreadId: Record; runningByThreadId: Record; autoTitle: boolean; modelsError: string | null; setParams: (params: InferenceParams) => void; setModels: (models: ChatModelSummary[]) => void; setLoras: (loras: ChatLoraSummary[]) => void; - setThreadWarming: (threadId: string, warming: boolean) => void; setThreadRunning: (threadId: string, running: boolean) => void; setAutoTitle: (enabled: boolean) => void; setModelsError: (error: string | null) => void; @@ -55,23 +53,12 @@ export const useChatRuntimeStore = create((set) => ({ params: DEFAULT_INFERENCE_PARAMS, models: [], loras: [], - warmingByThreadId: {}, runningByThreadId: {}, autoTitle: loadBool(AUTO_TITLE_KEY, false), modelsError: null, setParams: (params) => set({ params }), setModels: (models) => set({ models }), setLoras: (loras) => set({ loras }), - setThreadWarming: (threadId, warming) => - set((state) => { - const next = { ...state.warmingByThreadId }; - if (warming) { - next[threadId] = true; - } else { - delete next[threadId]; - } - return { warmingByThreadId: next }; - }), setThreadRunning: (threadId, running) => set((state) => { const next = { ...state.runningByThreadId };