use lsp_types::{
Diagnostic, DiagnosticSeverity, FoldingRange, FoldingRangeKind, Hover, HoverContents,
MarkupContent, MarkupKind, Position, Range, SemanticToken, SemanticTokenModifier,
SemanticTokenType, TextEdit,
};
use sql_dialect_fmt_formatter::{format, FormatOptions};
use sql_dialect_fmt_highlight::{semantic, HighlightKind};
pub fn token_types() -> Vec<SemanticTokenType> {
semantic::SemanticTokenType::LEGEND
.iter()
.map(|ty| SemanticTokenType::new(ty.name()))
.collect()
}
pub fn token_modifiers() -> Vec<SemanticTokenModifier> {
semantic::SemanticTokenModifiers::LEGEND
.iter()
.map(|&name| SemanticTokenModifier::new(name))
.collect()
}
pub struct LineIndex<'a> {
text: &'a str,
line_starts: Vec<usize>,
}
impl<'a> LineIndex<'a> {
pub fn new(text: &'a str) -> Self {
let mut line_starts = vec![0];
line_starts.extend(
text.bytes()
.enumerate()
.filter(|&(_, b)| b == b'\n')
.map(|(i, _)| i + 1),
);
LineIndex { text, line_starts }
}
pub fn position(&self, offset: usize) -> Position {
let offset = offset.min(self.text.len());
let line = match self.line_starts.binary_search(&offset) {
Ok(line) => line,
Err(next) => next - 1,
};
let line_start = self.line_starts[line];
let col: usize = self.text[line_start..offset]
.chars()
.map(char::len_utf16)
.sum();
Position::new(line as u32, col as u32)
}
pub fn end(&self) -> Position {
self.position(self.text.len())
}
pub fn offset(&self, position: Position) -> usize {
let line = position.line as usize;
let Some(&line_start) = self.line_starts.get(line) else {
return self.text.len();
};
let mut remaining = position.character as usize; let mut offset = line_start;
for ch in self.text[line_start..].chars() {
let width = ch.len_utf16();
if remaining < width || ch == '\n' {
break;
}
remaining -= width;
offset += ch.len_utf8();
}
offset
}
}
pub fn format_edits(text: &str, options: &FormatOptions) -> Vec<TextEdit> {
let formatted = format(text, options);
if formatted == text {
return Vec::new();
}
let index = LineIndex::new(text);
vec![TextEdit {
range: Range::new(Position::new(0, 0), index.end()),
new_text: formatted,
}]
}
pub fn diagnostics(text: &str) -> Vec<Diagnostic> {
let index = LineIndex::new(text);
let to_range = |span: std::ops::Range<usize>| {
Range::new(index.position(span.start), index.position(span.end))
};
let make = |range: Range, message: String| Diagnostic {
range,
severity: Some(DiagnosticSeverity::ERROR),
source: Some("sql-dialect-fmt".to_string()),
message,
..Default::default()
};
let lex_errors = sql_dialect_fmt_highlight::highlight(text).errors;
let parse = sql_dialect_fmt_parser::parse(text);
let mut diagnostics: Vec<_> = lex_errors
.into_iter()
.map(|err| make(to_range(err.range()), err.message))
.chain(
parse
.errors()
.iter()
.map(|err| make(to_range(err.range()), err.message.clone())),
)
.collect();
diagnostics.extend(embedded_language_diagnostics(text, &index));
diagnostics
}
fn embedded_language_diagnostics(text: &str, index: &LineIndex<'_>) -> Vec<Diagnostic> {
let mut diagnostics = Vec::new();
let mut expect_language_name = false;
let mut language_name: Option<(&str, std::ops::Range<usize>)> = None;
let mut saw_as_after_language = false;
for token in sql_dialect_fmt_highlight::highlight(text).tokens {
match token.kind {
HighlightKind::Whitespace | HighlightKind::Comment => {}
HighlightKind::Punctuation if token.text == ";" => {
expect_language_name = false;
language_name = None;
saw_as_after_language = false;
}
HighlightKind::DollarString => {
if saw_as_after_language {
if let Some((word, range)) = language_name.take() {
if !is_supported_embedded_language(word) {
diagnostics.push(Diagnostic {
range: Range::new(index.position(range.start), index.position(range.end)),
severity: Some(DiagnosticSeverity::WARNING),
source: Some("sql-dialect-fmt".to_string()),
message: format!(
"unsupported embedded language {word}; expected SQL, JAVASCRIPT, PYTHON, JAVA, or SCALA"
),
..Default::default()
});
}
}
}
expect_language_name = false;
saw_as_after_language = false;
}
HighlightKind::Keyword if token.text.eq_ignore_ascii_case("language") => {
expect_language_name = true;
language_name = None;
saw_as_after_language = false;
}
HighlightKind::Keyword | HighlightKind::Identifier | HighlightKind::Type
if expect_language_name =>
{
language_name = Some((token.text, token.range));
expect_language_name = false;
}
HighlightKind::Keyword
if language_name.is_some() && token.text.eq_ignore_ascii_case("as") =>
{
saw_as_after_language = true;
}
_ => {
expect_language_name = false;
}
}
}
diagnostics
}
fn is_supported_embedded_language(word: &str) -> bool {
["SQL", "JAVASCRIPT", "PYTHON", "JAVA", "SCALA"]
.iter()
.any(|candidate| candidate.eq_ignore_ascii_case(word))
}
pub fn hover(text: &str, position: Position) -> Option<Hover> {
let index = LineIndex::new(text);
let info = sql_dialect_fmt_hover::hover_at(text, index.offset(position))?;
let mut value = format!("**{}**\n\n{}", info.title, info.body);
if let Some(url) = info.docs_url {
value.push_str(&format!("\n\n[Snowflake docs]({url})"));
}
Some(Hover {
contents: HoverContents::Markup(MarkupContent {
kind: MarkupKind::Markdown,
value,
}),
range: Some(Range::new(
index.position(info.range.start),
index.position(info.range.end),
)),
})
}
pub fn folding_ranges(text: &str) -> Vec<FoldingRange> {
let index = LineIndex::new(text);
let root = sql_dialect_fmt_parser::parse(text).syntax();
root.children()
.filter_map(|stmt| {
let mut tokens = stmt
.descendants_with_tokens()
.filter_map(|el| el.into_token())
.filter(|t| !t.kind().is_trivia());
let first = tokens.next()?;
let last = tokens.last().unwrap_or_else(|| first.clone());
let start = index.position(first.text_range().start().into()).line;
let end = index.position(last.text_range().end().into()).line;
(end > start).then_some(FoldingRange {
start_line: start,
end_line: end,
kind: Some(FoldingRangeKind::Region),
..FoldingRange::default()
})
})
.collect()
}
pub fn apply_change(text: &str, range: Option<Range>, new_text: &str) -> String {
let Some(range) = range else {
return new_text.to_string();
};
let index = LineIndex::new(text);
let a = index.offset(range.start);
let b = index.offset(range.end);
let (start, end) = (a.min(b), a.max(b));
let mut out = String::with_capacity(text.len() - (end - start) + new_text.len());
out.push_str(&text[..start]);
out.push_str(new_text);
out.push_str(&text[end..]);
out
}
pub fn semantic_tokens(text: &str) -> Vec<SemanticToken> {
semantic::semantic_tokens_lsp(text)
.into_iter()
.map(
|[delta_line, delta_start, length, token_type, token_modifiers_bitset]| SemanticToken {
delta_line,
delta_start,
length,
token_type,
token_modifiers_bitset,
},
)
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn line_index_maps_offsets_to_utf16_positions() {
let text = "SELECT a\nFROM 芋;\n"; let index = LineIndex::new(text);
assert_eq!(index.position(0), Position::new(0, 0));
assert_eq!(index.position(7), Position::new(0, 7)); let from = text.find("FROM").unwrap();
assert_eq!(index.position(from), Position::new(1, 0));
let semicolon = text.find(';').unwrap();
assert_eq!(index.position(semicolon), Position::new(1, 6)); }
#[test]
fn formatting_replaces_the_whole_document() {
let edits = format_edits("select a,b from t", &FormatOptions::default());
assert_eq!(edits.len(), 1);
assert_eq!(edits[0].new_text, "SELECT a, b\nFROM t;\n");
assert_eq!(edits[0].range.start, Position::new(0, 0));
}
#[test]
fn already_formatted_input_yields_no_edits() {
let formatted = "SELECT a, b\nFROM t;\n";
assert!(format_edits(formatted, &FormatOptions::default()).is_empty());
}
#[test]
fn clean_sql_has_no_diagnostics() {
assert!(diagnostics("select 1").is_empty());
}
#[test]
fn broken_sql_reports_a_diagnostic() {
let diags = diagnostics("select from where");
assert!(!diags.is_empty());
assert_eq!(diags[0].severity, Some(DiagnosticSeverity::ERROR));
}
#[test]
fn lexer_errors_reach_diagnostics() {
let text = "SELECT 'oops";
let diags = diagnostics(text);
let lex_diag = diags
.iter()
.find(|d| d.message.contains("unterminated string"))
.expect("an unterminated-string diagnostic");
assert_eq!(lex_diag.severity, Some(DiagnosticSeverity::ERROR));
let quote = text.find('\'').unwrap() as u32;
assert_eq!(lex_diag.range.start, Position::new(0, quote));
assert!(lex_diag.range.end.character > lex_diag.range.start.character);
}
#[test]
fn parser_diagnostic_range_covers_the_token() {
let text = "MERGE tgt USING src ON a = b";
let diags = diagnostics(text);
let into = diags
.iter()
.find(|d| d.message == "expected INTO")
.expect("an INTO diagnostic");
let col = text.find("tgt").unwrap() as u32;
assert_eq!(into.range.start, Position::new(0, col));
assert_eq!(into.range.end, Position::new(0, col + 3));
}
#[test]
fn clean_sql_still_has_no_lexer_or_parser_diagnostics() {
assert!(diagnostics("SELECT a FROM t").is_empty());
}
#[test]
fn unsupported_embedded_language_is_a_warning() {
let text = "CREATE FUNCTION f() RETURNS STRING LANGUAGE RUBY AS $$x$$;";
let diags = diagnostics(text);
let language = diags
.iter()
.find(|d| d.message.contains("unsupported embedded language RUBY"))
.expect("unsupported-language diagnostic");
assert_eq!(language.severity, Some(DiagnosticSeverity::WARNING));
let col = text.find("RUBY").unwrap() as u32;
assert_eq!(language.range.start, Position::new(0, col));
assert_eq!(language.range.end, Position::new(0, col + 4));
}
#[test]
fn embedded_language_warning_does_not_fire_for_plain_columns_or_dynamic_sql() {
for text in [
"SELECT language FROM t;",
"EXECUTE IMMEDIATE $$ SELECT 1 $$;",
] {
assert!(
diagnostics(text)
.iter()
.all(|diag| !diag.message.contains("unsupported embedded language")),
"{text}"
);
}
}
#[test]
fn offset_is_the_inverse_of_position() {
let text = "SELECT a\nFROM 芋;\n";
let index = LineIndex::new(text);
for offset in [
0usize,
7,
text.find("FROM").unwrap(),
text.find(';').unwrap(),
] {
assert_eq!(index.offset(index.position(offset)), offset);
}
}
#[test]
fn hover_describes_a_type() {
let src = "select x::varchar from t";
let col = src.find("varchar").unwrap() as u32;
let hover = hover(src, Position::new(0, col)).expect("hover");
assert!(hover.range.is_some());
match hover.contents {
HoverContents::Markup(m) => assert!(m.value.to_lowercase().contains("varchar")),
_ => panic!("expected markup"),
}
}
#[test]
fn apply_change_splices_an_incremental_edit() {
let text = "hello\nworld\n";
let range = Range::new(Position::new(1, 0), Position::new(1, 5));
assert_eq!(apply_change(text, Some(range), "snow"), "hello\nsnow\n");
}
#[test]
fn apply_change_with_no_range_replaces_whole_document() {
assert_eq!(apply_change("old", None, "new text"), "new text");
}
#[test]
fn folding_ranges_cover_multiline_statements() {
let ranges = folding_ranges("select a,\nb\nfrom t;\n\nselect 1;");
assert_eq!(ranges.len(), 1); assert_eq!(ranges[0].start_line, 0);
assert_eq!(ranges[0].end_line, 2);
}
#[test]
fn semantic_tokens_tag_keywords() {
let tokens = semantic_tokens("select a from t");
assert!(!tokens.is_empty());
assert_eq!(tokens[0].delta_line, 0);
assert_eq!(tokens[0].delta_start, 0);
assert_eq!(tokens[0].length, 6);
assert_eq!(tokens[0].token_type, 0);
assert_eq!(
tokens[0].token_modifiers_bitset,
semantic::SemanticTokenModifiers::DEFAULT_LIBRARY.bits()
);
}
#[test]
fn server_legend_equals_the_highlighter_legend() {
let advertised: Vec<String> = token_types()
.iter()
.map(|t| t.as_str().to_string())
.collect();
let expected: Vec<String> = semantic::SemanticTokenType::LEGEND
.iter()
.map(|t| t.name().to_string())
.collect();
assert_eq!(advertised, expected);
assert_eq!(advertised.last().map(String::as_str), Some("function"));
let mods: Vec<String> = token_modifiers()
.iter()
.map(|m| m.as_str().to_string())
.collect();
assert_eq!(mods, vec!["documentation", "defaultLibrary"]);
}
#[test]
fn semantic_tokens_are_monotonic_and_never_panic_on_multiline() {
let tokens = semantic_tokens("select 1 /* a\nb */ from t");
assert!(tokens.iter().all(|t| t.length > 0));
}
}