use crate::{Dialect, SyntaxKind};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum KeywordDialect {
Shared,
SnowflakeOnly,
#[allow(dead_code)]
DatabricksOnly,
}
impl KeywordDialect {
#[inline]
#[must_use]
pub fn reserved_in(self, dialect: Dialect) -> bool {
match self {
KeywordDialect::Shared => true,
KeywordDialect::SnowflakeOnly => matches!(dialect, Dialect::Snowflake),
KeywordDialect::DatabricksOnly => matches!(dialect, Dialect::Databricks),
}
}
}
const KEYWORDS: &[(&str, SyntaxKind, KeywordDialect)] = {
use KeywordDialect::{Shared, SnowflakeOnly};
use SyntaxKind::*;
&[
("after", AFTER_KW, Shared),
("all", ALL_KW, Shared),
("alter", ALTER_KW, Shared),
("and", AND_KW, Shared),
("any", ANY_KW, Shared),
("as", AS_KW, Shared),
("asc", ASC_KW, Shared),
("begin", BEGIN_KW, Shared),
("between", BETWEEN_KW, Shared),
("by", BY_KW, Shared),
("call", CALL_KW, Shared),
("called", CALLED_KW, SnowflakeOnly),
("caller", CALLER_KW, SnowflakeOnly),
("case", CASE_KW, Shared),
("cast", CAST_KW, Shared),
("commit", COMMIT_KW, Shared),
("connect", CONNECT_KW, SnowflakeOnly),
("copy", COPY_KW, SnowflakeOnly),
("create", CREATE_KW, Shared),
("cross", CROSS_KW, Shared),
("current", CURRENT_KW, Shared),
("cursor", CURSOR_KW, SnowflakeOnly),
("declare", DECLARE_KW, Shared),
("delete", DELETE_KW, Shared),
("desc", DESC_KW, Shared),
("describe", DESCRIBE_KW, Shared),
("distinct", DISTINCT_KW, Shared),
("do", DO_KW, Shared),
("drop", DROP_KW, Shared),
("else", ELSE_KW, Shared),
("elseif", ELSEIF_KW, SnowflakeOnly),
("end", END_KW, Shared),
("except", EXCEPT_KW, Shared),
("exception", EXCEPTION_KW, SnowflakeOnly),
("execute", EXECUTE_KW, Shared),
("exists", EXISTS_KW, Shared),
("false", FALSE_KW, Shared),
("fetch", FETCH_KW, Shared),
("first", FIRST_KW, Shared),
("flatten", FLATTEN_KW, SnowflakeOnly),
("following", FOLLOWING_KW, Shared),
("for", FOR_KW, Shared),
("from", FROM_KW, Shared),
("full", FULL_KW, Shared),
("function", FUNCTION_KW, Shared),
("grant", GRANT_KW, Shared),
("grants", GRANTS_KW, SnowflakeOnly),
("group", GROUP_KW, Shared),
("handler", HANDLER_KW, SnowflakeOnly),
("having", HAVING_KW, Shared),
("if", IF_KW, Shared),
("ilike", ILIKE_KW, SnowflakeOnly),
("immediate", IMMEDIATE_KW, SnowflakeOnly),
("imports", IMPORTS_KW, SnowflakeOnly),
("in", IN_KW, Shared),
("inner", INNER_KW, Shared),
("input", INPUT_KW, Shared),
("insert", INSERT_KW, Shared),
("intersect", INTERSECT_KW, Shared),
("into", INTO_KW, Shared),
("is", IS_KW, Shared),
("java", JAVA_KW, Shared),
("javascript", JAVASCRIPT_KW, SnowflakeOnly),
("join", JOIN_KW, Shared),
("language", LANGUAGE_KW, Shared),
("last", LAST_KW, Shared),
("lateral", LATERAL_KW, Shared),
("left", LEFT_KW, Shared),
("let", LET_KW, Shared),
("like", LIKE_KW, Shared),
("limit", LIMIT_KW, Shared),
("loop", LOOP_KW, Shared),
("matched", MATCHED_KW, Shared),
("merge", MERGE_KW, Shared),
("minus", MINUS_KW, Shared),
("natural", NATURAL_KW, Shared),
("not", NOT_KW, Shared),
("null", NULL_KW, Shared),
("nulls", NULLS_KW, Shared),
("offset", OFFSET_KW, Shared),
("on", ON_KW, Shared),
("or", OR_KW, Shared),
("order", ORDER_KW, Shared),
("out", OUT_KW, Shared),
("outer", OUTER_KW, Shared),
("output", OUTPUT_KW, Shared),
("over", OVER_KW, Shared),
("overwrite", OVERWRITE_KW, Shared),
("owner", OWNER_KW, SnowflakeOnly),
("packages", PACKAGES_KW, SnowflakeOnly),
("partition", PARTITION_KW, Shared),
("pivot", PIVOT_KW, Shared),
("preceding", PRECEDING_KW, Shared),
("prior", PRIOR_KW, SnowflakeOnly),
("procedure", PROCEDURE_KW, Shared),
("python", PYTHON_KW, Shared),
("qualify", QUALIFY_KW, Shared),
("range", RANGE_KW, Shared),
("recursive", RECURSIVE_KW, Shared),
("regexp", REGEXP_KW, SnowflakeOnly),
("repeat", REPEAT_KW, Shared),
("replace", REPLACE_KW, Shared),
("resultset", RESULTSET_KW, SnowflakeOnly),
("return", RETURN_KW, Shared),
("returns", RETURNS_KW, Shared),
("revoke", REVOKE_KW, Shared),
("right", RIGHT_KW, Shared),
("rlike", RLIKE_KW, SnowflakeOnly),
("rollback", ROLLBACK_KW, Shared),
("row", ROW_KW, Shared),
("rows", ROWS_KW, Shared),
("runtime_version", RUNTIME_VERSION_KW, SnowflakeOnly),
("sample", SAMPLE_KW, SnowflakeOnly),
("scala", SCALA_KW, SnowflakeOnly),
("schedule", SCHEDULE_KW, SnowflakeOnly),
("secure", SECURE_KW, SnowflakeOnly),
("select", SELECT_KW, Shared),
("set", SET_KW, Shared),
("show", SHOW_KW, Shared),
("sql", SQL_KW, Shared),
("start", START_KW, Shared),
("strict", STRICT_KW, SnowflakeOnly),
("table", TABLE_KW, Shared),
("tablesample", TABLESAMPLE_KW, Shared),
("task", TASK_KW, SnowflakeOnly),
("temp", TEMP_KW, Shared),
("temporary", TEMPORARY_KW, Shared),
("then", THEN_KW, Shared),
("top", TOP_KW, SnowflakeOnly),
("transient", TRANSIENT_KW, SnowflakeOnly),
("true", TRUE_KW, Shared),
("truncate", TRUNCATE_KW, Shared),
("try_cast", TRY_CAST_KW, SnowflakeOnly),
("unbounded", UNBOUNDED_KW, Shared),
("undrop", UNDROP_KW, SnowflakeOnly),
("union", UNION_KW, Shared),
("unpivot", UNPIVOT_KW, Shared),
("until", UNTIL_KW, Shared),
("update", UPDATE_KW, Shared),
("use", USE_KW, Shared),
("using", USING_KW, Shared),
("values", VALUES_KW, Shared),
("view", VIEW_KW, Shared),
("volatile", VOLATILE_KW, SnowflakeOnly),
("warehouse", WAREHOUSE_KW, SnowflakeOnly),
("when", WHEN_KW, Shared),
("where", WHERE_KW, Shared),
("while", WHILE_KW, Shared),
("window", WINDOW_KW, Shared),
("with", WITH_KW, Shared),
("within", WITHIN_KW, Shared),
]
};
const MAX_KEYWORD_LEN: usize = 16;
#[inline]
fn lower_for_lookup(ident: &str, buf: &mut [u8; MAX_KEYWORD_LEN]) -> Option<usize> {
let bytes = ident.as_bytes();
if bytes.is_empty() || bytes.len() > MAX_KEYWORD_LEN {
return None;
}
for (slot, &b) in buf.iter_mut().zip(bytes) {
*slot = b.to_ascii_lowercase();
}
Some(bytes.len())
}
#[inline]
fn lookup(ident: &str) -> Option<(SyntaxKind, KeywordDialect)> {
let mut buf = [0u8; MAX_KEYWORD_LEN];
let len = lower_for_lookup(ident, &mut buf)?;
let lower = std::str::from_utf8(&buf[..len]).ok()?;
KEYWORDS
.binary_search_by(|(text, _, _)| text.cmp(&lower))
.ok()
.map(|index| {
let (_, kind, dialect) = KEYWORDS[index];
(kind, dialect)
})
}
#[must_use]
pub fn keyword_kind(ident: &str) -> Option<SyntaxKind> {
keyword_kind_for(ident, Dialect::Snowflake)
}
#[must_use]
pub fn keyword_kind_for(ident: &str, dialect: Dialect) -> Option<SyntaxKind> {
let (kind, kw_dialect) = lookup(ident)?;
kw_dialect.reserved_in(dialect).then_some(kind)
}
pub fn keyword_texts() -> impl ExactSizeIterator<Item = &'static str> {
KEYWORDS.iter().map(|(text, _, _)| *text)
}
#[cfg(test)]
mod tests {
use super::{keyword_kind, keyword_kind_for, keyword_texts, KeywordDialect, KEYWORDS};
use crate::{Dialect, SyntaxKind};
#[test]
fn keyword_lookup_is_case_insensitive() {
assert_eq!(keyword_kind("select"), Some(SyntaxKind::SELECT_KW));
assert_eq!(keyword_kind("SeLeCt"), Some(SyntaxKind::SELECT_KW));
assert_eq!(keyword_kind("QUALIFY"), Some(SyntaxKind::QUALIFY_KW));
assert_eq!(keyword_kind("javascript"), Some(SyntaxKind::JAVASCRIPT_KW));
assert_eq!(keyword_kind("try_cast"), Some(SyntaxKind::TRY_CAST_KW));
assert_eq!(keyword_kind("TASK"), Some(SyntaxKind::TASK_KW));
assert_eq!(
keyword_kind("runtime_version"),
Some(SyntaxKind::RUNTIME_VERSION_KW)
);
assert_eq!(keyword_kind("definitely_not_a_keyword"), None);
}
#[test]
fn every_keyword_variant_is_mapped() {
let range_count = SyntaxKind::__KW_END as u16 - SyntaxKind::__KW_START as u16 - 1;
assert_eq!(
KEYWORDS.len() as u16,
range_count,
"KEYWORDS table is out of sync with the SyntaxKind keyword block"
);
let mut seen = std::collections::HashSet::new();
for (text, kind, _) in KEYWORDS {
assert_eq!(
keyword_kind(text),
Some(*kind),
"keyword_kind({text:?}) is wrong"
);
assert_eq!(
keyword_kind(&text.to_uppercase()),
Some(*kind),
"keyword_kind is not case-insensitive for {text:?}"
);
assert!(
kind.is_keyword(),
"{kind:?} should be inside the keyword range"
);
assert!(seen.insert(*kind), "duplicate keyword kind for {text:?}");
assert!(
text.bytes().all(|b| !b.is_ascii_uppercase()),
"KEYWORDS text must be lowercase: {text:?}"
);
}
}
#[test]
fn keywords_are_sorted_for_binary_search() {
for window in KEYWORDS.windows(2) {
let (left, _, _) = window[0];
let (right, _, _) = window[1];
assert!(
left < right,
"KEYWORDS must be sorted: {left:?} >= {right:?}"
);
}
}
#[test]
fn keyword_texts_exposes_the_lookup_table_order() {
let texts: Vec<_> = keyword_texts().collect();
assert_eq!(texts.len(), KEYWORDS.len());
assert_eq!(texts.first(), Some(&"after"));
assert_eq!(texts.last(), Some(&"within"));
for (text, (table_text, _, _)) in texts.iter().zip(KEYWORDS) {
assert_eq!(text, table_text);
}
}
#[test]
fn every_keyword_has_a_dialect_classification() {
for (text, kind, dialect) in KEYWORDS {
assert!(
matches!(
dialect,
KeywordDialect::Shared
| KeywordDialect::SnowflakeOnly
| KeywordDialect::DatabricksOnly
),
"{text:?} ({kind:?}) has no dialect classification"
);
}
}
#[test]
fn snowflake_classification_is_byte_identical_to_legacy_keyword_kind() {
for (text, kind, _) in KEYWORDS {
assert_eq!(
keyword_kind_for(text, Dialect::Snowflake),
Some(*kind),
"Snowflake reservation changed for {text:?}"
);
assert_eq!(
keyword_kind(text),
keyword_kind_for(text, Dialect::Snowflake)
);
}
}
#[test]
fn shared_keywords_are_reserved_in_every_dialect() {
for word in ["select", "from", "where", "join", "group", "order", "case"] {
assert!(
keyword_kind_for(word, Dialect::Databricks).is_some(),
"{word}"
);
assert!(
keyword_kind_for(word, Dialect::Snowflake).is_some(),
"{word}"
);
}
}
#[test]
fn snowflake_only_keywords_are_identifiers_under_databricks() {
for word in [
"task",
"flatten",
"warehouse",
"schedule",
"transient",
"volatile",
"secure",
"undrop",
"elseif",
"cursor",
"resultset",
"connect",
"prior",
"top",
"copy",
"owner",
] {
assert!(
keyword_kind_for(word, Dialect::Snowflake).is_some(),
"{word} should be reserved in Snowflake"
);
assert_eq!(
keyword_kind_for(word, Dialect::Databricks),
None,
"{word} must be a plain identifier under Databricks"
);
}
}
}