From d1f024bf5e69d679d597d62e7f9f71862f211e7a Mon Sep 17 00:00:00 2001 From: Lev Kokotov Date: Mon, 14 Sep 2026 10:12:46 -0700 Subject: [PATCH] fix: normalize segfault --- src/make.rs | 5 +- src/normalize.rs | 265 +++++++++++++++++++++++++++++++++++--- tests/e2e_hasura.rs | 33 ++++- tests/normalize.rs | 249 +++++++++++++++++++++++++++++++++++ tests/postgres_regress.rs | 18 ++- 5 files changed, 547 insertions(+), 23 deletions(-) create mode 100644 tests/normalize.rs diff --git a/src/make.rs b/src/make.rs index 83e7f3c..db97f2d 100644 --- a/src/make.rs +++ b/src/make.rs @@ -137,7 +137,10 @@ impl<'mem> MemoryToken<'mem> { raw_stmt } - pub fn make_returning_clause(self, exprs: Unique<'mem, &NodeList>) -> Unique<'mem, &'mem nodes::ReturningClause> { + pub fn make_returning_clause( + self, + exprs: Unique<'mem, &NodeList>, + ) -> Unique<'mem, &'mem nodes::ReturningClause> { let mut return_clause = self.make_node::(); return_clause.as_mut().set_exprs(exprs); return_clause diff --git a/src/normalize.rs b/src/normalize.rs index b2645ce..69d5f0e 100644 --- a/src/normalize.rs +++ b/src/normalize.rs @@ -1,26 +1,193 @@ -use crate::transform::{Transform, TransformClosure}; -use crate::{DeparseResult, Node, NodeMut, Owned, deparse, make, nodes, parse}; +use std::{collections::HashSet, ptr}; + +use crate::list::NodeList; +use crate::list_mut::NodeListMut; +use crate::transform::{self, Assignable, Transform}; +use crate::{ConstValue, DeparseResult, Node, NodeMut, Owned, deparse, make, nodes, parse}; pub fn normalize(query: &nodes::RawStmt) -> Owned { make::owned(|mem| { let mut copy = mem.make_unique(query); - let mut param_count = 0; - TransformClosure::new(|mut node| match &mut *node { + Normalizer { + mem, + param_count: 0, + ordinals: HashSet::new(), + } + .transform_raw_stmt(copy.as_mut()); + copy + }) +} + +struct Normalizer<'mem> { + mem: make::MemoryToken<'mem>, + param_count: i32, + // Compare node identities only: the copied AST owns these constants for + // the entire traversal. Equal integer values elsewhere remain expressions. + ordinals: HashSet<*const nodes::A_Const>, +} + +impl<'mem> Normalizer<'mem> { + fn preserve_ordinal(&mut self, node: Node<'_>) { + if let Node::A_Const(value) = node + && matches!(value.val(), Some(ConstValue::Integer(_))) + { + self.ordinals.insert(ptr::from_ref(value)); + } + } + + // GROUP BY GROUPING SETS ((1, 2), ROLLUP(1)) refers to output columns, + // so the integers must survive even inside nested grouping lists. + // GROUP BY ROW(1, 2) instead becomes GROUP BY ROW($1, $2). + fn preserve_group_ordinals(&mut self, node: Node<'_>) { + match node { + Node::GroupingSet(group) => { + for node in group.content() { + self.preserve_group_ordinals(node); + } + } + Node::RowExpr(row) if row.row_format == nodes::CoercionForm::COERCE_IMPLICIT_CAST => { + // GROUP BY (1, 2) is a grouping list; ROW(1, 2) is an expression. + for node in row.args() { + self.preserve_group_ordinals(node); + } + } + _ => self.preserve_ordinal(node), + } + } + + // These SQL constructs mix expression arguments with syntax constants. + // Copy only the expression argument so it can be replaced independently. + fn normalize_first_arg(&mut self, mut args: NodeListMut<'mem, '_, NodeList>) { + if let Some(arg) = args.get(0) { + let mut arg = self.mem.make_unique(arg); + self.transform_node(Assignable::new(&mut arg)); + args.set(0, arg); + } + } +} + +impl<'mem> Transform<'mem> for Normalizer<'mem> { + fn transform_node<'mutref>(&mut self, mut node: Assignable<'mem, 'mutref>) { + match &mut *node { + NodeMut::A_Const(value) if self.ordinals.contains(&ptr::from_ref(&**value)) => {} NodeMut::A_Const(_) => { - param_count += 1; - node.replace(mem.make_param_ref(param_count).uncast()); - None + self.param_count += 1; + node.replace(self.mem.make_param_ref(self.param_count).uncast()); } NodeMut::ParamRef(p) => { - param_count += 1; - p.set_number(param_count); - None + self.param_count += 1; + p.set_number(self.param_count); } - _ => Some(node), - }) - .transform_raw_stmt(copy.as_mut()); - copy - }) + _ => transform::transform_node(node.into_inner(), self), + } + } + + // SQL: SELECT 42 GROUP BY 1 ORDER BY 1 + // Normalized: SELECT $1 GROUP BY 1 ORDER BY 1 + // The ordinal 1 identifies the first output column; $2 would be a constant + // expression instead. Expressions like ORDER BY a + 1 still become a + $n. + // Likewise, sum(a ORDER BY 1) and OVER (ORDER BY 1) use constant expressions, + // not output-column positions, so those integers must become parameters. + fn transform_select_stmt<'mutref>(&mut self, node: nodes::SelectStmtMut<'mem, 'mutref>) { + for group in node.group_clause() { + self.preserve_group_ordinals(group); + } + for sort in node.sort_clause() { + self.preserve_ordinal(sort.node()); + } + // Use the generated traversal to retain its parameter numbering order. + // Only SELECT's sort clause has ordinals: aggregate and window ORDER BY + // integers are ordinary expressions and must still be normalized. + transform::transform_select_stmt(node, self); + } + + // Transaction options are syntax constants, not expressions. The native + // deparser reads their values as A_Const nodes; replacing an isolation + // level with ParamRef makes it dereference an invalid string pointer. + // SQL: BEGIN ISOLATION LEVEL REPEATABLE READ READ ONLY + // Keep both options unchanged. REPEATABLE READ is stored as a string + // constant and READ ONLY as integer 1; neither can become $1 or $2. + fn transform_transaction_stmt<'mutref>( + &mut self, + _node: nodes::TransactionStmtMut<'mem, 'mutref>, + ) { + } + + // SET arguments also require literal values, including SET TRANSACTION. + // Override the typed visitor so this also covers SET nested in ALTER ROLE, + // ALTER DATABASE and function options. + // SQL: SET TRANSACTION ISOLATION LEVEL SERIALIZABLE + // Keep SERIALIZABLE intact for the same reason as BEGIN's isolation level. + // SET TIME ZONE 'UTC' and ALTER ROLE app SET statement_timeout = 1000 + // also retain their setting values, which are not parameter expressions. + fn transform_variable_set_stmt<'mutref>( + &mut self, + _node: nodes::VariableSetStmtMut<'mem, 'mutref>, + ) { + } + + // Type modifiers encode precision, length and interval units. They are + // not expression parameters, even when stored as A_Const nodes. + // SQL: SELECT 'abc'::varchar(10), INTERVAL '1.234' SECOND(2) + // Normalized: SELECT $1::varchar(10), $2::interval SECOND(2) + // varchar($n) is invalid syntax. Replacing interval's internal unit/precision + // constants can also make the deparser read the wrong units or precision. + fn transform_type_name<'mutref>(&mut self, _node: nodes::TypeNameMut<'mem, 'mutref>) {} + + // JSON_TABLE paths must remain string constants; the deparser directly + // reads their A_Const string values. + // SQL: SELECT * FROM JSON_TABLE('[1]'::jsonb, '$[*]' + // COLUMNS (value int PATH '$')) AS jt + // Only '[1]' becomes $1; '$[*]' and '$' must remain string constants. + // Replacing a path with ParamRef makes the deparser read an invalid pointer. + fn transform_json_table_path_spec<'mutref>( + &mut self, + _node: nodes::JsonTablePathSpecMut<'mem, 'mutref>, + ) { + } + + // CYCLE mark/default values use the AexprConst grammar, not expressions. + // SQL: WITH RECURSIVE t(id) AS (SELECT 1) + // CYCLE id SET cycle TO 'yes' DEFAULT 'no' USING path SELECT * FROM t + // Normalize SELECT 1 to SELECT $1, but keep 'yes' and 'no': parameters in + // those positions are rejected by the parser/deparser. + fn transform_cte_cycle_clause<'mutref>( + &mut self, + _node: nodes::CTECycleClauseMut<'mem, 'mutref>, + ) { + } + + fn transform_xml_expr<'mutref>(&mut self, mut node: nodes::XmlExprMut<'mem, 'mutref>) { + if node.op == nodes::XmlExprOp::IS_XMLROOT { + // VERSION and STANDALONE are syntax constants, unlike the XML value. + // SQL: XMLROOT(''::xml, VERSION '1.0', STANDALONE YES) + // Normalized: XMLROOT($1::xml, VERSION '1.0', STANDALONE YES) + // The deparser reads VERSION's null flag and STANDALONE's integer + // flag directly; replacing them can silently change the options. + self.normalize_first_arg(node.args_mut()); + } else { + transform::transform_xml_expr(node, self); + } + } + + fn transform_func_call<'mutref>(&mut self, mut node: nodes::FuncCallMut<'mem, 'mutref>) { + let mut names = node.funcname().iter().filter_map(Node::as_str); + let normalization_syntax = node.funcformat == nodes::CoercionForm::COERCE_SQL_SYNTAX + && names.next() == Some("pg_catalog") + && matches!(names.next(), Some("normalize" | "is_normalized")) + && names.next().is_none(); + if normalization_syntax { + // NFC/NFD/NFKC/NFKD is a keyword stored as an A_Const argument. + // NORMALIZE('hello', NFC) becomes NORMALIZE($1, NFC), and + // 'hello' IS NFC NORMALIZED becomes $1 IS NFC NORMALIZED. + // Turning NFC into $2 breaks the deparser's constant-value read. + // An ordinary call pg_catalog.normalize('hello', 'NFC') instead + // becomes pg_catalog.normalize($1, $2), since both are expressions. + self.normalize_first_arg(node.args_mut()); + } else { + transform::transform_func_call(node, self); + } + } } pub fn normalize_str(query: &str) -> crate::Result { @@ -45,3 +212,71 @@ fn test_normalize_does_the_thing() { assert!(normalize_str("").is_err()); } + +#[test] +fn test_normalize_transaction_options() { + for query in [ + "BEGIN ISOLATION LEVEL REPEATABLE READ READ ONLY", + "BEGIN ISOLATION LEVEL READ UNCOMMITTED READ WRITE", + "START TRANSACTION ISOLATION LEVEL SERIALIZABLE READ ONLY DEFERRABLE", + "BEGIN READ WRITE NOT DEFERRABLE", + "SET TRANSACTION ISOLATION LEVEL READ COMMITTED READ ONLY", + "SET SESSION CHARACTERISTICS AS TRANSACTION ISOLATION LEVEL SERIALIZABLE", + ] { + let ast = parse(query).expect("valid transaction options"); + let stmt = ast.first().expect("one statement"); + let original = deparse(stmt).expect("original statement deparses"); + let normalized = normalize_str(query).expect("normalized statement deparses"); + assert_eq!(normalized.as_str(), original.as_str(), "{query}"); + } +} + +#[test] +fn test_normalize_set_arguments() { + for query in [ + "SET client_encoding = 'UTF8'", + "SET client_min_messages TO WARNING", + "SET TIME ZONE 'UTC'", + "SET statement_timeout = 1000", + "ALTER ROLE postgres SET default_transaction_isolation TO 'repeatable read'", + "ALTER DATABASE postgres SET statement_timeout = 1000", + "CREATE FUNCTION f() RETURNS int LANGUAGE SQL SET statement_timeout = 1000 AS 'SELECT 1'", + ] { + let ast = parse(query).expect("valid SET arguments"); + let stmt = ast.first().expect("one statement"); + let original = deparse(stmt).expect("original statement deparses"); + let normalized = normalize_str(query).expect("normalized statement deparses"); + assert_eq!(normalized.as_str(), original.as_str(), "{query}"); + } +} + +#[test] +fn test_normalize_query_expressions() { + for (query, expected) in [ + ("SELECT 42, 'hello', $7", "SELECT $1, $2, $3"), + ("INSERT INTO t VALUES (42)", "INSERT INTO t VALUES ($1)"), + ( + "UPDATE t SET id = 42 WHERE id = 1", + "UPDATE t SET id = $1 WHERE id = $2", + ), + ("DELETE FROM t WHERE id = 42", "DELETE FROM t WHERE id = $1"), + ("EXPLAIN SELECT 42", "EXPLAIN SELECT $1"), + ] { + let ast = parse(query).expect("valid query"); + let stmt = ast.first().expect("one statement"); + let original = deparse(stmt).expect("original statement deparses"); + let normalized = normalize(stmt); + assert_eq!( + deparse(&*normalized) + .expect("normalized query deparses") + .as_str(), + expected + ); + assert_eq!( + deparse(stmt) + .expect("original query still deparses") + .as_str(), + original.as_str() + ); + } +} diff --git a/tests/e2e_hasura.rs b/tests/e2e_hasura.rs index 944cb23..a897a7c 100644 --- a/tests/e2e_hasura.rs +++ b/tests/e2e_hasura.rs @@ -1,5 +1,5 @@ use flate2::read::GzDecoder; -use pg_raw_parse::{deparse, parse}; +use pg_raw_parse::{deparse, normalize::normalize, parse}; use std::fs::File; use std::io::{BufRead, BufReader}; use std::path::Path; @@ -14,6 +14,7 @@ fn hasura_e2e_statements_parse_and_deparse() { let mut query = String::new(); let mut records = 0; let mut statements = 0; + let mut reparse_exceptions = 0; for (line_number, line) in reader.lines().enumerate() { let line = line.unwrap_or_else(|error| { @@ -36,29 +37,49 @@ fn hasura_e2e_statements_parse_and_deparse() { } records += 1; - parse_and_deparse_record(&query, records, &mut statements); + reparse_exceptions += parse_and_deparse_record(&query, records, &mut statements); query.clear(); } if !query.trim().is_empty() { records += 1; - parse_and_deparse_record(&query, records, &mut statements); + reparse_exceptions += parse_and_deparse_record(&query, records, &mut statements); } assert_eq!(records, 113_289, "unexpected Hasura SQL record count"); assert_eq!(statements, 136_459, "unexpected Hasura SQL statement count"); + assert_eq!( + reparse_exceptions, 38, + "unexpected pre-existing reparse exceptions" + ); eprintln!("Hasura corpus: {records} records, {statements} statements"); } -fn parse_and_deparse_record(query: &str, record: usize, statements: &mut usize) { +fn parse_and_deparse_record(query: &str, record: usize, statements: &mut usize) -> usize { + let mut reparse_exceptions = 0; let tree = parse(query) .unwrap_or_else(|error| panic!("failed to parse Hasura record {record}: {error}\n{query}")); - for statement in tree.stmts() { + for statement in tree.iter() { *statements += 1; let _ = format!("{statement:?}"); - deparse(statement).unwrap_or_else(|error| { + let original = deparse(statement).unwrap_or_else(|error| { panic!("failed to deparse Hasura statement {statements}: {error}\n{query}") }); + let normalized = normalize(statement); + let normalized = deparse(&*normalized).unwrap_or_else(|error| { + panic!("failed to deparse normalized Hasura statement {statements}: {error}\n{query}") + }); + // The corpus includes SQL that the raw parser accepts but the original + // deparser cannot round-trip. Still check normalization/deparsing for + // those inputs, and pin the exception count above. + if parse(original.as_str()).is_err() { + reparse_exceptions += 1; + } else { + parse(normalized.as_str()).unwrap_or_else(|error| { + panic!("failed to reparse normalized Hasura statement {statements}: {error}\noriginal: {query}\nnormalized: {}", normalized.as_str()) + }); + } } + reparse_exceptions } diff --git a/tests/normalize.rs b/tests/normalize.rs new file mode 100644 index 0000000..41f5f12 --- /dev/null +++ b/tests/normalize.rs @@ -0,0 +1,249 @@ +use pg_raw_parse::{deparse, normalize::normalize, parse}; + +fn assert_normalized(query: &str, expected: &str) { + let original = parse(query).expect("valid input SQL"); + let stmt = original.first().expect("one input statement"); + let before = deparse(stmt).expect("original SQL deparses"); + let expected = parse(expected).expect("valid expected SQL"); + let expected = + deparse(expected.first().expect("one expected statement")).expect("expected SQL deparses"); + + let normalized = normalize(stmt); + let actual = deparse(&*normalized).expect("normalized SQL deparses"); + assert_eq!(actual.as_str(), expected.as_str(), "{query}"); + parse(actual.as_str()).expect("normalized SQL reparses"); + assert_eq!( + deparse(stmt).expect("original SQL still deparses").as_str(), + before.as_str(), + "normalizing must not mutate the original AST" + ); +} + +#[test] +fn type_modifiers_remain_literals() { + for (query, expected) in [ + ( + "SELECT 'abc'::varchar(10), 42::numeric(10, 2)", + "SELECT $1::varchar(10), $2::numeric(10, 2)", + ), + ( + "SELECT 'a'::char, B'101'::bit(3)", + "SELECT $1::char, $2::bit(3)", + ), + ( + "SELECT '2026-09-14'::timestamp(3) with time zone", + "SELECT $1::timestamp(3) with time zone", + ), + ( + "SELECT '12:00'::time(2) with time zone", + "SELECT $1::time(2) with time zone", + ), + ( + "SELECT INTERVAL '1.234' SECOND(2), 42", + "SELECT $1::interval SECOND(2), $2", + ), + ( + "SELECT INTERVAL '1-2' YEAR TO MONTH", + "SELECT $1::interval YEAR TO MONTH", + ), + ( + "CREATE TABLE t (a varchar(20), b numeric(10, 2), c interval DAY TO SECOND(3))", + "CREATE TABLE t (a varchar(20), b numeric(10, 2), c interval DAY TO SECOND(3))", + ), + ] { + assert_normalized(query, expected); + } +} + +#[test] +fn json_table_paths_remain_strings() { + for (query, expected) in [ + ( + "SELECT * FROM JSON_TABLE('[1,2]'::jsonb, '$[*]' COLUMNS (value int PATH '$')) AS jt", + "SELECT * FROM JSON_TABLE($1::jsonb, '$[*]' COLUMNS (value int PATH '$')) AS jt", + ), + ( + "SELECT * FROM JSON_TABLE('{\"a\":[1]}'::jsonb, '$' AS root COLUMNS (NESTED PATH '$.a[*]' AS nested COLUMNS (value int PATH '$' DEFAULT 0 ON EMPTY))) AS jt", + "SELECT * FROM JSON_TABLE($1::jsonb, '$' AS root COLUMNS (NESTED PATH '$.a[*]' AS nested COLUMNS (value int PATH '$' DEFAULT $2 ON EMPTY))) AS jt", + ), + ] { + assert_normalized(query, expected); + } +} + +#[test] +fn xmlroot_options_remain_constants() { + for (query, expected) in [ + ( + "SELECT XMLROOT(''::xml, VERSION '1.0', STANDALONE YES), 42", + "SELECT XMLROOT($1::xml, VERSION '1.0', STANDALONE YES), $2", + ), + ( + "SELECT XMLROOT(''::xml, VERSION NO VALUE, STANDALONE NO)", + "SELECT XMLROOT($1::xml, VERSION NO VALUE, STANDALONE NO)", + ), + ( + "SELECT XMLROOT(''::xml, VERSION '1.0', STANDALONE NO VALUE)", + "SELECT XMLROOT($1::xml, VERSION '1.0', STANDALONE NO VALUE)", + ), + ( + "SELECT XMLROOT(XMLCONCAT(''::xml, ''::xml), VERSION NO VALUE)", + "SELECT XMLROOT(XMLCONCAT($1::xml, $2::xml), VERSION NO VALUE)", + ), + ] { + assert_normalized(query, expected); + } +} + +#[test] +fn unicode_normalization_forms_remain_keywords() { + for (query, expected) in [ + ( + "SELECT NORMALIZE('hello', NFC), 42", + "SELECT NORMALIZE($1, NFC), $2", + ), + ( + "SELECT NORMALIZE('hello', NFD), NORMALIZE('world', NFKC)", + "SELECT NORMALIZE($1, NFD), NORMALIZE($2, NFKC)", + ), + ( + "SELECT 'hello' IS NFKD NORMALIZED", + "SELECT $1 IS NFKD NORMALIZED", + ), + ( + "SELECT 'hello' IS NOT NFC NORMALIZED", + "SELECT $1 IS NOT NFC NORMALIZED", + ), + ("SELECT NORMALIZE('hello')", "SELECT NORMALIZE($1)"), + ( + "SELECT pg_catalog.normalize('hello', 'NFC')", + "SELECT pg_catalog.normalize($1, $2)", + ), + ( + "SELECT pg_catalog.is_normalized('hello', 'NFC')", + "SELECT pg_catalog.is_normalized($1, $2)", + ), + ] { + assert_normalized(query, expected); + } +} + +#[test] +fn cycle_mark_values_remain_constants() { + for (query, expected) in [ + ( + "WITH RECURSIVE t(id) AS (SELECT 1 UNION ALL SELECT id + 1 FROM t) CYCLE id SET cycle USING path SELECT * FROM t", + "WITH RECURSIVE t(id) AS (SELECT $1 UNION ALL SELECT id + $2 FROM t) CYCLE id SET cycle USING path SELECT * FROM t", + ), + ( + "WITH RECURSIVE t(id) AS (SELECT 1) CYCLE id SET cycle TO 'yes' DEFAULT 'no' USING path SELECT 42 FROM t", + "WITH RECURSIVE t(id) AS (SELECT $2) CYCLE id SET cycle TO 'yes' DEFAULT 'no' USING path SELECT $1 FROM t", + ), + ] { + assert_normalized(query, expected); + } +} + +#[test] +fn utility_statements_still_normalize_expressions() { + for (query, expected) in [ + ( + "EXPLAIN (ANALYZE FALSE, COSTS FALSE) SELECT 42", + "EXPLAIN (ANALYZE FALSE, COSTS FALSE) SELECT $1", + ), + ( + "CREATE VIEW v AS SELECT 'abc'::varchar(10)", + "CREATE VIEW v AS SELECT $1::varchar(10)", + ), + ( + "COPY (SELECT 42) TO STDOUT WITH (FORMAT csv, DELIMITER ',')", + "COPY (SELECT $1) TO STDOUT WITH (FORMAT csv, DELIMITER ',')", + ), + ( + "CREATE TABLE t (id int DEFAULT 42 CHECK (id > 0))", + "CREATE TABLE t (id int DEFAULT $1 CHECK (id > $2))", + ), + ] { + assert_normalized(query, expected); + } +} + +#[test] +fn group_and_order_by_ordinals_remain_integers() { + for (query, expected) in [ + ( + "SELECT a, sum(b) FROM t WHERE c = 42 GROUP BY 1 ORDER BY 1 LIMIT 10", + "SELECT a, sum(b) FROM t WHERE c = $1 GROUP BY 1 ORDER BY 1 LIMIT $2", + ), + ( + "SELECT 42, $7 FROM t GROUP BY 1, 2 ORDER BY 2 DESC NULLS LAST, 1", + "SELECT $1, $2 FROM t GROUP BY 1, 2 ORDER BY 2 DESC NULLS LAST, 1", + ), + ( + "SELECT 42 GROUP BY (1) ORDER BY (1)", + "SELECT $1 GROUP BY (1) ORDER BY (1)", + ), + ( + "SELECT 42 GROUP BY 0 ORDER BY -1", + "SELECT $1 GROUP BY 0 ORDER BY -1", + ), + ] { + assert_normalized(query, expected); + } +} + +#[test] +fn grouping_sets_and_nested_query_ordinals_remain_integers() { + for (query, expected) in [ + ( + "SELECT a, b FROM t WHERE c = 42 GROUP BY GROUPING SETS ((1, 2), ROLLUP(1, 2), CUBE(1), ()) ORDER BY 2", + "SELECT a, b FROM t WHERE c = $1 GROUP BY GROUPING SETS ((1, 2), ROLLUP(1, 2), CUBE(1), ()) ORDER BY 2", + ), + ( + "SELECT (SELECT 5 ORDER BY 1), 6 ORDER BY 2", + "SELECT (SELECT $1 ORDER BY 1), $2 ORDER BY 2", + ), + ( + "WITH t AS (SELECT 42 AS a GROUP BY 1 ORDER BY 1) SELECT * FROM t ORDER BY 1", + "WITH t AS (SELECT $1 AS a GROUP BY 1 ORDER BY 1) SELECT * FROM t ORDER BY 1", + ), + ( + "EXPLAIN SELECT 42 GROUP BY 1 ORDER BY 1", + "EXPLAIN SELECT $1 GROUP BY 1 ORDER BY 1", + ), + ( + "SELECT 1 UNION ALL SELECT 2 ORDER BY 1", + "SELECT $1 UNION ALL SELECT $2 ORDER BY 1", + ), + ] { + assert_normalized(query, expected); + } +} + +#[test] +fn numeric_expressions_are_not_ordinals() { + for (query, expected) in [ + ( + "SELECT a FROM t GROUP BY a + 1 ORDER BY a + 2 LIMIT 3", + "SELECT a FROM t GROUP BY a + $1 ORDER BY a + $2 LIMIT $3", + ), + ( + "SELECT a FROM t GROUP BY ROW(1, 2) ORDER BY (3, 4)", + "SELECT a FROM t GROUP BY ROW($1, $2) ORDER BY ($3, $4)", + ), + ( + "SELECT sum(a ORDER BY 1), row_number() OVER (ORDER BY 2) FROM t ORDER BY 1", + "SELECT sum(a ORDER BY $1), row_number() OVER (ORDER BY $2) FROM t ORDER BY 1", + ), + ( + "SELECT percentile_cont(0.5) WITHIN GROUP (ORDER BY 1) FROM t ORDER BY 1", + "SELECT percentile_cont($1) WITHIN GROUP (ORDER BY $2) FROM t ORDER BY 1", + ), + ( + "SELECT a FROM t GROUP BY 'literal' ORDER BY 1.5", + "SELECT a FROM t GROUP BY $1 ORDER BY $2", + ), + ] { + assert_normalized(query, expected); + } +} diff --git a/tests/postgres_regress.rs b/tests/postgres_regress.rs index 0fbf8b1..a86a63c 100644 --- a/tests/postgres_regress.rs +++ b/tests/postgres_regress.rs @@ -1,4 +1,4 @@ -use pg_raw_parse::{deparse_stmts, parse, raw}; +use pg_raw_parse::{deparse, deparse_stmts, normalize::normalize, parse, raw}; use std::ffi::{CStr, CString}; use std::fs; use std::path::{Path, PathBuf}; @@ -282,6 +282,22 @@ fn postgres_regression_sql_parses_and_round_trips() { path.display() ); } + + for stmt in tree.iter() { + let normalized = normalize(stmt); + let normalized = deparse(&*normalized).unwrap_or_else(|error| { + panic!( + "failed to deparse normalized {} at byte {location}: {error}\n{query}", + path.display() + ) + }); + parse(normalized.as_str()).unwrap_or_else(|error| { + panic!( + "failed to reparse normalized {} at byte {location}: {error}\noriginal: {query}\nnormalized: {}", + path.display(), normalized.as_str() + ) + }); + } } }