test(batch_prompts): add unit tests for batch prompts
This commit is contained in:
@@ -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
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user