style: apply Black formatting to new test files
This commit is contained in:
+23
-9
@@ -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
@@ -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
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user