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:
@@ -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,
|
||||
|
||||
@@ -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}",)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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}",)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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}"),
|
||||
|
||||
Reference in New Issue
Block a user