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 }