From 5fdd7d0e7596bc9050c98a5cbecdc447db1d85d4 Mon Sep 17 00:00:00 2001 From: Jon Saad-Falcon <41205309+jonsaadfalcon@users.noreply.github.com> Date: Fri, 13 Mar 2026 21:42:26 -0700 Subject: [PATCH] feat: model catalogue with download/delete and Qwen3.5 auto-pull (#54) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * feat: model catalogue with download/delete and auto-pull Qwen3.5 - Desktop boot: start server immediately with fallback model (qwen3:0.6b), then pull preferred model (qwen3.5:4b) and remaining Qwen3.5 variants that fit in RAM in the background. No more broken "Select model" state. - Backend: add POST /v1/models/pull and DELETE /v1/models/{name} endpoints so the frontend can trigger model downloads and deletions via Ollama. - Frontend: redesign CommandPalette (Cmd+K) with two tabs — "Installed" shows pulled models with select/delete, "Download Models" shows a catalogue of popular models plus a custom model input field. - Fix ollama_has_model() to use exact tag matching instead of prefix matching, preventing false positives. * fix: streaming, model switching, second-largest default, and tests - Streaming: use direct engine streaming for non-tool requests so tokens arrive in real-time instead of being batched by the agent bridge. Add error handling to _handle_stream so engine errors surface as content chunks instead of silent failures. - Model selection: pick the second-largest Qwen3.5 model that fits (leaves headroom for OS/apps) instead of the absolute largest. - Model switching: abort in-flight stream when the user changes models mid-generation, preventing stale-model errors. Improve error messages in catch blocks. - Tests: add tests/server/test_model_management.py with 11 tests covering model pull/delete endpoints, streaming error resilience, direct-engine streaming bypass, and model listing. All 100 server tests pass. --- desktop/src-tauri/Cargo.lock | 2 +- desktop/src-tauri/src/lib.rs | 124 +++++++- frontend/src/components/Chat/InputArea.tsx | 24 +- frontend/src/components/CommandPalette.tsx | 323 +++++++++++++++++---- frontend/src/lib/api.ts | 22 ++ src/openjarvis/server/routes.py | 119 +++++++- tests/server/test_model_management.py | 261 +++++++++++++++++ 7 files changed, 804 insertions(+), 71 deletions(-) create mode 100644 tests/server/test_model_management.py diff --git a/desktop/src-tauri/Cargo.lock b/desktop/src-tauri/Cargo.lock index ba5f0965..55112ab4 100644 --- a/desktop/src-tauri/Cargo.lock +++ b/desktop/src-tauri/Cargo.lock @@ -2573,7 +2573,7 @@ dependencies = [ [[package]] name = "openjarvis-desktop" -version = "1.0.0" +version = "0.1.0" dependencies = [ "reqwest 0.12.28", "serde", diff --git a/desktop/src-tauri/src/lib.rs b/desktop/src-tauri/src/lib.rs index d1f26918..43bdddb2 100644 --- a/desktop/src-tauri/src/lib.rs +++ b/desktop/src-tauri/src/lib.rs @@ -8,7 +8,74 @@ use tokio::sync::Mutex; const OLLAMA_PORT: u16 = 11434; const JARVIS_PORT: u16 = 8222; -const DEFAULT_MODEL: &str = "qwen3:0.6b"; + +/// Tiny fallback model — always available even on low-RAM machines. +const FALLBACK_MODEL: &str = "qwen3:0.6b"; + +/// Qwen3.5 model variants, ordered smallest to largest. +/// Each entry is (ollama_tag, approximate_download_size_gb, min_ram_gb). +const QWEN35_MODELS: &[(&str, f64, f64)] = &[ + ("qwen3.5:0.8b", 1.0, 4.0), + ("qwen3.5:2b", 2.7, 6.0), + ("qwen3.5:4b", 3.4, 8.0), + ("qwen3.5:9b", 6.6, 12.0), + ("qwen3.5:27b", 17.0, 24.0), + ("qwen3.5:35b", 24.0, 32.0), + ("qwen3.5:122b", 81.0, 96.0), +]; + +/// Get total system RAM in GB. +fn total_ram_gb() -> f64 { + #[cfg(target_os = "macos")] + { + use std::process::Command; + if let Ok(output) = Command::new("sysctl").args(["-n", "hw.memsize"]).output() { + if let Ok(s) = String::from_utf8(output.stdout) { + if let Ok(bytes) = s.trim().parse::() { + return bytes as f64 / (1024.0 * 1024.0 * 1024.0); + } + } + } + } + #[cfg(target_os = "linux")] + { + if let Ok(contents) = std::fs::read_to_string("/proc/meminfo") { + for line in contents.lines() { + if line.starts_with("MemTotal:") { + if let Some(kb_str) = line.split_whitespace().nth(1) { + if let Ok(kb) = kb_str.parse::() { + return kb as f64 / (1024.0 * 1024.0); + } + } + } + } + } + } + 8.0 +} + +/// Return the list of Qwen3.5 models that fit on this machine, smallest first. +fn models_that_fit() -> Vec<&'static str> { + let ram = total_ram_gb(); + QWEN35_MODELS + .iter() + .filter(|(_, _, min_ram)| ram >= *min_ram) + .map(|(tag, _, _)| *tag) + .collect() +} + +/// Pick the second-largest Qwen3.5 model that fits on this machine. +/// This leaves headroom for the OS / other apps while still providing +/// a capable model. Falls back to the largest if only one fits, or +/// to FALLBACK_MODEL if none fit. +fn preferred_model() -> &'static str { + let fitting = models_that_fit(); + match fitting.len() { + 0 => FALLBACK_MODEL, + 1 => fitting[0], + n => fitting[n - 2], // second-largest + } +} /// Resolve full path to a binary by checking common locations. /// macOS .app bundles don't inherit the shell PATH, so we probe manually. @@ -212,7 +279,11 @@ async fn ollama_has_model(model: &str) -> bool { return models.iter().any(|m| { m.get("name") .and_then(|n| n.as_str()) - .map(|n| n.starts_with(model.split(':').next().unwrap_or(model))) + .map(|n| { + n == model + || n.strip_suffix(":latest") == Some(model) + || model.strip_suffix(":latest") == Some(n) + }) .unwrap_or(false) }); } @@ -285,19 +356,22 @@ async fn boot_backend(backend: SharedBackend, status: SharedStatus) { s.detail = "Inference engine ready.".into(); } - // Phase 2: Ensure a default model exists + // Phase 2: Ensure at least one model exists. + // Pull the small fallback first so the server can start immediately, + // then pull the preferred model and remaining Qwen3.5 variants in + // the background (Phase 4). { let mut s = status.lock().await; s.phase = "model".into(); - s.detail = format!("Checking for model {}...", DEFAULT_MODEL); + s.detail = format!("Checking for model {}...", FALLBACK_MODEL); } - if !ollama_has_model(DEFAULT_MODEL).await { + if !ollama_has_model(FALLBACK_MODEL).await { { let mut s = status.lock().await; - s.detail = format!("Downloading {}... (this may take a minute)", DEFAULT_MODEL); + s.detail = format!("Downloading {}... (this may take a minute)", FALLBACK_MODEL); } - if let Err(e) = pull_model(DEFAULT_MODEL).await { + if let Err(e) = pull_model(FALLBACK_MODEL).await { let mut s = status.lock().await; s.error = Some(format!("Failed to download model: {}", e)); return; @@ -331,8 +405,21 @@ async fn boot_backend(backend: SharedBackend, status: SharedStatus) { return; } + // Use the preferred model if it's already pulled, otherwise fallback. + let pref = preferred_model(); + let startup_model = if ollama_has_model(pref).await { + pref + } else { + FALLBACK_MODEL + }; + let mut cmd = tokio::process::Command::new(&uv_bin); - cmd.args(["run", "jarvis", "serve", "--port", &JARVIS_PORT.to_string(), "--agent", "simple"]) + cmd.args([ + "run", "jarvis", "serve", + "--port", &JARVIS_PORT.to_string(), + "--model", startup_model, + "--agent", "simple", + ]) .stdout(std::process::Stdio::null()) .stderr(std::process::Stdio::null()) .current_dir(project_root.as_ref().unwrap()); @@ -368,6 +455,27 @@ async fn boot_backend(backend: SharedBackend, status: SharedStatus) { s.phase = "ready".into(); s.detail = "All systems ready.".into(); } + + // Phase 4: Pull preferred model + remaining Qwen3.5 variants in the + // background. The server is already running with the fallback, so + // the user can chat immediately. As each model finishes pulling it + // appears in the /v1/models list (Ollama serves it automatically). + let fitting = models_that_fit(); + let pref_bg = preferred_model().to_string(); + tokio::spawn(async move { + // Pull the preferred model first so it becomes available quickly. + if !ollama_has_model(&pref_bg).await { + let _ = pull_model(&pref_bg).await; + } + // Then pull remaining models that fit. + for model in fitting { + if model != pref_bg && model != FALLBACK_MODEL { + if !ollama_has_model(model).await { + let _ = pull_model(model).await; + } + } + } + }); } // --------------------------------------------------------------------------- diff --git a/frontend/src/components/Chat/InputArea.tsx b/frontend/src/components/Chat/InputArea.tsx index 8acc118a..aaa244d5 100644 --- a/frontend/src/components/Chat/InputArea.tsx +++ b/frontend/src/components/Chat/InputArea.tsx @@ -26,6 +26,22 @@ export function InputArea() { const { state: speechState, available: speechAvailable, startRecording, stopRecording } = useSpeech(); + // Abort in-flight stream when the user switches models mid-generation. + // This prevents errors from trying to continue a stream with a stale model. + const prevModelRef = useRef(selectedModel); + useEffect(() => { + if (prevModelRef.current !== selectedModel && streamState.isStreaming) { + abortRef.current?.abort(); + if (timerRef.current) { + clearInterval(timerRef.current); + timerRef.current = null; + } + resetStream(); + abortRef.current = null; + } + prevModelRef.current = selectedModel; + }, [selectedModel, streamState.isStreaming, resetStream]); + const micDisabled = !speechEnabled || !speechAvailable || streamState.isStreaming; const micReason: 'not-enabled' | 'no-backend' | 'streaming' | undefined = !speechEnabled ? 'not-enabled' @@ -189,9 +205,13 @@ export function InputArea() { } } } catch (err: any) { - if (err.name !== 'AbortError') { + if (err.name === 'AbortError') { + // User cancelled or model switch — keep whatever was accumulated + if (!accumulatedContent) accumulatedContent = '(Generation stopped)'; + } else { + const errMsg = err?.message || String(err); accumulatedContent = - accumulatedContent || 'Error: Failed to get response.'; + accumulatedContent || `Error: ${errMsg}`; } } finally { if (!accumulatedContent) { diff --git a/frontend/src/components/CommandPalette.tsx b/frontend/src/components/CommandPalette.tsx index f58d170a..4bb79d6f 100644 --- a/frontend/src/components/CommandPalette.tsx +++ b/frontend/src/components/CommandPalette.tsx @@ -1,22 +1,53 @@ import { useState, useRef, useEffect } from 'react'; -import { Search, Cpu, X } from 'lucide-react'; +import { Search, Cpu, X, Download, Loader2, Trash2, Check } from 'lucide-react'; import { useAppStore } from '../lib/store'; +import { pullModel, deleteModel, fetchModels } from '../lib/api'; + +/** Popular models that users can download from the catalogue. */ +const CATALOGUE_MODELS = [ + { id: 'qwen3.5:0.8b', size: '~1 GB', desc: 'Qwen 3.5 0.8B — fast, lightweight' }, + { id: 'qwen3.5:2b', size: '~2.7 GB', desc: 'Qwen 3.5 2B' }, + { id: 'qwen3.5:4b', size: '~3.4 GB', desc: 'Qwen 3.5 4B — recommended default' }, + { id: 'qwen3.5:9b', size: '~6.6 GB', desc: 'Qwen 3.5 9B' }, + { id: 'qwen3.5:27b', size: '~17 GB', desc: 'Qwen 3.5 27B' }, + { id: 'qwen3.5:35b', size: '~24 GB', desc: 'Qwen 3.5 35B' }, + { id: 'qwen3.5:122b', size: '~81 GB', desc: 'Qwen 3.5 122B — largest' }, + { id: 'llama3.3:latest', size: '~4.9 GB', desc: 'Llama 3.3 8B' }, + { id: 'mistral:latest', size: '~4.1 GB', desc: 'Mistral 7B' }, + { id: 'gemma3:latest', size: '~3.3 GB', desc: 'Gemma 3 4B' }, + { id: 'deepseek-r1:7b', size: '~4.7 GB', desc: 'DeepSeek R1 7B' }, + { id: 'phi4:latest', size: '~9.1 GB', desc: 'Phi-4 14B' }, +]; + +type Tab = 'installed' | 'catalogue'; export function CommandPalette() { const [query, setQuery] = useState(''); const [selectedIdx, setSelectedIdx] = useState(0); + const [tab, setTab] = useState('installed'); + const [pulling, setPulling] = useState(null); + const [pullError, setPullError] = useState(null); + const [pullSuccess, setPullSuccess] = useState(null); + const [deleting, setDeleting] = useState(null); + const [customModel, setCustomModel] = useState(''); const inputRef = useRef(null); const models = useAppStore((s) => s.models); const selectedModel = useAppStore((s) => s.selectedModel); const setSelectedModel = useAppStore((s) => s.setSelectedModel); + const setModels = useAppStore((s) => s.setModels); const setCommandPaletteOpen = useAppStore((s) => s.setCommandPaletteOpen); - const filtered = query - ? models.filter((m) => - m.id.toLowerCase().includes(query.toLowerCase()), - ) - : models; + const installedIds = new Set(models.map((m) => m.id)); + + const filtered = tab === 'installed' + ? (query + ? models.filter((m) => m.id.toLowerCase().includes(query.toLowerCase())) + : models) + : CATALOGUE_MODELS.filter((m) => + !installedIds.has(m.id) && + (!query || m.id.toLowerCase().includes(query.toLowerCase()) || m.desc.toLowerCase().includes(query.toLowerCase())) + ); useEffect(() => { inputRef.current?.focus(); @@ -24,13 +55,66 @@ export function CommandPalette() { useEffect(() => { setSelectedIdx(0); - }, [query]); + }, [query, tab]); + + // Clear success message after a delay + useEffect(() => { + if (pullSuccess) { + const t = setTimeout(() => setPullSuccess(null), 3000); + return () => clearTimeout(t); + } + }, [pullSuccess]); const handleSelect = (modelId: string) => { setSelectedModel(modelId); setCommandPaletteOpen(false); }; + const refreshModels = async () => { + try { + const m = await fetchModels(); + setModels(m); + } catch {} + }; + + const handlePull = async (modelId: string) => { + setPulling(modelId); + setPullError(null); + try { + await pullModel(modelId); + setPullSuccess(modelId); + await refreshModels(); + // Auto-select the newly pulled model + setSelectedModel(modelId); + } catch (e: any) { + setPullError(e.message || 'Download failed'); + } finally { + setPulling(null); + } + }; + + const handleDelete = async (modelId: string) => { + if (!confirm(`Delete model ${modelId}? You can re-download it later.`)) return; + setDeleting(modelId); + try { + await deleteModel(modelId); + await refreshModels(); + if (selectedModel === modelId) { + const remaining = models.filter((m) => m.id !== modelId); + if (remaining.length > 0) setSelectedModel(remaining[0].id); + } + } catch {} finally { + setDeleting(null); + } + }; + + const handleCustomPull = async () => { + const name = customModel.trim(); + if (!name) return; + await handlePull(name); + setCustomModel(''); + }; + const handleKeyDown = (e: React.KeyboardEvent) => { if (e.key === 'Escape') { setCommandPaletteOpen(false); @@ -40,9 +124,9 @@ export function CommandPalette() { } else if (e.key === 'ArrowUp') { e.preventDefault(); setSelectedIdx((i) => Math.max(i - 1, 0)); - } else if (e.key === 'Enter' && filtered.length > 0) { + } else if (e.key === 'Enter' && tab === 'installed' && filtered.length > 0) { e.preventDefault(); - handleSelect(filtered[selectedIdx].id); + handleSelect((filtered[selectedIdx] as any).id); } }; @@ -64,6 +148,27 @@ export function CommandPalette() { }} onClick={(e) => e.stopPropagation()} > + {/* Tabs */} +
+ {(['installed', 'catalogue'] as Tab[]).map((t) => ( + + ))} +
+ {/* Search input */}
setQuery(e.target.value)} onKeyDown={handleKeyDown} - placeholder="Search models..." + placeholder={tab === 'installed' ? 'Search installed models...' : 'Search models to download...'} className="flex-1 bg-transparent outline-none text-sm" style={{ color: 'var(--color-text)' }} /> @@ -89,49 +194,163 @@ export function CommandPalette() {
+ {/* Status messages */} + {pullError && ( +
+ {pullError} +
+ )} + {pullSuccess && ( +
+ Downloaded {pullSuccess} successfully +
+ )} + {/* Results */}
- {filtered.length === 0 ? ( -
- {models.length === 0 ? 'No models available — is the server running?' : 'No matching models'} -
+ {tab === 'installed' ? ( + /* ── Installed models tab ── */ + filtered.length === 0 ? ( +
+ {models.length === 0 + ? 'No models available — switch to "Download Models" to get started' + : 'No matching models'} +
+ ) : ( + (filtered as typeof models).map((model, idx) => { + const isActive = model.id === selectedModel; + const isSelected = idx === selectedIdx; + const isDeleting = deleting === model.id; + return ( +
setSelectedIdx(idx)} + > + + +
+ ); + }) + ) ) : ( - filtered.map((model, idx) => { - const isActive = model.id === selectedModel; - const isSelected = idx === selectedIdx; - return ( -
+ {isPulling ? ( + <> Downloading... + ) : justInstalled ? ( + <> Installed + ) : ( + <> Download + )} + - {isActive && ( - - Active - - )} - - ); - }) + ); + })} + + {/* Custom model input */} +
+
+ Or enter any Ollama model name: +
+
+ setCustomModel(e.target.value)} + onKeyDown={(e) => { if (e.key === 'Enter') { e.preventDefault(); handleCustomPull(); } }} + placeholder="e.g. codellama:7b" + className="flex-1 text-sm px-3 py-1.5 rounded-lg outline-none" + style={{ + background: 'var(--color-bg-secondary)', + color: 'var(--color-text)', + border: '1px solid var(--color-border)', + }} + /> + +
+
+ )} @@ -140,9 +359,15 @@ export function CommandPalette() { className="flex items-center gap-4 px-4 py-2 text-[11px]" style={{ borderTop: '1px solid var(--color-border)', color: 'var(--color-text-tertiary)' }} > - ↑↓ Navigate - Enter Select - Esc Close + {tab === 'installed' ? ( + <> + ↑↓ Navigate + Enter Select + Esc Close + + ) : ( + Models are downloaded from the Ollama registry + )} diff --git a/frontend/src/lib/api.ts b/frontend/src/lib/api.ts index c3f3dc3d..44ccda90 100644 --- a/frontend/src/lib/api.ts +++ b/frontend/src/lib/api.ts @@ -100,6 +100,28 @@ export async function fetchModels(): Promise { return data.data || []; } +export async function pullModel(modelName: string): Promise { + const res = await fetch(`${getBase()}/v1/models/pull`, { + method: 'POST', + headers: { 'Content-Type': 'application/json' }, + body: JSON.stringify({ model: modelName }), + }); + if (!res.ok) { + const detail = await res.text().catch(() => res.statusText); + throw new Error(`Failed to pull model: ${detail}`); + } +} + +export async function deleteModel(modelName: string): Promise { + const res = await fetch(`${getBase()}/v1/models/${encodeURIComponent(modelName)}`, { + method: 'DELETE', + }); + if (!res.ok) { + const detail = await res.text().catch(() => res.statusText); + throw new Error(`Failed to delete model: ${detail}`); + } +} + export async function fetchSavings(): Promise { const res = await fetch(`${getBase()}/v1/savings`); if (!res.ok) throw new Error(`Failed to fetch savings: ${res.status}`); diff --git a/src/openjarvis/server/routes.py b/src/openjarvis/server/routes.py index 3a3fe8a9..5fc4e62c 100644 --- a/src/openjarvis/server/routes.py +++ b/src/openjarvis/server/routes.py @@ -48,7 +48,11 @@ async def chat_completions(request_body: ChatCompletionRequest, request: Request if request_body.stream: bus = getattr(request.app.state, "bus", None) - if agent is not None and bus is not None: + # Use the agent stream bridge only when tools are present (the + # bridge runs agent.run() synchronously and word-splits the result, + # so it can't stream tokens in real-time). For plain chat, stream + # directly from the engine for true token-by-token output. + if agent is not None and bus is not None and request_body.tools: return await _handle_agent_stream(agent, bus, model, request_body) return await _handle_stream(engine, model, request_body) @@ -181,21 +185,42 @@ async def _handle_stream(engine, model: str, req: ChatCompletionRequest): ) yield f"data: {first_chunk.model_dump_json()}\n\n" - # Stream content - async for token in engine.stream( - messages, - model=model, - temperature=req.temperature, - max_tokens=req.max_tokens, - ): - chunk = ChatCompletionChunk( + try: + # Stream content + async for token in engine.stream( + messages, + model=model, + temperature=req.temperature, + max_tokens=req.max_tokens, + ): + chunk = ChatCompletionChunk( + id=chunk_id, + model=model, + choices=[StreamChoice( + delta=DeltaMessage(content=token), + )], + ) + yield f"data: {chunk.model_dump_json()}\n\n" + except Exception as exc: + # Surface errors as a content chunk so the frontend can + # display them instead of silently failing. + import logging + logging.getLogger("openjarvis.server").error( + "Stream error: %s", exc, exc_info=True, + ) + error_chunk = ChatCompletionChunk( id=chunk_id, model=model, choices=[StreamChoice( - delta=DeltaMessage(content=token), + delta=DeltaMessage( + content=f"\n\nError during generation: {exc}", + ), + finish_reason="stop", )], ) - yield f"data: {chunk.model_dump_json()}\n\n" + yield f"data: {error_chunk.model_dump_json()}\n\n" + yield "data: [DONE]\n\n" + return # Send finish chunk finish_chunk = ChatCompletionChunk( @@ -226,6 +251,78 @@ async def list_models(request: Request) -> ModelListResponse: ) +@router.post("/v1/models/pull") +async def pull_model(request: Request): + """Pull / download a model from the Ollama registry.""" + body = await request.json() + model_name = body.get("model", "").strip() + if not model_name: + raise HTTPException(status_code=400, detail="'model' field is required") + + engine = request.app.state.engine + engine_name = getattr(request.app.state, "engine_name", "") + # Only Ollama supports pulling + if engine_name != "ollama" and getattr(engine, "engine_id", "") != "ollama": + raise HTTPException( + status_code=501, + detail="Model pulling is only supported with the Ollama engine", + ) + + import httpx as _httpx + + host = getattr(engine, "_host", "http://localhost:11434") + client = _httpx.Client(base_url=host, timeout=600.0) + try: + resp = client.post( + "/api/pull", + json={"name": model_name, "stream": False}, + ) + resp.raise_for_status() + except (_httpx.ConnectError, _httpx.TimeoutException) as exc: + raise HTTPException(status_code=502, detail=f"Ollama unreachable: {exc}") + except _httpx.HTTPStatusError as exc: + raise HTTPException( + status_code=exc.response.status_code, + detail=f"Ollama error: {exc.response.text[:300]}", + ) + finally: + client.close() + + return {"status": "ok", "model": model_name} + + +@router.delete("/v1/models/{model_name:path}") +async def delete_model(model_name: str, request: Request): + """Delete a model from Ollama.""" + engine = request.app.state.engine + engine_name = getattr(request.app.state, "engine_name", "") + if engine_name != "ollama" and getattr(engine, "engine_id", "") != "ollama": + raise HTTPException(status_code=501, detail="Only supported with Ollama engine") + + import httpx as _httpx + + host = getattr(engine, "_host", "http://localhost:11434") + client = _httpx.Client(base_url=host, timeout=30.0) + try: + resp = client.request( + "DELETE", + "/api/delete", + json={"name": model_name}, + ) + resp.raise_for_status() + except (_httpx.ConnectError, _httpx.TimeoutException) as exc: + raise HTTPException(status_code=502, detail=f"Ollama unreachable: {exc}") + except _httpx.HTTPStatusError as exc: + raise HTTPException( + status_code=exc.response.status_code, + detail=f"Ollama error: {exc.response.text[:300]}", + ) + finally: + client.close() + + return {"status": "deleted", "model": model_name} + + @router.get("/v1/savings") async def savings(request: Request): """Return savings summary compared to cloud providers. diff --git a/tests/server/test_model_management.py b/tests/server/test_model_management.py new file mode 100644 index 00000000..7c5f99ef --- /dev/null +++ b/tests/server/test_model_management.py @@ -0,0 +1,261 @@ +"""Tests for model pull / delete API endpoints and streaming resilience.""" + +from __future__ import annotations + +import json +from unittest.mock import MagicMock, patch + +import pytest + +fastapi = pytest.importorskip("fastapi") +from fastapi.testclient import TestClient # noqa: E402 + +from openjarvis.server.app import create_app # noqa: E402 + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + + +def _make_engine(engine_id="mock", models=None): + engine = MagicMock() + engine.engine_id = engine_id + engine.health.return_value = True + engine.list_models.return_value = models or ["test-model"] + engine.generate.return_value = { + "content": "Hello", + "usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8}, + "model": "test-model", + "finish_reason": "stop", + } + + async def mock_stream(messages, *, model, temperature=0.7, max_tokens=1024, **kw): + for token in ["Hello", " ", "world"]: + yield token + + engine.stream = mock_stream + return engine + + +def _make_ollama_engine(models=None): + """Create a mock engine that looks like OllamaEngine.""" + engine = _make_engine(engine_id="ollama", models=models) + engine._host = "http://localhost:11434" + return engine + + +def _app(engine, engine_name="mock"): + return create_app(engine, "test-model", engine_name=engine_name) + + +# --------------------------------------------------------------------------- +# Model pull endpoint +# --------------------------------------------------------------------------- + + +class TestModelPull: + def test_pull_requires_model_field(self): + engine = _make_ollama_engine() + client = TestClient(_app(engine, engine_name="ollama")) + resp = client.post("/v1/models/pull", json={}) + assert resp.status_code == 400 + assert "model" in resp.json()["detail"].lower() + + def test_pull_rejects_non_ollama_engine(self): + engine = _make_engine(engine_id="vllm") + client = TestClient(_app(engine, engine_name="vllm")) + resp = client.post("/v1/models/pull", json={"model": "foo"}) + assert resp.status_code == 501 + + def test_pull_success(self): + engine = _make_ollama_engine() + client = TestClient(_app(engine, engine_name="ollama")) + + + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.raise_for_status = MagicMock() + + with patch("httpx.Client") as MockClient: + instance = MockClient.return_value + instance.post.return_value = mock_resp + instance.close = MagicMock() + + resp = client.post("/v1/models/pull", json={"model": "qwen3.5:4b"}) + + assert resp.status_code == 200 + data = resp.json() + assert data["status"] == "ok" + assert data["model"] == "qwen3.5:4b" + + def test_pull_ollama_unreachable(self): + engine = _make_ollama_engine() + client = TestClient(_app(engine, engine_name="ollama")) + + import httpx + + with patch("httpx.Client") as MockClient: + instance = MockClient.return_value + instance.post.side_effect = httpx.ConnectError("refused") + instance.close = MagicMock() + + resp = client.post("/v1/models/pull", json={"model": "foo"}) + + assert resp.status_code == 502 + + +# --------------------------------------------------------------------------- +# Model delete endpoint +# --------------------------------------------------------------------------- + + +class TestModelDelete: + def test_delete_rejects_non_ollama(self): + engine = _make_engine(engine_id="vllm") + client = TestClient(_app(engine, engine_name="vllm")) + resp = client.delete("/v1/models/test-model") + assert resp.status_code == 501 + + def test_delete_success(self): + engine = _make_ollama_engine() + client = TestClient(_app(engine, engine_name="ollama")) + + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.raise_for_status = MagicMock() + + with patch("httpx.Client") as MockClient: + instance = MockClient.return_value + instance.request.return_value = mock_resp + instance.close = MagicMock() + + resp = client.delete("/v1/models/qwen3:0.6b") + + assert resp.status_code == 200 + data = resp.json() + assert data["status"] == "deleted" + assert data["model"] == "qwen3:0.6b" + + +# --------------------------------------------------------------------------- +# Streaming resilience +# --------------------------------------------------------------------------- + + +class TestStreamingResilience: + """Verify streaming handles errors gracefully.""" + + def test_stream_error_returns_error_chunk(self): + """When the engine raises during streaming, error is sent as content.""" + engine = _make_engine() + + async def failing_stream(messages, *, model, **kw): + yield "partial" + raise RuntimeError("model not found") + + engine.stream = failing_stream + app = create_app(engine, "test-model") + client = TestClient(app) + + resp = client.post("/v1/chat/completions", json={ + "model": "bad-model", + "messages": [{"role": "user", "content": "Hello"}], + "stream": True, + }) + assert resp.status_code == 200 + + # Should contain partial content + error message + [DONE] + text = resp.text + assert "partial" in text + assert "model not found" in text + assert "[DONE]" in text + + def test_stream_tokens_arrive(self): + """Verify tokens stream through correctly (not batched).""" + engine = _make_engine() + app = create_app(engine, "test-model") + client = TestClient(app) + + resp = client.post("/v1/chat/completions", json={ + "model": "test-model", + "messages": [{"role": "user", "content": "Hello"}], + "stream": True, + }) + assert resp.status_code == 200 + + # Collect tokens + tokens = [] + for line in resp.text.strip().split("\n"): + if line.startswith("data:") and "[DONE]" not in line: + data = json.loads(line[5:].strip()) + content = data.get("choices", [{}])[0].get("delta", {}).get("content") + if content: + tokens.append(content) + + assert tokens == ["Hello", " ", "world"] + + def test_stream_without_agent_uses_direct_engine(self): + """When no tools in request, streaming should use engine.stream directly + even if an agent is configured (for real token-by-token output).""" + from openjarvis.agents._stubs import AgentResult + + engine = _make_engine() + agent = MagicMock() + agent.agent_id = "simple" + agent.run.return_value = AgentResult( + content="agent response", turns=1, + ) + + app = create_app(engine, "test-model", agent=agent) + client = TestClient(app) + + resp = client.post("/v1/chat/completions", json={ + "model": "test-model", + "messages": [{"role": "user", "content": "Hello"}], + "stream": True, + # No tools — should use direct engine stream + }) + assert resp.status_code == 200 + + # Should get engine tokens, not agent response + tokens = [] + for line in resp.text.strip().split("\n"): + if line.startswith("data:") and "[DONE]" not in line: + data = json.loads(line[5:].strip()) + content = data.get("choices", [{}])[0].get("delta", {}).get("content") + if content: + tokens.append(content) + + assert "".join(tokens) == "Hello world" + # Agent.run should NOT have been called + agent.run.assert_not_called() + + +# --------------------------------------------------------------------------- +# Models endpoint +# --------------------------------------------------------------------------- + + +class TestModelsEndpointExtended: + def test_models_list_multiple(self): + engine = _make_engine( + models=["qwen3.5:4b", "qwen3.5:9b", "qwen3:0.6b"], + ) + app = create_app(engine, "qwen3.5:4b") + client = TestClient(app) + resp = client.get("/v1/models") + assert resp.status_code == 200 + ids = [m["id"] for m in resp.json()["data"]] + assert "qwen3.5:4b" in ids + assert "qwen3.5:9b" in ids + assert "qwen3:0.6b" in ids + + def test_models_empty_engine(self): + """When engine.list_models() returns empty, endpoint still succeeds.""" + engine = _make_engine(models=[]) + app = create_app(engine, "test-model") + client = TestClient(app) + resp = client.get("/v1/models") + assert resp.status_code == 200 + # The endpoint returns whatever list_models() gives + assert resp.json()["object"] == "list"