fix(ollama): correctness — event-loop safety, HTTP errors, endpoint, types, timeout

fix(ollama): correctness — event-loop safety, HTTP errors, endpoint, types, timeout
This commit is contained in:
James Veitch
2026-06-29 01:39:45 +01:00
committed by GitHub
10 changed files with 393 additions and 59 deletions
+2 -2
View File
@@ -13,8 +13,8 @@ A collection of workflow efficiency and quality-of-life nodes built out of neces
| **Circuit Breaker** | Halts the current ComfyUI queue run gracefully without crashing the server. Wire the `status` toggle to a boolean condition to skip the rest of the queue when a condition isn't met. |
| **Ollama Client** | Configures a connection to an Ollama server (default: `http://localhost:11434`). Threads the host URL through the graph as an `OLLAMA_CLIENT` socket. |
| **Ollama Model Selector** | Fetches the live model list from Ollama and presents it as a dropdown. Outputs the selected model name. |
| **Ollama Load Model** | Loads a model into Ollama's memory using `/api/show` with `keep_alive=-1`. |
| **Ollama Unload Model** | Evicts a model from Ollama's memory using `/api/show` with `keep_alive=0`. |
| **Ollama Load Model** | Loads a model into Ollama's memory using `/api/generate` with `keep_alive=-1`. |
| **Ollama Unload Model** | Evicts a model from Ollama's memory using `/api/generate` with `keep_alive=0`. |
| **Ollama Chat Completion** | Sends a prompt (and optional conversation history) to Ollama `/api/chat` and returns the response text plus the updated history. |
| **Ollama Option — \*** | Seven composable option nodes (Temperature, Seed, Max Tokens, Top P, Top K, Repeat Penalty, Extra Body) that merge into an `OLLAMA_OPTIONS` dict wired into Chat Completion. |
| **Ollama Debug History** | Serialises an `OLLAMA_HISTORY` list to a pretty-printed JSON string for inspection. |
Binary file not shown.

Before

Width:  |  Height:  |  Size: 32 KiB

After

Width:  |  Height:  |  Size: 32 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 32 KiB

After

Width:  |  Height:  |  Size: 36 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 41 KiB

After

Width:  |  Height:  |  Size: 50 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 55 KiB

After

Width:  |  Height:  |  Size: 57 KiB

+2 -2
View File
@@ -9,8 +9,8 @@ A collection of workflow efficiency and quality-of-life nodes built out of neces
| **Circuit Breaker** | Halts the current ComfyUI queue run gracefully without crashing the server. Wire the `status` toggle to a boolean condition to skip the rest of the queue when a condition isn't met. |
| **Ollama Client** | Configures a connection to an Ollama server (default: `http://localhost:11434`). Threads the host URL through the graph as an `OLLAMA_CLIENT` socket. |
| **Ollama Model Selector** | Fetches the live model list from Ollama and presents it as a dropdown. Outputs the selected model name. |
| **Ollama Load Model** | Loads a model into Ollama's memory using `/api/show` with `keep_alive=-1`. |
| **Ollama Unload Model** | Evicts a model from Ollama's memory using `/api/show` with `keep_alive=0`. |
| **Ollama Load Model** | Loads a model into Ollama's memory using `/api/generate` with `keep_alive=-1`. |
| **Ollama Unload Model** | Evicts a model from Ollama's memory using `/api/generate` with `keep_alive=0`. |
| **Ollama Chat Completion** | Sends a prompt (and optional conversation history) to Ollama `/api/chat` and returns the response text plus the updated history. |
| **Ollama Option — \*** | Seven composable option nodes (Temperature, Seed, Max Tokens, Top P, Top K, Repeat Penalty, Extra Body) that merge into an `OLLAMA_OPTIONS` dict wired into Chat Completion. |
| **Ollama Debug History** | Serialises an `OLLAMA_HISTORY` list to a pretty-printed JSON string for inspection. |
+39 -34
View File
@@ -29,16 +29,18 @@ _STORAGE_STATE = {
"localStorage": [
{
"name": "workflow",
"value": json.dumps({
"last_node_id": 0,
"last_link_id": 0,
"nodes": [],
"links": [],
"groups": [],
"config": {},
"extra": {},
"version": 0.4,
}),
"value": json.dumps(
{
"last_node_id": 0,
"last_link_id": 0,
"nodes": [],
"links": [],
"groups": [],
"config": {},
"extra": {},
"version": 0.4,
}
),
},
{"name": "Comfy.OpenWorkflowsPaths", "value": "[]"},
{"name": "Comfy.ActiveWorkflowIndex", "value": "0"},
@@ -456,28 +458,28 @@ async def scene_ollama_lifecycle(page: Page, out: Path) -> None:
async () => {
const graph = window.app.graph;
// OllamaClient — far left
// OllamaClient — far left, vertically centred relative to the chain
const client = LiteGraph.createNode("OllamaClient");
client.pos = [40, 140];
client.pos = [60, 300];
graph.add(client);
const hostWidget = client.widgets && client.widgets.find(w => w.name === "host");
if (hostWidget) hostWidget.value = "http://localhost:11434";
// OllamaLoadModel — centre-left
// OllamaLoadModel — generous gap right of client
const load = LiteGraph.createNode("OllamaLoadModel");
load.pos = [300, 40];
load.pos = [380, 80];
graph.add(load);
// OllamaChatCompletion — centre-right
// OllamaChatCompletion — wide node, plenty of space to the right of load
const chat = LiteGraph.createNode("OllamaChatCompletion");
chat.pos = [580, 40];
chat.pos = [720, 40];
graph.add(chat);
const twPrompt = chat.widgets && chat.widgets.find(w => w.name === "prompt");
if (twPrompt) twPrompt.value = "Describe this image in one sentence.";
// OllamaUnloadModel — far right
// OllamaUnloadModel — far right, vertically offset to match chat's outputs
const unload = LiteGraph.createNode("OllamaUnloadModel");
unload.pos = [900, 120];
unload.pos = [1200, 280];
graph.add(unload);
// Populate model dropdowns from live Ollama
@@ -487,28 +489,27 @@ async def scene_ollama_lifecycle(page: Page, out: Path) -> None:
const data = await resp.json();
const models = data.models || [];
if (models.length) {
for (const node of [load, chat]) {
const w = node.widgets && node.widgets.find(w => w.name === "model");
for (const n of [load, chat]) {
const w = n.widgets && n.widgets.find(w => w.name === "model");
if (w) { w.options = w.options || {}; w.options.values = models; w.value = models[0]; }
}
}
}
} catch(e) {}
// Wire: client → load.client, client → chat.client, client → unload.client
// Wire: client → load, chat, unload
client.connect(0, load, 0);
client.connect(0, chat, 0);
client.connect(0, unload, 0);
// Wire: load.model_name (output 0) → chat.model_name input (optional)
// This forces Load to run before Chat.
// Wire: load.model_name → chat.model_name (forces Load before Chat)
const chatModelNameSlot = chat.inputs
? chat.inputs.findIndex(i => i.name === "model_name")
: -1;
if (chatModelNameSlot >= 0) load.connect(0, chat, chatModelNameSlot);
// Wire: chat.model_name (output 2) → unload.model (input 1)
// Wire: chat.response (output 0) → unload.passthrough (optional input)
// Wire: chat.model_name (output 2) → unload.model
// Wire: chat.response (output 0) → unload.passthrough
const unloadModelSlot = unload.inputs
? unload.inputs.findIndex(i => i.name === "model")
: 1;
@@ -518,24 +519,26 @@ async def scene_ollama_lifecycle(page: Page, out: Path) -> None:
chat.connect(2, unload, unloadModelSlot >= 0 ? unloadModelSlot : 1);
if (unloadPassSlot >= 0) chat.connect(0, unload, unloadPassSlot);
await new Promise(r => setTimeout(r, 1200));
// Let nodes render and auto-size before reading bounding box
await new Promise(r => setTimeout(r, 1500));
window.app.canvas.setDirty(true, true);
window.app.canvas.draw(true, true);
await new Promise(r => setTimeout(r, 300));
const nodes = [client, load, chat, unload];
const minX = Math.min(...nodes.map(n => n.pos[0])) - 20;
const minY = Math.min(...nodes.map(n => n.pos[1])) - 20;
const maxX = Math.max(...nodes.map(n => n.pos[0] + n.size[0])) + 20;
const maxY = Math.max(...nodes.map(n => n.pos[1] + n.size[1])) + 20;
const minX = Math.min(...nodes.map(n => n.pos[0])) - 30;
const minY = Math.min(...nodes.map(n => n.pos[1])) - 30;
const maxX = Math.max(...nodes.map(n => n.pos[0] + n.size[0])) + 30;
const maxY = Math.max(...nodes.map(n => n.pos[1] + n.size[1])) + 30;
return { pos: [minX, minY], size: [maxX - minX, maxY - minY] };
}
"""
)
await asyncio.sleep(2.0)
await asyncio.sleep(1.5)
await _redraw(page)
await _frame_node(page, info["pos"], info["size"], scale=0.85)
await _capture(page, out, info["pos"], info["size"], scale=0.85)
await _frame_node(page, info["pos"], info["size"], scale=0.72)
await _capture(page, out, info["pos"], info["size"], scale=0.72)
async def scene_ollama_options(page: Page, out: Path) -> None:
@@ -609,7 +612,9 @@ async def main() -> int:
async with async_playwright() as pw:
browser = await pw.chromium.launch(headless=True)
context = await browser.new_context(viewport=VIEWPORT, storage_state=_STORAGE_STATE)
context = await browser.new_context(
viewport=VIEWPORT, storage_state=_STORAGE_STATE
)
page = await context.new_page()
print(f"Opening {COMFYUI_URL} …")
+38 -14
View File
@@ -31,12 +31,17 @@ class OllamaClientType(str):
def _run_async(coro):
"""Run an async coroutine synchronously using a fresh event loop."""
loop = asyncio.new_event_loop()
"""Run an async coroutine synchronously, safe inside a running event loop."""
try:
return loop.run_until_complete(coro)
finally:
loop.close()
asyncio.get_running_loop()
# Called from within a running loop (e.g. ComfyUI's async executor).
# Spin up a worker thread with its own loop to avoid "loop already running".
import concurrent.futures
with concurrent.futures.ThreadPoolExecutor(max_workers=1) as pool:
return pool.submit(asyncio.run, coro).result()
except RuntimeError:
return asyncio.run(coro)
async def _fetch_models(host: str) -> list[str]:
@@ -56,7 +61,7 @@ async def _fetch_models(host: str) -> list[str]:
return []
async def _post_json(url: str, payload: dict) -> dict:
async def _post_json(url: str, payload: dict, *, timeout: float = 120.0) -> dict:
"""POST JSON to url, return parsed response dict."""
import aiohttp
@@ -65,17 +70,22 @@ async def _post_json(url: str, payload: dict) -> dict:
async with session.post(
url,
json=payload,
timeout=aiohttp.ClientTimeout(total=120),
timeout=aiohttp.ClientTimeout(total=timeout),
) as resp:
if resp.status >= 400:
body = await resp.text()
raise RuntimeError(
f"Ollama returned HTTP {resp.status} for {url}: {body[:300]}"
)
return await resp.json()
except aiohttp.ClientConnectionError as exc:
raise RuntimeError(f"Cannot reach Ollama at {url}: {exc}") from exc
# Populated at import time; empty list if Ollama is unreachable.
_DEFAULT_MODELS: list[str] = _run_async(_fetch_models("http://localhost:11434")) or [
"(start Ollama to see models)"
]
# Static placeholder — avoids a network call at import time (which always fails
# in CI and slows cold-start). The JS refresh button calls /dv/ollama/models
# at runtime to populate the live list.
_DEFAULT_MODELS: list[str] = ["(⟳ click Refresh models)"]
# ---------------------------------------------------------------------------
@@ -172,7 +182,11 @@ class OllamaLoadModel:
if not model.strip():
raise ValueError("model name cannot be empty")
_run_async(
_post_json(f"{client}/api/show", {"model": model, "keep_alive": "-1"})
_post_json(
f"{client}/api/generate",
{"model": model, "keep_alive": -1, "stream": False},
timeout=300.0,
)
)
return (model,)
@@ -206,7 +220,13 @@ class OllamaUnloadModel:
def unload_model(self, client, model: str, passthrough: str = ""):
if not model.strip():
raise ValueError("model name cannot be empty")
_run_async(_post_json(f"{client}/api/show", {"model": model, "keep_alive": 0}))
_run_async(
_post_json(
f"{client}/api/generate",
{"model": model, "keep_alive": 0, "stream": False},
timeout=30.0,
)
)
return (model, passthrough)
@@ -231,6 +251,7 @@ class OllamaChatCompletion:
# Wire OllamaLoadModel.model_name here to guarantee load runs
# before chat and to override the dropdown with the wired value.
"model_name": ("STRING", {"forceInput": True}),
"timeout_secs": ("INT", {"default": 300, "min": 30, "max": 3600}),
},
}
@@ -248,6 +269,7 @@ class OllamaChatCompletion:
history=None,
options=None,
model_name=None,
timeout_secs=300,
):
effective_model = (
model_name.strip() if model_name and model_name.strip() else model
@@ -265,7 +287,9 @@ class OllamaChatCompletion:
}
if options:
payload["options"] = options
result = _run_async(_post_json(f"{client}/api/chat", payload))
result = _run_async(
_post_json(f"{client}/api/chat", payload, timeout=float(timeout_secs))
)
response_text = result.get("message", {}).get("content", "")
updated = list(history)
updated.append({"role": "user", "content": prompt})
+20 -4
View File
@@ -39,12 +39,28 @@ async function refreshModelDropdown(node, host) {
}
/**
* Locate the host string for a node — either from a connected OllamaClient
* widget value or from the default.
* Locate the host string for a node.
*
* First checks the node's own widgets (OllamaClient has a "host" widget).
* Otherwise traverses graph links to find a connected OllamaClient node and
* reads its "host" widget — this is the common case for downstream nodes
* (ModelSelector, LoadModel, ChatCompletion) that receive the client socket.
*/
function getHostFromNode(node) {
const clientWidget = node.widgets?.find(w => w.name === "host");
if (clientWidget) return clientWidget.value;
const ownHostWidget = node.widgets?.find(w => w.name === "host");
if (ownHostWidget) return ownHostWidget.value;
for (const input of node.inputs ?? []) {
if (!input.link) continue;
const link = node.graph?.links[input.link];
if (!link) continue;
const sourceNode = node.graph?.getNodeById(link.origin_id);
if (sourceNode?.type === "OllamaClient") {
const hostWidget = sourceNode.widgets?.find(w => w.name === "host");
if (hostWidget?.value) return hostWidget.value;
}
}
return "http://localhost:11434";
}
+292 -3
View File
@@ -15,6 +15,7 @@ BDD coverage:
features/us6_history_inspection.feature
"""
import asyncio
import json
import pytest
@@ -36,9 +37,161 @@ from comfydv.ollama import (
OllamaOptionTopP,
OllamaUnloadModel,
_fetch_models,
_post_json,
_run_async,
)
# ---------------------------------------------------------------------------
# Infrastructure tests (Issues 1–3: event loop safety, HTTP errors, timeout)
# ---------------------------------------------------------------------------
class TestInfrastructure:
# ---- Issue 1: _run_async event loop safety --------------------------------
def test_run_async_works_from_running_loop(self):
"""Issue 1: _run_async must work when called from inside a running event loop.
ComfyUI drives nodes inside its own asyncio event loop. Calling
_run_async (which currently spawns a new loop) from within that context
raises "Cannot run the event loop while another loop is running".
"""
async def _simple():
return 42
async def _caller():
# Simulates a synchronous ComfyUI node method being invoked from
# within the server's async event loop.
return _run_async(_simple())
result = asyncio.run(_caller())
assert result == 42
# ---- Issue 2: _post_json HTTP error checking ------------------------------
def test_post_json_raises_on_http_4xx(self, monkeypatch):
"""Issue 2: _post_json must raise RuntimeError on 4xx responses.
Currently it silently returns the response body as a dict.
"""
import aiohttp
class FakeResponse:
status = 422
async def text(self):
return "Unprocessable"
async def json(self):
return {"error": "Unprocessable"}
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
class FakeSession:
def post(self, *args, **kwargs):
return FakeResponse()
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
monkeypatch.setattr(aiohttp, "ClientSession", lambda: FakeSession())
with pytest.raises(RuntimeError, match="HTTP 422"):
_run_async(_post_json("http://localhost/test", {}))
def test_post_json_raises_on_http_5xx(self, monkeypatch):
"""Issue 2: _post_json must raise RuntimeError on 5xx responses.
Currently it silently returns the response body as a dict.
"""
import aiohttp
class FakeResponse:
status = 500
async def text(self):
return "Internal Server Error"
async def json(self):
return {"error": "Internal Server Error"}
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
class FakeSession:
def post(self, *args, **kwargs):
return FakeResponse()
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
monkeypatch.setattr(aiohttp, "ClientSession", lambda: FakeSession())
with pytest.raises(RuntimeError, match="HTTP 500"):
_run_async(_post_json("http://localhost/test", {}))
# ---- Issue 3: _post_json timeout parameter --------------------------------
def test_post_json_timeout_forwarded(self, monkeypatch):
"""Issue 3: _post_json must accept a timeout kwarg and forward it to aiohttp.
Currently _post_json has no timeout parameter, so the call raises TypeError.
"""
import aiohttp
captured = {}
class FakeResponse:
status = 200
async def text(self):
return "{}"
async def json(self):
return {}
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
class FakeSession:
def post(self, url, *, json=None, timeout=None):
captured["timeout"] = timeout
return FakeResponse()
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
monkeypatch.setattr(aiohttp, "ClientSession", lambda: FakeSession())
_run_async(_post_json("http://localhost/test", {}, timeout=42.0))
assert captured.get("timeout") is not None, (
"_post_json did not forward a timeout object to aiohttp"
)
assert captured["timeout"].total == 42.0, (
f"Expected ClientTimeout(total=42.0) but got total={captured['timeout'].total}"
)
# ---------------------------------------------------------------------------
# US1 — Ollama Connection (us1_ollama_connection.feature)
# ---------------------------------------------------------------------------
@@ -128,6 +281,98 @@ class TestUS3ModelLifecycle:
with pytest.raises(ValueError, match="cannot be empty"):
OllamaUnloadModel().unload_model(client="http://localhost:11434", model="")
# ---- Issue 4: Load/Unload API endpoint ------------------------------------
def test_load_uses_api_generate_not_api_show(self, monkeypatch):
"""Issue 4: load_model should POST to /api/generate, not /api/show."""
captured = {}
async def fake_post(url, payload, *, timeout=120.0):
captured["url"] = url
return {}
import comfydv.ollama as ollama_mod
monkeypatch.setattr(ollama_mod, "_post_json", fake_post)
OllamaLoadModel().load_model(
client="http://localhost:11434", model="test-model"
)
assert captured.get("url", "").endswith("/api/generate"), (
f"Expected URL ending in /api/generate but got: {captured.get('url')}"
)
def test_unload_uses_api_generate_not_api_show(self, monkeypatch):
"""Issue 4: unload_model should POST to /api/generate, not /api/show."""
captured = {}
async def fake_post(url, payload, *, timeout=120.0):
captured["url"] = url
return {}
import comfydv.ollama as ollama_mod
monkeypatch.setattr(ollama_mod, "_post_json", fake_post)
OllamaUnloadModel().unload_model(
client="http://localhost:11434", model="test-model"
)
assert captured.get("url", "").endswith("/api/generate"), (
f"Expected URL ending in /api/generate but got: {captured.get('url')}"
)
# ---- Issue 5: keep_alive integer type -------------------------------------
def test_load_keep_alive_is_integer_negative_one(self, monkeypatch):
"""Issue 5: load_model must send keep_alive as integer -1, not string '-1'."""
captured = {}
async def fake_post(url, payload, *, timeout=120.0):
captured["payload"] = payload
return {}
import comfydv.ollama as ollama_mod
monkeypatch.setattr(ollama_mod, "_post_json", fake_post)
OllamaLoadModel().load_model(
client="http://localhost:11434", model="test-model"
)
payload = captured.get("payload", {})
assert payload.get("keep_alive") == -1, (
f"Expected keep_alive=-1 (int) but got {payload.get('keep_alive')!r}"
)
assert isinstance(payload.get("keep_alive"), int), (
f"keep_alive must be int, got {type(payload.get('keep_alive'))}"
)
def test_unload_keep_alive_is_integer_zero(self, monkeypatch):
"""Issue 5: unload_model must send keep_alive as integer 0."""
captured = {}
async def fake_post(url, payload, *, timeout=120.0):
captured["payload"] = payload
return {}
import comfydv.ollama as ollama_mod
monkeypatch.setattr(ollama_mod, "_post_json", fake_post)
OllamaUnloadModel().unload_model(
client="http://localhost:11434", model="test-model"
)
payload = captured.get("payload", {})
assert payload.get("keep_alive") == 0, (
f"Expected keep_alive=0 (int) but got {payload.get('keep_alive')!r}"
)
assert isinstance(payload.get("keep_alive"), int), (
f"keep_alive must be int, got {type(payload.get('keep_alive'))}"
)
@pytest.mark.integration
def test_load_model_returns_name(self, ollama_host, skip_if_no_ollama):
"""Scenario: Load Model loads model into Ollama memory."""
@@ -204,7 +449,7 @@ class TestUS4ChatCompletion:
"""model_name kwarg takes precedence over the COMBO widget value."""
captured = {}
async def fake_post(url, payload):
async def fake_post(url, payload, *, timeout=120.0):
captured["model"] = payload.get("model")
return {"message": {"content": "ok"}}
@@ -221,6 +466,50 @@ class TestUS4ChatCompletion:
assert effective == "wired-value"
assert captured.get("model") == "wired-value"
# ---- Issue 6: Chat timeout widget -----------------------------------------
def test_chat_has_timeout_secs_input(self):
"""Issue 6: OllamaChatCompletion must expose a timeout_secs input widget.
Currently INPUT_TYPES() does not include 'timeout_secs'.
"""
input_types = OllamaChatCompletion.INPUT_TYPES()
all_inputs = {
**input_types.get("required", {}),
**input_types.get("optional", {}),
}
assert "timeout_secs" in all_inputs, (
"OllamaChatCompletion.INPUT_TYPES() must include 'timeout_secs' "
"(in required or optional)"
)
def test_chat_timeout_forwarded_to_http(self, monkeypatch):
"""Issue 6: timeout_secs kwarg must be forwarded to _post_json's timeout param.
Currently the chat() method does not accept timeout_secs, so this raises
TypeError before reaching _post_json.
"""
captured = {}
async def fake_post(url, payload, *, timeout=120.0):
captured["timeout"] = timeout
return {"message": {"content": "ok"}}
import comfydv.ollama as ollama_mod
monkeypatch.setattr(ollama_mod, "_post_json", fake_post)
OllamaChatCompletion().chat(
client="http://x",
model="m",
prompt="p",
timeout_secs=600,
)
assert captured.get("timeout") == 600.0, (
f"Expected timeout forwarded as 600.0 but got {captured.get('timeout')!r}"
)
@pytest.mark.integration
def test_single_turn_returns_non_empty_response(
self, ollama_host, skip_if_no_ollama
@@ -341,8 +630,8 @@ class TestUS5ComposableOptions:
history=[],
options=opts2,
)
r1, _ = OllamaChatCompletion().chat(**kwargs)
r2, _ = OllamaChatCompletion().chat(**kwargs)
r1, _, _model = OllamaChatCompletion().chat(**kwargs)
r2, _, _model = OllamaChatCompletion().chat(**kwargs)
assert r1 == r2