use super::{FormatOptions, KeywordCase, format_parsed, format_sql};
use crate::tokenizer::TokenKind;
use crate::{BuiltinDialect, ParseOptions, parse_with_builtin_options, tokenize_with_builtin};
const D: BuiltinDialect = BuiltinDialect::Ansi;
fn fmt(sql: &str) -> String {
format_sql(sql, D, &FormatOptions::default()).expect("formats")
}
const CORPUS: &[&str] = &[
"SELECT 1",
"SELECT a, b, c FROM t",
"SELECT DISTINCT a FROM t",
"SELECT ALL a FROM t",
"SELECT a FROM t WHERE a = 1 AND b > 2",
"SELECT a, count(*) FROM t GROUP BY a HAVING count(*) > 1",
"SELECT a FROM t ORDER BY a DESC, b ASC",
"SELECT a FROM t LIMIT 10",
"SELECT a FROM t WHERE a IN (1, 2, 3) ORDER BY a LIMIT 5",
"SELECT a FROM t1 JOIN t2 ON t1.id = t2.id WHERE t1.x > 0",
"SELECT a FROM t1 LEFT JOIN t2 ON t1.id = t2.id",
"SELECT a, b FROM (SELECT a, b FROM inner_t WHERE a > 0) sub",
"SELECT a FROM t WHERE a IN (SELECT id FROM u WHERE u.active)",
"SELECT a FROM t WHERE EXISTS (SELECT 1 FROM u WHERE u.k = t.k)",
"SELECT a FROM t WHERE NOT EXISTS (SELECT 1 FROM u WHERE u.k = t.k)",
"SELECT a, (SELECT max(b) FROM u WHERE u.k = t.k) AS mx FROM t",
"SELECT a FROM t WHERE a > (SELECT 1)",
"SELECT x FROM (SELECT a AS x FROM u WHERE a IN (SELECT id FROM v WHERE v.ok)) s",
"SELECT a FROM t JOIN (SELECT k, max(b) AS mb FROM u GROUP BY k) m ON m.k = t.k",
"SELECT a FROM t WHERE name = '(SELECT b FROM u WHERE b > 0)' AND a IN (SELECT b FROM u WHERE b > 0)",
"SELECT a FROM t1 UNION SELECT a FROM t2",
"SELECT a FROM t1 UNION ALL SELECT a FROM t2 INTERSECT SELECT a FROM t3",
"WITH cte AS (SELECT a FROM t) SELECT a FROM cte",
"WITH a AS (SELECT 1), b AS (SELECT 2) SELECT * FROM a, b",
"SELECT CAST(a AS INT) FROM t",
"SELECT a + b * c - d FROM t",
"SELECT a FROM t WHERE a IS NOT NULL AND b IS NULL",
"SELECT CASE WHEN a > 0 THEN 'pos' ELSE 'neg' END FROM t",
"SELECT \"Quoted\", a FROM \"Table\"",
"INSERT INTO t (a, b) VALUES (1, 2)",
"UPDATE t SET a = 1 WHERE b = 2",
"DELETE FROM t WHERE a = 1",
"CREATE TABLE t (a INT, b VARCHAR(10))",
"SELECT a FROM t; SELECT b FROM u",
];
fn parse_stmts(sql: &str) -> Vec<crate::ast::Statement> {
parse_with_builtin_options(sql, D, ParseOptions::default())
.expect("parses")
.into_statements()
}
#[test]
fn format_output_parses_back_structurally_equal() {
for &sql in CORPUS {
let before = parse_stmts(sql);
let formatted = fmt(sql);
let after = parse_stmts(&formatted);
assert_eq!(
before, after,
"structural drift for {sql:?}\n--- formatted ---\n{formatted}\n---"
);
}
}
#[test]
fn format_output_always_reparses() {
for &sql in CORPUS {
let formatted = fmt(sql);
assert!(
parse_with_builtin_options(&formatted, D, ParseOptions::default()).is_ok(),
"formatted output failed to reparse for {sql:?}:\n{formatted}"
);
}
}
fn spelling_tokens(sql: &str) -> Vec<String> {
tokenize_with_builtin(sql, D)
.expect("tokenizes")
.into_iter()
.filter(|t| {
matches!(
t.kind,
TokenKind::Word | TokenKind::Number | TokenKind::String | TokenKind::QuotedIdent
)
})
.map(|t| sql[t.span.start() as usize..t.span.end() as usize].to_owned())
.collect()
}
#[test]
fn format_preserves_spelling_tokens() {
for &sql in CORPUS {
let formatted = fmt(sql);
assert_eq!(
spelling_tokens(sql),
spelling_tokens(&formatted),
"spelling drift for {sql:?}\n{formatted}"
);
}
}
#[test]
fn simple_select_lays_out_one_clause_per_line() {
let out = fmt("select a, b from t where a = 1 group by a");
assert_eq!(out, "SELECT a, b\nFROM t\nWHERE a = 1\nGROUP BY a");
}
#[test]
fn wide_projection_breaks_one_item_per_line() {
let sql = "SELECT alpha, bravo, charlie, delta, echo, foxtrot, golf, hotel, india FROM t";
let out = format_sql(sql, D, &FormatOptions::default().with_max_line_length(40)).unwrap();
assert!(out.contains("SELECT\n alpha,\n bravo,"), "got:\n{out}");
}
#[test]
fn indent_width_is_honoured() {
let sql = "SELECT alpha, bravo, charlie, delta, echo, foxtrot, golf, hotel FROM t";
let opts = FormatOptions::default()
.with_max_line_length(30)
.with_indent_width(4);
let out = format_sql(sql, D, &opts).unwrap();
assert!(out.contains("SELECT\n alpha,"), "got:\n{out}");
}
#[test]
fn keyword_case_lower_recases_all_keywords() {
let opts = FormatOptions::default().with_keyword_case(KeywordCase::Lower);
let out = format_sql("SELECT a FROM t WHERE a IS NOT NULL", D, &opts).unwrap();
assert_eq!(out, "select a\nfrom t\nwhere a is not null");
}
#[test]
fn keyword_case_preserve_follows_dominant_source_case() {
let opts = FormatOptions::default().with_keyword_case(KeywordCase::Preserve);
let out = format_sql("select a from t where a = 1", D, &opts).unwrap();
assert!(out.starts_with("select a"), "got:\n{out}");
assert!(out.contains("from t"));
}
#[test]
fn comments_survive_the_format() {
let cases = [
"-- leading\nSELECT 1",
"SELECT a -- trailing\nFROM t",
"SELECT a FROM t WHERE a = 1\n-- before group\nGROUP BY a",
"SELECT count(/* inside */) FROM t",
"SELECT /* block */ a FROM t",
];
for sql in cases {
let parsed =
parse_with_builtin_options(sql, D, ParseOptions::default().with_trivia_capture(true))
.expect("parses");
let comment_count = parsed
.trivia()
.iter()
.filter(|t| {
matches!(
t.kind(),
crate::tokenizer::TriviaKind::LineComment
| crate::tokenizer::TriviaKind::BlockComment
)
})
.count();
let out = format_parsed(&parsed, D, &FormatOptions::default());
let out_comments = tokenize_out_comment_count(&out);
assert_eq!(
comment_count, out_comments,
"comment dropped for {sql:?}\n--- out ---\n{out}"
);
assert!(
parse_with_builtin_options(&out, D, ParseOptions::default()).is_ok(),
"output with comments failed to reparse for {sql:?}:\n{out}"
);
}
}
fn tokenize_out_comment_count(sql: &str) -> usize {
let (_tokens, trivia) = crate::tokenize_with_builtin_trivia(sql, D).expect("output tokenizes");
trivia
.all()
.iter()
.filter(|t| {
matches!(
t.kind(),
crate::tokenizer::TriviaKind::LineComment
| crate::tokenizer::TriviaKind::BlockComment
)
})
.count()
}
#[test]
fn comment_before_group_by_renders_before_the_keyword() {
let sql = "SELECT a FROM t WHERE a = 1\n-- note\nGROUP BY a";
let parsed =
parse_with_builtin_options(sql, D, ParseOptions::default().with_trivia_capture(true))
.expect("parses");
let out = format_parsed(&parsed, D, &FormatOptions::default());
let note = out.find("-- note").expect("comment present");
let group = out.find("GROUP BY").expect("group by present");
assert!(note < group, "comment should precede GROUP BY:\n{out}");
}
#[test]
fn empty_input_formats_to_empty() {
assert_eq!(fmt(""), "");
}
const COMMENT_CORPUS: &[&str] = &[
"-- header\nSELECT a FROM t WHERE a = 1",
"SELECT 1;\n-- divider\nSELECT 2",
"SELECT a FROM t WHERE a = 1 -- filter\n",
"SELECT a FROM t\n-- trailing note",
"SELECT 1;\n-- after",
"/* head */ SELECT a FROM t",
"SELECT a + /* mid */ b FROM t",
"SELECT count(/* why */) FROM t",
"SELECT x FROM (SELECT a /* inner */ FROM u) s",
"INSERT INTO t (a) VALUES (1) /* tail */",
"SELECT a FROM t WHERE b = /* mid */ 2",
"SELECT a FROM t WHERE b = 2 AND c /* z */ IN (1, 2)",
"SELECT a /* on a */\n, b FROM t",
"SELECT a\n, /* on b */ b FROM t",
"SELECT a, -- keep with a\nb FROM t",
"SELECT a, b FROM t GROUP BY a, -- keep with a\nb",
"SELECT x FROM (SELECT a /* pick */ FROM u WHERE a > 0) s",
"SELECT x FROM (SELECT a, -- keep with a\nb FROM u WHERE a > 0) s",
"SELECT a FROM t WHERE a IN (SELECT id /* only ids */ FROM u WHERE u.active)",
"SELECT a FROM t WHERE EXISTS (SELECT 1 FROM u -- probe\nWHERE u.k = t.k)",
"SELECT a FROM t WHERE a > (SELECT 1 /* floor */)",
];
#[test]
fn comment_formatting_is_byte_stable_idempotent() {
for &sql in COMMENT_CORPUS {
let once = fmt(sql);
let twice = fmt(&once);
assert_eq!(
once, twice,
"format is not idempotent for {sql:?}\n--- once ---\n{once}\n--- twice ---\n{twice}\n---"
);
}
}
#[test]
fn corpus_formatting_is_byte_stable_idempotent() {
for &sql in CORPUS {
let once = fmt(sql);
let twice = fmt(&once);
assert_eq!(
once, twice,
"format is not idempotent for {sql:?}\n--- once ---\n{once}\n--- twice ---\n{twice}\n---"
);
}
}
#[test]
fn comment_corpus_never_drops_and_reparses() {
for &sql in COMMENT_CORPUS {
let parsed =
parse_with_builtin_options(sql, D, ParseOptions::default().with_trivia_capture(true))
.expect("parses");
let comment_count = parsed
.trivia()
.iter()
.filter(|t| {
matches!(
t.kind(),
crate::tokenizer::TriviaKind::LineComment
| crate::tokenizer::TriviaKind::BlockComment
)
})
.count();
let out = format_parsed(&parsed, D, &FormatOptions::default());
assert_eq!(
comment_count,
tokenize_out_comment_count(&out),
"comment dropped for {sql:?}\n--- out ---\n{out}"
);
assert!(
parse_with_builtin_options(&out, D, ParseOptions::default()).is_ok(),
"output failed to reparse for {sql:?}:\n{out}"
);
}
}
#[test]
fn statement_boundary_comments_hold_position() {
assert!(fmt("-- header\nSELECT a FROM t").starts_with("-- header\nSELECT"));
let between = fmt("SELECT 1;\n-- divider\nSELECT 2");
assert!(
between.contains("-- divider\nSELECT 2"),
"divider must precede the second statement:\n{between}"
);
assert!(
fmt("SELECT a FROM t -- tail\n").contains("FROM t -- tail"),
"trailing comment stays with the statement"
);
}
#[test]
fn fragment_interior_comment_hoists_adjacent() {
assert!(fmt("SELECT a + /* c */ b FROM t").contains("SELECT a + b /* c */"));
assert!(fmt("SELECT count(/* c */) FROM t").contains("count() /* c */"));
assert!(fmt("SELECT a FROM t WHERE b = /* c */ 2").contains("WHERE b = 2 /* c */"));
}
#[test]
fn item_trailing_comment_renders_before_the_comma() {
assert_eq!(
fmt("SELECT a /* on a */\n, b FROM t"),
"SELECT a /* on a */, b\nFROM t"
);
}