style: apply Black formatting to new test files

This commit is contained in:
Vito Sansevero
2026-02-07 12:36:11 -08:00
parent e1bf06cb3d
commit 15d5a9197a
4 changed files with 115 additions and 64 deletions
+23 -9
View File
@@ -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)
+11 -2
View File
@@ -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
+79 -47
View File
@@ -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")
+2 -6
View File
@@ -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):