feat: modernize typing to PEP 585, add ruff config, refactor combiner mode handling

Changes:
- Add [tool.ruff] config with target-version py310 and a curated rule set
  (E/F/W, I, UP, B, SIM, RUF) so future drift is caught in CI lint.
- PEP 585 sweep across all node modules: drop legacy typing.Dict / List /
  Tuple / Optional in favor of dict / list / tuple / `X | None`. Annotate
  class-level mutable defaults as ClassVar to satisfy RUF012.
- PromptCombinerNode: replace the magic-string mode chain with a CombineMode
  StrEnum + match statement. The dropdown choices in INPUT_TYPES are now
  derived from the same Literal alias used in the function signature, so the
  UI and the type contract can't drift apart.
- Smoke tests for PromptCombinerNode (14 tests) covering enum mapping, all
  three modes, edge cases, and the unknown-mode error path. Brings combiner
  coverage from 28% to 96%.
- Tidy preexisting issues surfaced by the new lint rules: B904 except chaining
  in style_presets, RUF013 implicit Optional, RUF059 unused unpack, SIM117
  nested-with consolidation in tests.

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
limbicnation
2026-04-30 06:14:44 +02:00
co-authored by Claude Opus 4.7
parent 8e19db8e72
commit 839188382c
12 changed files with 250 additions and 164 deletions
+20 -27
View File
@@ -11,9 +11,9 @@ Handles:
import logging
import subprocess
import time
import threading
from typing import Any, Dict, List, Optional, Tuple
import time
from typing import Any, ClassVar
logger = logging.getLogger(__name__)
@@ -45,20 +45,20 @@ class OllamaClient:
"""
CHUNK_TIMEOUT = 30
DEFAULT_MODELS = ["qwen3:8b", "qwen3:4b", "llama3.2:latest"]
LORA_KEYWORDS = ["lora", "limbicnation", "fine", "style", "prompt"]
DEFAULT_MODELS: ClassVar[list[str]] = ["qwen3:8b", "qwen3:4b", "llama3.2:latest"]
LORA_KEYWORDS: ClassVar[list[str]] = ["lora", "limbicnation", "fine", "style", "prompt"]
def __init__(self, logger_prefix: str = "OllamaClient"):
self.logger_prefix = logger_prefix
# Instance-level cache (was class-level, causing race conditions)
self._cached_models: Optional[List[str]] = None
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]:
def discover_models(self) -> list[str]:
"""
Fetch available Ollama models with instance-level caching.
Prioritizes LoRA-enhanced models.
@@ -72,14 +72,12 @@ class OllamaClient:
try:
result = ollama.list()
models = [
m.get("model", "") for m in result.get("models", []) if "model" in m
]
models = [m.get("model", "") for m in result.get("models", []) if "model" in m]
if not models:
return self.DEFAULT_MODELS
def sort_key(name: str) -> Tuple[int, str]:
def sort_key(name: str) -> tuple[int, str]:
name_lower = name.lower()
is_lora = any(kw in name_lower for kw in self.LORA_KEYWORDS)
return (0 if is_lora else 1, name)
@@ -105,7 +103,7 @@ class OllamaClient:
self._log(f"Unexpected error fetching models: {e}")
return self.DEFAULT_MODELS
def check_health(self, model: str) -> Tuple[bool, str, bool]:
def check_health(self, model: str) -> tuple[bool, str, bool]:
"""
Quick health check: is Ollama running and is the model loaded?
@@ -129,10 +127,7 @@ 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
)
is_loaded = any(model == rm or model.startswith(rm.split(":")[0]) for rm in running_models)
if is_loaded:
return (True, f"Model '{model}' is loaded in VRAM", True)
else:
@@ -153,9 +148,9 @@ class OllamaClient:
temperature: float,
top_p: float,
timeout: int,
pbar: Optional[Any] = None,
seed: Optional[int] = None,
) -> Optional[str]:
pbar: Any | None = None,
seed: int | None = None,
) -> str | None:
"""
Stream ollama.generate() with per-chunk and total timeout enforcement.
@@ -168,14 +163,14 @@ class OllamaClient:
if not OLLAMA_API_AVAILABLE:
return None
chunks: List[str] = []
chunks: list[str] = []
start = time.monotonic()
first_chunk_timeout = min(timeout * 0.6, 90)
chunk_timeout = self.CHUNK_TIMEOUT
got_first_chunk = False
try:
options: Dict[str, Any] = {"temperature": temperature, "top_p": top_p}
options: dict[str, Any] = {"temperature": temperature, "top_p": top_p}
if seed is not None:
options["seed"] = seed
@@ -186,7 +181,7 @@ class OllamaClient:
options=options,
)
result_holder: Dict[str, Any] = {}
result_holder: dict[str, Any] = {}
def _iter_next(it):
try:
@@ -210,9 +205,7 @@ 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)
@@ -267,7 +260,7 @@ class OllamaClient:
model: str,
prompt: str,
timeout: int,
) -> Tuple[bool, str]:
) -> tuple[bool, str]:
"""
Fallback generation via subprocess call to ollama CLI.
@@ -297,9 +290,9 @@ class OllamaClient:
except FileNotFoundError:
return (False, "Ollama not found. Install from: https://ollama.ai")
except Exception as e:
return (False, f"Error: {str(e)}")
return (False, f"Error: {e!s}")
def create_progress_bar(self, unique_id: Optional[str] = None) -> Optional[Any]:
def create_progress_bar(self, unique_id: str | None = None) -> Any | None:
"""Create a ComfyUI progress bar if available."""
if COMFY_PROGRESS_AVAILABLE and unique_id is not None:
try:
+5 -7
View File
@@ -4,7 +4,7 @@ Generates negative prompts from positive prompts using style-aware templates.
"""
import logging
from typing import Any, Dict, Tuple
from typing import Any, ClassVar
from .adapters.ollama_client import OllamaClient
from .prompt_generator_node import extract_final_prompt
@@ -33,7 +33,7 @@ Focus on common artifacts for this style: {style_hints}
Negative prompt:"""
STYLE_HINTS = {
STYLE_HINTS: ClassVar[dict[str, str]] = {
"cinematic": "blurry, overexposed, underexposed, shaky cam, lens flare abuse, bad CGI",
"anime": "3d render, realistic, western cartoon, bad anatomy, extra limbs, deformed",
"photorealistic": "painting, illustration, cartoon, oversaturated, artificial look",
@@ -46,7 +46,7 @@ Negative prompt:"""
}
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Any]:
def INPUT_TYPES(cls) -> dict[str, Any]:
styles = list(cls.STYLE_HINTS.keys())
return {
"required": {
@@ -114,7 +114,7 @@ Negative prompt:"""
temperature: float = 0.3,
top_p: float = 0.9,
timeout: int = 60,
) -> Tuple[str]:
) -> tuple[str]:
"""
Generate a negative prompt from a positive prompt.
@@ -154,9 +154,7 @@ Negative prompt:"""
if output is None:
# Fallback to subprocess
success, output = client.generate_subprocess(
model, negative_prompt_text, timeout
)
success, output = client.generate_subprocess(model, negative_prompt_text, timeout)
if not success:
return (f"[NegativePrompt] Generation failed: {output}",)
+45 -27
View File
@@ -3,7 +3,20 @@ Prompt Combiner Node for ComfyUI
Merges multiple prompt strings with configurable blending strategies.
"""
from typing import Any, Dict, List, Tuple
from enum import StrEnum
from typing import Any, Literal, get_args
class CombineMode(StrEnum):
"""Supported strategies for combining multiple prompts."""
BLEND = "blend"
CONCAT = "concat"
WEIGHTED_AVERAGE = "weighted_average"
# Single source of truth for the choices exposed in INPUT_TYPES and accepted by combine().
ModeLiteral = Literal["blend", "concat", "weighted_average"]
class PromptCombinerNode:
@@ -17,7 +30,7 @@ class PromptCombinerNode:
"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Any]:
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"prompt_1": (
@@ -29,8 +42,8 @@ class PromptCombinerNode:
},
),
"mode": (
["blend", "concat", "weighted_average"],
{"default": "blend"},
list(get_args(ModeLiteral)),
{"default": CombineMode.BLEND.value},
),
},
"optional": {
@@ -117,7 +130,7 @@ class PromptCombinerNode:
def combine(
self,
prompt_1: str,
mode: str,
mode: ModeLiteral,
prompt_2: str = "",
prompt_3: str = "",
prompt_4: str = "",
@@ -126,7 +139,7 @@ class PromptCombinerNode:
weight_3: float = 1.0,
weight_4: float = 1.0,
separator: str = ", ",
) -> Tuple[str]:
) -> tuple[str]:
"""
Combine multiple prompts using the selected mode.
@@ -140,16 +153,16 @@ class PromptCombinerNode:
Returns:
Tuple containing the combined prompt string
"""
# Collect non-empty prompts with their weights
prompts: List[Tuple[str, float]] = []
for p, w in [
(prompt_1, weight_1),
(prompt_2, weight_2),
(prompt_3, weight_3),
(prompt_4, weight_4),
]:
if p and p.strip():
prompts.append((p.strip(), w))
prompts: list[tuple[str, float]] = [
(p.strip(), w)
for p, w in (
(prompt_1, weight_1),
(prompt_2, weight_2),
(prompt_3, weight_3),
(prompt_4, weight_4),
)
if p and p.strip()
]
if not prompts:
return ("[PromptCombiner] At least one prompt is required.",)
@@ -157,16 +170,21 @@ class PromptCombinerNode:
if len(prompts) == 1:
return (prompts[0][0],)
if mode == "blend":
return (self._blend(prompts),)
elif mode == "concat":
return (self._concat(prompts, separator),)
elif mode == "weighted_average":
return (self._weighted_average(prompts),)
else:
return (f"[PromptCombiner] Unknown mode: {mode}",)
try:
selected = CombineMode(mode)
except ValueError:
valid = ", ".join(m.value for m in CombineMode)
return (f"[PromptCombiner] Unknown mode {mode!r}. Valid: {valid}",)
def _blend(self, prompts: List[Tuple[str, float]]) -> str:
match selected:
case CombineMode.BLEND:
return (self._blend(prompts),)
case CombineMode.CONCAT:
return (self._concat(prompts, separator),)
case CombineMode.WEIGHTED_AVERAGE:
return (self._weighted_average(prompts),)
def _blend(self, prompts: list[tuple[str, float]]) -> str:
"""
Blend prompts using ComfyUI-style emphasis markers.
Higher weight = more parentheses emphasis.
@@ -186,12 +204,12 @@ class PromptCombinerNode:
parts.append(text)
return ", ".join(parts)
def _concat(self, prompts: List[Tuple[str, float]], separator: str) -> str:
def _concat(self, prompts: list[tuple[str, float]], separator: str) -> str:
"""Simple concatenation with separator."""
texts = [p[0] for p in prompts]
return separator.join(texts)
def _weighted_average(self, prompts: List[Tuple[str, float]]) -> str:
def _weighted_average(self, prompts: list[tuple[str, float]]) -> str:
"""
Weighted text combination using emphasis markers.
Higher-weighted prompts receive stronger ComfyUI emphasis parentheses.
+23 -36
View File
@@ -5,7 +5,7 @@ Generate detailed Stable Diffusion prompts using Qwen3-8B via Ollama
import re
from pathlib import Path
from typing import Any, Dict, Optional, Tuple
from typing import Any, ClassVar
from .adapters.ollama_client import OllamaClient
@@ -18,7 +18,7 @@ except ImportError:
YAML_AVAILABLE = False
try:
from jinja2 import Environment, BaseLoader
from jinja2 import BaseLoader, Environment
JINJA2_AVAILABLE = True
except ImportError:
@@ -53,9 +53,7 @@ def extract_final_prompt(text: str) -> str:
return text
# Remove Qwen3 thinking blocks: "Thinking...\n...\n...done thinking.\n"
text = re.sub(
r"Thinking\.\.\.[\s\S]*?\.\.\.done thinking\.[\s]*", "", text, flags=re.DOTALL
)
text = re.sub(r"Thinking\.\.\.[\s\S]*?\.\.\.done thinking\.[\s]*", "", text, flags=re.DOTALL)
# Remove common prefixes like "**Prompt:**" or "**Stable Diffusion Prompt:**"
text = re.sub(r"\*\*(?:Stable Diffusion )?Prompt:\*\*\s*", "", text)
@@ -84,7 +82,7 @@ class PromptGeneratorNode:
CHUNK_TIMEOUT = 30
# Default style templates (fallback when YAML not available)
DEFAULT_STYLES = {
DEFAULT_STYLES: ClassVar[dict[str, dict[str, str]]] = {
"cinematic": {
"name": "Cinematic",
"description": "Dramatic lighting and composition for film-quality images",
@@ -214,7 +212,11 @@ Format the response as a single, detailed sci-fi prompt.""",
"video_wan": {
"name": "Video (WanVideo)",
"description": "Minimalist template optimized for WanVideo LoRA",
"template": "Generate a video prompt for: {{ description }}{% if emphasis %} with focus on {{ emphasis }}{% endif %}{% if mood %}, mood is {{ mood }}{% endif %}",
"template": (
"Generate a video prompt for: {{ description }}"
"{% if emphasis %} with focus on {{ emphasis }}{% endif %}"
"{% if mood %}, mood is {{ mood }}{% endif %}"
),
},
"still_image": {
"name": "Still Image (Photography)",
@@ -252,18 +254,16 @@ Format the response as a single, detailed photography prompt.""",
template_path = Path(__file__).parent.parent / "config" / "templates.yaml"
if YAML_AVAILABLE and template_path.exists():
try:
with open(template_path, "r") as f:
with open(template_path) as f:
templates = yaml.safe_load(f)
if templates:
return list(templates.keys())
except Exception as e:
print(
f"[PromptGenerator] Warning: Failed to load style list from templates.yaml: {e}"
)
print(f"[PromptGenerator] Warning: Failed to load style list from templates.yaml: {e}")
return list(cls.DEFAULT_STYLES.keys())
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Any]:
def INPUT_TYPES(cls) -> dict[str, Any]:
"""Define input parameters for the node."""
available_models = cls._get_available_models()
available_styles = cls._get_style_list()
@@ -280,18 +280,12 @@ Format the response as a single, detailed photography prompt.""",
),
"style": (
available_styles,
{
"default": available_styles[0]
if available_styles
else "cinematic"
},
{"default": available_styles[0] if available_styles else "cinematic"},
),
"model": (
available_models,
{
"default": available_models[0]
if available_models
else "qwen3:8b",
"default": available_models[0] if available_models else "qwen3:8b",
"tooltip": "Select Ollama model. LoRA-enhanced models appear first.",
},
),
@@ -362,18 +356,16 @@ Format the response as a single, detailed photography prompt.""",
CATEGORY = "text/generation"
OUTPUT_NODE = False
def _load_templates(self) -> Dict[str, Any]:
def _load_templates(self) -> dict[str, Any]:
"""Load style templates from YAML file or use defaults."""
template_path = Path(__file__).parent.parent / "config" / "templates.yaml"
if YAML_AVAILABLE and template_path.exists():
try:
with open(template_path, "r") as f:
with open(template_path) as f:
templates = yaml.safe_load(f)
if templates:
print(
f"[PromptGenerator] Loaded {len(templates)} styles from templates.yaml"
)
print(f"[PromptGenerator] Loaded {len(templates)} styles from templates.yaml")
return templates
except Exception as e:
print(f"[PromptGenerator] Warning: Failed to load templates.yaml: {e}")
@@ -385,16 +377,13 @@ Format the response as a single, detailed photography prompt.""",
self,
style: str,
description: str,
emphasis: Optional[str] = None,
mood: Optional[str] = None,
emphasis: str | None = None,
mood: str | None = None,
) -> str:
"""Render a Jinja2 template with the given variables."""
template_data = self.style_templates.get(style)
if template_data is None:
print(
f"[PromptGenerator] Warning: style '{style}' not found in templates, "
f"falling back to 'cinematic'"
)
print(f"[PromptGenerator] Warning: style '{style}' not found in templates, falling back to 'cinematic'")
template_data = self.DEFAULT_STYLES.get("cinematic")
# Handle YAML format with 'template' key
@@ -444,8 +433,8 @@ Format the response as a single, detailed photography prompt.""",
include_reasoning: bool = False,
model: str = "qwen3:8b",
timeout: int = 120,
unique_id: str = None,
) -> Tuple[str]:
unique_id: str | None = None,
) -> tuple[str]:
"""
Generate a detailed image prompt using Ollama with streaming and progress.
@@ -497,9 +486,7 @@ Format the response as a single, detailed photography prompt.""",
if is_healthy and not is_model_loaded:
# Cold start: add 30% buffer
effective_timeout = min(int(timeout * 1.3), 600)
print(
f"[PromptGenerator] Cold start detected, effective timeout: {effective_timeout}s"
)
print(f"[PromptGenerator] Cold start detected, effective timeout: {effective_timeout}s")
output = client.generate_streaming(
model=model,
+5 -5
View File
@@ -4,7 +4,7 @@ Refines a raw prompt through iterative LLM passes for higher quality output.
"""
import logging
from typing import Any, Dict, Tuple, Optional
from typing import Any
from .adapters.ollama_client import OllamaClient
from .prompt_generator_node import extract_final_prompt
@@ -36,7 +36,7 @@ Original prompt: {prompt}
Refined prompt:"""
@classmethod
def INPUT_TYPES(cls) -> Dict[str, Any]:
def INPUT_TYPES(cls) -> dict[str, Any]:
return {
"required": {
"prompt": (
@@ -122,8 +122,8 @@ Refined prompt:"""
top_p: float = 0.9,
seed: int = -1,
timeout: int = 120,
unique_id: Optional[str] = None,
) -> Tuple[str]:
unique_id: str | None = None,
) -> tuple[str]:
"""
Refine a prompt through iterative LLM passes.
@@ -147,7 +147,7 @@ Refined prompt:"""
current_prompt = prompt.strip()
# Determine effective seed
effective_seed: Optional[int] = None if seed == -1 else seed
effective_seed: int | None = None if seed == -1 else seed
for i in range(passes):
logger.info("Pass %d/%d with model='%s'", i + 1, passes, model)
+1 -3
View File
@@ -3,8 +3,6 @@ Style Applier Node for ComfyUI
Applies style keywords to prompts using the shared StylePreset system (9 styles).
"""
from typing import Tuple
class StyleApplierNode:
"""
@@ -66,7 +64,7 @@ class StyleApplierNode:
position: str = "suffix",
emphasis: str = "medium",
include_technical: bool = True,
) -> Tuple[str, str]:
) -> tuple[str, str]:
"""Apply style keywords to a prompt."""
from style_presets import StylePreset
+14
View File
@@ -18,3 +18,17 @@ Icon = "https://raw.githubusercontent.com/Limbicnation/ComfyUI-PromptGenerator/m
[tool.comfy.python]
requires-python = ">=3.10"
[tool.ruff]
target-version = "py310"
line-length = 120
[tool.ruff.lint]
select = [
"E", "F", "W", # pycodestyle / pyflakes
"I", # isort
"UP", # pyupgrade — drives PEP 585 / 604 rewrites
"B", # flake8-bugbear
"SIM", # flake8-simplify
"RUF", # ruff-specific
]
+21 -36
View File
@@ -14,7 +14,6 @@ Usage in ComfyUI nodes:
"""
from dataclasses import dataclass, field
from typing import Dict, List, Optional, Tuple
from enum import Enum
@@ -36,32 +35,20 @@ class StyleMode(Enum):
class StyleKeywords:
"""Container for style-specific keywords and descriptors."""
primary: List[str] = field(default_factory=list)
lighting: List[str] = field(default_factory=list)
technical: List[str] = field(default_factory=list)
composition: List[str] = field(default_factory=list)
texture: List[str] = field(default_factory=list)
primary: list[str] = field(default_factory=list)
lighting: list[str] = field(default_factory=list)
technical: list[str] = field(default_factory=list)
composition: list[str] = field(default_factory=list)
texture: list[str] = field(default_factory=list)
def to_prompt_string(self, separator: str = ", ") -> str:
"""Convert all keywords to a single prompt string."""
all_keywords = (
self.primary
+ self.lighting
+ self.technical
+ self.composition
+ self.texture
)
all_keywords = self.primary + self.lighting + self.technical + self.composition + self.texture
return separator.join(all_keywords)
def to_list(self) -> List[str]:
def to_list(self) -> list[str]:
"""Return all keywords as a flat list (ComfyUI-compatible)."""
return (
self.primary
+ self.lighting
+ self.technical
+ self.composition
+ self.texture
)
return self.primary + self.lighting + self.technical + self.composition + self.texture
@dataclass
@@ -74,7 +61,7 @@ class StyleDefinition:
# Style definitions - single source of truth
STYLE_DEFINITIONS: Dict[StyleMode, StyleDefinition] = {
STYLE_DEFINITIONS: dict[StyleMode, StyleDefinition] = {
StyleMode.CINEMATIC: StyleDefinition(
name="Cinematic",
description="Film-like visuals with dramatic lighting and anamorphic qualities",
@@ -333,12 +320,12 @@ def get_style_keywords(style: str) -> StyleKeywords:
try:
mode = StyleMode(style.lower())
return STYLE_DEFINITIONS[mode].keywords
except ValueError:
except ValueError as err:
available = [s.value for s in StyleMode]
raise ValueError(f"Unknown style '{style}'. Available: {', '.join(available)}")
raise ValueError(f"Unknown style '{style}'. Available: {', '.join(available)}") from err
def get_available_styles() -> List[str]:
def get_available_styles() -> list[str]:
"""Get list of available style names."""
return [mode.value for mode in StyleMode]
@@ -350,17 +337,15 @@ class StylePreset:
self._styles = STYLE_DEFINITIONS
@staticmethod
def get_style_choices() -> Tuple[str, ...]:
def get_style_choices() -> tuple[str, ...]:
"""Get available style choices as tuple (ComfyUI dropdown format)."""
return tuple(mode.value for mode in StyleMode)
def get_style_keywords(self, style: str) -> List[str]:
def get_style_keywords(self, style: str) -> list[str]:
"""Get style keywords as a list."""
return get_style_keywords(style).to_list()
def get_style_prompt(
self, style: str, emphasis: Optional[str] = None, include_technical: bool = True
) -> str:
def get_style_prompt(self, style: str, emphasis: str | None = None, include_technical: bool = True) -> str:
"""Get a formatted style prompt string."""
keywords = get_style_keywords(style)
@@ -380,11 +365,11 @@ class StylePreset:
__all__ = [
"StyleMode",
"StyleKeywords",
"StyleDefinition",
"StylePreset",
"get_style_keywords",
"get_available_styles",
"STYLE_DEFINITIONS",
"StyleDefinition",
"StyleKeywords",
"StyleMode",
"StylePreset",
"get_available_styles",
"get_style_keywords",
]
+2 -7
View File
@@ -18,9 +18,7 @@ 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):
@@ -72,7 +70,4 @@ class TestExtractFinalPrompt:
'"A mystical forest at twilight, dramatic lighting, 8k"\n'
"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"
+2 -2
View File
@@ -2,7 +2,7 @@
Unit tests for OllamaClient adapter.
"""
from unittest.mock import patch, MagicMock
from unittest.mock import MagicMock, patch
from nodes.adapters.ollama_client import OllamaClient
@@ -79,7 +79,7 @@ class TestOllamaClientHealth:
"""Should report unhealthy when ollama package missing."""
with patch("nodes.adapters.ollama_client.OLLAMA_API_AVAILABLE", False):
client = OllamaClient()
healthy, msg, loaded = client.check_health("qwen3:8b")
healthy, msg, _loaded = client.check_health("qwen3:8b")
assert healthy is False
assert "not available" in msg
+100
View File
@@ -0,0 +1,100 @@
"""Smoke tests for PromptCombinerNode (pure logic, no I/O)."""
from __future__ import annotations
import pytest
from nodes.prompt_combiner_node import CombineMode, ModeLiteral, PromptCombinerNode
@pytest.fixture
def node() -> PromptCombinerNode:
return PromptCombinerNode()
class TestCombineModeEnum:
def test_enum_values_are_canonical_strings(self) -> None:
assert CombineMode.BLEND.value == "blend"
assert CombineMode.CONCAT.value == "concat"
assert CombineMode.WEIGHTED_AVERAGE.value == "weighted_average"
def test_input_types_dropdown_matches_literal(self) -> None:
from typing import get_args
choices = PromptCombinerNode.INPUT_TYPES()["required"]["mode"][0]
assert choices == list(get_args(ModeLiteral))
class TestCombineEmptyInputs:
def test_no_prompts_returns_helpful_message(self, node: PromptCombinerNode) -> None:
(out,) = node.combine(prompt_1="", mode="blend")
assert "At least one prompt" in out
def test_whitespace_only_prompts_treated_as_empty(self, node: PromptCombinerNode) -> None:
(out,) = node.combine(prompt_1=" ", mode="blend", prompt_2="\n\t")
assert "At least one prompt" in out
def test_single_prompt_returned_unmodified(self, node: PromptCombinerNode) -> None:
(out,) = node.combine(prompt_1="a forest at dawn", mode="blend")
assert out == "a forest at dawn"
class TestBlendMode:
def test_high_weight_gets_double_parens(self, node: PromptCombinerNode) -> None:
(out,) = node.combine(prompt_1="forest", weight_1=1.6, prompt_2="dawn", weight_2=1.0, mode="blend")
assert "((forest))" in out
assert "dawn" in out
def test_moderate_weight_gets_single_parens(self, node: PromptCombinerNode) -> None:
(out,) = node.combine(prompt_1="forest", weight_1=1.3, prompt_2="dawn", weight_2=1.0, mode="blend")
assert "(forest)" in out
def test_low_weight_gets_brackets(self, node: PromptCombinerNode) -> None:
(out,) = node.combine(prompt_1="forest", weight_1=0.4, prompt_2="dawn", weight_2=1.0, mode="blend")
assert "[forest]" in out
def test_zero_weight_excluded(self, node: PromptCombinerNode) -> None:
(out,) = node.combine(prompt_1="forest", weight_1=0.0, prompt_2="dawn", weight_2=1.0, mode="blend")
assert "forest" not in out
assert "dawn" in out
class TestConcatMode:
def test_default_separator(self, node: PromptCombinerNode) -> None:
(out,) = node.combine(prompt_1="a", mode="concat", prompt_2="b", prompt_3="c")
assert out == "a, b, c"
def test_custom_separator(self, node: PromptCombinerNode) -> None:
(out,) = node.combine(prompt_1="a", mode="concat", prompt_2="b", separator=" | ")
assert out == "a | b"
class TestWeightedAverageMode:
def test_dominant_weight_gets_triple_parens(self, node: PromptCombinerNode) -> None:
# Triple parens requires ratio >= 2.0, i.e. dominant weight at least 2x avg.
# weights 2.0, 0.0 → avg 1.0 → hero ratio 2.0 → triple parens
(out,) = node.combine(
prompt_1="hero",
weight_1=2.0,
prompt_2="extras",
weight_2=0.0,
mode="weighted_average",
)
assert "(((hero)))" in out
def test_zero_total_weight_falls_back_to_join(self, node: PromptCombinerNode) -> None:
(out,) = node.combine(
prompt_1="a",
weight_1=0.0,
prompt_2="b",
weight_2=0.0,
mode="weighted_average",
)
assert out == "a, b"
class TestUnknownMode:
def test_unknown_mode_returns_friendly_error(self, node: PromptCombinerNode) -> None:
(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
+12 -14
View File
@@ -3,7 +3,7 @@ Unit tests for style presets and template loading.
"""
from pathlib import Path
from unittest.mock import patch, mock_open
from unittest.mock import mock_open, patch
from nodes.prompt_generator_node import PromptGeneratorNode
@@ -34,9 +34,7 @@ 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):
@@ -53,21 +51,21 @@ class TestStylePresets:
"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"),
),
patch.object(Path, "exists", return_value=True),
patch("yaml.safe_load", return_value=mock_yaml),
):
with patch.object(Path, "exists", return_value=True):
with patch("yaml.safe_load", return_value=mock_yaml):
node = PromptGeneratorNode()
assert "test_style" in node.style_templates
node = PromptGeneratorNode()
assert "test_style" in node.style_templates
def test_render_template_jinja2(self):
"""_render_template should substitute Jinja2 variables."""
node = PromptGeneratorNode()
result = node._render_template(
"cinematic", "a forest", emphasis="lighting", mood="mysterious"
)
result = node._render_template("cinematic", "a forest", emphasis="lighting", mood="mysterious")
assert "a forest" in result
assert "lighting" in result
assert "mysterious" in result