|
|
|
@@ -1,7 +1,7 @@
|
|
|
|
|
"""Tests for Sampler Combo node."""
|
|
|
|
|
|
|
|
|
|
import pytest
|
|
|
|
|
from unittest.mock import patch, MagicMock
|
|
|
|
|
from unittest.mock import patch
|
|
|
|
|
from kikotools.tools.sampler_combo.node import SamplerComboNode
|
|
|
|
|
from kikotools.tools.sampler_combo.logic import (
|
|
|
|
|
validate_sampler_settings,
|
|
|
|
@@ -47,7 +47,7 @@ class TestSamplerComboLogic:
|
|
|
|
|
"""Test getting sampler combo with valid inputs."""
|
|
|
|
|
result = get_sampler_combo("euler", "normal", 20, 7.0)
|
|
|
|
|
assert result == ("euler", "normal", 20, 7.0)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
result = get_sampler_combo("dpmpp_2m", "karras", 25, 8.5)
|
|
|
|
|
assert result == ("dpmpp_2m", "karras", 25, 8.5)
|
|
|
|
|
|
|
|
|
@@ -55,7 +55,7 @@ class TestSamplerComboLogic:
|
|
|
|
|
"""Test that invalid inputs return safe defaults."""
|
|
|
|
|
result = get_sampler_combo("invalid", "normal", 20, 7.0)
|
|
|
|
|
assert result == ("euler", "normal", 20, 7.0)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
result = get_sampler_combo("euler", "invalid", 20, 7.0)
|
|
|
|
|
assert result == ("euler", "normal", 20, 7.0)
|
|
|
|
|
|
|
|
|
@@ -64,14 +64,14 @@ class TestSamplerComboLogic:
|
|
|
|
|
# Test steps clamping
|
|
|
|
|
result = get_sampler_combo("euler", "normal", 0, 7.0)
|
|
|
|
|
assert result[2] >= 1 # steps should be at least 1
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
result = get_sampler_combo("euler", "normal", 1500, 7.0)
|
|
|
|
|
assert result[2] <= 1000 # steps should be at most 1000
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# Test CFG clamping
|
|
|
|
|
result = get_sampler_combo("euler", "normal", 20, -5.0)
|
|
|
|
|
assert result[3] >= 0.0 # CFG should be at least 0
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
result = get_sampler_combo("euler", "normal", 20, 50.0)
|
|
|
|
|
assert result[3] <= 30.0 # CFG should be at most 30
|
|
|
|
|
|
|
|
|
@@ -81,10 +81,10 @@ class TestSamplerComboLogic:
|
|
|
|
|
assert isinstance(suggestions, list)
|
|
|
|
|
assert len(suggestions) > 0
|
|
|
|
|
assert "normal" in suggestions
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
suggestions = get_compatible_scheduler_suggestions("ddim")
|
|
|
|
|
assert "ddim_uniform" in suggestions
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# Test unknown sampler returns defaults
|
|
|
|
|
suggestions = get_compatible_scheduler_suggestions("unknown_sampler")
|
|
|
|
|
assert "normal" in suggestions
|
|
|
|
@@ -98,7 +98,7 @@ class TestSamplerComboLogic:
|
|
|
|
|
assert isinstance(default_steps, int)
|
|
|
|
|
assert min_steps <= default_steps <= max_steps
|
|
|
|
|
assert min_steps > 0
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# Test unknown sampler returns defaults
|
|
|
|
|
min_steps, max_steps, default_steps = get_recommended_steps_range("unknown")
|
|
|
|
|
assert min_steps == 10
|
|
|
|
@@ -113,7 +113,7 @@ class TestSamplerComboLogic:
|
|
|
|
|
assert isinstance(default_cfg, float)
|
|
|
|
|
assert min_cfg <= default_cfg <= max_cfg
|
|
|
|
|
assert min_cfg >= 0.0
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# Test unknown sampler returns defaults
|
|
|
|
|
min_cfg, max_cfg, default_cfg = get_recommended_cfg_range("unknown")
|
|
|
|
|
assert min_cfg == 1.0
|
|
|
|
@@ -142,34 +142,34 @@ class TestSamplerComboNode:
|
|
|
|
|
def test_input_types_structure(self):
|
|
|
|
|
"""Test that INPUT_TYPES returns correct structure."""
|
|
|
|
|
input_types = SamplerComboNode.INPUT_TYPES()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
assert "required" in input_types
|
|
|
|
|
required = input_types["required"]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# Check all required inputs are present
|
|
|
|
|
assert "sampler_name" in required
|
|
|
|
|
assert "scheduler" in required
|
|
|
|
|
assert "steps" in required
|
|
|
|
|
assert "cfg" in required
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# Check sampler input structure
|
|
|
|
|
sampler_input = required["sampler_name"]
|
|
|
|
|
assert sampler_input[0] == SAMPLERS
|
|
|
|
|
assert isinstance(sampler_input[1], dict)
|
|
|
|
|
assert "default" in sampler_input[1]
|
|
|
|
|
assert "tooltip" in sampler_input[1]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# Check scheduler input structure
|
|
|
|
|
scheduler_input = required["scheduler"]
|
|
|
|
|
assert scheduler_input[0] == SCHEDULERS
|
|
|
|
|
assert isinstance(scheduler_input[1], dict)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# Check steps input structure
|
|
|
|
|
steps_input = required["steps"]
|
|
|
|
|
assert steps_input[0] == "INT"
|
|
|
|
|
assert steps_input[1]["min"] == 1
|
|
|
|
|
assert steps_input[1]["max"] == 1000
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# Check CFG input structure
|
|
|
|
|
cfg_input = required["cfg"]
|
|
|
|
|
assert cfg_input[0] == "FLOAT"
|
|
|
|
@@ -179,7 +179,12 @@ class TestSamplerComboNode:
|
|
|
|
|
def test_return_types_structure(self):
|
|
|
|
|
"""Test that return types are correctly defined."""
|
|
|
|
|
assert SamplerComboNode.RETURN_TYPES == (SAMPLERS, SCHEDULERS, "INT", "FLOAT")
|
|
|
|
|
assert SamplerComboNode.RETURN_NAMES == ("sampler_name", "scheduler", "steps", "cfg")
|
|
|
|
|
assert SamplerComboNode.RETURN_NAMES == (
|
|
|
|
|
"sampler_name",
|
|
|
|
|
"scheduler",
|
|
|
|
|
"steps",
|
|
|
|
|
"cfg",
|
|
|
|
|
)
|
|
|
|
|
assert SamplerComboNode.FUNCTION == "get_sampler_combo"
|
|
|
|
|
assert SamplerComboNode.CATEGORY == "ComfyAssets"
|
|
|
|
|
|
|
|
|
@@ -187,22 +192,25 @@ class TestSamplerComboNode:
|
|
|
|
|
"""Test get_sampler_combo with valid inputs."""
|
|
|
|
|
result = self.node.get_sampler_combo("euler", "normal", 20, 7.0)
|
|
|
|
|
assert result == ("euler", "normal", 20, 7.0)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
result = self.node.get_sampler_combo("dpmpp_2m", "karras", 15, 8.5)
|
|
|
|
|
assert result == ("dpmpp_2m", "karras", 15, 8.5)
|
|
|
|
|
|
|
|
|
|
def test_get_sampler_combo_invalid_inputs_returns_defaults(self):
|
|
|
|
|
"""Test that invalid inputs return safe defaults."""
|
|
|
|
|
with patch.object(self.node, 'handle_error') as mock_error:
|
|
|
|
|
with patch.object(self.node, "handle_error") as mock_error:
|
|
|
|
|
mock_error.side_effect = ValueError("Invalid settings")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
try:
|
|
|
|
|
result = self.node.get_sampler_combo("invalid", "normal", 20, 7.0)
|
|
|
|
|
except ValueError:
|
|
|
|
|
pass # Expected when handle_error raises
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# Test with exception handling bypassed
|
|
|
|
|
with patch('kikotools.tools.sampler_combo.node.validate_sampler_settings', return_value=False):
|
|
|
|
|
with patch(
|
|
|
|
|
"kikotools.tools.sampler_combo.node.validate_sampler_settings",
|
|
|
|
|
return_value=False,
|
|
|
|
|
):
|
|
|
|
|
result = self.node.get_sampler_combo("invalid", "normal", 20, 7.0)
|
|
|
|
|
assert result == ("euler", "normal", 20, 7.0)
|
|
|
|
|
|
|
|
|
@@ -221,7 +229,7 @@ class TestSamplerComboNode:
|
|
|
|
|
suggestions = self.node.get_scheduler_suggestions("euler")
|
|
|
|
|
assert isinstance(suggestions, list)
|
|
|
|
|
assert len(suggestions) > 0
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
suggestions = self.node.get_scheduler_suggestions("ddim")
|
|
|
|
|
assert "ddim_uniform" in suggestions
|
|
|
|
|
|
|
|
|
@@ -277,7 +285,7 @@ class TestSamplerComboNode:
|
|
|
|
|
assert "SamplerComboNode" in str_repr
|
|
|
|
|
assert "samplers=" in str_repr
|
|
|
|
|
assert "schedulers=" in str_repr
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
repr_str = repr(self.node)
|
|
|
|
|
assert "SamplerComboNode" in repr_str
|
|
|
|
|
assert "category=" in repr_str
|
|
|
|
@@ -286,10 +294,11 @@ class TestSamplerComboNode:
|
|
|
|
|
def test_node_inheritance(self):
|
|
|
|
|
"""Test that node properly inherits from base class."""
|
|
|
|
|
from kikotools.base.base_node import ComfyAssetsBaseNode
|
|
|
|
|
|
|
|
|
|
assert isinstance(self.node, ComfyAssetsBaseNode)
|
|
|
|
|
assert hasattr(self.node, 'validate_inputs')
|
|
|
|
|
assert hasattr(self.node, 'handle_error')
|
|
|
|
|
assert hasattr(self.node, 'log_info')
|
|
|
|
|
assert hasattr(self.node, "validate_inputs")
|
|
|
|
|
assert hasattr(self.node, "handle_error")
|
|
|
|
|
assert hasattr(self.node, "log_info")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
class TestSamplerComboIntegration:
|
|
|
|
@@ -298,7 +307,7 @@ class TestSamplerComboIntegration:
|
|
|
|
|
def test_full_workflow_valid_settings(self):
|
|
|
|
|
"""Test complete workflow with valid settings."""
|
|
|
|
|
node = SamplerComboNode()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# Test with different sampler/scheduler combinations
|
|
|
|
|
test_cases = [
|
|
|
|
|
("euler", "normal", 20, 7.0),
|
|
|
|
@@ -306,7 +315,7 @@ class TestSamplerComboIntegration:
|
|
|
|
|
("euler_ancestral", "exponential", 25, 9.0),
|
|
|
|
|
("ddim", "ddim_uniform", 30, 6.0),
|
|
|
|
|
]
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
for sampler, scheduler, steps, cfg in test_cases:
|
|
|
|
|
result = node.get_sampler_combo(sampler, scheduler, steps, cfg)
|
|
|
|
|
assert result == (sampler, scheduler, steps, cfg)
|
|
|
|
@@ -314,19 +323,16 @@ class TestSamplerComboIntegration:
|
|
|
|
|
def test_recommendation_compatibility(self):
|
|
|
|
|
"""Test that recommendations are compatible with actual functionality."""
|
|
|
|
|
node = SamplerComboNode()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
for sampler in SAMPLERS[:5]: # Test first 5 samplers
|
|
|
|
|
suggestions = node.get_scheduler_suggestions(sampler)
|
|
|
|
|
steps_rec = node.get_steps_recommendation(sampler)
|
|
|
|
|
cfg_rec = node.get_cfg_recommendation(sampler)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# Test that recommendations work with the node
|
|
|
|
|
for scheduler in suggestions[:2]: # Test first 2 suggestions
|
|
|
|
|
result = node.get_sampler_combo(
|
|
|
|
|
sampler,
|
|
|
|
|
scheduler,
|
|
|
|
|
steps_rec["default"],
|
|
|
|
|
cfg_rec["default"]
|
|
|
|
|
sampler, scheduler, steps_rec["default"], cfg_rec["default"]
|
|
|
|
|
)
|
|
|
|
|
assert result[0] == sampler
|
|
|
|
|
assert result[1] == scheduler
|
|
|
|
@@ -336,9 +342,11 @@ class TestSamplerComboIntegration:
|
|
|
|
|
def test_error_recovery(self):
|
|
|
|
|
"""Test error recovery with malformed inputs."""
|
|
|
|
|
node = SamplerComboNode()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
# These should all return safe defaults due to error handling
|
|
|
|
|
with patch('kikotools.tools.sampler_combo.logic.validate_sampler_settings',
|
|
|
|
|
side_effect=Exception("Simulated error")):
|
|
|
|
|
with patch(
|
|
|
|
|
"kikotools.tools.sampler_combo.logic.validate_sampler_settings",
|
|
|
|
|
side_effect=Exception("Simulated error"),
|
|
|
|
|
):
|
|
|
|
|
result = node.get_sampler_combo("euler", "normal", 20, 7.0)
|
|
|
|
|
assert result == ("euler", "normal", 20, 7.0) # Safe defaults
|
|
|
|
|
assert result == ("euler", "normal", 20, 7.0) # Safe defaults
|
|
|
|
|