Merge pull request #107 from ComfyAssets/chore/ci-cd-improvements
Chore/ci cd improvements
This commit is contained in:
@@ -1,5 +1,11 @@
|
||||
# ComfyUI Prompt Manager
|
||||
|
||||

|
||||

|
||||

|
||||

|
||||

|
||||
|
||||
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
|
||||
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user