use anyhow::{Context, Result};
use tree_sitter::{InputEdit, Language, Parser, Point, Tree};
use super::highlight::{Capture, HighlightQuery};
use super::indent::IndentQuery;
use super::injection::InjectionEngine;
use super::textobject::TextObjectQuery;
use super::{bracket, highlight, indent, textobject};
pub struct Engine {
parser: Parser,
tree: Option<Tree>,
source: String,
parsed_version: Option<u64>,
highlight: HighlightQuery,
indent: Option<IndentQuery>,
textobject: Option<TextObjectQuery>,
injection: Option<InjectionEngine>,
pub warnings: Vec<String>,
}
impl Engine {
pub(super) fn new(
language: Language,
highlights_src: &str,
textobjects_src: Option<&str>,
indents_src: Option<&str>,
injection: Option<InjectionEngine>,
) -> Result<Self> {
let mut parser = Parser::new();
parser
.set_language(&language)
.context("setting parser language (ABI mismatch?)")?;
let highlight = highlight::HighlightQuery::compile(&language, highlights_src)
.context("compiling highlights query")?;
let mut warnings = Vec::new();
let textobject = match textobjects_src {
Some(src) => match textobject::TextObjectQuery::compile(&language, src) {
Ok(q) => Some(q),
Err(e) => {
warnings.push(format!(
"textobjects.scm compile failed, syntactic text objects disabled: {e}"
));
None
}
},
None => None,
};
let indent = match indents_src {
Some(src) => match indent::IndentQuery::compile(&language, src) {
Ok(q) => Some(q),
Err(e) => {
warnings.push(format!(
"indents.scm compile failed, auto-indent disabled: {e}"
));
None
}
},
None => None,
};
Ok(Self {
parser,
tree: None,
source: String::new(),
parsed_version: None,
highlight,
indent,
textobject,
injection,
warnings,
})
}
pub fn refresh(&mut self, source: &str, version: u64) {
if self.parsed_version == Some(version) {
return;
}
let old_tree = match self.tree.as_mut() {
Some(tree) if !self.source.is_empty() => {
let edit = compute_input_edit(&self.source, source);
tree.edit(&edit);
Some(&*tree)
}
_ => None,
};
self.tree = self.parser.parse(source, old_tree);
self.source = source.to_string();
self.parsed_version = Some(version);
}
pub fn captures_in_rows(&self, start_row: usize, end_row: usize) -> Vec<Capture> {
let Some(tree) = &self.tree else {
return Vec::new();
};
let mut out = self
.highlight
.captures_in_rows(&self.source, tree, start_row, end_row);
if let Some(inj) = self.injection.as_ref() {
out.extend(inj.captures_in_rows(&self.source, tree, start_row, end_row));
}
out.sort_by_key(|c| (c.start_row, c.start_col));
out
}
pub fn indent_begins_at(&self, row: usize) -> bool {
let Some(tree) = &self.tree else {
return false;
};
let Some(q) = self.indent.as_ref() else {
return false;
};
q.begins_at(&self.source, tree, row)
}
pub fn indent_scopes_in_rows(&self, start_row: usize, end_row: usize) -> Vec<(usize, usize)> {
let Some(tree) = &self.tree else {
return Vec::new();
};
let mut out = Vec::new();
if let Some(q) = self.indent.as_ref() {
out.extend(q.scopes_in_rows(&self.source, tree, start_row, end_row));
}
if let Some(inj) = self.injection.as_ref() {
out.extend(inj.indent_scopes_in_rows(&self.source, tree, start_row, end_row));
}
out.sort_by(|a, b| a.0.cmp(&b.0).then(b.1.cmp(&a.1)));
out
}
pub fn find_text_object(
&self,
target: &str,
cursor_row: usize,
cursor_col_chars: usize,
) -> Option<(usize, usize, usize, usize)> {
let tree = self.tree.as_ref()?;
let q = self.textobject.as_ref()?;
q.find(&self.source, tree, target, cursor_row, cursor_col_chars)
}
pub fn matching_bracket(&self, row: usize, col_chars: usize) -> Option<(usize, usize)> {
let tree = self.tree.as_ref()?;
bracket::matching(&self.source, tree, row, col_chars)
}
}
pub(super) fn byte_to_char_col(source: &str, row: usize, byte_col: usize) -> usize {
let line = source.lines().nth(row).unwrap_or("");
let take = byte_col.min(line.len());
line[..take].chars().count()
}
pub(super) fn char_to_byte_col(source: &str, row: usize, char_col: usize) -> usize {
let line = source.lines().nth(row).unwrap_or("");
line.char_indices()
.nth(char_col)
.map(|(b, _)| b)
.unwrap_or(line.len())
}
fn compute_input_edit(old: &str, new: &str) -> InputEdit {
let old_bytes = old.as_bytes();
let new_bytes = new.as_bytes();
let common_prefix = old_bytes
.iter()
.zip(new_bytes.iter())
.take_while(|(a, b)| a == b)
.count();
let max_suffix = old_bytes
.len()
.min(new_bytes.len())
.saturating_sub(common_prefix);
let common_suffix = old_bytes
.iter()
.rev()
.zip(new_bytes.iter().rev())
.take(max_suffix)
.take_while(|(a, b)| a == b)
.count();
let start_byte = common_prefix;
let old_end_byte = old_bytes.len() - common_suffix;
let new_end_byte = new_bytes.len() - common_suffix;
InputEdit {
start_byte,
old_end_byte,
new_end_byte,
start_position: byte_to_point(old_bytes, start_byte),
old_end_position: byte_to_point(old_bytes, old_end_byte),
new_end_position: byte_to_point(new_bytes, new_end_byte),
}
}
fn byte_to_point(bytes: &[u8], offset: usize) -> Point {
let offset = offset.min(bytes.len());
let mut row = 0usize;
let mut line_start = 0usize;
for (i, &b) in bytes[..offset].iter().enumerate() {
if b == b'\n' {
row += 1;
line_start = i + 1;
}
}
Point {
row,
column: offset - line_start,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn byte_to_char_handles_ascii() {
let src = "let x = 1\nprintln!(\"hi\")";
assert_eq!(byte_to_char_col(src, 0, 0), 0);
assert_eq!(byte_to_char_col(src, 0, 4), 4);
assert_eq!(byte_to_char_col(src, 1, 9), 9);
}
#[test]
fn byte_to_char_handles_multibyte() {
let src = "あ x";
assert_eq!(byte_to_char_col(src, 0, 3), 1);
assert_eq!(byte_to_char_col(src, 0, 5), 3);
}
#[test]
fn input_edit_single_byte_insertion() {
let edit = compute_input_edit("abc", "abXc");
assert_eq!(edit.start_byte, 2);
assert_eq!(edit.old_end_byte, 2);
assert_eq!(edit.new_end_byte, 3);
assert_eq!(edit.start_position, Point { row: 0, column: 2 });
assert_eq!(edit.new_end_position, Point { row: 0, column: 3 });
}
#[test]
fn input_edit_no_change_is_noop_range() {
let edit = compute_input_edit("hello", "hello");
assert_eq!(edit.start_byte, 5);
assert_eq!(edit.old_end_byte, 5);
assert_eq!(edit.new_end_byte, 5);
}
#[test]
fn input_edit_multi_line_replacement() {
let edit = compute_input_edit("fn a() {\n 1\n}\n", "fn a() {\n 42\n}\n");
assert_eq!(edit.start_byte, 11);
assert_eq!(edit.old_end_byte, 15 - 3);
assert_eq!(edit.new_end_byte, 16 - 3);
assert_eq!(edit.start_position, Point { row: 1, column: 2 });
}
#[test]
fn input_edit_full_replacement() {
let edit = compute_input_edit("abc", "xyz");
assert_eq!(edit.start_byte, 0);
assert_eq!(edit.old_end_byte, 3);
assert_eq!(edit.new_end_byte, 3);
}
}