fix: address PR #9 review issues

- fix(scripts): tomllib fallback for Python 3.10 compat
- fix(ci): remove verify-tag-exists from test.yml (publish.yml already guards)
- fix(tests): ensure local nodes/ import in test_node_registration.py
- fix(adapter): shared module-level model cache across OllamaClient instances
- fix(adapter): exact model tag match in check_health() to avoid prefix conflation
- fix(refiner): per-pass seed increment so multi-pass refinement varies
- fix(combiner): validate mode before single-prompt fast path
- feat(generator): strip DeepSeek <think> and markdown <details> blocks
- test: add prompt_refiner seed tests, cache sharing tests, exact match tests
- chore(deps): add dev extras with tomli fallback
This commit is contained in:
limbicnation
2026-04-30 08:05:17 +02:00
parent fc5d02362f
commit ed5b681f06
12 changed files with 190 additions and 47 deletions
+1 -20
View File
@@ -72,23 +72,4 @@ jobs:
run: |
comfy --skip-prompt node validate
verify-tag-exists:
if: github.event_name == 'push' && github.ref == 'refs/heads/main'
runs-on: ubuntu-latest
steps:
- name: Checkout code (with full tag history)
uses: actions/checkout@v4
with:
fetch-depth: 0
- name: Set up Python
uses: actions/setup-python@v5
with:
python-version: '3.11'
- name: Require matching git tag for current pyproject.toml version
run: |
VERSION=$(python -c 'import tomllib, pathlib; print(tomllib.loads(pathlib.Path("pyproject.toml").read_text())["project"]["version"])')
if ! git rev-parse -q --verify "refs/tags/v${VERSION}" >/dev/null; then
echo "::error::pyproject.toml is at v${VERSION} but no matching git tag exists. Tag the release before merging the bump to main."
exit 1
fi
echo "OK: git tag v${VERSION} exists."
+16 -12
View File
@@ -17,6 +17,11 @@ from typing import Any, ClassVar
logger = logging.getLogger(__name__)
# Module-level shared cache (thread-safe across OllamaClient instances)
_MODEL_CACHE: list[str] | None = None
_CACHE_TIME: float = 0.0
_CACHE_LOCK = threading.Lock()
# Optional imports with graceful degradation
try:
import ollama
@@ -50,22 +55,20 @@ class OllamaClient:
def __init__(self, logger_prefix: str = "OllamaClient"):
self.logger_prefix = logger_prefix
# Instance-level cache (was class-level, causing race conditions)
self._cached_models: list[str] | None = None
self._cache_time = 0.0
self._cache_lock = threading.Lock()
def _log(self, message: str) -> None:
logger.info("[%s] %s", self.logger_prefix, message)
def discover_models(self) -> list[str]:
"""
Fetch available Ollama models with instance-level caching.
Fetch available Ollama models with module-level caching.
Prioritizes LoRA-enhanced models.
"""
with self._cache_lock:
if self._cached_models and (time.time() - self._cache_time) < 60:
return self._cached_models
global _MODEL_CACHE, _CACHE_TIME
with _CACHE_LOCK:
if _MODEL_CACHE and (time.time() - _CACHE_TIME) < 60:
return _MODEL_CACHE.copy()
if not OLLAMA_API_AVAILABLE:
return self.DEFAULT_MODELS
@@ -84,9 +87,9 @@ class OllamaClient:
models = sorted(models, key=sort_key)
with self._cache_lock:
self._cached_models = models
self._cache_time = time.time()
with _CACHE_LOCK:
_MODEL_CACHE = models
_CACHE_TIME = time.time()
self._log(f"Found {len(models)} Ollama models")
return models
@@ -127,7 +130,8 @@ class OllamaClient:
try:
ps_response = ollama.ps()
running_models = [m.model for m in ps_response.models]
is_loaded = any(model == rm or model.startswith(rm.split(":")[0]) for rm in running_models)
# Exact match only; do not conflate different tags (e.g. qwen3:8b vs qwen3:4b)
is_loaded = model in running_models
if is_loaded:
return (True, f"Model '{model}' is loaded in VRAM", True)
else:
+3 -3
View File
@@ -167,15 +167,15 @@ class PromptCombinerNode:
if not prompts:
return ("[PromptCombiner] At least one prompt is required.",)
if len(prompts) == 1:
return (prompts[0][0],)
try:
selected = CombineMode(mode)
except ValueError:
valid = ", ".join(m.value for m in CombineMode)
return (f"[PromptCombiner] Unknown mode {mode!r}. Valid: {valid}",)
if len(prompts) == 1:
return (prompts[0][0],)
match selected:
case CombineMode.BLEND:
return (self._blend(prompts),)
+5 -1
View File
@@ -52,8 +52,12 @@ def extract_final_prompt(text: str) -> str:
if not text:
return text
# Remove Qwen3 thinking blocks: "Thinking...\n...\n...done thinking.\n"
# DeepSeek / generic XML thinking blocks
text = re.sub(r"<think[\s\S]*?</think>", "", text, flags=re.DOTALL)
# Qwen3 thinking blocks: "Thinking...\n...\n...done thinking.\n"
text = re.sub(r"Thinking\.\.\.[\s\S]*?\.\.\.done thinking\.[\s]*", "", text, flags=re.DOTALL)
# Markdown details blocks
text = re.sub(r"<details>\s*<summary>.*?</summary>[\s\S]*?</details>", "", text, flags=re.DOTALL)
# Remove common prefixes like "**Prompt:**" or "**Stable Diffusion Prompt:**"
text = re.sub(r"\*\*(?:Stable Diffusion )?Prompt:\*\*\s*", "", text)
+4 -1
View File
@@ -159,6 +159,9 @@ Refined prompt:"""
# Build refinement prompt
refinement = self.REFINEMENT_PROMPT.format(prompt=current_prompt)
# Derive per-pass seed so multi-pass refinement isn't a no-op
pass_seed = None if effective_seed is None else effective_seed + i
# Generate refined version
output = client.generate_streaming(
model=model,
@@ -167,7 +170,7 @@ Refined prompt:"""
top_p=top_p,
timeout=timeout,
pbar=pbar,
seed=effective_seed,
seed=pass_seed,
)
if output is None:
+3
View File
@@ -16,6 +16,9 @@ PublisherId = "limbicnation"
DisplayName = "Prompt Generator"
Icon = "https://raw.githubusercontent.com/Limbicnation/ComfyUI-PromptGenerator/main/icon.png"
[project.optional-dependencies]
dev = ["tomli>=1.1.0 ; python_version < '3.11'", "pytest", "pytest-cov", "ruff"]
[tool.comfy.python]
requires-python = ">=3.10"
+4 -1
View File
@@ -11,7 +11,10 @@ import subprocess
import sys
from pathlib import Path
import tomllib
try:
import tomllib
except ModuleNotFoundError: # pragma: no cover
import tomli as tomllib # type: ignore[no-redef]
def current_version() -> str:
+25
View File
@@ -71,3 +71,28 @@ class TestExtractFinalPrompt:
"None"
)
assert extract_final_prompt(text) == "A mystical forest at twilight, dramatic lighting, 8k"
def test_deepseek_think_block(self):
"""DeepSeek / generic XML <think> blocks should be stripped."""
text = "<think>Some reasoning here</think>Final prompt here"
assert extract_final_prompt(text) == "Final prompt here"
def test_deepseek_think_multiline(self):
"""Multiline <think> blocks should be stripped."""
text = "<think>\nReasoning line 1\nReasoning line 2\n</think>\nFinal prompt"
assert extract_final_prompt(text) == "Final prompt"
def test_markdown_details_block(self):
"""Markdown <details> blocks should be stripped."""
text = "<details><summary>Reasoning</summary>Hidden reasoning</details>Final prompt"
assert extract_final_prompt(text) == "Final prompt"
def test_mixed_reasoning_formats(self):
"""Multiple reasoning formats in one string should all be removed."""
text = (
"<think>DeepSeek reasoning</think>\n"
"Thinking...\nQwen reasoning\n...done thinking.\n"
"<details><summary>Summary</summary>Details</details>\n"
"Final clean prompt"
)
assert extract_final_prompt(text) == "Final clean prompt"
+10 -1
View File
@@ -9,13 +9,22 @@ from __future__ import annotations
import importlib.util
import pkgutil
import sys
from pathlib import Path
from types import ModuleType
from typing import Any
import pytest
import nodes
# Guarantee we import the local nodes/ package, not a system-wide ComfyUI nodes package.
_REPO_ROOT = Path(__file__).resolve().parents[2]
_NODES_PKG = _REPO_ROOT / "nodes"
if str(_NODES_PKG) not in sys.path:
sys.path.insert(0, str(_NODES_PKG))
import nodes # noqa: E402
assert nodes.__file__.startswith(str(_REPO_ROOT)), f"Imported wrong nodes package: {nodes.__file__}"
EXPECTED_NODES: frozenset[str] = frozenset(
{
+55 -8
View File
@@ -33,8 +33,8 @@ class TestOllamaClientDiscovery:
try:
client_module.ollama = mock_ollama
client_module.OLLAMA_API_AVAILABLE = True
OllamaClient._cached_models = None
OllamaClient._cache_time = 0
client_module._MODEL_CACHE = None
client_module._CACHE_TIME = 0
client = OllamaClient()
models = client.discover_models()
# Both LoRA models should be first (sorted alphabetically within LoRA group)
@@ -44,8 +44,8 @@ class TestOllamaClientDiscovery:
finally:
client_module.ollama = orig_ollama
client_module.OLLAMA_API_AVAILABLE = orig_available
OllamaClient._cached_models = None
OllamaClient._cache_time = 0
client_module._MODEL_CACHE = None
client_module._CACHE_TIME = 0
def test_caching(self):
"""Second call should return cached results within 60s."""
@@ -59,8 +59,8 @@ class TestOllamaClientDiscovery:
try:
client_module.ollama = mock_ollama
client_module.OLLAMA_API_AVAILABLE = True
OllamaClient._cached_models = None
OllamaClient._cache_time = 0
client_module._MODEL_CACHE = None
client_module._CACHE_TIME = 0
client = OllamaClient()
_ = client.discover_models()
_ = client.discover_models()
@@ -68,8 +68,33 @@ class TestOllamaClientDiscovery:
finally:
client_module.ollama = orig_ollama
client_module.OLLAMA_API_AVAILABLE = orig_available
OllamaClient._cached_models = None
OllamaClient._cache_time = 0
client_module._MODEL_CACHE = None
client_module._CACHE_TIME = 0
def test_cache_shared_across_instances(self):
"""Module-level cache should be shared between OllamaClient instances."""
mock_models = [{"model": "qwen3:8b"}]
mock_ollama = MagicMock()
mock_ollama.list.return_value = {"models": mock_models}
import nodes.adapters.ollama_client as client_module
orig_ollama = getattr(client_module, "ollama", None)
orig_available = getattr(client_module, "OLLAMA_API_AVAILABLE", False)
try:
client_module.ollama = mock_ollama
client_module.OLLAMA_API_AVAILABLE = True
client_module._MODEL_CACHE = None
client_module._CACHE_TIME = 0
client1 = OllamaClient()
client2 = OllamaClient()
_ = client1.discover_models()
_ = client2.discover_models()
mock_ollama.list.assert_called_once() # Shared cache
finally:
client_module.ollama = orig_ollama
client_module.OLLAMA_API_AVAILABLE = orig_available
client_module._MODEL_CACHE = None
client_module._CACHE_TIME = 0
class TestOllamaClientHealth:
@@ -129,6 +154,28 @@ class TestOllamaClientHealth:
client_module.ollama = orig_ollama
client_module.OLLAMA_API_AVAILABLE = orig_available
def test_exact_tag_match_no_prefix_conflation(self):
"""qwen3:4b should NOT report loaded when only qwen3:8b is in VRAM."""
mock_ps = MagicMock()
mock_ps.models = [MagicMock(model="qwen3:8b")]
mock_ollama = MagicMock()
mock_ollama.ps.return_value = mock_ps
mock_ollama.list.return_value = {"models": []}
import nodes.adapters.ollama_client as client_module
orig_ollama = getattr(client_module, "ollama", None)
orig_available = getattr(client_module, "OLLAMA_API_AVAILABLE", False)
try:
client_module.ollama = mock_ollama
client_module.OLLAMA_API_AVAILABLE = True
client = OllamaClient()
healthy, _msg, loaded = client.check_health("qwen3:4b")
assert healthy is True
assert loaded is False
finally:
client_module.ollama = orig_ollama
client_module.OLLAMA_API_AVAILABLE = orig_available
class TestOllamaClientSubprocess:
"""Test suite for subprocess fallback."""
+4
View File
@@ -98,3 +98,7 @@ class TestUnknownMode:
(out,) = node.combine(prompt_1="x", prompt_2="y", mode="not_a_real_mode") # type: ignore[arg-type]
assert "Unknown mode" in out
assert "blend" in out # mentions valid options
def test_unknown_mode_with_single_prompt_still_validated(self, node: PromptCombinerNode) -> None:
(out,) = node.combine(prompt_1="x", mode="not_a_real_mode") # type: ignore[arg-type]
assert "Unknown mode" in out
+60
View File
@@ -0,0 +1,60 @@
"""Unit tests for PromptRefinerNode seed behavior."""
from unittest.mock import patch
from nodes.prompt_refiner_node import PromptRefinerNode
class TestPromptRefinerSeed:
"""Test suite for per-pass seed derivation."""
def test_per_pass_seed_increments(self):
"""Each pass should receive an incremented seed when seed != -1."""
node = PromptRefinerNode()
captured_seeds: list[int | None] = []
def _capture_seed(*, model, prompt, temperature, top_p, timeout, pbar, seed):
captured_seeds.append(seed)
return f"refined with seed={seed}"
with patch.object(node, "REFINEMENT_PROMPT", "{prompt}"), patch(
"nodes.prompt_refiner_node.OllamaClient.generate_streaming",
side_effect=_capture_seed,
):
result = node.refine(
prompt="a forest",
model="qwen3:8b",
passes=3,
seed=42,
temperature=0.5,
top_p=0.9,
timeout=30,
)
assert captured_seeds == [42, 43, 44]
assert "refined with seed=44" in result[0]
def test_random_seed_none_for_all_passes(self):
"""When seed is -1, all passes should receive None (random)."""
node = PromptRefinerNode()
captured_seeds: list[int | None] = []
def _capture_seed(*, model, prompt, temperature, top_p, timeout, pbar, seed):
captured_seeds.append(seed)
return "refined"
with patch.object(node, "REFINEMENT_PROMPT", "{prompt}"), patch(
"nodes.prompt_refiner_node.OllamaClient.generate_streaming",
side_effect=_capture_seed,
):
node.refine(
prompt="a forest",
model="qwen3:8b",
passes=2,
seed=-1,
temperature=0.5,
top_p=0.9,
timeout=30,
)
assert captured_seeds == [None, None]