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