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:
parent
2e9f756ca6
commit
1eb07f6ad2
4 changed files with 102 additions and 129 deletions
|
|
@ -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);
|
||||
}
|
||||
},
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 };
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue