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.
This commit is contained in:
limbicnation
2026-05-08 03:43:58 +02:00
parent b5e1731ff7
commit bf54b54ca9
6 changed files with 306 additions and 24 deletions
+89 -11
View File
@@ -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,
+8 -3
View File
@@ -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}",)
+10 -4
View File
@@ -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)
+8 -3
View File
@@ -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}",)
+188 -1
View File
@@ -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(<nil>) (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
+3 -2
View File
@@ -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}"),