Merge pull request #207 from unslothai/feature/attachment-restore

fix(chat): persist + hydrate user attachments in IndexedDB history
This commit is contained in:
Wasim Yousef Said 2026-02-21 22:09:53 -08:00 committed by GitHub
commit 4f2e434bc3
12 changed files with 684 additions and 140 deletions

View file

@ -107,6 +107,23 @@ class InferenceBackend:
# Apply inference optimization
FastVisionModel.for_inference(model)
# FastVisionModel may return a raw tokenizer (e.g. GemmaTokenizerFast)
# instead of a proper Processor for some models (e.g. Gemma-3).
# In that case, load the real processor from the base model.
from transformers import ProcessorMixin
if not (isinstance(processor, ProcessorMixin) or hasattr(processor, "image_processor")):
processor_source = config.base_model if config.is_lora else config.identifier
logger.warning(
f"FastVisionModel returned {type(processor).__name__} (no image_processor) "
f"for '{model_name}' — loading proper processor from '{processor_source}'"
)
from transformers import AutoProcessor
processor = AutoProcessor.from_pretrained(
processor_source,
token=hf_token if hf_token and hf_token.strip() else None,
)
logger.info(f"Loaded {type(processor).__name__} from {processor_source}")
self.models[model_name]["model"] = model
self.models[model_name]["tokenizer"] = processor
self.models[model_name]["processor"] = processor
@ -574,59 +591,78 @@ class InferenceBackend:
model_info = self.models[self.active_model_name]
is_vision = model_info.get("is_vision", False)
tokenizer = model_info.get("tokenizer") or model_info.get("processor")
# Unwrap processor → raw tokenizer for VLMs on the text path
tokenizer = getattr(tokenizer, "tokenizer", tokenizer)
top_k = self._normalize_top_k(top_k)
if is_vision:
# Vision model generation
yield from self._generate_vision_response(
messages, system_prompt, image,
temperature, top_p, top_k, min_p, max_new_tokens, repetition_penalty,
cancel_event=cancel_event,
if is_vision and image:
# Vision model generation (only when an image is actually provided)
# Check that the stored processor can actually handle images.
# FastVisionModel may return a raw tokenizer (e.g. GemmaTokenizerFast)
# instead of a proper ProcessorMixin for some models (e.g. Gemma-3).
from transformers import ProcessorMixin
processor = model_info.get("processor")
has_image_processing = (
processor is not None
and (isinstance(processor, ProcessorMixin) or hasattr(processor, "image_processor"))
)
else:
# Text model: Use training pipeline approach
# Messages are already in ChatML format from eval.py
# Step 1: Apply get_chat_template if model is in mapper
try:
from utils.datasets import MODEL_TO_TEMPLATE_MAPPER, get_tokenizer_chat_template
model_name_lower = self.active_model_name.lower()
# Check if model has a registered template
if model_name_lower in MODEL_TO_TEMPLATE_MAPPER:
template_name = MODEL_TO_TEMPLATE_MAPPER[model_name_lower]
logger.info(f"Applying chat template '{template_name}' for {self.active_model_name}")
# This modifies the tokenizer with the correct template
tokenizer = get_chat_template(
tokenizer,
self.active_model_name
)
else:
logger.info(f"No registered template for {self.active_model_name}, using tokenizer default")
except Exception as e:
logger.warning(f"Could not apply get_chat_template: {e}")
# Step 2: Format with tokenizer.apply_chat_template()
try:
formatted_prompt = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True
if has_image_processing:
yield from self._generate_vision_response(
messages, system_prompt, image,
temperature, top_p, top_k, min_p, max_new_tokens, repetition_penalty,
cancel_event=cancel_event,
)
return
else:
logger.warning(
f"Model '{self.active_model_name}' is marked as vision but its processor "
f"({type(processor).__name__}) has no image_processor — "
f"falling back to text-only generation (image will be ignored)."
)
logger.debug(f"Formatted prompt: {formatted_prompt[:200]}...")
except Exception as e:
logger.error(f"Error applying chat template: {e}")
# Fallback to manual formatting
formatted_prompt = self.format_chat_prompt(messages, system_prompt)
# Step 3: Generate
yield from self.generate_stream(
formatted_prompt, temperature, top_p, top_k, min_p, max_new_tokens, repetition_penalty,
cancel_event=cancel_event,
_adapter_state=_adapter_state,
# Text path: Use training pipeline approach
# Messages are already in ChatML format from eval.py
# Step 1: Apply get_chat_template if model is in mapper
try:
from utils.datasets import MODEL_TO_TEMPLATE_MAPPER, get_tokenizer_chat_template
model_name_lower = self.active_model_name.lower()
# Check if model has a registered template
if model_name_lower in MODEL_TO_TEMPLATE_MAPPER:
template_name = MODEL_TO_TEMPLATE_MAPPER[model_name_lower]
logger.info(f"Applying chat template '{template_name}' for {self.active_model_name}")
# This modifies the tokenizer with the correct template
tokenizer = get_chat_template(
tokenizer,
chat_template=template_name,
)
else:
logger.info(f"No registered template for {self.active_model_name}, using tokenizer default")
except Exception as e:
logger.warning(f"Could not apply get_chat_template: {e}")
# Step 2: Format with tokenizer.apply_chat_template()
try:
formatted_prompt = tokenizer.apply_chat_template(
messages,
tokenize=False,
add_generation_prompt=True
)
logger.debug(f"Formatted prompt: {formatted_prompt[:200]}...")
except Exception as e:
logger.error(f"Error applying chat template: {e}")
# Fallback to manual formatting
formatted_prompt = self.format_chat_prompt(messages, system_prompt)
# Step 3: Generate
yield from self.generate_stream(
formatted_prompt, temperature, top_p, top_k, min_p, max_new_tokens, repetition_penalty,
cancel_event=cancel_event,
_adapter_state=_adapter_state,
)
def _generate_vision_response(self, messages, system_prompt, image,
temperature, top_p, top_k, min_p, max_new_tokens,
@ -635,6 +671,9 @@ class InferenceBackend:
model_info = self.models[self.active_model_name]
model = model_info["model"]
processor = model_info["processor"]
# FastVisionModel may return a raw tokenizer (e.g. GemmaTokenizerFast)
# instead of a Processor for some models. Safe unwrap for tokenize-only ops.
raw_tokenizer = getattr(processor, "tokenizer", processor)
# Extract user message
user_message = ""
@ -658,7 +697,7 @@ class InferenceBackend:
}
]
input_text = processor.apply_chat_template(vision_messages, add_generation_prompt=True)
input_text = processor.apply_chat_template(vision_messages, add_generation_prompt=True, tokenize=False)
inputs = processor(
image,
input_text,
@ -668,7 +707,7 @@ class InferenceBackend:
else:
# Text-only for vision model
formatted_prompt = self.format_chat_prompt(messages, system_prompt)
inputs = processor.tokenizer(formatted_prompt, return_tensors="pt").to(self.device)
inputs = raw_tokenizer(formatted_prompt, return_tensors="pt").to(self.device)
# Stream with TextIteratorStreamer + background thread
try:
@ -676,7 +715,7 @@ class InferenceBackend:
import threading
streamer = TextIteratorStreamer(
processor.tokenizer,
raw_tokenizer,
skip_prompt=True,
skip_special_tokens=True,
timeout=0.2,
@ -766,7 +805,11 @@ class InferenceBackend:
model_info = self.models[self.active_model_name]
model = model_info["model"]
# For VLMs the stored "tokenizer" is actually the processor.
# Unwrap to get the real tokenizer so TextIteratorStreamer's
# skip_prompt / skip_special_tokens work correctly.
tokenizer = model_info["tokenizer"]
tokenizer = getattr(tokenizer, "tokenizer", tokenizer)
try:
inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
@ -876,6 +919,7 @@ class InferenceBackend:
chat_template_info = self.models[self.active_model_name].get("chat_template_info", {})
tokenizer = self.models[self.active_model_name]["tokenizer"]
tokenizer = getattr(tokenizer, "tokenizer", tokenizer)
chat_messages = []
@ -1116,24 +1160,13 @@ class InferenceBackend:
return img
def _clean_generated_text(self, text: str) -> str:
import re
text = re.sub(r'<\|start_header_id\|>.*?<\|end_header_id\|>', '', text)
text = re.sub(r'<\|eot_id\|>', '', text)
text = re.sub(r'<\|begin_of_text\|>', '', text)
text = re.sub(r'\[INST\].*?\[/INST\]', '', text)
text = re.sub(r'<s>|</s>', '', text)
# Clean ChatML tokens (used by Qwen2-VL and similar models)
text = re.sub(r'<\|im_start\|>.*?<\|im_end\|>', '', text)
text = re.sub(r'<\|im_end\|>', '', text)
text = re.sub(r'<\|im_start\|>', '', text)
text = re.sub(r'^\s*(assistant|user|system):\s*', '', text, flags=re.IGNORECASE)
text = text.strip()
return text
"""Strip leaked special tokens using the tokenizer's own token list."""
tokenizer = self.models.get(self.active_model_name, {}).get("tokenizer")
if tokenizer:
for token in getattr(tokenizer, "all_special_tokens", []):
if token in text:
text = text.replace(token, "")
return text.strip()
def _load_chat_template_info(self, model_name: str):
if model_name not in self.models or not self.models[model_name].get("tokenizer"):

View file

@ -1,7 +1,20 @@
import { Input } from "@/components/ui/input";
import { Spinner } from "@/components/ui/spinner";
import { useDebouncedValue, useHfModelSearch, useInfiniteScroll } from "@/hooks";
import {
Tooltip,
TooltipContent,
TooltipTrigger,
} from "@/components/ui/tooltip";
import {
useDebouncedValue,
useGpuInfo,
useHfModelSearch,
useInfiniteScroll,
useRecommendedModelVram,
} from "@/hooks";
import { cn, formatCompact } from "@/lib/utils";
import type { VramFitStatus } from "@/lib/vram";
import { checkVramFit, estimateLoadingVram } from "@/lib/vram";
import { Search01Icon } from "@hugeicons/core-free-icons";
import { HugeiconsIcon } from "@hugeicons/react";
import { useMemo, useState, type ReactNode } from "react";
@ -28,27 +41,77 @@ function ModelRow({
meta,
selected,
onClick,
vramStatus,
vramEst,
gpuGb,
}: {
label: string;
meta?: string;
selected?: boolean;
onClick: () => void;
vramStatus?: VramFitStatus | null;
vramEst?: number;
gpuGb?: number;
}) {
return (
const exceeds = vramStatus === "exceeds";
const showVramTooltip =
vramEst != null && vramEst > 0 && gpuGb != null && gpuGb > 0;
const vramTooltipText =
showVramTooltip && vramStatus
? exceeds
? `Needs ~${vramEst}GB VRAM (GPU: ${gpuGb}GB)`
: vramStatus === "tight"
? `~${vramEst}GB VRAM (tight fit on ${gpuGb}GB)`
: `~${vramEst}GB VRAM`
: null;
const content = (
<button
type="button"
onClick={onClick}
className={cn(
"flex w-full items-center justify-between gap-2 rounded-md px-2.5 py-1.5 text-left text-sm transition-colors hover:bg-accent",
selected && "bg-accent/60",
exceeds && "opacity-50",
)}
>
<span className="min-w-0 flex-1 truncate">{label}</span>
{meta ? (
<span className="shrink-0 text-[10px] text-muted-foreground">{meta}</span>
) : null}
<span
className={cn(
"min-w-0 flex-1 truncate",
exceeds && "line-through decoration-muted-foreground/50",
)}
>
{label}
</span>
<span className="flex items-center gap-1.5 shrink-0">
{vramStatus === "exceeds" && (
<span className="text-[9px] font-medium text-red-400">OOM</span>
)}
{vramStatus === "tight" && (
<span className="text-[9px] font-medium text-amber-400">TIGHT</span>
)}
{vramStatus === "fits" && (
<span className="text-[9px] font-medium text-emerald-500/90">FIT</span>
)}
{meta ? (
<span className="text-[10px] text-muted-foreground">{meta}</span>
) : null}
</span>
</button>
);
if (vramTooltipText) {
return (
<Tooltip>
<TooltipTrigger asChild>{content}</TooltipTrigger>
<TooltipContent side="left" className="max-w-xs break-all">
{label}
<span className="block text-[10px] mt-1">{vramTooltipText}</span>
</TooltipContent>
</Tooltip>
);
}
return content;
}
export function HubModelPicker({
@ -60,6 +123,7 @@ export function HubModelPicker({
value?: string;
onSelect: (id: string, meta: ModelSelectorChangeMeta) => void;
}) {
const gpu = useGpuInfo();
const [query, setQuery] = useState("");
const debouncedQuery = useDebouncedValue(query);
const { results, isLoading, isLoadingMore, fetchMore } = useHfModelSearch(
@ -71,6 +135,9 @@ export function HubModelPicker({
[models, value],
);
const { paramCountById: recommendedParamCountById } =
useRecommendedModelVram(recommendedIds);
const showHfSection = debouncedQuery.trim().length > 0;
const recommendedSet = useMemo(() => new Set(recommendedIds), [recommendedIds]);
@ -94,6 +161,49 @@ export function HubModelPicker({
[results],
);
const vramMap = useMemo(() => {
const map = new Map<
string,
{ est: number; status: VramFitStatus | null; detail: string | null }
>();
for (const r of results) {
const detail = r.totalParams
? formatCompact(r.totalParams)
: r.downloads != null
? `${formatCompact(r.downloads)}`
: null;
if (r.totalParams) {
const est = estimateLoadingVram(r.totalParams, "qlora");
const status = gpu.available
? checkVramFit(est, gpu.memoryTotalGb)
: null;
map.set(r.id, { est, status, detail });
} else {
map.set(r.id, { est: 0, status: null, detail });
}
}
return map;
}, [results, gpu]);
const recommendedVramMap = useMemo(() => {
const map = new Map<
string,
{ est: number; status: VramFitStatus | null; detail: string | null }
>();
for (const id of recommendedIds) {
const totalParams = recommendedParamCountById.get(id);
if (totalParams) {
const est = estimateLoadingVram(totalParams, "qlora");
const status = gpu.available
? checkVramFit(est, gpu.memoryTotalGb)
: null;
const detail = formatCompact(totalParams);
map.set(id, { est, status, detail });
}
}
return map;
}, [recommendedIds, recommendedParamCountById, gpu]);
const { scrollRef, sentinelRef } = useInfiniteScroll(fetchMore, results.length);
return (
@ -124,14 +234,23 @@ export function HubModelPicker({
No default models.
</div>
) : (
recommendedIds.map((id) => (
<ModelRow
key={id}
label={id}
selected={value === id}
onClick={() => onSelect(id, { source: "hub", isLora: false })}
/>
))
recommendedIds.map((id) => {
const vram = recommendedVramMap.get(id);
return (
<ModelRow
key={id}
label={id}
meta={vram?.detail ?? undefined}
selected={value === id}
onClick={() =>
onSelect(id, { source: "hub", isLora: false })
}
vramStatus={vram?.status ?? null}
vramEst={vram?.est}
gpuGb={gpu.available ? gpu.memoryTotalGb : undefined}
/>
);
})
)}
</>
) : null}
@ -144,15 +263,23 @@ export function HubModelPicker({
No matching models.
</div>
) : (
hfIds.map((id) => (
<ModelRow
key={id}
label={id}
meta={metricsById.get(id)}
selected={value === id}
onClick={() => onSelect(id, { source: "hub", isLora: false })}
/>
))
hfIds.map((id) => {
const vram = vramMap.get(id);
return (
<ModelRow
key={id}
label={id}
meta={metricsById.get(id)}
selected={value === id}
onClick={() =>
onSelect(id, { source: "hub", isLora: false })
}
vramStatus={vram?.status ?? null}
vramEst={vram?.est}
gpuGb={gpu.available ? gpu.memoryTotalGb : undefined}
/>
);
})
)}
<div ref={sentinelRef} className="h-px" />
{isLoadingMore ? (

View file

@ -57,7 +57,7 @@ export function Navbar() {
};
return (
<header className="top-0 z-40 h-16 w-full">
<header className="relative top-0 z-40 h-16 w-full">
<div className="mx-auto flex h-full max-w-7xl items-center justify-between px-4 sm:px-6">
{/* Left: logo */}
<div

View file

@ -65,6 +65,15 @@ function findLatestUserImageBase64(messages: RunMessages): string | undefined {
continue;
}
// Image in message.content (e.g. compare view appends content with image parts)
for (const part of message.content ?? []) {
if (part.type === "image" && "image" in part) {
const encoded = extractImageBase64(part.image);
if (encoded) return encoded;
}
}
// Image in message.attachments (e.g. chat composer)
if ("attachments" in message && (message.attachments?.length ?? 0) > 0) {
for (const attachment of message.attachments ?? []) {
for (const part of attachment.content ?? []) {

View file

@ -5,6 +5,7 @@ import {
} from "@/components/assistant-ui/model-selector";
import { Thread } from "@/components/assistant-ui/thread";
import { Button } from "@/components/ui/button";
import { Spinner } from "@/components/ui/spinner";
import { SidebarProvider, SidebarTrigger, useSidebar } from "@/components/ui/sidebar";
import {
Sheet,
@ -283,7 +284,8 @@ export function ChatPage(): ReactElement {
const modelsFromStore = useChatRuntimeStore((state) => state.models);
const lorasFromStore = useChatRuntimeStore((state) => state.loras);
const modelsError = useChatRuntimeStore((state) => state.modelsError);
const { refresh, selectModel, ejectModel } = useChatModelRuntime();
const { refresh, selectModel, ejectModel, loadingModel } =
useChatModelRuntime();
const refreshRef = useRef(refresh);
const selectModelRef = useRef(selectModel);
@ -518,6 +520,17 @@ export function ChatPage(): ReactElement {
contentDataTour="chat-model-selector-popover"
className="max-w-[62vw] sm:max-w-none"
/>
{loadingModel ? (
<div
className="flex items-center gap-1.5 text-muted-foreground"
title={`Loading ${loadingModel.displayName}. This may include downloading.`}
>
<Spinner className="size-3.5 shrink-0" />
<span className="text-xs">
Downloading model
</span>
</div>
) : null}
</div>
{modelsError && (
<div className="ml-2 text-xs text-destructive truncate max-w-[28rem]">

View file

@ -1,4 +1,4 @@
import { useCallback } from "react";
import { useCallback, useState } from "react";
import { toast } from "sonner";
import {
getInferenceStatus,
@ -116,6 +116,11 @@ export function useChatModelRuntime() {
const setCheckpoint = useChatRuntimeStore((state) => state.setCheckpoint);
const clearCheckpoint = useChatRuntimeStore((state) => state.clearCheckpoint);
const [loadingModel, setLoadingModel] = useState<{
id: string;
displayName: string;
} | null>(null);
const refresh = useCallback(async () => {
setModelsError(null);
try {
@ -157,6 +162,7 @@ export function useChatModelRuntime() {
const displayName = model?.name || lora?.name || modelId;
setModelsError(null);
setLoadingModel({ id: modelId, displayName });
try {
async function performLoad(): Promise<void> {
if (params.checkpoint) {
@ -176,19 +182,20 @@ export function useChatModelRuntime() {
await refresh();
}
let description = "Base model selected.";
if (isLora) {
description = "Fine-tuned (LoRA) selected.";
}
const loadPromise = performLoad().finally(() => {
setLoadingModel(null);
});
await toast.promise(performLoad(), {
loading: `Loading ${displayName}`,
await toast.promise(loadPromise, {
loading: "Loading model…",
success: `${displayName} loaded`,
error: (err) =>
err instanceof Error ? err.message : "Failed to load model",
description,
description:
"This may include downloading. Large models can take a while.",
});
} catch (error) {
setLoadingModel(null);
const message =
error instanceof Error ? error.message : "Failed to load model";
setModelsError(message);
@ -227,5 +234,6 @@ export function useChatModelRuntime() {
refresh,
selectModel,
ejectModel,
loadingModel,
};
}

View file

@ -35,6 +35,14 @@ const DEFAULT_SUGGESTIONS = [
"Format a comparison of 3 databases as a markdown table with pros and cons",
];
type TitleResponse = {
choices?: Array<{
message?: {
content?: string;
};
}>;
};
class VisionImageAdapter implements AttachmentAdapter {
accept = "image/jpeg,image/png,image/webp,image/gif";
@ -216,7 +224,7 @@ async function generateTitleWithModel(payload: {
}),
});
const body = (await response.json().catch(() => null)) as any;
const body = (await response.json().catch(() => null)) as TitleResponse | null;
if (!response.ok) return null;
const raw: string | undefined = body?.choices?.[0]?.message?.content;
if (!raw) return null;
@ -233,27 +241,42 @@ function fallbackTitleFromUserText(userText: string): string {
return cleaned.slice(0, max) + (cleaned.length > max ? "..." : "");
}
function cloneContent(content: ThreadMessage["content"]): ThreadMessage["content"] {
return Array.isArray(content)
? JSON.parse(JSON.stringify(content))
: [];
}
function cloneAttachments(
attachments: readonly CompleteAttachment[] | undefined,
): readonly CompleteAttachment[] {
if (!Array.isArray(attachments)) {
return [];
}
return JSON.parse(JSON.stringify(attachments));
}
function toThreadMessage(m: MessageRecord): ThreadMessage {
const base = {
id: m.id,
createdAt: new Date(m.createdAt),
content:
Array.isArray(m.content) && m.content.length > 0
? m.content
: [{ type: "text" as const, text: "" }],
};
const content =
Array.isArray(m.content) && m.content.length > 0
? cloneContent(m.content)
: [{ type: "text" as const, text: "" }];
if (m.role === "user") {
return {
...base,
id: m.id,
createdAt: new Date(m.createdAt),
role: "user" as const,
attachments: [],
content: content as Extract<ThreadMessage, { role: "user" }>["content"],
attachments: cloneAttachments(m.attachments),
metadata: { custom: {} },
};
}
return {
...base,
id: m.id,
createdAt: new Date(m.createdAt),
role: "assistant" as const,
content: content as Extract<ThreadMessage, { role: "assistant" }>["content"],
status: { type: "complete" as const, reason: "unknown" as const },
metadata: {
custom: (m.metadata as Record<string, unknown>) ?? {},
@ -441,9 +464,9 @@ function ThreadHistoryProvider({
async append({ message }: ExportedMessageRepositoryItem) {
const { remoteId } = await aui.threadListItem().initialize();
const content = Array.isArray(message.content)
? JSON.parse(JSON.stringify(message.content))
: [];
const content = cloneContent(message.content);
const attachments =
message.role === "user" ? cloneAttachments(message.attachments) : [];
const custom = message.metadata?.custom;
const existing = await db.messages.get(message.id);
const createdAt =
@ -455,6 +478,7 @@ function ThreadHistoryProvider({
threadId: remoteId,
role: message.role,
content,
...(attachments.length > 0 && { attachments }),
...(custom && Object.keys(custom).length > 0 && { metadata: custom }),
createdAt,
});

View file

@ -1,25 +1,102 @@
import { TooltipIconButton } from "@/components/assistant-ui/tooltip-icon-button";
import { Button } from "@/components/ui/button";
import { useAui } from "@assistant-ui/react";
import { ArrowUpIcon, SquareIcon } from "lucide-react";
import { ArrowUpIcon, MicIcon, PlusIcon, SquareIcon, XIcon } from "lucide-react";
import {
type KeyboardEvent,
type MutableRefObject,
type ReactElement,
type ReactNode,
createContext,
useCallback,
useContext,
useEffect,
useRef,
useState,
} from "react";
export type CompareMessagePart =
| { type: "text"; text: string }
| { type: "image"; image: string };
export interface CompareHandle {
append: (content: { type: "text"; text: string }[]) => void;
append: (content: CompareMessagePart[]) => void;
cancel: () => void;
isRunning: () => boolean;
}
const IMAGE_ACCEPT = "image/jpeg,image/png,image/webp,image/gif";
const MAX_IMAGE_SIZE = 20 * 1024 * 1024;
function fileToBase64DataURL(file: File): Promise<string> {
return new Promise((resolve, reject) => {
const reader = new FileReader();
reader.onload = () => resolve(reader.result as string);
reader.onerror = () => reject(new Error("Failed to read image file"));
reader.readAsDataURL(file);
});
}
function useDictation(
setText: (value: string | ((prev: string) => string)) => void,
) {
const [isDictating, setIsDictating] = useState(false);
const recognitionRef = useRef<SpeechRecognition | null>(null);
const start = useCallback(() => {
const SpeechRecognitionAPI =
typeof window !== "undefined" &&
(window.SpeechRecognition ?? (window as unknown as { webkitSpeechRecognition?: typeof SpeechRecognition }).webkitSpeechRecognition);
if (!SpeechRecognitionAPI) {
return;
}
const recognition = new SpeechRecognitionAPI() as SpeechRecognition;
recognition.continuous = true;
recognition.interimResults = true;
recognition.lang = "en-US";
recognition.onresult = (event: SpeechRecognitionEvent) => {
const last = event.resultIndex;
const result = event.results[last];
if (!result?.isFinal) return;
const transcript = result[0]?.transcript?.trim();
if (transcript) {
setText((prev) => (prev ? `${prev} ${transcript}` : transcript));
}
};
recognition.onerror = () => {
setIsDictating(false);
};
recognition.onend = () => {
setIsDictating(false);
};
recognition.start();
recognitionRef.current = recognition;
setIsDictating(true);
}, [setText]);
const stop = useCallback(() => {
if (recognitionRef.current) {
recognitionRef.current.stop();
recognitionRef.current = null;
}
setIsDictating(false);
}, []);
useEffect(() => {
return () => {
if (recognitionRef.current) {
recognitionRef.current.abort();
}
};
}, []);
const supported =
typeof window !== "undefined" &&
!!(window.SpeechRecognition ?? (window as unknown as { webkitSpeechRecognition?: unknown }).webkitSpeechRecognition);
return { isDictating, start, stop, supported };
}
export type CompareHandles = MutableRefObject<Record<string, CompareHandle>>;
const CompareHandlesContext = createContext<CompareHandles | null>(null);
@ -66,6 +143,37 @@ export function RegisterCompareHandle({
return null;
}
type PendingImage = { id: string; file: File };
function PendingImageThumb({
file,
onRemove,
}: {
file: File;
onRemove: () => void;
}): ReactElement {
const [src, setSrc] = useState<string | null>(null);
useEffect(() => {
const url = URL.createObjectURL(file);
setSrc(url);
return () => URL.revokeObjectURL(url);
}, [file]);
if (!src) return <div className="size-14 animate-pulse rounded-[14px] bg-muted" />;
return (
<div className="relative size-14 shrink-0 overflow-hidden rounded-[14px] border border-foreground/20 bg-muted">
<img src={src} alt={file.name} className="h-full w-full object-cover" />
<button
type="button"
onClick={onRemove}
className="absolute top-1 right-1 flex size-5 items-center justify-center rounded-full bg-white text-muted-foreground shadow-sm hover:bg-destructive hover:text-destructive-foreground"
aria-label="Remove attachment"
>
<XIcon className="size-3" />
</button>
</div>
);
}
export function SharedComposer({
handlesRef,
}: {
@ -73,7 +181,14 @@ export function SharedComposer({
}): ReactElement {
const [text, setText] = useState("");
const [running, setRunning] = useState(false);
const [pendingImages, setPendingImages] = useState<PendingImage[]>([]);
const [dragging, setDragging] = useState(false);
const textareaRef = useRef<HTMLTextAreaElement>(null);
const fileInputRef = useRef<HTMLInputElement>(null);
const { isDictating, start: startDictation, stop: stopDictation, supported: dictationSupported } = useDictation(
setText,
);
useEffect(() => {
const id = setInterval(() => {
@ -84,23 +199,50 @@ export function SharedComposer({
return () => clearInterval(id);
}, [handlesRef]);
function send() {
const msg = text.trim();
if (!msg) {
return;
const addFiles = useCallback((files: FileList | null) => {
if (!files?.length) return;
const next: PendingImage[] = [];
for (let i = 0; i < files.length; i++) {
const file = files[i];
if (!file?.type.match(/^image\/(jpeg|png|webp|gif)$/i)) continue;
if (file.size > MAX_IMAGE_SIZE) continue;
next.push({ id: crypto.randomUUID(), file });
}
setPendingImages((prev) => [...prev, ...next]);
}, []);
const removePendingImage = useCallback((id: string) => {
setPendingImages((prev) => prev.filter((p) => p.id !== id));
}, []);
async function send() {
const msg = text.trim();
if (!msg && pendingImages.length === 0) return;
const content: CompareMessagePart[] = [];
for (const { file } of pendingImages) {
try {
const image = await fileToBase64DataURL(file);
content.push({ type: "image", image });
} catch {
// skip failed image
}
}
if (msg) {
content.push({ type: "text", text: msg });
}
if (content.length === 0) return;
const content: { type: "text"; text: string }[] = [
{ type: "text", text: msg },
];
for (const handle of Object.values(handlesRef.current)) {
handle.append(content);
}
setText("");
setPendingImages([]);
textareaRef.current?.focus();
}
function stop() {
if (isDictating) stopDictation();
for (const handle of Object.values(handlesRef.current)) {
handle.cancel();
}
@ -115,8 +257,33 @@ export function SharedComposer({
}
}
const canSend = (text.trim().length > 0 || pendingImages.length > 0) && !running;
return (
<div className="shadow-border ring-1 ring-border relative flex w-full flex-col rounded-2xl bg-background px-1 pt-2 transition-shadow">
<div
className={`shadow-border ring-1 ring-border relative flex w-full flex-col rounded-2xl bg-background px-1 pt-2 transition-shadow outline-none ${dragging ? "ring-ring bg-accent/50" : ""}`}
onDragOver={(e) => {
e.preventDefault();
setDragging(true);
}}
onDragLeave={() => setDragging(false)}
onDrop={(e) => {
e.preventDefault();
setDragging(false);
addFiles(e.dataTransfer.files);
}}
>
{pendingImages.length > 0 && (
<div className="mb-2 flex w-full flex-row flex-wrap items-center gap-2 px-1.5 pt-0.5 pb-1">
{pendingImages.map(({ id, file }) => (
<PendingImageThumb
key={id}
file={file}
onRemove={() => removePendingImage(id)}
/>
))}
</div>
)}
<textarea
ref={textareaRef}
value={text}
@ -126,30 +293,85 @@ export function SharedComposer({
className="mb-1 max-h-32 min-h-14 w-full resize-none bg-transparent px-4 pt-2 pb-3 text-sm outline-none placeholder:text-muted-foreground"
rows={1}
/>
<div className="relative mx-2 mb-2 flex items-center justify-end">
{running ? (
<Button
type="button"
variant="default"
size="icon"
className="size-8 rounded-full"
onClick={stop}
>
<SquareIcon className="size-3 fill-current" />
</Button>
) : (
<div className="relative mx-2 mb-2 flex items-center justify-between">
<div className="flex items-center gap-1">
<input
ref={fileInputRef}
type="file"
accept={IMAGE_ACCEPT}
multiple
className="hidden"
onChange={(e) => {
addFiles(e.target.files);
e.target.value = "";
}}
/>
<TooltipIconButton
tooltip="Send message"
tooltip="Add attachment"
side="bottom"
variant="default"
variant="ghost"
size="icon"
className="size-8 rounded-full"
onClick={send}
disabled={!text.trim()}
className="size-8 rounded-full text-muted-foreground hover:bg-muted-foreground/15"
onClick={() => fileInputRef.current?.click()}
aria-label="Add attachment"
>
<ArrowUpIcon className="size-4" />
<PlusIcon className="size-5 stroke-[1.5px]" />
</TooltipIconButton>
)}
</div>
<div className="flex items-center gap-1">
{dictationSupported && (
<>
{!isDictating ? (
<TooltipIconButton
tooltip="Dictate"
side="bottom"
variant="ghost"
size="icon"
className="size-8 rounded-full text-muted-foreground hover:bg-muted-foreground/15"
onClick={startDictation}
aria-label="Dictate"
>
<MicIcon className="size-4" />
</TooltipIconButton>
) : (
<TooltipIconButton
tooltip="Stop dictation"
side="bottom"
variant="ghost"
size="icon"
className="size-8 rounded-full text-destructive"
onClick={stopDictation}
aria-label="Stop dictation"
>
<SquareIcon className="size-3 animate-pulse fill-current" />
</TooltipIconButton>
)}
</>
)}
{running ? (
<Button
type="button"
variant="default"
size="icon"
className="size-8 rounded-full"
onClick={stop}
>
<SquareIcon className="size-3 fill-current" />
</Button>
) : (
<TooltipIconButton
tooltip="Send message"
side="bottom"
variant="default"
size="icon"
className="size-8 rounded-full"
onClick={send}
disabled={!canSend}
>
<ArrowUpIcon className="size-4" />
</TooltipIconButton>
)}
</div>
</div>
</div>
);

View file

@ -18,6 +18,7 @@ export interface MessageRecord {
threadId: string;
role: import("@assistant-ui/react").ThreadMessage["role"];
content: import("@assistant-ui/react").ThreadMessage["content"];
attachments?: import("@assistant-ui/react").ThreadMessage["attachments"];
metadata?: Record<string, unknown>;
createdAt: number;
}

View file

@ -3,6 +3,7 @@ export { useGpuInfo } from "./use-gpu-info";
export { useGpuUtilization } from "./use-gpu-utilization";
export { useHardwareInfo } from "./use-hardware-info";
export { useHfModelSearch } from "./use-hf-model-search";
export { useRecommendedModelVram } from "./use-recommended-model-vram";
export { useHfDatasetSearch } from "./use-hf-dataset-search";
export { useHfDatasetSplits } from "./use-hf-dataset-splits";
export { useHfTokenValidation } from "./use-hf-token-validation";

View file

@ -0,0 +1,57 @@
import { modelInfo } from "@huggingface/hub";
import { useEffect, useState } from "react";
/**
* Fetches Hugging Face model info (safetensors total param count) for a list of
* model IDs. Used to show VRAM fit (FIT / TIGHT / OOM) for recommended/default
* models in the chat model dropdown.
*/
export function useRecommendedModelVram(ids: string[]) {
const [paramCountById, setParamCountById] = useState<
Map<string, number>
>(new Map());
const [isLoading, setIsLoading] = useState(false);
const stableKey = [...ids].filter(Boolean).sort().join(",");
useEffect(() => {
const stableIds = stableKey ? stableKey.split(",") : [];
if (stableIds.length === 0) {
setParamCountById(new Map());
setIsLoading(false);
return;
}
let canceled = false;
void (async () => {
setIsLoading(true);
const next = new Map<string, number>();
await Promise.all(
stableIds.map(async (id) => {
if (canceled) return;
try {
const info = await modelInfo({
name: id,
additionalFields: ["safetensors"],
});
const raw = info as { safetensors?: { total?: number } };
const total = raw.safetensors?.total;
if (typeof total === "number" && total > 0) {
next.set(id, total);
}
} catch {
// Model not on HF or no safetensors; skip
}
}),
);
if (!canceled) {
setParamCountById(next);
setIsLoading(false);
}
})();
return () => {
canceled = true;
};
}, [stableKey]);
return { paramCountById, isLoading };
}

View file

@ -0,0 +1,49 @@
/**
* Minimal Web Speech API (Speech Recognition) types for browsers that support it.
* Full types: @types/dom-speech-recognition
*/
interface SpeechRecognitionResultList {
readonly length: number;
item(index: number): SpeechRecognitionResult;
[index: number]: SpeechRecognitionResult;
}
interface SpeechRecognitionResult {
readonly length: number;
readonly isFinal: boolean;
item(index: number): SpeechRecognitionAlternative;
[index: number]: SpeechRecognitionAlternative;
}
interface SpeechRecognitionAlternative {
readonly transcript: string;
readonly confidence: number;
}
interface SpeechRecognitionEvent extends Event {
readonly resultIndex: number;
readonly results: SpeechRecognitionResultList;
}
interface SpeechRecognition extends EventTarget {
continuous: boolean;
interimResults: boolean;
lang: string;
onresult: ((event: SpeechRecognitionEvent) => void) | null;
onerror: ((event: Event) => void) | null;
onend: (() => void) | null;
start(): void;
stop(): void;
abort(): void;
}
interface SpeechRecognitionConstructor {
new (): SpeechRecognition;
}
interface Window {
SpeechRecognition?: SpeechRecognitionConstructor;
webkitSpeechRecognition?: SpeechRecognitionConstructor;
}
declare var SpeechRecognition: SpeechRecognitionConstructor | undefined;