Devcontainer setup

This commit is contained in:
Adam Outler
2025-10-23 23:33:04 +00:00
parent 3b7830b922
commit edd5bd27b0
6 changed files with 430 additions and 422 deletions
+2 -2
View File
@@ -210,7 +210,7 @@ HEALTHCHECK --interval=30s --timeout=10s --start-period=60s --retries=3 \
FROM runner AS netalertx-devcontainer FROM runner AS netalertx-devcontainer
ENV INSTALL_DIR=/app ENV INSTALL_DIR=/app
ENV PYTHONPATH=/workspaces/NetAlertX/test:/workspaces/NetAlertX/server:/app:/app/server:/opt/venv/lib/python3.12/site-packages ENV PYTHONPATH=/workspaces/NetAlertX/test:/workspaces/NetAlertX/server:/app:/app/server:/opt/venv/lib/python3.12/site-packages:/usr/lib/python3.12/site-packages
ENV PATH=/services:${PATH} ENV PATH=/services:${PATH}
ENV PHP_INI_SCAN_DIR=/services/config/php/conf.d:/etc/php83/conf.d ENV PHP_INI_SCAN_DIR=/services/config/php/conf.d:/etc/php83/conf.d
ENV LISTEN_ADDR=0.0.0.0 ENV LISTEN_ADDR=0.0.0.0
@@ -231,7 +231,7 @@ RUN mkdir /workspaces && \
install -d -o netalertx -g netalertx -m 777 /services/run/logs && \ install -d -o netalertx -g netalertx -m 777 /services/run/logs && \
install -d -o netalertx -g netalertx -m 777 /app/run/tmp/client_body && \ install -d -o netalertx -g netalertx -m 777 /app/run/tmp/client_body && \
sed -i -e 's|:/app:|:/workspaces:|' /etc/passwd && \ sed -i -e 's|:/app:|:/workspaces:|' /etc/passwd && \
find /opt/venv -type d -exec chmod o+rw {} \; find /opt/venv -type d -exec chmod o+rwx {} \;
USER netalertx USER netalertx
ENTRYPOINT ["/bin/sh","-c","sleep infinity"] ENTRYPOINT ["/bin/sh","-c","sleep infinity"]
+1 -1
View File
@@ -43,7 +43,7 @@
} }
}, },
"postCreateCommand": "pip install pytest docker", "postCreateCommand": "/opt/venv/bin/pip3 install pytest docker debugpy",
"postStartCommand": "${containerWorkspaceFolder}/.devcontainer/scripts/setup.sh", "postStartCommand": "${containerWorkspaceFolder}/.devcontainer/scripts/setup.sh",
"customizations": { "customizations": {
@@ -7,7 +7,7 @@
FROM runner AS netalertx-devcontainer FROM runner AS netalertx-devcontainer
ENV INSTALL_DIR=/app ENV INSTALL_DIR=/app
ENV PYTHONPATH=/workspaces/NetAlertX/test:/workspaces/NetAlertX/server:/app:/app/server:/opt/venv/lib/python3.12/site-packages ENV PYTHONPATH=/workspaces/NetAlertX/test:/workspaces/NetAlertX/server:/app:/app/server:/opt/venv/lib/python3.12/site-packages:/usr/lib/python3.12/site-packages
ENV PATH=/services:${PATH} ENV PATH=/services:${PATH}
ENV PHP_INI_SCAN_DIR=/services/config/php/conf.d:/etc/php83/conf.d ENV PHP_INI_SCAN_DIR=/services/config/php/conf.d:/etc/php83/conf.d
ENV LISTEN_ADDR=0.0.0.0 ENV LISTEN_ADDR=0.0.0.0
@@ -28,7 +28,7 @@ RUN mkdir /workspaces && \
install -d -o netalertx -g netalertx -m 777 /services/run/logs && \ install -d -o netalertx -g netalertx -m 777 /services/run/logs && \
install -d -o netalertx -g netalertx -m 777 /app/run/tmp/client_body && \ install -d -o netalertx -g netalertx -m 777 /app/run/tmp/client_body && \
sed -i -e 's|:/app:|:/workspaces:|' /etc/passwd && \ sed -i -e 's|:/app:|:/workspaces:|' /etc/passwd && \
find /opt/venv -type d -exec chmod o+rw {} \; find /opt/venv -type d -exec chmod o+rwx {} \;
USER netalertx USER netalertx
ENTRYPOINT ["/bin/sh","-c","sleep infinity"] ENTRYPOINT ["/bin/sh","-c","sleep infinity"]
+5 -3
View File
@@ -57,8 +57,9 @@ NETALERTX_DOCKER_ERROR_CHECK=0
# Run all pre-startup checks to validate container environment and dependencies # Run all pre-startup checks to validate container environment and dependencies
echo "Startup pre-checks" if [ ${NETALERTX_DEBUG != 1} ]; then
for script in ${SYSTEM_SERVICES_SCRIPTS}/check-*.sh; do echo "Startup pre-checks"
for script in ${SYSTEM_SERVICES_SCRIPTS}/check-*.sh; do
script_name=$(basename "$script" | sed 's/^check-//;s/\.sh$//;s/-/ /g') script_name=$(basename "$script" | sed 's/^check-//;s/\.sh$//;s/-/ /g')
echo " --> ${script_name}" echo " --> ${script_name}"
@@ -70,7 +71,8 @@ for script in ${SYSTEM_SERVICES_SCRIPTS}/check-*.sh; do
echo exit code ${NETALERTX_DOCKER_ERROR_CHECK} from ${script} echo exit code ${NETALERTX_DOCKER_ERROR_CHECK} from ${script}
exit ${NETALERTX_DOCKER_ERROR_CHECK} exit ${NETALERTX_DOCKER_ERROR_CHECK}
fi fi
done done
fi
# Exit after checks if in check-only mode (for testing) # Exit after checks if in check-only mode (for testing)
if [ "${NETALERTX_CHECK_ONLY:-0}" -eq 1 ]; then if [ "${NETALERTX_CHECK_ONLY:-0}" -eq 1 ]; then
+125 -121
View File
@@ -5,26 +5,25 @@ Tests the fix for Issue #1210 - compound conditions with multiple AND/OR clauses
""" """
import sys import sys
import unittest import pytest
from unittest.mock import MagicMock from unittest.mock import MagicMock
# Mock the logger module before importing SafeConditionBuilder # Mock the logger module before importing SafeConditionBuilder
sys.modules['logger'] = MagicMock() sys.modules['logger'] = MagicMock()
# Add parent directory to path for imports # Add parent directory to path for imports
sys.path.insert(0, '/tmp/netalertx_hotfix/server/db') sys.path.insert(0, '/workspaces/NetAlertX')
from sql_safe_builder import SafeConditionBuilder from server.db.sql_safe_builder import SafeConditionBuilder
class TestCompoundConditions(unittest.TestCase): @pytest.fixture
"""Test compound condition parsing functionality.""" def builder():
def setUp(self):
"""Create a fresh builder instance for each test.""" """Create a fresh builder instance for each test."""
self.builder = SafeConditionBuilder() return SafeConditionBuilder()
def test_user_failing_filter_six_and_clauses(self):
def test_user_failing_filter_six_and_clauses(builder):
"""Test the exact user-reported failing filter from Issue #1210.""" """Test the exact user-reported failing filter from Issue #1210."""
condition = ( condition = (
"AND devLastIP NOT LIKE '192.168.50.%' " "AND devLastIP NOT LIKE '192.168.50.%' "
@@ -35,292 +34,297 @@ class TestCompoundConditions(unittest.TestCase):
"AND devLastIP NOT LIKE '192.168.70.4'" "AND devLastIP NOT LIKE '192.168.70.4'"
) )
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# Should successfully parse # Should successfully parse
self.assertIsNotNone(sql) assert sql is not None
self.assertIsNotNone(params) assert params is not None
# Should have 6 parameters (one per clause) # Should have 6 parameters (one per clause)
self.assertEqual(len(params), 6) assert len(params) == 6
# Should contain all 6 AND operators # Should contain all 6 AND operators
self.assertEqual(sql.count('AND'), 6) assert sql.count('AND') == 6
# Should contain all 6 NOT LIKE operators # Should contain all 6 NOT LIKE operators
self.assertEqual(sql.count('NOT LIKE'), 6) assert sql.count('NOT LIKE') == 6
# Should have 6 parameter placeholders # Should have 6 parameter placeholders
self.assertEqual(sql.count(':param_'), 6) assert sql.count(':param_') == 6
# Verify all IP patterns are in parameters # Verify all IP patterns are in parameters
param_values = list(params.values()) param_values = list(params.values())
self.assertIn('192.168.50.%', param_values) assert '192.168.50.%' in param_values
self.assertIn('192.168.60.%', param_values) assert '192.168.60.%' in param_values
self.assertIn('192.168.70.2', param_values) assert '192.168.70.2' in param_values
self.assertIn('192.168.70.5', param_values) assert '192.168.70.5' in param_values
self.assertIn('192.168.70.3', param_values) assert '192.168.70.3' in param_values
self.assertIn('192.168.70.4', param_values) assert '192.168.70.4' in param_values
def test_multiple_and_clauses_simple(self):
def test_multiple_and_clauses_simple(builder):
"""Test multiple AND clauses with simple equality operators.""" """Test multiple AND clauses with simple equality operators."""
condition = "AND devName = 'Device1' AND devVendor = 'Apple' AND devFavorite = '1'" condition = "AND devName = 'Device1' AND devVendor = 'Apple' AND devFavorite = '1'"
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# Should have 3 parameters # Should have 3 parameters
self.assertEqual(len(params), 3) assert len(params) == 3
# Should have 3 AND operators # Should have 3 AND operators
self.assertEqual(sql.count('AND'), 3) assert sql.count('AND') == 3
# Verify all values are parameterized # Verify all values are parameterized
param_values = list(params.values()) param_values = list(params.values())
self.assertIn('Device1', param_values) assert 'Device1' in param_values
self.assertIn('Apple', param_values) assert 'Apple' in param_values
self.assertIn('1', param_values) assert '1' in param_values
def test_multiple_or_clauses(self):
def test_multiple_or_clauses(builder):
"""Test multiple OR clauses.""" """Test multiple OR clauses."""
condition = "OR devName = 'Device1' OR devName = 'Device2' OR devName = 'Device3'" condition = "OR devName = 'Device1' OR devName = 'Device2' OR devName = 'Device3'"
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# Should have 3 parameters # Should have 3 parameters
self.assertEqual(len(params), 3) assert len(params) == 3
# Should have 3 OR operators # Should have 3 OR operators
self.assertEqual(sql.count('OR'), 3) assert sql.count('OR') == 3
# Verify all device names are parameterized # Verify all device names are parameterized
param_values = list(params.values()) param_values = list(params.values())
self.assertIn('Device1', param_values) assert 'Device1' in param_values
self.assertIn('Device2', param_values) assert 'Device2' in param_values
self.assertIn('Device3', param_values) assert 'Device3' in param_values
def test_mixed_and_or_clauses(self): def test_mixed_and_or_clauses(builder):
"""Test mixed AND/OR logical operators.""" """Test mixed AND/OR logical operators."""
condition = "AND devName = 'Device1' OR devName = 'Device2' AND devFavorite = '1'" condition = "AND devName = 'Device1' OR devName = 'Device2' AND devFavorite = '1'"
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# Should have 3 parameters # Should have 3 parameters
self.assertEqual(len(params), 3) assert len(params) == 3
# Should preserve the logical operator order # Should preserve the logical operator order
self.assertIn('AND', sql) assert 'AND' in sql
self.assertIn('OR', sql) assert 'OR' in sql
# Verify all values are parameterized # Verify all values are parameterized
param_values = list(params.values()) param_values = list(params.values())
self.assertIn('Device1', param_values) assert 'Device1' in param_values
self.assertIn('Device2', param_values) assert 'Device2' in param_values
self.assertIn('1', param_values) assert '1' in param_values
def test_single_condition_backward_compatibility(self):
def test_single_condition_backward_compatibility(builder):
"""Test that single conditions still work (backward compatibility).""" """Test that single conditions still work (backward compatibility)."""
condition = "AND devName = 'TestDevice'" condition = "AND devName = 'TestDevice'"
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# Should have 1 parameter # Should have 1 parameter
self.assertEqual(len(params), 1) assert len(params) == 1
# Should match expected format # Should match expected format
self.assertIn('AND devName = :param_', sql) assert 'AND devName = :param_' in sql
# Parameter should contain the value # Parameter should contain the value
self.assertIn('TestDevice', params.values()) assert 'TestDevice' in params.values()
def test_single_condition_like_operator(self):
def test_single_condition_like_operator(builder):
"""Test single LIKE condition for backward compatibility.""" """Test single LIKE condition for backward compatibility."""
condition = "AND devComments LIKE '%important%'" condition = "AND devComments LIKE '%important%'"
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# Should have 1 parameter # Should have 1 parameter
self.assertEqual(len(params), 1) assert len(params) == 1
# Should contain LIKE operator # Should contain LIKE operator
self.assertIn('LIKE', sql) assert 'LIKE' in sql
# Parameter should contain the pattern # Parameter should contain the pattern
self.assertIn('%important%', params.values()) assert '%important%' in params.values()
def test_compound_with_like_patterns(self):
def test_compound_with_like_patterns(builder):
"""Test compound conditions with LIKE patterns.""" """Test compound conditions with LIKE patterns."""
condition = "AND devLastIP LIKE '192.168.%' AND devVendor LIKE '%Apple%'" condition = "AND devLastIP LIKE '192.168.%' AND devVendor LIKE '%Apple%'"
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# Should have 2 parameters # Should have 2 parameters
self.assertEqual(len(params), 2) assert len(params) == 2
# Should have 2 LIKE operators # Should have 2 LIKE operators
self.assertEqual(sql.count('LIKE'), 2) assert sql.count('LIKE') == 2
# Verify patterns are parameterized # Verify patterns are parameterized
param_values = list(params.values()) param_values = list(params.values())
self.assertIn('192.168.%', param_values) assert '192.168.%' in param_values
self.assertIn('%Apple%', param_values) assert '%Apple%' in param_values
def test_compound_with_inequality_operators(self):
def test_compound_with_inequality_operators(builder):
"""Test compound conditions with various inequality operators.""" """Test compound conditions with various inequality operators."""
condition = "AND eve_DateTime > '2024-01-01' AND eve_DateTime < '2024-12-31'" condition = "AND eve_DateTime > '2024-01-01' AND eve_DateTime < '2024-12-31'"
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# Should have 2 parameters # Should have 2 parameters
self.assertEqual(len(params), 2) assert len(params) == 2
# Should have both operators # Should have both operators
self.assertIn('>', sql) assert '>' in sql
self.assertIn('<', sql) assert '<' in sql
# Verify dates are parameterized # Verify dates are parameterized
param_values = list(params.values()) param_values = list(params.values())
self.assertIn('2024-01-01', param_values) assert '2024-01-01' in param_values
self.assertIn('2024-12-31', param_values) assert '2024-12-31' in param_values
def test_empty_condition(self):
def test_empty_condition(builder):
"""Test empty condition string.""" """Test empty condition string."""
condition = "" condition = ""
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# Should return empty results # Should return empty results
self.assertEqual(sql, "") assert sql == ""
self.assertEqual(params, {}) assert params == {}
def test_whitespace_only_condition(self):
def test_whitespace_only_condition(builder):
"""Test condition with only whitespace.""" """Test condition with only whitespace."""
condition = " \t\n " condition = " \t\n "
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# Should return empty results # Should return empty results
self.assertEqual(sql, "") assert sql == ""
self.assertEqual(params, {}) assert params == {}
def test_invalid_column_name_rejected(self):
def test_invalid_column_name_rejected(builder):
"""Test that invalid column names are rejected.""" """Test that invalid column names are rejected."""
condition = "AND malicious_column = 'value'" condition = "AND malicious_column = 'value'"
with self.assertRaises(ValueError): with pytest.raises(ValueError):
self.builder.build_safe_condition(condition) builder.build_safe_condition(condition)
def test_invalid_operator_rejected(self):
def test_invalid_operator_rejected(builder):
"""Test that invalid operators are rejected.""" """Test that invalid operators are rejected."""
condition = "AND devName EXECUTE 'DROP TABLE'" condition = "AND devName EXECUTE 'DROP TABLE'"
with self.assertRaises(ValueError): with pytest.raises(ValueError):
self.builder.build_safe_condition(condition) builder.build_safe_condition(condition)
def test_sql_injection_attempt_blocked(self):
def test_sql_injection_attempt_blocked(builder):
"""Test that SQL injection attempts are blocked.""" """Test that SQL injection attempts are blocked."""
condition = "AND devName = 'value'; DROP TABLE devices; --" condition = "AND devName = 'value'; DROP TABLE devices; --"
# Should either reject or sanitize the dangerous input # Should either reject or sanitize the dangerous input
# The semicolon and comment should not appear in the final SQL # The semicolon and comment should not appear in the final SQL
try: try:
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# If it doesn't raise an error, it should sanitize the input # If it doesn't raise an error, it should sanitize the input
self.assertNotIn('DROP', sql.upper()) assert 'DROP' not in sql.upper()
self.assertNotIn(';', sql) assert ';' not in sql
except ValueError: except ValueError:
# Rejection is also acceptable # Rejection is also acceptable
pass pass
def test_quoted_string_with_spaces(self):
def test_quoted_string_with_spaces(builder):
"""Test that quoted strings with spaces are handled correctly.""" """Test that quoted strings with spaces are handled correctly."""
condition = "AND devName = 'My Device Name' AND devComments = 'Has spaces here'" condition = "AND devName = 'My Device Name' AND devComments = 'Has spaces here'"
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# Should have 2 parameters # Should have 2 parameters
self.assertEqual(len(params), 2) assert len(params) == 2
# Verify values with spaces are preserved # Verify values with spaces are preserved
param_values = list(params.values()) param_values = list(params.values())
self.assertIn('My Device Name', param_values) assert 'My Device Name' in param_values
self.assertIn('Has spaces here', param_values) assert 'Has spaces here' in param_values
def test_compound_condition_with_not_equal(self):
def test_compound_condition_with_not_equal(builder):
"""Test compound conditions with != operator.""" """Test compound conditions with != operator."""
condition = "AND devName != 'Device1' AND devVendor != 'Unknown'" condition = "AND devName != 'Device1' AND devVendor != 'Unknown'"
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# Should have 2 parameters # Should have 2 parameters
self.assertEqual(len(params), 2) assert len(params) == 2
# Should have != operators (or converted to <>) # Should have != operators (or converted to <>)
self.assertTrue('!=' in sql or '<>' in sql) assert '!=' in sql or '<>' in sql
# Verify values are parameterized # Verify values are parameterized
param_values = list(params.values()) param_values = list(params.values())
self.assertIn('Device1', param_values) assert 'Device1' in param_values
self.assertIn('Unknown', param_values) assert 'Unknown' in param_values
def test_very_long_compound_condition(self):
def test_very_long_compound_condition(builder):
"""Test handling of very long compound conditions (10+ clauses).""" """Test handling of very long compound conditions (10+ clauses)."""
clauses = [] clauses = []
for i in range(10): for i in range(10):
clauses.append(f"AND devName != 'Device{i}'") clauses.append(f"AND devName != 'Device{i}'")
condition = " ".join(clauses) condition = " ".join(clauses)
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# Should have 10 parameters # Should have 10 parameters
self.assertEqual(len(params), 10) assert len(params) == 10
# Should have 10 AND operators # Should have 10 AND operators
self.assertEqual(sql.count('AND'), 10) assert sql.count('AND') == 10
# Verify all device names are parameterized # Verify all device names are parameterized
param_values = list(params.values()) param_values = list(params.values())
for i in range(10): for i in range(10):
self.assertIn(f'Device{i}', param_values) assert f'Device{i}' in param_values
class TestParameterGeneration(unittest.TestCase): def test_parameters_have_unique_names(builder):
"""Test parameter generation and naming."""
def setUp(self):
"""Create a fresh builder instance for each test."""
self.builder = SafeConditionBuilder()
def test_parameters_have_unique_names(self):
"""Test that all parameters get unique names.""" """Test that all parameters get unique names."""
condition = "AND devName = 'A' AND devName = 'B' AND devName = 'C'" condition = "AND devName = 'A' AND devName = 'B' AND devName = 'C'"
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# All parameter names should be unique # All parameter names should be unique
param_names = list(params.keys()) param_names = list(params.keys())
self.assertEqual(len(param_names), len(set(param_names))) assert len(param_names) == len(set(param_names))
def test_parameter_values_match_condition(self):
def test_parameter_values_match_condition(builder):
"""Test that parameter values correctly match the condition values.""" """Test that parameter values correctly match the condition values."""
condition = "AND devLastIP NOT LIKE '192.168.1.%' AND devLastIP NOT LIKE '10.0.0.%'" condition = "AND devLastIP NOT LIKE '192.168.1.%' AND devLastIP NOT LIKE '10.0.0.%'"
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# Should have exactly the values from the condition # Should have exactly the values from the condition
param_values = sorted(params.values()) param_values = sorted(params.values())
expected_values = sorted(['192.168.1.%', '10.0.0.%']) expected_values = sorted(['192.168.1.%', '10.0.0.%'])
self.assertEqual(param_values, expected_values) assert param_values == expected_values
def test_parameters_referenced_in_sql(self):
def test_parameters_referenced_in_sql(builder):
"""Test that all parameters are actually referenced in the SQL.""" """Test that all parameters are actually referenced in the SQL."""
condition = "AND devName = 'Device1' AND devVendor = 'Apple'" condition = "AND devName = 'Device1' AND devVendor = 'Apple'"
sql, params = self.builder.build_safe_condition(condition) sql, params = builder.build_safe_condition(condition)
# Every parameter should appear in the SQL # Every parameter should appear in the SQL
for param_name in params.keys(): for param_name in params.keys():
self.assertIn(f':{param_name}', sql) assert f':{param_name}' in sql
if __name__ == '__main__':
unittest.main()
+90 -88
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 TestSafeConditionBuilder: class SafeConditionBuilder:
""" """
Test version of SafeConditionBuilder with mock logger. Test version of SafeConditionBuilder with mock logger.
""" """
@@ -152,84 +152,90 @@ class TestSafeConditionBuilder:
return "", {} return "", {}
class TestSafeConditionBuilderSecurity(unittest.TestCase): @pytest.fixture
"""Test cases for the SafeConditionBuilder security functionality.""" def builder():
"""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(self): def test_initialization(builder):
"""Test that SafeConditionBuilder initializes correctly.""" """Test that SafeConditionBuilder initializes correctly."""
self.assertIsInstance(self.builder, TestSafeConditionBuilder) assert isinstance(builder, SafeConditionBuilder)
self.assertEqual(self.builder.param_counter, 0) assert builder.param_counter == 0
self.assertEqual(self.builder.parameters, {}) assert 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 = self.builder._sanitize_string("normal string") result = builder._sanitize_string("normal string")
self.assertEqual(result, "normal string") assert result == "normal string"
# Test s-quote replacement # Test s-quote replacement
result = self.builder._sanitize_string("test{s-quote}value") result = builder._sanitize_string("test{s-quote}value")
self.assertEqual(result, "test'value") assert result == "test'value"
# Test control character removal # Test control character removal
result = self.builder._sanitize_string("test\x00\x01string") result = builder._sanitize_string("test\x00\x01string")
self.assertEqual(result, "teststring") assert result == "teststring"
# Test excessive whitespace # Test excessive whitespace
result = self.builder._sanitize_string(" test string ") result = builder._sanitize_string(" test string ")
self.assertEqual(result, "test string") assert 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
self.assertTrue(self.builder._validate_column_name('eve_MAC')) assert builder._validate_column_name('eve_MAC')
self.assertTrue(self.builder._validate_column_name('devName')) assert builder._validate_column_name('devName')
self.assertTrue(self.builder._validate_column_name('eve_EventType')) assert builder._validate_column_name('eve_EventType')
# Invalid columns # Invalid columns
self.assertFalse(self.builder._validate_column_name('malicious_column')) assert not builder._validate_column_name('malicious_column')
self.assertFalse(self.builder._validate_column_name('drop_table')) assert not builder._validate_column_name('drop_table')
self.assertFalse(self.builder._validate_column_name('user_input')) assert not 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
self.assertTrue(self.builder._validate_operator('=')) assert builder._validate_operator('=')
self.assertTrue(self.builder._validate_operator('LIKE')) assert builder._validate_operator('LIKE')
self.assertTrue(self.builder._validate_operator('IN')) assert builder._validate_operator('IN')
# Invalid operators # Invalid operators
self.assertFalse(self.builder._validate_operator('UNION')) assert not builder._validate_operator('UNION')
self.assertFalse(self.builder._validate_operator('DROP')) assert not builder._validate_operator('DROP')
self.assertFalse(self.builder._validate_operator('EXEC')) assert not 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 = self.builder._build_simple_condition('AND', 'devName', '=', 'TestDevice') sql, params = builder._build_simple_condition('AND', 'devName', '=', 'TestDevice')
self.assertIn('AND devName = :param_', sql) assert 'AND devName = :param_' in sql
self.assertEqual(len(params), 1) assert len(params) == 1
self.assertIn('TestDevice', params.values()) assert 'TestDevice' in 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 self.assertRaises(ValueError) as context: with pytest.raises(ValueError) as exc_info:
self.builder._build_simple_condition('AND', 'invalid_column', '=', 'value') builder._build_simple_condition('AND', 'invalid_column', '=', 'value')
self.assertIn('Invalid column name', str(context.exception)) assert 'Invalid column name' in str(exc_info.value)
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 self.assertRaises(ValueError) as context: with pytest.raises(ValueError) as exc_info:
self.builder._build_simple_condition('AND', 'devName', 'UNION', 'value') builder._build_simple_condition('AND', 'devName', 'UNION', 'value')
self.assertIn('Invalid operator', str(context.exception)) assert 'Invalid operator' in str(exc_info.value)
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; --",
@@ -240,37 +246,39 @@ class TestSafeConditionBuilderSecurity(unittest.TestCase):
] ]
for injection in injection_attempts: for injection in injection_attempts:
with self.subTest(injection=injection): with pytest.raises(ValueError):
with self.assertRaises(ValueError): builder.build_safe_condition(f"AND devName = '{injection}'")
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 = self.builder.get_safe_condition_legacy("AND devName = 'TestDevice'") sql, params = builder.get_safe_condition_legacy("AND devName = 'TestDevice'")
self.assertIn('devName', sql) assert 'devName' in sql
self.assertIn('TestDevice', params.values()) assert 'TestDevice' in params.values()
# Test empty condition # Test empty condition
sql, params = self.builder.get_safe_condition_legacy("") sql, params = builder.get_safe_condition_legacy("")
self.assertEqual(sql, "") assert sql == ""
self.assertEqual(params, {}) assert params == {}
# Test invalid condition returns empty # Test invalid condition returns empty
sql, params = self.builder.get_safe_condition_legacy("INVALID SQL INJECTION") sql, params = builder.get_safe_condition_legacy("INVALID SQL INJECTION")
self.assertEqual(sql, "") assert sql == ""
self.assertEqual(params, {}) assert params == {}
def test_parameter_generation(self):
def test_parameter_generation(builder):
"""Test that parameters are generated correctly.""" """Test that parameters are generated correctly."""
# Test multiple parameters # Test single parameter
sql1, params1 = self.builder.build_safe_condition("AND devName = 'Device1'") sql, params = builder.build_safe_condition("AND devName = 'Device1'")
sql2, params2 = self.builder.build_safe_condition("AND devName = 'Device2'")
# Each should have unique parameter names # Should have 1 parameter
self.assertNotEqual(list(params1.keys())[0], list(params2.keys())[0]) assert len(params) == 1
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>",
@@ -280,32 +288,32 @@ class TestSafeConditionBuilderSecurity(unittest.TestCase):
] ]
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 = self.builder.build_safe_condition(f"AND devName = '{payload}'") sql, params = builder.build_safe_condition(f"AND devName = '{payload}'")
# If processed, should be parameterized # If processed, should be parameterized
self.assertIn(':', sql) assert ':' in sql
self.assertIn(payload, params.values()) assert payload in 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 = [
"Ülrich's Device", "Ülrichs Device",
"Café Network", "Café Network",
"测试设备", "测试设备",
"Устройство" "Устройство"
] ]
for unicode_str in unicode_strings: for unicode_str in unicode_strings:
with self.subTest(unicode_str=unicode_str): sql, params = builder.build_safe_condition(f"AND devName = '{unicode_str}'")
sql, params = self.builder.build_safe_condition(f"AND devName = '{unicode_str}'") assert unicode_str in params.values()
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
@@ -316,16 +324,10 @@ class TestSafeConditionBuilderSecurity(unittest.TestCase):
] ]
for case in edge_cases: for case in edge_cases:
with self.subTest(case=case):
try: try:
sql, params = self.builder.get_safe_condition_legacy(case) sql, params = builder.get_safe_condition_legacy(case)
# Should either return valid result or empty safe result # Should either return valid result or empty safe result
self.assertIsInstance(sql, str) assert isinstance(sql, str)
self.assertIsInstance(params, dict) assert isinstance(params, dict)
except Exception: except Exception:
self.fail(f"Unexpected exception for edge case: {case}") pytest.fail(f"Unexpected exception for edge case: {case}")
if __name__ == '__main__':
# Run the test suite
unittest.main(verbosity=2)