From 5ae7985bc2c96eb403f1e21361029d96875f8f59 Mon Sep 17 00:00:00 2001 From: Vito Sansevero Date: Wed, 27 Aug 2025 09:32:09 -0700 Subject: [PATCH] test(batch_prompts): add unit tests for batch prompts --- tests/unit/tools/test_batch_prompts.py | 336 ++++++++++++++++++++ tests/unit/tools/test_local_image_loader.py | 290 +++++++++++++++++ 2 files changed, 626 insertions(+) create mode 100644 tests/unit/tools/test_batch_prompts.py create mode 100644 tests/unit/tools/test_local_image_loader.py diff --git a/tests/unit/tools/test_batch_prompts.py b/tests/unit/tools/test_batch_prompts.py new file mode 100644 index 0000000..5081de2 --- /dev/null +++ b/tests/unit/tools/test_batch_prompts.py @@ -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 diff --git a/tests/unit/tools/test_local_image_loader.py b/tests/unit/tools/test_local_image_loader.py new file mode 100644 index 0000000..02999f3 --- /dev/null +++ b/tests/unit/tools/test_local_image_loader.py @@ -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