use crate::dialect::BuiltinDialect;
use crate::tokenize_with_builtin;
use crate::tokenizer::TokenKind;
#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
pub enum KeywordCase {
#[default]
Upper,
Lower,
Preserve,
}
impl KeywordCase {
pub fn from_name(name: &str) -> Option<KeywordCase> {
if name.eq_ignore_ascii_case("upper") {
Some(KeywordCase::Upper)
} else if name.eq_ignore_ascii_case("lower") {
Some(KeywordCase::Lower)
} else if name.eq_ignore_ascii_case("preserve") {
Some(KeywordCase::Preserve)
} else {
None
}
}
}
fn lowercase_keywords(case: KeywordCase, source: &str, dialect: BuiltinDialect) -> bool {
match case {
KeywordCase::Upper => false,
KeywordCase::Lower => true,
KeywordCase::Preserve => dominant_is_lower(source, dialect),
}
}
fn dominant_is_lower(source: &str, dialect: BuiltinDialect) -> bool {
let Ok(tokens) = tokenize_with_builtin(source, dialect) else {
return false;
};
let mut lower = 0usize;
let mut upper = 0usize;
for token in &tokens {
if !matches!(token.kind, TokenKind::Keyword(_)) {
continue;
}
let Some(text) = slice(source, token.span.start(), token.span.end()) else {
continue;
};
let has_alpha = text.chars().any(|c| c.is_ascii_alphabetic());
if !has_alpha {
continue;
}
if text.chars().all(|c| !c.is_ascii_uppercase()) {
lower += 1;
} else if text.chars().all(|c| !c.is_ascii_lowercase()) {
upper += 1;
}
}
lower > upper
}
pub fn apply(rendered: String, source: &str, dialect: BuiltinDialect, case: KeywordCase) -> String {
if case == KeywordCase::Upper {
return rendered;
}
let to_lower = lowercase_keywords(case, source, dialect);
let Ok(tokens) = tokenize_with_builtin(&rendered, dialect) else {
return rendered;
};
let mut out = String::with_capacity(rendered.len());
let mut cursor = 0usize;
for token in &tokens {
if !matches!(token.kind, TokenKind::Keyword(_)) {
continue;
}
let start = token.span.start() as usize;
let end = token.span.end() as usize;
if start < cursor || end > rendered.len() {
continue;
}
out.push_str(&rendered[cursor..start]);
let keyword = &rendered[start..end];
if to_lower {
out.push_str(&keyword.to_ascii_lowercase());
} else {
out.push_str(&keyword.to_ascii_uppercase());
}
cursor = end;
}
out.push_str(&rendered[cursor..]);
out
}
fn slice(source: &str, start: u32, end: u32) -> Option<&str> {
source.get(start as usize..end as usize)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn upper_is_identity() {
let rendered = "SELECT a FROM t".to_owned();
assert_eq!(
apply(
rendered.clone(),
"select a from t",
BuiltinDialect::Ansi,
KeywordCase::Upper
),
rendered
);
}
#[test]
fn lower_recases_only_keywords() {
assert_eq!(
apply(
"SELECT Col FROM Tbl WHERE Col IS NOT NULL".to_owned(),
"",
BuiltinDialect::Ansi,
KeywordCase::Lower
),
"select Col from Tbl where Col is not null"
);
}
#[test]
fn lower_leaves_quoted_identifiers_and_strings_untouched() {
assert_eq!(
apply(
"SELECT \"FROM\" FROM t WHERE x = 'SELECT'".to_owned(),
"",
BuiltinDialect::Ansi,
KeywordCase::Lower
),
"select \"FROM\" from t where x = 'SELECT'"
);
}
#[test]
fn preserve_follows_dominant_lowercase_source() {
assert_eq!(
apply(
"SELECT Col FROM Tbl".to_owned(),
"select col from tbl where col = 1",
BuiltinDialect::Ansi,
KeywordCase::Preserve
),
"select Col from Tbl"
);
}
#[test]
fn preserve_follows_dominant_uppercase_source() {
assert_eq!(
apply(
"SELECT Col FROM Tbl".to_owned(),
"SELECT col FROM tbl WHERE col = 1",
BuiltinDialect::Ansi,
KeywordCase::Preserve
),
"SELECT Col FROM Tbl"
);
}
}