Merge branch 'nightly' into feature/canvas-lab

This commit is contained in:
Shine1i 2026-02-23 21:54:35 +01:00
commit 8739a01f56
76 changed files with 629 additions and 382 deletions

View file

@ -26,25 +26,22 @@ import {
type FC,
type PropsWithChildren,
useEffect,
useMemo,
useState,
} from "react";
import { useShallow } from "zustand/shallow";
const useFileSrc = (file: File | undefined): string | undefined => {
const objectUrl = useMemo(
() => (file ? URL.createObjectURL(file) : undefined),
[file],
);
const [objectUrl, setObjectUrl] = useState<string | undefined>(undefined);
useEffect(() => {
if (!objectUrl) {
return undefined;
if (!file) {
setObjectUrl(undefined);
return;
}
return () => {
URL.revokeObjectURL(objectUrl);
};
}, [objectUrl]);
const url = URL.createObjectURL(file);
setObjectUrl(url);
return () => URL.revokeObjectURL(url);
}, [file]);
return objectUrl;
};

View file

@ -151,12 +151,9 @@ export function HubModelPicker({
const metricsById = useMemo(
() =>
new Map(
results.map((result) => [
result.id,
result.totalParams
? formatCompact(result.totalParams)
: `${formatCompact(result.downloads)}`,
]),
results
.filter((result) => result.totalParams)
.map((result) => [result.id, formatCompact(result.totalParams!)]),
),
[results],
);
@ -167,11 +164,7 @@ export function HubModelPicker({
{ 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;
const detail = r.totalParams ? formatCompact(r.totalParams) : null;
if (r.totalParams) {
const est = estimateLoadingVram(r.totalParams, "qlora");
const status = gpu.available

View file

@ -218,9 +218,9 @@ function InlineSidebar({
className={cn(
"bg-sidebar text-sidebar-foreground h-full overflow-hidden rounded-2xl corner-squircle transition-[width] duration-200 ease-linear",
!collapsed &&
(side === "left"
? "border-r border-0 border-sidebar-border"
: "border-l border-0 border-sidebar-border"),
(side === "left"
? "border-r border-0 border-sidebar-border"
: "border-l border-0 border-sidebar-border"),
collapsed ? "w-0" : "w-(--sidebar-width)",
)}
>
@ -301,9 +301,18 @@ export function ChatPage(): ReactElement {
const handleCheckpointChange = useCallback(
(value: string, meta?: { isLora: boolean }) => {
void selectModel({ id: value, isLora: meta?.isLora });
const currentCheckpoint =
useChatRuntimeStore.getState().params.checkpoint;
if (!value || value === currentCheckpoint) return;
setView({ mode: "single", newThreadNonce: crypto.randomUUID() });
void (async () => {
if (currentCheckpoint) {
await ejectModel();
}
await selectModel({ id: value, isLora: meta?.isLora });
})();
},
[selectModel],
[selectModel, ejectModel],
);
const handleEject = useCallback(() => {
void ejectModel();
@ -349,6 +358,41 @@ export function ChatPage(): ReactElement {
setViewBeforeCompare(null);
}, [viewBeforeCompare]);
const handleThreadSelect = useCallback(
(nextView: ChatView) => {
setView(nextView);
const threadId =
nextView.mode === "single" ? nextView.threadId : undefined;
const pairId =
nextView.mode === "compare" ? nextView.pairId : undefined;
void (async () => {
let thread: import("./types").ThreadRecord | undefined;
if (threadId) {
thread = await db.threads.get(threadId);
} else if (pairId) {
thread = await db.threads
.where("pairId")
.equals(pairId)
.first();
}
const threadModelId = thread?.modelId;
if (!threadModelId) return;
const currentCheckpoint =
useChatRuntimeStore.getState().params.checkpoint;
if (threadModelId === currentCheckpoint) return;
if (currentCheckpoint) {
await ejectModel();
}
await selectModel({ id: threadModelId });
})();
},
[ejectModel, selectModel],
);
const models = useMemo<ModelOption[]>(
() =>
modelsFromStore.map((model) => ({
@ -475,99 +519,99 @@ export function ChatPage(): ReactElement {
return (
<div className="h-[calc(100dvh-4rem)] bg-background overflow-hidden">
<GuidedTour {...tour.tourProps} />
<SidebarProvider
defaultOpen={true}
open={sidebarOpen}
onOpenChange={setSidebarOpen}
className="!min-h-0 h-full w-full max-w-7xl mx-auto px-2 sm:px-4"
style={
{
"--sidebar-width": "14rem",
"--sidebar-width-icon": "3rem",
} as CSSProperties
}
>
<InlineSidebar>
<ThreadSidebar
view={view}
onSelect={setView}
onNewThread={handleNewThread}
onNewCompare={handleNewCompare}
showCompare={canCompare}
/>
</InlineSidebar>
<GuidedTour {...tour.tourProps} />
<SidebarProvider
defaultOpen={true}
open={sidebarOpen}
onOpenChange={setSidebarOpen}
className="!min-h-0 h-full w-full max-w-7xl mx-auto px-2 sm:px-4"
style={
{
"--sidebar-width": "14rem",
"--sidebar-width-icon": "3rem",
} as CSSProperties
}
>
<InlineSidebar>
<ThreadSidebar
view={view}
onSelect={handleThreadSelect}
onNewThread={handleNewThread}
onNewCompare={handleNewCompare}
showCompare={canCompare}
/>
</InlineSidebar>
<div className="flex min-h-0 min-w-0 flex-1 flex-col">
<div className="flex h-11 shrink-0 items-center px-1.5 sm:px-2">
<div className="flex items-center gap-1">
<SidebarTrigger />
<TopBarActions
onNewThread={handleNewThread}
onNewCompare={handleNewCompare}
showCompare={canCompare}
/>
<ModelSelector
models={models}
loraModels={loraModels}
value={inferenceParams.checkpoint}
onValueChange={handleCheckpointChange}
onEject={handleEject}
variant="ghost"
open={modelSelectorOpen}
onOpenChange={handleModelSelectorOpenChange}
triggerDataTour="chat-model-selector"
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]">
{modelsError}
<div className="flex min-h-0 min-w-0 flex-1 flex-col">
<div className="flex h-11 shrink-0 items-center px-1.5 sm:px-2">
<div className="flex items-center gap-1">
<SidebarTrigger />
<TopBarActions
onNewThread={handleNewThread}
onNewCompare={handleNewCompare}
showCompare={canCompare}
/>
<ModelSelector
models={models}
loraModels={loraModels}
value={inferenceParams.checkpoint}
onValueChange={handleCheckpointChange}
onEject={handleEject}
variant="ghost"
open={modelSelectorOpen}
onOpenChange={handleModelSelectorOpenChange}
triggerDataTour="chat-model-selector"
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]">
{modelsError}
</div>
)}
<div className="flex-1" />
<button
type="button"
onClick={() => setSettingsOpen((o) => !o)}
className="flex h-9 w-9 items-center justify-center rounded-md text-muted-foreground transition-colors hover:bg-accent hover:text-foreground"
title="Inference settings"
data-tour="chat-settings"
>
<HugeiconsIcon icon={Settings04Icon} className="size-5" />
</button>
</div>
{view.mode === "single" ? (
<SingleContent
key={view.threadId ?? view.newThreadNonce ?? "new"}
threadId={view.threadId}
newThreadNonce={view.newThreadNonce}
/>
) : (
<CompareContent key={view.pairId} pairId={view.pairId} />
)}
<div className="flex-1" />
<button
type="button"
onClick={() => setSettingsOpen((o) => !o)}
className="flex h-9 w-9 items-center justify-center rounded-md text-muted-foreground transition-colors hover:bg-accent hover:text-foreground"
title="Inference settings"
data-tour="chat-settings"
>
<HugeiconsIcon icon={Settings04Icon} className="size-5" />
</button>
</div>
{view.mode === "single" ? (
<SingleContent
key={view.threadId ?? view.newThreadNonce ?? "new"}
threadId={view.threadId}
newThreadNonce={view.newThreadNonce}
/>
) : (
<CompareContent key={view.pairId} pairId={view.pairId} />
)}
</div>
<ChatSettingsPanel
open={settingsOpen}
params={inferenceParams}
onParamsChange={setInferenceParams}
autoTitle={autoTitle}
onAutoTitleChange={setAutoTitle}
/>
</SidebarProvider>
<ChatSettingsPanel
open={settingsOpen}
params={inferenceParams}
onParamsChange={setInferenceParams}
autoTitle={autoTitle}
onAutoTitleChange={setAutoTitle}
/>
</SidebarProvider>
</div>
);
}

View file

@ -19,6 +19,20 @@ db.version(2)
})
.upgrade((tx) => tx.table("messages").clear());
db.version(3)
.stores({
threads: "id, modelType, pairId, archived, createdAt",
messages: "id, threadId, createdAt",
})
.upgrade((tx) =>
tx
.table("threads")
.toCollection()
.modify((thread) => {
if (!thread.modelId) thread.modelId = "";
}),
);
export { db };
export function useLiveQuery<T>(

View file

@ -165,8 +165,10 @@ export function useChatModelRuntime() {
setLoadingModel({ id: modelId, displayName });
try {
async function performLoad(): Promise<void> {
if (params.checkpoint) {
await unloadModel({ model_path: params.checkpoint });
const currentCheckpoint =
useChatRuntimeStore.getState().params.checkpoint;
if (currentCheckpoint) {
await unloadModel({ model_path: currentCheckpoint });
}
const loadResponse = await loadModel({

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>) ?? {},
@ -300,10 +323,13 @@ function createDexieAdapter(
},
async initialize(threadId: string) {
const currentModelId =
useChatRuntimeStore.getState().params.checkpoint ?? "";
await db.threads.add({
id: threadId,
title: "New Chat",
modelType,
modelId: currentModelId,
pairId,
archived: false,
createdAt: Date.now(),
@ -441,9 +467,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 +481,7 @@ function ThreadHistoryProvider({
threadId: remoteId,
role: message.role,
content,
...(attachments.length > 0 && { attachments }),
...(custom && Object.keys(custom).length > 0 && { metadata: custom }),
createdAt,
});

View file

@ -8,6 +8,7 @@ export interface ThreadRecord {
id: string;
title: string;
modelType: ModelType;
modelId?: string;
pairId?: string;
archived: boolean;
createdAt: number;
@ -18,6 +19,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

@ -244,10 +244,6 @@ export function ModelSelectionStep() {
<span className="text-xs text-muted-foreground shrink-0">
{sizeLabel}
</span>
) : r?.downloads != null ? (
<span className="text-[10px] text-muted-foreground shrink-0">
{formatCompact(r.downloads)}
</span>
) : null}
</ComboboxItem>
);

View file

@ -210,11 +210,7 @@ export function ModelSection() {
{ est: number; status: VramFitStatus | null; detail: string | null }
>();
for (const r of hfResults) {
const detail = r.totalParams
? formatCompact(r.totalParams)
: r.downloads != null
? `\u2193${formatCompact(r.downloads)}`
: null;
const detail = r.totalParams ? formatCompact(r.totalParams) : null;
if (r.totalParams) {
const est = estimateLoadingVram(r.totalParams, method);
const status = gpu.available

View file

@ -30,7 +30,7 @@ import {
ZapIcon,
} from "@hugeicons/core-free-icons";
import { HugeiconsIcon } from "@hugeicons/react";
import { useState, type ReactElement, type ReactNode } from "react";
import { useEffect, useState, type ReactElement, type ReactNode } from "react";
import { Link, useNavigate } from "@tanstack/react-router";
import { useShallow } from "zustand/react/shallow";
import { useGpuUtilization } from "@/hooks";
@ -83,6 +83,13 @@ export function ProgressSection(): ReactElement {
const { stopTrainingRun } = useTrainingActions();
const gpu = useGpuUtilization(runtime.isTrainingRunning);
const [stopDialogOpen, setStopDialogOpen] = useState(false);
const [stopRequested, setStopRequested] = useState(false);
useEffect(() => {
if (!runtime.isTrainingRunning) {
setStopRequested(false);
}
}, [runtime.isTrainingRunning]);
const pct =
runtime.totalSteps > 0
@ -209,11 +216,12 @@ export function ProgressSection(): ReactElement {
data-tour="studio-training-stop"
variant="destructive"
size="sm"
className="h-7 cursor-pointer px-3 text-xs"
className={`h-7 px-3 text-xs ${stopRequested ? "cursor-not-allowed opacity-60" : "cursor-pointer"}`}
onClick={() => setStopDialogOpen(true)}
disabled={!runtime.isTrainingRunning}
disabled={!runtime.isTrainingRunning || stopRequested}
>
<HugeiconsIcon icon={StopIcon} className="size-3" /> Stop
<HugeiconsIcon icon={StopIcon} className="size-3" />
{stopRequested ? "Stopping…" : "Stop"}
</Button>
<AlertDialogContent overlayClassName="bg-background/40 supports-backdrop-filter:backdrop-blur-[1px]">
<AlertDialogHeader>
@ -226,12 +234,24 @@ export function ProgressSection(): ReactElement {
<AlertDialogCancel>Continue Training</AlertDialogCancel>
<AlertDialogAction
variant="destructive"
onClick={() => void stopTrainingRun(false)}
onClick={() => {
setStopRequested(true);
setStopDialogOpen(false);
void stopTrainingRun(false).then((ok) => {
if (!ok) setStopRequested(false);
});
}}
>
Cancel Training
</AlertDialogAction>
<AlertDialogAction
onClick={() => void stopTrainingRun(true)}
onClick={() => {
setStopRequested(true);
setStopDialogOpen(false);
void stopTrainingRun(true).then((ok) => {
if (!ok) setStopRequested(false);
});
}}
>
Stop and Save
</AlertDialogAction>

View file

@ -1,6 +1,6 @@
import type { PipelineType } from "@huggingface/hub";
import { listModels } from "@huggingface/hub";
import { useCallback } from "react";
import { useCallback, useMemo } from "react";
import { useHfPaginatedSearch } from "./use-hf-paginated-search";
export interface HfModelResult {
@ -64,6 +64,53 @@ function mapModel(raw: unknown): HfModelResult | null {
};
}
/** Number of unsloth results to pull up-front before yielding general results. */
const UNSLOTH_PREFETCH = 20;
/**
* Creates a merged async generator that yields unsloth-owned models first,
* then general results (with deduplication).
*/
async function* mergedModelIterator(
query: string,
task?: PipelineType,
accessToken?: string,
): AsyncGenerator<unknown> {
const common = {
additionalFields: ["safetensors", "tags"] as ("safetensors" | "tags")[],
fetch: withPopularitySort,
...(accessToken ? { credentials: { accessToken } } : {}),
};
// Fire both iterators immediately (parallel network requests on first pull)
const unslothIter = listModels({
search: { query, owner: "unsloth", ...(task ? { task } : {}) },
...common,
});
const generalIter = listModels({
search: { query, ...(task ? { task } : {}) },
...common,
});
// Phase 1: pull & yield unsloth models first
const seen = new Set<string>();
let count = 0;
for await (const model of unslothIter) {
const m = model as { name?: string };
if (m.name) seen.add(m.name);
yield model;
count++;
if (count >= UNSLOTH_PREFETCH) break;
}
// Phase 2: yield general results, skipping already-seen unsloth models
for await (const model of generalIter) {
const m = model as { name?: string };
if (m.name && seen.has(m.name)) continue;
yield model;
}
}
export function useHfModelSearch(
query: string,
options?: { task?: PipelineType; accessToken?: string },
@ -71,18 +118,36 @@ export function useHfModelSearch(
const { task, accessToken } = options ?? {};
const createIter = useCallback(
() =>
listModels({
search: {
...(query.trim() ? { query } : { owner: "unsloth" }),
...(task ? { task } : {}),
},
additionalFields: ["safetensors", "tags"],
fetch: withPopularitySort,
...(accessToken ? { credentials: { accessToken } } : {}),
}) as AsyncGenerator<unknown>,
() => {
const trimmed = query.trim();
if (!trimmed) {
// No query → show default unsloth models
return listModels({
search: { owner: "unsloth", ...(task ? { task } : {}) },
additionalFields: ["safetensors", "tags"],
fetch: withPopularitySort,
...(accessToken ? { credentials: { accessToken } } : {}),
}) as AsyncGenerator<unknown>;
}
// Dual-query: unsloth first, then general
return mergedModelIterator(trimmed, task, accessToken) as AsyncGenerator<unknown>;
},
[query, task, accessToken],
);
return useHfPaginatedSearch(createIter, mapModel);
const search = useHfPaginatedSearch(createIter, mapModel);
// Secondary sort guarantee: unsloth models always float to the top
const results = useMemo(
() =>
[...search.results].sort((a, b) => {
const aFirst = a.id.startsWith("unsloth/") ? 0 : 1;
const bFirst = b.id.startsWith("unsloth/") ? 0 : 1;
return aFirst - bFirst;
}),
[search.results],
);
return { ...search, results };
}