Merge pull request #100 from unslothai/feature/chat-compare

feat: chat compare + inference stream cancel fix
This commit is contained in:
Wasim Yousef Said 2026-02-15 12:29:54 -08:00 committed by GitHub
commit 28364f3314
16 changed files with 534 additions and 1466 deletions

View file

@ -489,6 +489,7 @@ class InferenceBackend:
def generate_with_adapter_control(
self,
use_adapter: Optional[Union[bool, str]] = None,
cancel_event=None,
**gen_kwargs,
) -> Generator[str, None, None]:
"""
@ -505,7 +506,7 @@ class InferenceBackend:
with self._generation_lock:
self._apply_adapter_state(use_adapter)
# Delegate to the lock-free generation path
yield from self._generate_chat_response_inner(**gen_kwargs)
yield from self._generate_chat_response_inner(cancel_event=cancel_event, **gen_kwargs)
def generate_chat_response(self,
messages: list,
@ -515,7 +516,8 @@ class InferenceBackend:
top_p: float = 0.9,
top_k: int = 40,
max_new_tokens: int = 256,
repetition_penalty: float = 1.1) -> Generator[str, None, None]:
repetition_penalty: float = 1.1,
cancel_event=None) -> Generator[str, None, None]:
"""
Generate response for text or vision models.
Acquires the generation lock. For adapter-controlled generation,
@ -531,6 +533,7 @@ class InferenceBackend:
top_k=top_k,
max_new_tokens=max_new_tokens,
repetition_penalty=repetition_penalty,
cancel_event=cancel_event,
)
def _generate_chat_response_inner(self,
@ -541,7 +544,8 @@ class InferenceBackend:
top_p: float = 0.9,
top_k: int = 40,
max_new_tokens: int = 256,
repetition_penalty: float = 1.1) -> Generator[str, None, None]:
repetition_penalty: float = 1.1,
cancel_event=None) -> Generator[str, None, None]:
"""
Inner generation logic (no lock). Called by both generate_chat_response
and generate_with_adapter_control.
@ -558,7 +562,8 @@ class InferenceBackend:
# Vision model generation
yield from self._generate_vision_response(
messages, system_prompt, image,
temperature, top_p, top_k, max_new_tokens, repetition_penalty
temperature, top_p, top_k, max_new_tokens, repetition_penalty,
cancel_event=cancel_event,
)
else:
# Text model: Use training pipeline approach
@ -600,12 +605,13 @@ class InferenceBackend:
# Step 3: Generate
yield from self.generate_stream(
formatted_prompt, temperature, top_p, top_k, max_new_tokens, repetition_penalty
formatted_prompt, temperature, top_p, top_k, max_new_tokens, repetition_penalty,
cancel_event=cancel_event,
)
def _generate_vision_response(self, messages, system_prompt, image,
temperature, top_p, top_k, max_new_tokens,
repetition_penalty) -> Generator[str, None, None]:
repetition_penalty, cancel_event=None) -> Generator[str, None, None]:
"""Handle vision model generation with true token-by-token streaming."""
model_info = self.models[self.active_model_name]
model = model_info["model"]
@ -651,7 +657,10 @@ class InferenceBackend:
import threading
streamer = TextIteratorStreamer(
processor.tokenizer, skip_prompt=True, skip_special_tokens=True
processor.tokenizer,
skip_prompt=True,
skip_special_tokens=True,
timeout=0.2,
)
generation_kwargs = dict(
@ -664,23 +673,50 @@ class InferenceBackend:
top_k=top_k,
)
err: dict[str, str] = {}
def generate_fn():
try:
model.generate(**generation_kwargs)
except Exception as e:
err["msg"] = str(e)
logger.error(f"Vision generation error in thread: {e}")
finally:
try:
streamer.end()
except Exception:
pass
thread = threading.Thread(target=generate_fn)
thread.start()
output = ""
for new_token in streamer:
if new_token:
output += new_token
cleaned = self._clean_generated_text(output)
yield cleaned
from queue import Empty
try:
while True:
if cancel_event is not None and cancel_event.is_set():
break
try:
new_token = next(streamer)
except StopIteration:
break
except Empty:
if not thread.is_alive():
break
continue
if new_token:
output += new_token
cleaned = self._clean_generated_text(output)
yield cleaned
finally:
if cancel_event is not None:
cancel_event.set()
thread.join(timeout=10)
if thread.is_alive():
logger.warning("Vision generation thread did not exit after cancel/join timeout")
thread.join()
if err.get("msg"):
yield f"Error: {err['msg']}"
except Exception as e:
logger.error(f"Vision generation error: {e}")
@ -693,7 +729,8 @@ class InferenceBackend:
top_p: float = 0.9,
top_k: int = 40,
max_new_tokens: int = 256,
repetition_penalty: float = 1.1) -> Generator[str, None, None]:
repetition_penalty: float = 1.1,
cancel_event=None) -> Generator[str, None, None]:
"""Generate streaming text response (text models only)."""
if not self.active_model_name:
yield "Error: No active model"
@ -709,7 +746,12 @@ class InferenceBackend:
from transformers import TextIteratorStreamer
import threading
streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True)
streamer = TextIteratorStreamer(
tokenizer,
skip_prompt=True,
skip_special_tokens=True,
timeout=0.2,
)
generation_kwargs = dict(
**inputs,
@ -723,24 +765,66 @@ class InferenceBackend:
eos_token_id=tokenizer.eos_token_id,
pad_token_id=tokenizer.eos_token_id if tokenizer.pad_token_id is None else tokenizer.pad_token_id,
)
if cancel_event is not None:
from transformers.generation.stopping_criteria import (
StoppingCriteria,
StoppingCriteriaList,
)
class _CancelCriteria(StoppingCriteria):
def __init__(self, ev):
self.ev = ev
def __call__(self, input_ids, scores, **kwargs):
return self.ev.is_set()
generation_kwargs["stopping_criteria"] = StoppingCriteriaList(
[_CancelCriteria(cancel_event)]
)
def generate_fn():
try:
model.generate(**generation_kwargs)
except Exception as e:
err["msg"] = str(e)
logger.error(f"Generation error: {e}")
finally:
try:
streamer.end()
except Exception:
pass
err: dict[str, str] = {}
thread = threading.Thread(target=generate_fn)
thread.start()
output = ""
for new_token in streamer:
if new_token:
output += new_token
cleaned = self._clean_generated_text(output)
yield cleaned
from queue import Empty
try:
while True:
if cancel_event is not None and cancel_event.is_set():
break
try:
new_token = next(streamer)
except StopIteration:
break
except Empty:
if not thread.is_alive():
break
continue
if new_token:
output += new_token
cleaned = self._clean_generated_text(output)
yield cleaned
finally:
if cancel_event is not None:
cancel_event.set()
thread.join(timeout=10)
if thread.is_alive():
logger.warning("Generation thread did not exit after cancel/join timeout")
thread.join()
if err.get("msg"):
yield f"Error: {err['msg']}"
except Exception as e:
logger.error(f"Error during generation: {e}")

View file

@ -5,11 +5,13 @@ import sys
import time
import uuid
from pathlib import Path
from fastapi import APIRouter, HTTPException
from fastapi import APIRouter, HTTPException, Request
from fastapi.responses import StreamingResponse, JSONResponse
from typing import Optional
import json
import logging
import asyncio
import threading
@ -304,7 +306,7 @@ def _extract_content_parts(
@router.post("/chat/completions")
async def openai_chat_completions(request: ChatCompletionRequest):
async def openai_chat_completions(payload: ChatCompletionRequest, request: Request):
"""
OpenAI-compatible chat completions endpoint.
@ -324,7 +326,7 @@ async def openai_chat_completions(request: ChatCompletionRequest):
# ── Parse messages (handles multimodal content parts) ─────
system_prompt, chat_messages, extracted_image_b64 = _extract_content_parts(
request.messages
payload.messages
)
# If no non-system messages were provided, error out
@ -336,7 +338,7 @@ async def openai_chat_completions(request: ChatCompletionRequest):
# ── Decode image (from content parts OR legacy field) ─────
# Content-part images take priority; fall back to legacy field
image_b64 = extracted_image_b64 or request.image_base64
image_b64 = extracted_image_b64 or payload.image_base64
image = None
if image_b64:
@ -366,31 +368,35 @@ async def openai_chat_completions(request: ChatCompletionRequest):
messages=chat_messages,
system_prompt=system_prompt,
image=image,
temperature=request.temperature,
top_p=request.top_p,
top_k=request.top_k,
max_new_tokens=request.max_tokens or 512,
repetition_penalty=request.repetition_penalty,
temperature=payload.temperature,
top_p=payload.top_p,
top_k=payload.top_k,
max_new_tokens=payload.max_tokens or 512,
repetition_penalty=payload.repetition_penalty,
)
# ── Choose generation path (adapter-controlled or standard) ──
if request.use_adapter is not None:
cancel_event = threading.Event()
if payload.use_adapter is not None:
# Compare mode: toggle adapter state atomically with generation
def generate():
return backend.generate_with_adapter_control(
use_adapter=request.use_adapter, **gen_kwargs
use_adapter=payload.use_adapter,
cancel_event=cancel_event,
**gen_kwargs,
)
else:
# Standard path: no adapter toggling
def generate():
return backend.generate_chat_response(**gen_kwargs)
return backend.generate_chat_response(cancel_event=cancel_event, **gen_kwargs)
model_name = backend.active_model_name or request.model
model_name = backend.active_model_name or payload.model
completion_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
created = int(time.time())
# ── Streaming response ────────────────────────────────────────
if request.stream:
if payload.stream:
async def stream_chunks():
try:
# First chunk: send the role
@ -409,6 +415,10 @@ async def openai_chat_completions(request: ChatCompletionRequest):
# text, so we diff to get incremental deltas.
prev_text = ""
for cumulative in generate():
if await request.is_disconnected():
cancel_event.set()
backend.reset_generation_state()
return
new_text = cumulative[len(prev_text):]
prev_text = cumulative
if not new_text:
@ -437,6 +447,10 @@ async def openai_chat_completions(request: ChatCompletionRequest):
yield f"data: {final_chunk.model_dump_json(exclude_none=True)}\n\n"
yield "data: [DONE]\n\n"
except asyncio.CancelledError:
cancel_event.set()
backend.reset_generation_state()
raise
except Exception as e:
backend.reset_generation_state()
logger.error(f"Error during OpenAI streaming: {e}", exc_info=True)
@ -477,4 +491,3 @@ async def openai_chat_completions(request: ChatCompletionRequest):
backend.reset_generation_state()
logger.error(f"Error during OpenAI completion: {e}", exc_info=True)
raise HTTPException(status_code=500, detail=str(e))

View file

@ -41,6 +41,13 @@ def is_local_path(path: str) -> bool:
if not path:
return False
# If it exists on disk, treat as local (covers relative paths like "outputs/foo").
try:
if Path(normalize_path(path)).expanduser().exists():
return True
except Exception:
pass
# Obvious HF patterns
if path.count('/') == 1 and not path.startswith(('/', '.', '~')):
return False # Looks like org/model format

View file

@ -2,7 +2,6 @@ import { createRouter } from "@tanstack/react-router";
import { Route as rootRoute } from "./routes/__root";
import { Route as chatRoute } from "./routes/chat";
import { Route as gridTestRoute } from "./routes/grid-test";
import { Route as homeRoute } from "./routes/home";
import { Route as loginRoute } from "./routes/login";
import { Route as onboardingRoute } from "./routes/onboarding";
import { Route as exportRoute } from "./routes/export";
@ -10,7 +9,6 @@ import { Route as signupRoute } from "./routes/signup";
import { Route as studioRoute } from "./routes/studio";
const routeTree = rootRoute.addChildren([
homeRoute,
onboardingRoute,
loginRoute,
signupRoute,

View file

@ -1,15 +0,0 @@
import { ComponentExample } from "@/components/component-example";
import { createRoute } from "@tanstack/react-router";
import { requireAuth } from "../auth-guards";
import { Route as rootRoute } from "./__root";
export const Route = createRoute({
getParentRoute: () => rootRoute,
path: "/",
beforeLoad: () => requireAuth(),
component: HomePage,
});
function HomePage() {
return <ComponentExample />;
}

View file

@ -7,9 +7,7 @@ import { MarkdownText } from "@/components/assistant-ui/markdown-text";
import { Reasoning, ReasoningGroup } from "@/components/assistant-ui/reasoning";
import { ToolFallback } from "@/components/assistant-ui/tool-fallback";
import { TooltipIconButton } from "@/components/assistant-ui/tooltip-icon-button";
import { AnimatedShinyText } from "@/components/ui/animated-shiny-text";
import { Button } from "@/components/ui/button";
import { useChatRuntimeStore } from "@/features/chat/stores/chat-runtime-store";
import { cn } from "@/lib/utils";
import {
ActionBarMorePrimitive,
@ -73,7 +71,6 @@ export const Thread: FC<{ hideComposer?: boolean; hideWelcome?: boolean }> = ({
<ThreadPrimitive.ViewportFooter className="aui-thread-viewport-footer sticky bottom-0 mt-auto flex w-full flex-col gap-4 overflow-visible bg-background pb-4 md:pb-4 before:pointer-events-none before:absolute before:inset-x-0 before:bottom-full before:h-20 before:bg-gradient-to-t before:from-background before:to-transparent">
<ThreadScrollToBottom />
<WarmupIndicator />
<AuiIf condition={({ thread }) => !thread.isEmpty}>
{!hideComposer && <ComposerAnimated />}
</AuiIf>
@ -83,28 +80,6 @@ export const Thread: FC<{ hideComposer?: boolean; hideWelcome?: boolean }> = ({
);
};
const WarmupIndicator: FC = () => {
const threadId = useAuiState(({ threads }) => threads.mainThreadId);
const isRunning = useAuiState(({ thread }) => thread.isRunning);
const isWarmingUp = useChatRuntimeStore((state) =>
Boolean(state.warmingByThreadId[threadId ?? "__default"]),
);
if (!isRunning || !isWarmingUp) {
return null;
}
return (
<div className="mx-auto -mb-2 w-full max-w-(--thread-max-width) px-2">
<div className="inline-flex items-center rounded-full border border-border/60 bg-background/90 px-3 py-1.5 text-xs text-muted-foreground shadow-sm">
<AnimatedShinyText className="text-xs">
Warming up model...
</AnimatedShinyText>
</div>
</div>
);
};
const ThreadScrollToBottom: FC = () => {
return (
<ThreadPrimitive.ScrollToBottom asChild={true}>

File diff suppressed because it is too large Load diff

View file

@ -1,5 +1,7 @@
import type { ChatModelAdapter } from "@assistant-ui/react";
import { toast } from "sonner";
import { streamChatCompletions } from "./chat-api";
import { db } from "../db";
import { useChatRuntimeStore } from "../stores/chat-runtime-store";
import {
hasClosedThinkTag,
@ -81,13 +83,33 @@ function findLatestUserImageBase64(messages: RunMessages): string | undefined {
return undefined;
}
async function resolveUseAdapter(
threadId: string | undefined,
): Promise<boolean | undefined> {
if (!threadId) {
return undefined;
}
try {
const thread = await db.threads.get(threadId);
if (!thread?.pairId) {
return undefined;
}
return thread.modelType === "lora";
} catch {
return undefined;
}
}
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", {
description: "Pick model in top bar, then retry.",
});
throw new Error("Load a model first.");
}
@ -104,10 +126,48 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
});
}
const imageBase64 = findLatestUserImageBase64(messages);
const useAdapter = await resolveUseAdapter(unstable_threadId);
const threadKey = unstable_threadId || "__default";
let waitingFirstChunk = true;
useChatRuntimeStore.getState().setThreadWarming(threadKey, true);
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) return;
if (abortSignal.aborted) return;
warmupToastShown = true;
toast.promise(firstTokenPromise, {
loading: "Warming up model",
success: "Generating",
error: (err) =>
err instanceof Error && err.message ? err.message : "Generation failed",
description: "Waiting for first token.",
duration: 900,
});
}, warmupDelayMs);
runtime.setThreadRunning(threadKey, true);
let cumulativeText = "";
let reasoningStartAt: number | null = null;
let reasoningDuration = 0;
@ -124,6 +184,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
top_k: params.topK,
repetition_penalty: params.repetitionPenalty,
image_base64: imageBase64,
...(useAdapter === undefined ? {} : { use_adapter: useAdapter }),
},
abortSignal,
);
@ -135,7 +196,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
}
if (waitingFirstChunk) {
waitingFirstChunk = false;
useChatRuntimeStore.getState().setThreadWarming(threadKey, false);
settleFirstTokenOk();
}
cumulativeText += delta;
@ -155,10 +216,30 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
};
}
}
} finally {
if (waitingFirstChunk) {
useChatRuntimeStore.getState().setThreadWarming(threadKey, false);
settleFirstTokenOk();
} catch (err) {
settleFirstTokenErr(err instanceof Error ? err : new Error("Generation failed"));
const isEarly = waitingFirstChunk;
if (!abortSignal.aborted && !(warmupToastShown && isEarly)) {
toast.error("Generation failed", {
description: err instanceof Error ? err.message : "Unknown error",
});
}
throw err;
} finally {
clearTimeout(warmupTimer);
if (waitingFirstChunk) {
if (warmupToastShown && !firstTokenSettled) {
if (abortSignal.aborted) {
settleFirstTokenErr(new Error("Cancelled"));
} else {
settleFirstTokenErr(new Error("No tokens received"));
}
} else {
settleFirstTokenOk();
}
}
runtime.setThreadRunning(threadKey, false);
}
},
};

View file

@ -174,7 +174,8 @@ function InlineSidebar({
function TopBarActions({
onNewThread,
onNewCompare,
}: { onNewThread: () => void; onNewCompare: () => void }) {
showCompare,
}: { onNewThread: () => void; onNewCompare: () => void; showCompare: boolean }) {
const { state } = useSidebar();
if (state !== "collapsed") {
return null;
@ -189,14 +190,16 @@ function TopBarActions({
</TooltipTrigger>
<TooltipContent side="bottom">New Chat</TooltipContent>
</Tooltip>
<Tooltip>
<TooltipTrigger asChild={true}>
<Button variant="ghost" size="icon-sm" onClick={onNewCompare}>
<HugeiconsIcon icon={ColumnInsertIcon} strokeWidth={2} />
</Button>
</TooltipTrigger>
<TooltipContent side="bottom">Compare</TooltipContent>
</Tooltip>
{showCompare ? (
<Tooltip>
<TooltipTrigger asChild={true}>
<Button variant="ghost" size="icon-sm" onClick={onNewCompare}>
<HugeiconsIcon icon={ColumnInsertIcon} strokeWidth={2} />
</Button>
</TooltipTrigger>
<TooltipContent side="bottom">Compare</TooltipContent>
</Tooltip>
) : null}
</>
);
}
@ -209,10 +212,17 @@ export function ChatPage(): ReactElement {
const [settingsOpen, setSettingsOpen] = useState(false);
const inferenceParams = useChatRuntimeStore((state) => state.params);
const setInferenceParams = useChatRuntimeStore((state) => state.setParams);
const autoTitle = useChatRuntimeStore((state) => state.autoTitle);
const setAutoTitle = useChatRuntimeStore((state) => state.setAutoTitle);
const modelsFromStore = useChatRuntimeStore((state) => state.models);
const lorasFromStore = useChatRuntimeStore((state) => state.loras);
const modelsError = useChatRuntimeStore((state) => state.modelsError);
const { refresh, selectModel, ejectModel } = useChatModelRuntime();
const canCompare = useMemo(() => {
const selected = inferenceParams.checkpoint;
if (!selected) return false;
return lorasFromStore.some((lora) => lora.id === selected);
}, [inferenceParams.checkpoint, lorasFromStore]);
const handleCheckpointChange = useCallback(
(value: string, meta?: { isLora: boolean }) => {
@ -275,6 +285,7 @@ export function ChatPage(): ReactElement {
onSelect={setView}
onNewThread={handleNewThread}
onNewCompare={handleNewCompare}
showCompare={canCompare}
/>
</InlineSidebar>
@ -285,6 +296,7 @@ export function ChatPage(): ReactElement {
<TopBarActions
onNewThread={handleNewThread}
onNewCompare={handleNewCompare}
showCompare={canCompare}
/>
<ModelSelector
models={models}
@ -326,6 +338,8 @@ export function ChatPage(): ReactElement {
open={settingsOpen}
params={inferenceParams}
onParamsChange={setInferenceParams}
autoTitle={autoTitle}
onAutoTitleChange={setAutoTitle}
/>
</SidebarProvider>
</div>

View file

@ -23,6 +23,7 @@ import {
DEFAULT_INFERENCE_PARAMS,
type InferenceParams,
} from "./types/runtime";
import { Switch } from "@/components/ui/switch";
export const defaultInferenceParams = DEFAULT_INFERENCE_PARAMS;
export type { InferenceParams } from "./types/runtime";
@ -143,12 +144,16 @@ interface ChatSettingsPanelProps {
open: boolean;
params: InferenceParams;
onParamsChange: (params: InferenceParams) => void;
autoTitle: boolean;
onAutoTitleChange: (enabled: boolean) => void;
}
export function ChatSettingsPanel({
open,
params,
onParamsChange,
autoTitle,
onAutoTitleChange,
}: ChatSettingsPanelProps) {
const [presets, setPresets] = useState<Preset[]>(BUILTIN_PRESETS);
const [activePreset, setActivePreset] = useState("Default");
@ -268,7 +273,7 @@ export function ChatSettingsPanel({
<CollapsibleSection
icon={SlidersHorizontalIcon}
label="Sampling"
defaultOpen={true}
defaultOpen={false}
>
<div className="flex flex-col gap-5">
<ParamSlider
@ -315,9 +320,18 @@ export function ChatSettingsPanel({
</CollapsibleSection>
<CollapsibleSection icon={Settings02Icon} label="Settings">
<p className="text-xs text-muted-foreground">
No additional settings yet.
</p>
<div className="flex items-center justify-between gap-3 py-1">
<div className="min-w-0">
<div className="text-xs font-medium">Auto title</div>
<div className="text-[11px] text-muted-foreground">
Generate short title after reply.
</div>
</div>
<Switch
checked={autoTitle}
onCheckedChange={onAutoTitleChange}
/>
</div>
</CollapsibleSection>
</div>
</div>

View file

@ -105,6 +105,9 @@ export function useChatModelRuntime() {
const message =
error instanceof Error ? error.message : "Failed to load models";
setModelsError(message);
toast.error("Failed to refresh models", {
description: message,
});
}
}, [setCheckpoint, setLoras, setModels, setModelsError]);
@ -122,30 +125,42 @@ export function useChatModelRuntime() {
const isLora =
explicitIsLora ?? model?.isLora ?? (lora ? true : false);
const displayName = model?.name || lora?.name || modelId;
const loadingToastId = toast.loading(`Loading ${displayName}...`);
setModelsError(null);
try {
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,
});
setCheckpoint(modelId);
await refresh();
}
await loadModel({
model_path: modelId,
hf_token: null,
max_seq_length: DEFAULT_MODEL_MAX_SEQ_LENGTH,
load_in_4bit: true,
is_lora: isLora,
});
let description = "Base model selected.";
if (isLora) {
description = "Fine-tuned (LoRA) selected.";
}
setCheckpoint(modelId);
await refresh();
toast.success(`${displayName} loaded`, { id: loadingToastId });
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";
setModelsError(message);
toast.error(message, { id: loadingToastId });
}
},
[loras, models, params.checkpoint, refresh, setCheckpoint, setModelsError],
@ -157,9 +172,19 @@ export function useChatModelRuntime() {
}
setModelsError(null);
try {
await unloadModel({ model_path: params.checkpoint });
clearCheckpoint();
await refresh();
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

@ -11,7 +11,6 @@ import {
SimpleTextAttachmentAdapter,
type ThreadHistoryAdapter,
type ThreadMessage,
type ThreadUserMessagePart,
WebSpeechDictationAdapter,
type unstable_RemoteThreadListAdapter,
useAui,
@ -23,8 +22,10 @@ import { createAssistantStream } from "assistant-stream";
import mammoth from "mammoth";
import { type ReactElement, type ReactNode, useEffect, useMemo } from "react";
import { extractText, getDocumentProxy } from "unpdf";
import { authFetch } from "@/features/auth";
import { createOpenAIStreamAdapter } from "./api/chat-adapter";
import { db } from "./db";
import { useChatRuntimeStore } from "./stores/chat-runtime-store";
import type { MessageRecord, ModelType } from "./types";
const DEFAULT_SUGGESTIONS = [
@ -149,6 +150,89 @@ class DocxAttachmentAdapter implements AttachmentAdapter {
}
}
function clip(input: string, maxLen: number): string {
const text = input.replace(/\s+/g, " ").trim();
if (text.length <= maxLen) return text;
return text.slice(0, maxLen).trimEnd();
}
function extractTextParts(m: ThreadMessage | undefined): string {
if (!m) return "";
const content = Array.isArray(m.content) ? m.content : [];
return content
.filter((p): p is Extract<typeof p, { type: "text" }> => p.type === "text")
.map((p) => p.text)
.join("")
.trim();
}
async function generateTitleWithModel(payload: {
userText: string;
}): Promise<string | null> {
const params = useChatRuntimeStore.getState().params;
if (!params.checkpoint) return null;
const user = clip(payload.userText, 256);
const parts: string[] = [user];
function normalizeTitle(raw: string): string | null {
let title = raw.split(/\r?\n/, 1)[0] ?? "";
title = title.replace(/^\s*title\s*:\s*/i, "");
title = title.replace(/[^\x20-\x7E]+/g, " ");
title = title.replace(/["'`]+/g, "");
title = title.replace(/[.!?:;,]+/g, " ");
title = title.replace(/\s+/g, " ").trim();
// Model echo fail-safe.
if (/\b(user|base|lora|assistant)\s*:/i.test(title)) {
return null;
}
const words = title.split(" ").filter(Boolean).slice(0, 6);
const joined = words.join(" ").trim();
if (!joined) return null;
return joined.length > 60 ? joined.slice(0, 60).trimEnd() : joined;
}
const response = await authFetch("/api/inference/chat/completions", {
method: "POST",
headers: { "Content-Type": "application/json" },
body: JSON.stringify({
model: params.checkpoint,
stream: false,
temperature: 0.2,
top_p: 0.9,
max_tokens: 24,
top_k: 40,
repetition_penalty: 1.05,
messages: [
{
role: "system",
content:
"Write 1 concise chat title for the user's message. Rules: 2-6 words, no quotes, no punctuation, ASCII only, do not echo input. Output title only.",
},
{ role: "user", content: parts.join("\n") },
],
}),
});
const body = (await response.json().catch(() => null)) as any;
if (!response.ok) return null;
const raw: string | undefined = body?.choices?.[0]?.message?.content;
if (!raw) return null;
return normalizeTitle(raw);
}
const inflightTitleByKey = new Set<string>();
function fallbackTitleFromUserText(userText: string): string {
const firstLine = (userText || "").split(/\r?\n/, 1)[0] ?? "";
const cleaned = firstLine.replace(/\s+/g, " ").trim();
const max = 48;
if (!cleaned) return "New Chat";
return cleaned.slice(0, max) + (cleaned.length > max ? "..." : "");
}
function toThreadMessage(m: MessageRecord): ThreadMessage {
const base = {
id: m.id,
@ -245,32 +329,83 @@ function createDexieAdapter(
},
async generateTitle(remoteId: string, messages: readonly ThreadMessage[]) {
const autoTitle = useChatRuntimeStore.getState().autoTitle;
const thread = await db.threads.get(remoteId);
const defaultTitle = "New Chat";
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 textParts =
firstUser?.content.filter(
(part): part is Extract<ThreadUserMessagePart, { type: "text" }> =>
part.type === "text",
) ?? [];
const text = textParts.map((part) => part.text).join("") || "New Chat";
const title = text.slice(0, 60) + (text.length > 60 ? "..." : "");
const userText = extractTextParts(firstUser) || defaultTitle;
await db.threads.update(remoteId, { title });
if (!autoTitle) {
const title = fallbackTitleFromUserText(userText);
await persistTitle(title);
return streamTitle(title);
}
const key = pairId ? `pair:${pairId}` : `thread:${remoteId}`;
if (inflightTitleByKey.has(key)) {
return streamTitle(thread.title || defaultTitle);
}
// Compare: wait until both threads done.
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 });
const running = useChatRuntimeStore.getState().runningByThreadId;
if (running[paired.id]) {
setTimeout(() => {
void createDexieAdapter(modelType, pairId).generateTitle(remoteId, messages);
}, 600);
return streamTitle(thread.title || defaultTitle);
}
}
}
return createAssistantStream((controller) => {
controller.appendText(title);
controller.close();
});
inflightTitleByKey.add(key);
try {
const title =
(await generateTitleWithModel({
userText,
})) ||
fallbackTitleFromUserText(userText);
await persistTitle(title);
return streamTitle(title);
} finally {
inflightTitleByKey.delete(key);
}
},
};
}
@ -287,10 +422,19 @@ function ThreadHistoryProvider({
if (!remoteId) {
return { messages: [] };
}
const msgs = await db.messages
.where("threadId")
.equals(remoteId)
.sortBy("createdAt");
const roleOrder: Record<string, number> = {
system: 0,
user: 1,
assistant: 2,
};
const msgs = await db.messages.where("threadId").equals(remoteId).toArray();
msgs.sort((a, b) => {
if (a.createdAt !== b.createdAt) return a.createdAt - b.createdAt;
const aOrder = roleOrder[a.role] ?? 99;
const bOrder = roleOrder[b.role] ?? 99;
if (aOrder !== bOrder) return aOrder - bOrder;
return a.id < b.id ? -1 : a.id > b.id ? 1 : 0;
});
return ExportedMessageRepository.fromArray(msgs.map(toThreadMessage));
},
@ -301,13 +445,18 @@ function ThreadHistoryProvider({
? JSON.parse(JSON.stringify(message.content))
: [];
const custom = message.metadata?.custom;
const existing = await db.messages.get(message.id);
const createdAt =
existing?.createdAt ??
message.createdAt?.getTime?.() ??
Date.now();
await db.messages.put({
id: message.id,
threadId: remoteId,
role: message.role,
content,
...(custom && Object.keys(custom).length > 0 && { metadata: custom }),
createdAt: message.createdAt?.getTime() ?? Date.now(),
createdAt,
});
},
}),

View file

@ -52,7 +52,9 @@ export function RegisterCompareHandle({
}
const currentHandles = handlesRef.current;
currentHandles[name] = {
append: (content) => aui.thread().append({ role: "user", content }),
// fixes occasional reorder on reload.
append: (content) =>
aui.thread().append({ role: "user", content, createdAt: new Date() } as never),
cancel: () => aui.thread().cancelRun(),
isRunning: () => aui.thread().getState().isRunning,
};

View file

@ -6,16 +6,44 @@ import {
type InferenceParams,
} from "../types/runtime";
const AUTO_TITLE_KEY = "unsloth_chat_auto_title";
function canUseStorage(): boolean {
return typeof window !== "undefined";
}
function loadBool(key: string, fallback: boolean): boolean {
if (!canUseStorage()) return fallback;
try {
const raw = localStorage.getItem(key);
if (raw === null) return fallback;
return raw === "true";
} catch {
return fallback;
}
}
function saveBool(key: string, value: boolean): void {
if (!canUseStorage()) return;
try {
localStorage.setItem(key, value ? "true" : "false");
} catch {
// ignore
}
}
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;
setCheckpoint: (modelId: string) => void;
clearCheckpoint: () => void;
@ -25,20 +53,26 @@ 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) =>
setThreadRunning: (threadId, running) =>
set((state) => {
const next = { ...state.warmingByThreadId };
if (warming) {
const next = { ...state.runningByThreadId };
if (running) {
next[threadId] = true;
} else {
delete next[threadId];
}
return { warmingByThreadId: next };
return { runningByThreadId: next };
}),
setAutoTitle: (autoTitle) =>
set(() => {
saveBool(AUTO_TITLE_KEY, autoTitle);
return { autoTitle };
}),
setModelsError: (modelsError) => set({ modelsError }),
setCheckpoint: (modelId) =>

View file

@ -62,11 +62,13 @@ export function ThreadSidebar({
onSelect,
onNewThread,
onNewCompare,
showCompare,
}: {
view: ChatView;
onSelect: (view: ChatView) => void;
onNewThread: () => void;
onNewCompare: () => void;
showCompare: boolean;
}) {
const allThreads = useLiveQuery(
() => db.threads.orderBy("createdAt").reverse().toArray(),
@ -112,12 +114,14 @@ export function ThreadSidebar({
<span>New Chat</span>
</SidebarMenuButton>
</SidebarMenuItem>
<SidebarMenuItem>
<SidebarMenuButton onClick={onNewCompare}>
<HugeiconsIcon icon={ColumnInsertIcon} />
<span>Compare</span>
</SidebarMenuButton>
</SidebarMenuItem>
{showCompare ? (
<SidebarMenuItem>
<SidebarMenuButton onClick={onNewCompare}>
<HugeiconsIcon icon={ColumnInsertIcon} />
<span>Compare</span>
</SidebarMenuButton>
</SidebarMenuItem>
) : null}
</SidebarMenu>
</SidebarGroupContent>
</SidebarGroup>

View file

@ -63,6 +63,7 @@ export interface OpenAIChatCompletionsRequest {
top_k: number;
repetition_penalty: number;
image_base64?: string;
use_adapter?: boolean | string | null;
}
export interface OpenAIChatDelta {