Restore test_safe_builder_unit.py to upstream version (remove local changes)

This commit is contained in:
Adam Outler
2025-10-24 20:32:50 +00:00
parent 7f74c2d6f3
commit 32f9111f66
+88 -90
View File
@@ -4,15 +4,15 @@ This test file has minimal dependencies to ensure it can run in any environment.
""" """
import sys import sys
import unittest
import re import re
import pytest
from unittest.mock import Mock, patch from unittest.mock import Mock, patch
# Mock the logger module to avoid dependency issues # Mock the logger module to avoid dependency issues
sys.modules['logger'] = Mock() sys.modules['logger'] = Mock()
# Standalone version of SafeConditionBuilder for testing # Standalone version of SafeConditionBuilder for testing
class SafeConditionBuilder: class TestSafeConditionBuilder:
""" """
Test version of SafeConditionBuilder with mock logger. Test version of SafeConditionBuilder with mock logger.
""" """
@@ -152,90 +152,84 @@ class SafeConditionBuilder:
return "", {} return "", {}
@pytest.fixture class TestSafeConditionBuilderSecurity(unittest.TestCase):
def builder(): """Test cases for the SafeConditionBuilder security functionality."""
"""Fixture to provide a fresh SafeConditionBuilder instance for each test."""
return SafeConditionBuilder()
def setUp(self):
"""Set up test fixtures before each test method."""
self.builder = TestSafeConditionBuilder()
def test_initialization(builder): def test_initialization(self):
"""Test that SafeConditionBuilder initializes correctly.""" """Test that SafeConditionBuilder initializes correctly."""
assert isinstance(builder, SafeConditionBuilder) self.assertIsInstance(self.builder, TestSafeConditionBuilder)
assert builder.param_counter == 0 self.assertEqual(self.builder.param_counter, 0)
assert builder.parameters == {} self.assertEqual(self.builder.parameters, {})
def test_sanitize_string(self):
def test_sanitize_string(builder):
"""Test string sanitization functionality.""" """Test string sanitization functionality."""
# Test normal string # Test normal string
result = builder._sanitize_string("normal string") result = self.builder._sanitize_string("normal string")
assert result == "normal string" self.assertEqual(result, "normal string")
# Test s-quote replacement # Test s-quote replacement
result = builder._sanitize_string("test{s-quote}value") result = self.builder._sanitize_string("test{s-quote}value")
assert result == "test'value" self.assertEqual(result, "test'value")
# Test control character removal # Test control character removal
result = builder._sanitize_string("test\x00\x01string") result = self.builder._sanitize_string("test\x00\x01string")
assert result == "teststring" self.assertEqual(result, "teststring")
# Test excessive whitespace # Test excessive whitespace
result = builder._sanitize_string(" test string ") result = self.builder._sanitize_string(" test string ")
assert result == "test string" self.assertEqual(result, "test string")
def test_validate_column_name(self):
def test_validate_column_name(builder):
"""Test column name validation against whitelist.""" """Test column name validation against whitelist."""
# Valid columns # Valid columns
assert builder._validate_column_name('eve_MAC') self.assertTrue(self.builder._validate_column_name('eve_MAC'))
assert builder._validate_column_name('devName') self.assertTrue(self.builder._validate_column_name('devName'))
assert builder._validate_column_name('eve_EventType') self.assertTrue(self.builder._validate_column_name('eve_EventType'))
# Invalid columns # Invalid columns
assert not builder._validate_column_name('malicious_column') self.assertFalse(self.builder._validate_column_name('malicious_column'))
assert not builder._validate_column_name('drop_table') self.assertFalse(self.builder._validate_column_name('drop_table'))
assert not builder._validate_column_name('user_input') self.assertFalse(self.builder._validate_column_name('user_input'))
def test_validate_operator(self):
def test_validate_operator(builder):
"""Test operator validation against whitelist.""" """Test operator validation against whitelist."""
# Valid operators # Valid operators
assert builder._validate_operator('=') self.assertTrue(self.builder._validate_operator('='))
assert builder._validate_operator('LIKE') self.assertTrue(self.builder._validate_operator('LIKE'))
assert builder._validate_operator('IN') self.assertTrue(self.builder._validate_operator('IN'))
# Invalid operators # Invalid operators
assert not builder._validate_operator('UNION') self.assertFalse(self.builder._validate_operator('UNION'))
assert not builder._validate_operator('DROP') self.assertFalse(self.builder._validate_operator('DROP'))
assert not builder._validate_operator('EXEC') self.assertFalse(self.builder._validate_operator('EXEC'))
def test_build_simple_condition_valid(self):
def test_build_simple_condition_valid(builder):
"""Test building valid simple conditions.""" """Test building valid simple conditions."""
sql, params = builder._build_simple_condition('AND', 'devName', '=', 'TestDevice') sql, params = self.builder._build_simple_condition('AND', 'devName', '=', 'TestDevice')
assert 'AND devName = :param_' in sql self.assertIn('AND devName = :param_', sql)
assert len(params) == 1 self.assertEqual(len(params), 1)
assert 'TestDevice' in params.values() self.assertIn('TestDevice', params.values())
def test_build_simple_condition_invalid_column(self):
def test_build_simple_condition_invalid_column(builder):
"""Test that invalid column names are rejected.""" """Test that invalid column names are rejected."""
with pytest.raises(ValueError) as exc_info: with self.assertRaises(ValueError) as context:
builder._build_simple_condition('AND', 'invalid_column', '=', 'value') self.builder._build_simple_condition('AND', 'invalid_column', '=', 'value')
assert 'Invalid column name' in str(exc_info.value) self.assertIn('Invalid column name', str(context.exception))
def test_build_simple_condition_invalid_operator(self):
def test_build_simple_condition_invalid_operator(builder):
"""Test that invalid operators are rejected.""" """Test that invalid operators are rejected."""
with pytest.raises(ValueError) as exc_info: with self.assertRaises(ValueError) as context:
builder._build_simple_condition('AND', 'devName', 'UNION', 'value') self.builder._build_simple_condition('AND', 'devName', 'UNION', 'value')
assert 'Invalid operator' in str(exc_info.value) self.assertIn('Invalid operator', str(context.exception))
def test_sql_injection_attempts(self):
def test_sql_injection_attempts(builder):
"""Test that various SQL injection attempts are blocked.""" """Test that various SQL injection attempts are blocked."""
injection_attempts = [ injection_attempts = [
"'; DROP TABLE Devices; --", "'; DROP TABLE Devices; --",
@@ -246,39 +240,37 @@ def test_sql_injection_attempts(builder):
] ]
for injection in injection_attempts: for injection in injection_attempts:
with pytest.raises(ValueError): with self.subTest(injection=injection):
builder.build_safe_condition(f"AND devName = '{injection}'") with self.assertRaises(ValueError):
self.builder.build_safe_condition(f"AND devName = '{injection}'")
def test_legacy_condition_compatibility(self):
def test_legacy_condition_compatibility(builder):
"""Test backward compatibility with legacy condition formats.""" """Test backward compatibility with legacy condition formats."""
# Test simple condition # Test simple condition
sql, params = builder.get_safe_condition_legacy("AND devName = 'TestDevice'") sql, params = self.builder.get_safe_condition_legacy("AND devName = 'TestDevice'")
assert 'devName' in sql self.assertIn('devName', sql)
assert 'TestDevice' in params.values() self.assertIn('TestDevice', params.values())
# Test empty condition # Test empty condition
sql, params = builder.get_safe_condition_legacy("") sql, params = self.builder.get_safe_condition_legacy("")
assert sql == "" self.assertEqual(sql, "")
assert params == {} self.assertEqual(params, {})
# Test invalid condition returns empty # Test invalid condition returns empty
sql, params = builder.get_safe_condition_legacy("INVALID SQL INJECTION") sql, params = self.builder.get_safe_condition_legacy("INVALID SQL INJECTION")
assert sql == "" self.assertEqual(sql, "")
assert params == {} self.assertEqual(params, {})
def test_parameter_generation(self):
def test_parameter_generation(builder):
"""Test that parameters are generated correctly.""" """Test that parameters are generated correctly."""
# Test single parameter # Test multiple parameters
sql, params = builder.build_safe_condition("AND devName = 'Device1'") sql1, params1 = self.builder.build_safe_condition("AND devName = 'Device1'")
sql2, params2 = self.builder.build_safe_condition("AND devName = 'Device2'")
# Should have 1 parameter # Each should have unique parameter names
assert len(params) == 1 self.assertNotEqual(list(params1.keys())[0], list(params2.keys())[0])
assert 'param_1' in params
def test_xss_prevention(self):
def test_xss_prevention(builder):
"""Test that XSS-like payloads in device names are handled safely.""" """Test that XSS-like payloads in device names are handled safely."""
xss_payloads = [ xss_payloads = [
"<script>alert('xss')</script>", "<script>alert('xss')</script>",
@@ -288,32 +280,32 @@ def test_xss_prevention(builder):
] ]
for payload in xss_payloads: for payload in xss_payloads:
with self.subTest(payload=payload):
# Should either process safely or reject # Should either process safely or reject
try: try:
sql, params = builder.build_safe_condition(f"AND devName = '{payload}'") sql, params = self.builder.build_safe_condition(f"AND devName = '{payload}'")
# If processed, should be parameterized # If processed, should be parameterized
assert ':' in sql self.assertIn(':', sql)
assert payload in params.values() self.assertIn(payload, params.values())
except ValueError: except ValueError:
# Rejection is also acceptable for safety # Rejection is also acceptable for safety
pass pass
def test_unicode_handling(self):
def test_unicode_handling(builder):
"""Test that Unicode characters are handled properly.""" """Test that Unicode characters are handled properly."""
unicode_strings = [ unicode_strings = [
"Ülrichs Device", "Ülrich's Device",
"Café Network", "Café Network",
"测试设备", "测试设备",
"Устройство" "Устройство"
] ]
for unicode_str in unicode_strings: for unicode_str in unicode_strings:
sql, params = builder.build_safe_condition(f"AND devName = '{unicode_str}'") with self.subTest(unicode_str=unicode_str):
assert unicode_str in params.values() sql, params = self.builder.build_safe_condition(f"AND devName = '{unicode_str}'")
self.assertIn(unicode_str, params.values())
def test_edge_cases(self):
def test_edge_cases(builder):
"""Test edge cases and boundary conditions.""" """Test edge cases and boundary conditions."""
edge_cases = [ edge_cases = [
"", # Empty string "", # Empty string
@@ -324,10 +316,16 @@ def test_edge_cases(builder):
] ]
for case in edge_cases: for case in edge_cases:
with self.subTest(case=case):
try: try:
sql, params = builder.get_safe_condition_legacy(case) sql, params = self.builder.get_safe_condition_legacy(case)
# Should either return valid result or empty safe result # Should either return valid result or empty safe result
assert isinstance(sql, str) self.assertIsInstance(sql, str)
assert isinstance(params, dict) self.assertIsInstance(params, dict)
except Exception: except Exception:
pytest.fail(f"Unexpected exception for edge case: {case}") self.fail(f"Unexpected exception for edge case: {case}")
if __name__ == '__main__':
# Run the test suite
unittest.main(verbosity=2)