Merge pull request #11 from darth-veitcher/feature/ollama-auth-headers-and-format-string-inline-display

feat: Ollama auth headers + response cache, FormatString inline display
This commit is contained in:
James Veitch
2026-07-04 13:45:07 +01:00
committed by GitHub
6 changed files with 707 additions and 52 deletions
+9
View File
@@ -6,6 +6,9 @@ from .ollama import (
OllamaClient,
OllamaChatCompletion,
OllamaDebugHistory,
OllamaHeaderBasicAuth,
OllamaHeaderBearerToken,
OllamaHeaderCustom,
OllamaHistoryLength,
OllamaLoadModel,
OllamaModelSelector,
@@ -43,6 +46,9 @@ NODE_CLASS_MAPPINGS = {
"OllamaOptionExtraBody": OllamaOptionExtraBody,
"OllamaDebugHistory": OllamaDebugHistory,
"OllamaHistoryLength": OllamaHistoryLength,
"OllamaHeaderBasicAuth": OllamaHeaderBasicAuth,
"OllamaHeaderBearerToken": OllamaHeaderBearerToken,
"OllamaHeaderCustom": OllamaHeaderCustom,
}
# A dictionary that contains the friendly/humanly readable titles for the nodes
@@ -65,6 +71,9 @@ NODE_DISPLAY_NAME_MAPPINGS = {
"OllamaOptionExtraBody": "Ollama Option — Extra Body",
"OllamaDebugHistory": "Ollama Debug History",
"OllamaHistoryLength": "Ollama History Length",
"OllamaHeaderBasicAuth": "Ollama Header — Basic Auth",
"OllamaHeaderBearerToken": "Ollama Header — Bearer Token",
"OllamaHeaderCustom": "Ollama Header — Custom",
}
WEB_DIRECTORY = "../js"
+20 -10
View File
@@ -18,7 +18,7 @@ import os
import random
import re
import sys
from typing import Any, Dict, List, Tuple
from typing import Any, Dict, List
from aiohttp import web
from jinja2 import exceptions, sandbox
@@ -68,6 +68,7 @@ class FormatString:
RETURN_TYPES = ("STRING", "STRING")
RETURN_NAMES = ("formatted_string", "saved_file_path")
OUTPUT_IS_LIST = (False, False)
OUTPUT_NODE = True
# Store configurations for each node instance
node_configs: Dict[str, Dict[str, Any]] = {}
@@ -294,12 +295,14 @@ class FormatString:
save_path: str,
unique_id: str = "",
**kwargs,
) -> Tuple[str, ...]:
) -> Dict[str, Any]:
"""
Format a string using the specified template type and variables.
This is the main method executed by the node. It formats the template using either
Python's str.format() or Jinja2 templating, and optionally saves the state to disk.
The node is OUTPUT_NODE=True so the formatted string also renders as a read-only
text display directly on the node (no separate Show Text node needed).
Args:
template_type (str): Either "Simple" or "Jinja2" to specify the template engine.
@@ -309,7 +312,8 @@ class FormatString:
**kwargs: Variable keyword arguments that provide values for template variables.
Returns:
Tuple[str, ...]: A tuple containing the formatted string, the save path,
Dict[str, Any]: ``{"ui": {"text": [formatted_string]}, "result": result}`` where
``result`` is a tuple containing the formatted string, the save path,
followed by the values of input variables (in order).
Example:
@@ -317,7 +321,7 @@ class FormatString:
from format_string import FormatString
# Simple template example
result = FormatString.format_string(
ret = FormatString.format_string(
template_type="Simple",
template="Hello {name}, you are {age} years old",
save_path="",
@@ -325,40 +329,43 @@ class FormatString:
name="Alice",
age="30"
)
print(result) # Outputs: ('Hello Alice, you are 30 years old', '', 'Alice', '30')
print(ret["result"]) # Outputs: ('Hello Alice, you are 30 years old', '', 'Alice', '30')
# Jinja2 template example
result = FormatString.format_string(
ret = FormatString.format_string(
template_type="Jinja2",
template="Hello {{ name }}, today is {{ datetime.now().strftime('%A') }}",
save_path="",
unique_id="124",
name="Bob"
)
print(result) # Outputs: ('Hello Bob, today is Wednesday', '', 'Bob')
print(ret["result"]) # Outputs: ('Hello Bob, today is Wednesday', '', 'Bob')
```
<!-- Example Test:
>>> # Test simple format
>>> result = FormatString.format_string(
>>> ret = FormatString.format_string(
... template_type="Simple",
... template="Hello {name}, you are {age} years old",
... save_path="",
... name="Alice",
... age="30"
... )
>>> result = ret["result"]
>>> assert result[0] == "Hello Alice, you are 30 years old"
>>> assert result[1] == ""
>>> assert result[2] == "Alice"
>>> assert result[3] == "30"
>>> assert ret["ui"]["text"][0] == "Hello Alice, you are 30 years old"
>>>
>>> # Test Jinja2 format with datetime (can't test exact output due to time dependency)
>>> result = FormatString.format_string(
>>> ret = FormatString.format_string(
... template_type="Jinja2",
... template="Name: {{ name }}",
... save_path="",
... name="Bob"
... )
>>> result = ret["result"]
>>> assert result[0] == "Name: Bob"
>>> assert result[1] == ""
>>> assert result[2] == "Bob"
@@ -452,7 +459,10 @@ class FormatString:
logger.debug(
"Full result tuple length: %d, expected: %d", len(result), len(keys) + 2
)
return result
return {
"ui": {"text": [formatted_string]},
"result": result,
}
@classmethod
def update_widget(
+217 -15
View File
@@ -1,11 +1,12 @@
"""
Ollama model integration nodes for ComfyUI.
14 nodes: client configuration, model discovery, load/unload, chat completion,
composable inference options, and history utilities.
17 nodes: client configuration, auth headers, model discovery, load/unload,
chat completion, composable inference options, and history utilities.
ADR-004: aiohttp (already in ComfyUI dep tree) for all Ollama HTTP calls.
ADR-005: OllamaClient node is the single source of the host URL.
ADR-005: OllamaClient node is the single source of the host URL (and, per
US7, of any auth headers — every downstream node reaches the same server).
"""
import asyncio
@@ -13,17 +14,100 @@ import json
import logging
import os
import sys
import threading
import time
logger = logging.getLogger(__name__)
# ---------------------------------------------------------------------------
# Custom socket type
# Local response cache
# ---------------------------------------------------------------------------
#
# OllamaChatCompletion is OUTPUT_NODE=True (needed for inline display), which
# means ComfyUI re-executes it on every queue run even when none of its
# inputs changed — unlike normal nodes, it isn't skipped by ComfyUI's own
# input-hash cache. This cache absorbs those redundant round-trips: identical
# (client, headers, model, messages, options) reuse the prior response
# instead of re-querying Ollama. Model discovery gets the same treatment
# (many nodes independently query /api/tags on graph load) but with a short
# TTL so a newly-pulled model still surfaces after a refresh.
class _TTLLRUCache:
"""Bounded cache, LRU-evicted, with an optional per-entry TTL."""
def __init__(self, maxsize: int, ttl_seconds: float | None = None):
self.maxsize = maxsize
self.ttl_seconds = ttl_seconds
self._data: dict = {}
self._lock = threading.Lock()
def get(self, key):
with self._lock:
entry = self._data.get(key)
if entry is None:
return None, False
expires_at, value = entry
if self.ttl_seconds is not None and time.monotonic() > expires_at:
del self._data[key]
return None, False
# Re-insert to mark as most-recently-used (dicts preserve insertion order).
del self._data[key]
self._data[key] = (expires_at, value)
return value, True
def set(self, key, value):
with self._lock:
expires_at = (
time.monotonic() + self.ttl_seconds
if self.ttl_seconds is not None
else float("inf")
)
self._data.pop(key, None)
self._data[key] = (expires_at, value)
while len(self._data) > self.maxsize:
oldest_key = next(iter(self._data))
del self._data[oldest_key]
def clear(self):
with self._lock:
self._data.clear()
def _cache_key(*parts) -> str:
"""Deterministic, hashable key from arbitrary JSON-serializable parts."""
return json.dumps(parts, sort_keys=True, default=str)
_MODEL_LIST_CACHE = _TTLLRUCache(maxsize=32, ttl_seconds=20.0)
_CHAT_RESPONSE_CACHE = _TTLLRUCache(maxsize=64, ttl_seconds=None)
# ---------------------------------------------------------------------------
# Custom socket types
# ---------------------------------------------------------------------------
class OllamaClientType(str):
"""Typed string carrying the Ollama host URL through the node graph."""
"""Typed string carrying the Ollama host URL through the node graph.
Also carries an optional ``.headers`` dict (US7 — basic auth / bearer
tokens for Ollama servers behind a reverse proxy). Since this is a plain
``str`` subclass, every existing ``f"{client}/api/..."`` call site keeps
working unchanged; only code that wants auth reads ``.headers``.
"""
def __new__(cls, host, headers=None):
obj = super().__new__(cls, host)
obj.headers = dict(headers) if headers else {}
return obj
def _client_headers(client) -> dict | None:
"""Extract auth headers stashed on an OllamaClientType, if any."""
headers = getattr(client, "headers", None)
return dict(headers) if headers else None
# ---------------------------------------------------------------------------
@@ -45,24 +129,45 @@ def _run_async(coro):
return asyncio.run(coro)
async def _fetch_models(host: str) -> list[str]:
"""GET {host}/api/tags — return list of model name strings."""
async def _fetch_models(host: str, headers: dict | None = None) -> list[str]:
"""GET {host}/api/tags — return list of model name strings.
Cached for _MODEL_LIST_CACHE.ttl_seconds per (host, headers) pair — node
creation, the Refresh button, and startup all otherwise re-issue this
same request in quick succession.
"""
cache_key = _cache_key("models", host, headers or {})
cached, hit = _MODEL_LIST_CACHE.get(cache_key)
if hit:
return cached
import aiohttp
try:
async with aiohttp.ClientSession() as session:
async with session.get(
f"{host}/api/tags",
headers=headers or None,
timeout=aiohttp.ClientTimeout(total=5),
) as resp:
data = await resp.json()
return [m["name"] for m in data.get("models", [])]
models = [m["name"] for m in data.get("models", [])]
except Exception as exc:
logger.warning("Could not fetch Ollama models from %s: %s", host, exc)
return []
if models:
_MODEL_LIST_CACHE.set(cache_key, models)
return models
async def _post_json(url: str, payload: dict, *, timeout: float = 120.0) -> dict:
async def _post_json(
url: str,
payload: dict,
*,
timeout: float = 120.0,
headers: dict | None = None,
) -> dict:
"""POST JSON to url, return parsed response dict."""
import aiohttp
@@ -71,6 +176,7 @@ async def _post_json(url: str, payload: dict, *, timeout: float = 120.0) -> dict
async with session.post(
url,
json=payload,
headers=headers or None,
timeout=aiohttp.ClientTimeout(total=timeout),
) as resp:
if resp.status >= 400:
@@ -149,7 +255,10 @@ class OllamaClient:
return {
"required": {
"host": ("STRING", {"default": "http://localhost:11434"}),
}
},
"optional": {
"headers": ("OLLAMA_HEADERS",),
},
}
RETURN_TYPES = ("OLLAMA_CLIENT",)
@@ -157,8 +266,83 @@ class OllamaClient:
FUNCTION = "create_client"
CATEGORY = "dv/ollama"
def create_client(self, host: str):
return (OllamaClientType(host),)
def create_client(self, host: str, headers: dict | None = None):
return (OllamaClientType(host, headers),)
# ---------------------------------------------------------------------------
# US7 — Composable auth headers
# ---------------------------------------------------------------------------
def _merge_header(headers, name, value):
result = dict(headers) if headers else {}
result[name] = value
return result
class OllamaHeaderBasicAuth:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"username": ("STRING", {"default": ""}),
"password": ("STRING", {"default": ""}),
},
"optional": {"headers": ("OLLAMA_HEADERS",)},
}
RETURN_TYPES = ("OLLAMA_HEADERS",)
RETURN_NAMES = ("headers",)
FUNCTION = "set_basic_auth"
CATEGORY = "dv/ollama/headers"
def set_basic_auth(self, username, password, headers=None):
import base64
token = base64.b64encode(f"{username}:{password}".encode()).decode()
return (_merge_header(headers, "Authorization", f"Basic {token}"),)
class OllamaHeaderBearerToken:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"token": ("STRING", {"default": ""}),
},
"optional": {"headers": ("OLLAMA_HEADERS",)},
}
RETURN_TYPES = ("OLLAMA_HEADERS",)
RETURN_NAMES = ("headers",)
FUNCTION = "set_bearer_token"
CATEGORY = "dv/ollama/headers"
def set_bearer_token(self, token, headers=None):
return (_merge_header(headers, "Authorization", f"Bearer {token}"),)
class OllamaHeaderCustom:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"name": ("STRING", {"default": ""}),
"value": ("STRING", {"default": ""}),
},
"optional": {"headers": ("OLLAMA_HEADERS",)},
}
RETURN_TYPES = ("OLLAMA_HEADERS",)
RETURN_NAMES = ("headers",)
FUNCTION = "set_custom_header"
CATEGORY = "dv/ollama/headers"
def set_custom_header(self, name, value, headers=None):
if not name.strip():
raise ValueError("header name cannot be empty")
return (_merge_header(headers, name, value),)
# ---------------------------------------------------------------------------
@@ -213,6 +397,7 @@ class OllamaLoadModel:
f"{client}/api/generate",
{"model": model, "keep_alive": -1, "stream": False},
timeout=300.0,
headers=_client_headers(client),
)
)
return (model,)
@@ -252,6 +437,7 @@ class OllamaUnloadModel:
f"{client}/api/generate",
{"model": model, "keep_alive": 0, "stream": False},
timeout=30.0,
headers=_client_headers(client),
)
)
return (model, passthrough)
@@ -326,10 +512,26 @@ class OllamaChatCompletion:
}
if options:
payload["options"] = options
result = _run_async(
_post_json(f"{client}/api/chat", payload, timeout=float(timeout_secs))
headers = _client_headers(client)
cache_key = _cache_key(
"chat", client, headers or {}, effective_model, messages, options or {}
)
response_text = result.get("message", {}).get("content", "")
cached, hit = _CHAT_RESPONSE_CACHE.get(cache_key)
if hit:
response_text = cached
else:
result = _run_async(
_post_json(
f"{client}/api/chat",
payload,
timeout=float(timeout_secs),
headers=headers,
)
)
response_text = result.get("message", {}).get("content", "")
_CHAT_RESPONSE_CACHE.set(cache_key, response_text)
updated = list(history)
updated.append({"role": "user", "content": prompt})
updated.append({"role": "assistant", "content": response_text})
+18
View File
@@ -66,6 +66,24 @@ def pytest_configure(config):
# ---------------------------------------------------------------------------
@pytest.fixture(autouse=True)
def _clear_ollama_caches():
"""Reset comfydv.ollama's module-level LRU caches around every test.
Several tests reuse identical client/model/prompt inputs across cases
with different monkeypatched responses — without this, a later test would
silently get an earlier test's cached result instead of exercising its
own fake.
"""
from comfydv.ollama import _CHAT_RESPONSE_CACHE, _MODEL_LIST_CACHE
_MODEL_LIST_CACHE.clear()
_CHAT_RESPONSE_CACHE.clear()
yield
_MODEL_LIST_CACHE.clear()
_CHAT_RESPONSE_CACHE.clear()
@pytest.fixture(scope="session")
def ollama_host():
return "http://localhost:11434"
+73 -16
View File
@@ -86,7 +86,7 @@ class TestSimpleFormatting:
save_path="",
unique_id="test1",
name=sample_data["name"],
)
)["result"]
assert len(result) == 3 # formatted_string, saved_file_path, name
assert result[0] == "Hello Alice"
assert result[1] == ""
@@ -101,7 +101,7 @@ class TestSimpleFormatting:
unique_id="test2",
name=sample_data["name"],
age=sample_data["age"],
)
)["result"]
assert len(result) == 4 # formatted_string, saved_file_path, name, age
assert result[0] == "Hello Alice, you are 30"
assert result[1] == ""
@@ -115,7 +115,7 @@ class TestSimpleFormatting:
template="Hello World",
save_path="",
unique_id="test3",
)
)["result"]
assert len(result) == 2 # formatted_string, saved_file_path
assert result[0] == "Hello World"
assert result[1] == ""
@@ -142,7 +142,7 @@ class TestJinja2Formatting:
save_path="",
unique_id="test5",
name=sample_data["name"],
)
)["result"]
assert len(result) == 3 # formatted_string, saved_file_path, name
assert result[0] == "Hello Alice"
assert result[1] == ""
@@ -156,7 +156,7 @@ class TestJinja2Formatting:
save_path="",
unique_id="test6",
name=sample_data["name"],
)
)["result"]
assert len(result) == 3
assert result[0] == "Hello ALICE"
assert result[1] == ""
@@ -171,7 +171,7 @@ class TestJinja2Formatting:
unique_id="test7",
first=sample_data["first"],
last=sample_data["last"],
)
)["result"]
assert len(result) == 4 # formatted_string, saved_file_path, first, last
assert result[0] == "JOHN doe"
assert result[1] == ""
@@ -185,7 +185,7 @@ class TestJinja2Formatting:
template="Time: {{ now() }}",
save_path="",
unique_id="test8",
)
)["result"]
assert len(result) == 2 # formatted_string, saved_file_path (no extracted vars)
assert result[0].startswith("Time: ")
assert result[1] == ""
@@ -198,13 +198,68 @@ class TestJinja2Formatting:
save_path="",
unique_id="test9",
value=sample_data["value"],
)
)["result"]
# value is not extracted as a variable because it's used in an expression
assert len(result) == 2 # Just formatted_string, saved_file_path
assert result[0] == "Result: 10"
assert result[1] == ""
class TestInlineDisplay:
"""FormatString is OUTPUT_NODE=True so formatted_string renders on the node itself."""
def test_is_output_node(self, format_string_class):
"""FormatString must have OUTPUT_NODE=True for inline display."""
assert getattr(format_string_class, "OUTPUT_NODE", False) is True
def test_returns_ui_result_dict(self, format_string_class, sample_data):
"""format_string() must return {'ui': ..., 'result': ...} not a bare tuple."""
ret = format_string_class.format_string(
template_type="Simple",
template="Hello {name}",
save_path="",
unique_id="test-inline-1",
name=sample_data["name"],
)
assert isinstance(ret, dict), f"Expected dict, got {type(ret)}"
assert "ui" in ret, "Missing 'ui' key"
assert "result" in ret, "Missing 'result' key"
def test_ui_contains_formatted_string(self, format_string_class, sample_data):
"""Formatted string must appear in ui['text'] for the inline display."""
ret = format_string_class.format_string(
template_type="Simple",
template="Hello {name}",
save_path="",
unique_id="test-inline-2",
name=sample_data["name"],
)
assert ret["ui"]["text"][0] == "Hello Alice"
def test_ui_reflects_jinja2_output(self, format_string_class, sample_data):
"""Jinja2-rendered output must also surface in the inline display."""
ret = format_string_class.format_string(
template_type="Jinja2",
template="Hello {{ name | upper }}",
save_path="",
unique_id="test-inline-3",
name=sample_data["name"],
)
assert ret["ui"]["text"][0] == "Hello ALICE"
def test_result_is_unchanged_tuple(self, format_string_class, sample_data):
"""result key must still be the same output tuple as before OUTPUT_NODE was added."""
ret = format_string_class.format_string(
template_type="Simple",
template="Hello {name}, you are {age}",
save_path="",
unique_id="test-inline-4",
name=sample_data["name"],
age=sample_data["age"],
)
assert ret["result"] == ("Hello Alice, you are 30", "", "Alice", "30")
class TestDynamicOutputs:
"""Test dynamic output configuration."""
@@ -296,7 +351,7 @@ class TestOutputConsistency:
format_string_class.update_widget("node1", "Simple", "Hello {name}")
result = format_string_class.format_string(
"Simple", "Hello {name}", "", "node1", name=sample_data["name"]
)
)["result"]
assert len(result) == len(format_string_class.RETURN_TYPES)
assert len(result) == len(format_string_class.RETURN_NAMES)
@@ -311,7 +366,7 @@ class TestOutputConsistency:
"node2",
name=sample_data["name"],
age=sample_data["age"],
)
)["result"]
assert len(result) == len(format_string_class.RETURN_TYPES)
assert len(result) == len(format_string_class.RETURN_NAMES)
@@ -319,7 +374,9 @@ class TestOutputConsistency:
def test_output_consistency_no_vars(self, format_string_class):
"""Test output consistency with no variables."""
format_string_class.update_widget("node3", "Simple", "Hello World")
result = format_string_class.format_string("Simple", "Hello World", "", "node3")
result = format_string_class.format_string(
"Simple", "Hello World", "", "node3"
)["result"]
assert len(result) == len(format_string_class.RETURN_TYPES)
assert len(result) == len(format_string_class.RETURN_NAMES)
@@ -408,7 +465,7 @@ class TestStatePersistence:
save_path="",
unique_id="test",
name=sample_data["name"],
)
)["result"]
# Should complete without error
assert result[1] == "" # saved_file_path should be empty (position 1)
@@ -425,7 +482,7 @@ class TestEdgeCases:
"""Test with empty template."""
result = format_string_class.format_string(
template_type="Simple", template="", save_path="", unique_id="test"
)
)["result"]
assert len(result) == 2
assert result[0] == ""
assert result[1] == ""
@@ -437,7 +494,7 @@ class TestEdgeCases:
template="{{ unclosed",
save_path="",
unique_id="test",
)
)["result"]
# Should return error message in formatted_string
assert len(result) == 2
assert "Error in Jinja2 template" in result[0]
@@ -450,7 +507,7 @@ class TestEdgeCases:
save_path="",
unique_id="test",
name="<Alice & Bob>",
)
)["result"]
assert "<Alice & Bob>" in result[0] # formatted_string is at position 0
def test_unicode_in_template(self, format_string_class):
@@ -461,7 +518,7 @@ class TestEdgeCases:
save_path="",
unique_id="test",
name="世界",
)
)["result"]
assert "你好 世界 🎉" in result[0] # formatted_string is at position 0
+370 -11
View File
@@ -25,6 +25,9 @@ from comfydv.ollama import (
OllamaClient,
OllamaClientType,
OllamaDebugHistory,
OllamaHeaderBasicAuth,
OllamaHeaderBearerToken,
OllamaHeaderCustom,
OllamaHistoryLength,
OllamaLoadModel,
OllamaModelSelector,
@@ -36,6 +39,7 @@ from comfydv.ollama import (
OllamaOptionTopK,
OllamaOptionTopP,
OllamaUnloadModel,
_MODEL_LIST_CACHE,
_fetch_models,
_post_json,
_run_async,
@@ -171,7 +175,7 @@ class TestInfrastructure:
return False
class FakeSession:
def post(self, url, *, json=None, timeout=None):
def post(self, url, *, json=None, timeout=None, headers=None):
captured["timeout"] = timeout
return FakeResponse()
@@ -287,7 +291,7 @@ class TestUS3ModelLifecycle:
"""Issue 4: load_model should POST to /api/generate, not /api/show."""
captured = {}
async def fake_post(url, payload, *, timeout=120.0):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
captured["url"] = url
return {}
@@ -307,7 +311,7 @@ class TestUS3ModelLifecycle:
"""Issue 4: unload_model should POST to /api/generate, not /api/show."""
captured = {}
async def fake_post(url, payload, *, timeout=120.0):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
captured["url"] = url
return {}
@@ -329,7 +333,7 @@ class TestUS3ModelLifecycle:
"""Issue 5: load_model must send keep_alive as integer -1, not string '-1'."""
captured = {}
async def fake_post(url, payload, *, timeout=120.0):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
captured["payload"] = payload
return {}
@@ -353,7 +357,7 @@ class TestUS3ModelLifecycle:
"""Issue 5: unload_model must send keep_alive as integer 0."""
captured = {}
async def fake_post(url, payload, *, timeout=120.0):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
captured["payload"] = payload
return {}
@@ -444,7 +448,7 @@ class TestUS4ChatCompletion:
"""Wiring OllamaLoadModel.model_name → OllamaChatCompletion.model works."""
captured = {}
async def fake_post(url, payload, *, timeout=120.0):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
captured["model"] = payload.get("model")
return {"message": {"content": "ok"}}
@@ -480,7 +484,7 @@ class TestUS4ChatCompletion:
def test_chat_returns_ui_result_dict(self, monkeypatch):
"""chat() must return {'ui': ..., 'result': ...} not a bare tuple."""
async def fake_post(url, payload, *, timeout=120.0):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
return {"message": {"content": "hello"}}
import comfydv.ollama as ollama_mod
@@ -495,7 +499,7 @@ class TestUS4ChatCompletion:
def test_chat_ui_contains_response_text(self, monkeypatch):
"""Response text must appear in ui['text'] for the inline display."""
async def fake_post(url, payload, *, timeout=120.0):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
return {"message": {"content": "hello world"}}
import comfydv.ollama as ollama_mod
@@ -508,7 +512,7 @@ class TestUS4ChatCompletion:
def test_chat_result_is_3_tuple(self, monkeypatch):
"""result key must be the 3-tuple (response, history, model_name)."""
async def fake_post(url, payload, *, timeout=120.0):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
return {"message": {"content": "hello"}}
import comfydv.ollama as ollama_mod
@@ -548,7 +552,7 @@ class TestUS4ChatCompletion:
"""
captured = {}
async def fake_post(url, payload, *, timeout=120.0):
async def fake_post(url, payload, *, timeout=120.0, headers=None):
captured["timeout"] = timeout
return {"message": {"content": "ok"}}
@@ -739,13 +743,365 @@ class TestUS6HistoryInspection:
assert length == 0
# ---------------------------------------------------------------------------
# US7 — Auth headers
# ---------------------------------------------------------------------------
class TestUS7AuthHeaders:
def test_client_without_headers_has_empty_dict(self):
(client,) = OllamaClient().create_client("http://localhost:11434")
assert client.headers == {}
def test_client_carries_headers(self):
(client,) = OllamaClient().create_client(
"http://localhost:11434", headers={"Authorization": "Bearer abc"}
)
assert client.headers == {"Authorization": "Bearer abc"}
assert client == "http://localhost:11434", (
"OllamaClientType must still compare equal to the plain host string"
)
def test_basic_auth_sets_authorization_header(self):
(headers,) = OllamaHeaderBasicAuth().set_basic_auth(
username="alice", password="hunter2"
)
assert headers["Authorization"].startswith("Basic ")
def test_basic_auth_encodes_username_password(self):
import base64
(headers,) = OllamaHeaderBasicAuth().set_basic_auth(
username="alice", password="hunter2"
)
token = headers["Authorization"].removeprefix("Basic ")
assert base64.b64decode(token).decode() == "alice:hunter2"
def test_bearer_token_sets_authorization_header(self):
(headers,) = OllamaHeaderBearerToken().set_bearer_token(token="sk-12345")
assert headers == {"Authorization": "Bearer sk-12345"}
def test_custom_header_sets_arbitrary_name(self):
(headers,) = OllamaHeaderCustom().set_custom_header(
name="X-Api-Key", value="abc123"
)
assert headers == {"X-Api-Key": "abc123"}
def test_custom_header_empty_name_raises(self):
with pytest.raises(ValueError, match="cannot be empty"):
OllamaHeaderCustom().set_custom_header(name="", value="abc123")
def test_headers_chain_and_merge(self):
"""Scenario: Bearer token + a custom header both reach the request."""
(h1,) = OllamaHeaderBearerToken().set_bearer_token(token="sk-12345")
(h2,) = OllamaHeaderCustom().set_custom_header(
name="X-Api-Key", value="abc123", headers=h1
)
assert h2 == {"Authorization": "Bearer sk-12345", "X-Api-Key": "abc123"}
def test_load_model_forwards_client_headers(self, monkeypatch):
captured = {}
async def fake_post(url, payload, *, timeout=120.0, headers=None):
captured["headers"] = headers
return {}
import comfydv.ollama as ollama_mod
monkeypatch.setattr(ollama_mod, "_post_json", fake_post)
(client,) = OllamaClient().create_client(
"http://localhost:11434", headers={"Authorization": "Bearer sk-1"}
)
OllamaLoadModel().load_model(client=client, model="test-model")
assert captured["headers"] == {"Authorization": "Bearer sk-1"}
def test_unload_model_forwards_client_headers(self, monkeypatch):
captured = {}
async def fake_post(url, payload, *, timeout=120.0, headers=None):
captured["headers"] = headers
return {}
import comfydv.ollama as ollama_mod
monkeypatch.setattr(ollama_mod, "_post_json", fake_post)
(client,) = OllamaClient().create_client(
"http://localhost:11434", headers={"Authorization": "Bearer sk-1"}
)
OllamaUnloadModel().unload_model(client=client, model="test-model")
assert captured["headers"] == {"Authorization": "Bearer sk-1"}
def test_chat_forwards_client_headers(self, monkeypatch):
captured = {}
async def fake_post(url, payload, *, timeout=120.0, headers=None):
captured["headers"] = headers
return {"message": {"content": "ok"}}
import comfydv.ollama as ollama_mod
monkeypatch.setattr(ollama_mod, "_post_json", fake_post)
(client,) = OllamaClient().create_client(
"http://localhost:11434", headers={"Authorization": "Bearer sk-1"}
)
OllamaChatCompletion().chat(client=client, model="m", prompt="hi")
assert captured["headers"] == {"Authorization": "Bearer sk-1"}
def test_plain_string_client_has_no_headers(self, monkeypatch):
"""Backward compat: a bare string client (no OllamaClient node) sends no headers."""
captured = {}
async def fake_post(url, payload, *, timeout=120.0, headers=None):
captured["headers"] = headers
return {"message": {"content": "ok"}}
import comfydv.ollama as ollama_mod
monkeypatch.setattr(ollama_mod, "_post_json", fake_post)
OllamaChatCompletion().chat(
client="http://localhost:11434", model="m", prompt="hi"
)
assert captured["headers"] is None
# ---------------------------------------------------------------------------
# Response cache — model discovery + chat completion
# ---------------------------------------------------------------------------
class TestResponseCache:
def test_fetch_models_second_call_is_cached(self, monkeypatch):
calls = {"n": 0}
class FakeResponse:
status = 200
async def json(self):
calls["n"] += 1
return {"models": [{"name": "llama3:latest"}]}
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
class FakeSession:
def get(self, url, *, headers=None, timeout=None):
return FakeResponse()
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
import aiohttp
monkeypatch.setattr(aiohttp, "ClientSession", lambda: FakeSession())
models1 = _run_async(_fetch_models("http://localhost:11434"))
models2 = _run_async(_fetch_models("http://localhost:11434"))
assert models1 == models2 == ["llama3:latest"]
assert calls["n"] == 1, "second call with identical inputs must hit the cache"
def test_fetch_models_different_host_not_cached(self, monkeypatch):
calls = {"n": 0}
class FakeResponse:
status = 200
async def json(self):
calls["n"] += 1
return {"models": [{"name": "llama3:latest"}]}
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
class FakeSession:
def get(self, url, *, headers=None, timeout=None):
return FakeResponse()
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
import aiohttp
monkeypatch.setattr(aiohttp, "ClientSession", lambda: FakeSession())
_run_async(_fetch_models("http://host-a:11434"))
_run_async(_fetch_models("http://host-b:11434"))
assert calls["n"] == 2, "different hosts must not share a cache entry"
def test_fetch_models_ttl_expiry_refetches(self, monkeypatch):
"""A newly-installed model must surface once the TTL lapses."""
calls = {"n": 0}
class FakeResponse:
status = 200
async def json(self):
calls["n"] += 1
return {"models": [{"name": "llama3:latest"}]}
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
class FakeSession:
def get(self, url, *, headers=None, timeout=None):
return FakeResponse()
async def __aenter__(self):
return self
async def __aexit__(self, *args):
return False
import aiohttp
monkeypatch.setattr(aiohttp, "ClientSession", lambda: FakeSession())
_run_async(_fetch_models("http://localhost:11434"))
assert calls["n"] == 1
# Simulate TTL lapsing by clearing the cache directly rather than
# sleeping in a unit test.
_MODEL_LIST_CACHE.clear()
_run_async(_fetch_models("http://localhost:11434"))
assert calls["n"] == 2
def test_chat_second_identical_call_is_cached(self, monkeypatch):
calls = {"n": 0}
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls["n"] += 1
return {"message": {"content": f"response #{calls['n']}"}}
import comfydv.ollama as ollama_mod
monkeypatch.setattr(ollama_mod, "_post_json", fake_post)
r1, _, _ = OllamaChatCompletion().chat(
client="http://x", model="m", prompt="hi"
)["result"]
r2, _, _ = OllamaChatCompletion().chat(
client="http://x", model="m", prompt="hi"
)["result"]
assert r1 == r2 == "response #1"
assert calls["n"] == 1, "identical chat inputs must reuse the cached response"
def test_chat_different_prompt_not_cached(self, monkeypatch):
calls = {"n": 0}
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls["n"] += 1
return {"message": {"content": f"response #{calls['n']}"}}
import comfydv.ollama as ollama_mod
monkeypatch.setattr(ollama_mod, "_post_json", fake_post)
OllamaChatCompletion().chat(client="http://x", model="m", prompt="hi")
OllamaChatCompletion().chat(client="http://x", model="m", prompt="bye")
assert calls["n"] == 2, "different prompts must not share a cache entry"
def test_chat_different_options_not_cached(self, monkeypatch):
calls = {"n": 0}
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls["n"] += 1
return {"message": {"content": f"response #{calls['n']}"}}
import comfydv.ollama as ollama_mod
monkeypatch.setattr(ollama_mod, "_post_json", fake_post)
OllamaChatCompletion().chat(
client="http://x", model="m", prompt="hi", options={"temperature": 0.0}
)
OllamaChatCompletion().chat(
client="http://x", model="m", prompt="hi", options={"temperature": 0.9}
)
assert calls["n"] == 2, "different options must not share a cache entry"
def test_chat_grown_history_not_cached(self, monkeypatch):
"""A follow-up turn (different message history) must not reuse turn 1's cache entry."""
calls = {"n": 0}
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls["n"] += 1
return {"message": {"content": f"response #{calls['n']}"}}
import comfydv.ollama as ollama_mod
monkeypatch.setattr(ollama_mod, "_post_json", fake_post)
_, h1, _ = OllamaChatCompletion().chat(
client="http://x", model="m", prompt="turn 1", history=[]
)["result"]
OllamaChatCompletion().chat(
client="http://x", model="m", prompt="turn 2", history=h1
)
assert calls["n"] == 2
def test_chat_different_client_headers_not_cached(self, monkeypatch):
calls = {"n": 0}
async def fake_post(url, payload, *, timeout=120.0, headers=None):
calls["n"] += 1
return {"message": {"content": f"response #{calls['n']}"}}
import comfydv.ollama as ollama_mod
monkeypatch.setattr(ollama_mod, "_post_json", fake_post)
(client_a,) = OllamaClient().create_client(
"http://x", headers={"Authorization": "Bearer a"}
)
(client_b,) = OllamaClient().create_client(
"http://x", headers={"Authorization": "Bearer b"}
)
OllamaChatCompletion().chat(client=client_a, model="m", prompt="hi")
OllamaChatCompletion().chat(client=client_b, model="m", prompt="hi")
assert calls["n"] == 2, (
"requests to the same host with different auth headers must not share "
"a cache entry"
)
# ---------------------------------------------------------------------------
# Node contract sanity checks (ComfyUI registration requirements)
# ---------------------------------------------------------------------------
class TestNodeContracts:
"""All 14 nodes must satisfy ComfyUI's node registration contract."""
"""All 17 nodes must satisfy ComfyUI's node registration contract."""
NODE_CLASSES = [
OllamaClient,
@@ -762,6 +1118,9 @@ class TestNodeContracts:
OllamaOptionExtraBody,
OllamaDebugHistory,
OllamaHistoryLength,
OllamaHeaderBasicAuth,
OllamaHeaderBearerToken,
OllamaHeaderCustom,
]
@pytest.mark.parametrize("node_cls", NODE_CLASSES, ids=lambda c: c.__name__)