diff --git a/tests/test_config.py b/tests/test_config.py index e2cbd13..05037ce 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -44,12 +44,17 @@ class TestGalleryConfig(unittest.TestCase): 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["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}}) + GalleryConfig.update_config( + {"monitoring": {"enabled": False, "processing_delay": 5.0}} + ) self.assertFalse(GalleryConfig.MONITORING_ENABLED) self.assertEqual(GalleryConfig.PROCESSING_DELAY, 5.0) @@ -58,12 +63,16 @@ class TestGalleryConfig(unittest.TestCase): self.assertEqual(GalleryConfig.PROMPT_TIMEOUT, 300) def test_update_database(self): - GalleryConfig.update_config({"database": {"auto_cleanup": False, "max_image_age_days": 30}}) + 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}}) + 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) @@ -128,15 +137,17 @@ class TestPromptManagerConfig(unittest.TestCase): self.assertEqual(PromptManagerConfig.MAX_SEARCH_RESULTS, 50) def test_update_propagates_to_gallery(self): - PromptManagerConfig.update_config({ - "gallery": {"monitoring": {"enabled": False}} - }) + 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 = tempfile.NamedTemporaryFile( + delete=False, suffix=".json", dir=tempfile.gettempdir() + ) tmp.close() try: # Modify a value @@ -179,10 +190,13 @@ class TestPromptManagerConfig(unittest.TestCase): 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 = tempfile.NamedTemporaryFile( + delete=False, suffix=".json", dir=tempfile.gettempdir() + ) tmp.close() try: PromptManagerConfig.save_to_file(tmp.name) diff --git a/tests/test_hashing.py b/tests/test_hashing.py index b868dfb..ed631fa 100644 --- a/tests/test_hashing.py +++ b/tests/test_hashing.py @@ -12,14 +12,23 @@ 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 +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"} + content = { + "text": "A prompt", + "category": "nature", + "tags": ["sky"], + "workflow_name": "wf1", + } h = generate_content_hash(content) self.assertEqual(len(h), 64) # SHA256 hex digest diff --git a/tests/test_metadata_extraction.py b/tests/test_metadata_extraction.py index bb3d396..ec7da95 100644 --- a/tests/test_metadata_extraction.py +++ b/tests/test_metadata_extraction.py @@ -83,10 +83,12 @@ class TestExtractMetadata(MetadataTestCase): self.assertIn("prompt", result) def test_extracts_text_encoder_nodes(self): - workflow = {"nodes": [ - {"type": "CLIPTextEncode", "inputs": {"text": "positive prompt"}}, - {"type": "KSampler", "inputs": {}}, - ]} + 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) @@ -126,10 +128,12 @@ class TestExtractMetadata(MetadataTestCase): def test_invalid_prompt_json_handled(self): workflow = {"nodes": [{"type": "CLIPTextEncode", "inputs": {}}]} - path = self._png({ - "workflow": json.dumps(workflow), - "prompt": "broken json {{", - }) + 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) @@ -168,31 +172,39 @@ class TestFindTextEncoderNodes(MetadataTestCase): """Test find_text_encoder_nodes method.""" def test_finds_clip_text_encode(self): - workflow = {"nodes": [ - {"type": "CLIPTextEncode", "inputs": {"text": "prompt"}}, - ]} + 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": {}}, - ]} + 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"}}, - ]} + 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": {}}, - ]} + workflow = { + "nodes": [ + {"type": "KSampler", "inputs": {}}, + {"type": "CheckpointLoader", "inputs": {}}, + ] + } nodes = self.extractor.find_text_encoder_nodes(workflow) self.assertEqual(len(nodes), 0) @@ -206,9 +218,13 @@ class TestFindTextEncoderNodes(MetadataTestCase): self.assertEqual(nodes[0]["node_id"], "1") def test_handles_nested_workflow(self): - workflow = {"workflow": {"nodes": [ - {"type": "CLIPTextEncode", "inputs": {}}, - ]}} + workflow = { + "workflow": { + "nodes": [ + {"type": "CLIPTextEncode", "inputs": {}}, + ] + } + } nodes = self.extractor.find_text_encoder_nodes(workflow) self.assertEqual(len(nodes), 1) @@ -233,20 +249,26 @@ class TestIsTextEncoderNode(MetadataTestCase): 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"})) + 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"} - )) + 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"} - )) + 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"})) @@ -264,23 +286,29 @@ class TestExtractPromptTextFromWorkflow(MetadataTestCase): """Test extract_prompt_text_from_workflow method.""" def test_extracts_text_input(self): - workflow = {"nodes": [ - {"type": "CLIPTextEncode", "inputs": {"text": "beautiful sunset"}}, - ]} + 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"}}, - ]} + 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}}, - ]} + workflow = { + "nodes": [ + {"type": "KSampler", "inputs": {"steps": 20}}, + ] + } text = self.extractor.extract_prompt_text_from_workflow(workflow) self.assertIsNone(text) @@ -289,17 +317,21 @@ class TestExtractPromptTextFromWorkflow(MetadataTestCase): self.assertIsNone(text) def test_handles_list_input_value(self): - workflow = {"nodes": [ - {"type": "CLIPTextEncode", "inputs": {"text": ["linked_value", 0]}}, - ]} + 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"}}, - ]} + 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") diff --git a/tests/test_prompt_tracker.py b/tests/test_prompt_tracker.py index f04a87c..2b8c6b2 100644 --- a/tests/test_prompt_tracker.py +++ b/tests/test_prompt_tracker.py @@ -47,15 +47,11 @@ class TestPromptTrackerBasic(unittest.TestCase): 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} - ) + 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} - ) + 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):