Merge pull request #100 from unslothai/feature/chat-compare
feat: chat compare + inference stream cancel fix
This commit is contained in:
commit
28364f3314
16 changed files with 534 additions and 1466 deletions
|
|
@ -489,6 +489,7 @@ class InferenceBackend:
|
|||
def generate_with_adapter_control(
|
||||
self,
|
||||
use_adapter: Optional[Union[bool, str]] = None,
|
||||
cancel_event=None,
|
||||
**gen_kwargs,
|
||||
) -> Generator[str, None, None]:
|
||||
"""
|
||||
|
|
@ -505,7 +506,7 @@ class InferenceBackend:
|
|||
with self._generation_lock:
|
||||
self._apply_adapter_state(use_adapter)
|
||||
# Delegate to the lock-free generation path
|
||||
yield from self._generate_chat_response_inner(**gen_kwargs)
|
||||
yield from self._generate_chat_response_inner(cancel_event=cancel_event, **gen_kwargs)
|
||||
|
||||
def generate_chat_response(self,
|
||||
messages: list,
|
||||
|
|
@ -515,7 +516,8 @@ class InferenceBackend:
|
|||
top_p: float = 0.9,
|
||||
top_k: int = 40,
|
||||
max_new_tokens: int = 256,
|
||||
repetition_penalty: float = 1.1) -> Generator[str, None, None]:
|
||||
repetition_penalty: float = 1.1,
|
||||
cancel_event=None) -> Generator[str, None, None]:
|
||||
"""
|
||||
Generate response for text or vision models.
|
||||
Acquires the generation lock. For adapter-controlled generation,
|
||||
|
|
@ -531,6 +533,7 @@ class InferenceBackend:
|
|||
top_k=top_k,
|
||||
max_new_tokens=max_new_tokens,
|
||||
repetition_penalty=repetition_penalty,
|
||||
cancel_event=cancel_event,
|
||||
)
|
||||
|
||||
def _generate_chat_response_inner(self,
|
||||
|
|
@ -541,7 +544,8 @@ class InferenceBackend:
|
|||
top_p: float = 0.9,
|
||||
top_k: int = 40,
|
||||
max_new_tokens: int = 256,
|
||||
repetition_penalty: float = 1.1) -> Generator[str, None, None]:
|
||||
repetition_penalty: float = 1.1,
|
||||
cancel_event=None) -> Generator[str, None, None]:
|
||||
"""
|
||||
Inner generation logic (no lock). Called by both generate_chat_response
|
||||
and generate_with_adapter_control.
|
||||
|
|
@ -558,7 +562,8 @@ class InferenceBackend:
|
|||
# Vision model generation
|
||||
yield from self._generate_vision_response(
|
||||
messages, system_prompt, image,
|
||||
temperature, top_p, top_k, max_new_tokens, repetition_penalty
|
||||
temperature, top_p, top_k, max_new_tokens, repetition_penalty,
|
||||
cancel_event=cancel_event,
|
||||
)
|
||||
else:
|
||||
# Text model: Use training pipeline approach
|
||||
|
|
@ -600,12 +605,13 @@ class InferenceBackend:
|
|||
|
||||
# Step 3: Generate
|
||||
yield from self.generate_stream(
|
||||
formatted_prompt, temperature, top_p, top_k, max_new_tokens, repetition_penalty
|
||||
formatted_prompt, temperature, top_p, top_k, max_new_tokens, repetition_penalty,
|
||||
cancel_event=cancel_event,
|
||||
)
|
||||
|
||||
def _generate_vision_response(self, messages, system_prompt, image,
|
||||
temperature, top_p, top_k, max_new_tokens,
|
||||
repetition_penalty) -> Generator[str, None, None]:
|
||||
repetition_penalty, cancel_event=None) -> Generator[str, None, None]:
|
||||
"""Handle vision model generation with true token-by-token streaming."""
|
||||
model_info = self.models[self.active_model_name]
|
||||
model = model_info["model"]
|
||||
|
|
@ -651,7 +657,10 @@ class InferenceBackend:
|
|||
import threading
|
||||
|
||||
streamer = TextIteratorStreamer(
|
||||
processor.tokenizer, skip_prompt=True, skip_special_tokens=True
|
||||
processor.tokenizer,
|
||||
skip_prompt=True,
|
||||
skip_special_tokens=True,
|
||||
timeout=0.2,
|
||||
)
|
||||
|
||||
generation_kwargs = dict(
|
||||
|
|
@ -664,23 +673,50 @@ class InferenceBackend:
|
|||
top_k=top_k,
|
||||
)
|
||||
|
||||
err: dict[str, str] = {}
|
||||
|
||||
def generate_fn():
|
||||
try:
|
||||
model.generate(**generation_kwargs)
|
||||
except Exception as e:
|
||||
err["msg"] = str(e)
|
||||
logger.error(f"Vision generation error in thread: {e}")
|
||||
finally:
|
||||
try:
|
||||
streamer.end()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
thread = threading.Thread(target=generate_fn)
|
||||
thread.start()
|
||||
|
||||
output = ""
|
||||
for new_token in streamer:
|
||||
if new_token:
|
||||
output += new_token
|
||||
cleaned = self._clean_generated_text(output)
|
||||
yield cleaned
|
||||
from queue import Empty
|
||||
try:
|
||||
while True:
|
||||
if cancel_event is not None and cancel_event.is_set():
|
||||
break
|
||||
try:
|
||||
new_token = next(streamer)
|
||||
except StopIteration:
|
||||
break
|
||||
except Empty:
|
||||
if not thread.is_alive():
|
||||
break
|
||||
continue
|
||||
if new_token:
|
||||
output += new_token
|
||||
cleaned = self._clean_generated_text(output)
|
||||
yield cleaned
|
||||
finally:
|
||||
if cancel_event is not None:
|
||||
cancel_event.set()
|
||||
thread.join(timeout=10)
|
||||
if thread.is_alive():
|
||||
logger.warning("Vision generation thread did not exit after cancel/join timeout")
|
||||
|
||||
thread.join()
|
||||
if err.get("msg"):
|
||||
yield f"Error: {err['msg']}"
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Vision generation error: {e}")
|
||||
|
|
@ -693,7 +729,8 @@ class InferenceBackend:
|
|||
top_p: float = 0.9,
|
||||
top_k: int = 40,
|
||||
max_new_tokens: int = 256,
|
||||
repetition_penalty: float = 1.1) -> Generator[str, None, None]:
|
||||
repetition_penalty: float = 1.1,
|
||||
cancel_event=None) -> Generator[str, None, None]:
|
||||
"""Generate streaming text response (text models only)."""
|
||||
if not self.active_model_name:
|
||||
yield "Error: No active model"
|
||||
|
|
@ -709,7 +746,12 @@ class InferenceBackend:
|
|||
from transformers import TextIteratorStreamer
|
||||
import threading
|
||||
|
||||
streamer = TextIteratorStreamer(tokenizer, skip_prompt=True, skip_special_tokens=True)
|
||||
streamer = TextIteratorStreamer(
|
||||
tokenizer,
|
||||
skip_prompt=True,
|
||||
skip_special_tokens=True,
|
||||
timeout=0.2,
|
||||
)
|
||||
|
||||
generation_kwargs = dict(
|
||||
**inputs,
|
||||
|
|
@ -723,24 +765,66 @@ class InferenceBackend:
|
|||
eos_token_id=tokenizer.eos_token_id,
|
||||
pad_token_id=tokenizer.eos_token_id if tokenizer.pad_token_id is None else tokenizer.pad_token_id,
|
||||
)
|
||||
if cancel_event is not None:
|
||||
from transformers.generation.stopping_criteria import (
|
||||
StoppingCriteria,
|
||||
StoppingCriteriaList,
|
||||
)
|
||||
|
||||
class _CancelCriteria(StoppingCriteria):
|
||||
def __init__(self, ev):
|
||||
self.ev = ev
|
||||
|
||||
def __call__(self, input_ids, scores, **kwargs):
|
||||
return self.ev.is_set()
|
||||
|
||||
generation_kwargs["stopping_criteria"] = StoppingCriteriaList(
|
||||
[_CancelCriteria(cancel_event)]
|
||||
)
|
||||
|
||||
def generate_fn():
|
||||
try:
|
||||
model.generate(**generation_kwargs)
|
||||
except Exception as e:
|
||||
err["msg"] = str(e)
|
||||
logger.error(f"Generation error: {e}")
|
||||
finally:
|
||||
try:
|
||||
streamer.end()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
err: dict[str, str] = {}
|
||||
thread = threading.Thread(target=generate_fn)
|
||||
thread.start()
|
||||
|
||||
output = ""
|
||||
for new_token in streamer:
|
||||
if new_token:
|
||||
output += new_token
|
||||
cleaned = self._clean_generated_text(output)
|
||||
yield cleaned
|
||||
from queue import Empty
|
||||
try:
|
||||
while True:
|
||||
if cancel_event is not None and cancel_event.is_set():
|
||||
break
|
||||
try:
|
||||
new_token = next(streamer)
|
||||
except StopIteration:
|
||||
break
|
||||
except Empty:
|
||||
if not thread.is_alive():
|
||||
break
|
||||
continue
|
||||
if new_token:
|
||||
output += new_token
|
||||
cleaned = self._clean_generated_text(output)
|
||||
yield cleaned
|
||||
finally:
|
||||
if cancel_event is not None:
|
||||
cancel_event.set()
|
||||
thread.join(timeout=10)
|
||||
if thread.is_alive():
|
||||
logger.warning("Generation thread did not exit after cancel/join timeout")
|
||||
|
||||
thread.join()
|
||||
if err.get("msg"):
|
||||
yield f"Error: {err['msg']}"
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Error during generation: {e}")
|
||||
|
|
|
|||
|
|
@ -5,11 +5,13 @@ import sys
|
|||
import time
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
from fastapi import APIRouter, HTTPException
|
||||
from fastapi import APIRouter, HTTPException, Request
|
||||
from fastapi.responses import StreamingResponse, JSONResponse
|
||||
from typing import Optional
|
||||
import json
|
||||
import logging
|
||||
import asyncio
|
||||
import threading
|
||||
|
||||
|
||||
|
||||
|
|
@ -304,7 +306,7 @@ def _extract_content_parts(
|
|||
|
||||
|
||||
@router.post("/chat/completions")
|
||||
async def openai_chat_completions(request: ChatCompletionRequest):
|
||||
async def openai_chat_completions(payload: ChatCompletionRequest, request: Request):
|
||||
"""
|
||||
OpenAI-compatible chat completions endpoint.
|
||||
|
||||
|
|
@ -324,7 +326,7 @@ async def openai_chat_completions(request: ChatCompletionRequest):
|
|||
|
||||
# ── Parse messages (handles multimodal content parts) ─────
|
||||
system_prompt, chat_messages, extracted_image_b64 = _extract_content_parts(
|
||||
request.messages
|
||||
payload.messages
|
||||
)
|
||||
|
||||
# If no non-system messages were provided, error out
|
||||
|
|
@ -336,7 +338,7 @@ async def openai_chat_completions(request: ChatCompletionRequest):
|
|||
|
||||
# ── Decode image (from content parts OR legacy field) ─────
|
||||
# Content-part images take priority; fall back to legacy field
|
||||
image_b64 = extracted_image_b64 or request.image_base64
|
||||
image_b64 = extracted_image_b64 or payload.image_base64
|
||||
image = None
|
||||
|
||||
if image_b64:
|
||||
|
|
@ -366,31 +368,35 @@ async def openai_chat_completions(request: ChatCompletionRequest):
|
|||
messages=chat_messages,
|
||||
system_prompt=system_prompt,
|
||||
image=image,
|
||||
temperature=request.temperature,
|
||||
top_p=request.top_p,
|
||||
top_k=request.top_k,
|
||||
max_new_tokens=request.max_tokens or 512,
|
||||
repetition_penalty=request.repetition_penalty,
|
||||
temperature=payload.temperature,
|
||||
top_p=payload.top_p,
|
||||
top_k=payload.top_k,
|
||||
max_new_tokens=payload.max_tokens or 512,
|
||||
repetition_penalty=payload.repetition_penalty,
|
||||
)
|
||||
|
||||
# ── Choose generation path (adapter-controlled or standard) ──
|
||||
if request.use_adapter is not None:
|
||||
cancel_event = threading.Event()
|
||||
|
||||
if payload.use_adapter is not None:
|
||||
# Compare mode: toggle adapter state atomically with generation
|
||||
def generate():
|
||||
return backend.generate_with_adapter_control(
|
||||
use_adapter=request.use_adapter, **gen_kwargs
|
||||
use_adapter=payload.use_adapter,
|
||||
cancel_event=cancel_event,
|
||||
**gen_kwargs,
|
||||
)
|
||||
else:
|
||||
# Standard path: no adapter toggling
|
||||
def generate():
|
||||
return backend.generate_chat_response(**gen_kwargs)
|
||||
return backend.generate_chat_response(cancel_event=cancel_event, **gen_kwargs)
|
||||
|
||||
model_name = backend.active_model_name or request.model
|
||||
model_name = backend.active_model_name or payload.model
|
||||
completion_id = f"chatcmpl-{uuid.uuid4().hex[:12]}"
|
||||
created = int(time.time())
|
||||
|
||||
# ── Streaming response ────────────────────────────────────────
|
||||
if request.stream:
|
||||
if payload.stream:
|
||||
async def stream_chunks():
|
||||
try:
|
||||
# First chunk: send the role
|
||||
|
|
@ -409,6 +415,10 @@ async def openai_chat_completions(request: ChatCompletionRequest):
|
|||
# text, so we diff to get incremental deltas.
|
||||
prev_text = ""
|
||||
for cumulative in generate():
|
||||
if await request.is_disconnected():
|
||||
cancel_event.set()
|
||||
backend.reset_generation_state()
|
||||
return
|
||||
new_text = cumulative[len(prev_text):]
|
||||
prev_text = cumulative
|
||||
if not new_text:
|
||||
|
|
@ -437,6 +447,10 @@ async def openai_chat_completions(request: ChatCompletionRequest):
|
|||
yield f"data: {final_chunk.model_dump_json(exclude_none=True)}\n\n"
|
||||
yield "data: [DONE]\n\n"
|
||||
|
||||
except asyncio.CancelledError:
|
||||
cancel_event.set()
|
||||
backend.reset_generation_state()
|
||||
raise
|
||||
except Exception as e:
|
||||
backend.reset_generation_state()
|
||||
logger.error(f"Error during OpenAI streaming: {e}", exc_info=True)
|
||||
|
|
@ -477,4 +491,3 @@ async def openai_chat_completions(request: ChatCompletionRequest):
|
|||
backend.reset_generation_state()
|
||||
logger.error(f"Error during OpenAI completion: {e}", exc_info=True)
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
|
|
|||
|
|
@ -41,6 +41,13 @@ def is_local_path(path: str) -> bool:
|
|||
if not path:
|
||||
return False
|
||||
|
||||
# If it exists on disk, treat as local (covers relative paths like "outputs/foo").
|
||||
try:
|
||||
if Path(normalize_path(path)).expanduser().exists():
|
||||
return True
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Obvious HF patterns
|
||||
if path.count('/') == 1 and not path.startswith(('/', '.', '~')):
|
||||
return False # Looks like org/model format
|
||||
|
|
|
|||
|
|
@ -2,7 +2,6 @@ import { createRouter } from "@tanstack/react-router";
|
|||
import { Route as rootRoute } from "./routes/__root";
|
||||
import { Route as chatRoute } from "./routes/chat";
|
||||
import { Route as gridTestRoute } from "./routes/grid-test";
|
||||
import { Route as homeRoute } from "./routes/home";
|
||||
import { Route as loginRoute } from "./routes/login";
|
||||
import { Route as onboardingRoute } from "./routes/onboarding";
|
||||
import { Route as exportRoute } from "./routes/export";
|
||||
|
|
@ -10,7 +9,6 @@ import { Route as signupRoute } from "./routes/signup";
|
|||
import { Route as studioRoute } from "./routes/studio";
|
||||
|
||||
const routeTree = rootRoute.addChildren([
|
||||
homeRoute,
|
||||
onboardingRoute,
|
||||
loginRoute,
|
||||
signupRoute,
|
||||
|
|
|
|||
|
|
@ -1,15 +0,0 @@
|
|||
import { ComponentExample } from "@/components/component-example";
|
||||
import { createRoute } from "@tanstack/react-router";
|
||||
import { requireAuth } from "../auth-guards";
|
||||
import { Route as rootRoute } from "./__root";
|
||||
|
||||
export const Route = createRoute({
|
||||
getParentRoute: () => rootRoute,
|
||||
path: "/",
|
||||
beforeLoad: () => requireAuth(),
|
||||
component: HomePage,
|
||||
});
|
||||
|
||||
function HomePage() {
|
||||
return <ComponentExample />;
|
||||
}
|
||||
|
|
@ -7,9 +7,7 @@ import { MarkdownText } from "@/components/assistant-ui/markdown-text";
|
|||
import { Reasoning, ReasoningGroup } from "@/components/assistant-ui/reasoning";
|
||||
import { ToolFallback } from "@/components/assistant-ui/tool-fallback";
|
||||
import { TooltipIconButton } from "@/components/assistant-ui/tooltip-icon-button";
|
||||
import { AnimatedShinyText } from "@/components/ui/animated-shiny-text";
|
||||
import { Button } from "@/components/ui/button";
|
||||
import { useChatRuntimeStore } from "@/features/chat/stores/chat-runtime-store";
|
||||
import { cn } from "@/lib/utils";
|
||||
import {
|
||||
ActionBarMorePrimitive,
|
||||
|
|
@ -73,7 +71,6 @@ export const Thread: FC<{ hideComposer?: boolean; hideWelcome?: boolean }> = ({
|
|||
|
||||
<ThreadPrimitive.ViewportFooter className="aui-thread-viewport-footer sticky bottom-0 mt-auto flex w-full flex-col gap-4 overflow-visible bg-background pb-4 md:pb-4 before:pointer-events-none before:absolute before:inset-x-0 before:bottom-full before:h-20 before:bg-gradient-to-t before:from-background before:to-transparent">
|
||||
<ThreadScrollToBottom />
|
||||
<WarmupIndicator />
|
||||
<AuiIf condition={({ thread }) => !thread.isEmpty}>
|
||||
{!hideComposer && <ComposerAnimated />}
|
||||
</AuiIf>
|
||||
|
|
@ -83,28 +80,6 @@ export const Thread: FC<{ hideComposer?: boolean; hideWelcome?: boolean }> = ({
|
|||
);
|
||||
};
|
||||
|
||||
const WarmupIndicator: FC = () => {
|
||||
const threadId = useAuiState(({ threads }) => threads.mainThreadId);
|
||||
const isRunning = useAuiState(({ thread }) => thread.isRunning);
|
||||
const isWarmingUp = useChatRuntimeStore((state) =>
|
||||
Boolean(state.warmingByThreadId[threadId ?? "__default"]),
|
||||
);
|
||||
|
||||
if (!isRunning || !isWarmingUp) {
|
||||
return null;
|
||||
}
|
||||
|
||||
return (
|
||||
<div className="mx-auto -mb-2 w-full max-w-(--thread-max-width) px-2">
|
||||
<div className="inline-flex items-center rounded-full border border-border/60 bg-background/90 px-3 py-1.5 text-xs text-muted-foreground shadow-sm">
|
||||
<AnimatedShinyText className="text-xs">
|
||||
Warming up model...
|
||||
</AnimatedShinyText>
|
||||
</div>
|
||||
</div>
|
||||
);
|
||||
};
|
||||
|
||||
const ThreadScrollToBottom: FC = () => {
|
||||
return (
|
||||
<ThreadPrimitive.ScrollToBottom asChild={true}>
|
||||
|
|
|
|||
File diff suppressed because it is too large
Load diff
|
|
@ -1,5 +1,7 @@
|
|||
import type { ChatModelAdapter } from "@assistant-ui/react";
|
||||
import { toast } from "sonner";
|
||||
import { streamChatCompletions } from "./chat-api";
|
||||
import { db } from "../db";
|
||||
import { useChatRuntimeStore } from "../stores/chat-runtime-store";
|
||||
import {
|
||||
hasClosedThinkTag,
|
||||
|
|
@ -81,13 +83,33 @@ function findLatestUserImageBase64(messages: RunMessages): string | undefined {
|
|||
return undefined;
|
||||
}
|
||||
|
||||
async function resolveUseAdapter(
|
||||
threadId: string | undefined,
|
||||
): Promise<boolean | undefined> {
|
||||
if (!threadId) {
|
||||
return undefined;
|
||||
}
|
||||
try {
|
||||
const thread = await db.threads.get(threadId);
|
||||
if (!thread?.pairId) {
|
||||
return undefined;
|
||||
}
|
||||
return thread.modelType === "lora";
|
||||
} catch {
|
||||
return undefined;
|
||||
}
|
||||
}
|
||||
|
||||
export function createOpenAIStreamAdapter(): ChatModelAdapter {
|
||||
return {
|
||||
async *run({ messages, abortSignal, unstable_threadId }) {
|
||||
const state = useChatRuntimeStore.getState();
|
||||
const { params } = state;
|
||||
const runtime = useChatRuntimeStore.getState();
|
||||
const { params } = runtime;
|
||||
|
||||
if (!params.checkpoint) {
|
||||
toast.error("No model loaded", {
|
||||
description: "Pick model in top bar, then retry.",
|
||||
});
|
||||
throw new Error("Load a model first.");
|
||||
}
|
||||
|
||||
|
|
@ -104,10 +126,48 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
|
|||
});
|
||||
}
|
||||
const imageBase64 = findLatestUserImageBase64(messages);
|
||||
const useAdapter = await resolveUseAdapter(unstable_threadId);
|
||||
|
||||
const threadKey = unstable_threadId || "__default";
|
||||
let waitingFirstChunk = true;
|
||||
useChatRuntimeStore.getState().setThreadWarming(threadKey, true);
|
||||
let firstTokenSettled = false;
|
||||
let resolveFirstToken: (() => void) | null = null;
|
||||
let rejectFirstToken: ((err: unknown) => void) | null = null;
|
||||
const firstTokenPromise = new Promise<void>((resolve, reject) => {
|
||||
resolveFirstToken = resolve;
|
||||
rejectFirstToken = reject;
|
||||
});
|
||||
// Avoid unhandled rejections if toast.promise never attached.
|
||||
void firstTokenPromise.catch(() => {});
|
||||
|
||||
function settleFirstTokenOk(): void {
|
||||
if (firstTokenSettled) return;
|
||||
firstTokenSettled = true;
|
||||
resolveFirstToken?.();
|
||||
}
|
||||
|
||||
function settleFirstTokenErr(err: unknown): void {
|
||||
if (firstTokenSettled) return;
|
||||
firstTokenSettled = true;
|
||||
rejectFirstToken?.(err);
|
||||
}
|
||||
|
||||
let warmupToastShown = false;
|
||||
const warmupDelayMs = 450;
|
||||
const warmupTimer = setTimeout(() => {
|
||||
if (!waitingFirstChunk) return;
|
||||
if (abortSignal.aborted) return;
|
||||
warmupToastShown = true;
|
||||
toast.promise(firstTokenPromise, {
|
||||
loading: "Warming up model",
|
||||
success: "Generating",
|
||||
error: (err) =>
|
||||
err instanceof Error && err.message ? err.message : "Generation failed",
|
||||
description: "Waiting for first token.",
|
||||
duration: 900,
|
||||
});
|
||||
}, warmupDelayMs);
|
||||
runtime.setThreadRunning(threadKey, true);
|
||||
let cumulativeText = "";
|
||||
let reasoningStartAt: number | null = null;
|
||||
let reasoningDuration = 0;
|
||||
|
|
@ -124,6 +184,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
|
|||
top_k: params.topK,
|
||||
repetition_penalty: params.repetitionPenalty,
|
||||
image_base64: imageBase64,
|
||||
...(useAdapter === undefined ? {} : { use_adapter: useAdapter }),
|
||||
},
|
||||
abortSignal,
|
||||
);
|
||||
|
|
@ -135,7 +196,7 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
|
|||
}
|
||||
if (waitingFirstChunk) {
|
||||
waitingFirstChunk = false;
|
||||
useChatRuntimeStore.getState().setThreadWarming(threadKey, false);
|
||||
settleFirstTokenOk();
|
||||
}
|
||||
|
||||
cumulativeText += delta;
|
||||
|
|
@ -155,10 +216,30 @@ export function createOpenAIStreamAdapter(): ChatModelAdapter {
|
|||
};
|
||||
}
|
||||
}
|
||||
} finally {
|
||||
if (waitingFirstChunk) {
|
||||
useChatRuntimeStore.getState().setThreadWarming(threadKey, false);
|
||||
settleFirstTokenOk();
|
||||
} catch (err) {
|
||||
settleFirstTokenErr(err instanceof Error ? err : new Error("Generation failed"));
|
||||
const isEarly = waitingFirstChunk;
|
||||
if (!abortSignal.aborted && !(warmupToastShown && isEarly)) {
|
||||
toast.error("Generation failed", {
|
||||
description: err instanceof Error ? err.message : "Unknown error",
|
||||
});
|
||||
}
|
||||
throw err;
|
||||
} finally {
|
||||
clearTimeout(warmupTimer);
|
||||
if (waitingFirstChunk) {
|
||||
if (warmupToastShown && !firstTokenSettled) {
|
||||
if (abortSignal.aborted) {
|
||||
settleFirstTokenErr(new Error("Cancelled"));
|
||||
} else {
|
||||
settleFirstTokenErr(new Error("No tokens received"));
|
||||
}
|
||||
} else {
|
||||
settleFirstTokenOk();
|
||||
}
|
||||
}
|
||||
runtime.setThreadRunning(threadKey, false);
|
||||
}
|
||||
},
|
||||
};
|
||||
|
|
|
|||
|
|
@ -174,7 +174,8 @@ function InlineSidebar({
|
|||
function TopBarActions({
|
||||
onNewThread,
|
||||
onNewCompare,
|
||||
}: { onNewThread: () => void; onNewCompare: () => void }) {
|
||||
showCompare,
|
||||
}: { onNewThread: () => void; onNewCompare: () => void; showCompare: boolean }) {
|
||||
const { state } = useSidebar();
|
||||
if (state !== "collapsed") {
|
||||
return null;
|
||||
|
|
@ -189,14 +190,16 @@ function TopBarActions({
|
|||
</TooltipTrigger>
|
||||
<TooltipContent side="bottom">New Chat</TooltipContent>
|
||||
</Tooltip>
|
||||
<Tooltip>
|
||||
<TooltipTrigger asChild={true}>
|
||||
<Button variant="ghost" size="icon-sm" onClick={onNewCompare}>
|
||||
<HugeiconsIcon icon={ColumnInsertIcon} strokeWidth={2} />
|
||||
</Button>
|
||||
</TooltipTrigger>
|
||||
<TooltipContent side="bottom">Compare</TooltipContent>
|
||||
</Tooltip>
|
||||
{showCompare ? (
|
||||
<Tooltip>
|
||||
<TooltipTrigger asChild={true}>
|
||||
<Button variant="ghost" size="icon-sm" onClick={onNewCompare}>
|
||||
<HugeiconsIcon icon={ColumnInsertIcon} strokeWidth={2} />
|
||||
</Button>
|
||||
</TooltipTrigger>
|
||||
<TooltipContent side="bottom">Compare</TooltipContent>
|
||||
</Tooltip>
|
||||
) : null}
|
||||
</>
|
||||
);
|
||||
}
|
||||
|
|
@ -209,10 +212,17 @@ export function ChatPage(): ReactElement {
|
|||
const [settingsOpen, setSettingsOpen] = useState(false);
|
||||
const inferenceParams = useChatRuntimeStore((state) => state.params);
|
||||
const setInferenceParams = useChatRuntimeStore((state) => state.setParams);
|
||||
const autoTitle = useChatRuntimeStore((state) => state.autoTitle);
|
||||
const setAutoTitle = useChatRuntimeStore((state) => state.setAutoTitle);
|
||||
const modelsFromStore = useChatRuntimeStore((state) => state.models);
|
||||
const lorasFromStore = useChatRuntimeStore((state) => state.loras);
|
||||
const modelsError = useChatRuntimeStore((state) => state.modelsError);
|
||||
const { refresh, selectModel, ejectModel } = useChatModelRuntime();
|
||||
const canCompare = useMemo(() => {
|
||||
const selected = inferenceParams.checkpoint;
|
||||
if (!selected) return false;
|
||||
return lorasFromStore.some((lora) => lora.id === selected);
|
||||
}, [inferenceParams.checkpoint, lorasFromStore]);
|
||||
|
||||
const handleCheckpointChange = useCallback(
|
||||
(value: string, meta?: { isLora: boolean }) => {
|
||||
|
|
@ -275,6 +285,7 @@ export function ChatPage(): ReactElement {
|
|||
onSelect={setView}
|
||||
onNewThread={handleNewThread}
|
||||
onNewCompare={handleNewCompare}
|
||||
showCompare={canCompare}
|
||||
/>
|
||||
</InlineSidebar>
|
||||
|
||||
|
|
@ -285,6 +296,7 @@ export function ChatPage(): ReactElement {
|
|||
<TopBarActions
|
||||
onNewThread={handleNewThread}
|
||||
onNewCompare={handleNewCompare}
|
||||
showCompare={canCompare}
|
||||
/>
|
||||
<ModelSelector
|
||||
models={models}
|
||||
|
|
@ -326,6 +338,8 @@ export function ChatPage(): ReactElement {
|
|||
open={settingsOpen}
|
||||
params={inferenceParams}
|
||||
onParamsChange={setInferenceParams}
|
||||
autoTitle={autoTitle}
|
||||
onAutoTitleChange={setAutoTitle}
|
||||
/>
|
||||
</SidebarProvider>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -23,6 +23,7 @@ import {
|
|||
DEFAULT_INFERENCE_PARAMS,
|
||||
type InferenceParams,
|
||||
} from "./types/runtime";
|
||||
import { Switch } from "@/components/ui/switch";
|
||||
|
||||
export const defaultInferenceParams = DEFAULT_INFERENCE_PARAMS;
|
||||
export type { InferenceParams } from "./types/runtime";
|
||||
|
|
@ -143,12 +144,16 @@ interface ChatSettingsPanelProps {
|
|||
open: boolean;
|
||||
params: InferenceParams;
|
||||
onParamsChange: (params: InferenceParams) => void;
|
||||
autoTitle: boolean;
|
||||
onAutoTitleChange: (enabled: boolean) => void;
|
||||
}
|
||||
|
||||
export function ChatSettingsPanel({
|
||||
open,
|
||||
params,
|
||||
onParamsChange,
|
||||
autoTitle,
|
||||
onAutoTitleChange,
|
||||
}: ChatSettingsPanelProps) {
|
||||
const [presets, setPresets] = useState<Preset[]>(BUILTIN_PRESETS);
|
||||
const [activePreset, setActivePreset] = useState("Default");
|
||||
|
|
@ -268,7 +273,7 @@ export function ChatSettingsPanel({
|
|||
<CollapsibleSection
|
||||
icon={SlidersHorizontalIcon}
|
||||
label="Sampling"
|
||||
defaultOpen={true}
|
||||
defaultOpen={false}
|
||||
>
|
||||
<div className="flex flex-col gap-5">
|
||||
<ParamSlider
|
||||
|
|
@ -315,9 +320,18 @@ export function ChatSettingsPanel({
|
|||
</CollapsibleSection>
|
||||
|
||||
<CollapsibleSection icon={Settings02Icon} label="Settings">
|
||||
<p className="text-xs text-muted-foreground">
|
||||
No additional settings yet.
|
||||
</p>
|
||||
<div className="flex items-center justify-between gap-3 py-1">
|
||||
<div className="min-w-0">
|
||||
<div className="text-xs font-medium">Auto title</div>
|
||||
<div className="text-[11px] text-muted-foreground">
|
||||
Generate short title after reply.
|
||||
</div>
|
||||
</div>
|
||||
<Switch
|
||||
checked={autoTitle}
|
||||
onCheckedChange={onAutoTitleChange}
|
||||
/>
|
||||
</div>
|
||||
</CollapsibleSection>
|
||||
</div>
|
||||
</div>
|
||||
|
|
|
|||
|
|
@ -105,6 +105,9 @@ export function useChatModelRuntime() {
|
|||
const message =
|
||||
error instanceof Error ? error.message : "Failed to load models";
|
||||
setModelsError(message);
|
||||
toast.error("Failed to refresh models", {
|
||||
description: message,
|
||||
});
|
||||
}
|
||||
}, [setCheckpoint, setLoras, setModels, setModelsError]);
|
||||
|
||||
|
|
@ -122,30 +125,42 @@ export function useChatModelRuntime() {
|
|||
const isLora =
|
||||
explicitIsLora ?? model?.isLora ?? (lora ? true : false);
|
||||
const displayName = model?.name || lora?.name || modelId;
|
||||
const loadingToastId = toast.loading(`Loading ${displayName}...`);
|
||||
|
||||
setModelsError(null);
|
||||
try {
|
||||
if (params.checkpoint) {
|
||||
await unloadModel({ model_path: params.checkpoint });
|
||||
async function performLoad(): Promise<void> {
|
||||
if (params.checkpoint) {
|
||||
await unloadModel({ model_path: params.checkpoint });
|
||||
}
|
||||
|
||||
await loadModel({
|
||||
model_path: modelId,
|
||||
hf_token: null,
|
||||
max_seq_length: DEFAULT_MODEL_MAX_SEQ_LENGTH,
|
||||
load_in_4bit: true,
|
||||
is_lora: isLora,
|
||||
});
|
||||
|
||||
setCheckpoint(modelId);
|
||||
await refresh();
|
||||
}
|
||||
|
||||
await loadModel({
|
||||
model_path: modelId,
|
||||
hf_token: null,
|
||||
max_seq_length: DEFAULT_MODEL_MAX_SEQ_LENGTH,
|
||||
load_in_4bit: true,
|
||||
is_lora: isLora,
|
||||
});
|
||||
let description = "Base model selected.";
|
||||
if (isLora) {
|
||||
description = "Fine-tuned (LoRA) selected.";
|
||||
}
|
||||
|
||||
setCheckpoint(modelId);
|
||||
await refresh();
|
||||
toast.success(`${displayName} loaded`, { id: loadingToastId });
|
||||
await toast.promise(performLoad(), {
|
||||
loading: `Loading ${displayName}`,
|
||||
success: `${displayName} loaded`,
|
||||
error: (err) =>
|
||||
err instanceof Error ? err.message : "Failed to load model",
|
||||
description,
|
||||
});
|
||||
} catch (error) {
|
||||
const message =
|
||||
error instanceof Error ? error.message : "Failed to load model";
|
||||
setModelsError(message);
|
||||
toast.error(message, { id: loadingToastId });
|
||||
}
|
||||
},
|
||||
[loras, models, params.checkpoint, refresh, setCheckpoint, setModelsError],
|
||||
|
|
@ -157,9 +172,19 @@ export function useChatModelRuntime() {
|
|||
}
|
||||
setModelsError(null);
|
||||
try {
|
||||
await unloadModel({ model_path: params.checkpoint });
|
||||
clearCheckpoint();
|
||||
await refresh();
|
||||
async function performUnload(): Promise<void> {
|
||||
await unloadModel({ model_path: params.checkpoint });
|
||||
clearCheckpoint();
|
||||
await refresh();
|
||||
}
|
||||
|
||||
await toast.promise(performUnload(), {
|
||||
loading: "Unloading model",
|
||||
success: "Model unloaded",
|
||||
error: (err) =>
|
||||
err instanceof Error ? err.message : "Failed to unload model",
|
||||
description: "Releases VRAM and resets inference state.",
|
||||
});
|
||||
} catch (error) {
|
||||
const message =
|
||||
error instanceof Error ? error.message : "Failed to unload model";
|
||||
|
|
|
|||
|
|
@ -11,7 +11,6 @@ import {
|
|||
SimpleTextAttachmentAdapter,
|
||||
type ThreadHistoryAdapter,
|
||||
type ThreadMessage,
|
||||
type ThreadUserMessagePart,
|
||||
WebSpeechDictationAdapter,
|
||||
type unstable_RemoteThreadListAdapter,
|
||||
useAui,
|
||||
|
|
@ -23,8 +22,10 @@ import { createAssistantStream } from "assistant-stream";
|
|||
import mammoth from "mammoth";
|
||||
import { type ReactElement, type ReactNode, useEffect, useMemo } from "react";
|
||||
import { extractText, getDocumentProxy } from "unpdf";
|
||||
import { authFetch } from "@/features/auth";
|
||||
import { createOpenAIStreamAdapter } from "./api/chat-adapter";
|
||||
import { db } from "./db";
|
||||
import { useChatRuntimeStore } from "./stores/chat-runtime-store";
|
||||
import type { MessageRecord, ModelType } from "./types";
|
||||
|
||||
const DEFAULT_SUGGESTIONS = [
|
||||
|
|
@ -149,6 +150,89 @@ class DocxAttachmentAdapter implements AttachmentAdapter {
|
|||
}
|
||||
}
|
||||
|
||||
function clip(input: string, maxLen: number): string {
|
||||
const text = input.replace(/\s+/g, " ").trim();
|
||||
if (text.length <= maxLen) return text;
|
||||
return text.slice(0, maxLen).trimEnd();
|
||||
}
|
||||
|
||||
function extractTextParts(m: ThreadMessage | undefined): string {
|
||||
if (!m) return "";
|
||||
const content = Array.isArray(m.content) ? m.content : [];
|
||||
return content
|
||||
.filter((p): p is Extract<typeof p, { type: "text" }> => p.type === "text")
|
||||
.map((p) => p.text)
|
||||
.join("")
|
||||
.trim();
|
||||
}
|
||||
|
||||
async function generateTitleWithModel(payload: {
|
||||
userText: string;
|
||||
}): Promise<string | null> {
|
||||
const params = useChatRuntimeStore.getState().params;
|
||||
if (!params.checkpoint) return null;
|
||||
|
||||
const user = clip(payload.userText, 256);
|
||||
const parts: string[] = [user];
|
||||
|
||||
function normalizeTitle(raw: string): string | null {
|
||||
let title = raw.split(/\r?\n/, 1)[0] ?? "";
|
||||
title = title.replace(/^\s*title\s*:\s*/i, "");
|
||||
title = title.replace(/[^\x20-\x7E]+/g, " ");
|
||||
title = title.replace(/["'`]+/g, "");
|
||||
title = title.replace(/[.!?:;,]+/g, " ");
|
||||
title = title.replace(/\s+/g, " ").trim();
|
||||
|
||||
// Model echo fail-safe.
|
||||
if (/\b(user|base|lora|assistant)\s*:/i.test(title)) {
|
||||
return null;
|
||||
}
|
||||
|
||||
const words = title.split(" ").filter(Boolean).slice(0, 6);
|
||||
const joined = words.join(" ").trim();
|
||||
if (!joined) return null;
|
||||
return joined.length > 60 ? joined.slice(0, 60).trimEnd() : joined;
|
||||
}
|
||||
|
||||
const response = await authFetch("/api/inference/chat/completions", {
|
||||
method: "POST",
|
||||
headers: { "Content-Type": "application/json" },
|
||||
body: JSON.stringify({
|
||||
model: params.checkpoint,
|
||||
stream: false,
|
||||
temperature: 0.2,
|
||||
top_p: 0.9,
|
||||
max_tokens: 24,
|
||||
top_k: 40,
|
||||
repetition_penalty: 1.05,
|
||||
messages: [
|
||||
{
|
||||
role: "system",
|
||||
content:
|
||||
"Write 1 concise chat title for the user's message. Rules: 2-6 words, no quotes, no punctuation, ASCII only, do not echo input. Output title only.",
|
||||
},
|
||||
{ role: "user", content: parts.join("\n") },
|
||||
],
|
||||
}),
|
||||
});
|
||||
|
||||
const body = (await response.json().catch(() => null)) as any;
|
||||
if (!response.ok) return null;
|
||||
const raw: string | undefined = body?.choices?.[0]?.message?.content;
|
||||
if (!raw) return null;
|
||||
return normalizeTitle(raw);
|
||||
}
|
||||
|
||||
const inflightTitleByKey = new Set<string>();
|
||||
|
||||
function fallbackTitleFromUserText(userText: string): string {
|
||||
const firstLine = (userText || "").split(/\r?\n/, 1)[0] ?? "";
|
||||
const cleaned = firstLine.replace(/\s+/g, " ").trim();
|
||||
const max = 48;
|
||||
if (!cleaned) return "New Chat";
|
||||
return cleaned.slice(0, max) + (cleaned.length > max ? "..." : "");
|
||||
}
|
||||
|
||||
function toThreadMessage(m: MessageRecord): ThreadMessage {
|
||||
const base = {
|
||||
id: m.id,
|
||||
|
|
@ -245,32 +329,83 @@ function createDexieAdapter(
|
|||
},
|
||||
|
||||
async generateTitle(remoteId: string, messages: readonly ThreadMessage[]) {
|
||||
const autoTitle = useChatRuntimeStore.getState().autoTitle;
|
||||
const thread = await db.threads.get(remoteId);
|
||||
const defaultTitle = "New Chat";
|
||||
|
||||
function streamTitle(title: string) {
|
||||
return createAssistantStream((c) => {
|
||||
c.appendText(title);
|
||||
c.close();
|
||||
});
|
||||
}
|
||||
|
||||
async function persistTitle(title: string): Promise<void> {
|
||||
await db.threads.update(remoteId, { title });
|
||||
if (!pairId) return;
|
||||
const paired = await db.threads
|
||||
.where("pairId")
|
||||
.equals(pairId)
|
||||
.filter((t) => t.id !== remoteId)
|
||||
.first();
|
||||
if (paired) await db.threads.update(paired.id, { title });
|
||||
}
|
||||
|
||||
if (!thread) {
|
||||
return streamTitle(defaultTitle);
|
||||
}
|
||||
|
||||
// Only generate once per thread/pair.
|
||||
if (thread.title && thread.title !== "New Chat") {
|
||||
return streamTitle(thread.title);
|
||||
}
|
||||
|
||||
const firstUser = messages.find((m) => m.role === "user");
|
||||
const textParts =
|
||||
firstUser?.content.filter(
|
||||
(part): part is Extract<ThreadUserMessagePart, { type: "text" }> =>
|
||||
part.type === "text",
|
||||
) ?? [];
|
||||
const text = textParts.map((part) => part.text).join("") || "New Chat";
|
||||
const title = text.slice(0, 60) + (text.length > 60 ? "..." : "");
|
||||
const userText = extractTextParts(firstUser) || defaultTitle;
|
||||
|
||||
await db.threads.update(remoteId, { title });
|
||||
if (!autoTitle) {
|
||||
const title = fallbackTitleFromUserText(userText);
|
||||
await persistTitle(title);
|
||||
return streamTitle(title);
|
||||
}
|
||||
|
||||
const key = pairId ? `pair:${pairId}` : `thread:${remoteId}`;
|
||||
if (inflightTitleByKey.has(key)) {
|
||||
return streamTitle(thread.title || defaultTitle);
|
||||
}
|
||||
|
||||
// Compare: wait until both threads done.
|
||||
if (pairId) {
|
||||
const paired = await db.threads
|
||||
.where("pairId")
|
||||
.equals(pairId)
|
||||
.filter((t) => t.id !== remoteId)
|
||||
.first();
|
||||
|
||||
if (paired) {
|
||||
await db.threads.update(paired.id, { title });
|
||||
const running = useChatRuntimeStore.getState().runningByThreadId;
|
||||
if (running[paired.id]) {
|
||||
setTimeout(() => {
|
||||
void createDexieAdapter(modelType, pairId).generateTitle(remoteId, messages);
|
||||
}, 600);
|
||||
return streamTitle(thread.title || defaultTitle);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return createAssistantStream((controller) => {
|
||||
controller.appendText(title);
|
||||
controller.close();
|
||||
});
|
||||
inflightTitleByKey.add(key);
|
||||
try {
|
||||
const title =
|
||||
(await generateTitleWithModel({
|
||||
userText,
|
||||
})) ||
|
||||
fallbackTitleFromUserText(userText);
|
||||
|
||||
await persistTitle(title);
|
||||
return streamTitle(title);
|
||||
} finally {
|
||||
inflightTitleByKey.delete(key);
|
||||
}
|
||||
},
|
||||
};
|
||||
}
|
||||
|
|
@ -287,10 +422,19 @@ function ThreadHistoryProvider({
|
|||
if (!remoteId) {
|
||||
return { messages: [] };
|
||||
}
|
||||
const msgs = await db.messages
|
||||
.where("threadId")
|
||||
.equals(remoteId)
|
||||
.sortBy("createdAt");
|
||||
const roleOrder: Record<string, number> = {
|
||||
system: 0,
|
||||
user: 1,
|
||||
assistant: 2,
|
||||
};
|
||||
const msgs = await db.messages.where("threadId").equals(remoteId).toArray();
|
||||
msgs.sort((a, b) => {
|
||||
if (a.createdAt !== b.createdAt) return a.createdAt - b.createdAt;
|
||||
const aOrder = roleOrder[a.role] ?? 99;
|
||||
const bOrder = roleOrder[b.role] ?? 99;
|
||||
if (aOrder !== bOrder) return aOrder - bOrder;
|
||||
return a.id < b.id ? -1 : a.id > b.id ? 1 : 0;
|
||||
});
|
||||
|
||||
return ExportedMessageRepository.fromArray(msgs.map(toThreadMessage));
|
||||
},
|
||||
|
|
@ -301,13 +445,18 @@ function ThreadHistoryProvider({
|
|||
? JSON.parse(JSON.stringify(message.content))
|
||||
: [];
|
||||
const custom = message.metadata?.custom;
|
||||
const existing = await db.messages.get(message.id);
|
||||
const createdAt =
|
||||
existing?.createdAt ??
|
||||
message.createdAt?.getTime?.() ??
|
||||
Date.now();
|
||||
await db.messages.put({
|
||||
id: message.id,
|
||||
threadId: remoteId,
|
||||
role: message.role,
|
||||
content,
|
||||
...(custom && Object.keys(custom).length > 0 && { metadata: custom }),
|
||||
createdAt: message.createdAt?.getTime() ?? Date.now(),
|
||||
createdAt,
|
||||
});
|
||||
},
|
||||
}),
|
||||
|
|
|
|||
|
|
@ -52,7 +52,9 @@ export function RegisterCompareHandle({
|
|||
}
|
||||
const currentHandles = handlesRef.current;
|
||||
currentHandles[name] = {
|
||||
append: (content) => aui.thread().append({ role: "user", content }),
|
||||
// fixes occasional reorder on reload.
|
||||
append: (content) =>
|
||||
aui.thread().append({ role: "user", content, createdAt: new Date() } as never),
|
||||
cancel: () => aui.thread().cancelRun(),
|
||||
isRunning: () => aui.thread().getState().isRunning,
|
||||
};
|
||||
|
|
|
|||
|
|
@ -6,16 +6,44 @@ import {
|
|||
type InferenceParams,
|
||||
} from "../types/runtime";
|
||||
|
||||
const AUTO_TITLE_KEY = "unsloth_chat_auto_title";
|
||||
|
||||
function canUseStorage(): boolean {
|
||||
return typeof window !== "undefined";
|
||||
}
|
||||
|
||||
function loadBool(key: string, fallback: boolean): boolean {
|
||||
if (!canUseStorage()) return fallback;
|
||||
try {
|
||||
const raw = localStorage.getItem(key);
|
||||
if (raw === null) return fallback;
|
||||
return raw === "true";
|
||||
} catch {
|
||||
return fallback;
|
||||
}
|
||||
}
|
||||
|
||||
function saveBool(key: string, value: boolean): void {
|
||||
if (!canUseStorage()) return;
|
||||
try {
|
||||
localStorage.setItem(key, value ? "true" : "false");
|
||||
} catch {
|
||||
// ignore
|
||||
}
|
||||
}
|
||||
|
||||
type ChatRuntimeStore = {
|
||||
params: InferenceParams;
|
||||
models: ChatModelSummary[];
|
||||
loras: ChatLoraSummary[];
|
||||
warmingByThreadId: Record<string, boolean>;
|
||||
runningByThreadId: Record<string, boolean>;
|
||||
autoTitle: boolean;
|
||||
modelsError: string | null;
|
||||
setParams: (params: InferenceParams) => void;
|
||||
setModels: (models: ChatModelSummary[]) => void;
|
||||
setLoras: (loras: ChatLoraSummary[]) => void;
|
||||
setThreadWarming: (threadId: string, warming: boolean) => void;
|
||||
setThreadRunning: (threadId: string, running: boolean) => void;
|
||||
setAutoTitle: (enabled: boolean) => void;
|
||||
setModelsError: (error: string | null) => void;
|
||||
setCheckpoint: (modelId: string) => void;
|
||||
clearCheckpoint: () => void;
|
||||
|
|
@ -25,20 +53,26 @@ export const useChatRuntimeStore = create<ChatRuntimeStore>((set) => ({
|
|||
params: DEFAULT_INFERENCE_PARAMS,
|
||||
models: [],
|
||||
loras: [],
|
||||
warmingByThreadId: {},
|
||||
runningByThreadId: {},
|
||||
autoTitle: loadBool(AUTO_TITLE_KEY, false),
|
||||
modelsError: null,
|
||||
setParams: (params) => set({ params }),
|
||||
setModels: (models) => set({ models }),
|
||||
setLoras: (loras) => set({ loras }),
|
||||
setThreadWarming: (threadId, warming) =>
|
||||
setThreadRunning: (threadId, running) =>
|
||||
set((state) => {
|
||||
const next = { ...state.warmingByThreadId };
|
||||
if (warming) {
|
||||
const next = { ...state.runningByThreadId };
|
||||
if (running) {
|
||||
next[threadId] = true;
|
||||
} else {
|
||||
delete next[threadId];
|
||||
}
|
||||
return { warmingByThreadId: next };
|
||||
return { runningByThreadId: next };
|
||||
}),
|
||||
setAutoTitle: (autoTitle) =>
|
||||
set(() => {
|
||||
saveBool(AUTO_TITLE_KEY, autoTitle);
|
||||
return { autoTitle };
|
||||
}),
|
||||
setModelsError: (modelsError) => set({ modelsError }),
|
||||
setCheckpoint: (modelId) =>
|
||||
|
|
|
|||
|
|
@ -62,11 +62,13 @@ export function ThreadSidebar({
|
|||
onSelect,
|
||||
onNewThread,
|
||||
onNewCompare,
|
||||
showCompare,
|
||||
}: {
|
||||
view: ChatView;
|
||||
onSelect: (view: ChatView) => void;
|
||||
onNewThread: () => void;
|
||||
onNewCompare: () => void;
|
||||
showCompare: boolean;
|
||||
}) {
|
||||
const allThreads = useLiveQuery(
|
||||
() => db.threads.orderBy("createdAt").reverse().toArray(),
|
||||
|
|
@ -112,12 +114,14 @@ export function ThreadSidebar({
|
|||
<span>New Chat</span>
|
||||
</SidebarMenuButton>
|
||||
</SidebarMenuItem>
|
||||
<SidebarMenuItem>
|
||||
<SidebarMenuButton onClick={onNewCompare}>
|
||||
<HugeiconsIcon icon={ColumnInsertIcon} />
|
||||
<span>Compare</span>
|
||||
</SidebarMenuButton>
|
||||
</SidebarMenuItem>
|
||||
{showCompare ? (
|
||||
<SidebarMenuItem>
|
||||
<SidebarMenuButton onClick={onNewCompare}>
|
||||
<HugeiconsIcon icon={ColumnInsertIcon} />
|
||||
<span>Compare</span>
|
||||
</SidebarMenuButton>
|
||||
</SidebarMenuItem>
|
||||
) : null}
|
||||
</SidebarMenu>
|
||||
</SidebarGroupContent>
|
||||
</SidebarGroup>
|
||||
|
|
|
|||
|
|
@ -63,6 +63,7 @@ export interface OpenAIChatCompletionsRequest {
|
|||
top_k: number;
|
||||
repetition_penalty: number;
|
||||
image_base64?: string;
|
||||
use_adapter?: boolean | string | null;
|
||||
}
|
||||
|
||||
export interface OpenAIChatDelta {
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue