const MASK: &str = "?";
pub(crate) const MAX_SQL_LENGTH: usize = 4000;
const MAX_NAMES: usize = 10;
const MAX_NAME_LENGTH: usize = 200;
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct SqlObjects {
pub operation: Option<String>,
pub procedures: Vec<String>,
pub relations: Vec<String>,
}
impl SqlObjects {
pub fn to_json(&self) -> String {
use crate::pii_scrubber::json_string;
let list = |names: &[String]| {
let quoted: Vec<String> = names.iter().map(|n| json_string(n)).collect();
format!("[{}]", quoted.join(","))
};
let operation = self
.operation
.as_deref()
.map(|op| format!("\"operation\":{},", json_string(op)))
.unwrap_or_default();
format!(
"{{{}\"procedures\":{},\"relations\":{}}}",
operation,
list(&self.procedures),
list(&self.relations)
)
}
}
fn is_word(c: u8) -> bool {
c == b'_' || c.is_ascii_alphanumeric()
}
pub(crate) fn mask(statement: &str) -> Option<String> {
if statement.trim().is_empty() {
return None;
}
let s = statement.as_bytes();
let mut out: Vec<u8> = Vec::with_capacity(s.len());
let mut i = 0;
while i < s.len() {
let c = s[i];
if c == b'\'' {
let mut j = i + 1;
while j < s.len() {
if s[j] == b'\'' {
if j + 1 < s.len() && s[j + 1] == b'\'' {
j += 2;
continue;
}
j += 1;
break;
}
j += 1;
}
out.extend_from_slice(MASK.as_bytes());
i = j;
} else if c == b'$' {
let mut j = i + 1;
while j < s.len() && (s[j] == b'_' || s[j].is_ascii_alphabetic()) {
j += 1;
}
if j < s.len() && s[j] == b'$' {
let tag = &s[i..=j];
let rest = &s[j + 1..];
let end = rest
.windows(tag.len())
.position(|w| w == tag)
.map(|p| j + 1 + p + tag.len())
.unwrap_or(s.len());
out.extend_from_slice(MASK.as_bytes());
i = end;
} else {
out.push(c);
i += 1;
}
} else if c.is_ascii_digit() {
let part_of_something =
i > 0 && (is_word(s[i - 1]) || s[i - 1] == b'$' || s[i - 1] == b'.');
match if part_of_something {
None
} else {
number_end(s, i)
} {
Some(end) => {
out.extend_from_slice(MASK.as_bytes());
i = end;
}
None => {
out.push(c);
i += 1;
}
}
} else {
out.push(c);
i += 1;
}
}
let masked = String::from_utf8_lossy(&out).into_owned();
if masked.chars().count() > MAX_SQL_LENGTH {
let truncated: String = masked.chars().take(MAX_SQL_LENGTH).collect();
return Some(format!("{truncated}..."));
}
Some(masked)
}
fn number_end(s: &[u8], i: usize) -> Option<usize> {
let mut k = i;
while k < s.len() && s[k].is_ascii_digit() {
k += 1;
}
let int_end = k;
if k + 1 < s.len() && s[k] == b'.' && s[k + 1].is_ascii_digit() {
let mut m = k + 1;
while m < s.len() && s[m].is_ascii_digit() {
m += 1;
}
if m >= s.len() || !is_word(s[m]) {
return Some(m);
}
}
if int_end >= s.len() || !is_word(s[int_end]) {
return Some(int_end);
}
None
}
struct Token {
text: String,
is_name: bool,
}
fn name_part_end(s: &[u8], i: usize) -> Option<usize> {
let c = *s.get(i)?;
if is_word(c) || c == b'$' || c == b'#' || c == b'@' {
let mut j = i;
while j < s.len() && (is_word(s[j]) || s[j] == b'$' || s[j] == b'#' || s[j] == b'@') {
j += 1;
}
return Some(j);
}
if c == b'"' || c == b'`' || c == b'[' {
let closer = if c == b'[' { b']' } else { c };
let mut j = i + 1;
while j < s.len() && s[j] != closer {
j += 1;
}
if j < s.len() && j > i + 1 {
return Some(j + 1);
}
}
None
}
fn name_end(s: &[u8], i: usize) -> Option<usize> {
let mut end = name_part_end(s, i)?;
while end < s.len() && s[end] == b'.' {
match name_part_end(s, end + 1) {
Some(next) => end = next,
None => break,
}
}
Some(end)
}
fn tokenize(sql: &str) -> Vec<Token> {
let s = sql.as_bytes();
let mut tokens = Vec::new();
let mut i = 0;
while i < s.len() {
if s[i].is_ascii_whitespace() || s[i] == 0x0b {
i += 1;
continue;
}
if let Some(end) = name_end(s, i) {
tokens.push(Token {
text: sql[i..end].to_string(),
is_name: true,
});
i = end;
continue;
}
let width = sql[i..].chars().next().map_or(1, char::len_utf8);
tokens.push(Token {
text: sql[i..i + width].to_string(),
is_name: false,
});
i += width;
}
tokens
}
fn is_full_name(name: &str) -> bool {
!name.is_empty() && name_end(name.as_bytes(), 0) == Some(name.len())
}
const OPERATIONS: [&str; 13] = [
"SELECT", "INSERT", "UPDATE", "DELETE", "MERGE", "WITH", "CALL", "EXEC", "EXECUTE", "CREATE",
"ALTER", "DROP", "TRUNCATE",
];
const BUILTINS: [&str; 21] = [
"count",
"sum",
"min",
"max",
"avg",
"now",
"coalesce",
"nullif",
"lower",
"upper",
"length",
"concat",
"cast",
"date_trunc",
"current_timestamp",
"current_date",
"row_number",
"rank",
"json_build_object",
"json_agg",
"array_agg",
];
const KEYWORDS_NOT_NAMES: [&str; 8] = [
"select",
"set",
"values",
"where",
"lateral",
"only",
"unnest",
"generate_series",
];
pub(crate) fn extract_objects(masked: &str) -> Option<SqlObjects> {
if masked.trim().is_empty() {
return None;
}
let all = tokenize(masked);
let mut tokens: Vec<&Token> = Vec::new();
let mut i = 0;
while i < all.len() {
let t = &all[i];
if t.is_name
&& ["extract", "substring", "trim", "overlay"].contains(&t.text.to_lowercase().as_str())
&& all.get(i + 1).is_some_and(|n| n.text == "(")
{
let mut close = None;
for (j, token) in all.iter().enumerate().skip(i + 2) {
if token.text == "(" {
break;
}
if token.text == ")" {
close = Some(j);
break;
}
}
if let Some(j) = close {
i = j + 1;
continue;
}
}
tokens.push(t);
i += 1;
}
let mut procedures: Vec<String> = Vec::new();
let mut relations: Vec<String> = Vec::new();
let mut i = 0;
while i + 1 < tokens.len() {
if tokens[i].is_name
&& ["call", "exec", "execute", "perform"]
.contains(&tokens[i].text.to_lowercase().as_str())
&& tokens[i + 1].is_name
{
let name = &tokens[i + 1].text;
let first = name.split('.').next().unwrap_or("").to_lowercase();
if !["immediate", "function", "procedure"].contains(&first.as_str()) {
procedures.push(name.clone());
i += 1;
}
}
i += 1;
}
let mut i = 0;
while i + 1 < tokens.len() {
let keyword = tokens[i].text.to_lowercase();
if !(tokens[i].is_name
&& ["from", "join", "into", "update", "table"].contains(&keyword.as_str())
&& tokens[i + 1].is_name)
{
i += 1;
continue;
}
let name = tokens[i + 1].text.clone();
let paren = tokens.get(i + 2).is_some_and(|t| t.text == "(");
i += 2;
if KEYWORDS_NOT_NAMES.contains(&name.to_lowercase().as_str()) {
continue;
}
if paren && (keyword == "from" || keyword == "join") {
procedures.push(name);
} else {
relations.push(name);
}
}
if tokens.len() >= 3
&& tokens[0].is_name
&& tokens[0].text.eq_ignore_ascii_case("select")
&& tokens[1].is_name
&& tokens[2].text == "("
&& !BUILTINS.contains(&tokens[1].text.to_lowercase().as_str())
&& !tokens
.iter()
.any(|t| t.is_name && t.text.eq_ignore_ascii_case("from"))
{
procedures.push(tokens[1].text.clone());
}
let operation = tokens.first().and_then(|t| {
let word: String = t
.text
.bytes()
.take_while(|b| is_word(*b))
.map(char::from)
.collect();
let upper = word.to_uppercase();
OPERATIONS.contains(&upper.as_str()).then_some(upper)
});
let objects = SqlObjects {
operation,
procedures: clean(procedures),
relations: clean(relations),
};
if objects.procedures.is_empty() && objects.relations.is_empty() && objects.operation.is_none()
{
return None;
}
Some(objects)
}
fn clean(names: Vec<String>) -> Vec<String> {
let mut cleaned: Vec<String> = Vec::new();
for raw in names {
let name: String = raw.trim().chars().take(MAX_NAME_LENGTH).collect();
if is_full_name(&name) && !cleaned.contains(&name) {
cleaned.push(name);
}
}
cleaned.truncate(MAX_NAMES);
cleaned
}
#[cfg(test)]
mod tests {
use super::*;
fn objects(sql: &str) -> Option<SqlObjects> {
extract_objects(sql)
}
#[test]
fn masks_strings_and_numbers_but_not_identifiers_or_placeholders() {
assert_eq!(
mask("SELECT * FROM orders2 WHERE email = 'a@b.co' AND id = 42 AND x = $1").unwrap(),
"SELECT * FROM orders2 WHERE email = ? AND id = ? AND x = $1"
);
assert_eq!(
mask("SELECT price * 1.5 FROM t WHERE a IN (1,2,3)").unwrap(),
"SELECT price * ? FROM t WHERE a IN (?,?,?)"
);
assert_eq!(mask("SELECT 1.5x FROM t").unwrap(), "SELECT ?.5x FROM t");
}
#[test]
fn masks_an_escaped_quote_a_cut_off_string_and_a_dollar_quoted_body() {
assert_eq!(mask("EXEC sp_x @t = 'it''s'").unwrap(), "EXEC sp_x @t = ?");
assert_eq!(
mask("SELECT 1 WHERE n = 'oops").unwrap(),
"SELECT ? WHERE n = ?"
);
assert_eq!(mask("DO $b$ BEGIN PERFORM 1; END $b$").unwrap(), "DO ?");
}
#[test]
fn keeps_multibyte_text_intact_and_is_idempotent_truncating_and_blank_safe() {
assert_eq!(
mask("SELECT \"naïve\" FROM t WHERE a = 'é'").unwrap(),
"SELECT \"naïve\" FROM t WHERE a = ?"
);
let once = mask("SELECT * FROM t WHERE a = 'x' AND b = 9").unwrap();
assert_eq!(mask(&once).unwrap(), once);
assert_eq!(
mask(&format!("SELECT {} b", "a, ".repeat(3000)))
.unwrap()
.chars()
.count(),
MAX_SQL_LENGTH + 3
);
assert_eq!(mask(" "), None);
}
#[test]
fn finds_a_stored_procedure_with_its_schema() {
assert_eq!(
objects("EXEC dbo.sp_refund_order @id = ?").unwrap(),
SqlObjects {
operation: Some("EXEC".into()),
procedures: vec!["dbo.sp_refund_order".into()],
relations: vec![]
}
);
assert_eq!(
objects("CALL refund_order(?, ?)").unwrap().procedures,
vec!["refund_order"]
);
assert_eq!(
objects("SELECT refund_order(?, ?)").unwrap().procedures,
vec!["refund_order"]
);
}
#[test]
fn finds_views_joined_tables_and_table_functions() {
assert_eq!(
objects("SELECT * FROM v_totals t JOIN public.customers c ON c.id = t.id")
.unwrap()
.relations,
vec!["v_totals", "public.customers"]
);
assert_eq!(
objects("SELECT * FROM get_open_orders(?) o")
.unwrap()
.procedures,
vec!["get_open_orders"]
);
}
#[test]
fn does_not_misread_column_lists_builtins_or_from_inside_extract() {
assert_eq!(
objects("INSERT INTO audit_log (a) VALUES (?)")
.unwrap()
.procedures,
Vec::<String>::new()
);
assert_eq!(
objects("SELECT count(*) FROM orders").unwrap().procedures,
Vec::<String>::new()
);
assert_eq!(
objects("SELECT 1 FROM orders WHERE extract(year FROM created_at) = ?")
.unwrap()
.relations,
vec!["orders"]
);
assert_eq!(objects("garbage"), None);
}
#[test]
fn keeps_quoted_and_bracketed_identifiers_whole() {
assert_eq!(
objects("UPDATE \"Order Items\" SET qty = ?")
.unwrap()
.relations,
vec!["\"Order Items\""]
);
assert_eq!(
objects("INSERT INTO [dbo].[audit_log] (a) VALUES (?)")
.unwrap()
.relations,
vec!["[dbo].[audit_log]"]
);
}
#[test]
fn serializes_to_json() {
let json = objects("EXEC dbo.sp_x @id = ?").unwrap().to_json();
assert_eq!(
json,
"{\"operation\":\"EXEC\",\"procedures\":[\"dbo.sp_x\"],\"relations\":[]}"
);
}
}