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:
commit
b62e4d18cd
74 changed files with 10967 additions and 1234 deletions
|
|
@ -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;
|
||||
}
|
||||
510
studio/frontend/src/components/assistant-ui/image.tsx
Normal file
510
studio/frontend/src/components/assistant-ui/image.tsx
Normal 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,
|
||||
};
|
||||
|
|
@ -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">
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 />
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -248,7 +248,7 @@ interface BackendInferenceDefaults {
|
|||
export interface BackendInferenceEnvelope {
|
||||
is_gguf?: boolean;
|
||||
context_length?: number | null;
|
||||
inference?: BackendInferenceDefaults;
|
||||
inference?: BackendInferenceDefaults | null;
|
||||
}
|
||||
|
||||
export function mergeBackendRecommendedInference({
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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 && (
|
||||
|
|
|
|||
|
|
@ -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 }),
|
||||
}));
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
);
|
||||
}
|
||||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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>
|
||||
|
|
|
|||
|
|
@ -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 };
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
},
|
||||
};
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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 } };
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue