use crate::{highlight, HighlightKind, HighlightToken};
use sql_dialect_fmt_text::{utf16_len, LineIndex};
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum SemanticTokenType {
Keyword,
Type,
Variable,
String,
Number,
Parameter,
Operator,
Comment,
Namespace,
Function,
}
impl SemanticTokenType {
pub const LEGEND: &'static [SemanticTokenType] = &[
SemanticTokenType::Keyword,
SemanticTokenType::Type,
SemanticTokenType::Variable,
SemanticTokenType::String,
SemanticTokenType::Number,
SemanticTokenType::Parameter,
SemanticTokenType::Operator,
SemanticTokenType::Comment,
SemanticTokenType::Namespace,
SemanticTokenType::Function,
];
pub const fn name(self) -> &'static str {
match self {
SemanticTokenType::Keyword => "keyword",
SemanticTokenType::Type => "type",
SemanticTokenType::Variable => "variable",
SemanticTokenType::String => "string",
SemanticTokenType::Number => "number",
SemanticTokenType::Parameter => "parameter",
SemanticTokenType::Operator => "operator",
SemanticTokenType::Comment => "comment",
SemanticTokenType::Namespace => "namespace",
SemanticTokenType::Function => "function",
}
}
pub const fn index(self) -> u32 {
match self {
SemanticTokenType::Keyword => 0,
SemanticTokenType::Type => 1,
SemanticTokenType::Variable => 2,
SemanticTokenType::String => 3,
SemanticTokenType::Number => 4,
SemanticTokenType::Parameter => 5,
SemanticTokenType::Operator => 6,
SemanticTokenType::Comment => 7,
SemanticTokenType::Namespace => 8,
SemanticTokenType::Function => 9,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Default)]
pub struct SemanticTokenModifiers(u32);
impl SemanticTokenModifiers {
pub const NONE: SemanticTokenModifiers = SemanticTokenModifiers(0);
pub const DOCUMENTATION: SemanticTokenModifiers = SemanticTokenModifiers(1 << 0);
pub const DEFAULT_LIBRARY: SemanticTokenModifiers = SemanticTokenModifiers(1 << 1);
pub const LEGEND: &'static [&'static str] = &["documentation", "defaultLibrary"];
pub const fn bits(self) -> u32 {
self.0
}
pub const fn contains(self, other: SemanticTokenModifiers) -> bool {
self.0 & other.0 == other.0
}
}
impl std::ops::BitOr for SemanticTokenModifiers {
type Output = SemanticTokenModifiers;
fn bitor(self, rhs: SemanticTokenModifiers) -> SemanticTokenModifiers {
SemanticTokenModifiers(self.0 | rhs.0)
}
}
pub fn semantic_token(kind: HighlightKind) -> Option<(SemanticTokenType, SemanticTokenModifiers)> {
use HighlightKind::*;
use SemanticTokenModifiers as M;
Some(match kind {
Keyword => (SemanticTokenType::Keyword, M::DEFAULT_LIBRARY),
Type => (SemanticTokenType::Type, M::DEFAULT_LIBRARY),
Identifier | QuotedIdentifier => (SemanticTokenType::Variable, M::NONE),
String | DollarString => (SemanticTokenType::String, M::NONE),
Number => (SemanticTokenType::Number, M::NONE),
Variable => (SemanticTokenType::Parameter, M::NONE),
Operator => (SemanticTokenType::Operator, M::NONE),
Comment => (SemanticTokenType::Comment, M::DOCUMENTATION),
Whitespace | Punctuation | Error => return None,
})
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct Injection {
pub language: InjectedLanguage,
pub range: std::ops::Range<usize>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum InjectedLanguage {
Sql,
JavaScript,
Python,
Java,
Scala,
}
impl InjectedLanguage {
pub const fn scope(self) -> &'static str {
match self {
InjectedLanguage::Sql => "source.snowflake-sql",
InjectedLanguage::JavaScript => "source.js",
InjectedLanguage::Python => "source.python",
InjectedLanguage::Java => "source.java",
InjectedLanguage::Scala => "source.scala",
}
}
fn from_language_word(word: &str) -> InjectedLanguage {
if word.eq_ignore_ascii_case("javascript") {
InjectedLanguage::JavaScript
} else if word.eq_ignore_ascii_case("python") {
InjectedLanguage::Python
} else if word.eq_ignore_ascii_case("java") {
InjectedLanguage::Java
} else if word.eq_ignore_ascii_case("scala") {
InjectedLanguage::Scala
} else {
InjectedLanguage::Sql
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct ResolvedToken {
pub range: std::ops::Range<usize>,
pub token_type: SemanticTokenType,
pub modifiers: SemanticTokenModifiers,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct LineToken {
pub line: u32,
pub start_char: u32,
pub length: u32,
pub token_type: u32,
pub modifiers: u32,
}
pub fn detect_injections(input: &str) -> Vec<Injection> {
let highlighted = highlight(input);
let mut injections = Vec::new();
let mut language: Option<InjectedLanguage> = None;
let mut expect_language_name = false;
let mut saw_as_after_language = false;
let mut saw_as = false;
let mut saw_execute = false;
let mut saw_execute_immediate = false;
for token in &highlighted.tokens {
match token.kind {
HighlightKind::Whitespace | HighlightKind::Comment => {}
HighlightKind::DollarString => {
if saw_as_after_language || saw_as || saw_execute_immediate {
injections.push(Injection {
language: language.take().unwrap_or(InjectedLanguage::Sql),
range: token.range.clone(),
});
}
expect_language_name = false;
saw_as_after_language = false;
saw_as = false;
saw_execute = false;
saw_execute_immediate = false;
}
HighlightKind::Punctuation if token.text == ";" => {
language = None;
expect_language_name = false;
saw_as_after_language = false;
saw_as = false;
saw_execute = false;
saw_execute_immediate = false;
}
HighlightKind::Keyword if token.text.eq_ignore_ascii_case("language") => {
expect_language_name = true;
saw_as_after_language = false;
saw_as = false;
saw_execute = false;
saw_execute_immediate = false;
}
HighlightKind::Keyword | HighlightKind::Identifier if expect_language_name => {
language = Some(InjectedLanguage::from_language_word(token.text));
expect_language_name = false;
saw_as_after_language = false;
saw_as = false;
saw_execute = false;
saw_execute_immediate = false;
}
HighlightKind::Keyword if token.text.eq_ignore_ascii_case("as") => {
saw_as_after_language = language.is_some();
saw_as = language.is_none();
expect_language_name = false;
saw_execute = false;
saw_execute_immediate = false;
}
HighlightKind::Keyword if token.text.eq_ignore_ascii_case("execute") => {
expect_language_name = false;
saw_as_after_language = false;
saw_as = false;
saw_execute = true;
saw_execute_immediate = false;
}
HighlightKind::Keyword
if token.text.eq_ignore_ascii_case("immediate") && saw_execute =>
{
expect_language_name = false;
saw_as_after_language = false;
saw_as = false;
saw_execute = false;
saw_execute_immediate = true;
}
_ => {
expect_language_name = false;
saw_as_after_language = false;
saw_as = false;
saw_execute = false;
saw_execute_immediate = false;
}
}
}
injections
}
pub fn resolve_tokens(input: &str) -> Vec<ResolvedToken> {
let highlighted = highlight(input);
highlighted
.tokens
.iter()
.enumerate()
.filter_map(|(index, token)| {
if is_cortex_or_aisql_function_token(&highlighted.tokens, index) {
return Some(ResolvedToken {
range: token.range.clone(),
token_type: SemanticTokenType::Function,
modifiers: SemanticTokenModifiers::DEFAULT_LIBRARY,
});
}
let (token_type, modifiers) = semantic_token(token.kind)?;
Some(ResolvedToken {
range: token.range.clone(),
token_type,
modifiers,
})
})
.collect()
}
fn is_cortex_or_aisql_function_token(tokens: &[HighlightToken<'_>], index: usize) -> bool {
let token = &tokens[index];
if !matches!(
token.kind,
HighlightKind::Identifier | HighlightKind::Keyword
) {
return false;
}
if next_significant(tokens, index).is_none_or(|next| tokens[next].text != "(") {
return false;
}
token.text.to_ascii_uppercase().starts_with("AI_")
|| is_snowflake_cortex_qualified_leaf(tokens, index)
}
fn is_snowflake_cortex_qualified_leaf(tokens: &[HighlightToken<'_>], index: usize) -> bool {
let Some(dot_before_fn) = prev_significant(tokens, index) else {
return false;
};
if tokens[dot_before_fn].text != "." {
return false;
}
let Some(cortex) = prev_significant(tokens, dot_before_fn) else {
return false;
};
if !tokens[cortex].text.eq_ignore_ascii_case("cortex") {
return false;
}
let Some(dot_before_cortex) = prev_significant(tokens, cortex) else {
return false;
};
if tokens[dot_before_cortex].text != "." {
return false;
}
let Some(snowflake) = prev_significant(tokens, dot_before_cortex) else {
return false;
};
tokens[snowflake].text.eq_ignore_ascii_case("snowflake")
}
fn prev_significant(tokens: &[HighlightToken<'_>], index: usize) -> Option<usize> {
tokens[..index]
.iter()
.enumerate()
.rev()
.find(|(_, token)| {
!matches!(
token.kind,
HighlightKind::Whitespace | HighlightKind::Comment
)
})
.map(|(index, _)| index)
}
fn next_significant(tokens: &[HighlightToken<'_>], index: usize) -> Option<usize> {
tokens
.iter()
.enumerate()
.skip(index + 1)
.find(|(_, token)| {
!matches!(
token.kind,
HighlightKind::Whitespace | HighlightKind::Comment
)
})
.map(|(index, _)| index)
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct SemanticTokens {
pub tokens: Vec<ResolvedToken>,
pub injections: Vec<Injection>,
}
pub fn semantic_tokens(input: &str) -> SemanticTokens {
SemanticTokens {
tokens: resolve_tokens(input),
injections: detect_injections(input),
}
}
pub fn line_tokens(input: &str) -> Vec<LineToken> {
let resolved = resolve_tokens(input);
let index = LineIndex::new(input);
let mut out = Vec::new();
for token in &resolved {
let mut piece_start = token.range.start;
for piece in input[token.range.clone()].split('\n') {
let length = utf16_len(piece);
if length > 0 {
let position = index.utf16_position(piece_start);
out.push(LineToken {
line: position.line,
start_char: position.character,
length,
token_type: token.token_type.index(),
modifiers: token.modifiers.bits(),
});
}
piece_start += piece.len() + 1; }
}
out
}
pub fn line_tokens_utf8(input: &str) -> Vec<LineToken> {
let resolved = resolve_tokens(input);
let index = LineIndex::new(input);
let mut out = Vec::new();
for token in &resolved {
let mut piece_start = token.range.start;
for piece in input[token.range.clone()].split('\n') {
let length = piece.len() as u32;
if length > 0 {
let position = index.utf8_position(piece_start);
out.push(LineToken {
line: position.line,
start_char: position.character,
length,
token_type: token.token_type.index(),
modifiers: token.modifiers.bits(),
});
}
piece_start += piece.len() + 1; }
}
out
}
pub fn delta_encode(tokens: &[LineToken]) -> Vec<[u32; 5]> {
let mut out = Vec::with_capacity(tokens.len());
let (mut prev_line, mut prev_col) = (0u32, 0u32);
for token in tokens {
let delta_line = token.line - prev_line;
let delta_start = if delta_line == 0 {
token.start_char - prev_col
} else {
token.start_char
};
out.push([
delta_line,
delta_start,
token.length,
token.token_type,
token.modifiers,
]);
(prev_line, prev_col) = (token.line, token.start_char);
}
out
}
pub fn semantic_tokens_lsp(input: &str) -> Vec<[u32; 5]> {
delta_encode(&line_tokens(input))
}
pub fn semantic_tokens_lsp_utf8(input: &str) -> Vec<[u32; 5]> {
delta_encode(&line_tokens_utf8(input))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn legend_indices_match_their_position() {
for (i, ty) in SemanticTokenType::LEGEND.iter().enumerate() {
assert_eq!(ty.index() as usize, i, "legend index drift for {ty:?}");
}
assert_eq!(SemanticTokenType::Keyword.name(), "keyword");
assert_eq!(SemanticTokenType::Parameter.name(), "parameter");
assert_eq!(SemanticTokenType::Namespace.name(), "namespace");
}
#[test]
fn modifier_bitset_is_a_real_bitset() {
let both = SemanticTokenModifiers::DOCUMENTATION | SemanticTokenModifiers::DEFAULT_LIBRARY;
assert_eq!(both.bits(), 0b11);
assert!(both.contains(SemanticTokenModifiers::DOCUMENTATION));
assert!(both.contains(SemanticTokenModifiers::DEFAULT_LIBRARY));
assert!(!SemanticTokenModifiers::NONE.contains(SemanticTokenModifiers::DOCUMENTATION));
assert_eq!(SemanticTokenModifiers::LEGEND.len(), 2);
}
#[test]
fn every_highlight_kind_maps_consistently() {
use HighlightKind::*;
for kind in [
Keyword,
Type,
Identifier,
QuotedIdentifier,
String,
DollarString,
Number,
Variable,
Operator,
Comment,
] {
assert!(semantic_token(kind).is_some(), "{kind:?} should map");
}
for kind in [Whitespace, Punctuation, Error] {
assert!(semantic_token(kind).is_none(), "{kind:?} should not map");
}
assert_eq!(
semantic_token(Keyword).unwrap().1,
SemanticTokenModifiers::DEFAULT_LIBRARY
);
assert_eq!(
semantic_token(Type).unwrap().1,
SemanticTokenModifiers::DEFAULT_LIBRARY
);
assert_eq!(
semantic_token(Comment).unwrap().1,
SemanticTokenModifiers::DOCUMENTATION
);
}
#[test]
fn resolve_drops_trivia_and_punctuation() {
let toks = resolve_tokens("SELECT a, 1 -- c\n");
let types: Vec<_> = toks.iter().map(|t| t.token_type).collect();
assert_eq!(
types,
vec![
SemanticTokenType::Keyword,
SemanticTokenType::Variable,
SemanticTokenType::Number,
SemanticTokenType::Comment,
]
);
}
#[test]
fn resolved_ranges_match_source_text() {
let sql = "SELECT $1::NUMBER FROM t";
for tok in resolve_tokens(sql) {
assert!(tok.range.end <= sql.len());
}
let sem = semantic_tokens(sql);
let by_text: Vec<_> = sem
.tokens
.iter()
.map(|t| (&sql[t.range.clone()], t.token_type))
.collect();
assert!(by_text.contains(&("$1", SemanticTokenType::Parameter)));
assert!(by_text.contains(&("NUMBER", SemanticTokenType::Type)));
assert!(by_text.contains(&("::", SemanticTokenType::Operator)));
}
#[test]
fn line_tokens_are_utf16_and_split_multiline() {
let sql = "SELECT '芋' /* a\nb */ x";
let lines = line_tokens(sql);
let string_tok = lines
.iter()
.find(|t| t.token_type == SemanticTokenType::String.index())
.unwrap();
assert_eq!(string_tok.length, 3);
let comment_lines: Vec<_> = lines
.iter()
.filter(|t| t.token_type == SemanticTokenType::Comment.index())
.map(|t| t.line)
.collect();
assert_eq!(comment_lines, vec![0, 1]);
}
#[test]
fn delta_encode_resets_column_on_newline() {
let sql = "SELECT a\nFROM t";
let encoded = semantic_tokens_lsp(sql);
assert_eq!(
encoded[0],
[0, 0, 6, 0, SemanticTokenModifiers::DEFAULT_LIBRARY.bits()]
);
assert_eq!(encoded[1], [0, 7, 1, 2, 0]);
assert_eq!(
encoded[2],
[1, 0, 4, 0, SemanticTokenModifiers::DEFAULT_LIBRARY.bits()]
);
assert_eq!(encoded[3], [0, 5, 1, 2, 0]);
}
#[test]
fn empty_input_yields_no_tokens() {
assert!(line_tokens("").is_empty());
assert!(semantic_tokens_lsp("").is_empty());
assert!(detect_injections("").is_empty());
assert!(semantic_tokens("").tokens.is_empty());
}
#[test]
fn detects_javascript_injection_from_language_clause() {
let sql = "CREATE FUNCTION f() RETURNS STRING LANGUAGE JAVASCRIPT AS $$ return 1; $$;";
let injections = detect_injections(sql);
assert_eq!(injections.len(), 1);
assert_eq!(injections[0].language, InjectedLanguage::JavaScript);
let body = &sql[injections[0].range.clone()];
assert!(body.starts_with("$$"));
assert!(body.ends_with("$$"));
assert!(body.contains("return 1;"));
}
#[test]
fn detects_python_and_scala_and_java_injections() {
for (word, expected) in [
("PYTHON", InjectedLanguage::Python),
("Java", InjectedLanguage::Java),
("scala", InjectedLanguage::Scala),
("SQL", InjectedLanguage::Sql),
] {
let sql = format!("CREATE FUNCTION f() RETURNS INT LANGUAGE {word} AS $$x$$;");
let injections = detect_injections(&sql);
assert_eq!(injections.len(), 1, "for {word}");
assert_eq!(injections[0].language, expected, "for {word}");
}
}
#[test]
fn body_without_language_clause_defaults_to_sql() {
let sql = "CREATE PROCEDURE p() RETURNS STRING AS $$ BEGIN RETURN 'ok'; END $$; \
EXECUTE IMMEDIATE $$ SELECT 1 $$;";
let injections = detect_injections(sql);
assert_eq!(injections.len(), 2);
assert_eq!(injections[0].language, InjectedLanguage::Sql);
assert_eq!(injections[1].language, InjectedLanguage::Sql);
assert_eq!(InjectedLanguage::Sql.scope(), "source.snowflake-sql");
assert_eq!(InjectedLanguage::JavaScript.scope(), "source.js");
}
#[test]
fn non_body_dollar_strings_do_not_get_injected() {
let sql = "SELECT $$plain text$$ AS value; \
CREATE PROCEDURE p() LANGUAGE PYTHON IMPORTS = ($$stage/file.py$$) AS $$body$$;";
let injections = detect_injections(sql);
assert_eq!(injections.len(), 1);
assert_eq!(injections[0].language, InjectedLanguage::Python);
assert_eq!(&sql[injections[0].range.clone()], "$$body$$");
}
#[test]
fn language_clause_does_not_leak_across_statements() {
let sql = "CREATE FUNCTION a() RETURNS INT LANGUAGE JAVASCRIPT AS $$1$$; \
EXECUTE IMMEDIATE $$ SELECT 2 $$;";
let injections = detect_injections(sql);
assert_eq!(injections.len(), 2);
assert_eq!(injections[0].language, InjectedLanguage::JavaScript);
assert_eq!(injections[1].language, InjectedLanguage::Sql);
}
#[test]
fn multiple_bodies_each_get_their_language() {
let sql = "CREATE FUNCTION a() LANGUAGE PYTHON AS $$py$$; \
CREATE FUNCTION b() LANGUAGE JAVASCRIPT AS $$js$$;";
let injections = detect_injections(sql);
assert_eq!(injections.len(), 2);
assert_eq!(injections[0].language, InjectedLanguage::Python);
assert_eq!(injections[1].language, InjectedLanguage::JavaScript);
}
#[test]
fn semantic_tokens_bundles_tokens_and_injections() {
let sql = "CREATE FUNCTION f() LANGUAGE JAVASCRIPT AS $$ return 1; $$";
let sem = semantic_tokens(sql);
assert_eq!(sem.injections.len(), 1);
assert_eq!(sem.injections[0].language, InjectedLanguage::JavaScript);
let dollar = sem
.tokens
.iter()
.find(|t| {
t.token_type == SemanticTokenType::String && sql[t.range.clone()].contains("$$")
})
.expect("dollar string token");
assert_eq!(dollar.range, sem.injections[0].range);
}
#[test]
fn never_panics_on_adversarial_input() {
for sql in [
"$$",
"$$ unterminated",
"LANGUAGE",
"LANGUAGE $$x$$",
";;;;",
"@~/stage/ $1 :name ? ->> |> => :: ->",
"'長芋' \"畑\" -- 芋\n/* 芋 */ $$芋$$",
"\r\n\r\n",
] {
let lines = line_tokens(sql);
let encoded = delta_encode(&lines);
assert_eq!(lines.len(), encoded.len());
let _ = semantic_tokens(sql);
let _ = semantic_tokens_lsp(sql);
}
}
}