From bf54b54ca9cbc429e8370d6385d69bc8f7eaf21a Mon Sep 17 00:00:00 2001 From: limbicnation Date: Fri, 8 May 2026 03:43:58 +0200 Subject: [PATCH] fix: handle Ollama llama-runner crashes with structured error categorization Replace the catch-all RuntimeError wrapper in OllamaClient.generate_streaming with a StreamResult dataclass that categorises failures (ok, timeout, transient, model_crash, server_error, unavailable). Llama-runner crashes (HTTP 500 with "runner terminated" / "exit status" / "load failed") are detected and surfaced as user-actionable messages in the ComfyUI prompt output instead of bubbling up as Python stacktraces. PromptGenerator, NegativePrompt, and PromptRefiner nodes now branch on result.kind: model_crash / server_error / unavailable surface the message directly (subprocess fallback would also fail), while timeout / transient fall through to the existing subprocess fallback path. Adds 6 unit tests covering each error class plus the success path. Updates the existing PromptRefiner seed tests to return StreamResult from mocks. --- nodes/adapters/ollama_client.py | 100 ++++++++++++++-- nodes/negative_prompt_node.py | 11 +- nodes/prompt_generator_node.py | 14 ++- nodes/prompt_refiner_node.py | 11 +- tests/unit/test_ollama_client.py | 189 +++++++++++++++++++++++++++++- tests/unit/test_prompt_refiner.py | 5 +- 6 files changed, 306 insertions(+), 24 deletions(-) diff --git a/nodes/adapters/ollama_client.py b/nodes/adapters/ollama_client.py index 1902edc..6e2a565 100644 --- a/nodes/adapters/ollama_client.py +++ b/nodes/adapters/ollama_client.py @@ -13,6 +13,7 @@ import logging import subprocess import threading import time +from dataclasses import dataclass, field from typing import Any, ClassVar logger = logging.getLogger(__name__) @@ -30,6 +31,30 @@ try: except ImportError: OLLAMA_API_AVAILABLE = False +try: + from ollama import ResponseError as OllamaResponseError +except ImportError: + OllamaResponseError = None # older ollama package or package missing + + +@dataclass +class StreamResult: + """Outcome of a streaming generation call. + + kind values: + ok - generation succeeded; text holds the response + timeout - per-chunk or total timeout exceeded + transient - network / connection issue; subprocess fallback may help + model_crash - llama runner subprocess died (HTTP 500); subprocess won't help + server_error - other non-2xx response from Ollama server + unavailable - ollama python package not installed + """ + + text: str | None = None + kind: str = "ok" + message: str = field(default="") + + try: import comfy.utils @@ -154,7 +179,7 @@ class OllamaClient: timeout: int, pbar: Any | None = None, seed: int | None = None, - ) -> str | None: + ) -> StreamResult: """ Stream ollama.generate() with per-chunk and total timeout enforcement. @@ -162,10 +187,15 @@ class OllamaClient: seed: Optional seed for deterministic generation (e.g. PromptRefiner) Returns: - Full response text, or None on failure. + StreamResult with `kind` indicating success or failure category. On + success, `text` holds the concatenated response. On failure, `message` + contains a user-facing description suitable for surfacing in the UI. """ if not OLLAMA_API_AVAILABLE: - return None + return StreamResult( + kind="unavailable", + message="Ollama python package not installed.", + ) chunks: list[str] = [] start = time.monotonic() @@ -203,7 +233,10 @@ class OllamaClient: elapsed = time.monotonic() - start if elapsed >= timeout: self._log(f"Total timeout ({timeout}s) reached") - break + return StreamResult( + kind="timeout", + message=f"Total timeout ({timeout}s) reached during generation.", + ) result_holder.clear() t = threading.Thread(target=_iter_next, args=(it,), daemon=True) @@ -216,7 +249,10 @@ class OllamaClient: if t.is_alive(): label = "first chunk" if not got_first_chunk else "chunk" self._log(f"Timeout waiting for {label} ({wait_time:.0f}s)") - return None + return StreamResult( + kind="timeout", + message=f"Timeout waiting for {label} ({wait_time:.0f}s).", + ) if "error" in result_holder: raise result_holder["error"] @@ -240,24 +276,66 @@ class OllamaClient: except ConnectionError as e: self._log(f"Connection error: {e}") - return None + return StreamResult(kind="transient", message=f"Ollama not reachable: {e}") except TimeoutError as e: self._log(f"Timeout error: {e}") - return None + return StreamResult(kind="timeout", message=f"Ollama timeout: {e}") except (TypeError, AttributeError) as e: # API contract mismatch — propagate so caller knows the integration broke raise RuntimeError(f"Ollama API contract mismatch: {e}") from e except Exception as e: - # Unknown failure — log and propagate with context - raise RuntimeError(f"Streaming generation failed: {e}") from e + return self._classify_streaming_exception(e, model) if not chunks: - return None + return StreamResult( + kind="transient", + message="Generation returned no content.", + ) elapsed = time.monotonic() - start full_text = "".join(chunks) self._log(f"Streaming complete: {len(full_text)} chars in {elapsed:.1f}s") - return full_text + return StreamResult(text=full_text, kind="ok") + + def _classify_streaming_exception(self, exc: Exception, model: str) -> StreamResult: + """Map an exception raised during streaming to a StreamResult. + + Detects llama-runner crashes (HTTP 500 with 'runner terminated' / + 'exit status' / 'load failed' in the body) and surfaces an actionable + message instead of a raw stacktrace. + """ + msg = str(exc).lower() + is_response_error = OllamaResponseError is not None and isinstance(exc, OllamaResponseError) + # Even without OllamaResponseError class, recognise duck-typed responses + status = getattr(exc, "status_code", None) + + crash_signals = ( + "runner terminated", + "runner process has terminated", + "exit status", + "load failed", + ) + if (is_response_error or status is not None) and status == 500 and any(s in msg for s in crash_signals): + friendly = ( + f"Ollama llama runner crashed loading '{model}'. " + "This usually means the model file is corrupt, incompatible, " + "or out of VRAM. Try: (1) restart Ollama " + "(`sudo systemctl restart ollama`), " + "(2) switch to a known-good model like 'qwen3:8b', " + "(3) re-pull or rebuild the model." + ) + self._log(f"Model crash detected: {exc}") + return StreamResult(kind="model_crash", message=friendly) + + if is_response_error or status is not None: + self._log(f"Server error {status}: {exc}") + return StreamResult( + kind="server_error", + message=f"Ollama returned status {status}: {exc}", + ) + + self._log(f"Streaming failed: {exc}") + return StreamResult(kind="transient", message=f"Streaming failed: {exc}") def generate_subprocess( self, diff --git a/nodes/negative_prompt_node.py b/nodes/negative_prompt_node.py index 6659d9e..62efa61 100644 --- a/nodes/negative_prompt_node.py +++ b/nodes/negative_prompt_node.py @@ -144,7 +144,7 @@ Negative prompt:""" logger.info("Generating negative for style='%s'", style) # Generate via streaming - output = client.generate_streaming( + result = client.generate_streaming( model=model, prompt=negative_prompt_text, temperature=temperature, @@ -152,8 +152,13 @@ Negative prompt:""" timeout=timeout, ) - if output is None: - # Fallback to subprocess + if result.kind == "ok" and result.text is not None: + output = result.text + elif result.kind in ("model_crash", "server_error", "unavailable"): + # Subprocess fallback would also fail; surface message directly. + return (f"[NegativePrompt] {result.message}",) + else: + # timeout / transient — try subprocess success, output = client.generate_subprocess(model, negative_prompt_text, timeout) if not success: return (f"[NegativePrompt] Generation failed: {output}",) diff --git a/nodes/prompt_generator_node.py b/nodes/prompt_generator_node.py index 766372b..b48ff2c 100644 --- a/nodes/prompt_generator_node.py +++ b/nodes/prompt_generator_node.py @@ -492,7 +492,7 @@ Format the response as a single, detailed photography prompt.""", effective_timeout = min(int(timeout * 1.3), 600) print(f"[PromptGenerator] Cold start detected, effective timeout: {effective_timeout}s") - output = client.generate_streaming( + result = client.generate_streaming( model=model, prompt=prompt, temperature=temperature, @@ -501,8 +501,8 @@ Format the response as a single, detailed photography prompt.""", pbar=pbar, ) - if output is not None: - output = output.strip() + if result.kind == "ok" and result.text is not None: + output = result.text.strip() if not include_reasoning: output = extract_final_prompt(output) @@ -515,7 +515,13 @@ Format the response as a single, detailed photography prompt.""", else: return ("[PromptGenerator] Generation returned empty result.",) - print("[PromptGenerator] Streaming failed, falling back to subprocess") + # Subprocess fallback would also fail for these classes; surface the + # message immediately so the user gets actionable guidance. + if result.kind in ("model_crash", "server_error", "unavailable"): + print(f"[PromptGenerator] {result.kind}: {result.message}") + return (f"[PromptGenerator] {result.message}",) + + print(f"[PromptGenerator] Streaming failed ({result.kind}), falling back to subprocess") # Fallback to subprocess (no temperature/top_p control) success, output = client.generate_subprocess(model, prompt, timeout) diff --git a/nodes/prompt_refiner_node.py b/nodes/prompt_refiner_node.py index 65863a3..0eea086 100644 --- a/nodes/prompt_refiner_node.py +++ b/nodes/prompt_refiner_node.py @@ -163,7 +163,7 @@ Refined prompt:""" pass_seed = None if effective_seed is None else effective_seed + i # Generate refined version - output = client.generate_streaming( + result = client.generate_streaming( model=model, prompt=refinement, temperature=temperature, @@ -173,8 +173,13 @@ Refined prompt:""" seed=pass_seed, ) - if output is None: - # Fallback to subprocess + if result.kind == "ok" and result.text is not None: + output = result.text + elif result.kind in ("model_crash", "server_error", "unavailable"): + # Subprocess fallback would also fail; surface directly. + return (f"[PromptRefiner] Pass {i + 1}: {result.message}",) + else: + # timeout / transient — try subprocess success, output = client.generate_subprocess(model, refinement, timeout) if not success: return (f"[PromptRefiner] Pass {i + 1} failed: {output}",) diff --git a/tests/unit/test_ollama_client.py b/tests/unit/test_ollama_client.py index a3e9245..674faa7 100644 --- a/tests/unit/test_ollama_client.py +++ b/tests/unit/test_ollama_client.py @@ -4,7 +4,7 @@ Unit tests for OllamaClient adapter. from unittest.mock import MagicMock, patch -from nodes.adapters.ollama_client import OllamaClient +from nodes.adapters.ollama_client import OllamaClient, StreamResult class TestOllamaClientDiscovery: @@ -219,3 +219,190 @@ class TestOllamaClientSubprocess: success, output = client.generate_subprocess("qwen3:8b", "prompt", 30) assert success is False assert "not found" in output.lower() or "Install" in output + + +class _FakeResponseError(Exception): + """Duck-typed stand-in for ollama.ResponseError with status_code attr.""" + + def __init__(self, message: str, status_code: int): + super().__init__(message) + self.status_code = status_code + + +class TestOllamaClientStreamingErrors: + """Test suite for generate_streaming error categorisation.""" + + @staticmethod + def _patch_module(mock_ollama, response_error_cls=None): + """Helper: swap module-level ollama + ResponseError, return restore fn.""" + import nodes.adapters.ollama_client as client_module + + orig = ( + getattr(client_module, "ollama", None), + getattr(client_module, "OLLAMA_API_AVAILABLE", False), + getattr(client_module, "OllamaResponseError", None), + ) + client_module.ollama = mock_ollama + client_module.OLLAMA_API_AVAILABLE = True + client_module.OllamaResponseError = response_error_cls + + def restore(): + client_module.ollama = orig[0] + client_module.OLLAMA_API_AVAILABLE = orig[1] + client_module.OllamaResponseError = orig[2] + + return restore + + def test_runner_crash_returns_model_crash_kind(self): + """500 + 'runner terminated' should be classified as model_crash.""" + crash_exc = _FakeResponseError( + "llama runner process has terminated: %!w() (status code: 500)", + status_code=500, + ) + + def _bad_iter(): + yield from () + raise crash_exc + + mock_ollama = MagicMock() + mock_ollama.generate.return_value = _bad_iter() + + restore = self._patch_module(mock_ollama, response_error_cls=_FakeResponseError) + try: + client = OllamaClient() + result = client.generate_streaming( + model="qwen3-limbicnation", + prompt="hello", + temperature=0.7, + top_p=0.9, + timeout=30, + ) + finally: + restore() + + assert isinstance(result, StreamResult) + assert result.kind == "model_crash" + assert "qwen3-limbicnation" in result.message + assert "restart Ollama" in result.message + assert "qwen3:8b" in result.message + assert result.text is None + + def test_runner_crash_via_duck_typed_error_when_class_unavailable(self): + """Even without ollama.ResponseError, status_code=500 + signal triggers model_crash.""" + crash_exc = _FakeResponseError("Load failed: llama runner terminated", status_code=500) + + def _bad_iter(): + yield from () + raise crash_exc + + mock_ollama = MagicMock() + mock_ollama.generate.return_value = _bad_iter() + + # Simulate older ollama package (no ResponseError export) + restore = self._patch_module(mock_ollama, response_error_cls=None) + try: + client = OllamaClient() + result = client.generate_streaming( + model="bad-model", + prompt="x", + temperature=0.7, + top_p=0.9, + timeout=30, + ) + finally: + restore() + + assert result.kind == "model_crash" + + def test_generic_server_error_returns_server_error(self): + """Non-crash 5xx (e.g. 503) should be classified as server_error.""" + exc = _FakeResponseError("service unavailable", status_code=503) + + def _bad_iter(): + yield from () + raise exc + + mock_ollama = MagicMock() + mock_ollama.generate.return_value = _bad_iter() + + restore = self._patch_module(mock_ollama, response_error_cls=_FakeResponseError) + try: + client = OllamaClient() + result = client.generate_streaming( + model="qwen3:8b", + prompt="x", + temperature=0.7, + top_p=0.9, + timeout=30, + ) + finally: + restore() + + assert result.kind == "server_error" + assert "503" in result.message + + def test_connection_error_returns_transient(self): + """ConnectionError raised during streaming should be transient.""" + + def _bad_iter(): + yield from () + raise ConnectionError("connection refused") + + mock_ollama = MagicMock() + mock_ollama.generate.return_value = _bad_iter() + + restore = self._patch_module(mock_ollama, response_error_cls=_FakeResponseError) + try: + client = OllamaClient() + result = client.generate_streaming( + model="qwen3:8b", + prompt="x", + temperature=0.7, + top_p=0.9, + timeout=30, + ) + finally: + restore() + + assert result.kind == "transient" + assert "not reachable" in result.message + + def test_success_returns_ok(self): + """Normal streaming should yield kind='ok' with concatenated text.""" + + def _good_iter(): + yield {"response": "hello "} + yield {"response": "world"} + + mock_ollama = MagicMock() + mock_ollama.generate.return_value = _good_iter() + + restore = self._patch_module(mock_ollama, response_error_cls=_FakeResponseError) + try: + client = OllamaClient() + result = client.generate_streaming( + model="qwen3:8b", + prompt="x", + temperature=0.7, + top_p=0.9, + timeout=30, + ) + finally: + restore() + + assert result.kind == "ok" + assert result.text == "hello world" + + def test_unavailable_when_api_missing(self): + """If ollama package is unavailable, kind should be 'unavailable'.""" + with patch("nodes.adapters.ollama_client.OLLAMA_API_AVAILABLE", False): + client = OllamaClient() + result = client.generate_streaming( + model="qwen3:8b", + prompt="x", + temperature=0.7, + top_p=0.9, + timeout=30, + ) + assert result.kind == "unavailable" + assert result.text is None diff --git a/tests/unit/test_prompt_refiner.py b/tests/unit/test_prompt_refiner.py index 578abc5..457de52 100644 --- a/tests/unit/test_prompt_refiner.py +++ b/tests/unit/test_prompt_refiner.py @@ -2,6 +2,7 @@ from unittest.mock import patch +from nodes.adapters.ollama_client import StreamResult from nodes.prompt_refiner_node import PromptRefinerNode @@ -15,7 +16,7 @@ class TestPromptRefinerSeed: def _capture_seed(*, model, prompt, temperature, top_p, timeout, pbar, seed): captured_seeds.append(seed) - return f"refined with seed={seed}" + return StreamResult(text=f"refined with seed={seed}", kind="ok") with ( patch.object(node, "REFINEMENT_PROMPT", "{prompt}"), @@ -44,7 +45,7 @@ class TestPromptRefinerSeed: def _capture_seed(*, model, prompt, temperature, top_p, timeout, pbar, seed): captured_seeds.append(seed) - return "refined" + return StreamResult(text="refined", kind="ok") with ( patch.object(node, "REFINEMENT_PROMPT", "{prompt}"),