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:
@@ -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."
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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),)
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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(
|
||||
{
|
||||
|
||||
@@ -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."""
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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]
|
||||
Reference in New Issue
Block a user