Files
ComfyAssets-ComfyUI_PromptM…/tests/test_basic.py
T
Vito Sansevero 45d23efcf1 test: add comprehensive database tests and CI pipeline
Phase 6 test infrastructure (6.1, 6.2, 6.5):
- Fix .coveragerc source path (src -> .)
- Fix pre-existing test failure (patch target + method rename)
- Add 52 database layer tests covering CRUD, junction table tags,
  search, pagination, statistics, image linking, and edge cases
- Add GitHub Actions CI workflow (Python 3.10-3.12)
- Total: 61 tests, all passing
2026-02-07 07:06:54 -08:00

269 lines
9.2 KiB
Python

"""
Basic tests for PromptManager functionality.
"""
import os
import tempfile
import unittest
from unittest.mock import Mock, patch
# Add the parent directory to the path for imports
import sys
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
from database.operations import PromptDatabase
from utils.hashing import generate_prompt_hash
from utils.validators import validate_prompt_text, validate_rating, validate_tags
class TestBasicFunctionality(unittest.TestCase):
"""Test basic functionality of PromptManager components."""
def setUp(self):
"""Set up test fixtures."""
# Use a temporary database for testing
self.temp_db = tempfile.NamedTemporaryFile(delete=False, suffix='.db')
self.temp_db.close()
self.db = PromptDatabase(self.temp_db.name)
def tearDown(self):
"""Clean up test fixtures."""
# Remove temporary database
if os.path.exists(self.temp_db.name):
os.unlink(self.temp_db.name)
def test_prompt_hash_generation(self):
"""Test prompt hash generation."""
text1 = "A beautiful landscape with mountains"
text2 = "A BEAUTIFUL LANDSCAPE WITH MOUNTAINS" # Different case
text3 = " A beautiful landscape with mountains " # Extra whitespace
text4 = "A different prompt"
hash1 = generate_prompt_hash(text1)
hash2 = generate_prompt_hash(text2)
hash3 = generate_prompt_hash(text3)
hash4 = generate_prompt_hash(text4)
# Same content should produce same hash (case and whitespace insensitive)
self.assertEqual(hash1, hash2)
self.assertEqual(hash1, hash3)
# Different content should produce different hash
self.assertNotEqual(hash1, hash4)
# Hash should be 64 characters (SHA256 hex)
self.assertEqual(len(hash1), 64)
def test_prompt_validation(self):
"""Test prompt text validation."""
# Valid prompts
self.assertTrue(validate_prompt_text("Valid prompt text"))
self.assertTrue(validate_prompt_text("Another valid prompt"))
# Invalid prompts
with self.assertRaises(ValueError):
validate_prompt_text("") # Empty
with self.assertRaises(ValueError):
validate_prompt_text(" ") # Whitespace only
with self.assertRaises(ValueError):
validate_prompt_text("x" * 10001) # Too long
with self.assertRaises(ValueError):
validate_prompt_text(123) # Not a string
def test_rating_validation(self):
"""Test rating validation."""
# Valid ratings
self.assertTrue(validate_rating(None))
self.assertTrue(validate_rating(1))
self.assertTrue(validate_rating(3))
self.assertTrue(validate_rating(5))
# Invalid ratings
with self.assertRaises(ValueError):
validate_rating(0) # Too low
with self.assertRaises(ValueError):
validate_rating(6) # Too high
with self.assertRaises(ValueError):
validate_rating("3") # Not an integer
def test_tags_validation(self):
"""Test tags validation."""
# Valid tags
self.assertTrue(validate_tags(None))
self.assertTrue(validate_tags([]))
self.assertTrue(validate_tags(["tag1", "tag2"]))
self.assertTrue(validate_tags("tag1, tag2, tag3"))
# Invalid tags
with self.assertRaises(ValueError):
validate_tags([""]) # Empty tag
with self.assertRaises(ValueError):
validate_tags(["x" * 51]) # Tag too long
with self.assertRaises(ValueError):
validate_tags(["tag"] * 21) # Too many tags
def test_database_save_and_retrieve(self):
"""Test basic database save and retrieve operations."""
# Save a prompt
prompt_id = self.db.save_prompt(
text="Test prompt for database",
category="test",
tags=["test", "database"],
rating=4,
notes="Test notes",
prompt_hash=generate_prompt_hash("Test prompt for database")
)
self.assertIsInstance(prompt_id, int)
self.assertGreater(prompt_id, 0)
# Retrieve the prompt
retrieved = self.db.get_prompt_by_id(prompt_id)
self.assertIsNotNone(retrieved)
self.assertEqual(retrieved['text'], "Test prompt for database")
self.assertEqual(retrieved['category'], "test")
self.assertEqual(retrieved['tags'], ["test", "database"])
self.assertEqual(retrieved['rating'], 4)
self.assertEqual(retrieved['notes'], "Test notes")
def test_duplicate_detection(self):
"""Test duplicate prompt detection."""
text = "Duplicate test prompt"
hash_val = generate_prompt_hash(text)
# Save original prompt
prompt_id1 = self.db.save_prompt(
text=text,
category="test",
prompt_hash=hash_val
)
# Try to save the same prompt again
existing = self.db.get_prompt_by_hash(hash_val)
self.assertIsNotNone(existing)
self.assertEqual(existing['id'], prompt_id1)
def test_search_functionality(self):
"""Test prompt search functionality."""
# Save multiple prompts
prompts = [
{
"text": "Beautiful landscape with mountains",
"category": "landscape",
"tags": ["nature", "mountains"],
"rating": 5
},
{
"text": "Portrait of a woman",
"category": "portrait",
"tags": ["people", "woman"],
"rating": 4
},
{
"text": "Abstract art piece",
"category": "abstract",
"tags": ["art", "abstract"],
"rating": 3
}
]
for prompt_data in prompts:
self.db.save_prompt(
text=prompt_data["text"],
category=prompt_data["category"],
tags=prompt_data["tags"],
rating=prompt_data["rating"],
prompt_hash=generate_prompt_hash(prompt_data["text"])
)
# Test text search
results = self.db.search_prompts(text="landscape")
self.assertEqual(len(results), 1)
self.assertIn("landscape", results[0]['text'])
# Test category search
results = self.db.search_prompts(category="portrait")
self.assertEqual(len(results), 1)
self.assertEqual(results[0]['category'], "portrait")
# Test rating search
results = self.db.search_prompts(rating_min=4)
self.assertEqual(len(results), 2) # Rating 4 and 5
# Test tag search
results = self.db.search_prompts(tags=["nature"])
self.assertEqual(len(results), 1)
self.assertIn("nature", results[0]['tags'])
class TestNodeIntegration(unittest.TestCase):
"""Test the actual node integration (mocked)."""
def setUp(self):
"""Set up test fixtures."""
# Use a temporary database for testing
self.temp_db = tempfile.NamedTemporaryFile(delete=False, suffix='.db')
self.temp_db.close()
def tearDown(self):
"""Clean up test fixtures."""
# Remove temporary database
if os.path.exists(self.temp_db.name):
os.unlink(self.temp_db.name)
@patch('prompt_manager_base.PromptDatabase')
def test_node_encode_function(self, mock_db_class):
"""Test the node's encode function with mocked dependencies."""
# Mock the database
mock_db = Mock()
mock_db_class.return_value = mock_db
mock_db.save_prompt.return_value = 1
mock_db.get_prompt_by_hash.return_value = None
# Mock CLIP model
mock_clip = Mock()
mock_clip.tokenize.return_value = "mock_tokens"
mock_clip.encode_from_tokens_scheduled.return_value = "mock_conditioning"
# Import and test the node
from prompt_manager import PromptManager
node = PromptManager()
# Test encoding
result = node.encode_prompt(
clip=mock_clip,
text="Test prompt",
search_text=""
)
# Verify CLIP was called correctly
mock_clip.tokenize.assert_called_once_with("Test prompt")
mock_clip.encode_from_tokens_scheduled.assert_called_once_with("mock_tokens")
# Verify result tuple (conditioning, prompt_text)
self.assertEqual(result[1], "Test prompt")
def test_node_input_types(self):
"""Test the node's input type definitions."""
from prompt_manager import PromptManager
input_types = PromptManager.INPUT_TYPES()
# Check required inputs
self.assertIn("text", input_types["required"])
self.assertIn("clip", input_types["required"])
# Check optional inputs
self.assertIn("search_text", input_types["optional"])
if __name__ == '__main__':
unittest.main()