1const MASK: &str = "?";
17pub(crate) const MAX_SQL_LENGTH: usize = 4000;
18const MAX_NAMES: usize = 10;
19const MAX_NAME_LENGTH: usize = 200;
20
21#[derive(Clone, Debug, PartialEq, Eq)]
25pub struct SqlObjects {
26 pub operation: Option<String>,
27 pub procedures: Vec<String>,
28 pub relations: Vec<String>,
29}
30
31impl SqlObjects {
32 pub fn to_json(&self) -> String {
33 use crate::pii_scrubber::json_string;
34 let list = |names: &[String]| {
35 let quoted: Vec<String> = names.iter().map(|n| json_string(n)).collect();
36 format!("[{}]", quoted.join(","))
37 };
38 let operation = self
39 .operation
40 .as_deref()
41 .map(|op| format!("\"operation\":{},", json_string(op)))
42 .unwrap_or_default();
43 format!(
44 "{{{}\"procedures\":{},\"relations\":{}}}",
45 operation,
46 list(&self.procedures),
47 list(&self.relations)
48 )
49 }
50}
51
52fn is_word(c: u8) -> bool {
53 c == b'_' || c.is_ascii_alphanumeric()
54}
55
56pub(crate) fn mask(statement: &str) -> Option<String> {
58 if statement.trim().is_empty() {
59 return None;
60 }
61
62 let s = statement.as_bytes();
63 let mut out: Vec<u8> = Vec::with_capacity(s.len());
64 let mut i = 0;
65 while i < s.len() {
66 let c = s[i];
67 if c == b'\'' {
68 let mut j = i + 1;
71 while j < s.len() {
72 if s[j] == b'\'' {
73 if j + 1 < s.len() && s[j + 1] == b'\'' {
74 j += 2;
75 continue;
76 }
77 j += 1;
78 break;
79 }
80 j += 1;
81 }
82 out.extend_from_slice(MASK.as_bytes());
83 i = j;
84 } else if c == b'$' {
85 let mut j = i + 1;
87 while j < s.len() && (s[j] == b'_' || s[j].is_ascii_alphabetic()) {
88 j += 1;
89 }
90 if j < s.len() && s[j] == b'$' {
91 let tag = &s[i..=j];
92 let rest = &s[j + 1..];
93 let end = rest
94 .windows(tag.len())
95 .position(|w| w == tag)
96 .map(|p| j + 1 + p + tag.len())
97 .unwrap_or(s.len());
98 out.extend_from_slice(MASK.as_bytes());
99 i = end;
100 } else {
101 out.push(c);
102 i += 1;
103 }
104 } else if c.is_ascii_digit() {
105 let part_of_something =
108 i > 0 && (is_word(s[i - 1]) || s[i - 1] == b'$' || s[i - 1] == b'.');
109 match if part_of_something {
110 None
111 } else {
112 number_end(s, i)
113 } {
114 Some(end) => {
115 out.extend_from_slice(MASK.as_bytes());
116 i = end;
117 }
118 None => {
119 out.push(c);
120 i += 1;
121 }
122 }
123 } else {
124 out.push(c);
125 i += 1;
126 }
127 }
128
129 let masked = String::from_utf8_lossy(&out).into_owned();
131 if masked.chars().count() > MAX_SQL_LENGTH {
132 let truncated: String = masked.chars().take(MAX_SQL_LENGTH).collect();
133 return Some(format!("{truncated}..."));
134 }
135 Some(masked)
136}
137
138fn number_end(s: &[u8], i: usize) -> Option<usize> {
142 let mut k = i;
143 while k < s.len() && s[k].is_ascii_digit() {
144 k += 1;
145 }
146 let int_end = k;
147 if k + 1 < s.len() && s[k] == b'.' && s[k + 1].is_ascii_digit() {
148 let mut m = k + 1;
149 while m < s.len() && s[m].is_ascii_digit() {
150 m += 1;
151 }
152 if m >= s.len() || !is_word(s[m]) {
153 return Some(m);
154 }
155 }
156 if int_end >= s.len() || !is_word(s[int_end]) {
157 return Some(int_end);
158 }
159 None
160}
161
162struct Token {
163 text: String,
164 is_name: bool,
165}
166
167fn name_part_end(s: &[u8], i: usize) -> Option<usize> {
170 let c = *s.get(i)?;
171 if is_word(c) || c == b'$' || c == b'#' || c == b'@' {
172 let mut j = i;
173 while j < s.len() && (is_word(s[j]) || s[j] == b'$' || s[j] == b'#' || s[j] == b'@') {
174 j += 1;
175 }
176 return Some(j);
177 }
178 if c == b'"' || c == b'`' || c == b'[' {
179 let closer = if c == b'[' { b']' } else { c };
180 let mut j = i + 1;
181 while j < s.len() && s[j] != closer {
182 j += 1;
183 }
184 if j < s.len() && j > i + 1 {
185 return Some(j + 1);
186 }
187 }
188 None
189}
190
191fn name_end(s: &[u8], i: usize) -> Option<usize> {
193 let mut end = name_part_end(s, i)?;
194 while end < s.len() && s[end] == b'.' {
195 match name_part_end(s, end + 1) {
196 Some(next) => end = next,
197 None => break,
198 }
199 }
200 Some(end)
201}
202
203fn tokenize(sql: &str) -> Vec<Token> {
204 let s = sql.as_bytes();
205 let mut tokens = Vec::new();
206 let mut i = 0;
207 while i < s.len() {
208 if s[i].is_ascii_whitespace() || s[i] == 0x0b {
209 i += 1;
210 continue;
211 }
212 if let Some(end) = name_end(s, i) {
213 tokens.push(Token {
214 text: sql[i..end].to_string(),
215 is_name: true,
216 });
217 i = end;
218 continue;
219 }
220 let width = sql[i..].chars().next().map_or(1, char::len_utf8);
223 tokens.push(Token {
224 text: sql[i..i + width].to_string(),
225 is_name: false,
226 });
227 i += width;
228 }
229 tokens
230}
231
232fn is_full_name(name: &str) -> bool {
233 !name.is_empty() && name_end(name.as_bytes(), 0) == Some(name.len())
234}
235
236const OPERATIONS: [&str; 13] = [
237 "SELECT", "INSERT", "UPDATE", "DELETE", "MERGE", "WITH", "CALL", "EXEC", "EXECUTE", "CREATE",
238 "ALTER", "DROP", "TRUNCATE",
239];
240const BUILTINS: [&str; 21] = [
241 "count",
242 "sum",
243 "min",
244 "max",
245 "avg",
246 "now",
247 "coalesce",
248 "nullif",
249 "lower",
250 "upper",
251 "length",
252 "concat",
253 "cast",
254 "date_trunc",
255 "current_timestamp",
256 "current_date",
257 "row_number",
258 "rank",
259 "json_build_object",
260 "json_agg",
261 "array_agg",
262];
263const KEYWORDS_NOT_NAMES: [&str; 8] = [
264 "select",
265 "set",
266 "values",
267 "where",
268 "lateral",
269 "only",
270 "unnest",
271 "generate_series",
272];
273
274pub(crate) fn extract_objects(masked: &str) -> Option<SqlObjects> {
277 if masked.trim().is_empty() {
278 return None;
279 }
280
281 let all = tokenize(masked);
283 let mut tokens: Vec<&Token> = Vec::new();
284 let mut i = 0;
285 while i < all.len() {
286 let t = &all[i];
287 if t.is_name
288 && ["extract", "substring", "trim", "overlay"].contains(&t.text.to_lowercase().as_str())
289 && all.get(i + 1).is_some_and(|n| n.text == "(")
290 {
291 let mut close = None;
292 for (j, token) in all.iter().enumerate().skip(i + 2) {
293 if token.text == "(" {
294 break;
295 }
296 if token.text == ")" {
297 close = Some(j);
298 break;
299 }
300 }
301 if let Some(j) = close {
302 i = j + 1;
303 continue;
304 }
305 }
306 tokens.push(t);
307 i += 1;
308 }
309
310 let mut procedures: Vec<String> = Vec::new();
311 let mut relations: Vec<String> = Vec::new();
312
313 let mut i = 0;
314 while i + 1 < tokens.len() {
315 if tokens[i].is_name
316 && ["call", "exec", "execute", "perform"]
317 .contains(&tokens[i].text.to_lowercase().as_str())
318 && tokens[i + 1].is_name
319 {
320 let name = &tokens[i + 1].text;
321 let first = name.split('.').next().unwrap_or("").to_lowercase();
322 if !["immediate", "function", "procedure"].contains(&first.as_str()) {
323 procedures.push(name.clone());
324 i += 1;
325 }
326 }
327 i += 1;
328 }
329
330 let mut i = 0;
331 while i + 1 < tokens.len() {
332 let keyword = tokens[i].text.to_lowercase();
333 if !(tokens[i].is_name
334 && ["from", "join", "into", "update", "table"].contains(&keyword.as_str())
335 && tokens[i + 1].is_name)
336 {
337 i += 1;
338 continue;
339 }
340 let name = tokens[i + 1].text.clone();
341 let paren = tokens.get(i + 2).is_some_and(|t| t.text == "(");
342 i += 2;
343 if KEYWORDS_NOT_NAMES.contains(&name.to_lowercase().as_str()) {
344 continue;
345 }
346 if paren && (keyword == "from" || keyword == "join") {
349 procedures.push(name);
350 } else {
351 relations.push(name);
352 }
353 }
354
355 if tokens.len() >= 3
356 && tokens[0].is_name
357 && tokens[0].text.eq_ignore_ascii_case("select")
358 && tokens[1].is_name
359 && tokens[2].text == "("
360 && !BUILTINS.contains(&tokens[1].text.to_lowercase().as_str())
361 && !tokens
362 .iter()
363 .any(|t| t.is_name && t.text.eq_ignore_ascii_case("from"))
364 {
365 procedures.push(tokens[1].text.clone());
366 }
367
368 let operation = tokens.first().and_then(|t| {
369 let word: String = t
370 .text
371 .bytes()
372 .take_while(|b| is_word(*b))
373 .map(char::from)
374 .collect();
375 let upper = word.to_uppercase();
376 OPERATIONS.contains(&upper.as_str()).then_some(upper)
377 });
378
379 let objects = SqlObjects {
380 operation,
381 procedures: clean(procedures),
382 relations: clean(relations),
383 };
384 if objects.procedures.is_empty() && objects.relations.is_empty() && objects.operation.is_none()
385 {
386 return None;
387 }
388 Some(objects)
389}
390
391fn clean(names: Vec<String>) -> Vec<String> {
392 let mut cleaned: Vec<String> = Vec::new();
393 for raw in names {
394 let name: String = raw.trim().chars().take(MAX_NAME_LENGTH).collect();
395 if is_full_name(&name) && !cleaned.contains(&name) {
396 cleaned.push(name);
397 }
398 }
399 cleaned.truncate(MAX_NAMES);
400 cleaned
401}
402
403#[cfg(test)]
404mod tests {
405 use super::*;
406
407 fn objects(sql: &str) -> Option<SqlObjects> {
408 extract_objects(sql)
409 }
410
411 #[test]
412 fn masks_strings_and_numbers_but_not_identifiers_or_placeholders() {
413 assert_eq!(
414 mask("SELECT * FROM orders2 WHERE email = 'a@b.co' AND id = 42 AND x = $1").unwrap(),
415 "SELECT * FROM orders2 WHERE email = ? AND id = ? AND x = $1"
416 );
417 assert_eq!(
418 mask("SELECT price * 1.5 FROM t WHERE a IN (1,2,3)").unwrap(),
419 "SELECT price * ? FROM t WHERE a IN (?,?,?)"
420 );
421 assert_eq!(mask("SELECT 1.5x FROM t").unwrap(), "SELECT ?.5x FROM t");
422 }
423
424 #[test]
425 fn masks_an_escaped_quote_a_cut_off_string_and_a_dollar_quoted_body() {
426 assert_eq!(mask("EXEC sp_x @t = 'it''s'").unwrap(), "EXEC sp_x @t = ?");
427 assert_eq!(
428 mask("SELECT 1 WHERE n = 'oops").unwrap(),
429 "SELECT ? WHERE n = ?"
430 );
431 assert_eq!(mask("DO $b$ BEGIN PERFORM 1; END $b$").unwrap(), "DO ?");
432 }
433
434 #[test]
435 fn keeps_multibyte_text_intact_and_is_idempotent_truncating_and_blank_safe() {
436 assert_eq!(
437 mask("SELECT \"naïve\" FROM t WHERE a = 'é'").unwrap(),
438 "SELECT \"naïve\" FROM t WHERE a = ?"
439 );
440 let once = mask("SELECT * FROM t WHERE a = 'x' AND b = 9").unwrap();
441 assert_eq!(mask(&once).unwrap(), once);
442 assert_eq!(
443 mask(&format!("SELECT {} b", "a, ".repeat(3000)))
444 .unwrap()
445 .chars()
446 .count(),
447 MAX_SQL_LENGTH + 3
448 );
449 assert_eq!(mask(" "), None);
450 }
451
452 #[test]
453 fn finds_a_stored_procedure_with_its_schema() {
454 assert_eq!(
455 objects("EXEC dbo.sp_refund_order @id = ?").unwrap(),
456 SqlObjects {
457 operation: Some("EXEC".into()),
458 procedures: vec!["dbo.sp_refund_order".into()],
459 relations: vec![]
460 }
461 );
462 assert_eq!(
463 objects("CALL refund_order(?, ?)").unwrap().procedures,
464 vec!["refund_order"]
465 );
466 assert_eq!(
467 objects("SELECT refund_order(?, ?)").unwrap().procedures,
468 vec!["refund_order"]
469 );
470 }
471
472 #[test]
473 fn finds_views_joined_tables_and_table_functions() {
474 assert_eq!(
475 objects("SELECT * FROM v_totals t JOIN public.customers c ON c.id = t.id")
476 .unwrap()
477 .relations,
478 vec!["v_totals", "public.customers"]
479 );
480 assert_eq!(
481 objects("SELECT * FROM get_open_orders(?) o")
482 .unwrap()
483 .procedures,
484 vec!["get_open_orders"]
485 );
486 }
487
488 #[test]
489 fn does_not_misread_column_lists_builtins_or_from_inside_extract() {
490 assert_eq!(
491 objects("INSERT INTO audit_log (a) VALUES (?)")
492 .unwrap()
493 .procedures,
494 Vec::<String>::new()
495 );
496 assert_eq!(
497 objects("SELECT count(*) FROM orders").unwrap().procedures,
498 Vec::<String>::new()
499 );
500 assert_eq!(
501 objects("SELECT 1 FROM orders WHERE extract(year FROM created_at) = ?")
502 .unwrap()
503 .relations,
504 vec!["orders"]
505 );
506 assert_eq!(objects("garbage"), None);
507 }
508
509 #[test]
510 fn keeps_quoted_and_bracketed_identifiers_whole() {
511 assert_eq!(
512 objects("UPDATE \"Order Items\" SET qty = ?")
513 .unwrap()
514 .relations,
515 vec!["\"Order Items\""]
516 );
517 assert_eq!(
518 objects("INSERT INTO [dbo].[audit_log] (a) VALUES (?)")
519 .unwrap()
520 .relations,
521 vec!["[dbo].[audit_log]"]
522 );
523 }
524
525 #[test]
526 fn serializes_to_json() {
527 let json = objects("EXEC dbo.sp_x @id = ?").unwrap().to_json();
528 assert_eq!(
529 json,
530 "{\"operation\":\"EXEC\",\"procedures\":[\"dbo.sp_x\"],\"relations\":[]}"
531 );
532 }
533}