lib

Core libraries for Radroots
git clone https://radroots.dev/git/lib.git
Log | Files | Refs | README

statement_policy.rs (7666B)


      1 //! Private service-controlled SQL admission policy.
      2 
      3 #[derive(Clone, Copy)]
      4 enum CreatePrefix {
      5     None,
      6     Create,
      7     CreateTemp,
      8     Other,
      9 }
     10 
     11 pub(crate) fn contains_forbidden_statement_control(sql: &str) -> bool {
     12     let bytes = sql.as_bytes();
     13     let mut cursor = 0;
     14     let mut statement_start = true;
     15     let mut create_prefix = CreatePrefix::None;
     16     let mut trigger_definition = false;
     17     let mut trigger_body = false;
     18     let mut trigger_case_depth = 0_u32;
     19     let mut trigger_end_seen = false;
     20 
     21     while cursor < bytes.len() {
     22         let byte = bytes[cursor];
     23         if statement_start && bytes[cursor..].starts_with(&[0xef, 0xbb, 0xbf]) {
     24             cursor += 3;
     25             continue;
     26         }
     27         if byte.is_ascii_whitespace() {
     28             cursor += 1;
     29             continue;
     30         }
     31         if byte == b'-' && bytes.get(cursor + 1) == Some(&b'-') {
     32             cursor += 2;
     33             while cursor < bytes.len() && !matches!(bytes[cursor], b'\n' | b'\r') {
     34                 cursor += 1;
     35             }
     36             continue;
     37         }
     38         if byte == b'/' && bytes.get(cursor + 1) == Some(&b'*') {
     39             cursor += 2;
     40             while cursor + 1 < bytes.len() && !(bytes[cursor] == b'*' && bytes[cursor + 1] == b'/')
     41             {
     42                 cursor += 1;
     43             }
     44             cursor = bytes.len().min(cursor + 2);
     45             continue;
     46         }
     47         if matches!(byte, b'\'' | b'"' | b'`' | b'[') {
     48             let terminator = if byte == b'[' { b']' } else { byte };
     49             statement_start = false;
     50             create_prefix = CreatePrefix::Other;
     51             cursor += 1;
     52             while cursor < bytes.len() {
     53                 if bytes[cursor] == terminator {
     54                     if bytes.get(cursor + 1) == Some(&terminator) {
     55                         cursor += 2;
     56                     } else {
     57                         cursor += 1;
     58                         break;
     59                     }
     60                 } else {
     61                     cursor += 1;
     62                 }
     63             }
     64             continue;
     65         }
     66         if byte == b';' {
     67             if trigger_definition && trigger_body && !trigger_end_seen {
     68                 cursor += 1;
     69                 continue;
     70             }
     71             statement_start = true;
     72             create_prefix = CreatePrefix::None;
     73             trigger_definition = false;
     74             trigger_body = false;
     75             trigger_case_depth = 0;
     76             trigger_end_seen = false;
     77             cursor += 1;
     78             continue;
     79         }
     80         if !is_identifier(byte) {
     81             statement_start = false;
     82             create_prefix = CreatePrefix::Other;
     83             cursor += 1;
     84             continue;
     85         }
     86 
     87         let start = cursor;
     88         cursor += 1;
     89         while cursor < bytes.len() && is_identifier(bytes[cursor]) {
     90             cursor += 1;
     91         }
     92         let token = &bytes[start..cursor];
     93 
     94         if trigger_definition {
     95             if !trigger_body && token.eq_ignore_ascii_case(b"begin") {
     96                 trigger_body = true;
     97             } else if trigger_body && token.eq_ignore_ascii_case(b"case") {
     98                 trigger_case_depth = trigger_case_depth.saturating_add(1);
     99             } else if trigger_body && token.eq_ignore_ascii_case(b"end") {
    100                 if trigger_case_depth == 0 {
    101                     trigger_end_seen = true;
    102                 } else {
    103                     trigger_case_depth -= 1;
    104                 }
    105             }
    106             continue;
    107         }
    108 
    109         if statement_start {
    110             if is_forbidden(token) {
    111                 return true;
    112             }
    113             statement_start = false;
    114             create_prefix = if token.eq_ignore_ascii_case(b"create") {
    115                 CreatePrefix::Create
    116             } else {
    117                 CreatePrefix::Other
    118             };
    119             continue;
    120         }
    121 
    122         create_prefix = match create_prefix {
    123             CreatePrefix::Create
    124                 if token.eq_ignore_ascii_case(b"temp")
    125                     || token.eq_ignore_ascii_case(b"temporary") =>
    126             {
    127                 CreatePrefix::CreateTemp
    128             }
    129             CreatePrefix::Create | CreatePrefix::CreateTemp
    130                 if token.eq_ignore_ascii_case(b"trigger") =>
    131             {
    132                 trigger_definition = true;
    133                 CreatePrefix::None
    134             }
    135             CreatePrefix::Create | CreatePrefix::CreateTemp => CreatePrefix::Other,
    136             other => other,
    137         };
    138     }
    139 
    140     false
    141 }
    142 
    143 fn is_identifier(byte: u8) -> bool {
    144     byte.is_ascii_alphanumeric() || byte == b'_'
    145 }
    146 
    147 fn is_forbidden(token: &[u8]) -> bool {
    148     [
    149         b"pragma".as_slice(),
    150         b"attach",
    151         b"detach",
    152         b"begin",
    153         b"commit",
    154         b"end",
    155         b"rollback",
    156         b"savepoint",
    157         b"release",
    158     ]
    159     .into_iter()
    160     .any(|forbidden| token.eq_ignore_ascii_case(forbidden))
    161 }
    162 
    163 #[cfg(test)]
    164 mod tests {
    165     use super::*;
    166 
    167     #[test]
    168     fn statement_control_inventory_is_closed_and_case_insensitive() {
    169         for forbidden in [
    170             "/* ignored before control */ PRAGMA trusted_schema = ON",
    171             "\u{feff}PRAGMA trusted_schema = ON",
    172             "ATTACH DATABASE 'x' AS extra",
    173             "detach database extra",
    174             " /* ignored */ BeGiN IMMEDIATE",
    175             "SELECT 1; -- ignored\n CoMmIt",
    176             "END TRANSACTION",
    177             "ROLLBACK TO escaped",
    178             "SAVEPOINT escaped",
    179             "RELEASE SAVEPOINT escaped",
    180             "CREATE TRIGGER audit_insert AFTER INSERT ON items BEGIN INSERT INTO audit_log (value) VALUES (NEW.value); END; /* after trigger */ COMMIT",
    181         ] {
    182             assert!(contains_forbidden_statement_control(forbidden));
    183         }
    184     }
    185 
    186     #[test]
    187     fn values_identifiers_comments_case_and_triggers_remain_available() {
    188         for allowed in [
    189             "SELECT 'pragma attach detach begin commit end rollback savepoint release'",
    190             "SELECT CASE WHEN value = 1 THEN 'commit' ELSE 'end' END FROM items",
    191             "SELECT 1 /* PRAGMA ATTACH COMMIT */",
    192             "CREATE TRIGGER audit_insert AFTER INSERT ON items BEGIN INSERT INTO audit_log (value) VALUES (CASE WHEN NEW.value = 1 THEN 'commit' ELSE 'end' END); END;",
    193             "SELECT pragmatic FROM items",
    194             "SELECT attachment FROM items",
    195             "SELECT detached FROM items",
    196             "SELECT beginner FROM items",
    197             "SELECT committed FROM items",
    198             "SELECT ending FROM items",
    199             "SELECT rolled_back FROM items",
    200             "SELECT savepoints FROM items",
    201             "SELECT released FROM items",
    202             "SELECT COUNT(*) FROM host_probe",
    203         ] {
    204             assert!(!contains_forbidden_statement_control(allowed));
    205         }
    206     }
    207 
    208     #[test]
    209     fn lexical_edges_remain_bounded_and_do_not_invent_statement_control() {
    210         for allowed in [
    211             "",
    212             "-",
    213             "-- unterminated comment",
    214             "/",
    215             "/* unterminated comment",
    216             "/* interior * is not a terminator */ SELECT 1",
    217             "''",
    218             "'unterminated value",
    219             "'escaped''quote'",
    220             "[]",
    221             "[unterminated identifier",
    222             "SELECT 1;",
    223             "CREATE TABLE items (value TEXT)",
    224             "CREATE TEMP TABLE items (value TEXT)",
    225             "CREATE TEMPORARY TRIGGER audit AFTER INSERT ON items BEGIN SELECT 1; END;",
    226             "CREATE TRIGGER incomplete; SELECT 1",
    227         ] {
    228             assert!(
    229                 !contains_forbidden_statement_control(allowed),
    230                 "lexical edge was misclassified: {allowed:?}"
    231             );
    232         }
    233     }
    234 }