use std::env;
use std::fs;
use std::path::{Path, PathBuf};
use sql_dialect_fmt_formatter::{format, FormatOptions};
use sql_dialect_fmt_lexer::tokenize_for_dialect;
use sql_dialect_fmt_parser::{parse_with_dialect, Dialect};
use sql_dialect_fmt_syntax::SyntaxKind;
const EXTERNAL_CORPUS_ENV: &str = "SQL_DIALECT_FMT_EXTERNAL_CORPUS";
const EXTERNAL_CORPUS_LIMIT_ENV: &str = "SQL_DIALECT_FMT_EXTERNAL_CORPUS_LIMIT";
const EXTERNAL_CORPUS_DIALECT_ENV: &str = "SQL_DIALECT_FMT_EXTERNAL_CORPUS_DIALECT";
#[derive(Debug)]
struct CorpusFailure {
file: PathBuf,
reason: String,
}
#[derive(Clone, Copy, Debug, Default)]
struct FileOutcome {
lex_clean: bool,
parse_clean: bool,
changed: bool,
}
impl std::fmt::Display for CorpusFailure {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, " {}: {}", self.file.display(), self.reason)
}
}
fn check_file(
file: &Path,
source: &str,
options: &FormatOptions,
) -> Result<FileOutcome, CorpusFailure> {
let fail = |reason: String| CorpusFailure {
file: file.to_path_buf(),
reason,
};
let dialect = options.dialect;
let _ = parse_with_dialect(source, dialect);
let formatted = format(source, options);
let reformatted = format(&formatted, options);
if formatted != reformatted {
return Err(fail(format!(
"not idempotent: format(format(x)) != format(x); first diff at byte {}",
first_diff(&formatted, &reformatted)
)));
}
let lex_clean = tokenize_for_dialect(source, dialect).errors.is_empty();
let parse_clean = parse_with_dialect(source, dialect).errors().is_empty();
if lex_clean && parse_clean {
let reparse = parse_with_dialect(&formatted, dialect);
if !reparse.errors().is_empty() {
return Err(fail(format!(
"formatted output does not reparse cleanly: {:?}",
reparse.errors()
)));
}
let before = significant_tokens(source, dialect);
let after = significant_tokens(&formatted, dialect);
if before != after {
return Err(fail(format!(
"significant tokens changed across formatting ({} -> {} tokens)",
before.len(),
after.len()
)));
}
}
Ok(FileOutcome {
lex_clean,
parse_clean,
changed: formatted != source,
})
}
fn run_corpus(files: &[PathBuf], label: &str, options: &FormatOptions) {
assert!(!files.is_empty(), "{label}: no .sql files found");
let mut failures = Vec::new();
let mut utf8_files = 0usize;
let mut lex_clean = 0usize;
let mut parse_clean = 0usize;
let mut changed = 0usize;
for file in files {
let Ok(bytes) = fs::read(file) else {
failures.push(CorpusFailure {
file: file.clone(),
reason: "could not read file".to_string(),
});
continue;
};
let Ok(source) = String::from_utf8(bytes) else {
continue;
};
utf8_files += 1;
match check_file(file, &source, options) {
Ok(outcome) => {
lex_clean += usize::from(outcome.lex_clean);
parse_clean += usize::from(outcome.parse_clean);
changed += usize::from(outcome.changed);
}
Err(failure) => failures.push(failure),
}
}
assert!(
failures.is_empty(),
"{label}: {} file(s) violated a formatter invariant:\n{}",
failures.len(),
failures
.iter()
.map(ToString::to_string)
.collect::<Vec<_>>()
.join("\n")
);
println!(
"corpus summary: dialect={} files={} utf8={} lex_clean={} parse_clean={} changed={}",
dialect_name(options.dialect),
files.len(),
utf8_files,
lex_clean,
parse_clean,
changed
);
}
fn sample_corpus_dir() -> PathBuf {
Path::new(env!("CARGO_MANIFEST_DIR"))
.join("tests")
.join("corpus_sample")
}
#[test]
fn sample_corpus_is_clean() {
let mut files = Vec::new();
collect_sql_files(&sample_corpus_dir(), &mut files);
files.sort();
assert!(
files.len() >= 6,
"expected the curated sample corpus to cover every statement family (>= 6 files), found {}",
files.len()
);
run_corpus(&files, "sample_corpus", &FormatOptions::default());
let options = FormatOptions::default();
let mut drifted = Vec::new();
for file in &files {
let source = fs::read_to_string(file).expect("read sample file");
let formatted = format(&source, &options);
if formatted != source {
drifted.push(CorpusFailure {
file: file.clone(),
reason: format!(
"sample is not in canonical form; first diff at byte {}",
first_diff(&source, &formatted)
),
});
}
}
assert!(
drifted.is_empty(),
"sample_corpus: {} sample(s) drifted from canonical form (re-run the formatter and commit \
the output):\n{}",
drifted.len(),
drifted
.iter()
.map(ToString::to_string)
.collect::<Vec<_>>()
.join("\n")
);
}
#[test]
#[ignore = "set SQL_DIALECT_FMT_EXTERNAL_CORPUS to one or more SQL files/directories"]
fn external_corpus_preserves_formatter_invariants() {
let roots = external_corpus_roots();
let mut files = Vec::new();
for root in env::split_paths(&roots) {
let root = resolve_corpus_root(root);
collect_sql_files(&root, &mut files);
}
files.sort();
files.dedup();
let limit = external_corpus_limit();
files.truncate(limit);
let options = FormatOptions::default().with_dialect(external_corpus_dialect());
run_corpus(&files, "external_corpus", &options);
}
fn external_corpus_roots() -> std::ffi::OsString {
env::var_os(EXTERNAL_CORPUS_ENV)
.unwrap_or_else(|| panic!("{EXTERNAL_CORPUS_ENV} must point at SQL files/directories"))
}
fn external_corpus_limit() -> usize {
env::var(EXTERNAL_CORPUS_LIMIT_ENV)
.ok()
.and_then(|value| value.parse::<usize>().ok())
.unwrap_or(usize::MAX)
}
fn external_corpus_dialect() -> Dialect {
match env::var(EXTERNAL_CORPUS_DIALECT_ENV)
.unwrap_or_else(|_| "snowflake".to_string())
.to_ascii_lowercase()
.as_str()
{
"snowflake" => Dialect::Snowflake,
"databricks" => Dialect::Databricks,
value => panic!(
"{EXTERNAL_CORPUS_DIALECT_ENV} must be `snowflake` or `databricks`, got {value:?}"
),
}
}
fn collect_sql_files(path: &Path, out: &mut Vec<PathBuf>) {
if path.is_file() {
if is_sql_file(path) {
out.push(path.to_path_buf());
}
return;
}
if !path.is_dir() {
return;
}
let Ok(entries) = fs::read_dir(path) else {
return;
};
for entry in entries.flatten() {
collect_sql_files(&entry.path(), out);
}
}
fn resolve_corpus_root(root: PathBuf) -> PathBuf {
if root.is_absolute() || root.exists() {
return root;
}
let manifest_dir = Path::new(env!("CARGO_MANIFEST_DIR"));
let workspace_root = manifest_dir
.parent()
.and_then(Path::parent)
.unwrap_or(manifest_dir);
let workspace_relative = workspace_root.join(&root);
if workspace_relative.exists() {
workspace_relative
} else {
root
}
}
fn is_sql_file(path: &Path) -> bool {
path.extension()
.and_then(|ext| ext.to_str())
.is_some_and(|ext| ext.eq_ignore_ascii_case("sql"))
}
fn significant_tokens(sql: &str, dialect: Dialect) -> Vec<String> {
tokenize_for_dialect(sql, dialect)
.tokens
.into_iter()
.filter(|token| !token.kind.is_trivia() && token.kind != SyntaxKind::SEMICOLON)
.map(|token| {
if token.kind == SyntaxKind::DOLLAR_STRING {
"<DOLLAR_STRING>".to_string()
} else {
token.text.to_ascii_uppercase()
}
})
.collect()
}
fn dialect_name(dialect: Dialect) -> &'static str {
match dialect {
Dialect::Snowflake => "snowflake",
Dialect::Databricks => "databricks",
_ => "unknown",
}
}
#[test]
fn significant_token_check_treats_formatted_dollar_bodies_as_opaque() {
let before = "EXECUTE IMMEDIATE $$begin select 1; end$$";
let after = "EXECUTE IMMEDIATE $$\nBEGIN\n SELECT 1;\nEND;\n$$";
assert_eq!(
significant_tokens(before, Dialect::Snowflake),
significant_tokens(after, Dialect::Snowflake)
);
}
fn first_diff(a: &str, b: &str) -> usize {
a.bytes()
.zip(b.bytes())
.position(|(x, y)| x != y)
.unwrap_or_else(|| a.len().min(b.len()))
}