fix(ollama): correctness — event-loop safety, HTTP errors, endpoint, types, timeout
fix(ollama): correctness — event-loop safety, HTTP errors, endpoint, types, timeout
@@ -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. |
|
||||
|
||||
|
Before Width: | Height: | Size: 32 KiB After Width: | Height: | Size: 32 KiB |
|
Before Width: | Height: | Size: 32 KiB After Width: | Height: | Size: 36 KiB |
|
Before Width: | Height: | Size: 41 KiB After Width: | Height: | Size: 50 KiB |
|
Before Width: | Height: | Size: 55 KiB After Width: | Height: | Size: 57 KiB |
@@ -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. |
|
||||
|
||||
@@ -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} …")
|
||||
|
||||
@@ -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})
|
||||
|
||||
@@ -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";
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
|
||||