Skip to main content

forge_ops_tracker/
sql_statement.rs

1// Finds the SQL behind a database error and reduces it to something safe to send: the names of the
2// stored procedures, tables and views it touched, and (only if `capture_sql_statement` is on) the
3// statement itself with every string and number replaced by "?". Ported from
4// gems/forge_ops_tracker's SqlStatement, which is itself ported from the server's own
5// SqlStatementMasker/SqlObjectExtractor: same rules everywhere, and the server applies them again
6// on arrival, so a difference here can only ever mean less is masked client-side, never that
7// something unmasked gets stored.
8//
9// Written as a hand-rolled scanner and tokenizer rather than regular expressions: the `regex`
10// crate deliberately has no lookaround or backreferences, which the shared pattern relies on (a
11// number is only a value when it isn't part of an identifier, and a dollar-quoted body ends at the
12// same tag that opened it). The rules are identical; only the mechanism differs.
13//
14// Deliberately not a SQL parser.
15
16const MASK: &str = "?";
17pub(crate) const MAX_SQL_LENGTH: usize = 4000;
18const MAX_NAMES: usize = 10;
19const MAX_NAME_LENGTH: usize = 200;
20
21/// What extraction found in a statement: its operation (the first keyword), the stored
22/// procedures/functions it called, and the tables/views it touched. A view and a table are written
23/// the same way in SQL text, so both land in `relations`.
24#[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
56/// Replaces every string literal and number in `statement` with `?`. `None` for a blank statement.
57pub(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            // A string literal; '' is an escaped quote. One cut off by truncation (no closing
69            // quote) is masked to the end of the statement, never left half-visible.
70            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            // A dollar-quoted body ($tag$ ... $tag$): PostgreSQL function bodies and DO blocks.
86            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            // A number, unless it's part of an identifier (orders2, sp_v2), a $1 placeholder, or
106            // the fraction of another number; those digits are left alone.
107            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    // Only ever cut at ASCII delimiters above, so this is always valid UTF-8.
130    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
138// Where the number starting at `i` ends, or None when it isn't a standalone number (digits
139// immediately followed by a letter or underscore). A decimal that fails that check falls back to
140// just its integer part, the same way the shared pattern's backtracking does.
141fn 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
167// Where the identifier part starting at `i` ends: bare (letters, digits, _ $ # @), "double
168// quoted", [bracketed] (SQL Server) or `backticked`.
169fn 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
191// Where the (optionally schema-qualified) name starting at `i` ends.
192fn 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        // Any other character is its own token; step by whole characters so a multi-byte one is
221        // never split.
222        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
274/// Takes an already-masked statement (so a keyword inside a string value can't be mistaken for
275/// SQL) and returns the names it touched, or `None` when nothing recognizable was found.
276pub(crate) fn extract_objects(masked: &str) -> Option<SqlObjects> {
277    if masked.trim().is_empty() {
278        return None;
279    }
280
281    // EXTRACT(year FROM col), SUBSTRING(x FROM 2), TRIM(BOTH FROM x): a FROM that isn't a table.
282    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        // FROM/JOIN some_function(...) is a set-returning function (often a stored one), not a
347        // table. INSERT INTO t (a, b) is just a column list, so INTO/UPDATE/TABLE never count.
348        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}