style: apply Black formatter (line-length=88) to all Python files
Automated formatting pass across 30 files to establish consistent code style enforced by CI. No logic changes.
This commit is contained in:
+62
-67
@@ -9,6 +9,7 @@ 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
|
||||
@@ -18,61 +19,61 @@ from utils.validators import validate_prompt_text, validate_rating, validate_tag
|
||||
|
||||
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 = 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
|
||||
@@ -80,17 +81,17 @@ class TestBasicFunctionality(unittest.TestCase):
|
||||
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
|
||||
@@ -98,17 +99,17 @@ class TestBasicFunctionality(unittest.TestCase):
|
||||
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
|
||||
@@ -118,38 +119,36 @@ class TestBasicFunctionality(unittest.TestCase):
|
||||
tags=["test", "database"],
|
||||
rating=4,
|
||||
notes="Test notes",
|
||||
prompt_hash=generate_prompt_hash("Test prompt for database")
|
||||
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")
|
||||
|
||||
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
|
||||
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)
|
||||
|
||||
self.assertEqual(existing["id"], prompt_id1)
|
||||
|
||||
def test_search_functionality(self):
|
||||
"""Test prompt search functionality."""
|
||||
# Save multiple prompts
|
||||
@@ -158,67 +157,67 @@ class TestBasicFunctionality(unittest.TestCase):
|
||||
"text": "Beautiful landscape with mountains",
|
||||
"category": "landscape",
|
||||
"tags": ["nature", "mountains"],
|
||||
"rating": 5
|
||||
"rating": 5,
|
||||
},
|
||||
{
|
||||
"text": "Portrait of a woman",
|
||||
"category": "portrait",
|
||||
"tags": ["people", "woman"],
|
||||
"rating": 4
|
||||
"rating": 4,
|
||||
},
|
||||
{
|
||||
"text": "Abstract art piece",
|
||||
"category": "abstract",
|
||||
"tags": ["art", "abstract"],
|
||||
"rating": 3
|
||||
}
|
||||
"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"])
|
||||
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'])
|
||||
|
||||
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")
|
||||
|
||||
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'])
|
||||
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 = 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')
|
||||
|
||||
@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
|
||||
@@ -238,11 +237,7 @@ class TestNodeIntegration(unittest.TestCase):
|
||||
node = PromptManager()
|
||||
|
||||
# Test encoding
|
||||
result = node.encode_prompt(
|
||||
clip=mock_clip,
|
||||
text="Test prompt",
|
||||
search_text=""
|
||||
)
|
||||
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")
|
||||
@@ -250,20 +245,20 @@ class TestNodeIntegration(unittest.TestCase):
|
||||
|
||||
# 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()
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Reference in New Issue
Block a user