Merge remote-tracking branch 'origin/main' into pr-5717

Resolves package.json conflict: keep main's biome 1->2 simplification
(`biome check` without trailing ".") and the test / test:watch scripts
this branch added for vitest.
This commit is contained in:
Daniel Han 2026-05-27 06:21:06 +00:00
commit b62e4d18cd
74 changed files with 10967 additions and 1234 deletions

View file

@ -0,0 +1,77 @@
// SPDX-License-Identifier: AGPL-3.0-only
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
"use client";
import {
type ReactNode,
createContext,
useCallback,
useContext,
useMemo,
useState,
} from "react";
export type GeneratedImageOverlayState = {
image: string;
title: string;
metadata: string;
filename?: string;
openaiImageGenerationCallId?: string;
openaiResponseId?: string;
openaiReasoningItem?: unknown;
threadId?: string | null;
};
type GeneratedImageOverlayContextValue = {
overlay: GeneratedImageOverlayState | null;
openOverlay: (overlay: GeneratedImageOverlayState) => void;
closeOverlay: () => void;
};
const GeneratedImageOverlayContext =
createContext<GeneratedImageOverlayContextValue | null>(null);
export function GeneratedImageOverlayProvider({
children,
threadId = null,
}: {
children: ReactNode;
threadId?: string | null;
}) {
const [overlay, setOverlay] = useState<GeneratedImageOverlayState | null>(
null,
);
const openOverlay = useCallback(
(nextOverlay: GeneratedImageOverlayState) => {
setOverlay({ ...nextOverlay, threadId: nextOverlay.threadId ?? threadId });
},
[threadId],
);
const closeOverlay = useCallback(() => {
setOverlay(null);
}, []);
const value = useMemo(
() => ({ overlay, openOverlay, closeOverlay }),
[closeOverlay, openOverlay, overlay],
);
return (
<GeneratedImageOverlayContext.Provider value={value}>
{children}
</GeneratedImageOverlayContext.Provider>
);
}
export function useGeneratedImageOverlay(): GeneratedImageOverlayContextValue {
const context = useContext(GeneratedImageOverlayContext);
if (!context) {
throw new Error(
"useGeneratedImageOverlay must be used within GeneratedImageOverlayProvider.",
);
}
return context;
}

View file

@ -0,0 +1,510 @@
// SPDX-License-Identifier: AGPL-3.0-only
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
//
// Portions adapted from assistant-ui packages/ui/src/components/assistant-ui/image.tsx
// MIT License, Copyright (c) 2025 AgentbaseAI Inc.
// Source: https://github.com/assistant-ui/assistant-ui/blob/main/packages/ui/src/components/assistant-ui/image.tsx
"use client";
import { cn } from "@/lib/utils";
import type {
ImageMessagePart,
ImageMessagePartComponent,
} from "@assistant-ui/react";
import { type VariantProps, cva } from "class-variance-authority";
import {
CopyIcon,
DownloadIcon,
ImageIcon,
ImageOffIcon,
Loader2Icon,
RefreshCwIcon,
ShieldAlertIcon,
} from "lucide-react";
import {
type ComponentProps,
type PropsWithChildren,
memo,
useEffect,
useRef,
useState,
} from "react";
import { createPortal } from "react-dom";
const extensionForMimeType = (mimeType?: string): string => {
switch (mimeType) {
case "image/png":
return "png";
case "image/jpeg":
case "image/jpg":
return "jpg";
case "image/webp":
return "webp";
case "image/gif":
return "gif";
case "image/svg+xml":
return "svg";
default:
return "png";
}
};
const DATA_URI_MIME_RE = /data:([^;]+)/;
const DATA_URI_BASE64_RE = /;base64/i;
const IMAGE_DATA_URI_MIME_RE = /^data:([^;,]+)/;
export const dataUriToBlob = (dataUri: string): Blob => {
const [meta, data] = dataUri.split(",");
const mime = meta?.match(DATA_URI_MIME_RE)?.[1] ?? "application/octet-stream";
if (!DATA_URI_BASE64_RE.test(meta ?? "")) {
return new Blob([decodeURIComponent(data ?? "")], { type: mime });
}
const bytes = atob(data ?? "");
const arr = new Uint8Array(bytes.length);
for (let i = 0; i < bytes.length; i += 1) {
arr[i] = bytes.charCodeAt(i);
}
return new Blob([arr], { type: mime });
};
const mimeFromImage = (image: string): string | undefined =>
image.match(IMAGE_DATA_URI_MIME_RE)?.[1];
export const downloadImagePart = (
part: Pick<ImageMessagePart, "image" | "filename">,
): void => {
if (typeof document === "undefined") {
return;
}
const ext = extensionForMimeType(mimeFromImage(part.image));
const filename = part.filename ?? `image.${ext}`;
const isDataUri = part.image.startsWith("data:");
const objectUrl = isDataUri
? URL.createObjectURL(dataUriToBlob(part.image))
: null;
const href = objectUrl ?? part.image;
const a = document.createElement("a");
a.href = href;
a.download = filename;
a.rel = "noopener";
document.body.appendChild(a);
a.click();
document.body.removeChild(a);
if (objectUrl) {
URL.revokeObjectURL(objectUrl);
}
};
export const copyImagePart = async (
part: Pick<ImageMessagePart, "image">,
): Promise<void> => {
if (
typeof navigator === "undefined" ||
!navigator.clipboard ||
typeof ClipboardItem === "undefined"
) {
throw new Error("Clipboard API is not available in this environment.");
}
const blob = part.image.startsWith("data:")
? dataUriToBlob(part.image)
: await fetch(part.image).then((r) => r.blob());
const mime = mimeFromImage(part.image) || blob.type || "image/png";
await navigator.clipboard.write([new ClipboardItem({ [mime]: blob })]);
};
const reportImageCopyError = (error: unknown): void => {
if (typeof window === "undefined") {
return;
}
window.dispatchEvent(
new CustomEvent("assistant-ui:image-copy-error", { detail: error }),
);
};
const imageVariants = cva(
"aui-image-root relative overflow-hidden rounded-lg",
{
variants: {
variant: {
outline: "border border-border",
ghost: "",
muted: "bg-muted/50",
},
size: {
sm: "max-w-64",
default: "max-w-96",
lg: "max-w-[512px]",
full: "w-full",
},
},
defaultVariants: {
variant: "outline",
size: "default",
},
},
);
export type ImageRootProps = ComponentProps<"div"> &
VariantProps<typeof imageVariants>;
function ImageRoot({
className,
variant,
size,
children,
...props
}: ImageRootProps) {
return (
<div
data-slot="image-root"
data-variant={variant}
data-size={size}
className={cn(imageVariants({ variant, size, className }))}
{...props}
>
{children}
</div>
);
}
type ImagePreviewProps = Omit<ComponentProps<"img">, "children"> & {
containerClassName?: string;
};
function ImagePreview({
className,
containerClassName,
onLoad,
onError,
alt = "Image content",
src,
...props
}: ImagePreviewProps) {
const imgRef = useRef<HTMLImageElement>(null);
const [loadedSrc, setLoadedSrc] = useState<string | undefined>(undefined);
const [errorSrc, setErrorSrc] = useState<string | undefined>(undefined);
const loaded = loadedSrc === src;
const error = errorSrc === src;
useEffect(() => {
if (
typeof src === "string" &&
imgRef.current?.complete &&
imgRef.current.naturalWidth > 0
) {
setLoadedSrc(src);
}
}, [src]);
return (
<div
data-slot="image-preview"
className={cn("relative min-h-32", containerClassName)}
>
{!(loaded || error) && (
<div
data-slot="image-preview-loading"
className="absolute inset-0 flex items-center justify-center bg-muted/50"
>
<ImageIcon className="size-8 animate-pulse text-muted-foreground" />
</div>
)}
{error ? (
<div
data-slot="image-preview-error"
className="flex min-h-32 items-center justify-center bg-muted/50 p-4"
>
<ImageOffIcon className="size-8 text-muted-foreground" />
</div>
) : (
<img
{...props}
ref={imgRef}
src={src}
alt={alt}
className={cn(
"block h-auto w-full object-contain",
!loaded && "invisible",
className,
)}
onLoad={(e) => {
if (typeof src === "string") {
setLoadedSrc(src);
}
onLoad?.(e);
}}
onError={(e) => {
if (typeof src === "string") {
setErrorSrc(src);
}
onError?.(e);
}}
/>
)}
</div>
);
}
function ImageFilename({
className,
children,
...props
}: ComponentProps<"span">) {
if (!children) {
return null;
}
return (
<span
data-slot="image-filename"
className={cn(
"block truncate px-2 py-1.5 text-muted-foreground text-xs",
className,
)}
{...props}
>
{children}
</span>
);
}
type ImageZoomProps = PropsWithChildren<{
src: string;
alt?: string;
}>;
function ImageZoom({ src, alt = "Image preview", children }: ImageZoomProps) {
const [isOpen, setIsOpen] = useState(false);
const handleOpen = () => setIsOpen(true);
const handleClose = () => setIsOpen(false);
useEffect(() => {
if (!isOpen) {
return;
}
const handleKeyDown = (e: KeyboardEvent) => {
if (e.key === "Escape") {
setIsOpen(false);
}
};
document.addEventListener("keydown", handleKeyDown);
return () => document.removeEventListener("keydown", handleKeyDown);
}, [isOpen]);
useEffect(() => {
if (!isOpen) {
return;
}
const originalOverflow = document.body.style.overflow;
document.body.style.overflow = "hidden";
return () => {
document.body.style.overflow = originalOverflow;
};
}, [isOpen]);
return (
<>
<button
type="button"
onClick={handleOpen}
className="aui-image-zoom-trigger w-full cursor-zoom-in border-0 bg-transparent p-0 text-left"
aria-label="Click to zoom image"
>
{children}
</button>
{isOpen &&
typeof document !== "undefined" &&
createPortal(
<button
type="button"
data-slot="image-zoom-overlay"
className="aui-image-zoom-overlay fade-in fixed inset-0 z-50 flex animate-in items-center justify-center border-0 bg-black/80 p-0 duration-200"
onClick={handleClose}
aria-label="Close zoomed image"
>
<img
data-slot="image-zoom-content"
src={src}
alt={alt}
className="aui-image-zoom-content fade-in zoom-in-95 max-h-[90vh] max-w-[90vw] animate-in cursor-zoom-out object-contain duration-200"
/>
</button>,
document.body,
)}
</>
);
}
function ImageGenerating({ className }: { className?: string }) {
return (
<div
data-slot="image-generating"
className={cn(
"flex min-h-32 items-center justify-center bg-muted/50 p-4",
className,
)}
>
<Loader2Icon className="size-8 animate-spin text-muted-foreground" />
<span className="sr-only">Generating image</span>
</div>
);
}
function ImageContentFilterError({
className,
reason,
}: {
className?: string;
reason?: string;
}) {
return (
<div
data-slot="image-content-filter-error"
className={cn(
"flex min-h-32 flex-col items-center justify-center gap-2 bg-muted/50 p-4 text-center",
className,
)}
>
<ShieldAlertIcon className="size-8 text-muted-foreground" />
<p className="font-medium text-sm">Image could not be generated</p>
{reason && <p className="text-muted-foreground text-xs">{reason}</p>}
</div>
);
}
export type ImageActionsProps = {
part: ImageMessagePart;
/**
* Wire to your own generation call to show a regenerate button. The button
* renders only when this is set and the part carries a `prompt`.
*/
onRegenerate?: () => void | Promise<void>;
className?: string;
};
function RegenerateButton({
onRegenerate,
}: {
onRegenerate: () => void | Promise<void>;
}) {
const [isRegenerating, setIsRegenerating] = useState(false);
return (
<button
type="button"
onClick={async () => {
setIsRegenerating(true);
try {
await onRegenerate();
} finally {
setIsRegenerating(false);
}
}}
disabled={isRegenerating}
data-slot="image-regenerate"
aria-label="Regenerate image"
className="inline-flex size-7 items-center justify-center rounded hover:bg-muted disabled:opacity-50"
>
<RefreshCwIcon
className={cn("size-4", isRegenerating && "animate-spin")}
/>
</button>
);
}
function ImageActions({ part, onRegenerate, className }: ImageActionsProps) {
return (
<div
data-slot="image-actions"
className={cn("flex items-center gap-1 p-1", className)}
>
<button
type="button"
onClick={() => downloadImagePart(part)}
data-slot="image-download"
aria-label="Download image"
className="inline-flex size-7 items-center justify-center rounded hover:bg-muted"
>
<DownloadIcon className="size-4" />
</button>
<button
type="button"
onClick={() => {
copyImagePart(part).catch((error) => {
reportImageCopyError(error);
});
}}
data-slot="image-copy"
aria-label="Copy image"
className="inline-flex size-7 items-center justify-center rounded hover:bg-muted"
>
<CopyIcon className="size-4" />
</button>
{onRegenerate && <RegenerateButton onRegenerate={onRegenerate} />}
</div>
);
}
const ImageImpl: ImageMessagePartComponent = (props) => {
const { image, filename, status } = props;
const alt = filename || "Image content";
if (status?.type === "running") {
return (
<ImageRoot>
<ImageGenerating />
<ImageFilename>{filename}</ImageFilename>
</ImageRoot>
);
}
if (status?.type === "incomplete" && status.reason === "content-filter") {
return (
<ImageRoot>
<ImageContentFilterError reason="The provider blocked this image." />
</ImageRoot>
);
}
return (
<ImageRoot>
<ImageZoom src={image} alt={alt}>
<ImagePreview src={image} alt={alt} />
</ImageZoom>
<ImageFilename>{filename}</ImageFilename>
</ImageRoot>
);
};
const Image = memo(ImageImpl) as unknown as ImageMessagePartComponent & {
Root: typeof ImageRoot;
Preview: typeof ImagePreview;
Filename: typeof ImageFilename;
Zoom: typeof ImageZoom;
Actions: typeof ImageActions;
Generating: typeof ImageGenerating;
ContentFilterError: typeof ImageContentFilterError;
};
Image.displayName = "Image";
Image.Root = ImageRoot;
Image.Preview = ImagePreview;
Image.Filename = ImageFilename;
Image.Zoom = ImageZoom;
Image.Actions = ImageActions;
Image.Generating = ImageGenerating;
Image.ContentFilterError = ImageContentFilterError;
export {
Image,
ImageRoot,
ImagePreview,
ImageFilename,
ImageZoom,
ImageActions,
ImageGenerating,
ImageContentFilterError,
imageVariants,
};

View file

@ -33,10 +33,24 @@ export const MessageTiming: FC<{
if (timing?.totalStreamTime === undefined) return null;
const serverTimings = (
const custom = (
message.metadata as Record<string, unknown> | undefined
)?.custom as { serverTimings?: Record<string, number> } | undefined;
const st = serverTimings?.serverTimings;
)?.custom as
| {
serverTimings?: Record<string, number>;
contextUsage?: {
cachedTokens?: number;
cacheWriteTokens?: number;
};
}
| undefined;
const st = custom?.serverTimings;
// `??` (not `||`) so an explicit cache_n=0 isn't replaced by a stale
// contextUsage.cachedTokens from a prior turn.
const cacheHits =
st?.cache_n ?? custom?.contextUsage?.cachedTokens ?? 0;
// Anthropic-only cache-write count.
const cacheWrites = custom?.contextUsage?.cacheWriteTokens ?? 0;
// Guard unphysical tok/s: llama.cpp emits predicted_ms=0 on no-op
// turns, blowing the rate up to Infinity. Require >=1 token AND a
@ -122,11 +136,19 @@ export const MessageTiming: FC<{
</span>
</div>
)}
{(st?.cache_n ?? 0) > 0 && (
{cacheHits > 0 && (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Cache hits</span>
<span className="font-mono tabular-nums">
{formatNumber(st!.cache_n)}
{formatNumber(cacheHits)}
</span>
</div>
)}
{cacheWrites > 0 && (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Cache writes</span>
<span className="font-mono tabular-nums">
{formatNumber(cacheWrites)}
</span>
</div>
)}
@ -146,7 +168,7 @@ export const MessageTiming: FC<{
</>
) : (
<>
{/* Client-side metrics (safetensors fallback) */}
{/* Client-side metrics (safetensors + external provider fallback) */}
{timing.firstTokenTime !== undefined && (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">First token</span>
@ -155,6 +177,22 @@ export const MessageTiming: FC<{
</span>
</div>
)}
{cacheHits > 0 && (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Cache hits</span>
<span className="font-mono tabular-nums">
{formatNumber(cacheHits)}
</span>
</div>
)}
{cacheWrites > 0 && (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Cache writes</span>
<span className="font-mono tabular-nums">
{formatNumber(cacheWrites)}
</span>
</div>
)}
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Total</span>
<span className="font-mono tabular-nums">

View file

@ -134,7 +134,7 @@ function ReasoningTrigger({
{active ? (
<span className="text-sm">Thinking...</span>
) : (
<span>Thought for {duration ?? 0} seconds</span>
<span>Thought for {duration ?? 0} {duration === 1 ? "second" : "seconds"}</span>
)}
</span>
<ChevronDownIcon

View file

@ -127,6 +127,12 @@ function Source({
// ── Source badge with hover card ─────────────────────────────
interface SourceData {
/**
* Stable per-citation key. Two Anthropic document citations into
* different spans of the same source share a ``url``, so React keys
* on ``id`` to keep each footnote distinct.
*/
id: string;
url: string;
title: string;
description?: string;
@ -190,8 +196,14 @@ const SourcesGroup: FC = () => {
"url" in part &&
part.url
) {
const url = part.url as string;
const partId =
typeof (part as { id?: unknown }).id === "string"
? ((part as { id: string }).id)
: url;
sources.push({
url: part.url as string,
id: partId,
url,
title: (part as { title?: string }).title || "",
description: (part as { metadata?: { description?: string } })
.metadata?.description,
@ -258,7 +270,7 @@ const SourcesGroup: FC = () => {
className="flex w-full flex-wrap gap-1 invisible absolute pointer-events-none"
>
{sources.map((source) => (
<span key={source.url} className="inline-block">
<span key={source.id} className="inline-block">
<Source href={source.url}>
<SourceIcon url={source.url} />
<SourceTitle>{source.title || extractDomain(source.url)}</SourceTitle>
@ -270,7 +282,7 @@ const SourcesGroup: FC = () => {
{/* Visible container */}
<div className="flex flex-wrap gap-1">
{displayedSources.map((source) => (
<SourceBadge key={source.url} source={source} />
<SourceBadge key={source.id} source={source} />
))}
{shouldCollapse && !expanded && (
<button

View file

@ -7,6 +7,11 @@ import {
UserMessageAttachments,
} from "@/components/assistant-ui/attachment";
import { CodeToggleIcon } from "@/components/assistant-ui/code-toggle-icon";
import {
GeneratedImageOverlayProvider,
useGeneratedImageOverlay,
} from "@/components/assistant-ui/generated-image-overlay-context";
import { downloadImagePart } from "@/components/assistant-ui/image";
import { MarkdownText } from "@/components/assistant-ui/markdown-text";
import { MessageTiming } from "@/components/assistant-ui/message-timing";
import { Reasoning, ReasoningGroup } from "@/components/assistant-ui/reasoning";
@ -39,13 +44,14 @@ import {
import { sentAudioNames } from "@/features/chat/api/chat-adapter";
import { parseExternalModelId } from "@/features/chat/external-providers";
import { getExternalReasoningCapabilities } from "@/features/chat/provider-capabilities";
import { useExternalProvidersStore } from "@/features/chat/stores/external-providers-store";
import { useChatRuntimeStore } from "@/features/chat/stores/chat-runtime-store";
import { useExternalProvidersStore } from "@/features/chat/stores/external-providers-store";
import { deleteThreadMessage } from "@/features/chat/utils/delete-thread-message";
import { applyQwenThinkingParams } from "@/features/chat/utils/qwen-params";
import { isTauri } from "@/lib/api-base";
import { deleteThreadMessage } from "@/features/chat/utils/delete-thread-message";
import { AUDIO_ACCEPT, MAX_AUDIO_SIZE, fileToBase64 } from "@/lib/audio-utils";
import { copyToClipboard } from "@/lib/copy-to-clipboard";
import { toast } from "@/lib/toast";
import { cn } from "@/lib/utils";
import {
ActionBarMorePrimitive,
@ -69,7 +75,6 @@ import {
DownloadIcon,
GlobeIcon,
HeadphonesIcon,
ImageIcon,
LightbulbIcon,
LightbulbOffIcon,
MicIcon,
@ -79,30 +84,31 @@ import {
TerminalIcon,
XIcon,
} from "lucide-react";
import { Copy01Icon, Delete02Icon, Edit03Icon, Tick02Icon } from "@hugeicons/core-free-icons";
import {
Copy01Icon,
Delete02Icon,
Edit03Icon,
Image03Icon,
Tick02Icon,
} from "@hugeicons/core-free-icons";
import { HugeiconsIcon } from "@hugeicons/react";
import {
type ChangeEvent,
type ComponentProps,
type CompositionEvent,
type FC,
type FormEvent,
type KeyboardEvent,
useCallback,
useEffect,
useRef,
useState,
} from "react";
import { toast } from "@/lib/toast";
export const Thread: FC<{
hideComposer?: boolean;
hideWelcome?: boolean;
targetThreadId?: string;
}> = ({
hideComposer,
hideWelcome,
targetThreadId,
}) => {
}> = ({ hideComposer, hideWelcome, targetThreadId }) => {
// Intent-aware autoscroll: replaces assistant-ui's built-in autoscroll
// to prevent the streaming-mutation race that makes the viewport snap
// back to the bottom while the user is scrolling up (see the hook for
@ -113,85 +119,205 @@ export const Thread: FC<{
const isComposerAttachPending = useAuiState(({ threads }) =>
targetThreadId ? threads.mainThreadId !== targetThreadId : false,
);
const activeThreadId = useChatRuntimeStore((s) => s.activeThreadId);
const threadId = targetThreadId ?? activeThreadId ?? null;
return (
<ThreadPrimitive.Root
className="aui-root aui-thread-root @container relative flex min-h-0 min-w-0 flex-1 basis-0 flex-col overflow-hidden"
style={{
["--thread-max-width" as string]: "48rem",
["--thread-content-max-width" as string]:
"calc(var(--thread-max-width) - 1.5rem)",
}}
>
<IntentAwareScrollProvider value={autoScrollContext}>
<ThreadPrimitive.Viewport
ref={viewportRef}
autoScroll={false}
scrollToBottomOnRunStart={false}
scrollToBottomOnInitialize={false}
scrollToBottomOnThreadSwitch={false}
className={cn(
"aui-thread-viewport aui-stream-viewport relative flex min-h-0 min-w-0 flex-1 basis-0 flex-col overflow-x-auto overflow-y-auto scroll-smooth px-5",
hideComposer ? "pt-4" : "pt-[48px]",
)}
>
{!hideWelcome && (
<AuiIf condition={({ thread }) => thread.isEmpty && !thread.isLoading}>
<ThreadWelcome hideComposer={hideComposer} />
</AuiIf>
)}
<GeneratedImageOverlayProvider key={threadId ?? "default"} threadId={threadId}>
<ThreadPrimitive.Root
className="aui-root aui-thread-root @container relative flex min-h-0 min-w-0 flex-1 basis-0 flex-col overflow-hidden"
style={{
["--thread-max-width" as string]: "48rem",
["--thread-content-max-width" as string]:
"calc(var(--thread-max-width) - 1.5rem)",
}}
>
<IntentAwareScrollProvider value={autoScrollContext}>
<ThreadPrimitive.Viewport
ref={viewportRef}
autoScroll={false}
scrollToBottomOnRunStart={false}
scrollToBottomOnInitialize={false}
scrollToBottomOnThreadSwitch={false}
className={cn(
"aui-thread-viewport aui-stream-viewport relative flex min-h-0 min-w-0 flex-1 basis-0 flex-col overflow-x-auto overflow-y-auto scroll-smooth px-5",
hideComposer ? "pt-4" : "pt-[48px]",
)}
>
{!hideWelcome && (
<AuiIf
condition={({ thread }) => thread.isEmpty && !thread.isLoading}
>
<ThreadWelcome hideComposer={hideComposer} threadId={threadId} />
</AuiIf>
)}
<ThreadPrimitive.Messages
components={{
UserMessage,
EditComposer,
AssistantMessage,
}}
/>
<ThreadPrimitive.Messages
components={{
UserMessage,
EditComposer,
AssistantMessage,
}}
/>
{/* Bottom slack so the last message has breathing room above the
{/* Bottom slack so the last message has breathing room above the
sticky scroll-to-bottom button (and the floating composer in
single mode). Without this, content would butt against the
sticky footer and feel cramped. */}
<AuiIf condition={({ thread }) => hideWelcome || !thread.isEmpty}>
<div
className={cn("shrink-0", hideComposer ? "h-16" : "h-40")}
aria-hidden={true}
/>
</AuiIf>
<AuiIf condition={({ thread }) => hideWelcome || !thread.isEmpty}>
<ThreadPrimitive.ViewportFooter
className={cn(
"aui-thread-viewport-footer pointer-events-none sticky z-20 flex w-full justify-center bg-transparent",
hideComposer ? "bottom-3" : "bottom-[140px]",
)}
>
<ThreadScrollToBottom />
</ThreadPrimitive.ViewportFooter>
</AuiIf>
</ThreadPrimitive.Viewport>
{!hideComposer && (
<AuiIf condition={({ thread }) => hideWelcome || !thread.isEmpty}>
<div className="aui-thread-composer-dock pointer-events-none absolute bottom-0 left-0 right-0 md:right-[10px] z-20">
<AuiIf condition={({ thread }) => hideWelcome || !thread.isEmpty}>
<div
className={cn("shrink-0", hideComposer ? "h-16" : "h-40")}
aria-hidden={true}
className="absolute inset-x-0 bottom-0 top-[10px] bg-background"
/>
<div className="relative px-5 pb-2">
<div className="pointer-events-auto mx-auto w-full max-w-(--thread-max-width)">
<ComposerAnimated disabled={isComposerAttachPending} />
</div>
<p className="composer-footer-note">
LLMs can make mistakes. Double-check responses.
</p>
</div>
</div>
</AuiIf>
</AuiIf>
<AuiIf condition={({ thread }) => hideWelcome || !thread.isEmpty}>
<ThreadPrimitive.ViewportFooter
className={cn(
"aui-thread-viewport-footer pointer-events-none sticky z-20 flex w-full justify-center bg-transparent",
// 150px (was 140px) to add a small gap above the composer
hideComposer ? "bottom-3" : "bottom-[150px]",
)}
>
<ThreadScrollToBottom />
</ThreadPrimitive.ViewportFooter>
</AuiIf>
</ThreadPrimitive.Viewport>
<GeneratedImageViewportOverlay hideComposer={hideComposer} />
{!hideComposer && (
<AuiIf condition={({ thread }) => hideWelcome || !thread.isEmpty}>
<ThreadComposerDock
disabled={isComposerAttachPending}
threadId={threadId}
/>
</AuiIf>
)}
</IntentAwareScrollProvider>
</ThreadPrimitive.Root>
</GeneratedImageOverlayProvider>
);
};
const GeneratedImageViewportOverlay: FC<{ hideComposer?: boolean }> = ({
hideComposer,
}) => {
const { overlay, closeOverlay } = useGeneratedImageOverlay();
useEffect(() => {
if (!overlay) {
return;
}
document.querySelector<HTMLTextAreaElement>(".aui-composer-input")?.focus();
}, [overlay]);
if (!overlay) {
return null;
}
return (
<div className="pointer-events-none absolute inset-0 z-30">
<button
type="button"
className="pointer-events-auto absolute inset-0 bg-background/65 backdrop-blur-[1px] dark:bg-background/55"
onClick={closeOverlay}
aria-label="Close generated image preview"
/>
<section
className={cn(
"pointer-events-none absolute inset-x-5 top-[48px] flex flex-col items-center",
hideComposer ? "bottom-4" : "bottom-[150px]",
)}
</IntentAwareScrollProvider>
</ThreadPrimitive.Root>
aria-label="Generated image preview"
>
<div className="pointer-events-auto relative flex min-h-0 w-full max-w-[1100px] flex-1 flex-col items-center justify-center gap-3 rounded-3xl bg-muted/10 p-3 ring-1 ring-border/20">
<div className="absolute inset-x-3 top-3 z-10 flex justify-end">
<div className="flex shrink-0 items-center gap-1 rounded-full bg-background/70 p-1 ring-1 ring-border/20 backdrop-blur-sm">
<Button
type="button"
variant="ghost"
size="icon-sm"
className="size-7 rounded-full"
onClick={() =>
downloadImagePart({
image: overlay.image,
filename: overlay.filename,
})
}
aria-label="Download generated image"
>
<DownloadIcon className="size-3.5" />
</Button>
<Button
type="button"
variant="ghost"
size="icon-sm"
className="size-7 rounded-full"
onClick={closeOverlay}
aria-label="Close generated image preview"
>
<XIcon className="size-3.5" />
</Button>
</div>
</div>
<div className="flex min-h-0 flex-1 items-center justify-center pt-1">
<img
src={overlay.image}
alt={overlay.title}
className="max-h-full max-w-full object-contain"
/>
</div>
<div
className="w-full max-w-[min(100%,46rem)] shrink-0 text-center"
title={overlay.title}
>
<p className="truncate text-xs font-semibold text-foreground/80">
Generated image
</p>
{overlay.metadata ? (
<p className="truncate text-[11px] font-medium text-muted-foreground">
{overlay.metadata}
</p>
) : null}
{hideComposer ? null : (
<p className="mx-auto mt-2 inline-flex rounded-full bg-primary/10 px-3 py-1 text-xs font-medium text-primary">
Type edits below, then send
</p>
)}
</div>
</div>
</section>
</div>
);
};
const ThreadComposerDock: FC<{
disabled?: boolean;
threadId?: string | null;
}> = ({ disabled, threadId }) => {
const { overlay } = useGeneratedImageOverlay();
return (
<div
className={cn(
"aui-thread-composer-dock pointer-events-none absolute bottom-0 left-0 right-0 md:right-[10px]",
overlay ? "z-40" : "z-20",
)}
>
<div
aria-hidden={true}
className="absolute inset-x-0 bottom-0 top-[10px] bg-background"
/>
<div className="relative px-5 pb-2">
<div className="pointer-events-auto mx-auto w-full max-w-(--thread-max-width)">
<ComposerAnimated disabled={disabled} threadId={threadId} />
</div>
<p className="composer-footer-note">
LLMs can make mistakes. Double-check responses.
</p>
</div>
</div>
);
};
@ -219,13 +345,17 @@ const ThreadScrollToBottom: FC = () => {
);
};
const ThreadWelcome: FC<{ hideComposer?: boolean }> = ({ hideComposer }) => {
const ThreadWelcome: FC<{
hideComposer?: boolean;
threadId?: string | null;
}> = ({ hideComposer, threadId }) => {
const [currentEmoji, setCurrentEmoji] = useState("large sloth drink.png");
useEffect(() => {
const hour = new Date().getHours();
if (hour >= 6 && hour < 12) setCurrentEmoji("large sloth drink.png");
else if (hour >= 12 && hour < 17) setCurrentEmoji("sloth magnify final.png");
else if (hour >= 12 && hour < 17)
setCurrentEmoji("sloth magnify final.png");
else if (hour >= 17 && hour < 21) setCurrentEmoji("sloth shy large.png");
else setCurrentEmoji("unsloth-gem.png");
}, []);
@ -240,11 +370,7 @@ const ThreadWelcome: FC<{ hideComposer?: boolean }> = ({ hideComposer }) => {
<div className="aui-thread-welcome-center flex w-full grow flex-col items-center justify-center pb-[48px]">
<div className="aui-thread-welcome-message flex w-full flex-col justify-center gap-6 px-4">
<div className="flex flex-col items-center gap-2 text-center">
<img
src={currentEmojiSrc}
alt="Sloth mascot"
className="size-20"
/>
<img src={currentEmojiSrc} alt="Sloth mascot" className="size-20" />
<h1 className="aui-thread-welcome-message-inner fade-in slide-in-from-bottom-1 animate-in font-heading font-semibold text-2xl tracking-[-0.02em] duration-200">
Chat with your model
</h1>
@ -252,18 +378,21 @@ const ThreadWelcome: FC<{ hideComposer?: boolean }> = ({ hideComposer }) => {
Run GGUFs, safetensors, vision and audio models
</p>
</div>
{!hideComposer && <ComposerAnimated />}
{!hideComposer && <ComposerAnimated threadId={threadId} />}
</div>
</div>
</div>
);
};
const ComposerAnimated: FC<{ disabled?: boolean }> = ({ disabled }) => {
const ComposerAnimated: FC<{
disabled?: boolean;
threadId?: string | null;
}> = ({ disabled, threadId }) => {
return (
<div className="relative mx-auto min-w-0 w-full max-w-(--thread-max-width)">
<div className="relative z-10 w-full">
<Composer disabled={disabled} />
<Composer disabled={disabled} threadId={threadId} />
</div>
</div>
);
@ -293,8 +422,21 @@ const PendingAudioChip: FC = () => {
);
};
const Composer: FC<{ disabled?: boolean }> = ({ disabled }) => {
const { inputProps, isComposing, isComposingRef } = useImeComposerInputHandlers();
const Composer: FC<{
disabled?: boolean;
threadId?: string | null;
}> = ({ disabled, threadId }) => {
const aui = useAui();
const { overlay, closeOverlay } = useGeneratedImageOverlay();
const setImageToolsEnabled = useChatRuntimeStore(
(s) => s.setImageToolsEnabled,
);
const activeThreadId = useChatRuntimeStore((s) => s.activeThreadId);
const setPendingImageEditReference = useChatRuntimeStore(
(s) => s.setPendingImageEditReference,
);
const { inputProps, isComposing, isComposingRef } =
useImeComposerInputHandlers();
const composerText = useAuiState(({ composer }) => composer.text);
const hasAttachments = useAuiState(
({ composer }) => composer.attachments.length > 0,
@ -304,22 +446,78 @@ const Composer: FC<{ disabled?: boolean }> = ({ disabled }) => {
(attachment) => attachment.status.type === "running",
),
);
const hasPendingAudio = useChatRuntimeStore((s) => Boolean(s.pendingAudioName));
const hasPendingAudio = useChatRuntimeStore((s) =>
Boolean(s.pendingAudioName),
);
const referenceThreadId = threadId ?? activeThreadId ?? null;
const hasSendableContent =
composerText.trim().length > 0 || hasAttachments || hasPendingAudio;
const shouldBlockSend = useCallback(
() =>
!hasSendableContent || isComposingRef.current || hasPendingAttachments,
[hasPendingAttachments, hasSendableContent, isComposingRef],
);
const handleSubmit = useCallback(
(event: FormEvent<HTMLFormElement>) => {
if (
disabled ||
!hasSendableContent ||
isComposingRef.current ||
hasPendingAttachments
) {
(event: Parameters<NonNullable<ComponentProps<"form">["onSubmit"]>>[0]) => {
if (disabled || shouldBlockSend()) {
event.preventDefault();
return;
}
if (overlay) {
const trimmed = composerText.trim();
if (!trimmed) {
event.preventDefault();
return;
}
if (!overlay.openaiImageGenerationCallId) {
event.preventDefault();
toast.error("This generated image cannot be edited", {
description:
"The original image reference is missing. Generate the image again, then retry the edit.",
});
closeOverlay();
return;
}
if ((overlay.threadId ?? null) !== referenceThreadId) {
event.preventDefault();
toast.error("This generated image belongs to another chat", {
description: "Open the original chat and retry the edit.",
});
closeOverlay();
return;
}
setImageToolsEnabled(true);
setPendingImageEditReference({
threadId: overlay.threadId ?? referenceThreadId,
openaiImageGenerationCallId: overlay.openaiImageGenerationCallId,
...(overlay.openaiResponseId
? { openaiResponseId: overlay.openaiResponseId }
: {}),
openaiReasoningItem: overlay.openaiReasoningItem,
});
flushResourcesSync(() => {
aui
.composer()
.setText(
`Use the selected generated image as the reference and apply this edit: ${trimmed}. Preserve everything else exactly.`,
);
});
closeOverlay();
}
},
[disabled, hasPendingAttachments, hasSendableContent, isComposingRef],
[
aui,
closeOverlay,
composerText,
disabled,
overlay,
referenceThreadId,
setImageToolsEnabled,
setPendingImageEditReference,
shouldBlockSend,
],
);
const composerContent = (
@ -328,13 +526,15 @@ const Composer: FC<{ disabled?: boolean }> = ({ disabled }) => {
<PendingAudioChip />
<ToolStatusDisplay />
<ComposerPrimitive.Input
placeholder="Send a message..."
placeholder={
overlay ? "Type your edits for your image" : "Send a message..."
}
className="aui-composer-input composer-input"
minRows={1}
maxRows={12}
autoFocus={!disabled}
disabled={disabled}
aria-label="Message input"
aria-label={overlay ? "Image edit instructions" : "Message input"}
// dir="auto": browser picks LTR/RTL from the first strong char;
// no effect on Latin / CJK / Devanagari.
dir="auto"
@ -342,11 +542,12 @@ const Composer: FC<{ disabled?: boolean }> = ({ disabled }) => {
/>
<ComposerAction
disabled={
disabled || !hasSendableContent || isComposing || hasPendingAttachments
}
blockSend={() =>
!hasSendableContent || isComposingRef.current || hasPendingAttachments
disabled ||
!hasSendableContent ||
isComposing ||
hasPendingAttachments
}
shouldBlockSend={shouldBlockSend}
/>
</>
);
@ -553,7 +754,6 @@ const ComposerAudioUpload: FC = () => {
);
};
const ReasoningToggle: FC = () => {
const modelLoaded = useChatRuntimeStore(
(s) => !!s.params.checkpoint && !s.modelLoading,
@ -565,8 +765,12 @@ const ReasoningToggle: FC = () => {
const setReasoningEnabled = useChatRuntimeStore((s) => s.setReasoningEnabled);
const reasoningStyle = useChatRuntimeStore((s) => s.reasoningStyle);
const reasoningEffort = useChatRuntimeStore((s) => s.reasoningEffort);
const supportsReasoningOff = useChatRuntimeStore((s) => s.supportsReasoningOff);
const reasoningEffortLevels = useChatRuntimeStore((s) => s.reasoningEffortLevels);
const supportsReasoningOff = useChatRuntimeStore(
(s) => s.supportsReasoningOff,
);
const reasoningEffortLevels = useChatRuntimeStore(
(s) => s.reasoningEffortLevels,
);
const setReasoningEffort = useChatRuntimeStore((s) => s.setReasoningEffort);
const lastOpenRouterChosenModel = useChatRuntimeStore(
(s) => s.lastOpenRouterChosenModel,
@ -619,7 +823,8 @@ const ReasoningToggle: FC = () => {
effectiveReasoningEnabled && reasoningEffort !== "none";
const disabled = !(modelLoaded && effectiveSupportsReasoning);
const formatEffortLabel = (level: typeof reasoningEffort): string => {
if (level !== "xhigh") return level.charAt(0).toUpperCase() + level.slice(1);
if (level !== "xhigh")
return level.charAt(0).toUpperCase() + level.slice(1);
const normalized = externalSelection?.modelId?.trim().toLowerCase() ?? "";
if (
normalized.startsWith("claude-opus-4-6") ||
@ -677,23 +882,25 @@ const ReasoningToggle: FC = () => {
{effectiveReasoningEffortLevels
.filter((level) => level !== "none")
.map((level) => (
<DropdownMenuItem
key={level}
onSelect={() => {
setReasoningEffort(level);
setReasoningEnabled(true);
applyQwenThinkingParams(true);
// Kimi's $web_search builtin forbids thinking, so
// enabling thinking flips the Search pill off.
if (isKimiExternal && toolsEnabled) {
setToolsEnabled(false);
}
}}
>
{formatEffortLabel(level)}
{effectiveReasoningVisualEnabled && reasoningEffort === level ? " \u2713" : ""}
</DropdownMenuItem>
))}
<DropdownMenuItem
key={level}
onSelect={() => {
setReasoningEffort(level);
setReasoningEnabled(true);
applyQwenThinkingParams(true);
// Kimi's $web_search builtin forbids thinking, so
// enabling thinking flips the Search pill off.
if (isKimiExternal && toolsEnabled) {
setToolsEnabled(false);
}
}}
>
{formatEffortLabel(level)}
{effectiveReasoningVisualEnabled && reasoningEffort === level
? " \u2713"
: ""}
</DropdownMenuItem>
))}
</DropdownMenuContent>
</DropdownMenu>
);
@ -808,8 +1015,7 @@ const WebSearchToggle: FC = () => {
? externalProviders.find((p) => p.id === externalSelection.providerId)
: undefined;
const isKimiExternal = selectedExternalProvider?.providerType === "kimi";
const disabled =
!modelLoaded || !(supportsTools || supportsBuiltinWebSearch);
const disabled = !modelLoaded || !(supportsTools || supportsBuiltinWebSearch);
return (
<button
@ -899,10 +1105,12 @@ const ImagesToggle: FC = () => {
className="composer-pill-btn"
data-active={imageToolsEnabled && !disabled ? "true" : "false"}
aria-label={
imageToolsEnabled ? "Disable image generation" : "Enable image generation"
imageToolsEnabled
? "Disable image generation"
: "Enable image generation"
}
>
<ImageIcon className="size-3.5" />
<HugeiconsIcon icon={Image03Icon} className="size-3.5" strokeWidth={2} />
<span>Images</span>
</button>
);
@ -966,10 +1174,10 @@ const ToolStatusDisplay: FC = () => {
);
};
const ComposerAction: FC<{ disabled?: boolean; blockSend?: () => boolean }> = ({
disabled,
blockSend,
}) => {
const ComposerAction: FC<{
disabled?: boolean;
shouldBlockSend?: () => boolean;
}> = ({ disabled, shouldBlockSend }) => {
return (
<div className="aui-composer-action-wrapper composer-action-wrapper">
<div className="flex items-center gap-0.5">
@ -1016,7 +1224,7 @@ const ComposerAction: FC<{ disabled?: boolean; blockSend?: () => boolean }> = ({
size="icon"
disabled={disabled}
onClick={(event) => {
if (blockSend?.()) {
if (shouldBlockSend?.()) {
event.preventDefault();
}
}}
@ -1284,7 +1492,11 @@ const UserActionBar: FC = () => {
<CopyButton />
<ActionBarPrimitive.Edit asChild={true}>
<TooltipIconButton tooltip="Edit" className="aui-user-action-edit">
<HugeiconsIcon icon={Edit03Icon} strokeWidth={1.75} className="size-icon" />
<HugeiconsIcon
icon={Edit03Icon}
strokeWidth={1.75}
className="size-icon"
/>
</TooltipIconButton>
</ActionBarPrimitive.Edit>
<DeleteMessageButton />

View file

@ -3,9 +3,14 @@
"use client";
import { type ToolCallMessagePartComponent, useAuiState } from "@assistant-ui/react";
import { ImageIcon, LoaderIcon } from "lucide-react";
import { memo, useEffect, useState } from "react";
import { Button } from "@/components/ui/button";
import { cn } from "@/lib/utils";
import type { ToolCallMessagePartComponent } from "@assistant-ui/react";
import { DownloadIcon, ImageIcon, PencilIcon } from "lucide-react";
import type { CSSProperties, MouseEvent } from "react";
import { memo, useCallback, useEffect, useRef, useState } from "react";
import { useGeneratedImageOverlay } from "./generated-image-overlay-context";
import { Image, downloadImagePart } from "./image";
import {
ToolFallbackContent,
ToolFallbackRoot,
@ -38,6 +43,9 @@ import {
interface ImageGenerationArgs {
prompt?: string;
kind?: string;
openai_image_generation_call_id?: unknown;
openai_response_id?: unknown;
openai_reasoning_item?: unknown;
}
interface ImageGenerationResult {
@ -46,6 +54,85 @@ interface ImageGenerationResult {
size?: string;
quality?: string;
background?: string;
prompt?: string;
}
type GeneratedImagePart = {
type: "image";
image: string;
filename?: string;
};
const CAPTION_COLLAPSED_LINES = 4;
const extensionForMime = (mime: string): string => {
switch (mime.toLowerCase()) {
case "image/jpeg":
case "image/jpg":
return "jpg";
case "image/webp":
return "webp";
case "image/gif":
return "gif";
case "image/svg+xml":
return "svg";
default:
return "png";
}
};
const imageFilenameFromPrompt = (prompt: string, mime: string): string => {
const slug = prompt
.trim()
.toLowerCase()
.replace(/[^a-z0-9]+/g, "-")
.replace(/^-|-$/g, "")
.slice(0, 48);
return `${slug || "generated-image"}.${extensionForMime(mime)}`;
};
const formatGeneratedImageLabel = (prompt: string): string => {
if (!prompt) {
return "Generated image";
}
return prompt.length > 80
? `Generated image: ${prompt.slice(0, 80)}`
: `Generated image: ${prompt}`;
};
const loadingDots = Array.from({ length: 64 }, (_, index) => {
const row = Math.floor(index / 8);
const col = index % 8;
return (
<span
key={index}
className="generated-image-loading-dot"
style={
{
"--dot-row": row,
"--dot-col": col,
} as CSSProperties
}
/>
);
});
function GeneratedImagePlaceholder({ label }: { label: string }) {
return (
<div
className={cn(
"generated-image-loading-card flex aspect-square w-[480px] max-w-full items-center justify-center rounded-2xl bg-muted/20 shadow-[0_0_12px_rgba(15,23,42,0.05),0_6px_18px_rgba(15,23,42,0.04)] dark:shadow-[0_0_12px_rgba(0,0,0,0.18),0_6px_18px_rgba(0,0,0,0.12)]",
)}
aria-busy="true"
aria-label={label}
aria-live="polite"
>
<span className="sr-only">{label}</span>
<div className="generated-image-loading-wave" aria-hidden={true}>
{loadingDots}
</div>
</div>
);
}
const ImageGenerationToolUIImpl: ToolCallMessagePartComponent = ({
@ -53,6 +140,7 @@ const ImageGenerationToolUIImpl: ToolCallMessagePartComponent = ({
result,
status,
}) => {
const { openOverlay } = useGeneratedImageOverlay();
const parsedArgs = (args as ImageGenerationArgs) ?? {};
const prompt = parsedArgs.prompt ?? "";
const isRunning = status?.type === "running";
@ -66,33 +154,131 @@ const ImageGenerationToolUIImpl: ToolCallMessagePartComponent = ({
const imageSrc = imageResult?.image_b64
? `data:${mime};base64,${imageResult.image_b64}`
: null;
const imageTitle =
imageResult?.prompt?.trim() || prompt.trim() || "Generated image";
const captionPrompt = imageResult?.prompt?.trim() || prompt.trim();
const promptLikelyNeedsExpansion = captionPrompt.length > 220;
const imageMetadata = [imageResult?.size, imageResult?.quality, mime]
.filter(Boolean)
.join(" · ");
const openaiImageGenerationCallId =
typeof parsedArgs.openai_image_generation_call_id === "string"
? parsedArgs.openai_image_generation_call_id
: undefined;
const openaiResponseId =
typeof parsedArgs.openai_response_id === "string"
? parsedArgs.openai_response_id
: undefined;
const imagePart: GeneratedImagePart | null = imageSrc
? {
type: "image",
image: imageSrc,
filename: imageFilenameFromPrompt(prompt, mime),
}
: null;
// Collapse the card once the model has resumed streaming prose
// after the image. Mirrors CodeExecutionToolUI so the inline image
// doesn't collapse mid-stream and the user can click to re-expand.
const hasText = useAuiState(({ message }) =>
message.content.some(
(p) =>
p.type === "text" &&
"text" in p &&
(p as { text: string }).text.length > 0,
),
);
const [open, setOpen] = useState(true);
useEffect(() => {
if (isRunning) {
setOpen(true);
} else if (hasText && !imageSrc) {
setOpen(false);
const [expandedCaptionPrompt, setExpandedCaptionPrompt] = useState<
string | null
>(null);
const [promptOverflow, setPromptOverflow] = useState<{
prompt: string;
canExpand: boolean;
} | null>(null);
const captionRef = useRef<HTMLDivElement | null>(null);
const isPendingImage = !imagePart && status?.type === "running";
const promptOverflowMeasured = promptOverflow?.prompt === captionPrompt;
const promptCanExpand = promptOverflowMeasured
? promptOverflow.canExpand
: false;
const promptExpanded = expandedCaptionPrompt === captionPrompt;
const updatePromptOverflow = useCallback(() => {
const captionElement = captionRef.current;
if (!captionElement || !captionPrompt) {
return;
}
}, [isRunning, hasText, imageSrc]);
const computedStyle = window.getComputedStyle(captionElement);
const lineHeight = Number.parseFloat(computedStyle.lineHeight);
const collapsedHeight =
(Number.isFinite(lineHeight) ? lineHeight : 20) *
CAPTION_COLLAPSED_LINES;
const hasOverflow = captionElement.scrollHeight > collapsedHeight + 1;
setPromptOverflow((current) =>
current?.prompt === captionPrompt && current.canExpand === hasOverflow
? current
: { prompt: captionPrompt, canExpand: hasOverflow },
);
}, [captionPrompt]);
useEffect(() => {
const captionElement = captionRef.current;
if (!captionElement || !captionPrompt) {
return;
}
const frame = window.requestAnimationFrame(updatePromptOverflow);
const resizeObserver =
typeof ResizeObserver === "undefined"
? null
: new ResizeObserver(updatePromptOverflow);
resizeObserver?.observe(captionElement);
window.addEventListener("resize", updatePromptOverflow);
return () => {
window.cancelAnimationFrame(frame);
resizeObserver?.disconnect();
window.removeEventListener("resize", updatePromptOverflow);
};
}, [captionPrompt, updatePromptOverflow]);
const shouldClampPrompt =
(promptOverflowMeasured ? promptCanExpand : promptLikelyNeedsExpansion) &&
!promptExpanded;
const runningLabel = "Generating image…";
const completedLabel = prompt
? prompt.length > 80
? `Generated image: ${prompt.slice(0, 80)}`
: `Generated image: ${prompt}`
: "Generated image";
const completedLabel = formatGeneratedImageLabel(prompt);
const showPreview = () => {
if (!imagePart) {
return;
}
openOverlay({
image: imagePart.image,
title: imageTitle,
metadata: imageMetadata,
filename: imagePart.filename,
openaiImageGenerationCallId,
openaiResponseId,
openaiReasoningItem: parsedArgs.openai_reasoning_item,
});
};
const stopOverlayActionPropagation = (
event: MouseEvent<HTMLButtonElement>,
) => {
event.preventDefault();
event.stopPropagation();
};
const handleDownload = (event: MouseEvent<HTMLButtonElement>) => {
stopOverlayActionPropagation(event);
if (imagePart) {
downloadImagePart(imagePart);
}
};
const handleEditClick = (event: MouseEvent<HTMLButtonElement>) => {
stopOverlayActionPropagation(event);
showPreview();
};
if (isPendingImage) {
return (
<div className="aui-tool-fallback-root w-full py-1">
<GeneratedImagePlaceholder label={runningLabel} />
</div>
);
}
return (
<ToolFallbackRoot open={open} onOpenChange={setOpen}>
@ -102,21 +288,77 @@ const ImageGenerationToolUIImpl: ToolCallMessagePartComponent = ({
icon={ImageIcon}
/>
<ToolFallbackContent>
{isRunning && !imageSrc ? (
<div className="flex items-center gap-2 text-sm text-muted-foreground">
<LoaderIcon className="size-3.5 animate-spin" />
<span>{runningLabel}</span>
</div>
) : imageSrc ? (
<figure className="m-0 flex flex-col gap-1.5">
<img
src={imageSrc}
alt={prompt || "Generated image"}
className="max-w-full rounded-md border border-border/60"
/>
{prompt ? (
<figcaption className="text-xs leading-snug text-muted-foreground">
{prompt}
{imagePart ? (
<figure className="m-0 flex flex-col gap-2">
<div className="group/generated-image relative aspect-square w-[480px] max-w-full overflow-hidden rounded-2xl bg-muted/25 shadow-lg shadow-foreground/5 dark:shadow-black/25">
<img
src={imagePart.image}
alt=""
aria-hidden={true}
className="pointer-events-none absolute inset-0 size-full scale-110 object-cover opacity-25 blur-2xl saturate-125"
/>
<div className="pointer-events-none absolute inset-0 bg-background/45" />
<button
type="button"
className="relative z-10 block size-full cursor-zoom-in overflow-hidden rounded-2xl focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-ring focus-visible:ring-offset-2 focus-visible:ring-offset-background"
onClick={showPreview}
aria-label="Open generated image preview"
>
<Image.Preview
src={imagePart.image}
alt={imageTitle}
containerClassName="flex size-full min-h-0 items-center justify-center bg-transparent"
className="size-full object-contain"
/>
</button>
<div className="pointer-events-none absolute inset-x-0 bottom-0 z-20 flex items-end justify-between gap-2 bg-gradient-to-t from-black/55 via-black/20 to-transparent p-3 opacity-100 transition-opacity sm:opacity-0 sm:group-hover/generated-image:opacity-100 sm:group-focus-within/generated-image:opacity-100">
<Button
type="button"
variant="dark"
size="sm"
className="pointer-events-auto h-8 rounded-full bg-black/70 text-white hover:bg-black/85"
onClick={handleEditClick}
>
<PencilIcon className="size-3.5" />
Edit
</Button>
<Button
type="button"
variant="dark"
size="icon-sm"
className="pointer-events-auto rounded-full bg-black/70 text-white hover:bg-black/85"
onClick={handleDownload}
aria-label="Download generated image"
>
<DownloadIcon className="size-4" />
</Button>
</div>
</div>
{captionPrompt ? (
<figcaption className="max-w-[480px] text-xs leading-5 text-muted-foreground">
<div
ref={captionRef}
className={cn(
"whitespace-pre-wrap break-words",
shouldClampPrompt && "max-h-20 overflow-hidden",
)}
>
{captionPrompt}
</div>
{promptCanExpand ? (
<button
type="button"
className="mt-2 inline-flex text-xs font-medium text-foreground/80 underline-offset-4 hover:text-foreground hover:underline focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-ring focus-visible:ring-offset-2 focus-visible:ring-offset-background"
onClick={() =>
setExpandedCaptionPrompt((value) =>
value === captionPrompt ? null : captionPrompt,
)
}
aria-expanded={promptExpanded}
>
{promptExpanded ? "Show less" : "Show more"}
</button>
) : null}
</figcaption>
) : null}
</figure>

View file

@ -24,6 +24,22 @@ const RE_TITLE = /Title:\s*(.+)/;
const RE_URL = /URL:\s*(.+)/;
const RE_SNIPPET = /Snippet:\s*(.+)/s;
/**
* Reject anything that is not a real http(s) URL. Web-search / web-fetch
* output is provider-controlled, so hostile ``javascript:`` / ``data:``
* lines must not reach the Source badge's <a href>.
*/
function isSafeHttpUrl(raw: string): boolean {
const value = raw.trim();
if (!value || /[\r\n]/.test(value)) return false;
try {
const parsed = new URL(value);
return parsed.protocol === "http:" || parsed.protocol === "https:";
} catch {
return false;
}
}
/** Parse the backend's "Title: ...\nURL: ...\nSnippet: ...\n---" format into structured sources. */
function parseSearchResults(raw: string): ParsedSource[] {
if (!raw) {
@ -35,13 +51,14 @@ function parseSearchResults(raw: string): ParsedSource[] {
const titleMatch = block.match(RE_TITLE);
const urlMatch = block.match(RE_URL);
const snippetMatch = block.match(RE_SNIPPET);
if (titleMatch && urlMatch) {
sources.push({
title: titleMatch[1].trim(),
url: urlMatch[1].trim(),
snippet: snippetMatch?.[1]?.trim() ?? "",
});
}
if (!titleMatch || !urlMatch) continue;
const url = urlMatch[1].trim();
if (!isSafeHttpUrl(url)) continue;
sources.push({
title: titleMatch[1].trim(),
url,
snippet: snippetMatch?.[1]?.trim() ?? "",
});
}
return sources;
}

View file

@ -1,7 +1,7 @@
// SPDX-License-Identifier: AGPL-3.0-only
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
import { getAuthToken } from "@/features/auth/session";
import { getAuthToken } from "@/features/auth";
import { apiUrl } from "@/lib/api-base";
import { toast } from "@/lib/toast";
import type { MessageTiming, ToolCallMessagePart } from "@assistant-ui/core";
@ -21,6 +21,7 @@ import { pickFriendlyContainerName } from "../lib/friendly-names";
import {
EXTERNAL_MAX_OUTPUT_TOKENS,
clampReasoningEffortToLevels,
getExternalMaxOutputTokens,
getExternalMinOutputTokens,
getExternalReasoningCapabilities,
getProviderCapabilities,
@ -28,13 +29,19 @@ import {
providerSupportsBuiltinImageGeneration,
providerSupportsBuiltinWebFetch,
providerSupportsBuiltinWebSearch,
providerSupportsFastMode,
} from "../provider-capabilities";
import { useChatRuntimeStore } from "../stores/chat-runtime-store";
import {
type PendingImageEditReference,
useChatRuntimeStore,
} from "../stores/chat-runtime-store";
import { useExternalProvidersStore } from "../stores/external-providers-store";
import { isMultimodalResponse } from "../types/api";
import type {
OpenAIChatCompletionsRequest,
OpenAIChatMessage,
OpenAIMessageContent,
OpenAIReasoningContentPart,
} from "../types/api";
import type { ChatModelSummary } from "../types/runtime";
import { getImageInputUnavailableReason } from "../utils/image-input-support";
@ -70,6 +77,13 @@ interface ServerUsage {
prompt_tokens: number;
completion_tokens: number;
total_tokens: number;
// External prompt-cache fields (see _build_usage_chunk in
// external_provider.py). cache_creation is Anthropic-only.
prompt_tokens_details?: {
cached_tokens?: number;
};
cache_creation_input_tokens?: number;
cache_read_input_tokens?: number;
}
/** Server-side timing data from llama-server's timings object. */
@ -143,6 +157,91 @@ async function updateStoredChatThreadEventually(
}
}
/**
* Return ``raw`` when it is a safe-to-navigate http(s) URL, or "" otherwise.
* Rejects non-string input, CR/LF (header injection), and non-http(s)
* schemes (``javascript:`` / ``data:`` / ``vbscript:``) so provider /
* tool-controlled strings cannot land in an <a href>.
*/
function isSafeNavigableSourceUrl(raw: unknown): string {
if (typeof raw !== "string") return "";
const value = raw.trim();
if (!value || /[\r\n]/.test(value)) return "";
try {
const parsed = new URL(value);
if (parsed.protocol === "http:" || parsed.protocol === "https:") {
return value;
}
} catch {
// Fall through.
}
return "";
}
/** Convert an Anthropic document citation dict into a Sources-panel source. */
function documentCitationToSource(
cit: Record<string, unknown>,
fallbackIdx: number,
): {
type: "source";
sourceType: "url";
id: string;
url: string;
title: string;
metadata?: { description: string };
} | null {
const source =
typeof cit.source === "string" && cit.source ? cit.source : "";
const docTitle =
(typeof cit.document_title === "string" && cit.document_title) ||
(typeof cit.title === "string" && cit.title) ||
"";
const docIndex =
typeof cit.document_index === "number" ? cit.document_index : undefined;
// Only treat ``source`` as a navigable URL when it is real http(s);
// search_result_location can carry a free-form id (e.g. ``kb-doc-42``)
// or a hostile ``javascript:`` / ``data:`` / ``vbscript:`` string.
// Fall back to a stable doc anchor otherwise.
const url =
isSafeNavigableSourceUrl(source) || `#anthropic-doc-${docIndex ?? fallbackIdx}`;
const title = docTitle || source || `Document ${fallbackIdx + 1}`;
const cited =
typeof cit.cited_text === "string" ? cit.cited_text.trim() : "";
// Trim the cited snippet so the Sources panel stays scannable.
const description =
cited.length > 240 ? `${cited.slice(0, 240)}...` : cited;
// Anthropic numbers inline [N] per citation, not per source URL.
// Fold citation type + position-bearing fields into the id so two
// distinct citations on the same source (or two search_result_locations
// with different search_result_index) keep separate Sources entries.
const citationType =
typeof cit.type === "string" ? String(cit.type) : "";
const positionParts = [
cit.search_result_index,
cit.start_char_index,
cit.end_char_index,
cit.start_page_number,
cit.end_page_number,
cit.start_block_index,
cit.end_block_index,
]
.filter((v) => typeof v === "number")
.map((v) => String(v))
.join(":");
const idAnchor = positionParts
? `${citationType}:${positionParts}`
: `${citationType}:${fallbackIdx}`;
const id = `${url}#${idAnchor}`;
return {
type: "source" as const,
sourceType: "url" as const,
id,
url,
title,
...(description ? { metadata: { description } } : {}),
};
}
/** Parse "Title: ...\nURL: ...\nSnippet: ..." blocks into source content parts. */
function parseSourcesFromResult(raw: string): {
type: "source";
@ -167,7 +266,11 @@ function parseSourcesFromResult(raw: string): {
const urlMatch = block.match(/URL:\s*(.+)/);
const snippetMatch = block.match(/Snippet:\s*(.+)/);
if (titleMatch && urlMatch) {
const url = urlMatch[1].trim();
// Drop blocks whose ``URL:`` is not safe http(s); provider/tool
// output is attacker-controllable so a hostile ``javascript:`` /
// ``data:`` line must not reach the Sources panel <a href>.
const url = isSafeNavigableSourceUrl(urlMatch[1]);
if (!url) continue;
const snippet = snippetMatch?.[1]?.trim();
sources.push({
type: "source" as const,
@ -302,37 +405,30 @@ function collectImageParts(
message: RunMessage,
): Array<{ type: "image_url"; image_url: { url: string } }> {
const parts: Array<{ type: "image_url"; image_url: { url: string } }> = [];
const pushImagePart = (part: { type: string }) => {
if (part.type !== "image" || !("image" in part)) {
return;
}
const src = (part as { image: string }).image;
if (!src) {
return;
}
parts.push({
type: "image_url",
image_url: {
url: src.startsWith("data:") ? src : `data:image/png;base64,${src}`,
},
});
};
for (const part of message.content ?? []) {
if (part.type === "image" && "image" in part) {
const src = (part as { image: string }).image;
if (src) {
parts.push({
type: "image_url",
image_url: {
url: src.startsWith("data:") ? src : `data:image/png;base64,${src}`,
},
});
}
}
pushImagePart(part);
}
if ("attachments" in message && (message.attachments?.length ?? 0) > 0) {
for (const attachment of message.attachments ?? []) {
for (const part of attachment.content ?? []) {
if (part.type === "image" && "image" in part) {
const src = (part as { image: string }).image;
if (src) {
parts.push({
type: "image_url",
image_url: {
url: src.startsWith("data:")
? src
: `data:image/png;base64,${src}`,
},
});
}
}
pushImagePart(part);
}
}
}
@ -340,6 +436,78 @@ function collectImageParts(
return parts;
}
function normalizeOpenAIReasoningItem(
value: unknown,
): OpenAIReasoningContentPart | null {
if (!value || typeof value !== "object") {
return null;
}
const item = value as Record<string, unknown>;
if (item.type !== "reasoning" || typeof item.id !== "string" || !item.id) {
return null;
}
const summary = Array.isArray(item.summary)
? item.summary.flatMap((part) => {
if (!part || typeof part !== "object") {
return [];
}
const summaryPart = part as Record<string, unknown>;
return summaryPart.type === "summary_text" &&
typeof summaryPart.text === "string"
? [{ type: "summary_text" as const, text: summaryPart.text }]
: [];
})
: [];
const normalized: OpenAIReasoningContentPart = {
type: "reasoning",
id: item.id,
summary,
};
if (
item.status === "in_progress" ||
item.status === "completed" ||
item.status === "incomplete"
) {
normalized.status = item.status;
}
return normalized;
}
function toOpenAIImageEditReferenceMessage(
reference: PendingImageEditReference,
): OpenAIChatMessage | null {
if (!reference.openaiImageGenerationCallId) {
return null;
}
const content: Exclude<OpenAIMessageContent, string> = [];
const reasoningItem = normalizeOpenAIReasoningItem(
reference.openaiReasoningItem,
);
if (reasoningItem) {
content.push(reasoningItem);
}
content.push({
type: "image_generation_call",
id: reference.openaiImageGenerationCallId,
...(reference.openaiResponseId
? { response_id: reference.openaiResponseId }
: {}),
});
return { role: "assistant", content };
}
// Refusal flag stamped on assistant metadata when the backend emits the
// `anthropic_refusal` _toolEvent. We drop the refused pair from the next
// request body (Anthropic guidance: leaving refusals in context keeps
// refusing). Metadata (not text) prevents content from spoofing a reset.
function isAnthropicRefusalMessage(message: RunMessage): boolean {
if (message.role !== "assistant") return false;
const metadata = (message as { metadata?: unknown }).metadata as
| { custom?: Record<string, unknown> }
| undefined;
return metadata?.custom?.anthropicRefusal === true;
}
function toOpenAIMessage(message: RunMessage): {
role: "system" | "user" | "assistant";
content: OpenAIMessageContent;
@ -360,16 +528,27 @@ function toOpenAIMessage(message: RunMessage): {
/data:audio\/[a-z0-9.+-]+;base64,[A-Za-z0-9+/=]+/g,
"[audio]",
);
if (isAnthropicRefusalMessage(message)) {
// Prune refused assistant turn from outbound history; the
// rendered transcript still shows the user-visible notice.
return null;
}
}
const imageParts = collectImageParts(message);
if (imageParts.length > 0) {
return {
role: message.role,
content: [{ type: "text", text: textContent }, ...imageParts],
content: [
...(textContent ? [{ type: "text" as const, text: textContent }] : []),
...imageParts,
],
};
}
if (!textContent) {
return null;
}
return { role: message.role, content: textContent };
}
@ -804,17 +983,52 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
// the user switches chats while waiting for model load / auto-load.
const resolvedThreadId =
(unstable_threadId ?? runtime.activeThreadId) || undefined;
const resolvedThreadKey = resolvedThreadId ?? null;
const pendingImageEditReferenceForRun = runtime.pendingImageEditReference;
const selectedImageEditReference =
(pendingImageEditReferenceForRun?.threadId ?? null) ===
resolvedThreadKey
? pendingImageEditReferenceForRun
: null;
const clearSelectedImageEditReference = () => {
if (!selectedImageEditReference) {
return;
}
const store = useChatRuntimeStore.getState();
const pending = store.pendingImageEditReference;
if (
pending?.openaiImageGenerationCallId ===
selectedImageEditReference.openaiImageGenerationCallId &&
pending.openaiResponseId ===
selectedImageEditReference.openaiResponseId &&
(pending.threadId ?? null) ===
(selectedImageEditReference.threadId ?? null)
) {
store.clearPendingImageEditReference();
}
};
// Wait for in-progress model load to finish before inferring
if (runtime.modelLoading) {
toast.info("Waiting for model to finish loading…");
await waitForModelReady(abortSignal);
try {
await waitForModelReady(abortSignal);
} catch (error) {
clearSelectedImageEditReference();
throw error;
}
}
if (!useChatRuntimeStore.getState().params.checkpoint) {
// Auto-load the smallest downloaded model
const { loaded, blockedByTrustRemoteCode } =
await autoLoadSmallestModel();
let loaded: boolean;
let blockedByTrustRemoteCode: boolean;
try {
({ loaded, blockedByTrustRemoteCode } = await autoLoadSmallestModel());
} catch (error) {
clearSelectedImageEditReference();
throw error;
}
if (!loaded) {
toast.error(
blockedByTrustRemoteCode
@ -826,6 +1040,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
: "Pick a model in the top bar, then retry.",
},
);
clearSelectedImageEditReference();
throw new Error("Load a model first.");
}
}
@ -833,7 +1048,13 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
// Re-read store after potential auto-load / model ready wait
runtime = useChatRuntimeStore.getState();
const { params } = runtime;
const { supportsTools, toolsEnabled, codeToolsEnabled, imageToolsEnabled } = runtime;
const {
supportsTools,
toolsEnabled,
codeToolsEnabled,
imageToolsEnabled,
webFetchToolsEnabled,
} = runtime;
const externalSelection = parseExternalModelId(params.checkpoint);
const isExternalRequest = externalSelection !== null;
if (
@ -844,6 +1065,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
description:
"Turn on Enable connections in Settings → Connections to use hosted models.",
});
clearSelectedImageEditReference();
throw new Error("Connections disabled.");
}
const externalProvider = isExternalRequest
@ -859,6 +1081,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
toast.error("Connection not found.", {
description: "Open Settings → Connections and add it again.",
});
clearSelectedImageEditReference();
throw new Error("Connection not found.");
}
// Local providers (llama.cpp / vLLM / Ollama) allow an empty key — only block hosted providers.
@ -869,36 +1092,34 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
toast.error("Missing API key for selected connection.", {
description: "Open Settings → Connections and set the API key again.",
});
clearSelectedImageEditReference();
throw new Error("Missing connection API key.");
}
const webSearchEnabledForThisTurn =
Boolean(
externalProvider &&
toolsEnabled &&
providerSupportsBuiltinWebSearch(externalProvider.providerType),
);
const codeExecEnabledForThisTurn =
Boolean(
externalProvider &&
externalSelection &&
codeToolsEnabled &&
providerSupportsBuiltinCodeExecution(
externalProvider.providerType,
externalSelection.modelId,
externalProvider.baseUrl,
),
);
// web_fetch shares the Search pill with web_search (no separate
// UI toggle), so it follows toolsEnabled. Anthropic is the only
// provider that ships it today; on others providerSupportsBuiltinWebFetch
// returns false and this stays inert.
const webFetchEnabledForThisTurn =
Boolean(
externalProvider &&
toolsEnabled &&
providerSupportsBuiltinWebFetch(externalProvider.providerType),
);
const webSearchEnabledForThisTurn = Boolean(
externalProvider &&
toolsEnabled &&
providerSupportsBuiltinWebSearch(externalProvider.providerType),
);
const codeExecEnabledForThisTurn = Boolean(
externalProvider &&
externalSelection &&
codeToolsEnabled &&
providerSupportsBuiltinCodeExecution(
externalProvider.providerType,
externalSelection.modelId,
externalProvider.baseUrl,
),
);
// Fetch pill is independent of Search (Anthropic bills web_fetch
// separately from web_search). Sourced from `webFetchToolsEnabled`;
// on providers without web_fetch the toggle is forced off in
// chat-page's runtime setState.
const webFetchEnabledForThisTurn = Boolean(
externalProvider &&
webFetchToolsEnabled &&
providerSupportsBuiltinWebFetch(externalProvider.providerType),
);
const providerShipsWebFetch = Boolean(
externalProvider &&
providerSupportsBuiltinWebFetch(externalProvider.providerType),
@ -918,11 +1139,58 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
),
);
const outboundMessages = messages
if (selectedImageEditReference && !imageGenerationEnabledForThisTurn) {
clearSelectedImageEditReference();
toast.error("Image editing is unavailable", {
description:
"Select an OpenAI image-generation model, then retry the edit.",
});
throw new Error("Image generation edit unavailable.");
}
// Two-pass build: a refused assistant turn also drops the user
// prompt that triggered it (leaving it in context re-triggers
// the classifier). Refusal flag rides assistant
// metadata.custom.anthropicRefusal, set out-of-band from the
// backend _toolEvent.
const survivingMessages: RunMessage[] = [];
for (const message of messages) {
if (isAnthropicRefusalMessage(message)) {
const last = survivingMessages.at(-1);
if (last && last.role === "user") {
survivingMessages.pop();
}
continue;
}
survivingMessages.push(message);
}
const outboundMessages = survivingMessages
.map(toOpenAIMessage)
.filter((message): message is NonNullable<typeof message> =>
Boolean(message),
);
if (selectedImageEditReference) {
const referenceMessage = toOpenAIImageEditReferenceMessage(
selectedImageEditReference,
);
if (!referenceMessage) {
clearSelectedImageEditReference();
toast.error("This generated image cannot be edited", {
description:
"The original image reference is missing. Generate the image again, then retry the edit.",
});
throw new Error("Generated image edit reference missing.");
}
let insertAt = outboundMessages.length;
for (let i = outboundMessages.length - 1; i >= 0; i -= 1) {
if (outboundMessages[i]?.role === "user") {
insertAt = i;
break;
}
}
outboundMessages.splice(insertAt, 0, referenceMessage);
}
const safeSystemPrompt =
typeof params.systemPrompt === "string" ? params.systemPrompt : "";
@ -941,24 +1209,51 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
const webLabel = providerShipsWebFetch
? "web search or web fetch"
: "web search";
if (!webSearchEnabledForThisTurn && !codeExecEnabledForThisTurn) {
// Treat search and fetch as a single "any web tool" axis so
// the guard only warns when neither pill is on; checking
// webSearchEnabledForThisTurn alone mis-fired when only Fetch
// was on and suppressed live web_fetch calls.
const anyWebEnabledForThisTurn =
webSearchEnabledForThisTurn || webFetchEnabledForThisTurn;
if (
!anyWebEnabledForThisTurn &&
!codeExecEnabledForThisTurn &&
!imageGenerationEnabledForThisTurn
) {
disabledToolGuard =
`You do not have ${webLabel} or code execution tools in this conversation. ` +
`You do not have ${webLabel}, code execution, or image generation tools in this conversation. ` +
"Answer from your own knowledge. " +
"If a request genuinely requires tool use, live data fetch or running code, " +
"If a request genuinely requires tool use, live data fetch, running code, or image generation, " +
"inform the user that you do not have access to these capabilities. " +
"Do not return tool-call syntax inside your response.";
} else if (!webSearchEnabledForThisTurn) {
} else if (!anyWebEnabledForThisTurn && !codeExecEnabledForThisTurn) {
disabledToolGuard =
`You do not have ${webLabel} or code execution tools in this conversation. ` +
"You may still use image generation tools when they are available and useful. " +
"If a request genuinely requires live data fetch or running code, " +
"inform the user that you do not have access to these capabilities. " +
"Do not return tool-call syntax inside your response.";
} else if (!anyWebEnabledForThisTurn) {
const availableTools = [
codeExecEnabledForThisTurn ? "code execution" : null,
imageGenerationEnabledForThisTurn ? "image generation" : null,
].filter(Boolean);
disabledToolGuard =
`You do not have ${webLabel} tools in this conversation. ` +
"You may still use code execution tools when they are available and useful. " +
(availableTools.length > 0
? `You may still use ${availableTools.join(" and ")} tools when they are available and useful. `
: "") +
"If a request genuinely requires live data fetch or web search tool use, " +
"inform the user that you do not have access to these capabilities. " +
"Do not return tool-call syntax inside your response.";
} else if (!codeExecEnabledForThisTurn) {
const availableTools = [
webLabel,
imageGenerationEnabledForThisTurn ? "image generation" : null,
].filter(Boolean);
disabledToolGuard =
"You do not have code execution tools in this conversation. " +
`You may still use ${webLabel} tools when they are available and useful. ` +
`You may still use ${availableTools.join(" and ")} tools when they are available and useful. ` +
"If a request genuinely requires running code or code execution tool use, " +
"inform the user that you do not have access to these capabilities. " +
"Do not return tool-call syntax inside your response.";
@ -988,8 +1283,10 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
});
}
}
const imageBase64 = findLatestUserImageBase64(messages);
const audioBase64 = findLatestUserAudioBase64(messages);
// Scan post-prune history so a refused user turn's image/audio
// doesn't gate or mis-attribute the next non-refused turn.
const imageBase64 = findLatestUserImageBase64(survivingMessages);
const audioBase64 = findLatestUserAudioBase64(survivingMessages);
// Block when ANY image is in the outbound payload (current or
// prior turns) and the loaded model can't process images. Keeps
@ -1018,6 +1315,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
const gatedThreadKey = resolvedThreadId || "__default";
runtime.setThreadRunning(gatedThreadKey, true);
runtime.setThreadRunning(gatedThreadKey, false);
clearSelectedImageEditReference();
throw new Error(imageGateReason);
}
}
@ -1025,7 +1323,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
if (audioBase64) {
const audioName = runtime.pendingAudioName;
if (audioName) {
const lastUserMsg = [...messages]
const lastUserMsg = [...survivingMessages]
.reverse()
.find((m) => m.role === "user");
if (lastUserMsg) sentAudioNames.set(lastUserMsg.id, audioName);
@ -1136,6 +1434,32 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
// Tool call content parts — accumulated and yielded cumulatively.
// result is set directly on the tool-call part when tool_end arrives.
const toolCallParts: ToolCallMessagePart[] = [];
const orderAssistantContent = (
textParts: ReturnType<typeof parseAssistantContent>,
) => {
const imageToolParts = toolCallParts.filter(
(part) => part.toolName === "image_generation",
);
const otherToolParts = toolCallParts.filter(
(part) => part.toolName !== "image_generation",
);
return [...otherToolParts, ...textParts, ...imageToolParts];
};
// Anthropic document_citations tool_event payload, converted to
// Sources-panel source parts at end-of-stream so the inline [N]
// markers have matching entries.
const documentCitationParts: Array<{
type: "source";
sourceType: "url";
id: string;
url: string;
title: string;
metadata?: { description: string };
}> = [];
// Latched on the `anthropic_refusal` tool event; stamped onto the
// final assistant metadata as `custom.anthropicRefusal` to drive
// the history-prune above.
let anthropicRefusalSeen = false;
let serverMetadata: {
usage?: ServerUsage;
timings?: ServerTimings;
@ -1314,8 +1638,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
) {
void updateStoredChatThreadEventually(t.id, {
openaiCodeExecContainerId: null,
})
.catch(() => {});
}).catch(() => {});
continue;
}
openaiCodeExecContainerId = t.openaiCodeExecContainerId;
@ -1359,8 +1682,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
openaiCodeExecContainerId = created.id;
void updateStoredChatThreadEventually(resolvedThreadId, {
openaiCodeExecContainerId: created.id,
})
.catch(() => {});
}).catch(() => {});
} catch {
// Fall back to backend's container_auto path on
// failure — keeps the chat moving; the next turn
@ -1382,18 +1704,17 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
...(externalCapabilities?.topP !== false
? { top_p: params.topP }
: {}),
// Clamp to the cross-provider output cap so a maxTokens value
// carried over from a local-model session does not blow past
// provider limits (e.g. Claude Opus 400s on >128k). Also
// floor to the provider's documented minimum — Kimi's
// thinking models need >=16k or the response truncates
// before the answer fits alongside reasoning_content.
// Floor at the provider's documented min (Kimi thinking
// needs >=16k); clamp at the per-model max.
max_tokens: Math.min(
Math.max(
params.maxTokens,
getExternalMinOutputTokens(externalProvider?.providerType),
),
EXTERNAL_MAX_OUTPUT_TOKENS,
getExternalMaxOutputTokens(
externalProvider?.providerType,
externalSelection?.modelId,
),
),
// Only forward sampling knobs the provider actually accepts; the
// backend's external-provider proxy is param-permissive and would
@ -1419,13 +1740,8 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
enable_tools: true,
enabled_tools: [
...(webSearchEnabledForThisTurn ? ["web_search"] : []),
// Pair web_fetch with the Search pill on any
// provider that ships it (Anthropic today). The
// common workflow is "search returns URLs, fetch
// reads them"; without web_fetch the model can
// surface a citation but cannot quote from the
// page body, which is the whole point of the
// tool. There is no separate UI toggle yet.
// web_fetch has its own Fetch pill, independent
// of Search. Anthropic-only today.
...(webFetchEnabledForThisTurn ? ["web_fetch"] : []),
...(codeExecEnabledForThisTurn ? ["code_execution"] : []),
// OpenAI Responses-API only: `image_generation`
@ -1473,11 +1789,23 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
// attaches `cache_control.ttl` when the value is one of
// "5m" / "1h" (see external_provider.py near line 1375),
// so unknown values are a no-op end-to-end.
...(supportsProviderPromptCacheTtl(externalProvider.providerType) &&
...(supportsProviderPromptCacheTtl(
externalProvider.providerType,
) &&
(externalProvider.enablePromptCaching ?? true) &&
isPromptCacheTtl(externalProvider.promptCacheTtl)
? { prompt_cache_ttl: externalProvider.promptCacheTtl }
: {}),
// Anthropic fast mode (Opus 4.6 / 4.7 only); backend
// silently drops on unsupported models as a second
// line of defence.
...(params.fastMode &&
providerSupportsFastMode(
externalProvider.providerType,
externalSelection.modelId,
)
? { fast_mode: true }
: {}),
...(externalReasoningCaps.supportsReasoning
? externalReasoningCaps.reasoningStyle === "reasoning_effort"
? externalReasoningEnabled
@ -1541,10 +1869,15 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
let retriedWithRefreshedKey = false;
while (true) {
try {
const stream = streamChatCompletions(
await buildRequestPayload(retriedWithRefreshedKey),
abortSignal,
);
let requestPayload: OpenAIChatCompletionsRequest;
try {
requestPayload = await buildRequestPayload(retriedWithRefreshedKey);
} catch (error) {
clearSelectedImageEditReference();
throw error;
}
clearSelectedImageEditReference();
const stream = streamChatCompletions(requestPayload, abortSignal);
for await (const chunk of stream) {
// Handle tool status events
@ -1583,6 +1916,27 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
}
continue;
}
if (toolEvent.type === "document_citations") {
// Convert Anthropic citations_delta footnotes into
// Sources-panel entries matching the inline [N] markers.
const cits = toolEvent.citations;
if (Array.isArray(cits)) {
cits.forEach((entry, idx) => {
if (!entry || typeof entry !== "object") return;
const part = documentCitationToSource(
entry as Record<string, unknown>,
idx,
);
if (
part &&
!documentCitationParts.some((p) => p.id === part.id)
) {
documentCitationParts.push(part);
}
});
}
continue;
}
if (toolEvent.type === "container_invalidated") {
if (resolvedThreadId) {
const field =
@ -1591,11 +1945,16 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
: "openaiCodeExecContainerId";
void updateStoredChatThreadEventually(resolvedThreadId, {
[field]: null,
})
.catch(() => {});
}).catch(() => {});
}
continue;
}
if (toolEvent.type === "anthropic_refusal") {
// Latch the backend refusal signal so the final
// message metadata can drive the prune.
anthropicRefusalSeen = true;
continue;
}
if (toolEvent.type === "tool_start") {
const id =
(toolEvent.tool_call_id as string) ||
@ -1630,6 +1989,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
size?: string;
quality?: string;
background?: string;
prompt?: string;
};
const imageB64 = toolEvent.image_b64 as string | undefined;
if (
@ -1651,6 +2011,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
size: toolEvent.size as string | undefined,
quality: toolEvent.quality as string | undefined,
background: toolEvent.background as string | undefined,
prompt: toolEvent.prompt as string | undefined,
};
} else if (imgIdx !== -1) {
const text = rawResult.slice(0, imgIdx);
@ -1668,16 +2029,29 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
} else {
parsedResult = rawResult;
}
const nextArgs =
toolEvent.arguments &&
typeof toolEvent.arguments === "object"
? (toolEvent.arguments as ToolCallMessagePart["args"])
: undefined;
const mergedArgs = nextArgs
? { ...(toolCallParts[idx].args ?? {}), ...nextArgs }
: toolCallParts[idx].args;
toolCallParts[idx] = {
...toolCallParts[idx],
args: mergedArgs,
argsText: mergedArgs
? JSON.stringify(mergedArgs)
: toolCallParts[idx].argsText,
result: parsedResult,
};
}
}
// Yield cumulative state so tool UI updates (tools first, text after)
// Yield cumulative state so tool UI updates. Search/code tools stay
// before the text, while generated images sit after the answer.
const textParts = parseAssistantContent(cumulativeText);
yield {
content: [...toolCallParts, ...textParts],
content: orderAssistantContent(textParts),
metadata: {
timing: buildTiming(
streamStartTime,
@ -1825,7 +2199,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
if (parts.length > 0 || toolCallParts.length > 0) {
yield {
content: [...toolCallParts, ...parts],
content: orderAssistantContent(parts),
metadata: {
timing: buildTiming(
streamStartTime,
@ -1881,18 +2255,31 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
const finalTokPerSec = meta?.timings?.predicted_per_second;
const serverPromptEvalTime = meta?.timings?.prompt_ms;
// Update context usage in store if we got valid server data
// Prefer llama-server timings; fall back to provider usage envelope.
const cachedTokens =
meta?.timings?.cache_n ??
meta?.usage?.prompt_tokens_details?.cached_tokens ??
meta?.usage?.cache_read_input_tokens ??
0;
// Anthropic-only (billed at the write premium).
const cacheWriteTokens = meta?.usage?.cache_creation_input_tokens ?? 0;
// Gate on the captured checkpoint still being active so a late
// completion from provider A doesn't populate the bar after the
// user switched to provider B mid-stream.
if (
meta?.usage &&
typeof meta.usage.prompt_tokens === "number" &&
typeof meta.usage.completion_tokens === "number" &&
typeof meta.usage.total_tokens === "number"
typeof meta.usage.total_tokens === "number" &&
useChatRuntimeStore.getState().params.checkpoint === params.checkpoint
) {
useChatRuntimeStore.getState().setContextUsage({
promptTokens: meta.usage.prompt_tokens,
completionTokens: meta.usage.completion_tokens,
totalTokens: meta.usage.total_tokens,
cachedTokens: meta.timings?.cache_n ?? 0,
cachedTokens,
cacheWriteTokens,
});
}
@ -1908,21 +2295,24 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
yield {
content: [
...toolCallParts,
...parseAssistantContent(cumulativeText),
...orderAssistantContent(parseAssistantContent(cumulativeText)),
...sourceParts,
...documentCitationParts,
],
metadata: {
timing: finalTiming,
custom: {
reasoningDuration,
// Persisted refusal flag driving the two-pass prune.
anthropicRefusal: anthropicRefusalSeen || undefined,
serverTimings: meta?.timings ?? undefined,
contextUsage: meta?.usage
? {
promptTokens: meta.usage.prompt_tokens,
completionTokens: meta.usage.completion_tokens,
totalTokens: meta.usage.total_tokens,
cachedTokens: meta.timings?.cache_n ?? 0,
cachedTokens,
cacheWriteTokens,
modelId: params.checkpoint,
}
: undefined,

View file

@ -57,6 +57,7 @@ import {
getProviderCapabilities,
providerSupportsBuiltinCodeExecution,
providerSupportsBuiltinImageGeneration,
providerSupportsBuiltinWebFetch,
providerSupportsBuiltinWebSearch,
} from "./provider-capabilities";
import { ChatRuntimeProvider } from "./runtime-provider";
@ -71,6 +72,7 @@ import {
CHAT_CODE_TOOLS_ENABLED_KEY,
CHAT_IMAGE_TOOLS_ENABLED_KEY,
CHAT_TOOLS_ENABLED_KEY,
CHAT_WEB_FETCH_TOOLS_ENABLED_KEY,
loadOptionalBool,
useChatRuntimeStore,
} from "./stores/chat-runtime-store";
@ -779,6 +781,9 @@ export function ChatPage(): ReactElement {
selection.modelId,
provider?.baseUrl,
);
const supportsBuiltinWebFetch = providerSupportsBuiltinWebFetch(
provider?.providerType,
);
// Kimi's k2.6/k2.5 default to thinking enabled on the server side
// (per https://platform.kimi.ai/docs/models). Mirror that default
// in the UI so the Think pill comes up clicked when the user picks
@ -801,6 +806,9 @@ export function ChatPage(): ReactElement {
const storedImageToolsEnabled = loadOptionalBool(
CHAT_IMAGE_TOOLS_ENABLED_KEY,
);
const storedWebFetchToolsEnabled = loadOptionalBool(
CHAT_WEB_FETCH_TOOLS_ENABLED_KEY,
);
const nextToolsEnabled = supportsBuiltinWebSearch
? isKimi
? false
@ -834,6 +842,7 @@ export function ChatPage(): ReactElement {
supportsBuiltinWebSearch,
supportsBuiltinCodeExecution,
supportsBuiltinImageGeneration,
supportsBuiltinWebFetch,
toolsEnabled: nextToolsEnabled,
codeToolsEnabled: supportsBuiltinCodeExecution
? (storedCodeToolsEnabled ?? false)
@ -841,6 +850,10 @@ export function ChatPage(): ReactElement {
imageToolsEnabled: supportsBuiltinImageGeneration
? (storedImageToolsEnabled ?? false)
: false,
// Default Fetch off (Anthropic bills per fetch); deliberate opt-in.
webFetchToolsEnabled: supportsBuiltinWebFetch
? (storedWebFetchToolsEnabled ?? false)
: false,
});
}, [externalProvidersForChat, inferenceParams.checkpoint]);
const canCompare = useMemo(() => {
@ -1008,6 +1021,9 @@ export function ChatPage(): ReactElement {
selectedExternal?.modelId,
selectedProvider?.baseUrl,
);
const supportsBuiltinWebFetch = providerSupportsBuiltinWebFetch(
selectedProvider?.providerType,
);
// See sibling useEffect above: Kimi's k2.x default to thinking
// enabled, so the Think pill comes up clicked. Search pill stays
// off by default; mutual exclusion flips them via the composer.
@ -1026,6 +1042,9 @@ export function ChatPage(): ReactElement {
const storedImageToolsEnabled = loadOptionalBool(
CHAT_IMAGE_TOOLS_ENABLED_KEY,
);
const storedWebFetchToolsEnabled = loadOptionalBool(
CHAT_WEB_FETCH_TOOLS_ENABLED_KEY,
);
const nextToolsEnabled = supportsBuiltinWebSearch
? isKimi
? false
@ -1037,6 +1056,10 @@ export function ChatPage(): ReactElement {
ggufMaxContextLength: null,
ggufNativeContextLength: null,
activeNativePathToken: null,
// Clear previous-model counters; the relaxed external-provider
// render gate would otherwise show stale stats until the next
// completion overwrites them.
contextUsage: null,
supportsReasoning: reasoningCaps.supportsReasoning,
reasoningAlwaysOn: reasoningCaps.reasoningAlwaysOn,
reasoningStyle: reasoningCaps.reasoningStyle,
@ -1063,6 +1086,7 @@ export function ChatPage(): ReactElement {
supportsBuiltinWebSearch,
supportsBuiltinCodeExecution,
supportsBuiltinImageGeneration,
supportsBuiltinWebFetch,
toolsEnabled: nextToolsEnabled,
codeToolsEnabled: supportsBuiltinCodeExecution
? (storedCodeToolsEnabled ?? false)
@ -1070,6 +1094,9 @@ export function ChatPage(): ReactElement {
imageToolsEnabled: supportsBuiltinImageGeneration
? (storedImageToolsEnabled ?? false)
: false,
webFetchToolsEnabled: supportsBuiltinWebFetch
? (storedWebFetchToolsEnabled ?? false)
: false,
...(stillOnOpenRouterFree ? {} : { lastOpenRouterChosenModel: null }),
});
return;
@ -1161,7 +1188,9 @@ export function ChatPage(): ReactElement {
if (!saved) return;
viewBeforeCompareRef.current = null;
navigate({ to: "/chat", search: saved });
// Restore context usage from the active thread's last assistant message.
// Restore usage from the last assistant message, but only if it
// matches the currently active checkpoint. Without this guard the
// relaxed render gate would show stale stats from another model.
const threadId =
saved.thread ?? useChatRuntimeStore.getState().activeThreadId;
if (threadId) {
@ -1175,7 +1204,29 @@ export function ChatPage(): ReactElement {
const usage = metadata?.contextUsage as ReturnType<
typeof useChatRuntimeStore.getState
>["contextUsage"];
if (usage) useChatRuntimeStore.getState().setContextUsage(usage);
if (!usage) return;
const store = useChatRuntimeStore.getState();
const activeCheckpoint = store.params.checkpoint;
const usageModelId =
(usage as { modelId?: unknown }).modelId;
// Scope by modelId when present; reject if no active checkpoint
// (model-scoped usage cannot be attributed to "nothing").
if (typeof usageModelId === "string" && usageModelId) {
if (!activeCheckpoint || usageModelId !== activeCheckpoint) {
return;
}
}
// For local turns, also require the restored count to fit in
// the active window. Skip when unknown (external provider).
const limit = store.ggufContextLength;
if (
typeof limit === "number" &&
limit > 0 &&
(usage.totalTokens ?? 0) > limit
) {
return;
}
store.setContextUsage(usage);
})
.catch((error) => {
if (!isExpectedBackgroundChatStorageError(error)) {
@ -1491,11 +1542,13 @@ export function ChatPage(): ReactElement {
) : null}
</div>
<div className="ml-auto flex items-center gap-2">
{view.mode === "single" && ggufContextLength && contextUsage ? (
{view.mode === "single" && contextUsage ? (
<ContextUsageBar
used={contextUsage.totalTokens}
// null on external providers; the bar handles that.
total={ggufContextLength}
cached={contextUsage.cachedTokens}
cacheWrites={contextUsage.cacheWriteTokens}
promptTokens={contextUsage.promptTokens}
completionTokens={contextUsage.completionTokens}
className="h-[34px]"

View file

@ -85,8 +85,10 @@ import {
import {
EXTERNAL_MAX_OUTPUT_TOKENS,
type ProviderCapabilities,
getExternalMaxOutputTokens,
getExternalMinOutputTokens,
providerSupportsBuiltinCodeExecution,
providerSupportsFastMode,
} from "./provider-capabilities";
import { useChatRuntimeStore } from "./stores/chat-runtime-store";
import type { InferenceParams } from "./types/runtime";
@ -552,6 +554,12 @@ export function ChatSettingsPanel({
activeExternalProvider.baseUrl,
) &&
activeExternalProvider.providerType === "openai";
const showFastModeControl =
activeExternalProvider != null &&
providerSupportsFastMode(
activeExternalProvider.providerType,
externalSelection?.modelId,
);
const activeThreadId = useChatRuntimeStore((s) => s.activeThreadId);
const openAiApiKeyForSection = activeExternalProvider
? getExternalProviderApiKey(activeExternalProvider.id) || null
@ -1152,6 +1160,28 @@ export function ChatSettingsPanel({
</Select>
</div>
) : null}
{showFastModeControl ? (
<div className="flex items-center justify-between gap-3 pt-3">
<div className="flex min-w-0 items-center gap-1.5">
<span className="min-w-0 text-[13px] font-medium leading-[1.25] tracking-nav text-nav-fg">
Fast mode
</span>
<InfoHint>
Beta. Up to 2.5x higher output tokens per second on
Claude Opus 4.6 and 4.7 at 6x standard Opus pricing.
Switching between fast and standard invalidates the
prompt cache and is incompatible with the Priority
service tier.
</InfoHint>
</div>
<Switch
className="panel-switch shrink-0"
checked={Boolean(params.fastMode)}
onCheckedChange={set("fastMode")}
aria-label="Fast mode"
/>
</div>
) : null}
</CollapsibleSection>
) : null}
@ -1280,7 +1310,10 @@ export function ChatSettingsPanel({
}
max={
isExternalModel
? EXTERNAL_MAX_OUTPUT_TOKENS
? getExternalMaxOutputTokens(
externalProviderType,
externalSelection?.modelId,
)
: isGguf && ggufContextLength
? ggufContextLength
: 32768

View file

@ -28,37 +28,66 @@ function getSeverityColor(percent: number): {
export const ContextUsageBar: FC<{
used: number;
total: number;
// null on external providers (no known window); bar then hides the ratio.
total?: number | null;
cached?: number;
// Anthropic-only (billed at the write premium).
cacheWrites?: number;
promptTokens?: number;
completionTokens?: number;
className?: string;
}> = ({ used, total, cached, promptTokens, completionTokens, className }) => {
if (total <= 0) return null;
}> = ({
used,
total,
cached,
cacheWrites,
promptTokens,
completionTokens,
className,
}) => {
const hasKnownLimit = typeof total === "number" && total > 0;
const hasUsageDetails =
promptTokens !== undefined ||
completionTokens !== undefined ||
(cached !== undefined && cached > 0) ||
(cacheWrites !== undefined && cacheWrites > 0);
const percent = Math.min((used / total) * 100, 100);
const severity = getSeverityColor(percent);
// Nothing to show: no limit and no per-turn counters.
if (!hasKnownLimit && used <= 0 && !hasUsageDetails) return null;
const percent = hasKnownLimit
? Math.min((used / (total as number)) * 100, 100)
: null;
const severity = getSeverityColor(percent ?? 0);
return (
<Tooltip>
<TooltipTrigger asChild>
<button
type="button"
aria-label={`Context usage: ${formatTokenCount(used)} of ${formatTokenCount(total)} tokens`}
aria-label={
hasKnownLimit
? `Context usage: ${formatTokenCount(used)} of ${formatTokenCount(total as number)} tokens`
: `Token usage: ${formatTokenCount(used)} tokens`
}
className={cn(
"flex items-center gap-2 rounded-[10px] px-2.5 py-1 font-mono text-chat-icon-fg text-[13px] tabular-nums transition-colors hover:bg-chat-icon-bg-hover hover:text-chat-icon-fg-hover",
className,
)}
>
<span>
{formatTokenCount(used)} / {formatTokenCount(total)}
{hasKnownLimit
? `${formatTokenCount(used)} / ${formatTokenCount(total as number)}`
: `${formatTokenCount(used)} tokens`}
</span>
<div className="h-1.5 w-16 rounded-full bg-black/10 dark:bg-white/15 overflow-hidden">
<div
className={cn("h-full rounded-full transition-all", severity.bar)}
style={{ width: `${percent}%` }}
/>
</div>
{hasKnownLimit && percent !== null ? (
<div className="h-1.5 w-16 rounded-full bg-black/10 dark:bg-white/15 overflow-hidden">
<div
className={cn("h-full rounded-full transition-all", severity.bar)}
style={{ width: `${percent}%` }}
/>
</div>
) : null}
</button>
</TooltipTrigger>
<TooltipContent
@ -68,12 +97,14 @@ export const ContextUsageBar: FC<{
className="[&_span>svg]:hidden!"
>
<div className="grid min-w-44 gap-1.5 text-xs">
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Context usage</span>
<span className={cn("font-mono tabular-nums font-medium", severity.text)}>
{percent.toFixed(1)}%
</span>
</div>
{hasKnownLimit && percent !== null ? (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Context usage</span>
<span className={cn("font-mono tabular-nums font-medium", severity.text)}>
{percent.toFixed(1)}%
</span>
</div>
) : null}
{promptTokens !== undefined && (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Prompt tokens</span>
@ -98,20 +129,32 @@ export const ContextUsageBar: FC<{
</span>
</div>
)}
{cacheWrites !== undefined && cacheWrites > 0 && (
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Cache writes</span>
<span className="font-mono tabular-nums">
{formatTokenCountFull(cacheWrites)}
</span>
</div>
)}
<div className="my-0.5 border-t border-border/40" />
<div className="flex items-center justify-between gap-4">
<span className="text-muted-foreground">Total</span>
<span className="text-muted-foreground">
{hasKnownLimit ? "Total" : "Total tokens"}
</span>
<span className="font-mono tabular-nums">
{formatTokenCountFull(used)} / {formatTokenCountFull(total)}
{hasKnownLimit
? `${formatTokenCountFull(used)} / ${formatTokenCountFull(total as number)}`
: formatTokenCountFull(used)}
</span>
</div>
{percent > 85 && (
{hasKnownLimit && percent !== null && percent > 85 ? (
<div className="mt-1 max-w-64 text-[11px] leading-snug text-muted-foreground/90">
Close to the context limit. Generation will stop at 100%.
Increase <span className="font-medium">Context Length</span> in
the chat Settings panel to keep going.
</div>
)}
) : null}
</div>
</TooltipContent>
</Tooltip>

View file

@ -2,6 +2,14 @@
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
export { ChatPage } from "./chat-page";
export {
getInferenceStatus,
listGgufVariants,
listLocalModels,
loadModel,
type LocalModelInfo,
} from "./api/chat-api";
export type { GgufVariantDetail } from "./types/api";
export {
ChatSettingsPanel,
defaultInferenceParams,

View file

@ -248,7 +248,7 @@ interface BackendInferenceDefaults {
export interface BackendInferenceEnvelope {
is_gguf?: boolean;
context_length?: number | null;
inference?: BackendInferenceDefaults;
inference?: BackendInferenceDefaults | null;
}
export function mergeBackendRecommendedInference({

View file

@ -71,18 +71,95 @@ export function clampReasoningEffortToLevels(
}
/**
* Output-token cap for any external provider request. Picked to stay below the
* tightest declared limit across the providers we ship (Anthropic Claude Opus
* tops out at 128k, GPT-5.x ~128k, Gemini 2.5 ~65k, DeepSeek 8k) while staying
* well above what a typical chat reply needs. The local-model path is not
* subject to this local backends honour whatever the loaded context allows.
*
* If a user's stored maxTokens (e.g. carried over from a prior local-model
* session with a 128k+ context) exceeds this, chat-adapter clamps the
* outbound request so the provider does not 400 on it.
* Fallback cap for unknown providers / models. Prefer
* `getExternalMaxOutputTokens(providerType, modelId)` for the real cap.
*/
export const EXTERNAL_MAX_OUTPUT_TOKENS = 32768;
/**
* Per-model max-output caps from each provider's docs:
* OpenAI: developers.openai.com/api/docs/models/gpt-5.5
* Anthropic: platform.claude.com/docs/en/about-claude/models
* Gemini: ai.google.dev/gemini-api/docs/models/gemini-3.1-pro-preview
* DeepSeek: api-docs.deepseek.com/quick_start/pricing (V4 family)
* Local-model path is unaffected.
*/
const EXTERNAL_MAX_OUTPUT_TOKENS_BY_MODEL: Array<{
providerType: string;
prefixes: readonly string[];
cap: number;
}> = [
// OpenAI
{ providerType: "openai", prefixes: ["gpt-5.5-pro", "gpt-5.5"], cap: 128000 },
{ providerType: "openai", prefixes: ["gpt-5.4-pro", "gpt-5.4"], cap: 65536 },
{ providerType: "openai", prefixes: ["gpt-5.3"], cap: 16384 },
// Anthropic
{
providerType: "anthropic",
prefixes: ["claude-opus-4-7"],
cap: 128000,
},
{
providerType: "anthropic",
prefixes: [
"claude-opus-4-6",
"claude-sonnet-4-6",
"claude-opus-4-5",
"claude-sonnet-4-5",
"claude-haiku-4-5",
],
cap: 64000,
},
// Gemini
{
providerType: "gemini",
prefixes: ["gemini-3", "gemini-pro", "gemini-flash"],
cap: 65536,
},
// DeepSeek (V4: deepseek-chat / deepseek-reasoner alias V4-flash).
{ providerType: "deepseek", prefixes: ["deepseek"], cap: 384000 },
];
/**
* Documented per-model output cap; unknown ids fall back to
* `EXTERNAL_MAX_OUTPUT_TOKENS` (32k). OpenRouter ids are
* `provider/model`; the prefix is stripped before matching.
*/
export function getExternalMaxOutputTokens(
providerType: string | null | undefined,
modelId: string | null | undefined,
): number {
if (!providerType || !modelId) return EXTERNAL_MAX_OUTPUT_TOKENS;
const normalized = modelId.trim().toLowerCase();
if (!normalized) return EXTERNAL_MAX_OUTPUT_TOKENS;
const stripped =
providerType === "openrouter" && normalized.includes("/")
? normalized.split("/").slice(-1)[0]
: normalized;
const effectiveProvider =
providerType === "openrouter"
? _inferProviderFromOpenrouterId(normalized) ?? providerType
: providerType;
for (const entry of EXTERNAL_MAX_OUTPUT_TOKENS_BY_MODEL) {
if (entry.providerType !== effectiveProvider) continue;
if (entry.prefixes.some((prefix) => stripped.startsWith(prefix))) {
return entry.cap;
}
}
return EXTERNAL_MAX_OUTPUT_TOKENS;
}
function _inferProviderFromOpenrouterId(
normalizedId: string,
): string | null {
// Map OpenRouter `provider/model` prefix to our internal providerType.
if (normalizedId.startsWith("openai/")) return "openai";
if (normalizedId.startsWith("anthropic/")) return "anthropic";
if (normalizedId.startsWith("google/")) return "gemini";
if (normalizedId.startsWith("deepseek/")) return "deepseek";
return null;
}
/**
* Whether the external provider offers a built-in web-search tool that the
* model invokes server-side. When `true`, the chat composer's Search button
@ -123,11 +200,9 @@ export function providerSupportsBuiltinWebSearch(
/**
* Whether the external provider exposes a server-side web_fetch tool
* that retrieves a single URL (text or PDF) and emits a document block.
* Only Anthropic ships one today (`web_fetch_20250910`); the chat
* composer pairs it with the Search pill because the typical workflow
* is "search returns URLs, fetch reads them" and the UI doesn't (yet)
* expose web_fetch as an independent toggle.
* (single URL, text or PDF) emitting a document block. Anthropic-only
* today (`web_fetch_20250910` / `web_fetch_20260209`). Gates the
* composer's standalone Fetch pill, independent of Search.
*/
export function providerSupportsBuiltinWebFetch(
providerType: string | null | undefined,
@ -135,6 +210,30 @@ export function providerSupportsBuiltinWebFetch(
return providerType === "anthropic";
}
/**
* Whether the active provider + model supports Anthropic fast-mode
* (`speed: "fast"` + `fast-mode-2026-02-01` header). Opus 4.6 / 4.7
* only per https://platform.claude.com/docs/en/build-with-claude/fast-mode.
* Backend silently drops on unsupported models as a second defence.
*/
const ANTHROPIC_FAST_MODE_MODEL_PREFIXES = [
"claude-opus-4-7",
"claude-opus-4-6",
] as const;
export function providerSupportsFastMode(
providerType: string | null | undefined,
modelId: string | null | undefined,
): boolean {
if (providerType !== "anthropic") return false;
if (!modelId) return false;
// Family boundary ("" or "-") required so IDs like "claude-opus-4-70"
// / "claude-opus-4-7b" do not match.
return ANTHROPIC_FAST_MODE_MODEL_PREFIXES.some(
(prefix) => modelId === prefix || modelId.startsWith(`${prefix}-`),
);
}
/**
* Whether the selected external provider/model exposes a server-side
* code-execution tool. Two providers ship one today:
@ -185,15 +284,21 @@ const OPENAI_CODE_EXECUTION_MODEL_PREFIXES = [
/**
* Strict check that a provider configuration points at OpenAI's
* managed cloud (api.openai.com), as opposed to a custom OpenAI-compat
* backend (ollama / llama.cpp / vLLM / generic "custom" preset). The
* shell tool ONLY exists on OpenAI cloud; sending it to anything else
* 400s the request. Mirror of the backend's
* `is_openai_cloud = "api.openai.com" in self.base_url` guard.
* managed cloud (api.openai.com) or Azure OpenAI Foundry
* (*.openai.azure.com), as opposed to a custom OpenAI-compat backend
* (ollama / llama.cpp / vLLM / generic "custom" preset). The shell and
* image-generation tools only exist on cloud backends; sending them to
* anything else 400s the request. Mirror of the backend's
* `_is_openai_family_cloud` host check.
*/
function isOpenAICloudBaseUrl(baseUrl: string | null | undefined): boolean {
if (!baseUrl) return true; // No override → uses the default openai.com base.
return baseUrl.trim().toLowerCase().includes("api.openai.com");
try {
const host = new URL(baseUrl).hostname.toLowerCase();
return host === "api.openai.com" || host.endsWith(".openai.azure.com");
} catch {
return false;
}
}
export function providerSupportsBuiltinCodeExecution(

View file

@ -826,17 +826,24 @@ function useStudioRuntimeAdapters(): StudioRuntimeAdapters {
completionTokens: number;
totalTokens: number;
cachedTokens: number;
cacheWriteTokens?: number;
modelId?: string;
}
| undefined;
const store = useChatRuntimeStore.getState();
if (
savedUsage &&
store.ggufContextLength &&
savedUsage.totalTokens <= store.ggufContextLength &&
(!savedUsage.modelId ||
savedUsage.modelId === store.params.checkpoint)
) {
// Window check applies only when a local GGUF window is known;
// external providers have ggufContextLength === null.
const withinLocalLimit =
!store.ggufContextLength ||
(savedUsage?.totalTokens ?? 0) <= store.ggufContextLength;
// Legacy unscoped usage (no modelId) is only trusted when a
// known local window bounds the totals, so we can't misattribute
// an old local turn to a newly-selected external provider.
const modelMatches = savedUsage?.modelId
? savedUsage.modelId === store.params.checkpoint
: typeof store.ggufContextLength === "number" &&
store.ggufContextLength > 0;
if (savedUsage && withinLocalLimit && modelMatches) {
store.setContextUsage(savedUsage);
}

View file

@ -21,7 +21,20 @@ import { isTauri } from "@/lib/api-base";
import { isMultimodalResponse } from "./types/api";
import { getImageInputUnavailableReason } from "./utils/image-input-support";
import { useAui } from "@assistant-ui/react";
import { ArrowUpIcon, GlobeIcon, HeadphonesIcon, ImageIcon, LightbulbIcon, LightbulbOffIcon, MicIcon, PlusIcon, SquareIcon, XIcon } from "lucide-react";
import {
ArrowUpIcon,
DownloadIcon,
GlobeIcon,
HeadphonesIcon,
LightbulbIcon,
LightbulbOffIcon,
MicIcon,
PlusIcon,
SquareIcon,
XIcon,
} from "lucide-react";
import { Image03Icon } from "@hugeicons/core-free-icons";
import { HugeiconsIcon } from "@hugeicons/react";
import { toast } from "@/lib/toast";
import { loadModel, validateModel } from "./api/chat-api";
import { parseExternalModelId, providerTypeSupportsVision } from "./external-providers";
@ -34,6 +47,7 @@ import {
getExternalReasoningCapabilities,
providerSupportsBuiltinCodeExecution,
providerSupportsBuiltinImageGeneration,
providerSupportsBuiltinWebFetch,
} from "./provider-capabilities";
import {
type CompositionEvent,
@ -336,6 +350,12 @@ export function SharedComposer({
const setImageToolsEnabled = useChatRuntimeStore(
(s) => s.setImageToolsEnabled,
);
const webFetchToolsEnabled = useChatRuntimeStore(
(s) => s.webFetchToolsEnabled,
);
const setWebFetchToolsEnabled = useChatRuntimeStore(
(s) => s.setWebFetchToolsEnabled,
);
const lastOpenRouterChosenModel = useChatRuntimeStore(
(s) => s.lastOpenRouterChosenModel,
);
@ -426,6 +446,9 @@ export function SharedComposer({
effectiveExternalModelId,
selectedExternalProvider?.baseUrl,
);
const supportsBuiltinWebFetch = providerSupportsBuiltinWebFetch(
selectedExternalProvider?.providerType,
);
const searchDisabled =
!modelLoaded || !(supportsTools || supportsBuiltinWebSearch);
const codeDisabled =
@ -437,6 +460,9 @@ export function SharedComposer({
// the pill row stays compact for providers without the capability.
const imageDisabled = !modelLoaded || !supportsBuiltinImageGeneration;
const showImagePill = supportsBuiltinImageGeneration;
// Fetch pill: Anthropic-only (web_fetch_20250910 / web_fetch_20260209).
const webFetchDisabled = !modelLoaded || !supportsBuiltinWebFetch;
const showWebFetchPill = supportsBuiltinWebFetch;
// Backwards-compatible alias for any other call site that may still
// reference `toolsDisabled` (rare; both pills used it before).
const toolsDisabled = codeDisabled;
@ -1102,10 +1128,31 @@ export function SharedComposer({
imageToolsEnabled ? "Disable image generation" : "Enable image generation"
}
>
<ImageIcon className="size-3.5" />
<HugeiconsIcon
icon={Image03Icon}
className="size-3.5"
strokeWidth={2}
/>
<span>Images</span>
</button>
)}
{showWebFetchPill && (
<button
type="button"
disabled={webFetchDisabled}
onClick={() => setWebFetchToolsEnabled(!webFetchToolsEnabled)}
className="composer-pill-btn"
data-active={
webFetchToolsEnabled && !webFetchDisabled ? "true" : "false"
}
aria-label={
webFetchToolsEnabled ? "Disable URL fetch" : "Enable URL fetch"
}
>
<DownloadIcon className="size-3.5" />
<span>Fetch</span>
</button>
)}
</div>
<div className="flex items-center gap-1">
{dictationSupported && (

View file

@ -14,7 +14,9 @@ import {
DEFAULT_INFERENCE_PARAMS,
type InferenceParams,
} from "../types/runtime";
import { isExternalModelId } from "../external-providers";
import { isExternalModelId, parseExternalModelId } from "../external-providers";
import { getExternalMaxOutputTokens } from "../provider-capabilities";
import { useExternalProvidersStore } from "./external-providers-store";
import {
loadChatSettingsWithLegacyImport,
savePersistedChatSettingsPatch,
@ -25,6 +27,8 @@ export const CHAT_REASONING_ENABLED_KEY = "unsloth_chat_reasoning_enabled";
export const CHAT_TOOLS_ENABLED_KEY = "unsloth_chat_tools_enabled";
export const CHAT_CODE_TOOLS_ENABLED_KEY = "unsloth_chat_code_tools_enabled";
export const CHAT_IMAGE_TOOLS_ENABLED_KEY = "unsloth_chat_image_tools_enabled";
export const CHAT_WEB_FETCH_TOOLS_ENABLED_KEY =
"unsloth_chat_web_fetch_tools_enabled";
// External provider selection is encoded into `params.checkpoint` as
// `external::<providerId>::<modelId>`. PersistedChatSettings deliberately
@ -62,6 +66,12 @@ function saveLastExternalCheckpoint(value: string | null): void {
}
export type ReasoningStyle = "enable_thinking" | "reasoning_effort";
export type PendingImageEditReference = {
threadId: string | null;
openaiImageGenerationCallId: string;
openaiResponseId?: string;
openaiReasoningItem?: unknown;
};
export type ReasoningEffort =
| "none"
| "minimal"
@ -262,9 +272,21 @@ type ChatRuntimeStore = {
* receive the tool because their runtime cannot dispatch it.
*/
supportsBuiltinImageGeneration: boolean;
/**
* Whether the active external provider exposes a server-side
* web_fetch tool (Anthropic's `web_fetch_20250910` /
* `web_fetch_20260209`). Gates the composer's Fetch pill,
* independent of Search.
*/
supportsBuiltinWebFetch: boolean;
toolsEnabled: boolean;
codeToolsEnabled: boolean;
imageToolsEnabled: boolean;
/**
* Fetch pill state, independent of `toolsEnabled` (Search). Only
* consulted when `providerSupportsBuiltinWebFetch` is true.
*/
webFetchToolsEnabled: boolean;
toolStatus: string | null;
generatingStatus: string | null;
autoHealToolCalls: boolean;
@ -286,11 +308,14 @@ type ChatRuntimeStore = {
settingsPanelOpen: boolean;
pendingAudioBase64: string | null;
pendingAudioName: string | null;
pendingImageEditReference: PendingImageEditReference | null;
contextUsage: {
promptTokens: number;
completionTokens: number;
totalTokens: number;
cachedTokens: number;
// Anthropic-only; optional so pre-cache-stats persisted entries load.
cacheWriteTokens?: number;
} | null;
modelLoading: boolean;
activeNativePathToken: string | null;
@ -324,6 +349,7 @@ type ChatRuntimeStore = {
setToolsEnabled: (enabled: boolean, options?: { persist?: boolean }) => void;
setCodeToolsEnabled: (enabled: boolean) => void;
setImageToolsEnabled: (enabled: boolean) => void;
setWebFetchToolsEnabled: (enabled: boolean) => void;
setToolStatus: (status: string | null) => void;
setGeneratingStatus: (status: string | null) => void;
setAutoHealToolCalls: (enabled: boolean) => void;
@ -336,6 +362,10 @@ type ChatRuntimeStore = {
setChatTemplateOverride: (template: string | null) => void;
setPendingAudio: (base64: string, name: string) => void;
clearPendingAudio: () => void;
setPendingImageEditReference: (
reference: PendingImageEditReference | null,
) => void;
clearPendingImageEditReference: () => void;
setContextUsage: (usage: ChatRuntimeStore["contextUsage"]) => void;
};
@ -377,6 +407,7 @@ const PERSISTED_INFERENCE_PARAM_KEYS = [
"maxTokens",
"systemPrompt",
"trustRemoteCode",
"fastMode",
] as const satisfies readonly PersistedInferenceParamKey[];
const SCALAR_SETTING_KEYS = [
@ -564,9 +595,11 @@ export const useChatRuntimeStore = create<ChatRuntimeStore>((set, get) => ({
supportsBuiltinWebSearch: false,
supportsBuiltinCodeExecution: false,
supportsBuiltinImageGeneration: false,
supportsBuiltinWebFetch: false,
toolsEnabled: loadBool(CHAT_TOOLS_ENABLED_KEY, false),
codeToolsEnabled: loadBool(CHAT_CODE_TOOLS_ENABLED_KEY, false),
imageToolsEnabled: loadBool(CHAT_IMAGE_TOOLS_ENABLED_KEY, false),
webFetchToolsEnabled: loadBool(CHAT_WEB_FETCH_TOOLS_ENABLED_KEY, false),
toolStatus: null,
generatingStatus: null,
autoHealToolCalls: true,
@ -587,6 +620,7 @@ export const useChatRuntimeStore = create<ChatRuntimeStore>((set, get) => ({
settingsPanelOpen: false,
pendingAudioBase64: null,
pendingAudioName: null,
pendingImageEditReference: null,
contextUsage: null,
modelLoading: false,
activeNativePathToken: null,
@ -640,7 +674,14 @@ export const useChatRuntimeStore = create<ChatRuntimeStore>((set, get) => ({
if (state.settingsHydrated && hasKeys(changedParams)) {
saveSettingsPatch({ inferenceParams: changedParams });
}
return { params };
// Mirror setCheckpoint: the local model load path can mutate
// params.checkpoint via setParams() before setCheckpoint runs,
// leaving stale per-turn counters under the new checkpoint.
const checkpointChanged = state.params.checkpoint !== params.checkpoint;
return {
params,
...(checkpointChanged ? { contextUsage: null } : {}),
};
}),
setCustomPresets: (customPresets) =>
set(() => {
@ -704,12 +745,37 @@ export const useChatRuntimeStore = create<ChatRuntimeStore>((set, get) => ({
// mount, and a stale persisted local id would race against the
// freshly-loaded model. See LAST_EXTERNAL_CHECKPOINT_KEY notes.
saveLastExternalCheckpoint(isExternalModelId(modelId) ? modelId : null);
// Clear stale per-turn usage when the model changes; the relaxed
// external-provider render gate would otherwise show old counters
// until the next completion overwrites them.
const checkpointChanged = state.params.checkpoint !== modelId;
// Clamp maxTokens to the new model's cap on switch into an
// external model so a value carried over from a prior local
// session does not render above the slider's max.
let nextMaxTokens = state.params.maxTokens;
if (checkpointChanged && isExternalModelId(modelId)) {
const parsed = parseExternalModelId(modelId);
const provider = parsed
? useExternalProvidersStore
.getState()
.providers.find((p) => p.id === parsed.providerId)
: null;
const cap = getExternalMaxOutputTokens(
provider?.providerType,
parsed?.modelId,
);
if (nextMaxTokens > cap) {
nextMaxTokens = cap;
}
}
return {
params: {
...state.params,
checkpoint: modelId,
maxTokens: nextMaxTokens,
},
activeGgufVariant: ggufVariant ?? null,
...(checkpointChanged ? { contextUsage: null } : {}),
};
}),
setActiveThreadId: (activeThreadId) =>
@ -744,9 +810,11 @@ export const useChatRuntimeStore = create<ChatRuntimeStore>((set, get) => ({
supportsBuiltinWebSearch: false,
supportsBuiltinCodeExecution: false,
supportsBuiltinImageGeneration: false,
supportsBuiltinWebFetch: false,
toolsEnabled: false,
codeToolsEnabled: false,
imageToolsEnabled: false,
webFetchToolsEnabled: false,
toolStatus: null,
kvCacheDtype: null,
loadedKvCacheDtype: null,
@ -759,6 +827,7 @@ export const useChatRuntimeStore = create<ChatRuntimeStore>((set, get) => ({
defaultChatTemplate: null,
chatTemplateOverride: null,
loadedChatTemplateOverride: null,
pendingImageEditReference: null,
}));
},
setReasoningEnabled: (reasoningEnabled, options) =>
@ -806,6 +875,11 @@ export const useChatRuntimeStore = create<ChatRuntimeStore>((set, get) => ({
saveBool(CHAT_IMAGE_TOOLS_ENABLED_KEY, imageToolsEnabled);
return { imageToolsEnabled };
}),
setWebFetchToolsEnabled: (webFetchToolsEnabled) =>
set(() => {
saveBool(CHAT_WEB_FETCH_TOOLS_ENABLED_KEY, webFetchToolsEnabled);
return { webFetchToolsEnabled };
}),
setToolStatus: (toolStatus) => set({ toolStatus }),
setGeneratingStatus: (generatingStatus) => set({ generatingStatus }),
setAutoHealToolCalls: (autoHealToolCalls) =>
@ -845,5 +919,9 @@ export const useChatRuntimeStore = create<ChatRuntimeStore>((set, get) => ({
set({ pendingAudioBase64: base64, pendingAudioName: name }),
clearPendingAudio: () =>
set({ pendingAudioBase64: null, pendingAudioName: null }),
setPendingImageEditReference: (pendingImageEditReference) =>
set({ pendingImageEditReference }),
clearPendingImageEditReference: () =>
set({ pendingImageEditReference: null }),
setContextUsage: (contextUsage) => set({ contextUsage }),
}));

View file

@ -143,6 +143,7 @@ export interface UnloadModelRequest {
export interface InferenceStatusResponse {
active_model: string | null;
model_identifier?: string | null;
is_vision: boolean;
is_gguf?: boolean;
gguf_variant?: string | null;
@ -158,7 +159,7 @@ export interface InferenceStatusResponse {
min_p?: number;
presence_penalty?: number;
trust_remote_code?: boolean;
};
} | null;
requires_trust_remote_code?: boolean;
supports_reasoning?: boolean;
reasoning_style?: "enable_thinking" | "reasoning_effort";
@ -192,12 +193,31 @@ export interface AudioGenerationResponse {
}>;
}
export type OpenAIMessageContent =
| string
| Array<
| { type: "text"; text: string }
| { type: "image_url"; image_url: { url: string } }
>;
export type OpenAIReasoningSummaryPart = {
type: "summary_text";
text: string;
};
export type OpenAIReasoningContentPart = {
type: "reasoning";
id: string;
summary: OpenAIReasoningSummaryPart[];
status?: "in_progress" | "completed" | "incomplete";
};
export type OpenAIImageGenerationCallContentPart = {
type: "image_generation_call";
id: string;
response_id?: string;
};
export type OpenAIMessageContentPart =
| { type: "text"; text: string }
| { type: "image_url"; image_url: { url: string } }
| OpenAIReasoningContentPart
| OpenAIImageGenerationCallContentPart;
export type OpenAIMessageContent = string | OpenAIMessageContentPart[];
export interface OpenAIChatMessage {
role: "system" | "user" | "assistant";
@ -262,6 +282,12 @@ export interface OpenAIChatCompletionsRequest {
* the Anthropic provider with `code_execution` in `enabled_tools`.
*/
anthropic_code_exec_container_id?: string | null;
/**
* Anthropic fast-mode toggle. Opus 4.6 / 4.7 only; backend drops
* silently on every other model + provider. See
* https://platform.claude.com/docs/en/build-with-claude/fast-mode
*/
fast_mode?: boolean | null;
}
export interface OpenAIChatDelta {

View file

@ -14,6 +14,12 @@ export interface InferenceParams {
checkpoint: string;
/** Allow loading models with custom code (e.g. NVIDIA Nemotron). Only enable for repos you trust. */
trustRemoteCode?: boolean;
/**
* Anthropic fast-mode toggle. Opus 4.6 / 4.7 only; higher OTPS at
* 6x standard Opus pricing. Default false.
* https://platform.claude.com/docs/en/build-with-claude/fast-mode
*/
fastMode?: boolean;
}
export const DEFAULT_INFERENCE_PARAMS: InferenceParams = {
@ -28,6 +34,7 @@ export const DEFAULT_INFERENCE_PARAMS: InferenceParams = {
systemPrompt: "",
checkpoint: "",
trustRemoteCode: false,
fastMode: false,
};
export interface ChatModelSummary {

View file

@ -140,6 +140,11 @@ function sanitizeInferenceParams(
if (typeof value.trustRemoteCode === "boolean") {
params.trustRemoteCode = value.trustRemoteCode;
}
// Mirror trustRemoteCode handling so the toggle survives reload
// and the /api/chat/settings round-trip.
if (typeof value.fastMode === "boolean") {
params.fastMode = value.fastMode;
}
return hasKeys(params) ? params : undefined;
}

View file

@ -3,6 +3,7 @@
import { Input } from "@/components/ui/input";
import type { ReactElement } from "react";
import { LocalRecipeModelSelector } from "../../dialogs/models/local-recipe-model-selector";
import type { ModelConfig, ModelProviderConfig } from "../../types";
import { InlineField } from "./inline-field";
@ -32,7 +33,9 @@ export function InlineModel(props: InlineModelProps): ReactElement {
className="nodrag h-8 w-full text-xs"
placeholder="https://api.example.com/v1"
value={props.config.endpoint}
onChange={(event) => props.onUpdate({ endpoint: event.target.value })}
onChange={(event) =>
props.onUpdate({ endpoint: event.target.value })
}
/>
</InlineField>
<InlineField label="API key">
@ -53,23 +56,32 @@ export function InlineModel(props: InlineModelProps): ReactElement {
}
// model_config branch - mirror the local-aware provider sync from the
// dialog path so inline edits do not leave stale "local" placeholders
// on external providers and fill the placeholder when switching to local.
// dialog path so inline edits clear stale local-only metadata without
// synthesizing the legacy "local" placeholder.
const localNames = props.localProviderNames ?? new Set<string>();
const modelConfig = props.config;
const handleProviderChange = (nextProvider: string) => {
const isLocal = localNames.has(nextProvider);
if (isLocal && !modelConfig.model.trim()) {
props.onUpdate({ provider: nextProvider, model: "local" });
return;
}
if (!isLocal && modelConfig.model === "local") {
props.onUpdate({ provider: nextProvider, model: "" });
return;
}
props.onUpdate({ provider: nextProvider });
};
const isLinkedToLocal = localNames.has(modelConfig.provider);
const handleProviderChange = (nextProvider: string) => {
const nextIsLocal = localNames.has(nextProvider);
if (isLinkedToLocal !== nextIsLocal) {
props.onUpdate({
provider: nextProvider,
model: "",
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: undefined,
});
return;
}
props.onUpdate({
provider: nextProvider,
...(nextIsLocal
? {}
: {
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: undefined,
}),
});
};
return (
<div className="grid gap-3 sm:grid-cols-2">
@ -82,12 +94,38 @@ export function InlineModel(props: InlineModelProps): ReactElement {
/>
</InlineField>
<InlineField label="Model">
<Input
className="nodrag h-8 w-full text-xs"
placeholder={isLinkedToLocal ? "local" : "gpt-4o-mini"}
value={modelConfig.model}
onChange={(event) => props.onUpdate({ model: event.target.value })}
/>
{isLinkedToLocal ? (
<LocalRecipeModelSelector
compact={true}
className="h-8 rounded-md text-xs"
value={
modelConfig.model.trim().toLowerCase() === "local"
? ""
: modelConfig.model
}
ggufVariant={modelConfig.gguf_variant}
onChange={(model, variant) =>
props.onUpdate({
model,
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: variant ?? undefined,
})
}
/>
) : (
<Input
className="nodrag h-8 w-full text-xs"
placeholder="gpt-4o-mini"
value={modelConfig.model}
onChange={(event) =>
props.onUpdate({
model: event.target.value,
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: undefined,
})
}
/>
)}
</InlineField>
<InlineField label="Temperature" className="sm:col-span-2">
<Input

View file

@ -0,0 +1,644 @@
// SPDX-License-Identifier: AGPL-3.0-only
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
import { Badge } from "@/components/ui/badge";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import {
Popover,
PopoverContent,
PopoverTrigger,
} from "@/components/ui/popover";
import { Spinner } from "@/components/ui/spinner";
import {
type GgufVariantDetail,
type LocalModelInfo,
listGgufVariants,
listLocalModels,
} from "@/features/chat";
import { cn } from "@/lib/utils";
import { Link } from "@tanstack/react-router";
import { ChevronDownIcon, ChevronRightIcon, RefreshCwIcon } from "lucide-react";
import {
type ComponentPropsWithoutRef,
type ReactElement,
forwardRef,
useCallback,
useEffect,
useMemo,
useState,
} from "react";
const GGUF_SUFFIX_PATTERN = /-GGUF(?:$|-)/i;
type LocalRecipeModelSelectorProps = {
value: string;
ggufVariant?: string | null;
onChange: (modelId: string, ggufVariant?: string | null) => void;
inputId?: string;
disabled?: boolean;
compact?: boolean;
className?: string;
};
function normalizeForSearch(value: string): string {
return value.toLowerCase().replace(/[\s_.-]/g, "");
}
function hasGgufSuffix(value: string | null | undefined): boolean {
return GGUF_SUFFIX_PATTERN.test(value ?? "");
}
function getModelLabel(model: LocalModelInfo): string {
return model.model_id?.trim() || model.display_name || model.id;
}
function isDirectGguf(model: LocalModelInfo): boolean {
return model.path.toLowerCase().endsWith(".gguf");
}
function isExpandableGguf(model: LocalModelInfo): boolean {
return (
!isDirectGguf(model) &&
(hasGgufSuffix(model.id) ||
hasGgufSuffix(model.display_name) ||
hasGgufSuffix(model.model_id))
);
}
function sourceLabel(model: LocalModelInfo): string {
switch (model.source) {
case "models_dir":
return "Models";
case "hf_cache":
return "HF cache";
case "lmstudio":
return "LM Studio";
case "custom":
return "Custom folder";
default:
return "Local";
}
}
type SelectedModelSummary = {
label: string;
source: string;
isGguf: boolean;
};
function getSelectedModelSummary(
value: string,
selectedModel: LocalModelInfo | null,
ggufVariant?: string | null,
): SelectedModelSummary {
if (!selectedModel) {
return {
label: value,
source: "Local model",
isGguf: Boolean(ggufVariant),
};
}
return {
label: getModelLabel(selectedModel),
source: sourceLabel(selectedModel),
isGguf: isDirectGguf(selectedModel) || isExpandableGguf(selectedModel),
};
}
function LocalGgufVariantList({
repoId,
selectedVariant,
onSelect,
}: {
repoId: string;
selectedVariant?: string | null;
onSelect: (variant: string) => void;
}): ReactElement {
const [variants, setVariants] = useState<GgufVariantDetail[] | null>(null);
const [defaultVariant, setDefaultVariant] = useState<string | null>(null);
const [loading, setLoading] = useState(true);
const [error, setError] = useState<string | null>(null);
useEffect(() => {
let cancelled = false;
listGgufVariants(repoId)
.then((response) => {
if (cancelled) {
return;
}
setVariants(response.variants);
setDefaultVariant(response.default_variant);
})
.catch((err) => {
if (cancelled) {
return;
}
setError(
err instanceof Error ? err.message : "Failed to load variants.",
);
})
.finally(() => {
if (!cancelled) {
setLoading(false);
}
});
return () => {
cancelled = true;
};
}, [repoId]);
const sortedVariants = useMemo(() => {
if (!variants) {
return null;
}
return [...variants].sort((a, b) => {
if (a.quant === defaultVariant) {
return -1;
}
if (b.quant === defaultVariant) {
return 1;
}
if (a.downloaded !== b.downloaded) {
return a.downloaded ? -1 : 1;
}
return a.quant.localeCompare(b.quant);
});
}, [defaultVariant, variants]);
if (loading) {
return (
<div className="flex items-center gap-2 px-4 py-2 text-xs text-muted-foreground">
<Spinner className="size-3" />
Loading quantizations...
</div>
);
}
if (error) {
return <div className="px-4 py-2 text-xs text-destructive">{error}</div>;
}
if (!sortedVariants || sortedVariants.length === 0) {
return (
<div className="px-4 py-2 text-xs text-muted-foreground">
No GGUF quantizations found for this model.
</div>
);
}
return (
<div className="ml-6 mt-1 rounded-lg bg-muted/25 p-1.5">
<div className="mb-1 px-2 text-[10px] font-medium uppercase tracking-wide text-muted-foreground">
Quantization
</div>
<div className="space-y-0.5">
{sortedVariants.map((variant) => {
const selected = selectedVariant === variant.quant;
return (
<button
key={variant.filename}
type="button"
onClick={() => onSelect(variant.quant)}
className={cn(
"flex w-full items-center gap-2 rounded-md px-2 py-1.5 text-left text-xs transition-colors hover:bg-muted/60",
selected && "bg-background text-foreground shadow-sm",
)}
>
<span className="min-w-0 flex-1 truncate font-mono">
{variant.quant}
</span>
{variant.quant === defaultVariant ? (
<Badge variant="secondary" className="h-4 px-1.5 text-[10px]">
recommended
</Badge>
) : null}
{variant.downloaded ? (
<Badge variant="outline" className="h-4 px-1.5 text-[10px]">
ready
</Badge>
) : null}
</button>
);
})}
</div>
</div>
);
}
type SelectorTriggerProps = ComponentPropsWithoutRef<"button"> & {
value: string;
selectedModel: LocalModelInfo | null;
ggufVariant?: string | null;
inputId?: string;
disabled: boolean;
compact: boolean;
className?: string;
};
const SelectorTrigger = forwardRef<HTMLButtonElement, SelectorTriggerProps>(
function SelectorTrigger(
{
value,
selectedModel,
ggufVariant,
inputId,
disabled,
compact,
className,
...triggerProps
},
ref,
): ReactElement {
const selected = getSelectedModelSummary(value, selectedModel, ggufVariant);
return (
<button
{...triggerProps}
ref={ref}
id={inputId}
type="button"
disabled={disabled}
className={cn(
"nodrag flex w-full min-w-0 items-center gap-2 rounded-xl border border-border/70 bg-background px-3 text-left transition-colors hover:bg-muted/40 disabled:pointer-events-none disabled:opacity-60",
compact ? "min-h-8 py-1.5 text-xs" : "min-h-10 py-2 text-sm",
className,
)}
>
<span className="min-w-0 flex-1">
<span
className={cn(
"block truncate font-medium",
!selected.label && "text-muted-foreground",
)}
>
{selected.label || "Choose a local model"}
</span>
{compact ? null : (
<span className="mt-0.5 flex min-w-0 items-center gap-1.5 text-[11px] text-muted-foreground">
<span className="truncate">
{selected.label
? selected.source
: "Select from local and cached models"}
</span>
{selected.isGguf ? <span>GGUF</span> : null}
{ggufVariant ? (
<span className="truncate font-mono">{ggufVariant}</span>
) : null}
</span>
)}
</span>
{compact && ggufVariant ? (
<Badge
variant="secondary"
className="h-4 px-1.5 font-mono text-[10px]"
>
{ggufVariant}
</Badge>
) : null}
<ChevronDownIcon className="size-4 shrink-0 text-muted-foreground" />
</button>
);
},
);
function LocalModelRow({
model,
selected,
expanded,
probing,
ggufVariant,
onSelectModel,
onSelectVariant,
}: {
model: LocalModelInfo;
selected: boolean;
expanded: boolean;
probing: boolean;
ggufVariant?: string | null;
onSelectModel: (model: LocalModelInfo) => void;
onSelectVariant: (modelId: string, variant: string) => void;
}): ReactElement {
const expandable = isExpandableGguf(model);
const directGguf = isDirectGguf(model);
return (
<div>
<button
type="button"
disabled={probing}
onClick={() => onSelectModel(model)}
className={cn(
"flex w-full items-center gap-2 rounded-lg px-2.5 py-2.5 text-left text-sm transition-colors hover:bg-muted/50",
selected && "bg-muted/70 text-foreground ring-1 ring-border/70",
)}
>
{expandable ? (
expanded ? (
<ChevronDownIcon className="size-3.5 shrink-0 text-muted-foreground" />
) : (
<ChevronRightIcon className="size-3.5 shrink-0 text-muted-foreground" />
)
) : (
<span className="size-3.5 shrink-0" />
)}
<span className="min-w-0 flex-1">
<span className="block truncate font-medium">
{getModelLabel(model)}
</span>
<span className="mt-0.5 block truncate text-[11px] text-muted-foreground">
{model.id}
</span>
</span>
<span className="flex shrink-0 items-center gap-1">
{probing ? (
<Spinner className="size-3 text-muted-foreground" />
) : null}
{expandable || directGguf ? (
<Badge variant="secondary" className="h-4 px-1.5 text-[10px]">
GGUF
</Badge>
) : null}
<Badge variant="outline" className="h-4 px-1.5 text-[10px]">
{sourceLabel(model)}
</Badge>
</span>
</button>
{expanded ? (
<LocalGgufVariantList
repoId={model.id}
selectedVariant={selected ? ggufVariant : null}
onSelect={(variant) => onSelectVariant(model.id, variant)}
/>
) : null}
</div>
);
}
function LocalModelResults({
loading,
error,
models,
value,
ggufVariant,
expandedModelId,
probingVariantModelId,
onRefresh,
onSelectModel,
onSelectVariant,
}: {
loading: boolean;
error: string | null;
models: LocalModelInfo[];
value: string;
ggufVariant?: string | null;
expandedModelId: string | null;
probingVariantModelId: string | null;
onRefresh: () => void;
onSelectModel: (model: LocalModelInfo) => void;
onSelectVariant: (modelId: string, variant: string) => void;
}): ReactElement {
if (loading) {
return (
<div className="flex items-center gap-2 px-3 py-3 text-xs text-muted-foreground">
<Spinner className="size-3" />
Scanning local models...
</div>
);
}
if (error) {
return (
<div className="space-y-2 px-3 py-3 text-xs">
<p className="text-destructive">{error}</p>
<Button type="button" variant="outline" size="xs" onClick={onRefresh}>
Try again
</Button>
</div>
);
}
if (models.length === 0) {
return (
<div className="space-y-2 px-3 py-3 text-xs text-muted-foreground">
<p className="font-medium text-foreground">No local models found.</p>
<p>
Download a model or add a scan folder from Chat, then refresh this
list.
</p>
<Link
to="/chat"
className="inline-flex font-medium text-primary underline-offset-4 hover:underline"
>
Open Chat model picker
</Link>
</div>
);
}
return (
<div className="space-y-1">
{models.map((model) => (
<LocalModelRow
key={model.id}
model={model}
selected={model.id === value}
expanded={expandedModelId === model.id}
probing={probingVariantModelId === model.id}
ggufVariant={ggufVariant}
onSelectModel={onSelectModel}
onSelectVariant={onSelectVariant}
/>
))}
</div>
);
}
export function LocalRecipeModelSelector({
value,
ggufVariant,
onChange,
inputId,
disabled = false,
compact = false,
className,
}: LocalRecipeModelSelectorProps): ReactElement {
const [open, setOpen] = useState(false);
const [query, setQuery] = useState("");
const [models, setModels] = useState<LocalModelInfo[]>([]);
const [loading, setLoading] = useState(false);
const [error, setError] = useState<string | null>(null);
const [expandedModelId, setExpandedModelId] = useState<string | null>(null);
const [probingVariantModelId, setProbingVariantModelId] = useState<
string | null
>(null);
const [refreshKey, setRefreshKey] = useState(0);
const requestModelRefresh = useCallback(() => {
setLoading(true);
setError(null);
setRefreshKey((key) => key + 1);
}, []);
const handleOpenChange = useCallback(
(nextOpen: boolean) => {
setOpen(nextOpen);
if (nextOpen) {
requestModelRefresh();
}
},
[requestModelRefresh],
);
useEffect(() => {
if (!open || refreshKey < 0) {
return;
}
let cancelled = false;
listLocalModels()
.then((response) => {
if (cancelled) {
return;
}
setModels(response.models);
})
.catch((err) => {
if (cancelled) {
return;
}
setError(
err instanceof Error ? err.message : "Failed to list local models.",
);
})
.finally(() => {
if (!cancelled) {
setLoading(false);
}
});
return () => {
cancelled = true;
};
}, [open, refreshKey]);
const selectedModel = useMemo(
() => models.find((model) => model.id === value) ?? null,
[models, value],
);
const filteredModels = useMemo(() => {
const needle = normalizeForSearch(query.trim());
if (!needle) {
return models;
}
return models.filter((model) => {
const haystack = normalizeForSearch(
`${model.id} ${model.display_name} ${model.model_id ?? ""} ${model.path}`,
);
return haystack.includes(needle);
});
}, [models, query]);
const selectModel = useCallback(
async (model: LocalModelInfo) => {
if (isExpandableGguf(model)) {
setExpandedModelId((current) =>
current === model.id ? null : model.id,
);
return;
}
if (!isDirectGguf(model)) {
setProbingVariantModelId(model.id);
try {
const response = await listGgufVariants(model.id);
if (response.variants.length > 0) {
setExpandedModelId(model.id);
return;
}
} catch {
// Non-GGUF local models commonly have no variant endpoint. Fall
// through to regular selection so users can still choose them.
} finally {
setProbingVariantModelId(null);
}
}
onChange(model.id, null);
setOpen(false);
},
[onChange],
);
const selectVariant = useCallback(
(modelId: string, variant: string) => {
onChange(modelId, variant);
setOpen(false);
},
[onChange],
);
return (
<Popover open={open} onOpenChange={handleOpenChange}>
<PopoverTrigger asChild={true}>
<SelectorTrigger
value={value}
selectedModel={selectedModel}
ggufVariant={ggufVariant}
inputId={inputId}
disabled={disabled}
compact={compact}
className={className}
/>
</PopoverTrigger>
<PopoverContent
align="start"
sideOffset={6}
className="menu-soft-surface nodrag nowheel gap-0 overflow-hidden p-0"
style={{
width:
"min(max(var(--radix-popover-trigger-width), 34rem), calc(100vw - 1rem))",
}}
>
<div className="flex flex-col">
<div className="border-b border-border/60 p-2.5">
<div className="flex items-center gap-2">
<Input
value={query}
onChange={(event) => setQuery(event.target.value)}
placeholder="Filter local models"
className="h-8 flex-1"
autoFocus={true}
/>
<Button
type="button"
variant="ghost"
size="icon-sm"
onClick={requestModelRefresh}
aria-label="Refresh local models"
>
<RefreshCwIcon className="size-3.5" />
</Button>
</div>
</div>
<div
className="nowheel max-h-[min(24rem,calc(100vh-12rem))] overflow-y-auto overscroll-contain p-1.5"
onWheelCapture={(event) => event.stopPropagation()}
>
<LocalModelResults
loading={loading}
error={error}
models={filteredModels}
value={value}
ggufVariant={ggufVariant}
expandedModelId={expandedModelId}
probingVariantModelId={probingVariantModelId}
onRefresh={requestModelRefresh}
onSelectModel={selectModel}
onSelectVariant={selectVariant}
/>
</div>
</div>
</PopoverContent>
</Popover>
);
}

View file

@ -1,12 +1,12 @@
// SPDX-License-Identifier: AGPL-3.0-only
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
import { Checkbox } from "@/components/ui/checkbox";
import {
Collapsible,
CollapsibleContent,
CollapsibleTrigger,
} from "@/components/ui/collapsible";
import { Checkbox } from "@/components/ui/checkbox";
import {
Combobox,
ComboboxContent,
@ -22,6 +22,7 @@ import type { ModelConfig } from "../../types";
import { CollapsibleSectionTriggerButton } from "../shared/collapsible-section-trigger";
import { FieldLabel } from "../shared/field-label";
import { NameField } from "../shared/name-field";
import { LocalRecipeModelSelector } from "./local-recipe-model-selector";
type ModelConfigDialogProps = {
config: ModelConfig;
@ -45,6 +46,7 @@ export function ModelConfigDialog({
const maxTokensId = `${config.id}-max-tokens`;
const timeoutId = `${config.id}-timeout`;
const extraBodyId = `${config.id}-inference-extra-body`;
const skipHealthCheckId = `${config.id}-skip-health-check`;
const providerAnchorRef = useRef<HTMLDivElement>(null);
const providerInputRef = useRef(config.provider);
// Sync providerInputRef with the current provider value. Updating a ref in
@ -61,16 +63,25 @@ export function ModelConfigDialog({
onUpdate({ [key]: value } as Partial<ModelConfig>);
};
// Apply provider selection while keeping the local-provider model autofill
// consistent across both dropdown selection and free-typed + blur input.
// Apply provider selection while clearing model identifiers that only make
// sense for the previous provider locality.
const applyProviderChange = (selectedProvider: string) => {
const isLocal = localProviderNames.has(selectedProvider);
if (isLocal && !config.model.trim()) {
onUpdate({ provider: selectedProvider, model: "local" });
const nextIsLocal = localProviderNames.has(selectedProvider);
if (isLinkedToLocal !== nextIsLocal) {
onUpdate({
provider: selectedProvider,
model: "",
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: undefined,
});
return;
}
if (!isLocal && config.model === "local") {
onUpdate({ provider: selectedProvider, model: "" });
if (!nextIsLocal) {
onUpdate({
provider: selectedProvider,
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: undefined,
});
return;
}
updateField("provider", selectedProvider);
@ -88,8 +99,8 @@ export function ModelConfigDialog({
Set up one reusable model choice for your AI steps
</p>
<p className="mt-1 text-xs text-muted-foreground">
Choose the provider connection, enter the exact model ID, then save any
generation defaults you want to reuse.
Choose the provider connection, enter the exact model ID, then save
any generation defaults you want to reuse.
</p>
</div>
<div className="grid gap-1.5">
@ -144,15 +155,48 @@ export function ModelConfigDialog({
<FieldLabel
label="Model ID"
htmlFor={modelId}
hint={isLinkedToLocal ? "Uses the model loaded in Chat. Any value works here." : "The exact model name sent to the connection."}
/>
<Input
id={modelId}
className="nodrag"
placeholder={isLinkedToLocal ? "local" : "gpt-4o-mini"}
value={config.model}
onChange={(event) => updateField("model", event.target.value)}
hint={
isLinkedToLocal
? "Choose the local model Recipes should load before Run or Validate."
: "The exact model name sent to the connection."
}
/>
{isLinkedToLocal ? (
<LocalRecipeModelSelector
inputId={modelId}
value={
config.model.trim().toLowerCase() === "local" ? "" : config.model
}
ggufVariant={config.gguf_variant}
onChange={(model, variant) =>
onUpdate({
model,
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: variant ?? undefined,
})
}
/>
) : (
<Input
id={modelId}
className="nodrag"
placeholder="gpt-4o-mini"
value={config.model}
onChange={(event) =>
onUpdate({
model: event.target.value,
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: undefined,
})
}
/>
)}
{isLinkedToLocal ? (
<p className="text-xs text-muted-foreground">
Recipes will load this model automatically. GGUF quantization is
saved with the preset.
</p>
) : null}
</div>
<div className="grid gap-3">
<div className="space-y-1">
@ -250,8 +294,12 @@ export function ModelConfigDialog({
}
/>
</div>
<label className="flex items-center gap-2 text-xs font-semibold uppercase text-muted-foreground">
<label
htmlFor={skipHealthCheckId}
className="flex items-center gap-2 text-xs font-semibold uppercase text-muted-foreground"
>
<Checkbox
id={skipHealthCheckId}
checked={config.skip_health_check ?? false}
onCheckedChange={(value) =>
updateField("skip_health_check", Boolean(value))

View file

@ -1,13 +1,14 @@
// SPDX-License-Identifier: AGPL-3.0-only
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { GithubIcon, PlayCircleIcon } from "@hugeicons/core-free-icons";
import { HugeiconsIcon } from "@hugeicons/react";
import { type ReactElement, useEffect, useMemo, useState } from "react";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { FieldLabel } from "../dialogs/shared/field-label";
import { LocalRecipeModelSelector } from "../dialogs/models/local-recipe-model-selector";
import { GithubRepoSeedForm } from "../dialogs/seed/seed-dialog";
import { FieldLabel } from "../dialogs/shared/field-label";
import type { ModelConfig, NodeConfig, SeedConfig } from "../types";
type GithubCrawlerEasyViewProps = {
@ -44,6 +45,21 @@ export function GithubCrawlerEasyView({
) ?? null,
[configs],
);
const localProviderNames = useMemo(() => {
const names = new Set<string>();
for (const config of Object.values(configs)) {
if (config.kind === "model_provider" && config.is_local === true) {
names.add(config.name);
}
}
return names;
}, [configs]);
const isModelLinkedToLocal = modelConfig
? localProviderNames.has(modelConfig.provider)
: false;
const modelValue = modelConfig?.model ?? "";
const localModelValue =
modelValue.trim().toLowerCase() === "local" ? "" : modelValue;
// Local buffer for the Rows input so the user can hold transient invalid
// state (empty while backspacing, partial digits, etc.) without the parent
@ -52,17 +68,40 @@ export function GithubCrawlerEasyView({
// blur we clamp back to a sane default if the user left it empty.
const [rowsText, setRowsText] = useState(String(rows));
useEffect(() => {
// eslint-disable-next-line react-hooks/set-state-in-effect -- keep the draft input in sync when the parent resets rows.
setRowsText(String(rows));
}, [rows]);
const handleSeedUpdate = (patch: Partial<SeedConfig>): void => {
if (!seedConfig) return;
if (!seedConfig) {
return;
}
updateConfig(seedConfig.id, patch);
};
const handleModelChange = (value: string): void => {
if (!modelConfig) return;
updateConfig(modelConfig.id, { model: value });
if (!modelConfig) {
return;
}
updateConfig(modelConfig.id, {
model: value,
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: undefined,
});
};
const handleLocalModelChange = (
model: string,
variant?: string | null,
): void => {
if (!modelConfig) {
return;
}
updateConfig(modelConfig.id, {
model,
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: variant ?? undefined,
});
};
if (!seedConfig) {
@ -96,9 +135,8 @@ export function GithubCrawlerEasyView({
<h2 className="text-base font-semibold">GitHub Crawler</h2>
<p className="text-xs text-muted-foreground">
Crawl real GitHub issues and PRs and turn each thread into a{" "}
<code>{"{User, Assistant}"}</code> training pair.
Defaults use the server's <code>GH_TOKEN</code> env var and the
bundled local model.
<code>{"{User, Assistant}"}</code> training pair. Defaults use the
server's <code>GH_TOKEN</code> env var and the bundled local model.
</p>
</div>
</div>
@ -149,15 +187,28 @@ export function GithubCrawlerEasyView({
<div className="grid gap-1.5">
<FieldLabel
label="Model"
hint="OpenAI-compatible model id. Local GGUFs run on the bundled llama-server."
/>
<Input
className="nodrag font-mono text-xs"
value={modelConfig?.model ?? ""}
onChange={(event) => handleModelChange(event.target.value)}
placeholder="unsloth/gemma-4-E2B-it-GGUF"
disabled={!modelConfig}
hint={
isModelLinkedToLocal
? "Choose the local model this recipe should load."
: "OpenAI-compatible model id."
}
/>
{isModelLinkedToLocal ? (
<LocalRecipeModelSelector
value={localModelValue}
ggufVariant={modelConfig?.gguf_variant}
onChange={handleLocalModelChange}
disabled={!modelConfig}
/>
) : (
<Input
className="nodrag font-mono text-xs"
value={modelValue}
onChange={(event) => handleModelChange(event.target.value)}
placeholder="unsloth/gemma-4-E2B-it-GGUF"
disabled={!modelConfig}
/>
)}
</div>
</div>
</section>

View file

@ -40,6 +40,11 @@ type TrackRecipeExecutionParams = {
onPreviewSuccess?: () => void;
};
export type TrackRecipeExecutionResult = {
success: boolean;
terminal: boolean;
};
function isTerminalStatus(status: RecipeExecutionStatus): boolean {
return status === "completed" || status === "error" || status === "cancelled";
}
@ -53,7 +58,8 @@ function normalizeCompletedProgress(input: {
} {
const { latestExecution, rows } = input;
const progressTotal =
typeof latestExecution.progress?.total === "number" && latestExecution.progress.total > 0
typeof latestExecution.progress?.total === "number" &&
latestExecution.progress.total > 0
? latestExecution.progress.total
: latestExecution.rows > 0
? latestExecution.rows
@ -92,7 +98,7 @@ export async function trackRecipeExecution({
onUpsert,
onSetPreviewErrors,
onPreviewSuccess,
}: TrackRecipeExecutionParams): Promise<boolean> {
}: TrackRecipeExecutionParams): Promise<TrackRecipeExecutionResult> {
let done = false;
let lastStatus: RecipeExecutionStatus = initialExecution.status;
let completedEventPayload: Record<string, unknown> | null = null;
@ -124,7 +130,9 @@ export async function trackRecipeExecution({
}
const eventType =
typeof event.payload.type === "string" ? event.payload.type : event.event;
typeof event.payload.type === "string"
? event.payload.type
: event.event;
if (eventType === "job.started") {
latestExecution = {
@ -163,7 +171,7 @@ export async function trackRecipeExecution({
error:
typeof event.payload.error === "string"
? event.payload.error
: latestExecution.error ?? `${label} failed.`,
: (latestExecution.error ?? `${label} failed.`),
};
onUpsert(latestExecution);
return;
@ -178,6 +186,19 @@ export async function trackRecipeExecution({
return;
}
if (eventType === "job.cancelled") {
lastStatus = "cancelled";
done = true;
latestExecution = {
...latestExecution,
status: "cancelled",
finishedAt: Date.now(),
error: latestExecution.error ?? "Run cancelled.",
};
onUpsert(latestExecution);
return;
}
if (changed) {
onUpsert(latestExecution);
}
@ -189,6 +210,9 @@ export async function trackRecipeExecution({
try {
while (!done) {
const status = await getRecipeJobStatus(jobId);
if (done && isTerminalStatus(lastStatus)) {
break;
}
const mappedStatus = mapJobStatus(status.status);
lastStatus = mappedStatus;
latestExecution = applyExecutionStatusSnapshot(latestExecution, status);
@ -200,18 +224,19 @@ export async function trackRecipeExecution({
}
}
} catch (error) {
const message = toErrorMessage(error, `${label} failed.`);
latestExecution = {
...latestExecution,
status: "error",
error: message,
finishedAt: Date.now(),
};
onUpsert(latestExecution);
if (notify) {
toastError(`${label} failed`, message);
const terminal = isTerminalStatus(lastStatus);
if (!terminal) {
const message = toErrorMessage(error, `${label} failed.`);
latestExecution = {
...latestExecution,
error: message,
};
onUpsert(latestExecution);
if (notify) {
toastError(`${label} failed`, message);
}
return { success: false, terminal: false };
}
return false;
} finally {
eventsAbortController.abort();
}
@ -220,7 +245,10 @@ export async function trackRecipeExecution({
for (let attempt = 0; attempt < 3; attempt += 1) {
try {
const finalStatus = await getRecipeJobStatus(jobId);
latestExecution = applyExecutionStatusSnapshot(latestExecution, finalStatus);
latestExecution = applyExecutionStatusSnapshot(
latestExecution,
finalStatus,
);
} catch {
break;
}
@ -229,19 +257,20 @@ export async function trackRecipeExecution({
}
}
const eventAnalysis = completedEventPayload
? completedEventPayload["analysis"]
: null;
const eventDataset = completedEventPayload
? completedEventPayload["dataset"]
: null;
const completedPayload = completedEventPayload as Record<
string,
unknown
> | null;
const eventAnalysis = completedPayload ? completedPayload.analysis : null;
const eventDataset = completedPayload ? completedPayload.dataset : null;
const eventProcessorArtifacts =
completedEventPayload &&
typeof completedEventPayload["processor_artifacts"] === "object" &&
completedEventPayload["processor_artifacts"] !== null
? (completedEventPayload["processor_artifacts"] as Record<string, unknown>)
completedPayload &&
typeof completedPayload.processor_artifacts === "object" &&
completedPayload.processor_artifacts !== null
? (completedPayload.processor_artifacts as Record<string, unknown>)
: null;
const shouldFetchPreviewDataset = kind === "preview" && !Array.isArray(eventDataset);
const shouldFetchPreviewDataset =
kind === "preview" && !Array.isArray(eventDataset);
const shouldFetchAnalysis =
!completedEventPayload ||
typeof eventAnalysis !== "object" ||
@ -262,9 +291,7 @@ export async function trackRecipeExecution({
? normalizeAnalysis(analysisResult.value)
: latestExecution.analysis;
const datasetResponse =
datasetResult.status === "fulfilled"
? datasetResult.value
: null;
datasetResult.status === "fulfilled" ? datasetResult.value : null;
const dataset = datasetResponse
? normalizeDatasetRows(datasetResponse.dataset)
: latestExecution.dataset;
@ -272,7 +299,10 @@ export async function trackRecipeExecution({
datasetResponse && typeof datasetResponse.total === "number"
? datasetResponse.total
: latestExecution.datasetTotal;
const completedProgress = normalizeCompletedProgress({ latestExecution, rows });
const completedProgress = normalizeCompletedProgress({
latestExecution,
rows,
});
latestExecution = {
...latestExecution,
@ -285,7 +315,8 @@ export async function trackRecipeExecution({
datasetPage: 1,
datasetPageSize: DATASET_PAGE_SIZE,
error: null,
processor_artifacts: eventProcessorArtifacts ?? latestExecution.processor_artifacts,
processor_artifacts:
eventProcessorArtifacts ?? latestExecution.processor_artifacts,
finishedAt: latestExecution.finishedAt ?? Date.now(),
};
onUpsert(latestExecution);
@ -299,7 +330,7 @@ export async function trackRecipeExecution({
toastSuccess("Full run completed.");
}
}
return true;
return { success: true, terminal: true };
}
if (lastStatus === "cancelled") {
@ -313,7 +344,7 @@ export async function trackRecipeExecution({
if (notify) {
toastError(`${label} cancelled`, "The execution was cancelled.");
}
return false;
return { success: false, terminal: true };
}
latestExecution = {
@ -326,5 +357,5 @@ export async function trackRecipeExecution({
if (notify) {
toastError(`${label} failed`, latestExecution.error ?? "Execution failed.");
}
return false;
return { success: false, terminal: true };
}

View file

@ -1,14 +1,11 @@
// SPDX-License-Identifier: AGPL-3.0-only
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
import { useCallback, useEffect, useState } from "react";
import { useShallow } from "zustand/react/shallow";
import { getInferenceStatus, loadModel } from "@/features/chat";
import { toast } from "@/lib/toast";
import { toastError } from "@/shared/toast";
import {
getInferenceStatus,
loadModel,
} from "@/features/chat/api/chat-api";
import { useCallback, useEffect, useState } from "react";
import { useShallow } from "zustand/react/shallow";
import {
cancelRecipeJob,
createRecipeJob,
@ -23,8 +20,8 @@ import type {
import {
DATASET_PAGE_SIZE,
executionLabel,
normalizeRunName,
normalizeDatasetRows,
normalizeRunName,
toErrorMessage,
withExecutionDefaults,
} from "../executions/execution-helpers";
@ -32,84 +29,243 @@ import {
findResumableExecution,
loadSortedRecipeExecutions,
} from "../executions/hydration";
import { createBaseExecutionRecord } from "../executions/runtime";
import {
buildExecutionPayload,
sanitizeExecutionRows,
} from "../executions/run-settings";
import { createBaseExecutionRecord } from "../executions/runtime";
import { trackRecipeExecution } from "../executions/tracker";
import {
type RecipeRunSettings,
useRecipeExecutionsStore,
} from "../stores/recipe-executions";
import type { RecipePayload, RecipePayloadResult } from "../utils/payload/types";
import type {
RecipePayload,
RecipePayloadResult,
} from "../utils/payload/types";
/**
* Auto-load the local model before running a recipe that uses it.
*
* Looks at payload.recipe.model_providers for any provider with is_local=true,
* finds the bound model_configs and asks the backend to load whichever model
* the first local-bound model_config points at. Skips when the inference
* server already has that exact model active. This removes the "open /chat
* first" prerequisite that users kept tripping on.
*/
async function ensureLocalModelLoaded(
payload: RecipePayload,
): Promise<string | null> {
const GGUF_MODEL_PATTERN = /gguf/i;
function collectUsedLlmModelAliases(payload: RecipePayload): Set<string> {
const columns = Array.isArray(payload.recipe.columns)
? payload.recipe.columns
: [];
const aliases = new Set<string>();
for (const column of columns) {
const columnType = column.column_type;
if (typeof columnType !== "string" || !columnType.startsWith("llm-")) {
continue;
}
const alias = column.model_alias;
if (typeof alias === "string" && alias.trim()) {
aliases.add(alias.trim());
}
}
return aliases;
}
type LocalModelSelection = {
target: string;
ggufVariant: string;
aliases: string[];
};
type LocalModelLoadPlan =
| { selection: LocalModelSelection; error: null; legacyAliases?: never }
| { selection: null; error: string; legacyAliases?: never }
| { selection: null; error: null; legacyAliases: string[] };
type RestorableLocalModelSnapshot = {
selection: LocalModelSelection | null;
unrestorableLabel: string | null;
};
function getLocalProviderNames(payload: RecipePayload): Set<string> {
const providers = Array.isArray(payload.recipe.model_providers)
? (payload.recipe.model_providers as Array<Record<string, unknown>>)
? (payload.recipe.model_providers as Record<string, unknown>[])
: [];
const localProviderNames = new Set<string>();
for (const p of providers) {
if (p.is_local === true && typeof p.name === "string") {
localProviderNames.add(p.name);
for (const provider of providers) {
if (provider.is_local === true && typeof provider.name === "string") {
localProviderNames.add(provider.name);
}
}
if (localProviderNames.size === 0) {
return null;
return localProviderNames;
}
function findUsedLocalModelConfigs(
payload: RecipePayload,
localProviderNames: Set<string>,
): Record<string, unknown>[] {
const usedAliases = collectUsedLlmModelAliases(payload);
if (usedAliases.size === 0) {
return [];
}
const modelConfigs = Array.isArray(payload.recipe.model_configs)
? (payload.recipe.model_configs as Array<Record<string, unknown>>)
? payload.recipe.model_configs
: [];
const boundConfig = modelConfigs.find(
(c) => typeof c.provider === "string" && localProviderNames.has(c.provider),
);
return modelConfigs.filter((config) => {
const provider = config.provider;
const alias = config.alias;
return (
typeof provider === "string" &&
localProviderNames.has(provider) &&
typeof alias === "string" &&
usedAliases.has(alias)
);
});
}
function readLocalModelSelection(
boundConfig: Record<string, unknown>,
): LocalModelLoadPlan {
const alias =
typeof boundConfig.alias === "string" ? boundConfig.alias : "local model";
const target =
typeof boundConfig?.model === "string" ? boundConfig.model.trim() : "";
typeof boundConfig.model === "string" ? boundConfig.model.trim() : "";
const ggufVariant =
typeof boundConfig.gguf_variant === "string"
? boundConfig.gguf_variant.trim()
: "";
if (!target) {
return null;
return {
selection: null,
error: `Model config ${alias}: choose a local model before validating or running this recipe.`,
};
}
if (target.toLowerCase() === "local") {
return { selection: null, error: null, legacyAliases: [alias] };
}
return { selection: { target, ggufVariant, aliases: [alias] }, error: null };
}
function getLocalModelLoadPlan(
boundConfigs: Record<string, unknown>[],
): LocalModelLoadPlan | null {
const selections = new Map<string, LocalModelSelection>();
const legacyAliases: string[] = [];
for (const boundConfig of boundConfigs) {
const next = readLocalModelSelection(boundConfig);
if (next.error) {
return next;
}
if (next.legacyAliases) {
legacyAliases.push(...next.legacyAliases);
continue;
}
const selection = next.selection;
if (!selection) {
continue;
}
const key = `${selection.target.toLowerCase()}\u0000${selection.ggufVariant}`;
const existing = selections.get(key);
if (existing) {
existing.aliases.push(...selection.aliases);
continue;
}
selections.set(key, selection);
}
if (legacyAliases.length > 0 && selections.size > 0) {
const aliases = [
...legacyAliases,
...[...selections.values()].flatMap((selection) => selection.aliases),
].join(", ");
return {
selection: null,
error: `Recipes found mixed legacy and selected local models. Reselect the same concrete local model for: ${aliases}.`,
};
}
if (legacyAliases.length > 0) {
return { selection: null, error: null, legacyAliases };
}
if (selections.size > 1) {
const aliases = [...selections.values()]
.flatMap((selection) => selection.aliases)
.join(", ");
return {
selection: null,
error: `Recipes supports one active local model per run. Select the same local model and GGUF variant for: ${aliases}.`,
};
}
const selection = [...selections.values()][0];
return selection ? { selection, error: null } : null;
}
function isDirectGgufTarget(target: string): boolean {
return target.toLowerCase().endsWith(".gguf");
}
function localSelectionMatchesActive(input: {
target: string;
ggufVariant: string;
activeModel: string | null | undefined;
activeVariant: string;
}): boolean {
const { target, ggufVariant, activeModel, activeVariant } = input;
if (!activeModel || activeModel.toLowerCase() !== target.toLowerCase()) {
return false;
}
return (
activeVariant === ggufVariant ||
(isDirectGgufTarget(target) && !ggufVariant)
);
}
async function isLocalModelAlreadyLoaded(
selection: LocalModelSelection,
): Promise<boolean> {
const { target, ggufVariant } = selection;
try {
const status = await getInferenceStatus();
if (
status.active_model &&
status.active_model.toLowerCase() === target.toLowerCase()
) {
return null;
}
return localSelectionMatchesActive({
target,
ggufVariant,
activeModel: status.model_identifier ?? status.active_model,
activeVariant: status.gguf_variant?.trim() ?? "",
});
} catch {
// Fall through to load attempt; the backend will re-error if needed.
return false;
}
}
const toastId = toast.loading(`Loading ${target}`, {
async function loadLocalModelSelection(
selection: LocalModelSelection,
): Promise<string | null> {
const { target, ggufVariant } = selection;
const modelLabel = ggufVariant ? `${target} (${ggufVariant})` : target;
const toastId = toast.loading(`Loading ${modelLabel}...`, {
description: "Starting the local inference server for this recipe.",
});
try {
const isGguf = /gguf/i.test(target);
const isGguf = GGUF_MODEL_PATTERN.test(target) || Boolean(ggufVariant);
await loadModel({
// biome-ignore lint/style/useNamingConvention: api schema
model_path: target,
// biome-ignore lint/style/useNamingConvention: api schema
hf_token: null,
// biome-ignore lint/style/useNamingConvention: api schema
max_seq_length: isGguf ? 0 : 4096,
// biome-ignore lint/style/useNamingConvention: api schema
load_in_4bit: true,
// biome-ignore lint/style/useNamingConvention: api schema
is_lora: false,
gguf_variant: null,
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: ggufVariant || null,
// biome-ignore lint/style/useNamingConvention: api schema
trust_remote_code: false,
// biome-ignore lint/style/useNamingConvention: api schema
chat_template_override: null,
// biome-ignore lint/style/useNamingConvention: api schema
cache_type_kv: null,
// biome-ignore lint/style/useNamingConvention: api schema
speculative_type: null,
});
toast.success(`Loaded ${target}`, { id: toastId, duration: 2000 });
toast.success(`Loaded ${modelLabel}`, { id: toastId, duration: 2000 });
return null;
} catch (error) {
toast.dismiss(toastId);
@ -117,6 +273,147 @@ async function ensureLocalModelLoaded(
}
}
function getLocalModelLoadPlanForPayload(
payload: RecipePayload,
): LocalModelLoadPlan | null {
const localProviderNames = getLocalProviderNames(payload);
if (localProviderNames.size === 0) {
return null;
}
const boundConfigs = findUsedLocalModelConfigs(payload, localProviderNames);
return getLocalModelLoadPlan(boundConfigs);
}
async function getActiveLocalModelSelection(): Promise<LocalModelSelection | null> {
try {
const status = await getInferenceStatus();
const target = status.active_model?.trim();
if (!target) {
return null;
}
return {
target,
ggufVariant: status.gguf_variant?.trim() ?? "",
aliases: ["previous Chat model"],
};
} catch {
return null;
}
}
async function getRestorableActiveLocalModelSelection(): Promise<RestorableLocalModelSnapshot> {
try {
const status = await getInferenceStatus();
const activeLabel = status.active_model?.trim() ?? null;
const target = (
status.model_identifier ?? (status.is_gguf ? null : status.active_model)
)?.trim();
if (!target) {
return {
selection: null,
unrestorableLabel: activeLabel,
};
}
return {
selection: {
target,
ggufVariant: status.gguf_variant?.trim() ?? "",
aliases: ["previous Chat model"],
},
unrestorableLabel: null,
};
} catch {
return { selection: null, unrestorableLabel: null };
}
}
function isSameLocalModelSelection(
left: LocalModelSelection | null,
right: LocalModelSelection,
): boolean {
return Boolean(
left &&
left.target.toLowerCase() === right.target.toLowerCase() &&
left.ggufVariant === right.ggufVariant,
);
}
async function ensureLocalModelLoaded(
payload: RecipePayload,
): Promise<string | null> {
const loadPlan = getLocalModelLoadPlanForPayload(payload);
if (!loadPlan) {
return null;
}
if (loadPlan.legacyAliases) {
const activeSelection = await getActiveLocalModelSelection();
return activeSelection
? null
: `Existing recipe uses legacy local model for ${loadPlan.legacyAliases.join(", ")}. Select a concrete local model or load one in Chat.`;
}
if (!loadPlan.selection) {
return loadPlan.error;
}
if (await isLocalModelAlreadyLoaded(loadPlan.selection)) {
return null;
}
return loadLocalModelSelection(loadPlan.selection);
}
async function prepareLocalModelForRun(payload: RecipePayload): Promise<{
error: string | null;
restorePrevious: (() => Promise<void>) | null;
}> {
const loadPlan = getLocalModelLoadPlanForPayload(payload);
if (!loadPlan) {
return { error: null, restorePrevious: null };
}
if (loadPlan.legacyAliases) {
const activeSelection = await getActiveLocalModelSelection();
return activeSelection
? { error: null, restorePrevious: null }
: {
error: `Existing recipe uses legacy local model for ${loadPlan.legacyAliases.join(", ")}. Select a concrete local model or load one in Chat.`,
restorePrevious: null,
};
}
if (!loadPlan.selection) {
return { error: loadPlan.error, restorePrevious: null };
}
if (await isLocalModelAlreadyLoaded(loadPlan.selection)) {
return { error: null, restorePrevious: null };
}
const previousSnapshot = await getRestorableActiveLocalModelSelection();
const previousSelection = previousSnapshot.selection;
const error = await loadLocalModelSelection(loadPlan.selection);
if (error) {
return { error, restorePrevious: null };
}
if (isSameLocalModelSelection(previousSelection, loadPlan.selection)) {
return { error: null, restorePrevious: null };
}
return {
error: null,
restorePrevious: previousSelection
? async () => {
const restoreError = await loadLocalModelSelection(previousSelection);
if (restoreError) {
toastError("Could not restore previous local model", restoreError);
}
}
: previousSnapshot.unrestorableLabel
? () => {
toast.warning("Previous local model was not restored", {
description: `${previousSnapshot.unrestorableLabel} was selected from a native file path. Reopen it in Chat to continue with that model.`,
});
return Promise.resolve();
}
: null,
};
}
type UseRecipeExecutionsParams = {
recipeId: string;
currentSignature: string;
@ -161,7 +458,11 @@ type UseRecipeExecutionsResult = {
};
function formatValidationMessages(input: {
errors: Array<{ message: string; path?: string | null; code?: string | null }>;
errors: Array<{
message: string;
path?: string | null;
code?: string | null;
}>;
}): string[] {
return input.errors.map((item) => {
const path = item.path?.trim();
@ -249,7 +550,8 @@ export function useRecipeExecutions({
(record: RecipeExecutionRecord): void => {
const normalizedRecord = withExecutionDefaults(record);
upsertExecution(normalizedRecord);
void saveRecipeExecution(normalizedRecord).catch((error) => {
saveRecipeExecution(normalizedRecord).catch((error) => {
// biome-ignore lint/suspicious/noConsole: background persistence failures should not interrupt the UI
console.error("Save recipe execution failed:", error);
});
},
@ -287,7 +589,7 @@ export function useRecipeExecutions({
return;
}
void trackRecipeExecution({
trackRecipeExecution({
label: executionLabel(resumable.kind),
kind: resumable.kind,
rows: resumable.rows,
@ -299,11 +601,12 @@ export function useRecipeExecutions({
onPreviewSuccess,
});
} catch (error) {
// biome-ignore lint/suspicious/noConsole: hydration failures are non-blocking diagnostics
console.error("Load recipe executions failed:", error);
}
}
void hydrate();
hydrate();
return () => {
cancelled = true;
@ -344,9 +647,11 @@ export function useRecipeExecutions({
rows: number;
settings: RecipeRunSettings;
runName: string | null;
restorePrevious?: (() => Promise<void>) | null;
}): Promise<boolean> => {
const { kind, payload, rows, settings, runName } = input;
const setLoading = kind === "preview" ? setPreviewLoading : setFullLoading;
const { kind, payload, rows, settings, runName, restorePrevious } = input;
const setLoading =
kind === "preview" ? setPreviewLoading : setFullLoading;
const label = executionLabel(kind);
setLoading(true);
@ -362,6 +667,8 @@ export function useRecipeExecutions({
onExecutionStart?.();
setRunDialogOpen(false);
let jobCreated = false;
let shouldRestorePrevious = false;
try {
const jobPayload = buildExecutionPayload({
payload,
@ -371,13 +678,14 @@ export function useRecipeExecutions({
runName,
});
const createdJob = await createRecipeJob(jobPayload);
jobCreated = true;
const executionWithJob = {
...baseExecution,
jobId: createdJob.job_id,
};
upsertAndPersist(executionWithJob);
return await trackRecipeExecution({
const tracked = await trackRecipeExecution({
label,
kind,
rows,
@ -388,6 +696,8 @@ export function useRecipeExecutions({
onSetPreviewErrors: setRunErrors,
onPreviewSuccess,
});
shouldRestorePrevious = tracked.terminal;
return tracked.success;
} catch (error) {
const message = toErrorMessage(error, `${label} request failed.`);
upsertAndPersist({
@ -398,8 +708,14 @@ export function useRecipeExecutions({
});
setRunErrors([message]);
toastError(`${label} failed`, message);
if (!jobCreated) {
shouldRestorePrevious = true;
}
return false;
} finally {
if (shouldRestorePrevious && restorePrevious) {
await restorePrevious();
}
setLoading(false);
}
},
@ -416,6 +732,48 @@ export function useRecipeExecutions({
],
);
const prepareLocalModelForExecution = useCallback(
async (
payload: RecipePayload,
): Promise<(() => Promise<void>) | null | false> => {
const { error, restorePrevious } = await prepareLocalModelForRun(payload);
if (!error) {
return restorePrevious;
}
setRunErrors([error]);
toastError("Local model failed to load", error);
return false;
},
[setRunErrors],
);
const validateExecutionPayload = useCallback(
async (
executionPayload: Parameters<typeof validateRecipe>[0],
): Promise<boolean> => {
try {
const validation = await validateRecipe(executionPayload);
if (validation.valid) {
return true;
}
const errors = formatValidationMessages({
errors: validation.errors,
});
const fallback = validation.raw_detail ?? "Validation failed.";
const nextErrors = errors.length > 0 ? errors : [fallback];
setRunErrors(nextErrors);
toastError("Validation failed", nextErrors[0]);
return false;
} catch (error) {
const message = toErrorMessage(error, "Validation failed.");
setRunErrors([message]);
toastError("Validation failed", message);
return false;
}
},
[setRunErrors],
);
const runWithValidation = useCallback(
async (
kind: RecipeExecutionKind,
@ -435,20 +793,11 @@ export function useRecipeExecutions({
return false;
}
// Flip to the Runs pane BEFORE we run ensureLocalModelLoaded + validate.
// Validation re-crawls the seed (multiple seconds for the github_repo
// reader) and the user otherwise stares at a "Running..." button with
// nothing else changing. runExecution() later no-ops this callback if
// the view has already been flipped, so we fire it once here.
// Flip to the Runs pane before validation starts. Validation can re-crawl
// the seed (multiple seconds for the github_repo reader), and runExecution()
// later no-ops this callback if the view has already been flipped.
onExecutionStart?.();
const localLoadError = await ensureLocalModelLoaded(payload);
if (localLoadError) {
setRunErrors([localLoadError]);
toastError("Local model failed to load", localLoadError);
return false;
}
const normalizedRows = sanitizeExecutionRows(rows, kind);
const executionPayload = buildExecutionPayload({
payload,
@ -458,20 +807,17 @@ export function useRecipeExecutions({
runName,
});
try {
const validation = await validateRecipe(executionPayload);
if (!validation.valid) {
const errors = formatValidationMessages({ errors: validation.errors });
const fallback = validation.raw_detail ?? "Validation failed.";
const nextErrors = errors.length > 0 ? errors : [fallback];
setRunErrors(nextErrors);
toastError("Validation failed", nextErrors[0]);
return false;
}
} catch (error) {
const message = toErrorMessage(error, "Validation failed.");
setRunErrors([message]);
toastError("Validation failed", message);
if (!(await validateExecutionPayload(executionPayload))) {
return false;
}
// Recipe and Chat share one singleton local inference backend. This
// direct load is a point-in-time handoff to job creation, not a lease:
// if Chat swaps models after this succeeds, the backend will reject or
// run against the active backend state. A future generation token should
// be validated across this load and the `/jobs` loaded-model gate.
const restorePrevious = await prepareLocalModelForExecution(payload);
if (restorePrevious === false) {
return false;
}
@ -481,26 +827,29 @@ export function useRecipeExecutions({
rows: normalizedRows,
settings: runSettings,
runName,
restorePrevious,
});
},
[
onExecutionStart,
prepareLocalModelForExecution,
readExecutablePayload,
runExecution,
runSettings,
setRunErrors,
validateExecutionPayload,
],
);
const runPreview = useCallback(async (): Promise<boolean> => {
const runPreview = useCallback((): Promise<boolean> => {
return runWithValidation("preview", previewRows, null);
}, [previewRows, runWithValidation]);
const runFull = useCallback(async (): Promise<boolean> => {
const runFull = useCallback((): Promise<boolean> => {
return runWithValidation("full", fullRows, fullRunName);
}, [fullRows, fullRunName, runWithValidation]);
const runFromDialog = useCallback(async (): Promise<boolean> => {
const runFromDialog = useCallback((): Promise<boolean> => {
setValidateResult(null);
if (runDialogKind === "preview") {
return runPreview();
@ -512,9 +861,10 @@ export function useRecipeExecutions({
setRunErrors([]);
const payload = readPayload();
if (!payload) {
const nextErrors = payloadResult.errors.length > 0
? payloadResult.errors
: [payloadErrorMessage];
const nextErrors =
payloadResult.errors.length > 0
? payloadResult.errors
: [payloadErrorMessage];
setValidateResult({
valid: false,
errors: nextErrors,
@ -525,24 +875,46 @@ export function useRecipeExecutions({
const rows = runDialogKind === "preview" ? previewRows : fullRows;
const normalizedRows = sanitizeExecutionRows(rows, runDialogKind);
const executionPayload = buildExecutionPayload({
payload,
kind: runDialogKind,
rows: normalizedRows,
settings: runSettings,
runName: runDialogKind === "full" ? normalizeRunName(fullRunName) : null,
});
setValidateLoading(true);
try {
const executionPayload = buildExecutionPayload({
payload,
kind: runDialogKind,
rows: normalizedRows,
settings: runSettings,
runName:
runDialogKind === "full" ? normalizeRunName(fullRunName) : null,
});
const validation = await validateRecipe(executionPayload);
const errors = formatValidationMessages({ errors: validation.errors });
if (!validation.valid) {
setValidateResult({
valid: false,
errors,
rawDetail: validation.raw_detail ?? null,
});
return false;
}
const localLoadError = await ensureLocalModelLoaded(payload);
if (localLoadError) {
setRunErrors([localLoadError]);
setValidateResult({
valid: false,
errors: [localLoadError],
rawDetail: null,
});
toastError("Local model failed to load", localLoadError);
return false;
}
setValidateResult({
valid: validation.valid,
valid: true,
errors,
rawDetail: validation.raw_detail ?? null,
});
return validation.valid;
return true;
} catch (error) {
const message = toErrorMessage(error, "Validation failed.");
setValidateResult({
@ -612,7 +984,12 @@ export function useRecipeExecutions({
const loadExecutionDatasetPage = useCallback(
async (id: string, page: number): Promise<void> => {
const execution = executions.find((entry) => entry.id === id);
if (!execution || execution.kind !== "full" || !execution.jobId || page < 1) {
if (
!execution ||
execution.kind !== "full" ||
!execution.jobId ||
page < 1
) {
return;
}
@ -625,7 +1002,9 @@ export function useRecipeExecutions({
});
const dataset = normalizeDatasetRows(response.dataset);
const total =
typeof response.total === "number" ? response.total : execution.datasetTotal;
typeof response.total === "number"
? response.total
: execution.datasetTotal;
upsertAndPersist({
...execution,
dataset,

View file

@ -97,7 +97,9 @@ export function applyRenameToConfig(
next = {
...base,
// biome-ignore lint/style/useNamingConvention: api schema
target_columns: targets.map((target) => (target === from ? to : target)),
target_columns: targets.map((target) =>
target === from ? to : target,
),
};
}
}
@ -137,14 +139,12 @@ export function applyRemovalToConfig(
}
if (config.kind === "model_config" && config.provider === ref) {
const base = next as ModelConfig;
// Clear the synthetic "local" placeholder when the provider that was
// a local provider is removed; otherwise the stale placeholder would
// pass validation against a future external provider and then fail
// at runtime against a real API ("model not found").
next = {
...base,
provider: "",
model: base.model === "local" ? "" : base.model,
model: "",
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: undefined,
};
}
if (config.kind === "llm" && config.model_alias === ref) {
@ -156,7 +156,9 @@ export function applyRemovalToConfig(
next = { ...base, tool_alias: "" };
}
if (config.kind === "validator") {
const targets = (config.target_columns ?? []).filter((target) => target !== ref);
const targets = (config.target_columns ?? []).filter(
(target) => target !== ref,
);
if (targets.length !== (config.target_columns ?? []).length) {
const base = next as typeof config;
next = {
@ -206,5 +208,7 @@ export function applyRemovalToConfigs(
if (!ref) {
return configs;
}
return applyConfigTransform(configs, (config) => applyRemovalToConfig(config, ref));
return applyConfigTransform(configs, (config) =>
applyRemovalToConfig(config, ref),
);
}

View file

@ -12,23 +12,23 @@ import {
applyNodeChanges,
} from "@xyflow/react";
import { create } from "zustand";
import type {
RecipeNode,
RecipeProcessorConfig,
LayoutDirection,
LlmType,
NodeConfig,
SeedSourceType,
SamplerType,
} from "../types";
import {
getBlockDefinition,
type BlockKind,
type BlockType,
type SeedBlockType,
getBlockDefinition,
} from "../blocks/registry";
import { deriveDisplayGraph } from "../utils/graph/derive-display-graph";
import type {
LayoutDirection,
LlmType,
NodeConfig,
RecipeNode,
RecipeProcessorConfig,
SamplerType,
SeedSourceType,
} from "../types";
import { applyRecipeConnection, isValidRecipeConnection } from "../utils/graph";
import { deriveDisplayGraph } from "../utils/graph/derive-display-graph";
import {
HANDLE_IDS,
normalizeRecipeHandleId,
@ -42,8 +42,8 @@ import {
} from "./helpers/model-infra-layout";
import { applyEdgeRemovals, applyNodeRemovals } from "./helpers/removals";
import {
applyRenameToConfigs,
applyLayoutDirectionToNodes,
applyRenameToConfigs,
buildNodeUpdate,
syncEdgesForConfigPatch,
syncSubcategoryConfigsForCategoryUpdate,
@ -97,7 +97,11 @@ type RecipeStudioState = {
position?: XYPosition,
openDialog?: boolean,
) => void;
addLlmNode: (type: LlmType, position?: XYPosition, openDialog?: boolean) => void;
addLlmNode: (
type: LlmType,
position?: XYPosition,
openDialog?: boolean,
) => void;
addModelProviderNode: (position?: XYPosition, openDialog?: boolean) => void;
addModelConfigNode: (position?: XYPosition, openDialog?: boolean) => void;
addToolProfileNode: (position?: XYPosition, openDialog?: boolean) => void;
@ -250,7 +254,10 @@ function connectSemantic(
};
}
function isModelSemanticEdge(edge: Edge, configs: Record<string, NodeConfig>): boolean {
function isModelSemanticEdge(
edge: Edge,
configs: Record<string, NodeConfig>,
): boolean {
const source = configs[edge.source];
const target = configs[edge.target];
return Boolean(
@ -315,12 +322,16 @@ export const useRecipeStudioStore = create<RecipeStudioState>((set, get) => ({
auxNodePositions: {},
llmAuxVisibility: state.llmAuxVisibility,
});
const { nodes } = getLayoutedElements(displayGraph.nodes, displayGraph.edges, {
direction: state.layoutDirection,
nodesep: isTopBottom ? 120 : 80,
ranksep: isTopBottom ? 140 : 80,
configs: state.configs,
});
const { nodes } = getLayoutedElements(
displayGraph.nodes,
displayGraph.edges,
{
direction: state.layoutDirection,
nodesep: isTopBottom ? 120 : 80,
ranksep: isTopBottom ? 140 : 80,
configs: state.configs,
},
);
const layoutedPositions = new Map(
nodes.map((node) => [node.id, node.position] as const),
);
@ -381,13 +392,7 @@ export const useRecipeStudioStore = create<RecipeStudioState>((set, get) => ({
(config) => config.kind === "seed",
);
if (!existing) {
return buildAddedNodeState(
state,
"seed",
type,
position,
openDialog,
);
return buildAddedNodeState(state, "seed", type, position, openDialog);
}
let nextSourceType: SeedSourceType = "hf";
if (type === "seed_local") {
@ -430,7 +435,10 @@ export const useRecipeStudioStore = create<RecipeStudioState>((set, get) => ({
[existing.id]: nextConfig,
},
nodes: updateNodeData(
state.nodes.map((node) => ({ ...node, selected: node.id === existing.id })),
state.nodes.map((node) => ({
...node,
selected: node.id === existing.id,
})),
existing.id,
nextConfig,
state.layoutDirection,
@ -444,7 +452,13 @@ export const useRecipeStudioStore = create<RecipeStudioState>((set, get) => ({
if (state.executionLocked) {
return state;
}
const added = buildAddedNodeState(state, "llm", type, position, openDialog);
const added = buildAddedNodeState(
state,
"llm",
type,
position,
openDialog,
);
const context = getAddedNodeContext(added);
if (!context) {
return added;
@ -495,9 +509,7 @@ export const useRecipeStudioStore = create<RecipeStudioState>((set, get) => ({
let { nodes, configs } = context;
let edges = state.edges;
const unboundModelConfigs = Object.values(configs).filter(
(config) =>
config.kind === "model_config" &&
!config.provider.trim(),
(config) => config.kind === "model_config" && !config.provider.trim(),
);
if (!position && unboundModelConfigs.length > 0) {
nodes = placeNodeNear(
@ -605,7 +617,7 @@ export const useRecipeStudioStore = create<RecipeStudioState>((set, get) => ({
let { nodes, configs } = context;
let edges = state.edges;
const unboundLlms = Object.values(configs).filter(
(config) => config.kind === "llm" && !(config.tool_alias?.trim()),
(config) => config.kind === "llm" && !config.tool_alias?.trim(),
);
if (!position && unboundLlms.length > 0) {
nodes = placeNodeNear(
@ -757,17 +769,15 @@ export const useRecipeStudioStore = create<RecipeStudioState>((set, get) => ({
if (cfg.kind !== "model_config" || cfg.provider !== providerName) {
continue;
}
if (nextIsLocal && !cfg.model.trim()) {
// external -> local: auto fill the placeholder model id so the
// config does not fail "model is required" validation.
configs = { ...configs, [cfgId]: { ...cfg, model: "local" } };
continue;
}
if (!nextIsLocal && cfg.model === "local") {
// local -> external: clear the placeholder so the user picks a
// real model id for the new external endpoint.
configs = { ...configs, [cfgId]: { ...cfg, model: "" } };
}
configs = {
...configs,
[cfgId]: {
...cfg,
model: "",
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: undefined,
},
};
}
}
}

View file

@ -264,6 +264,8 @@ export type ModelConfig = {
kind: "model_config";
name: string;
model: string;
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant?: string;
provider: string;
// biome-ignore lint/style/useNamingConvention: api schema
inference_temperature?: string;

View file

@ -11,7 +11,6 @@ import {
isSemanticTargetHandle,
normalizeRecipeHandleId,
} from "../handles";
import { isSemanticRelation } from "./relations";
import {
isCategoryConfig,
isExpressionConfig,
@ -21,6 +20,7 @@ import {
VALIDATOR_OXC_CODE_LANGS,
VALIDATOR_SQL_CODE_LANGS,
} from "../validators/code-lang";
import { isSemanticRelation } from "./relations";
function buildTemplateWithRef(template: string, ref: string): string {
if (template.includes(ref)) {
@ -157,7 +157,10 @@ function isCompetingIncomingEdge(
return source.kind === "sampler" && source.sampler_type === "datetime";
}
function isModelSemanticRelation(source: NodeConfig, target: NodeConfig): boolean {
function isModelSemanticRelation(
source: NodeConfig,
target: NodeConfig,
): boolean {
return (
(source.kind === "model_provider" && target.kind === "model_config") ||
(source.kind === "model_config" && target.kind === "llm") ||
@ -181,7 +184,9 @@ function canApplyCodeLangToValidator(
if (normalized === "python") {
return true;
}
return VALIDATOR_SQL_CODE_LANGS.includes(normalized as typeof validator.code_lang);
return VALIDATOR_SQL_CODE_LANGS.includes(
normalized as typeof validator.code_lang,
);
}
function countHandleUsage(
@ -333,12 +338,8 @@ export function applyRecipeConnection(
if (!isValidRecipeConnection(connection, configs)) {
return { edges };
}
const initialSource = connection.source
? configs[connection.source]
: null;
const initialTarget = connection.target
? configs[connection.target]
: null;
const initialSource = connection.source ? configs[connection.source] : null;
const initialTarget = connection.target ? configs[connection.target] : null;
if (!(initialSource && initialTarget)) {
return { edges };
}
@ -386,17 +387,36 @@ export function applyRecipeConnection(
nextBaseEdges,
);
if (source.kind === "model_provider" && target.kind === "model_config") {
// Keep the model_config.model field in sync with provider mode when the
// link is changed via graph drag (the model-config dialog path has its
// own applyProviderChange helper that does the same thing).
// Keep model_config.provider in sync when a graph drag changes the link.
// Local providers now require an explicit selected load id; do not synthesize
// the legacy "local" placeholder. External relinks clear local-only GGUF
// metadata, while legacy placeholders are normalized back to empty.
const isSourceLocal = source.is_local === true;
let nextModel = target.model;
if (isSourceLocal && !nextModel.trim()) {
nextModel = "local";
} else if (!isSourceLocal && nextModel === "local") {
nextModel = "";
}
const next = { ...target, provider: source.name, model: nextModel };
const isLegacyLocalPlaceholder =
target.model.trim().toLowerCase() === "local";
const previousProviderName = target.provider.trim();
const previousProvider = Object.values(configs).find(
(config) =>
config.kind === "model_provider" &&
config.name === previousProviderName,
);
const wasLinkedToLocal =
previousProvider?.kind === "model_provider" &&
previousProvider.is_local === true;
const shouldClearModel =
isLegacyLocalPlaceholder ||
(isSourceLocal ? !wasLinkedToLocal : wasLinkedToLocal);
const next = {
...target,
provider: source.name,
...(shouldClearModel ? { model: "" } : {}),
...(shouldClearModel || !isSourceLocal
? {
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: undefined,
}
: {}),
};
return { edges: nextEdges, configs: { ...configs, [target.id]: next } };
}
if (source.kind === "model_config" && target.kind === "llm") {
@ -435,10 +455,9 @@ export function applyRecipeConnection(
// biome-ignore lint/style/useNamingConvention: api schema
target_columns: [source.name],
// biome-ignore lint/style/useNamingConvention: api schema
code_lang:
(
canUseCodeLangForTarget ? nextCodeLang : target.code_lang
) as typeof target.code_lang,
code_lang: (canUseCodeLangForTarget
? nextCodeLang
: target.code_lang) as typeof target.code_lang,
};
return { edges: nextEdges, configs: { ...configs, [target.id]: next } };
}

View file

@ -1,15 +1,8 @@
// SPDX-License-Identifier: AGPL-3.0-only
// Copyright 2026-present the Unsloth AI Inc. team. All rights reserved. See /studio/LICENSE.AGPL-3.0
import type {
ModelConfig,
ModelProviderConfig,
} from "../../../types";
import {
isRecord,
readNumberString,
readString,
} from "../helpers";
import type { ModelConfig, ModelProviderConfig } from "../../../types";
import { isRecord, readNumberString, readString } from "../helpers";
export function parseModelProvider(
provider: Record<string, unknown>,
@ -53,6 +46,8 @@ export function parseModelConfig(
kind: "model_config",
name,
model: readString(model.model) ?? "",
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: readString(model.gguf_variant) ?? undefined,
provider: readString(model.provider) ?? "",
// biome-ignore lint/style/useNamingConvention: api schema
inference_temperature: readNumberString(inference.temperature),

View file

@ -54,55 +54,62 @@ export function buildModelProvider(
};
}
export function buildModelConfig(
function assignFiniteNumber(
target: Record<string, unknown>,
key: string,
rawValue: string | undefined,
transform: (value: number) => number = (value) => value,
): void {
const trimmed = rawValue?.trim();
if (!trimmed) {
return;
}
const parsed = Number(trimmed);
if (Number.isFinite(parsed)) {
target[key] = transform(parsed);
}
}
function buildInferenceParameters(
config: ModelConfig,
errors: string[],
): Record<string, unknown> {
const inference: Record<string, unknown> = {};
const temp = config.inference_temperature?.trim();
const topP = config.inference_top_p?.trim();
const maxTokens = config.inference_max_tokens?.trim();
const timeout = config.inference_timeout?.trim();
assignFiniteNumber(inference, "temperature", config.inference_temperature);
assignFiniteNumber(inference, "top_p", config.inference_top_p);
assignFiniteNumber(inference, "max_tokens", config.inference_max_tokens);
assignFiniteNumber(
inference,
"timeout",
config.inference_timeout,
Math.trunc,
);
const extraBody = parseJsonObject(
config.inference_extra_body,
`Model ${config.name} inference extra_body`,
errors,
);
if (temp) {
const parsed = Number(temp);
if (Number.isFinite(parsed)) {
inference.temperature = parsed;
}
}
if (topP) {
const parsed = Number(topP);
if (Number.isFinite(parsed)) {
// biome-ignore lint/style/useNamingConvention: api schema
inference.top_p = parsed;
}
}
if (maxTokens) {
const parsed = Number(maxTokens);
if (Number.isFinite(parsed)) {
// biome-ignore lint/style/useNamingConvention: api schema
inference.max_tokens = parsed;
}
}
if (timeout) {
const parsed = Number(timeout);
if (Number.isFinite(parsed)) {
inference.timeout = Math.trunc(parsed);
}
}
if (extraBody) {
// biome-ignore lint/style/useNamingConvention: api schema
inference.extra_body = extraBody;
}
return inference;
}
export function buildModelConfig(
config: ModelConfig,
errors: string[],
): Record<string, unknown> {
const inference = buildInferenceParameters(config, errors);
const ggufVariant = config.gguf_variant?.trim();
return {
alias: config.name,
model: config.model,
// biome-ignore lint/style/useNamingConvention: api schema
gguf_variant: ggufVariant || undefined,
provider: config.provider || undefined,
// biome-ignore lint/style/useNamingConvention: api schema
inference_parameters:

View file

@ -54,7 +54,9 @@ export function validateTimedeltaConfigs(
}
const reference = config.reference_column_name?.trim() ?? "";
if (!reference) {
errors.push(`Timedelta ${config.name}: reference datetime column required.`);
errors.push(
`Timedelta ${config.name}: reference datetime column required.`,
);
continue;
}
const parent = nameToConfig.get(reference);
@ -63,7 +65,9 @@ export function validateTimedeltaConfigs(
parent.kind !== "sampler" ||
parent.sampler_type !== "datetime"
) {
errors.push(`Timedelta ${config.name}: reference '${reference}' must be datetime.`);
errors.push(
`Timedelta ${config.name}: reference '${reference}' must be datetime.`,
);
}
}
}
@ -91,9 +95,18 @@ export function validateModelConfigProviders(
const provider = config.provider.trim();
const alias = config.name;
const isLocal = localProviderNames.has(provider);
// Local providers do not require a real model id - the loaded Chat
// model is used regardless of what gets sent in the payload.
if (!isLocal && modelAliases.has(alias) && !config.model.trim()) {
const isUsed = modelAliases.has(alias);
const model = config.model.trim();
const isLegacyLocalPlaceholder = model.toLowerCase() === "local";
if (!isLocal && isUsed && isLegacyLocalPlaceholder) {
errors.push(`Model config ${alias}: model is required.`);
continue;
}
if (isLocal && isUsed && !model) {
errors.push(`Model config ${alias}: choose a local model.`);
}
if (!isLocal && isUsed && !model) {
errors.push(`Model config ${alias}: model is required.`);
}
if (provider && !modelProviderNames.has(provider)) {
@ -121,7 +134,9 @@ export function validateUsedProviders(
errors.push(`Model provider ${provider.name}: endpoint is required.`);
}
if (!provider.provider_type.trim()) {
errors.push(`Model provider ${provider.name}: provider_type is required.`);
errors.push(
`Model provider ${provider.name}: provider_type is required.`,
);
}
}
}
@ -145,7 +160,9 @@ export function validateValidatorConfigs(
continue;
}
if (targetConfig.kind !== "llm" || targetConfig.llm_type !== "code") {
errors.push(`Validator ${config.name}: target '${target}' must be LLM Code.`);
errors.push(
`Validator ${config.name}: target '${target}' must be LLM Code.`,
);
continue;
}
if (

View file

@ -1188,6 +1188,53 @@
border-color: var(--border) !important;
}
.generated-image-loading-card {
position: relative;
overflow: hidden;
contain: paint;
}
.generated-image-loading-wave {
position: relative;
display: grid;
grid-template-columns: repeat(8, minmax(0, 1fr));
gap: 14px;
width: min(66%, 18rem);
padding: 1.5rem;
border-radius: 1.5rem;
}
.generated-image-loading-dot {
width: 7px;
height: 7px;
border-radius: 9999px;
background: color-mix(in oklch, var(--muted-foreground) 82%, var(--primary));
opacity: 0.12;
transform: translate3d(0, 4px, 0) scale(0.72);
animation: generated-image-dot-wave 1850ms var(--ease-out-quart) infinite;
animation-delay: calc((var(--dot-row) * 72ms) + (var(--dot-col) * 72ms));
will-change: transform, opacity;
}
@keyframes generated-image-dot-wave {
0%,
22%,
100% {
opacity: 0.1;
transform: translate3d(0, 4px, 0) scale(0.72);
}
46% {
opacity: 0.46;
transform: translate3d(0, -3px, 0) scale(0.96);
}
66% {
opacity: 0.2;
transform: translate3d(0, 0, 0) scale(0.82);
}
}
/*
* prefers-reduced-motion: honour the OS-level "reduce motion" preference.
* Tailwind animate-in/out, Radix open/close transforms, infinite shine/pulse
@ -1197,11 +1244,11 @@
* end state. Hover colour changes become instant rather than fading, which is
* the documented WCAG outcome (motion is "minimised, not removed").
*
* .animate-spin is the exception: loading spinners are essential progress
* .animate-spin and generated image loading dots are the exceptions: loading
* indicators across Studio (tool execution loaders, sonner toasts, Tauri
* startup / update screens, the <Spinner /> primitive). Freezing them
* removes the only visual signal that work is in flight, so they keep
* animating but at a slower, less aggressive 1.5s cadence.
* startup / update screens, the <Spinner /> primitive, and image generation
* cards). Freezing them removes the only visual signal that work is in flight,
* so they keep animating.
*/
@media (prefers-reduced-motion: reduce) {
*,
@ -1217,4 +1264,9 @@
animation-duration: 1.5s !important;
animation-iteration-count: infinite !important;
}
.generated-image-loading-dot {
animation-duration: 1850ms !important;
animation-iteration-count: infinite !important;
}
}