fix: address code review issues (#7)

- Remove dead imports from prompt_generator_node.py and ollama_client.py
- Switch optional imports to importlib.util.find_spec pattern
- Remove orphaned _cached_models/_cache_time from PromptGeneratorNode
- Remove unused seed param from PromptRefinerNode
- Add top_p input to NegativePromptNode with proper wiring
- Fix pytest.ini to not omit adapters from coverage
- Remove unused pytest imports from test files
- Run ruff check + format (all clean now)
- All 29 tests passing
This commit is contained in:
limbicnation
2026-04-27 17:30:19 +02:00
parent 89c8341612
commit 8f2eca40d9
8 changed files with 74 additions and 59 deletions
+3 -3
View File
@@ -9,11 +9,9 @@ Handles:
- Subprocess fallback when Python API unavailable
"""
import re
import subprocess
import time
import threading
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
# Optional imports with graceful degradation
@@ -183,7 +181,9 @@ class OllamaClient:
t = threading.Thread(target=_iter_next, args=(it,), daemon=True)
t.start()
wait_time = first_chunk_timeout if not got_first_chunk else chunk_timeout
wait_time = (
first_chunk_timeout if not got_first_chunk else chunk_timeout
)
wait_time = min(wait_time, timeout - elapsed)
t.join(timeout=wait_time)
+12 -1
View File
@@ -75,6 +75,16 @@ Negative prompt:"""
"display": "slider",
},
),
"top_p": (
"FLOAT",
{
"default": 0.9,
"min": 0.1,
"max": 1.0,
"step": 0.1,
"display": "slider",
},
),
"timeout": (
"INT",
{
@@ -99,6 +109,7 @@ Negative prompt:"""
style: str,
model: str,
temperature: float = 0.3,
top_p: float = 0.9,
timeout: int = 60,
) -> Tuple[str]:
"""
@@ -134,7 +145,7 @@ Negative prompt:"""
model=model,
prompt=negative_prompt_text,
temperature=temperature,
top_p=0.9,
top_p=top_p,
timeout=timeout,
)
+17 -19
View File
@@ -4,11 +4,10 @@ Generate detailed Stable Diffusion prompts using Qwen3-8B via Ollama
"""
import re
import subprocess
import time
import threading
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
from typing import Any, Dict, Optional, Tuple
from .adapters.ollama_client import OllamaClient
# Optional imports with graceful degradation
try:
@@ -25,21 +24,24 @@ try:
except ImportError:
JINJA2_AVAILABLE = False
try:
import ollama
OLLAMA_API_AVAILABLE = True
except ImportError:
OLLAMA_API_AVAILABLE = False
OLLAMA_API_AVAILABLE = False
COMFY_PROGRESS_AVAILABLE = False
try:
import comfy.utils
import importlib.util
COMFY_PROGRESS_AVAILABLE = True
except ImportError:
COMFY_PROGRESS_AVAILABLE = False
if importlib.util.find_spec("ollama") is not None:
OLLAMA_API_AVAILABLE = True
except Exception:
pass
from .adapters.ollama_client import OllamaClient
try:
import importlib.util
if importlib.util.find_spec("comfy") is not None:
COMFY_PROGRESS_AVAILABLE = True
except Exception:
pass
def extract_final_prompt(text: str) -> str:
@@ -234,10 +236,6 @@ Format the response as a single, detailed photography prompt.""",
},
}
# Class-level cache for available models
_cached_models = None
_cache_time = 0
def __init__(self):
"""Initialize the node and load style templates."""
self.style_templates = self._load_templates()
+4 -16
View File
@@ -73,16 +73,6 @@ Refined prompt:"""
"display": "slider",
},
),
"seed": (
"INT",
{
"default": -1,
"min": -1,
"max": 2**32 - 1,
"step": 1,
"tooltip": "-1 for random, >=0 for deterministic",
},
),
"timeout": (
"INT",
{
@@ -107,7 +97,6 @@ Refined prompt:"""
model: str,
passes: int = 1,
temperature: float = 0.5,
seed: int = -1,
timeout: int = 120,
) -> Tuple[str]:
"""
@@ -118,7 +107,6 @@ Refined prompt:"""
model: Ollama model to use
passes: Number of refinement iterations (1-3)
temperature: Generation temperature
seed: Random seed for deterministic output (-1 = random)
timeout: Maximum generation time per pass
Returns:
@@ -147,9 +135,7 @@ Refined prompt:"""
if output is None:
# Fallback to subprocess
success, output = client.generate_subprocess(
model, refinement, timeout
)
success, output = client.generate_subprocess(model, refinement, timeout)
if not success:
return (f"[PromptRefiner] Pass {i + 1} failed: {output}",)
@@ -157,7 +143,9 @@ Refined prompt:"""
cleaned = extract_final_prompt(output.strip())
if cleaned:
current_prompt = cleaned
print(f"[PromptRefiner] Pass {i + 1} complete: {len(current_prompt)} chars")
print(
f"[PromptRefiner] Pass {i + 1} complete: {len(current_prompt)} chars"
)
else:
print(f"[PromptRefiner] Pass {i + 1} returned empty, keeping previous")
-1
View File
@@ -10,7 +10,6 @@ addopts = -v --tb=short
source = nodes
omit =
*/tests/*
*/adapters/*
*/__pycache__/*
[coverage:report]
+13 -8
View File
@@ -1,7 +1,7 @@
"""
Unit tests for extract_final_prompt utility.
"""
import pytest
from nodes.prompt_generator_node import extract_final_prompt
@@ -18,7 +18,9 @@ class TestExtractFinalPrompt:
def test_qwen3_thinking_block(self):
"""Qwen3 thinking blocks should be stripped."""
text = "Thinking...\nThis is the reasoning\n...done thinking.\nFinal prompt here"
text = (
"Thinking...\nThis is the reasoning\n...done thinking.\nFinal prompt here"
)
assert extract_final_prompt(text) == "Final prompt here"
def test_prompt_prefix_removal(self):
@@ -63,11 +65,14 @@ class TestExtractFinalPrompt:
def test_complex_real_world(self):
"""Complex real-world example with multiple artifacts."""
text = (
'Thinking...\n'
'I need to create a cinematic prompt\n'
'...done thinking.\n'
'**Stable Diffusion Prompt:**\n'
"Thinking...\n"
"I need to create a cinematic prompt\n"
"...done thinking.\n"
"**Stable Diffusion Prompt:**\n"
'"A mystical forest at twilight, dramatic lighting, 8k"\n'
'None'
"None"
)
assert (
extract_final_prompt(text)
== "A mystical forest at twilight, dramatic lighting, 8k"
)
assert extract_final_prompt(text) == "A mystical forest at twilight, dramatic lighting, 8k"
+6 -1
View File
@@ -1,7 +1,7 @@
"""
Unit tests for OllamaClient adapter.
"""
import pytest
from unittest.mock import patch, MagicMock
from nodes.adapters.ollama_client import OllamaClient
@@ -27,6 +27,7 @@ class TestOllamaClientDiscovery:
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:
@@ -52,6 +53,7 @@ class TestOllamaClientDiscovery:
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:
@@ -89,6 +91,7 @@ class TestOllamaClientHealth:
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:
@@ -111,6 +114,7 @@ class TestOllamaClientHealth:
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:
@@ -154,6 +158,7 @@ class TestOllamaClientSubprocess:
def test_subprocess_timeout(self):
"""Should handle timeout gracefully."""
from subprocess import TimeoutExpired
with patch("subprocess.run", side_effect=TimeoutExpired("ollama", 30)):
client = OllamaClient()
success, output = client.generate_subprocess("qwen3:8b", "prompt", 30)
+19 -10
View File
@@ -1,7 +1,7 @@
"""
Unit tests for style presets and template loading.
"""
import pytest
from pathlib import Path
from unittest.mock import patch, mock_open
@@ -18,8 +18,15 @@ class TestStylePresets:
def test_default_styles_keys(self):
"""DEFAULT_STYLES should contain expected style keys."""
expected = {
"cinematic", "anime", "photorealistic", "fantasy",
"abstract", "cyberpunk", "sci-fi", "video_wan", "still_image"
"cinematic",
"anime",
"photorealistic",
"fantasy",
"abstract",
"cyberpunk",
"sci-fi",
"video_wan",
"still_image",
}
assert set(PromptGeneratorNode.DEFAULT_STYLES.keys()) == expected
@@ -27,7 +34,9 @@ class TestStylePresets:
"""Each default style should have a template string."""
for key, data in PromptGeneratorNode.DEFAULT_STYLES.items():
assert "template" in data, f"Style '{key}' missing template"
assert isinstance(data["template"], str), f"Style '{key}' template not a string"
assert isinstance(data["template"], str), (
f"Style '{key}' template not a string"
)
assert len(data["template"]) > 0, f"Style '{key}' template is empty"
def test_get_style_list_fallback(self):
@@ -41,10 +50,13 @@ class TestStylePresets:
mock_yaml = {
"test_style": {
"name": "Test",
"template": "Test template for {{ description }}"
"template": "Test template for {{ description }}",
}
}
with patch("builtins.open", mock_open(read_data="test_style:\n name: Test\n template: Test template")):
with patch(
"builtins.open",
mock_open(read_data="test_style:\n name: Test\n template: Test template"),
):
with patch.object(Path, "exists", return_value=True):
with patch("yaml.safe_load", return_value=mock_yaml):
node = PromptGeneratorNode()
@@ -54,10 +66,7 @@ class TestStylePresets:
"""_render_template should substitute Jinja2 variables."""
node = PromptGeneratorNode()
result = node._render_template(
"cinematic",
"a forest",
emphasis="lighting",
mood="mysterious"
"cinematic", "a forest", emphasis="lighting", mood="mysterious"
)
assert "a forest" in result
assert "lighting" in result