Team Ai
Apppublic

joelgilbert/NL2SQL

sourceHugging Facemitupdated 11mo agoView on Hugging Face
0likes
test_security.py207 linesDownload Raw Back to tests
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