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.
This commit is contained in:
Shine1i 2026-02-15 18:59:03 +01:00
commit 1eb07f6ad2
4 changed files with 102 additions and 129 deletions

View file

@ -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<void>((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);
}
},
};

View file

@ -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<void> {
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<void> {
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";

View file

@ -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<void> {
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);
}

View file

@ -36,14 +36,12 @@ type ChatRuntimeStore = {
params: InferenceParams;
models: ChatModelSummary[];
loras: ChatLoraSummary[];
warmingByThreadId: Record<string, boolean>;
runningByThreadId: Record<string, boolean>;
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<ChatRuntimeStore>((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 };