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:
@@ -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)
|
||||
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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")
|
||||
|
||||
|
||||
@@ -10,7 +10,6 @@ addopts = -v --tb=short
|
||||
source = nodes
|
||||
omit =
|
||||
*/tests/*
|
||||
*/adapters/*
|
||||
*/__pycache__/*
|
||||
|
||||
[coverage:report]
|
||||
|
||||
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user