use std::cell::RefCell;
use anyhow::{Context, Result};
use tree_sitter::{InputEdit, Language, Parser, Point, Tree};
use super::fold::FoldQuery;
use super::highlight::{Capture, HighlightQuery};
use super::indent::IndentQuery;
use super::injection::InjectionEngine;
use super::textobject::TextObjectQuery;
use super::{bracket, fold, highlight, indent, textobject};
pub struct Engine {
parser: Parser,
tree: Option<Tree>,
source: String,
line_starts: Vec<usize>,
parsed_version: Option<u64>,
highlight: HighlightQuery,
indent: Option<IndentQuery>,
textobject: Option<TextObjectQuery>,
fold: Option<FoldQuery>,
injection: Option<InjectionEngine>,
capture_cache: RefCell<Option<CaptureCache>>,
fold_cache: RefCell<FoldRegionsCache>,
pub warnings: Vec<String>,
}
type FoldRegionsCache = Option<(Option<u64>, Vec<(usize, usize)>)>;
struct CaptureCache {
version: Option<u64>,
start_row: usize,
end_row: usize,
captures: Vec<Capture>,
}
impl Engine {
pub(super) fn new(
language: Language,
highlights_src: &str,
textobjects_src: Option<&str>,
indents_src: Option<&str>,
folds_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,
};
let fold = match folds_src {
Some(src) => match fold::FoldQuery::compile(&language, src) {
Ok(q) => Some(q),
Err(e) => {
warnings.push(format!(
"folds.scm compile failed, syntax folding disabled: {e}"
));
None
}
},
None => None,
};
Ok(Self {
parser,
tree: None,
source: String::new(),
line_starts: vec![0],
parsed_version: None,
highlight,
indent,
textobject,
fold,
injection,
capture_cache: RefCell::new(None),
fold_cache: RefCell::new(None),
warnings,
})
}
pub fn is_current(&self, version: u64) -> bool {
self.parsed_version == Some(version)
}
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.line_starts = line_start_offsets(&self.source);
self.parsed_version = Some(version);
self.capture_cache.borrow_mut().take();
self.fold_cache.borrow_mut().take();
}
pub fn captures_in_rows(&self, start_row: usize, end_row: usize) -> Vec<Capture> {
if let Some(c) = self.capture_cache.borrow().as_ref()
&& c.version == self.parsed_version
&& c.start_row == start_row
&& c.end_row == end_row
{
return c.captures.clone();
}
let Some(tree) = &self.tree else {
return Vec::new();
};
let mut out = self.highlight.captures_in_rows(
&self.source,
&self.line_starts,
tree,
start_row,
end_row,
);
if let Some(inj) = self.injection.as_ref() {
out.extend(inj.captures_in_rows(
&self.source,
&self.line_starts,
tree,
start_row,
end_row,
));
}
out.sort_by_key(|c| (c.start_row, c.start_col));
*self.capture_cache.borrow_mut() = Some(CaptureCache {
version: self.parsed_version,
start_row,
end_row,
captures: out.clone(),
});
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 has_fold_query(&self) -> bool {
self.fold.is_some()
}
pub fn fold_regions(&self) -> Vec<(usize, usize)> {
if let Some((v, r)) = self.fold_cache.borrow().as_ref()
&& *v == self.parsed_version
{
return r.clone();
}
let regions = match (&self.tree, &self.fold) {
(Some(tree), Some(q)) => fold::normalize_regions(q.regions(&self.source, tree)),
_ => Vec::new(),
};
*self.fold_cache.borrow_mut() = Some((self.parsed_version, regions.clone()));
regions
}
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)
}
#[cfg(test)]
pub(crate) fn all_text_objects(&self, target: &str) -> Vec<(usize, usize, usize, usize)> {
let (Some(tree), Some(q)) = (self.tree.as_ref(), self.textobject.as_ref()) else {
return Vec::new();
};
q.all(&self.source, tree, target)
}
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()
}
fn line_start_offsets(source: &str) -> Vec<usize> {
let mut starts = Vec::with_capacity(source.len() / 32 + 2);
starts.push(0);
for (i, b) in source.bytes().enumerate() {
if b == b'\n' {
starts.push(i + 1);
}
}
starts.push(source.len());
starts
}
pub(super) fn byte_to_char_col_indexed(
source: &str,
line_starts: &[usize],
row: usize,
byte_col: usize,
) -> usize {
let Some(&start) = line_starts.get(row) else {
return 0;
};
let end = line_starts.get(row + 1).copied().unwrap_or(source.len());
let line = &source[start..end.min(source.len())];
let take = byte_col.min(line.len());
match line.get(..take) {
Some(prefix) => prefix.chars().count(),
None => line.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);
}
use crate::config::Config;
use crate::syntax::Loader;
use std::path::Path;
use std::time::Instant;
fn engine_for_path(sample: &str, source: &str) -> Option<(Loader, Engine)> {
let cfg = Config::load(None).ok()?;
let spec = cfg.languages.by_path(Path::new(sample))?.clone();
let mut loader = Loader::new(cfg.grammar_dir.clone(), cfg.query_dir.clone());
let mut engine = loader.engine_for(&spec).ok()?;
engine.refresh(source, 1);
Some((loader, engine))
}
fn median_us(mut samples: Vec<f64>) -> f64 {
samples.sort_by(|a, b| a.partial_cmp(b).unwrap());
samples[samples.len() / 2]
}
#[test]
#[ignore]
fn perf_highlight_scaling() {
let unit = "fn compute(x: i32) -> i32 {\n let y = x * 2 + 1;\n y - 3\n}\n\n";
const WINDOW: usize = 50;
for lines_target in [2_000usize, 20_000, 100_000] {
let reps = lines_target / 5 + 1;
let source = unit.repeat(reps);
let total_rows = source.lines().count();
let Some((_loader, engine)) = engine_for_path("bench.rs", &source) else {
eprintln!("skip: rust grammar not installed");
return;
};
let mut miss = Vec::new();
for i in 0..400 {
let scroll = (i * 37) % total_rows.saturating_sub(WINDOW).max(1);
let t = Instant::now();
let caps = engine.captures_in_rows(scroll, scroll + WINDOW);
miss.push(t.elapsed().as_secs_f64() * 1e6);
std::hint::black_box(caps);
}
let mut hit = Vec::new();
let scroll = total_rows / 2;
let _ = engine.captures_in_rows(scroll, scroll + WINDOW); for _ in 0..2_000 {
let t = Instant::now();
let caps = engine.captures_in_rows(scroll, scroll + WINDOW);
hit.push(t.elapsed().as_secs_f64() * 1e6);
std::hint::black_box(caps);
}
eprintln!(
"rust {:>7} rows | miss(scroll/edit) median {:>7.1} us | hit(repaint) median {:>6.2} us",
total_rows,
median_us(miss),
median_us(hit),
);
}
}
#[test]
#[ignore]
fn perf_injection_markdown() {
const WINDOW: usize = 50;
let block = "```rust\nfn demo(n: usize) -> usize {\n let mut acc = 0;\n for i in 0..n { acc += i * 2; }\n acc\n}\n```\n\nSome prose paragraph between code blocks to mimic a real doc.\n\n";
let source = block.repeat(400);
let total_rows = source.lines().count();
let Some((_loader, mut engine)) = engine_for_path("bench.md", &source) else {
eprintln!("skip: markdown grammar not installed");
return;
};
let scroll = total_rows / 2;
let t = Instant::now();
let _ = engine.captures_in_rows(scroll, scroll + WINDOW);
let cold = t.elapsed().as_secs_f64() * 1e6;
let mut warm = Vec::new();
for v in 2..402u64 {
engine.refresh(&source, v);
let t = Instant::now();
let caps = engine.captures_in_rows(scroll, scroll + WINDOW);
warm.push(t.elapsed().as_secs_f64() * 1e6);
std::hint::black_box(caps);
}
eprintln!(
"markdown {} rows | cold(first paint) {:.1} us | warm(typing, sub-tree cache) median {:.1} us",
total_rows,
cold,
median_us(warm),
);
}
}