diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 6097946..8e58268 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -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." + diff --git a/nodes/adapters/ollama_client.py b/nodes/adapters/ollama_client.py index 126a38a..1902edc 100644 --- a/nodes/adapters/ollama_client.py +++ b/nodes/adapters/ollama_client.py @@ -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: diff --git a/nodes/prompt_combiner_node.py b/nodes/prompt_combiner_node.py index 17c97fc..a585687 100644 --- a/nodes/prompt_combiner_node.py +++ b/nodes/prompt_combiner_node.py @@ -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),) diff --git a/nodes/prompt_generator_node.py b/nodes/prompt_generator_node.py index ca50a87..766372b 100644 --- a/nodes/prompt_generator_node.py +++ b/nodes/prompt_generator_node.py @@ -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"", "", 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"
\s*.*?[\s\S]*?
", "", text, flags=re.DOTALL) # Remove common prefixes like "**Prompt:**" or "**Stable Diffusion Prompt:**" text = re.sub(r"\*\*(?:Stable Diffusion )?Prompt:\*\*\s*", "", text) diff --git a/nodes/prompt_refiner_node.py b/nodes/prompt_refiner_node.py index 75ea3a5..65863a3 100644 --- a/nodes/prompt_refiner_node.py +++ b/nodes/prompt_refiner_node.py @@ -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: diff --git a/pyproject.toml b/pyproject.toml index 9543c2a..2a0dadb 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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" diff --git a/scripts/check_version_tag.py b/scripts/check_version_tag.py index 43ee28f..b906705 100644 --- a/scripts/check_version_tag.py +++ b/scripts/check_version_tag.py @@ -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: diff --git a/tests/unit/test_extract_final_prompt.py b/tests/unit/test_extract_final_prompt.py index 44cb3d9..1dea645 100644 --- a/tests/unit/test_extract_final_prompt.py +++ b/tests/unit/test_extract_final_prompt.py @@ -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 blocks should be stripped.""" + text = "Some reasoning hereFinal prompt here" + assert extract_final_prompt(text) == "Final prompt here" + + def test_deepseek_think_multiline(self): + """Multiline blocks should be stripped.""" + text = "\nReasoning line 1\nReasoning line 2\n\nFinal prompt" + assert extract_final_prompt(text) == "Final prompt" + + def test_markdown_details_block(self): + """Markdown
blocks should be stripped.""" + text = "
ReasoningHidden reasoning
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 = ( + "DeepSeek reasoning\n" + "Thinking...\nQwen reasoning\n...done thinking.\n" + "
SummaryDetails
\n" + "Final clean prompt" + ) + assert extract_final_prompt(text) == "Final clean prompt" diff --git a/tests/unit/test_node_registration.py b/tests/unit/test_node_registration.py index 1e56b14..1e67be5 100644 --- a/tests/unit/test_node_registration.py +++ b/tests/unit/test_node_registration.py @@ -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( { diff --git a/tests/unit/test_ollama_client.py b/tests/unit/test_ollama_client.py index 8107fd1..a3e9245 100644 --- a/tests/unit/test_ollama_client.py +++ b/tests/unit/test_ollama_client.py @@ -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.""" diff --git a/tests/unit/test_prompt_combiner.py b/tests/unit/test_prompt_combiner.py index 930f5a4..1527eb6 100644 --- a/tests/unit/test_prompt_combiner.py +++ b/tests/unit/test_prompt_combiner.py @@ -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 diff --git a/tests/unit/test_prompt_refiner.py b/tests/unit/test_prompt_refiner.py new file mode 100644 index 0000000..7573361 --- /dev/null +++ b/tests/unit/test_prompt_refiner.py @@ -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]