joelgilbert/NL2SQL
0
1"""2Tests for SQL validation and security features.3"""4 5import pytest6from security.validator import query_validator7 8 9class TestQueryValidator:10 """Test query validation functionality."""11 12 def test_select_only_readonly_mode(self):13 """Test that only SELECT queries pass in readonly mode."""14 15 # Valid SELECT query16 sql = "SELECT * FROM users WHERE active = true"17 is_safe, issues = query_validator.validate_query(sql, mode="readonly")18 assert is_safe is True19 assert len(issues) == 020 21 # Invalid UPDATE query in readonly mode22 sql = "UPDATE users SET active = false WHERE id = 1"23 is_safe, issues = query_validator.validate_query(sql, mode="readonly")24 assert is_safe is False25 assert any("read-only" in issue.lower() for issue in issues)26 27 def test_destructive_operations(self):28 """Test detection of destructive operations."""29 30 destructive_queries = [31 "DROP TABLE users",32 "TRUNCATE TABLE logs",33 "ALTER TABLE users ADD COLUMN test VARCHAR(255)",34 "CREATE TABLE new_table (id INT)"35 ]36 37 for sql in destructive_queries:38 issues = query_validator.check_destructive_operations(sql)39 assert len(issues) > 040 41 def test_sql_injection_patterns(self):42 """Test SQL injection detection."""43 44 # Potential injection attempts45 injection_attempts = [46 "SELECT * FROM users; DROP TABLE users;",47 "SELECT * FROM users UNION SELECT * FROM passwords",48 "SELECT * FROM users WHERE id = 1 OR 1=1 --",49 ]50 51 for sql in injection_attempts:52 has_injection = query_validator.check_sql_injection(sql)53 assert has_injection is True54 55 # Clean query56 clean_sql = "SELECT name, email FROM users WHERE active = true"57 has_injection = query_validator.check_sql_injection(clean_sql)58 assert has_injection is False59 60 def test_where_clause_detection(self):61 """Test WHERE clause detection."""62 63 # Query with WHERE clause64 sql = "DELETE FROM users WHERE inactive_days > 365"65 assert query_validator.has_where_clause(sql) is True66 67 # Query without WHERE clause68 sql = "DELETE FROM users"69 assert query_validator.has_where_clause(sql) is False70 71 def test_empty_query(self):72 """Test validation of empty queries."""73 74 is_safe, issues = query_validator.validate_query("", mode="readonly")75 assert is_safe is False76 assert any("empty" in issue.lower() for issue in issues)77 78 def test_complexity_estimation(self):79 """Test query complexity scoring."""80 81 # Simple query82 simple_sql = "SELECT * FROM users"83 complexity = query_validator.estimate_query_complexity(simple_sql)84 assert complexity == 085 86 # Complex query with JOINs and aggregations87 complex_sql = """88 SELECT 89 u.region,90 COUNT(*) as order_count,91 SUM(o.total) as revenue92 FROM users u93 JOIN orders o ON u.id = o.user_id94 GROUP BY u.region95 ORDER BY revenue DESC96 """97 complexity = query_validator.estimate_query_complexity(complex_sql)98 assert complexity > 599 100 101class TestDBAAuth:102 """Test DBA authentication."""103 104 def test_authentication(self):105 """Test DBA password authentication."""106 from security.auth import DBAAuth107 108 # Create mock session state109 session_state = {}110 111 # Test with wrong password112 result = DBAAuth.authenticate("wrong_password", session_state)113 assert result is False114 assert session_state.get('dba_authenticated') is not True115 116 # Note: Testing with correct password requires actual env variable117 # In production, use proper test fixtures118 119 def test_session_management(self):120 """Test session state management."""121 from security.auth import DBAAuth122 123 session_state = {}124 125 # Not authenticated initially126 assert DBAAuth.is_authenticated(session_state) is False127 128 # Logout should handle non-existent state129 DBAAuth.logout(session_state)130 assert session_state.get('dba_authenticated') is False131 132 133class TestAuditLogger:134 """Test audit logging functionality."""135 136 def test_query_logging(self):137 """Test query attempt logging."""138 from security.audit_logger import AuditLogger139 import tempfile140 import os141 142 # Use temp directory for tests143 with tempfile.TemporaryDirectory() as tmpdir:144 logger = AuditLogger(log_dir=tmpdir)145 146 # Log a query147 log_id = logger.log_query_attempt(148 user_id="test_user",149 question="Show me all users",150 sql="SELECT * FROM users",151 mode="readonly"152 )153 154 assert log_id is not None155 assert log_id.startswith("20") # Starts with year156 157 # Update result158 logger.log_query_result(159 log_id=log_id,160 success=True,161 error=None,162 execution_time=0.5,163 row_count=10164 )165 166 # Verify log file exists167 log_files = os.listdir(tmpdir)168 assert len(log_files) > 0169 170 def test_statistics(self):171 """Test audit statistics calculation."""172 from security.audit_logger import AuditLogger173 import tempfile174 175 with tempfile.TemporaryDirectory() as tmpdir:176 logger = AuditLogger(log_dir=tmpdir)177 178 # Log some queries179 for i in range(5):180 log_id = logger.log_query_attempt(181 user_id=f"user_{i}",182 question=f"Query {i}",183 sql=f"SELECT * FROM table_{i}",184 mode="readonly"185 )186 187 # 3 successful, 2 failed188 logger.log_query_result(189 log_id=log_id,190 success=(i < 3),191 error="Error" if i >= 3 else None,192 execution_time=0.5,193 row_count=10 if i < 3 else 0194 )195 196 # Get statistics197 stats = logger.get_statistics()198 199 assert stats["total_queries"] == 5200 assert stats["successful_queries"] == 3201 assert stats["failed_queries"] == 2202 assert stats["success_rate"] == 60.0203 204 205if __name__ == "__main__":206 pytest.main([__file__, "-v"])207 