use crate::sql::Planner;
use crate::Value;
pub(crate) fn normalize_select_literals(sql: &str) -> Option<(String, Vec<Value>)> {
let bytes = sql.as_bytes();
let n = bytes.len();
let start = skip_ws(bytes, 0);
if !keyword_at(bytes, start, b"SELECT") {
return None;
}
let mut out = String::with_capacity(n + 16);
let mut params: Vec<Value> = Vec::new();
let mut i = 0usize;
let mut depth: i32 = 0;
let mut in_where = false;
let mut in_limit = false;
let mut in_offset = false;
while i < n {
let b = bytes[i];
if b == b'-' && i + 1 < n && bytes[i + 1] == b'-' {
return None;
}
if b == b'/' && i + 1 < n && bytes[i + 1] == b'*' {
return None;
}
if b == b'$' {
return None;
}
if b == b'\'' {
let end = skip_single_quoted(bytes, i)?; if in_where {
let raw = &sql[i + 1..end - 1]; let content = raw.replace("''", "'");
params.push(Planner::single_quoted_content_to_value(&content));
out.push('$');
out.push_str(¶ms.len().to_string());
} else {
out.push_str(&sql[i..end]);
}
i = end;
continue;
}
if (b == b'E' || b == b'e' || b == b'X' || b == b'x' || b == b'B' || b == b'b')
&& i + 1 < n
&& bytes[i + 1] == b'\''
{
return None;
}
if (b == b'U' || b == b'u') && i + 1 < n && bytes[i + 1] == b'&' {
return None;
}
if b.is_ascii_alphabetic() || b == b'_' {
let word_end = read_word(bytes, i);
let word = &sql[i..word_end];
if depth == 0 {
if word.eq_ignore_ascii_case("where") {
in_where = true;
in_limit = false;
in_offset = false;
out.push_str(word);
i = word_end;
continue;
}
if is_where_terminator_kw(word) {
in_where = false;
if word.eq_ignore_ascii_case("limit") {
in_limit = true;
in_offset = false;
} else if word.eq_ignore_ascii_case("offset") {
in_offset = true;
in_limit = false;
} else {
in_limit = false;
in_offset = false;
}
}
}
if in_where {
if word.eq_ignore_ascii_case("in") {
let after = skip_ws(bytes, word_end);
if after >= n || bytes[after] != b'(' {
return None;
}
if let Some((consumed_to, elems)) = lex_pure_literal_list(sql, bytes, after) {
if elems.is_empty() || elems.len() > MAX_IN_LIST {
return None;
}
let count = elems.len();
out.push_str(word);
out.push_str(" (");
for (k, value) in elems.into_iter().enumerate() {
if k > 0 {
out.push(',');
}
params.push(value);
out.push('$');
out.push_str(¶ms.len().to_string());
}
let last_placeholder = params.len();
for _ in count..count.next_power_of_two() {
out.push(',');
out.push('$');
out.push_str(&last_placeholder.to_string());
}
out.push(')');
i = consumed_to;
continue;
}
}
if is_predicate_bail_kw(word) {
return None;
}
}
out.push_str(word);
i = word_end;
continue;
}
if b == b'"' {
let end = skip_double_quoted(bytes, i)?;
out.push_str(&sql[i..end]);
i = end;
continue;
}
if b.is_ascii_digit() {
let end = read_number(bytes, i)?;
if in_where {
let text = &sql[i..end];
let value = Planner::number_literal_to_value(text).ok()?;
params.push(value);
out.push('$');
out.push_str(¶ms.len().to_string());
} else if depth == 0 && (in_limit || in_offset) && bytes[i..end].iter().all(u8::is_ascii_digit) {
let text = &sql[i..end];
let value = Planner::number_literal_to_value(text).ok()?;
params.push(value);
out.push('$');
out.push_str(¶ms.len().to_string());
in_limit = false;
in_offset = false;
} else {
out.push_str(&sql[i..end]);
}
i = end;
continue;
}
if b == b'.' && i + 1 < n && bytes[i + 1].is_ascii_digit() && (in_where || in_limit || in_offset) {
return None;
}
if b == b':' && i + 1 < n && bytes[i + 1] == b':' && in_where {
let ty_start = i + 2;
if ty_start >= n || !(bytes[ty_start].is_ascii_alphabetic() || bytes[ty_start] == b'_') {
return None;
}
let ty_end = read_word(bytes, ty_start);
let ty = &sql[ty_start..ty_end];
let after_ty = skip_ws(bytes, ty_end);
if !is_whitelisted_cast_type(ty) || (after_ty < n && bytes[after_ty] == b'(') {
return None;
}
out.push_str("::");
out.push_str(ty);
i = ty_end;
continue;
}
if b == b'(' {
depth += 1;
} else if b == b')' {
depth -= 1;
if depth < 0 {
return None; }
} else if b == b';' {
if skip_ws(bytes, i + 1) != n {
return None;
}
out.push_str(&sql[i..]);
i = n;
break;
}
out.push(b as char);
i += 1;
}
if !params.is_empty() {
Some((out, params))
} else {
None
}
}
fn skip_ws(bytes: &[u8], mut i: usize) -> usize {
while i < bytes.len() && bytes[i].is_ascii_whitespace() {
i += 1;
}
i
}
fn is_word_byte(b: u8) -> bool {
b.is_ascii_alphanumeric() || b == b'_'
}
fn keyword_at(bytes: &[u8], i: usize, kw: &[u8]) -> bool {
if i + kw.len() > bytes.len() {
return false;
}
if i > 0 && is_word_byte(bytes[i - 1]) {
return false;
}
if !bytes[i..i + kw.len()].eq_ignore_ascii_case(kw) {
return false;
}
let after = i + kw.len();
after >= bytes.len() || !is_word_byte(bytes[after])
}
fn read_word(bytes: &[u8], mut i: usize) -> usize {
while i < bytes.len() && is_word_byte(bytes[i]) {
i += 1;
}
i
}
fn read_number(bytes: &[u8], start: usize) -> Option<usize> {
let n = bytes.len();
let mut i = start;
while i < n && bytes[i].is_ascii_digit() {
i += 1;
}
if i < n && bytes[i] == b'.' {
i += 1;
while i < n && bytes[i].is_ascii_digit() {
i += 1;
}
if i < n && bytes[i] == b'.' {
return None;
}
}
if i < n && (bytes[i] == b'e' || bytes[i] == b'E') {
let mut j = i + 1;
if j < n && (bytes[j] == b'+' || bytes[j] == b'-') {
j += 1;
}
if j < n && bytes[j].is_ascii_digit() {
i = j;
while i < n && bytes[i].is_ascii_digit() {
i += 1;
}
}
}
if i < n && is_word_byte(bytes[i]) {
return None;
}
Some(i)
}
fn skip_single_quoted(bytes: &[u8], i: usize) -> Option<usize> {
let n = bytes.len();
let mut j = i + 1;
while j < n {
if bytes[j] == b'\'' {
if j + 1 < n && bytes[j + 1] == b'\'' {
j += 2; continue;
}
return Some(j + 1);
}
j += 1;
}
None
}
fn skip_double_quoted(bytes: &[u8], i: usize) -> Option<usize> {
let n = bytes.len();
let mut j = i + 1;
while j < n {
if bytes[j] == b'"' {
if j + 1 < n && bytes[j + 1] == b'"' {
j += 2;
continue;
}
return Some(j + 1);
}
j += 1;
}
None
}
fn is_where_terminator_kw(word: &str) -> bool {
const KWS: [&str; 11] = [
"group",
"order",
"limit",
"offset",
"having",
"window",
"fetch",
"for",
"union",
"intersect",
"except",
];
KWS.iter().any(|k| word.eq_ignore_ascii_case(k))
}
fn is_predicate_bail_kw(word: &str) -> bool {
const KWS: [&str; 8] = ["exists", "any", "all", "some", "values", "array", "case", "select"];
KWS.iter().any(|k| word.eq_ignore_ascii_case(k))
}
const MAX_IN_LIST: usize = 128;
fn is_whitelisted_cast_type(ty: &str) -> bool {
const TYPES: [&str; 20] = [
"uuid",
"int2",
"int4",
"int8",
"smallint",
"integer",
"int",
"bigint",
"text",
"varchar",
"date",
"timestamp",
"timestamptz",
"numeric",
"decimal",
"float4",
"float8",
"real",
"boolean",
"bool",
];
TYPES.iter().any(|k| ty.eq_ignore_ascii_case(k))
}
fn lex_pure_literal_list(sql: &str, bytes: &[u8], open_idx: usize) -> Option<(usize, Vec<Value>)> {
let n = bytes.len();
let mut j = open_idx + 1;
let mut elems = Vec::new();
loop {
j = skip_ws(bytes, j);
if j >= n {
return None;
}
match bytes[j] {
b'\'' => {
let end = skip_single_quoted(bytes, j)?;
let content = sql[j + 1..end - 1].replace("''", "'");
elems.push(Planner::single_quoted_content_to_value(&content));
j = end;
}
b'+' | b'-' => {
let sign_start = j;
if j + 1 >= n || !bytes[j + 1].is_ascii_digit() {
return None;
}
let end = read_number(bytes, j + 1)?;
elems.push(Planner::number_literal_to_value(&sql[sign_start..end]).ok()?);
j = end;
}
b if b.is_ascii_digit() => {
let end = read_number(bytes, j)?;
elems.push(Planner::number_literal_to_value(&sql[j..end]).ok()?);
j = end;
}
_ => return None,
}
j = skip_ws(bytes, j);
if j >= n {
return None;
}
match bytes[j] {
b',' => {
j += 1;
}
b')' => return Some((j + 1, elems)),
_ => return None,
}
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
fn norm(sql: &str) -> Option<(String, Vec<Value>)> {
normalize_select_literals(sql)
}
#[test]
fn simple_int_eq() {
let (s, p) = norm("SELECT abalance FROM t50 WHERE aid = 5").unwrap();
assert_eq!(s, "SELECT abalance FROM t50 WHERE aid = $1");
assert_eq!(p, vec![Value::Int4(5)]);
}
#[test]
fn multi_predicate() {
let (s, p) = norm("SELECT x FROM t WHERE a = 5 AND b = 'hi'").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE a = $1 AND b = $2");
assert_eq!(p, vec![Value::Int4(5), Value::String("hi".to_string())]);
}
#[test]
fn string_with_doubled_quote() {
let (s, p) = norm("SELECT x FROM t WHERE name = 'O''Brien'").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE name = $1");
assert_eq!(p, vec![Value::String("O'Brien".to_string())]);
}
#[test]
fn float_literal() {
let (_s, p) = norm("SELECT x FROM t WHERE ratio = 3.5").unwrap();
assert_eq!(p, vec![Value::Float8(3.5)]);
}
#[test]
fn negative_stays_operator() {
let (s, p) = norm("SELECT x FROM t WHERE v = -7").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE v = -$1");
assert_eq!(p, vec![Value::Int4(7)]);
}
#[test]
fn literal_in_select_list_not_touched() {
let (s, p) = norm("SELECT 1, x FROM t WHERE a = 2").unwrap();
assert_eq!(s, "SELECT 1, x FROM t WHERE a = $1");
assert_eq!(p, vec![Value::Int4(2)]);
}
#[test]
fn limit_offset_parameterized() {
let (s, p) = norm("SELECT x FROM t WHERE a = 5 LIMIT 10 OFFSET 3").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE a = $1 LIMIT $2 OFFSET $3");
assert_eq!(p, vec![Value::Int4(5), Value::Int4(10), Value::Int4(3)]);
}
#[test]
fn limit_offset_only_no_where() {
let (s, p) = norm("SELECT x FROM t LIMIT 20 OFFSET 40").unwrap();
assert_eq!(s, "SELECT x FROM t LIMIT $1 OFFSET $2");
assert_eq!(p, vec![Value::Int4(20), Value::Int4(40)]);
}
#[test]
fn offset_before_limit_parameterized() {
let (s, p) = norm("SELECT x FROM t OFFSET 40 LIMIT 20").unwrap();
assert_eq!(s, "SELECT x FROM t OFFSET $1 LIMIT $2");
assert_eq!(p, vec![Value::Int4(40), Value::Int4(20)]);
}
#[test]
fn offset_only_no_limit_parameterized() {
let (s, p) = norm("SELECT x FROM t OFFSET 5").unwrap();
assert_eq!(s, "SELECT x FROM t OFFSET $1");
assert_eq!(p, vec![Value::Int4(5)]);
}
#[test]
fn limit_expression_takes_only_first_bare_integer() {
let (s, p) = norm("SELECT x FROM t LIMIT 10+5").unwrap();
assert_eq!(s, "SELECT x FROM t LIMIT $1+5");
assert_eq!(p, vec![Value::Int4(10)]);
let (s2, p2) = norm("SELECT x FROM t LIMIT 5,10").unwrap();
assert_eq!(s2, "SELECT x FROM t LIMIT $1,10");
assert_eq!(p2, vec![Value::Int4(5)]);
}
#[test]
fn limit_non_integer_left_inline() {
let (s, p) = norm("SELECT x FROM t WHERE a = 1 LIMIT 5.5").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE a = $1 LIMIT 5.5");
assert_eq!(p, vec![Value::Int4(1)]);
}
#[test]
fn limit_all_untouched() {
let (s, p) = norm("SELECT x FROM t WHERE a = 5 LIMIT ALL").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE a = $1 LIMIT ALL");
assert_eq!(p, vec![Value::Int4(5)]);
}
#[test]
fn limit_inside_subquery_not_touched() {
let (s, p) = norm("SELECT * FROM (SELECT y FROM z LIMIT 5) t WHERE a = 1").unwrap();
assert_eq!(s, "SELECT * FROM (SELECT y FROM z LIMIT 5) t WHERE a = $1");
assert_eq!(p, vec![Value::Int4(1)]);
}
#[test]
fn existing_limit_placeholder_still_bails() {
assert!(norm("SELECT x FROM t WHERE a = 5 LIMIT $1").is_none());
}
#[test]
fn qualified_column_not_a_number() {
let (s, p) = norm("SELECT x FROM t WHERE t.col1 = 9").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE t.col1 = $1");
assert_eq!(p, vec![Value::Int4(9)]);
}
#[test]
fn in_list_padded_to_power_of_two() {
let (s, p) = norm("SELECT x FROM t WHERE a IN (1,2,3)").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE a IN ($1,$2,$3,$3)");
assert_eq!(p, vec![Value::Int4(1), Value::Int4(2), Value::Int4(3)]);
}
#[test]
fn in_list_exact_power_of_two_not_padded() {
let (s, p) = norm("SELECT x FROM t WHERE a IN (1,2)").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE a IN ($1,$2)");
assert_eq!(p.len(), 2);
}
#[test]
fn in_list_arity_buckets_share_shape() {
let (s1, _) = norm("SELECT x FROM t WHERE a IN (1,2,3)").unwrap();
let (s2, _) = norm("SELECT x FROM t WHERE a IN (7,8,9)").unwrap();
assert_eq!(s1, s2);
let (s5, p5) = norm("SELECT x FROM t WHERE a IN (1,2,3,4,5)").unwrap();
let (s7, _) = norm("SELECT x FROM t WHERE a IN (1,2,3,4,5,6,7)").unwrap();
assert_eq!(s5.matches('$').count(), 8);
assert_eq!(s7.matches('$').count(), 8);
assert_eq!(
p5.len(),
5,
"params carry the true arity, padding repeats the last placeholder"
);
}
#[test]
fn in_list_strings_and_signs() {
let (s, p) = norm("SELECT x FROM t WHERE name IN ('a', 'o''b') AND v IN (-1, +2)").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE name IN ($1,$2) AND v IN ($3,$4)");
assert_eq!(
p,
vec![
Value::String("a".into()),
Value::String("o'b".into()),
Value::Int4(-1),
Value::Int4(2),
]
);
}
#[test]
fn not_in_padded_same_as_in() {
let (s, _) = norm("SELECT x FROM t WHERE a NOT IN (1,2,3)").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE a NOT IN ($1,$2,$3,$3)");
}
#[test]
fn in_list_over_cap_bails() {
let list: Vec<String> = (1..=129).map(|i| i.to_string()).collect();
let sql = format!("SELECT x FROM t WHERE a IN ({})", list.join(","));
assert!(norm(&sql).is_none());
}
#[test]
fn in_without_paren_bails_position() {
assert!(norm("SELECT x FROM t WHERE POSITION('x' IN name) = 1").is_none());
}
#[test]
fn in_with_null_or_columns_falls_back_unpadded() {
let (s, p) = norm("SELECT x FROM t WHERE a IN (1, NULL)").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE a IN ($1, NULL)");
assert_eq!(p, vec![Value::Int4(1)]);
let (s2, p2) = norm("SELECT x FROM t WHERE a IN (1, b)").unwrap();
assert_eq!(s2, "SELECT x FROM t WHERE a IN ($1, b)");
assert_eq!(p2.len(), 1);
}
#[test]
fn in_subquery_still_bails() {
assert!(norm("SELECT x FROM t WHERE a IN (SELECT b FROM u)").is_none());
assert!(norm("SELECT x FROM t WHERE a NOT IN (SELECT b FROM u)").is_none());
}
#[test]
fn between_parameterizes_both_bounds() {
let (s, p) = norm("SELECT x FROM t WHERE a BETWEEN 1 AND 10").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE a BETWEEN $1 AND $2");
assert_eq!(p, vec![Value::Int4(1), Value::Int4(10)]);
let (s2, _) = norm("SELECT x FROM t WHERE a NOT BETWEEN 1 AND 10").unwrap();
assert_eq!(s2, "SELECT x FROM t WHERE a NOT BETWEEN $1 AND $2");
}
#[test]
fn whitelisted_cast_passes_through() {
let (s, p) = norm("SELECT x FROM t WHERE id = '11111111-1111-1111-1111-111111111111'::uuid").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE id = $1::uuid");
assert_eq!(p.len(), 1);
let (s2, _) = norm("SELECT x FROM t WHERE ts = '2026-01-01'::timestamp AND n = 5::int8").unwrap();
assert_eq!(s2, "SELECT x FROM t WHERE ts = $1::timestamp AND n = $2::int8");
let (s3, p3) = norm("SELECT x FROM t WHERE c::int = 5").unwrap();
assert_eq!(s3, "SELECT x FROM t WHERE c::int = $1");
assert_eq!(p3, vec![Value::Int4(5)]);
}
#[test]
fn non_whitelisted_or_precision_cast_bails() {
assert!(norm("SELECT x FROM t WHERE v = '1 day'::interval").is_none());
assert!(norm("SELECT x FROM t WHERE v = 'abc'::varchar(10)").is_none());
assert!(norm("SELECT x FROM t WHERE v = '{}'::jsonb").is_none());
}
#[test]
fn bails_on_subquery() {
assert!(norm("SELECT x FROM t WHERE a = (SELECT max(b) FROM u WHERE b = 3)").is_none());
}
#[test]
fn bails_on_comment() {
assert!(norm("SELECT x FROM t WHERE a = 5 -- c").is_none());
assert!(norm("SELECT x FROM t WHERE a = 5 /* c */").is_none());
}
#[test]
fn bails_on_escaped_and_dollar() {
assert!(norm(r"SELECT x FROM t WHERE b = E'\x41'").is_none());
assert!(norm("SELECT x FROM t WHERE a = $1").is_none());
assert!(norm("SELECT x FROM t WHERE b = $$hi$$").is_none());
}
#[test]
fn bails_without_where() {
assert!(norm("SELECT 1").is_none());
assert!(norm("SELECT x FROM t").is_none());
}
#[test]
fn bails_on_non_select() {
assert!(norm("UPDATE t SET x = 1 WHERE a = 2").is_none());
assert!(norm("WITH c AS (SELECT 1) SELECT * FROM c WHERE a = 2").is_none());
assert!(norm("INSERT INTO t VALUES (1)").is_none());
}
#[test]
fn bails_on_second_statement() {
assert!(norm("SELECT x FROM t WHERE a = 5; DROP TABLE t").is_none());
}
#[test]
fn trailing_semicolon_ok() {
let (s, p) = norm("SELECT x FROM t WHERE a = 5;").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE a = $1;");
assert_eq!(p, vec![Value::Int4(5)]);
}
#[test]
fn leading_dot_decimal_bails() {
assert!(norm("SELECT x FROM t WHERE r = .5").is_none());
}
#[test]
fn string_in_select_list_not_touched() {
let (s, p) = norm("SELECT 'label', x FROM t WHERE a = 1").unwrap();
assert_eq!(s, "SELECT 'label', x FROM t WHERE a = $1");
assert_eq!(p, vec![Value::Int4(1)]);
}
#[test]
fn like_predicate_ok() {
let (s, p) = norm("SELECT x FROM t WHERE name LIKE 'a%'").unwrap();
assert_eq!(s, "SELECT x FROM t WHERE name LIKE $1");
assert_eq!(p, vec![Value::String("a%".to_string())]);
}
#[test]
fn big_int_types_as_int8() {
let (_s, p) = norm("SELECT x FROM t WHERE a = 5000000000").unwrap();
assert_eq!(p, vec![Value::Int8(5_000_000_000)]);
}
#[test]
fn does_not_panic_on_random_ish_input() {
for s in [
"SELECT",
"SELECT ''''",
"SELECT x FROM t WHERE a = '",
"SELECT x FROM t WHERE = = =",
"SELECT x FROM t WHERE a = 1.2.3",
"select x from t where a=1and b=2",
"SELECT x FROM t WHERE a = 1)))",
] {
let _ = norm(s);
}
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used)]
mod differential {
use super::normalize_select_literals;
use crate::{EmbeddedDatabase, Tuple};
fn repr(rows: &[Tuple]) -> Vec<String> {
let mut out: Vec<String> = rows
.iter()
.map(|t| t.values.iter().map(|v| format!("{v:?}")).collect::<Vec<_>>().join("|"))
.collect();
out.sort();
out
}
fn seed() -> EmbeddedDatabase {
let db = EmbeddedDatabase::new_in_memory().unwrap();
db.execute(
"CREATE TABLE d (id INT PRIMARY KEY, n INT, big INT8, r FLOAT8, name TEXT, flag BOOLEAN, note TEXT)",
)
.unwrap();
db.execute("CREATE INDEX d_name ON d(name)").unwrap();
db.execute("CREATE INDEX d_n ON d(n)").unwrap();
let rows = [
"(1, 10, 5000000000, 1.5, 'alice', true, 'x')",
"(2, 20, 6000000000, -2.25, 'o''brien', false, NULL)",
"(3, 10, 7000000000, 0.0, 'carol', true, 'y')",
"(4, 40, 8000000000, 3.14159, 'dave', false, 'z')",
"(5, 20, 9000000000, -100.5, '', true, 'alice')",
"(6, 30, 1000000000, 42.0, 'Eve', false, 'x')",
];
for r in rows {
db.execute(&format!("INSERT INTO d VALUES {r}")).unwrap();
}
db
}
#[test]
fn raw_equals_normalized_over_corpus() {
let db = seed();
let corpus = [
"SELECT id FROM d WHERE n = 10",
"SELECT id FROM d WHERE n = 20 AND flag = true",
"SELECT id, name FROM d WHERE name = 'alice'",
"SELECT id FROM d WHERE name = 'o''brien'",
"SELECT id FROM d WHERE name = ''",
"SELECT id FROM d WHERE big = 7000000000",
"SELECT id FROM d WHERE r = 3.14159",
"SELECT id FROM d WHERE r = -2.25",
"SELECT id FROM d WHERE n = -5",
"SELECT id FROM d WHERE n > 15 AND n < 35",
"SELECT id FROM d WHERE n >= 20 AND r <= 0.0",
"SELECT id FROM d WHERE name <> 'alice'",
"SELECT id FROM d WHERE name LIKE 'a%'",
"SELECT id FROM d WHERE note = 'x' AND flag = false",
"SELECT id FROM d WHERE n = 20 OR name = 'dave'",
"SELECT id FROM d WHERE (n = 10 OR n = 20) AND flag = true",
"SELECT count(*) FROM d WHERE n = 10",
"SELECT id FROM d WHERE n = 10 ORDER BY id DESC",
"SELECT id FROM d WHERE n = 20 LIMIT 1",
"SELECT id FROM d WHERE id = 3",
"SELECT id FROM d WHERE name = 'Eve'",
"SELECT id FROM d WHERE big > 5000000000 AND big < 9000000000",
"SELECT id, n, r FROM d WHERE flag = true AND n = 10",
"SELECT id FROM d WHERE abs(r) = 2.25",
"SELECT id FROM d WHERE n = 10;",
"SELECT id FROM d WHERE n IN (10)",
"SELECT id FROM d WHERE n IN (10, 20)",
"SELECT id FROM d WHERE n IN (10, 20, 30)",
"SELECT id FROM d WHERE n IN (10, 20, 30, 40, 50)",
"SELECT id FROM d WHERE n IN (10, 10, 20)",
"SELECT id FROM d WHERE n NOT IN (10, 30)",
"SELECT id FROM d WHERE n NOT IN (10, 10)",
"SELECT id FROM d WHERE name IN ('alice', 'o''brien', '')",
"SELECT id FROM d WHERE n IN (-5, 10, 20)",
"SELECT id FROM d WHERE big IN (5000000000, 7000000000)",
"SELECT id FROM d WHERE n IN (1, NULL)",
"SELECT id FROM d WHERE n NOT IN (10, NULL)",
"SELECT id FROM d WHERE n BETWEEN 10 AND 25",
"SELECT id FROM d WHERE n NOT BETWEEN 10 AND 25",
"SELECT id FROM d WHERE r BETWEEN -3.0 AND 1.6",
"SELECT id FROM d WHERE big = 7000000000::int8",
"SELECT id FROM d WHERE n = 10 AND name IN ('alice','carol') OR r BETWEEN 40.0 AND 43.0",
"SELECT id FROM d WHERE n = 20 ORDER BY id LIMIT 1 OFFSET 1",
"SELECT id FROM d WHERE flag = true ORDER BY id LIMIT 2 OFFSET 1",
"SELECT id FROM d ORDER BY id LIMIT 3 OFFSET 2",
"SELECT id FROM d ORDER BY id DESC LIMIT 100 OFFSET 0",
"SELECT id FROM d ORDER BY id LIMIT 0",
"SELECT id FROM d ORDER BY id LIMIT 2 OFFSET 999",
"SELECT id FROM d WHERE n = 10 LIMIT ALL",
"SELECT id FROM d WHERE id IN (SELECT id FROM d LIMIT 3)",
];
let mut normalized_count = 0;
for raw in corpus {
let raw_rows = repr(
&db.query_raw_unnormalized(raw)
.unwrap_or_else(|e| panic!("raw failed for {raw}: {e}")),
);
if let Some((nsql, params)) = normalize_select_literals(raw) {
normalized_count += 1;
let norm_rows = db
.query_params(&nsql, ¶ms)
.unwrap_or_else(|e| panic!("normalized failed for {raw} -> {nsql}: {e}"));
assert_eq!(
raw_rows,
repr(&norm_rows),
"\nMISMATCH\n raw: {raw}\n norm: {nsql}\n params: {params:?}"
);
}
}
assert!(
normalized_count >= 45,
"expected most of the corpus to normalize, got {normalized_count}"
);
}
#[test]
fn property_fuzz_raw_equals_normalized() {
let db = seed();
let mut rng_state: u64 = 0x5eed_1e57_0BAD_F00D;
let mut next = move || {
rng_state = rng_state
.wrapping_mul(6364136223846793005)
.wrapping_add(1442695040888963407);
(rng_state >> 33) as u32
};
let int_cols = ["n", "id"];
let ops = ["=", ">", "<", ">=", "<=", "<>"];
let names = ["alice", "o''brien", "", "carol", "dave", "Eve", "zzz"];
let mut term = |next: &mut dyn FnMut() -> u32| -> String {
match next() % 5 {
0 => {
let col = int_cols[(next() % 2) as usize];
let op = ops[(next() % 6) as usize];
let v = (next() % 60) as i64 - 10;
format!("{col} {op} {v}")
}
1 => {
let col = int_cols[(next() % 2) as usize];
let arity = 1 + (next() % 9) as usize;
let vals: Vec<String> = (0..arity).map(|_| ((next() % 60) as i64 - 10).to_string()).collect();
let neg = if next() % 4 == 0 { "NOT " } else { "" };
format!("{col} {neg}IN ({})", vals.join(", "))
}
2 => {
let lo = (next() % 40) as i64 - 10;
let hi = lo + (next() % 30) as i64;
format!("n BETWEEN {lo} AND {hi}")
}
3 => {
let name = names[(next() % names.len() as u32) as usize];
format!("name = '{name}'")
}
_ => {
let v = ((next() % 800) as f64 / 100.0) - 4.0;
format!("r > {v:.2}")
}
}
};
let mut normalized_count = 0usize;
for _ in 0..300 {
let nterms = 1 + (next() % 3) as usize;
let mut pred = term(&mut next);
for _ in 1..nterms {
let joiner = if next() % 2 == 0 { "AND" } else { "OR" };
pred = format!("{pred} {joiner} {}", term(&mut next));
}
let raw = format!("SELECT id FROM d WHERE {pred}");
let raw_rows = repr(
&db.query_raw_unnormalized(&raw)
.unwrap_or_else(|e| panic!("raw failed for {raw}: {e}")),
);
if let Some((nsql, params)) = normalize_select_literals(&raw) {
normalized_count += 1;
let norm_rows = db
.query_params(&nsql, ¶ms)
.unwrap_or_else(|e| panic!("normalized failed for {raw} -> {nsql}: {e}"));
assert_eq!(
raw_rows,
repr(&norm_rows),
"\nFUZZ MISMATCH\n raw: {raw}\n norm: {nsql}\n params: {params:?}"
);
}
}
assert!(
normalized_count >= 250,
"fuzz corpus should mostly normalize, got {normalized_count}/300"
);
}
}