Merge pull request #107 from ComfyAssets/chore/ci-cd-improvements

Chore/ci cd improvements
This commit is contained in:
Vito
2026-02-07 12:43:49 -08:00
committed by GitHub
6 changed files with 1290 additions and 0 deletions
+6
View File
@@ -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
+211
View File
@@ -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()
+148
View File
@@ -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()
+374
View File
@@ -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()
+260
View File
@@ -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()
+291
View File
@@ -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()