diff --git a/README.md b/README.md index 1f619e9..c2dd2d7 100644 --- a/README.md +++ b/README.md @@ -1,5 +1,11 @@ # ComfyUI Prompt Manager +![License](https://img.shields.io/badge/license-MIT-blue.svg) +![Python](https://img.shields.io/badge/python-3.10%2B-blue.svg) +![ComfyUI](https://img.shields.io/badge/ComfyUI-compatible-green.svg) +![Tests](https://github.com/ComfyAssets/ComfyUI_PromptManager/workflows/Tests/badge.svg) +![Code Quality](https://github.com/ComfyAssets/ComfyUI_PromptManager/workflows/Code%20Quality/badge.svg) + A comprehensive ComfyUI custom node that extends the standard text encoder with persistent prompt storage, advanced search capabilities, automatic image gallery system, and powerful ComfyUI workflow metadata analysis using SQLite. ## 📋 A Note on v2 Development diff --git a/tests/test_config.py b/tests/test_config.py new file mode 100644 index 0000000..05037ce --- /dev/null +++ b/tests/test_config.py @@ -0,0 +1,211 @@ +""" +Tests for GalleryConfig and PromptManagerConfig. + +The config module imports `from server import PromptServer` at module level, +which is ComfyUI's server. We mock it before import to run tests standalone. +""" + +import json +import os +import sys +import tempfile +import unittest +from unittest.mock import MagicMock + +# Mock ComfyUI's server module before importing config +_mock_server = MagicMock() +_mock_server.PromptServer.instance.routes = MagicMock() +sys.modules["server"] = _mock_server + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from py.config import GalleryConfig, PromptManagerConfig + + +class TestGalleryConfig(unittest.TestCase): + """Test GalleryConfig class.""" + + def setUp(self): + """Save original values to restore after each test.""" + self._orig = GalleryConfig.get_config() + + def tearDown(self): + """Restore original config values.""" + GalleryConfig.update_config(self._orig) + + def test_get_config_structure(self): + config = GalleryConfig.get_config() + self.assertIn("monitoring", config) + self.assertIn("tracking", config) + self.assertIn("database", config) + self.assertIn("web_interface", config) + self.assertIn("performance", config) + + def test_default_values(self): + config = GalleryConfig.get_config() + self.assertTrue(config["monitoring"]["enabled"]) + self.assertEqual( + config["monitoring"]["extensions"], + [".png", ".jpg", ".jpeg", ".webp", ".gif"], + ) + self.assertEqual(config["tracking"]["prompt_timeout"], 600) + self.assertEqual(config["web_interface"]["images_per_page"], 20) + + def test_update_monitoring(self): + GalleryConfig.update_config( + {"monitoring": {"enabled": False, "processing_delay": 5.0}} + ) + self.assertFalse(GalleryConfig.MONITORING_ENABLED) + self.assertEqual(GalleryConfig.PROCESSING_DELAY, 5.0) + + def test_update_tracking(self): + GalleryConfig.update_config({"tracking": {"prompt_timeout": 300}}) + self.assertEqual(GalleryConfig.PROMPT_TIMEOUT, 300) + + def test_update_database(self): + GalleryConfig.update_config( + {"database": {"auto_cleanup": False, "max_image_age_days": 30}} + ) + self.assertFalse(GalleryConfig.AUTO_CLEANUP_MISSING_FILES) + self.assertEqual(GalleryConfig.MAX_IMAGE_AGE_DAYS, 30) + + def test_update_web_interface(self): + GalleryConfig.update_config( + {"web_interface": {"images_per_page": 50, "thumbnail_size": 512}} + ) + self.assertEqual(GalleryConfig.IMAGES_PER_PAGE, 50) + self.assertEqual(GalleryConfig.THUMBNAIL_SIZE, 512) + + def test_update_performance(self): + GalleryConfig.update_config({"performance": {"max_concurrent_processing": 5}}) + self.assertEqual(GalleryConfig.MAX_CONCURRENT_PROCESSING, 5) + + def test_partial_update_preserves_other_values(self): + original_timeout = GalleryConfig.PROMPT_TIMEOUT + GalleryConfig.update_config({"monitoring": {"enabled": False}}) + self.assertEqual(GalleryConfig.PROMPT_TIMEOUT, original_timeout) + + def test_empty_update_changes_nothing(self): + before = GalleryConfig.get_config() + GalleryConfig.update_config({}) + after = GalleryConfig.get_config() + self.assertEqual(before, after) + + +class TestPromptManagerConfig(unittest.TestCase): + """Test PromptManagerConfig class.""" + + def setUp(self): + self._orig = PromptManagerConfig.get_config() + + def tearDown(self): + PromptManagerConfig.update_config(self._orig) + + def test_get_config_structure(self): + config = PromptManagerConfig.get_config() + self.assertIn("database", config) + self.assertIn("web_ui", config) + self.assertIn("performance", config) + self.assertIn("gallery", config) + + def test_gallery_config_nested(self): + config = PromptManagerConfig.get_config() + # Gallery section should match GalleryConfig output + self.assertIn("monitoring", config["gallery"]) + self.assertIn("tracking", config["gallery"]) + + def test_default_database_values(self): + config = PromptManagerConfig.get_config() + self.assertEqual(config["database"]["default_path"], "prompts.db") + self.assertTrue(config["database"]["enable_duplicate_detection"]) + self.assertTrue(config["database"]["enable_auto_save"]) + + def test_default_web_ui_values(self): + config = PromptManagerConfig.get_config() + self.assertEqual(config["web_ui"]["webui_display_mode"], "newtab") + + def test_update_database(self): + PromptManagerConfig.update_config({"database": {"default_path": "custom.db"}}) + self.assertEqual(PromptManagerConfig.DEFAULT_DB_PATH, "custom.db") + + def test_update_web_ui(self): + PromptManagerConfig.update_config({"web_ui": {"webui_display_mode": "popup"}}) + self.assertEqual(PromptManagerConfig.WEBUI_DISPLAY_MODE, "popup") + + def test_update_performance(self): + PromptManagerConfig.update_config({"performance": {"max_search_results": 50}}) + self.assertEqual(PromptManagerConfig.MAX_SEARCH_RESULTS, 50) + + def test_update_propagates_to_gallery(self): + PromptManagerConfig.update_config( + {"gallery": {"monitoring": {"enabled": False}}} + ) + self.assertFalse(GalleryConfig.MONITORING_ENABLED) + # Restore + GalleryConfig.MONITORING_ENABLED = True + + def test_save_and_load_file(self): + tmp = tempfile.NamedTemporaryFile( + delete=False, suffix=".json", dir=tempfile.gettempdir() + ) + tmp.close() + try: + # Modify a value + PromptManagerConfig.MAX_SEARCH_RESULTS = 42 + PromptManagerConfig.save_to_file(tmp.name) + + # Reset and reload + PromptManagerConfig.MAX_SEARCH_RESULTS = 100 + PromptManagerConfig.load_from_file(tmp.name) + self.assertEqual(PromptManagerConfig.MAX_SEARCH_RESULTS, 42) + finally: + os.unlink(tmp.name) + + def test_load_nonexistent_file_uses_defaults(self): + # Should not raise, just log and continue + original = PromptManagerConfig.MAX_SEARCH_RESULTS + PromptManagerConfig.load_from_file("/nonexistent/path/config.json") + self.assertEqual(PromptManagerConfig.MAX_SEARCH_RESULTS, original) + + def test_load_invalid_json_file(self): + tmp = tempfile.NamedTemporaryFile(delete=False, suffix=".json", mode="w") + tmp.write("not valid json {{") + tmp.close() + try: + original = PromptManagerConfig.MAX_SEARCH_RESULTS + PromptManagerConfig.load_from_file(tmp.name) + # Should not crash, keeps existing values + self.assertEqual(PromptManagerConfig.MAX_SEARCH_RESULTS, original) + finally: + os.unlink(tmp.name) + + def test_save_creates_directory(self): + tmp_dir = tempfile.mkdtemp() + nested_path = os.path.join(tmp_dir, "subdir", "config.json") + try: + PromptManagerConfig.save_to_file(nested_path) + self.assertTrue(os.path.exists(nested_path)) + with open(nested_path) as f: + data = json.load(f) + self.assertIn("database", data) + finally: + import shutil + + shutil.rmtree(tmp_dir) + + def test_saved_file_is_valid_json(self): + tmp = tempfile.NamedTemporaryFile( + delete=False, suffix=".json", dir=tempfile.gettempdir() + ) + tmp.close() + try: + PromptManagerConfig.save_to_file(tmp.name) + with open(tmp.name) as f: + data = json.load(f) + self.assertIsInstance(data, dict) + finally: + os.unlink(tmp.name) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_hashing.py b/tests/test_hashing.py new file mode 100644 index 0000000..ed631fa --- /dev/null +++ b/tests/test_hashing.py @@ -0,0 +1,148 @@ +""" +Tests for hashing utilities. + +Tests generate_content_hash and is_duplicate_prompt which are +currently untested. generate_prompt_hash is covered in test_basic.py +but we add edge cases here. +""" + +import os +import sys +import unittest + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from utils.hashing import ( + generate_prompt_hash, + generate_content_hash, + is_duplicate_prompt, +) + + +class TestGenerateContentHash(unittest.TestCase): + """Test generate_content_hash function.""" + + def test_basic_content(self): + content = { + "text": "A prompt", + "category": "nature", + "tags": ["sky"], + "workflow_name": "wf1", + } + h = generate_content_hash(content) + self.assertEqual(len(h), 64) # SHA256 hex digest + + def test_normalization_case_insensitive(self): + c1 = {"text": "Hello World", "category": "Test"} + c2 = {"text": "hello world", "category": "test"} + self.assertEqual(generate_content_hash(c1), generate_content_hash(c2)) + + def test_normalization_strips_whitespace(self): + c1 = {"text": " hello ", "category": " test "} + c2 = {"text": "hello", "category": "test"} + self.assertEqual(generate_content_hash(c1), generate_content_hash(c2)) + + def test_tag_order_independent(self): + c1 = {"text": "prompt", "tags": ["beta", "alpha"]} + c2 = {"text": "prompt", "tags": ["alpha", "beta"]} + self.assertEqual(generate_content_hash(c1), generate_content_hash(c2)) + + def test_different_text_different_hash(self): + c1 = {"text": "prompt A"} + c2 = {"text": "prompt B"} + self.assertNotEqual(generate_content_hash(c1), generate_content_hash(c2)) + + def test_different_category_different_hash(self): + c1 = {"text": "same", "category": "cat1"} + c2 = {"text": "same", "category": "cat2"} + self.assertNotEqual(generate_content_hash(c1), generate_content_hash(c2)) + + def test_different_tags_different_hash(self): + c1 = {"text": "same", "tags": ["tag1"]} + c2 = {"text": "same", "tags": ["tag2"]} + self.assertNotEqual(generate_content_hash(c1), generate_content_hash(c2)) + + def test_empty_dict(self): + h = generate_content_hash({}) + self.assertEqual(len(h), 64) + + def test_missing_fields_use_defaults(self): + c1 = {"text": "hello"} + c2 = {"text": "hello", "category": None, "tags": [], "workflow_name": None} + self.assertEqual(generate_content_hash(c1), generate_content_hash(c2)) + + def test_none_category_same_as_empty(self): + c1 = {"text": "test", "category": None} + c2 = {"text": "test", "category": ""} + self.assertEqual(generate_content_hash(c1), generate_content_hash(c2)) + + def test_none_workflow_same_as_empty(self): + c1 = {"text": "test", "workflow_name": None} + c2 = {"text": "test", "workflow_name": ""} + self.assertEqual(generate_content_hash(c1), generate_content_hash(c2)) + + def test_empty_tags_filtered(self): + c1 = {"text": "test", "tags": ["valid", "", " "]} + c2 = {"text": "test", "tags": ["valid"]} + self.assertEqual(generate_content_hash(c1), generate_content_hash(c2)) + + def test_deterministic(self): + content = {"text": "stable", "tags": ["a", "b"]} + h1 = generate_content_hash(content) + h2 = generate_content_hash(content) + self.assertEqual(h1, h2) + + +class TestIsDuplicatePrompt(unittest.TestCase): + """Test is_duplicate_prompt function.""" + + def test_identical_texts_are_duplicates(self): + self.assertTrue(is_duplicate_prompt("hello", "hello")) + + def test_case_insensitive(self): + self.assertTrue(is_duplicate_prompt("Hello World", "hello world")) + + def test_whitespace_insensitive(self): + self.assertTrue(is_duplicate_prompt(" hello ", "hello")) + + def test_different_texts_not_duplicates(self): + self.assertFalse(is_duplicate_prompt("text A", "text B")) + + def test_threshold_parameter_accepted(self): + # threshold is unused but should not cause errors + self.assertTrue(is_duplicate_prompt("same", "same", threshold=0.5)) + + def test_empty_strings_are_duplicates(self): + self.assertTrue(is_duplicate_prompt("", "")) + + def test_unicode_comparison(self): + self.assertTrue(is_duplicate_prompt("日本語", "日本語")) + self.assertFalse(is_duplicate_prompt("日本語", "中文")) + + +class TestGeneratePromptHashEdgeCases(unittest.TestCase): + """Additional edge cases for generate_prompt_hash.""" + + def test_non_string_raises_type_error(self): + with self.assertRaises(TypeError): + generate_prompt_hash(123) + + def test_none_raises_type_error(self): + with self.assertRaises(TypeError): + generate_prompt_hash(None) + + def test_empty_string(self): + h = generate_prompt_hash("") + self.assertEqual(len(h), 64) + + def test_unicode_hash(self): + h = generate_prompt_hash("émojis 🎨 日本語") + self.assertEqual(len(h), 64) + + def test_very_long_text(self): + h = generate_prompt_hash("x" * 100000) + self.assertEqual(len(h), 64) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_metadata_extraction.py b/tests/test_metadata_extraction.py new file mode 100644 index 0000000..ec7da95 --- /dev/null +++ b/tests/test_metadata_extraction.py @@ -0,0 +1,374 @@ +""" +Tests for ComfyUI metadata extraction. + +Tests all methods of ComfyUIMetadataExtractor using real minimal +PNG files with injected text chunks to simulate ComfyUI output. +""" + +import json +import os +import sys +import tempfile +import unittest + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from PIL import Image +from PIL.PngImagePlugin import PngInfo + +from utils.metadata_extractor import ComfyUIMetadataExtractor + + +def _create_png(text_chunks=None, size=(64, 64)): + """Create a minimal PNG file with optional text chunks. + + Args: + text_chunks: dict of key->value to embed as PNG text metadata. + size: tuple (width, height). + + Returns: + Path to the temporary PNG file. + """ + tmp = tempfile.NamedTemporaryFile(delete=False, suffix=".png") + tmp.close() + + img = Image.new("RGB", size, color=(128, 128, 128)) + pnginfo = PngInfo() + if text_chunks: + for key, value in text_chunks.items(): + pnginfo.add_text(key, value) + img.save(tmp.name, pnginfo=pnginfo) + return tmp.name + + +class MetadataTestCase(unittest.TestCase): + """Base class that cleans up temp files.""" + + def setUp(self): + self.extractor = ComfyUIMetadataExtractor() + self.temp_files = [] + + def tearDown(self): + for f in self.temp_files: + if os.path.exists(f): + os.unlink(f) + + def _png(self, text_chunks=None, size=(64, 64)): + path = _create_png(text_chunks, size) + self.temp_files.append(path) + return path + + +class TestExtractMetadata(MetadataTestCase): + """Test extract_metadata method.""" + + def test_returns_none_for_plain_png(self): + path = self._png() + result = self.extractor.extract_metadata(path) + self.assertIsNone(result) + + def test_extracts_workflow_data(self): + workflow = {"nodes": [{"type": "CLIPTextEncode", "inputs": {"text": "hello"}}]} + path = self._png({"workflow": json.dumps(workflow)}) + result = self.extractor.extract_metadata(path) + self.assertIsNotNone(result) + self.assertIn("workflow", result) + self.assertEqual(result["workflow"]["nodes"][0]["type"], "CLIPTextEncode") + + def test_extracts_prompt_data(self): + prompt = {"1": {"class_type": "KSampler", "inputs": {}}} + path = self._png({"prompt": json.dumps(prompt)}) + result = self.extractor.extract_metadata(path) + self.assertIsNotNone(result) + self.assertIn("prompt", result) + + def test_extracts_text_encoder_nodes(self): + workflow = { + "nodes": [ + {"type": "CLIPTextEncode", "inputs": {"text": "positive prompt"}}, + {"type": "KSampler", "inputs": {}}, + ] + } + path = self._png({"workflow": json.dumps(workflow)}) + result = self.extractor.extract_metadata(path) + self.assertIn("text_encoder_nodes", result) + self.assertEqual(len(result["text_encoder_nodes"]), 1) + + def test_extracts_simple_metadata_fields(self): + path = self._png({"steps": "20", "seed": "12345", "model": "sd_v1.5"}) + result = self.extractor.extract_metadata(path) + self.assertIsNotNone(result) + # "20" and "12345" are valid JSON, so they get parsed to int + self.assertEqual(result["steps"], 20) + self.assertEqual(result["seed"], 12345) + # "sd_v1.5" fails JSON parse, stored as string + self.assertEqual(result["model"], "sd_v1.5") + + def test_json_metadata_fields_parsed(self): + path = self._png({"cfg_scale": json.dumps(7.5)}) + result = self.extractor.extract_metadata(path) + self.assertEqual(result["cfg_scale"], 7.5) + + def test_file_info_always_present(self): + workflow = {"nodes": [{"type": "CLIPTextEncode", "inputs": {}}]} + path = self._png({"workflow": json.dumps(workflow)}, size=(128, 256)) + result = self.extractor.extract_metadata(path) + fi = result["file_info"] + self.assertEqual(fi["dimensions"], [128, 256]) + self.assertEqual(fi["format"], "PNG") + self.assertEqual(fi["mode"], "RGB") + self.assertIn("size", fi) + self.assertIn("created_time", fi) + + def test_invalid_workflow_json_handled(self): + path = self._png({"workflow": "not valid json {{"}) + result = self.extractor.extract_metadata(path) + # Should not contain workflow since JSON parse failed + self.assertIsNone(result) + + def test_invalid_prompt_json_handled(self): + workflow = {"nodes": [{"type": "CLIPTextEncode", "inputs": {}}]} + path = self._png( + { + "workflow": json.dumps(workflow), + "prompt": "broken json {{", + } + ) + result = self.extractor.extract_metadata(path) + # Workflow should still be extracted even if prompt fails + self.assertIn("workflow", result) + self.assertNotIn("prompt", result) + + def test_nonexistent_file_returns_none(self): + result = self.extractor.extract_metadata("/nonexistent/path/image.png") + self.assertIsNone(result) + + +class TestGetFileInfo(MetadataTestCase): + """Test get_file_info method.""" + + def test_returns_correct_dimensions(self): + path = self._png(size=(320, 240)) + with Image.open(path) as img: + info = self.extractor.get_file_info(path, img) + self.assertEqual(info["dimensions"], [320, 240]) + + def test_returns_file_size(self): + path = self._png() + with Image.open(path) as img: + info = self.extractor.get_file_info(path, img) + self.assertGreater(info["size"], 0) + self.assertEqual(info["size"], os.path.getsize(path)) + + def test_returns_format_and_mode(self): + path = self._png() + with Image.open(path) as img: + info = self.extractor.get_file_info(path, img) + self.assertEqual(info["format"], "PNG") + self.assertEqual(info["mode"], "RGB") + + +class TestFindTextEncoderNodes(MetadataTestCase): + """Test find_text_encoder_nodes method.""" + + def test_finds_clip_text_encode(self): + workflow = { + "nodes": [ + {"type": "CLIPTextEncode", "inputs": {"text": "prompt"}}, + ] + } + nodes = self.extractor.find_text_encoder_nodes(workflow) + self.assertEqual(len(nodes), 1) + + def test_finds_sdxl_encoder(self): + workflow = { + "nodes": [ + {"type": "CLIPTextEncodeSDXL", "inputs": {}}, + ] + } + nodes = self.extractor.find_text_encoder_nodes(workflow) + self.assertEqual(len(nodes), 1) + + def test_finds_prompt_manager_node(self): + workflow = { + "nodes": [ + {"type": "PromptManager", "inputs": {"text": "test"}}, + ] + } + nodes = self.extractor.find_text_encoder_nodes(workflow) + self.assertEqual(len(nodes), 1) + + def test_ignores_non_encoder_nodes(self): + workflow = { + "nodes": [ + {"type": "KSampler", "inputs": {}}, + {"type": "CheckpointLoader", "inputs": {}}, + ] + } + nodes = self.extractor.find_text_encoder_nodes(workflow) + self.assertEqual(len(nodes), 0) + + def test_handles_dict_format_nodes(self): + workflow = { + "1": {"class_type": "CLIPTextEncode", "inputs": {"text": "test"}}, + "2": {"class_type": "KSampler", "inputs": {}}, + } + nodes = self.extractor.find_text_encoder_nodes(workflow) + self.assertEqual(len(nodes), 1) + self.assertEqual(nodes[0]["node_id"], "1") + + def test_handles_nested_workflow(self): + workflow = { + "workflow": { + "nodes": [ + {"type": "CLIPTextEncode", "inputs": {}}, + ] + } + } + nodes = self.extractor.find_text_encoder_nodes(workflow) + self.assertEqual(len(nodes), 1) + + def test_empty_workflow(self): + nodes = self.extractor.find_text_encoder_nodes({}) + self.assertEqual(nodes, []) + + def test_non_dict_returns_empty(self): + nodes = self.extractor.find_text_encoder_nodes("not a dict") + self.assertEqual(nodes, []) + + def test_none_nodes_data(self): + workflow = {"nodes": None} + nodes = self.extractor.find_text_encoder_nodes(workflow) + self.assertEqual(nodes, []) + + +class TestIsTextEncoderNode(MetadataTestCase): + """Test is_text_encoder_node method.""" + + def test_clip_text_encode(self): + self.assertTrue(self.extractor.is_text_encoder_node({"type": "CLIPTextEncode"})) + + def test_class_type_field(self): + self.assertTrue( + self.extractor.is_text_encoder_node({"class_type": "CLIPTextEncode"}) + ) + + def test_case_insensitive(self): + self.assertTrue(self.extractor.is_text_encoder_node({"type": "cliptextencode"})) + + def test_title_with_text_keyword(self): + self.assertTrue( + self.extractor.is_text_encoder_node( + {"type": "Unknown", "title": "My Text Encoder"} + ) + ) + + def test_title_with_prompt_keyword(self): + self.assertTrue( + self.extractor.is_text_encoder_node( + {"type": "Unknown", "title": "Prompt Input"} + ) + ) + + def test_non_encoder_node(self): + self.assertFalse(self.extractor.is_text_encoder_node({"type": "KSampler"})) + + def test_non_dict_returns_false(self): + self.assertFalse(self.extractor.is_text_encoder_node("string")) + self.assertFalse(self.extractor.is_text_encoder_node(None)) + self.assertFalse(self.extractor.is_text_encoder_node(42)) + + def test_empty_dict_returns_false(self): + self.assertFalse(self.extractor.is_text_encoder_node({})) + + +class TestExtractPromptTextFromWorkflow(MetadataTestCase): + """Test extract_prompt_text_from_workflow method.""" + + def test_extracts_text_input(self): + workflow = { + "nodes": [ + {"type": "CLIPTextEncode", "inputs": {"text": "beautiful sunset"}}, + ] + } + text = self.extractor.extract_prompt_text_from_workflow(workflow) + self.assertEqual(text, "beautiful sunset") + + def test_extracts_prompt_input(self): + workflow = { + "nodes": [ + {"type": "CLIPTextEncode", "inputs": {"prompt": "a cat"}}, + ] + } + text = self.extractor.extract_prompt_text_from_workflow(workflow) + self.assertEqual(text, "a cat") + + def test_returns_none_when_no_text(self): + workflow = { + "nodes": [ + {"type": "KSampler", "inputs": {"steps": 20}}, + ] + } + text = self.extractor.extract_prompt_text_from_workflow(workflow) + self.assertIsNone(text) + + def test_returns_none_for_empty_workflow(self): + text = self.extractor.extract_prompt_text_from_workflow({}) + self.assertIsNone(text) + + def test_handles_list_input_value(self): + workflow = { + "nodes": [ + {"type": "CLIPTextEncode", "inputs": {"text": ["linked_value", 0]}}, + ] + } + text = self.extractor.extract_prompt_text_from_workflow(workflow) + self.assertEqual(text, "linked_value") + + def test_skips_empty_text_fields(self): + workflow = { + "nodes": [ + {"type": "CLIPTextEncode", "inputs": {"text": ""}}, + {"type": "PromptManager", "inputs": {"text": "actual prompt"}}, + ] + } + text = self.extractor.extract_prompt_text_from_workflow(workflow) + self.assertEqual(text, "actual prompt") + + +class TestGetGenerationParameters(MetadataTestCase): + """Test get_generation_parameters method.""" + + def test_extracts_top_level_params(self): + metadata = {"steps": 20, "cfg_scale": 7.5, "sampler": "euler", "seed": 12345} + params = self.extractor.get_generation_parameters(metadata) + self.assertEqual(params["steps"], 20) + self.assertEqual(params["cfg_scale"], 7.5) + self.assertEqual(params["sampler"], "euler") + self.assertEqual(params["seed"], 12345) + + def test_empty_metadata(self): + params = self.extractor.get_generation_parameters({}) + self.assertEqual(params, {}) + + def test_only_relevant_fields_extracted(self): + metadata = {"steps": 20, "irrelevant_field": "ignored", "workflow": {}} + params = self.extractor.get_generation_parameters(metadata) + self.assertIn("steps", params) + self.assertNotIn("irrelevant_field", params) + + def test_with_workflow_present(self): + metadata = {"steps": 20, "workflow": {"nodes": []}} + params = self.extractor.get_generation_parameters(metadata) + self.assertIn("steps", params) + + def test_dimension_fields(self): + metadata = {"width": 512, "height": 768, "batch_size": 4} + params = self.extractor.get_generation_parameters(metadata) + self.assertEqual(params["width"], 512) + self.assertEqual(params["height"], 768) + self.assertEqual(params["batch_size"], 4) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_prompt_tracker.py b/tests/test_prompt_tracker.py new file mode 100644 index 0000000..2b8c6b2 --- /dev/null +++ b/tests/test_prompt_tracker.py @@ -0,0 +1,260 @@ +""" +Tests for prompt tracking system. + +Tests PromptTracker, PromptExecutionContext, and the singleton +get_prompt_tracker. Uses a mock db_manager to avoid database deps. +""" + +import os +import sys +import time +import threading +import unittest +from unittest.mock import MagicMock, patch + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from utils.prompt_tracker import PromptTracker, PromptExecutionContext + + +def _make_tracker(prompt_timeout=600, cleanup_interval=300): + """Create a PromptTracker with mocked db_manager and config.""" + mock_db = MagicMock() + mock_db.get_prompt_by_hash.return_value = None + + # Patch GalleryConfig import to avoid needing ComfyUI server + with patch.object(PromptTracker, "__init__", lambda self, db: None): + tracker = PromptTracker.__new__(PromptTracker) + + # Manual init to avoid the config import and cleanup thread + from utils.logging_config import get_logger + + tracker.logger = get_logger("prompt_manager.prompt_tracker") + tracker.db_manager = mock_db + tracker._local = threading.local() + tracker.active_prompts = {} + tracker.lock = threading.Lock() + tracker.cleanup_interval = cleanup_interval + tracker.prompt_timeout = prompt_timeout + # Don't start cleanup thread in tests + return tracker + + +class TestPromptTrackerBasic(unittest.TestCase): + """Test basic PromptTracker operations.""" + + def setUp(self): + self.tracker = _make_tracker() + + def test_set_current_prompt_returns_execution_id(self): + exec_id = self.tracker.set_current_prompt("test prompt", {"prompt_id": 42}) + self.assertTrue(exec_id.startswith("exec_")) + + def test_set_current_prompt_stores_in_active(self): + exec_id = self.tracker.set_current_prompt("test", {"prompt_id": 1}) + self.assertIn(exec_id, self.tracker.active_prompts) + + def test_get_current_prompt_returns_set_prompt(self): + self.tracker.set_current_prompt("hello world", {"prompt_id": 5}) + current = self.tracker.get_current_prompt() + self.assertIsNotNone(current) + self.assertEqual(current["text"], "hello world") + self.assertEqual(current["id"], 5) + + def test_get_current_prompt_returns_none_when_empty(self): + result = self.tracker.get_current_prompt() + self.assertIsNone(result) + + def test_clear_current_prompt(self): + self.tracker.set_current_prompt("to clear", {"prompt_id": 1}) + self.tracker.clear_current_prompt() + result = self.tracker.get_current_prompt() + self.assertIsNone(result) + + def test_clear_removes_from_active(self): + exec_id = self.tracker.set_current_prompt("test", {"prompt_id": 1}) + self.tracker.clear_current_prompt() + self.assertNotIn(exec_id, self.tracker.active_prompts) + + def test_clear_when_no_prompt_does_not_raise(self): + # Should not raise + self.tracker.clear_current_prompt() + + +class TestPromptTrackerTimeout(unittest.TestCase): + """Test timeout behavior.""" + + def test_expired_prompt_returns_none(self): + tracker = _make_tracker(prompt_timeout=0.1) + tracker.set_current_prompt("expires fast", {"prompt_id": 1}) + time.sleep(0.2) + result = tracker.get_current_prompt() + self.assertIsNone(result) + + def test_extend_timeout_resets_timestamp(self): + tracker = _make_tracker(prompt_timeout=1) + exec_id = tracker.set_current_prompt("test", {"prompt_id": 1}) + old_ts = tracker.active_prompts[exec_id]["timestamp"] + time.sleep(0.1) + tracker.extend_prompt_timeout(exec_id) + new_ts = tracker.active_prompts[exec_id]["timestamp"] + self.assertGreater(new_ts, old_ts) + + def test_extend_nonexistent_id_does_not_raise(self): + tracker = _make_tracker() + tracker.extend_prompt_timeout("nonexistent_id") + + +class TestPromptTrackerMultiplePrompts(unittest.TestCase): + """Test multiple prompt tracking.""" + + def setUp(self): + self.tracker = _make_tracker() + + def test_get_active_prompts(self): + self.tracker.set_current_prompt("p1", {"prompt_id": 1}) + self.tracker.set_current_prompt("p2", {"prompt_id": 2}) + active = self.tracker.get_active_prompts() + self.assertEqual(len(active), 2) + + def test_get_active_prompts_returns_copy(self): + self.tracker.set_current_prompt("p1", {"prompt_id": 1}) + active = self.tracker.get_active_prompts() + active.clear() + # Original should be untouched + self.assertEqual(len(self.tracker.active_prompts), 1) + + def test_clear_all_active_prompts(self): + self.tracker.set_current_prompt("p1", {"prompt_id": 1}) + self.tracker.set_current_prompt("p2", {"prompt_id": 2}) + cleared = self.tracker.clear_all_active_prompts() + self.assertEqual(cleared, 2) + self.assertEqual(len(self.tracker.active_prompts), 0) + + def test_clear_all_returns_zero_when_empty(self): + cleared = self.tracker.clear_all_active_prompts() + self.assertEqual(cleared, 0) + + +class TestPromptTrackerFallback(unittest.TestCase): + """Test fallback prompt lookup (for cross-thread access).""" + + def test_fallback_finds_recent_prompt(self): + tracker = _make_tracker() + # Set prompt in active_prompts directly (simulating another thread) + tracker.active_prompts["exec_123"] = { + "id": 42, + "execution_id": "exec_123", + "text": "from other thread", + "timestamp": time.time(), + "thread_id": 0, + "additional_data": {}, + } + # Current thread has no local prompt, but fallback should find it + result = tracker.get_current_prompt() + self.assertIsNotNone(result) + self.assertEqual(result["text"], "from other thread") + + def test_fallback_uses_db_hash_lookup(self): + tracker = _make_tracker() + tracker.db_manager.get_prompt_by_hash.return_value = {"id": 99, "text": "found"} + exec_id = tracker.set_current_prompt("lookup test") + self.assertEqual(tracker.active_prompts[exec_id]["id"], 99) + + def test_fallback_generates_temp_id(self): + tracker = _make_tracker() + tracker.db_manager.get_prompt_by_hash.return_value = None + exec_id = tracker.set_current_prompt("not in db") + prompt_id = tracker.active_prompts[exec_id]["id"] + self.assertTrue(str(prompt_id).startswith("temp_")) + + +class TestPromptTrackerStatus(unittest.TestCase): + """Test get_status method.""" + + def test_status_empty_tracker(self): + tracker = _make_tracker() + status = tracker.get_status() + self.assertEqual(status["active_prompts_count"], 0) + self.assertIsNone(status["current_prompt_id"]) + self.assertIsNone(status["current_execution_id"]) + self.assertEqual(status["prompt_timeout"], 600) + self.assertEqual(status["cleanup_interval"], 300) + + def test_status_with_active_prompt(self): + tracker = _make_tracker() + tracker.set_current_prompt("active", {"prompt_id": 7}) + status = tracker.get_status() + self.assertEqual(status["active_prompts_count"], 1) + self.assertEqual(status["current_prompt_id"], 7) + self.assertIsNotNone(status["current_execution_id"]) + + +class TestGenerateExecutionId(unittest.TestCase): + """Test execution ID generation.""" + + def test_format(self): + tracker = _make_tracker() + exec_id = tracker.generate_execution_id() + self.assertTrue(exec_id.startswith("exec_")) + parts = exec_id.split("_") + self.assertEqual(len(parts), 3) # exec, uuid8, timestamp + + def test_unique(self): + tracker = _make_tracker() + ids = {tracker.generate_execution_id() for _ in range(100)} + self.assertEqual(len(ids), 100) + + +class TestPromptExecutionContext(unittest.TestCase): + """Test the context manager.""" + + def test_context_manager_sets_prompt(self): + tracker = _make_tracker() + with PromptExecutionContext(tracker, "context test", prompt_id=10) as exec_id: + self.assertIsNotNone(exec_id) + current = tracker.get_current_prompt() + self.assertEqual(current["text"], "context test") + + def test_context_manager_does_not_clear_on_exit(self): + tracker = _make_tracker() + with PromptExecutionContext(tracker, "stays active", prompt_id=1): + pass + # Prompt should still be in active_prompts (designed behavior) + self.assertEqual(len(tracker.active_prompts), 1) + + def test_context_manager_with_exception(self): + tracker = _make_tracker() + try: + with PromptExecutionContext(tracker, "error test", prompt_id=1): + raise ValueError("test error") + except ValueError: + pass + # Should not crash, prompt stays tracked + self.assertEqual(len(tracker.active_prompts), 1) + + +class TestCleanupExpiredPrompts(unittest.TestCase): + """Test the cleanup logic (not the thread, just the logic).""" + + def test_expired_prompts_cleaned(self): + tracker = _make_tracker(prompt_timeout=0.1) + tracker.set_current_prompt("old", {"prompt_id": 1}) + time.sleep(0.2) + + # Manually run cleanup logic + current_time = time.time() + expired_ids = [] + with tracker.lock: + for exec_id, prompt_data in tracker.active_prompts.items(): + if current_time - prompt_data["timestamp"] > tracker.prompt_timeout: + expired_ids.append(exec_id) + for exec_id in expired_ids: + tracker.active_prompts.pop(exec_id, None) + + self.assertEqual(len(tracker.active_prompts), 0) + self.assertEqual(len(expired_ids), 1) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_validators.py b/tests/test_validators.py new file mode 100644 index 0000000..6826119 --- /dev/null +++ b/tests/test_validators.py @@ -0,0 +1,291 @@ +""" +Comprehensive tests for input validation utilities. + +Tests all 7 validator functions in utils/validators.py including +edge cases, boundary conditions, and security-related inputs. +""" + +import os +import sys +import unittest + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from utils.validators import ( + validate_prompt_text, + validate_rating, + validate_tags, + validate_category, + validate_workflow_name, + sanitize_input, + parse_tags_string, +) + + +class TestValidateCategory(unittest.TestCase): + """Test validate_category function.""" + + def test_none_is_valid(self): + self.assertTrue(validate_category(None)) + + def test_empty_string_is_valid(self): + self.assertTrue(validate_category("")) + + def test_whitespace_only_is_valid(self): + self.assertTrue(validate_category(" ")) + + def test_normal_category(self): + self.assertTrue(validate_category("landscape")) + + def test_category_with_spaces(self): + self.assertTrue(validate_category("nature photos")) + + def test_category_with_hyphens_underscores(self): + self.assertTrue(validate_category("sci-fi_art")) + + def test_unicode_category(self): + self.assertTrue(validate_category("日本語カテゴリ")) + + def test_max_length_boundary(self): + self.assertTrue(validate_category("x" * 100)) + + def test_exceeds_max_length(self): + with self.assertRaises(ValueError) as ctx: + validate_category("x" * 101) + self.assertIn("100", str(ctx.exception)) + + def test_non_string_raises(self): + with self.assertRaises(ValueError): + validate_category(123) + + def test_list_raises(self): + with self.assertRaises(ValueError): + validate_category(["category"]) + + def test_control_characters_rejected(self): + with self.assertRaises(ValueError) as ctx: + validate_category("bad\x00category") + self.assertIn("control characters", str(ctx.exception)) + + def test_tab_character_rejected(self): + with self.assertRaises(ValueError): + validate_category("bad\tcategory") + + def test_newline_rejected(self): + with self.assertRaises(ValueError): + validate_category("bad\ncategory") + + +class TestValidateWorkflowName(unittest.TestCase): + """Test validate_workflow_name function.""" + + def test_none_is_valid(self): + self.assertTrue(validate_workflow_name(None)) + + def test_empty_string_is_valid(self): + self.assertTrue(validate_workflow_name("")) + + def test_whitespace_only_is_valid(self): + self.assertTrue(validate_workflow_name(" ")) + + def test_normal_name(self): + self.assertTrue(validate_workflow_name("My Workflow")) + + def test_name_with_special_chars(self): + self.assertTrue(validate_workflow_name("workflow-v2_final (copy)")) + + def test_max_length_boundary(self): + self.assertTrue(validate_workflow_name("x" * 200)) + + def test_exceeds_max_length(self): + with self.assertRaises(ValueError) as ctx: + validate_workflow_name("x" * 201) + self.assertIn("200", str(ctx.exception)) + + def test_non_string_raises(self): + with self.assertRaises(ValueError): + validate_workflow_name(42) + + def test_dict_raises(self): + with self.assertRaises(ValueError): + validate_workflow_name({"name": "workflow"}) + + +class TestSanitizeInput(unittest.TestCase): + """Test sanitize_input function.""" + + def test_normal_text_unchanged(self): + self.assertEqual(sanitize_input("hello world"), "hello world") + + def test_strips_outer_whitespace(self): + self.assertEqual(sanitize_input(" hello "), "hello") + + def test_removes_null_bytes(self): + self.assertEqual(sanitize_input("hello\x00world"), "helloworld") + + def test_normalizes_windows_line_endings(self): + self.assertEqual(sanitize_input("line1\r\nline2"), "line1\nline2") + + def test_normalizes_old_mac_line_endings(self): + self.assertEqual(sanitize_input("line1\rline2"), "line1\nline2") + + def test_strips_whitespace_per_line(self): + result = sanitize_input(" line1 \n line2 ") + self.assertEqual(result, "line1\nline2") + + def test_limits_consecutive_empty_lines(self): + text = "a\n\n\n\n\nb" + result = sanitize_input(text) + # Max 2 consecutive empty lines = 3 newlines between content + # 4+ newlines (3+ empty lines) should never appear + self.assertNotIn("\n\n\n\n", result) + self.assertIn("a", result) + self.assertIn("b", result) + + def test_non_string_returns_empty(self): + self.assertEqual(sanitize_input(None), "") + self.assertEqual(sanitize_input(123), "") + self.assertEqual(sanitize_input([]), "") + + def test_empty_string(self): + self.assertEqual(sanitize_input(""), "") + + def test_preserves_single_newlines(self): + result = sanitize_input("line1\nline2\nline3") + self.assertEqual(result, "line1\nline2\nline3") + + def test_mixed_line_endings(self): + result = sanitize_input("a\r\nb\rc\nd") + self.assertEqual(result, "a\nb\nc\nd") + + +class TestParseTagsString(unittest.TestCase): + """Test parse_tags_string function.""" + + def test_simple_comma_separated(self): + result = parse_tags_string("tag1, tag2, tag3") + self.assertEqual(result, ["tag1", "tag2", "tag3"]) + + def test_extra_whitespace(self): + result = parse_tags_string(" tag1 , tag2 ") + self.assertEqual(result, ["tag1", "tag2"]) + + def test_empty_string_returns_empty_list(self): + self.assertEqual(parse_tags_string(""), []) + + def test_none_returns_empty_list(self): + self.assertEqual(parse_tags_string(None), []) + + def test_non_string_returns_empty_list(self): + self.assertEqual(parse_tags_string(123), []) + + def test_single_tag(self): + self.assertEqual(parse_tags_string("solo"), ["solo"]) + + def test_deduplication(self): + result = parse_tags_string("tag1, tag1, tag2, tag1") + self.assertEqual(result, ["tag1", "tag2"]) + + def test_empty_segments_skipped(self): + result = parse_tags_string("tag1,,tag2,,,tag3") + self.assertEqual(result, ["tag1", "tag2", "tag3"]) + + def test_max_20_tags(self): + tags = ", ".join(f"tag{i}" for i in range(25)) + result = parse_tags_string(tags) + self.assertEqual(len(result), 20) + + def test_strips_null_bytes_from_tags(self): + result = parse_tags_string("clean\x00tag, normal") + self.assertIn("cleantag", result) + self.assertIn("normal", result) + + +class TestValidatePromptTextEdgeCases(unittest.TestCase): + """Additional edge cases beyond test_basic.py coverage.""" + + def test_exactly_max_length(self): + self.assertTrue(validate_prompt_text("x" * 10000)) + + def test_one_over_max_length(self): + with self.assertRaises(ValueError): + validate_prompt_text("x" * 10001) + + def test_single_character(self): + self.assertTrue(validate_prompt_text("a")) + + def test_unicode_text(self): + self.assertTrue(validate_prompt_text("日本語プロンプト")) + + def test_multiline_text(self): + self.assertTrue(validate_prompt_text("line1\nline2\nline3")) + + def test_boolean_raises(self): + with self.assertRaises(ValueError): + validate_prompt_text(True) + + def test_none_raises(self): + with self.assertRaises(ValueError): + validate_prompt_text(None) + + +class TestValidateRatingEdgeCases(unittest.TestCase): + """Additional edge cases beyond test_basic.py coverage.""" + + def test_float_raises(self): + with self.assertRaises(ValueError): + validate_rating(3.5) + + def test_negative_raises(self): + with self.assertRaises(ValueError): + validate_rating(-1) + + def test_zero_raises(self): + with self.assertRaises(ValueError): + validate_rating(0) + + def test_boolean_true_is_int(self): + # In Python, bool is a subclass of int: True == 1 + # This tests the actual behavior + self.assertTrue(validate_rating(True)) + + +class TestValidateTagsEdgeCases(unittest.TestCase): + """Additional edge cases beyond test_basic.py coverage.""" + + def test_empty_list_is_valid(self): + self.assertTrue(validate_tags([])) + + def test_comma_string_parsed(self): + self.assertTrue(validate_tags("tag1, tag2")) + + def test_control_char_in_tag_rejected(self): + with self.assertRaises(ValueError) as ctx: + validate_tags(["bad\x00tag"]) + self.assertIn("control characters", str(ctx.exception)) + + def test_tab_in_tag_rejected(self): + with self.assertRaises(ValueError): + validate_tags(["bad\ttag"]) + + def test_non_string_tag_in_list_rejected(self): + with self.assertRaises(ValueError): + validate_tags([123]) + + def test_dict_input_rejected(self): + with self.assertRaises(ValueError): + validate_tags({"tag": "value"}) + + def test_exactly_20_tags(self): + self.assertTrue(validate_tags([f"tag{i}" for i in range(20)])) + + def test_exactly_50_char_tag(self): + self.assertTrue(validate_tags(["x" * 50])) + + def test_whitespace_only_tag_rejected(self): + with self.assertRaises(ValueError): + validate_tags([" "]) + + +if __name__ == "__main__": + unittest.main()