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:
co-authored by
Claude Opus 4.7
parent
8e19db8e72
commit
839188382c
@@ -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:
|
||||
|
||||
@@ -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}",)
|
||||
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
@@ -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",
|
||||
]
|
||||
|
||||
@@ -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,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
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user