diff --git a/nodes/adapters/ollama_client.py b/nodes/adapters/ollama_client.py index b32278c..126a38a 100644 --- a/nodes/adapters/ollama_client.py +++ b/nodes/adapters/ollama_client.py @@ -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: diff --git a/nodes/negative_prompt_node.py b/nodes/negative_prompt_node.py index 77e535e..6659d9e 100644 --- a/nodes/negative_prompt_node.py +++ b/nodes/negative_prompt_node.py @@ -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}",) diff --git a/nodes/prompt_combiner_node.py b/nodes/prompt_combiner_node.py index 892f533..17c97fc 100644 --- a/nodes/prompt_combiner_node.py +++ b/nodes/prompt_combiner_node.py @@ -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. diff --git a/nodes/prompt_generator_node.py b/nodes/prompt_generator_node.py index d70e2e1..ca50a87 100644 --- a/nodes/prompt_generator_node.py +++ b/nodes/prompt_generator_node.py @@ -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, diff --git a/nodes/prompt_refiner_node.py b/nodes/prompt_refiner_node.py index 38bfcd8..75ea3a5 100644 --- a/nodes/prompt_refiner_node.py +++ b/nodes/prompt_refiner_node.py @@ -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) diff --git a/nodes/style_applier_node.py b/nodes/style_applier_node.py index 0438e7f..d6e3a07 100644 --- a/nodes/style_applier_node.py +++ b/nodes/style_applier_node.py @@ -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 diff --git a/pyproject.toml b/pyproject.toml index cdfd388..04c46c5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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 +] diff --git a/style_presets.py b/style_presets.py index 049b5ff..c265240 100644 --- a/style_presets.py +++ b/style_presets.py @@ -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", ] diff --git a/tests/unit/test_extract_final_prompt.py b/tests/unit/test_extract_final_prompt.py index fc4b112..44cb3d9 100644 --- a/tests/unit/test_extract_final_prompt.py +++ b/tests/unit/test_extract_final_prompt.py @@ -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" diff --git a/tests/unit/test_ollama_client.py b/tests/unit/test_ollama_client.py index 846728f..8107fd1 100644 --- a/tests/unit/test_ollama_client.py +++ b/tests/unit/test_ollama_client.py @@ -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 diff --git a/tests/unit/test_prompt_combiner.py b/tests/unit/test_prompt_combiner.py new file mode 100644 index 0000000..930f5a4 --- /dev/null +++ b/tests/unit/test_prompt_combiner.py @@ -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 diff --git a/tests/unit/test_style_presets.py b/tests/unit/test_style_presets.py index 31832d9..a3fda26 100644 --- a/tests/unit/test_style_presets.py +++ b/tests/unit/test_style_presets.py @@ -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