diff --git a/studio/backend/core/inference/inference.py b/studio/backend/core/inference/inference.py index 23203613f2..6be7e63306 100644 --- a/studio/backend/core/inference/inference.py +++ b/studio/backend/core/inference/inference.py @@ -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}") diff --git a/studio/backend/routes/inference.py b/studio/backend/routes/inference.py index ae8fc31a46..3bb84d5c6b 100644 --- a/studio/backend/routes/inference.py +++ b/studio/backend/routes/inference.py @@ -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)) - diff --git a/studio/backend/utils/paths/path_utils.py b/studio/backend/utils/paths/path_utils.py index 7743952b6b..856e478bc2 100644 --- a/studio/backend/utils/paths/path_utils.py +++ b/studio/backend/utils/paths/path_utils.py @@ -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 diff --git a/studio/frontend/src/app/router.tsx b/studio/frontend/src/app/router.tsx index 27e0ee9fdf..de8564d035 100644 --- a/studio/frontend/src/app/router.tsx +++ b/studio/frontend/src/app/router.tsx @@ -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, diff --git a/studio/frontend/src/app/routes/home.tsx b/studio/frontend/src/app/routes/home.tsx deleted file mode 100644 index bf2f3a6b58..0000000000 --- a/studio/frontend/src/app/routes/home.tsx +++ /dev/null @@ -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 ; -} diff --git a/studio/frontend/src/components/assistant-ui/thread.tsx b/studio/frontend/src/components/assistant-ui/thread.tsx index 9c94e4eb22..c109c92ead 100644 --- a/studio/frontend/src/components/assistant-ui/thread.tsx +++ b/studio/frontend/src/components/assistant-ui/thread.tsx @@ -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 }> = ({ - !thread.isEmpty}> {!hideComposer && } @@ -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 ( -
-
- - Warming up model... - -
-
- ); -}; - const ThreadScrollToBottom: FC = () => { return ( diff --git a/studio/frontend/src/components/component-example.tsx b/studio/frontend/src/components/component-example.tsx deleted file mode 100644 index 3171e1abc0..0000000000 --- a/studio/frontend/src/components/component-example.tsx +++ /dev/null @@ -1,1318 +0,0 @@ -import * as React from "react"; - -import { Example, ExampleWrapper } from "@/components/example"; -import { - Accordion, - AccordionContent, - AccordionItem, - AccordionTrigger, -} from "@/components/ui/accordion"; -import { Alert, AlertDescription, AlertTitle } from "@/components/ui/alert"; -import { - AlertDialog, - AlertDialogAction, - AlertDialogCancel, - AlertDialogContent, - AlertDialogDescription, - AlertDialogFooter, - AlertDialogHeader, - AlertDialogMedia, - AlertDialogTitle, - AlertDialogTrigger, -} from "@/components/ui/alert-dialog"; -import { - Avatar, - AvatarBadge, - AvatarFallback, - AvatarGroup, - AvatarGroupCount, - AvatarImage, -} from "@/components/ui/avatar"; -import { Badge } from "@/components/ui/badge"; -import { - Breadcrumb, - BreadcrumbItem, - BreadcrumbLink, - BreadcrumbList, - BreadcrumbPage, - BreadcrumbSeparator, -} from "@/components/ui/breadcrumb"; -import { Button } from "@/components/ui/button"; -import { - Card, - CardAction, - CardContent, - CardDescription, - CardFooter, - CardHeader, - CardTitle, -} from "@/components/ui/card"; -import { Checkbox } from "@/components/ui/checkbox"; -import { - Combobox, - ComboboxContent, - ComboboxEmpty, - ComboboxInput, - ComboboxItem, - ComboboxList, -} from "@/components/ui/combobox"; -import { - Dialog, - DialogContent, - DialogDescription, - DialogFooter, - DialogHeader, - DialogTitle, - DialogTrigger, -} from "@/components/ui/dialog"; -import { - DropdownMenu, - DropdownMenuCheckboxItem, - DropdownMenuContent, - DropdownMenuGroup, - DropdownMenuItem, - DropdownMenuLabel, - DropdownMenuPortal, - DropdownMenuRadioGroup, - DropdownMenuRadioItem, - DropdownMenuSeparator, - DropdownMenuShortcut, - DropdownMenuSub, - DropdownMenuSubContent, - DropdownMenuSubTrigger, - DropdownMenuTrigger, -} from "@/components/ui/dropdown-menu"; -import { Field, FieldGroup, FieldLabel } from "@/components/ui/field"; -import { Input } from "@/components/ui/input"; -import { Label } from "@/components/ui/label"; -import { - Pagination, - PaginationContent, - PaginationItem, - PaginationLink, - PaginationNext, - PaginationPrevious, -} from "@/components/ui/pagination"; -import { Progress } from "@/components/ui/progress"; -import { RadioGroup, RadioGroupItem } from "@/components/ui/radio-group"; -import { - Select, - SelectContent, - SelectGroup, - SelectItem, - SelectTrigger, - SelectValue, -} from "@/components/ui/select"; -import { - Sheet, - SheetContent, - SheetDescription, - SheetFooter, - SheetHeader, - SheetTitle, - SheetTrigger, -} from "@/components/ui/sheet"; -import { Skeleton } from "@/components/ui/skeleton"; -import { Slider } from "@/components/ui/slider"; -import { Switch } from "@/components/ui/switch"; -import { - Table, - TableBody, - TableCell, - TableHead, - TableHeader, - TableRow, -} from "@/components/ui/table"; -import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; -import { Textarea } from "@/components/ui/textarea"; -import { Toggle } from "@/components/ui/toggle"; -import { ToggleGroup, ToggleGroupItem } from "@/components/ui/toggle-group"; -import { - Tooltip, - TooltipContent, - TooltipTrigger, -} from "@/components/ui/tooltip"; -import { - AlertCircleIcon, - BluetoothIcon, - CodeIcon, - ComputerIcon, - CreditCardIcon, - DownloadIcon, - EyeIcon, - File01Icon, - FileIcon, - FloppyDiskIcon, - FolderIcon, - FolderOpenIcon, - HelpCircleIcon, - InformationCircleIcon, - KeyboardIcon, - LanguageCircleIcon, - LayoutIcon, - LogoutIcon, - MailIcon, - MoonIcon, - MoreHorizontalCircle01Icon, - MoreVerticalCircle01Icon, - NotificationIcon, - PaintBoardIcon, - PanelRightIcon, - PlusSignIcon, - SearchIcon, - SettingsIcon, - ShieldIcon, - SunIcon, - TextBoldIcon, - TextItalicIcon, - TextUnderlineIcon, - UserIcon, -} from "@hugeicons/core-free-icons"; -import { HugeiconsIcon } from "@hugeicons/react"; - -export function ComponentExample() { - return ( - - - - - - - - - - - - - - ); -} - -function CardExample() { - return ( - - -
- mymind on Unsplash - - Observability Plus is replacing Monitoring - - Switch to the improved way to explore your data, with natural - language. Monitoring will no longer be available on the Pro plan in - November, 2025 - - - - - - - - - - - - - Allow accessory to connect? - - Do you want to allow the USB accessory to connect to this - device? - - - - Don't allow - Allow - - - - - Warning - - - - - ); -} - -const frameworks = [ - "Next.js", - "SvelteKit", - "Nuxt.js", - "Remix", - "Astro", -] as const; - -function FormExample() { - const [notifications, setNotifications] = React.useState({ - email: true, - sms: false, - push: true, - }); - const [theme, setTheme] = React.useState("light"); - - return ( - - - - User Information - Please fill in your details below - - - - - - - - File - - - New File - ⌘N - - - - New Folder - ⇧⌘N - - - - - Open Recent - - - - - Recent Projects - - - Project Alpha - - - - Project Beta - - - - - More Projects - - - - - - Project Gamma - - - - Project Delta - - - - - - - - - - Browse... - - - - - - - - - Save - ⌘S - - - - Export - ⇧⌘E - - - - - View - - setNotifications({ - ...notifications, - email: checked === true, - }) - } - > - - Show Sidebar - - - setNotifications({ - ...notifications, - sms: checked === true, - }) - } - > - - Show Status Bar - - - - - Theme - - - - - Appearance - - - - Light - - - - Dark - - - - System - - - - - - - - - - Account - - - Profile - ⇧⌘P - - - - Billing - - - - - Settings - - - - - Preferences - - - Keyboard Shortcuts - - - - Language - - - - - Notifications - - - - - - Notification Types - - - setNotifications({ - ...notifications, - push: checked === true, - }) - } - > - - Push Notifications - - - setNotifications({ - ...notifications, - email: checked === true, - }) - } - > - - Email Notifications - - - - - - - - - - - Privacy & Security - - - - - - - - - - - Help & Support - - - - Documentation - - - - - - - Sign Out - ⇧⌘Q - - - - - - - -
- -
- - Name - - - - Role - - -
- - - Framework - - - - - No frameworks found. - - {(item) => ( - - {item} - - )} - - - - - - Comments -