test(batch_prompts): add unit tests for batch prompts

This commit is contained in:
Vito Sansevero
2025-08-27 09:32:09 -07:00
parent b37bc763dc
commit 5ae7985bc2
2 changed files with 626 additions and 0 deletions
+336
View File
@@ -0,0 +1,336 @@
"""Unit tests for Batch Prompts node."""
import pytest
import tempfile
import os
from pathlib import Path
from kikotools.tools.batch_prompts.logic import (
load_prompts_from_file,
get_prompt_at_index,
get_next_prompt,
get_prompt_preview,
get_batch_info,
validate_prompt_file,
format_prompt_for_display,
split_prompt_into_positive_negative,
create_batch_queue,
)
from kikotools.tools.batch_prompts.node import BatchPromptsNode
class TestBatchPromptsLogic:
"""Test batch prompts logic functions."""
def test_load_prompts_from_file(self, tmp_path):
"""Test loading prompts from a file with --- separators."""
# Create test file
test_file = tmp_path / "test_prompts.txt"
test_content = """First prompt here
with multiple lines
---
Second prompt
also multiline
---
Third prompt"""
test_file.write_text(test_content)
# Load prompts
prompts = load_prompts_from_file(str(test_file))
assert len(prompts) == 3
assert "First prompt here\nwith multiple lines" in prompts[0]
assert "Second prompt\nalso multiline" in prompts[1]
assert "Third prompt" in prompts[2]
def test_load_prompts_empty_sections(self, tmp_path):
"""Test loading prompts with empty sections."""
test_file = tmp_path / "test_prompts.txt"
test_content = """First prompt
---
---
Second prompt
---
"""
test_file.write_text(test_content)
prompts = load_prompts_from_file(str(test_file))
# Should only get non-empty prompts
assert len(prompts) == 2
assert "First prompt" in prompts[0]
assert "Second prompt" in prompts[1]
def test_get_prompt_at_index(self):
"""Test getting prompt at specific index."""
prompts = ["Prompt 1", "Prompt 2", "Prompt 3"]
# Normal access
prompt, idx = get_prompt_at_index(prompts, 1, wrap=False)
assert prompt == "Prompt 2"
assert idx == 1
# With wrapping
prompt, idx = get_prompt_at_index(prompts, 4, wrap=True)
assert prompt == "Prompt 2" # 4 % 3 = 1
assert idx == 1
# Without wrapping, clamp to last
prompt, idx = get_prompt_at_index(prompts, 5, wrap=False)
assert prompt == "Prompt 3"
assert idx == 2
def test_get_next_prompt(self):
"""Test getting next prompt in sequence."""
prompts = ["Prompt 1", "Prompt 2", "Prompt 3"]
# Normal next
prompt, idx = get_next_prompt(prompts, 0, wrap=True)
assert prompt == "Prompt 2"
assert idx == 1
# Wrap around
prompt, idx = get_next_prompt(prompts, 2, wrap=True)
assert prompt == "Prompt 1"
assert idx == 0
# No wrap
prompt, idx = get_next_prompt(prompts, 2, wrap=False)
assert prompt == "Prompt 3"
assert idx == 2
def test_get_prompt_preview(self):
"""Test prompt preview truncation."""
short_prompt = "Short prompt"
long_prompt = "This is a very long prompt " * 10
# Short prompt unchanged
preview = get_prompt_preview(short_prompt, 100)
assert preview == short_prompt
# Long prompt truncated
preview = get_prompt_preview(long_prompt, 50)
assert len(preview) == 53 # 50 + "..."
assert preview.endswith("...")
def test_split_prompt_positive_negative(self):
"""Test splitting prompts into positive and negative."""
# With negative
prompt = "Beautiful landscape\nNegative: blurry, dark"
pos, neg = split_prompt_into_positive_negative(prompt)
assert pos == "Beautiful landscape"
assert neg == "blurry, dark"
# Without negative
prompt = "Just a positive prompt"
pos, neg = split_prompt_into_positive_negative(prompt)
assert pos == "Just a positive prompt"
assert neg == ""
# Case insensitive
prompt = "Positive part\nnegative: negative part"
pos, neg = split_prompt_into_positive_negative(prompt)
assert pos == "Positive part"
assert neg == "negative part"
def test_get_batch_info(self):
"""Test batch information generation."""
prompts = ["P1", "P2", "P3", "P4", "P5"]
info = get_batch_info(prompts, 2)
assert info["current_index"] == 2
assert info["total_prompts"] == 5
assert info["progress"] == "3/5"
assert info["percentage"] == 40.0
assert info["remaining"] == 2
assert info["is_complete"] == False
# Last prompt
info = get_batch_info(prompts, 4)
assert info["is_complete"] == True
assert info["remaining"] == 0
def test_validate_prompt_file(self, tmp_path):
"""Test prompt file validation."""
# Valid file
valid_file = tmp_path / "valid.txt"
valid_file.write_text("content")
is_valid, error = validate_prompt_file(str(valid_file))
assert is_valid
assert error == ""
# Non-existent file
is_valid, error = validate_prompt_file("/nonexistent/file.txt")
assert not is_valid
assert "not found" in error
# Empty path
is_valid, error = validate_prompt_file("")
assert not is_valid
assert "No file path" in error
def test_format_prompt_for_display(self):
"""Test prompt display formatting."""
prompt = "Test prompt"
formatted = format_prompt_for_display(prompt, 2, 5)
assert "[Prompt 3/5]" in formatted
assert "Test prompt" in formatted
assert "---" in formatted
def test_create_batch_queue(self):
"""Test batch queue creation."""
prompts = ["P1", "P2", "P3", "P4", "P5"]
# Batch size 2
batches = create_batch_queue(prompts, batch_size=2, randomize=False)
assert len(batches) == 3
assert batches[0] == [0, 1]
assert batches[1] == [2, 3]
assert batches[2] == [4]
# Batch size 1
batches = create_batch_queue(prompts, batch_size=1, randomize=False)
assert len(batches) == 5
assert all(len(b) == 1 for b in batches)
class TestBatchPromptsNode:
"""Test BatchPromptsNode class."""
def test_node_input_types(self):
"""Test node input type definitions."""
input_types = BatchPromptsNode.INPUT_TYPES()
assert "required" in input_types
assert "prompt_file" in input_types["required"]
assert "index" in input_types["required"]
assert "auto_increment" in input_types["required"]
assert "wrap_around" in input_types["required"]
assert "split_negative" in input_types["required"]
assert "optional" in input_types
assert "reload_file" in input_types["optional"]
assert "show_preview" in input_types["optional"]
def test_node_return_types(self):
"""Test node return type definitions."""
assert BatchPromptsNode.RETURN_TYPES == (
"STRING",
"STRING",
"STRING",
"STRING",
"INT",
"INT",
"STRING",
)
assert BatchPromptsNode.RETURN_NAMES == (
"positive",
"negative",
"full_prompt",
"next_prompt",
"current_index",
"total_prompts",
"batch_info",
)
assert BatchPromptsNode.FUNCTION == "process_batch_prompts"
assert "ComfyAssets" in BatchPromptsNode.CATEGORY
def test_process_batch_prompts(self, tmp_path):
"""Test processing batch prompts."""
# Create test file
test_file = tmp_path / "test_prompts.txt"
test_content = """Beautiful sunset
Negative: dark, blurry
---
Mountain landscape
Negative: fog, rain
---
Ocean view"""
test_file.write_text(test_content)
node = BatchPromptsNode()
# Process first prompt
result = node.process_batch_prompts(
prompt_file=str(test_file),
index=0,
auto_increment=False,
wrap_around=True,
split_negative=True,
reload_file=False,
show_preview=False,
)
positive, negative, full, next_prompt, idx, total, info = result
assert positive == "Beautiful sunset"
assert negative == "dark, blurry"
assert "Beautiful sunset" in full
assert "Mountain landscape" in next_prompt
assert idx == 0
assert total == 3
assert "1 of 3" in info
def test_process_without_negative_split(self, tmp_path):
"""Test processing without splitting negative prompts."""
test_file = tmp_path / "test_prompts.txt"
test_content = """Full prompt with Negative: included"""
test_file.write_text(test_content)
node = BatchPromptsNode()
result = node.process_batch_prompts(
prompt_file=str(test_file),
index=0,
auto_increment=False,
wrap_around=True,
split_negative=False,
reload_file=False,
show_preview=False,
)
positive, negative, full, _, _, _, _ = result
assert positive == "Full prompt with Negative: included"
assert negative == ""
def test_wrap_around_behavior(self, tmp_path):
"""Test wrap around behavior."""
test_file = tmp_path / "test_prompts.txt"
test_content = """Prompt 1
---
Prompt 2"""
test_file.write_text(test_content)
node = BatchPromptsNode()
# Test with wrap
result = node.process_batch_prompts(
prompt_file=str(test_file),
index=2, # Beyond end
auto_increment=False,
wrap_around=True,
split_negative=False,
reload_file=False,
show_preview=False,
)
positive, _, _, _, idx, _, _ = result
assert positive == "Prompt 1" # Wrapped to index 0
assert idx == 0
# Test without wrap
result = node.process_batch_prompts(
prompt_file=str(test_file),
index=2, # Beyond end
auto_increment=False,
wrap_around=False,
split_negative=False,
reload_file=True, # Force reload
show_preview=False,
)
positive, _, _, _, idx, _, _ = result
assert positive == "Prompt 2" # Clamped to last
assert idx == 1
+290
View File
@@ -0,0 +1,290 @@
"""Unit tests for Local Image Loader tool."""
import json
import os
import tempfile
from pathlib import Path
from unittest.mock import patch
import pytest
import torch
from PIL import Image, PngImagePlugin
from kikotools.tools.local_image_loader.logic import (
create_empty_tensor,
get_supported_extensions,
load_image_from_path,
scan_directory,
)
from kikotools.tools.local_image_loader.node import LocalImageLoaderNode
class TestLocalImageLoaderLogic:
"""Test the logic functions for local image loader."""
def test_get_supported_extensions(self):
"""Test getting supported file extensions."""
extensions = get_supported_extensions()
assert "image" in extensions
assert "video" in extensions
assert "audio" in extensions
assert ".jpg" in extensions["image"]
assert ".png" in extensions["image"]
assert ".mp4" in extensions["video"]
assert ".mp3" in extensions["audio"]
def test_create_empty_tensor(self):
"""Test creating an empty tensor."""
tensor = create_empty_tensor()
assert isinstance(tensor, torch.Tensor)
assert tensor.shape == (1, 1, 1, 4)
assert torch.all(tensor == 0)
def test_load_image_from_path_rgb(self):
"""Test loading an RGB image from file."""
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp:
# Create a test image
img = Image.new("RGB", (100, 100), color="red")
img.save(tmp.name)
try:
tensor, metadata = load_image_from_path(tmp.name)
# Check tensor
assert isinstance(tensor, torch.Tensor)
assert tensor.shape == (1, 100, 100, 3)
assert tensor.min() >= 0.0
assert tensor.max() <= 1.0
# Check metadata
assert metadata["width"] == 100
assert metadata["height"] == 100
assert metadata["filename"] == os.path.basename(tmp.name)
assert "mode" in metadata
assert "format" in metadata
finally:
os.unlink(tmp.name)
def test_load_image_from_path_rgba(self):
"""Test loading an RGBA image from file."""
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp:
# Create a test image with alpha
img = Image.new("RGBA", (50, 50), color=(255, 0, 0, 128))
img.save(tmp.name)
try:
tensor, metadata = load_image_from_path(tmp.name)
# Check tensor
assert isinstance(tensor, torch.Tensor)
assert tensor.shape == (1, 50, 50, 4) # RGBA has 4 channels
assert tensor.min() >= 0.0
assert tensor.max() <= 1.0
# Check metadata
assert metadata["width"] == 50
assert metadata["height"] == 50
finally:
os.unlink(tmp.name)
def test_load_image_from_path_with_metadata(self):
"""Test loading an image with embedded metadata."""
with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp:
# Create image with metadata
img = Image.new("RGB", (100, 100), color="blue")
# Add some metadata
metadata_to_save = {
"parameters": "test parameters",
"prompt": json.dumps({"text": "test prompt"}),
"workflow": json.dumps({"nodes": []}),
}
pnginfo = PngImagePlugin.PngInfo()
for key, value in metadata_to_save.items():
pnginfo.add_text(key, value)
img.save(tmp.name, pnginfo=pnginfo)
try:
tensor, metadata = load_image_from_path(tmp.name)
# Check embedded metadata
assert metadata.get("parameters") == "test parameters"
assert metadata.get("prompt") == {"text": "test prompt"}
assert metadata.get("workflow") == {"nodes": []}
finally:
os.unlink(tmp.name)
def test_load_image_from_nonexistent_path(self):
"""Test loading image from nonexistent path raises error."""
with pytest.raises(FileNotFoundError):
load_image_from_path("/nonexistent/path/image.png")
def test_scan_directory_images_only(self):
"""Test scanning directory for images only."""
with tempfile.TemporaryDirectory() as tmpdir:
# Create test files
Path(tmpdir, "image1.jpg").touch()
Path(tmpdir, "image2.png").touch()
Path(tmpdir, "video.mp4").touch()
Path(tmpdir, "audio.mp3").touch()
Path(tmpdir, "document.txt").touch()
Path(tmpdir, "subdir").mkdir()
items = scan_directory(tmpdir, show_videos=False, show_audio=False)
# Should have 1 directory and 2 images
assert len(items) == 3
# Check types
types = [item["type"] for item in items]
assert "dir" in types
assert types.count("image") == 2
def test_scan_directory_with_videos_audio(self):
"""Test scanning directory with videos and audio enabled."""
with tempfile.TemporaryDirectory() as tmpdir:
# Create test files
Path(tmpdir, "image.jpg").touch()
Path(tmpdir, "video.mp4").touch()
Path(tmpdir, "audio.mp3").touch()
items = scan_directory(tmpdir, show_videos=True, show_audio=True)
assert len(items) == 3
types = [item["type"] for item in items]
assert "image" in types
assert "video" in types
assert "audio" in types
def test_scan_directory_sorting(self):
"""Test directory scanning with different sort options."""
with tempfile.TemporaryDirectory() as tmpdir:
# Create files with different names
Path(tmpdir, "zebra.jpg").touch()
Path(tmpdir, "apple.jpg").touch()
Path(tmpdir, "banana.jpg").touch()
# Sort by name ascending
items = scan_directory(tmpdir, sort_by="name", sort_order="asc")
names = [item["name"] for item in items if item["type"] == "image"]
assert names == ["apple.jpg", "banana.jpg", "zebra.jpg"]
# Sort by name descending
items = scan_directory(tmpdir, sort_by="name", sort_order="desc")
names = [item["name"] for item in items if item["type"] == "image"]
assert names == ["zebra.jpg", "banana.jpg", "apple.jpg"]
def test_scan_nonexistent_directory(self):
"""Test scanning nonexistent directory raises error."""
with pytest.raises(NotADirectoryError):
scan_directory("/nonexistent/directory")
class TestLocalImageLoaderNode:
"""Test the Local Image Loader node."""
def test_input_types(self):
"""Test node input types definition."""
input_types = LocalImageLoaderNode.INPUT_TYPES()
assert "required" in input_types
assert "hidden" in input_types
assert "unique_id" in input_types["hidden"]
def test_node_properties(self):
"""Test node properties."""
assert LocalImageLoaderNode.RETURN_TYPES == (
"IMAGE",
"STRING",
"STRING",
"STRING",
)
assert LocalImageLoaderNode.RETURN_NAMES == (
"image",
"video_path",
"audio_path",
"info",
)
assert LocalImageLoaderNode.FUNCTION == "load_media"
assert LocalImageLoaderNode.CATEGORY == "🫶 ComfyAssets/💾 Images"
@patch("kikotools.tools.local_image_loader.node.load_selections")
def test_load_media_no_selection(self, mock_load_selections):
"""Test loading media with no selection returns empty values."""
mock_load_selections.return_value = {}
node = LocalImageLoaderNode()
image, video_path, audio_path, info = node.load_media("test_id")
# Check empty returns
assert isinstance(image, torch.Tensor)
assert image.shape == (1, 1, 1, 4)
assert torch.all(image == 0)
assert video_path == ""
assert audio_path == ""
assert info == ""
@patch("kikotools.tools.local_image_loader.node.load_selections")
@patch("kikotools.tools.local_image_loader.node.load_image_from_path")
def test_load_media_with_image_selection(
self, mock_load_image, mock_load_selections
):
"""Test loading media with image selection."""
# Setup mocks
mock_load_selections.return_value = {
"test_id": {"image": {"path": "/path/to/image.jpg"}}
}
test_tensor = torch.ones(1, 100, 100, 3)
test_metadata = {"width": 100, "height": 100, "filename": "image.jpg"}
mock_load_image.return_value = (test_tensor, test_metadata)
# Mock os.path.exists
with patch("os.path.exists", return_value=True):
node = LocalImageLoaderNode()
image, video_path, audio_path, info = node.load_media("test_id")
# Check returns
assert torch.equal(image, test_tensor)
assert video_path == ""
assert audio_path == ""
assert json.loads(info) == test_metadata
@patch("kikotools.tools.local_image_loader.node.load_selections")
def test_load_media_with_video_audio_selection(self, mock_load_selections):
"""Test loading media with video and audio selection."""
mock_load_selections.return_value = {
"test_id": {
"video": {"path": "/path/to/video.mp4"},
"audio": {"path": "/path/to/audio.mp3"},
}
}
with patch("os.path.exists", return_value=True):
node = LocalImageLoaderNode()
image, video_path, audio_path, info = node.load_media("test_id")
# Check returns
assert isinstance(image, torch.Tensor)
assert image.shape == (1, 1, 1, 4) # Empty tensor
assert video_path == "/path/to/video.mp4"
assert audio_path == "/path/to/audio.mp3"
assert info == ""
def test_is_changed(self):
"""Test IS_CHANGED method."""
with patch("os.path.exists", return_value=False):
result = LocalImageLoaderNode.IS_CHANGED()
assert result == float("inf")
with (
patch("os.path.exists", return_value=True),
patch("os.path.getmtime", return_value=12345.0),
):
result = LocalImageLoaderNode.IS_CHANGED()
assert result == 12345.0