use anyhow::Result;
use tree_sitter::{Language, Query, QueryCursor, QueryPredicateArg, StreamingIterator, Tree};
use super::engine::{byte_to_char_col, char_to_byte_col};
pub struct TextObjectQuery {
query: Query,
capture_names: Vec<String>,
}
impl TextObjectQuery {
pub(super) fn compile(language: &Language, src: &str) -> Result<Self> {
let query = Query::new(language, src)?;
let capture_names = query
.capture_names()
.iter()
.map(|s| s.to_string())
.collect();
Ok(Self {
query,
capture_names,
})
}
fn for_each_range(
&self,
source: &str,
tree: &Tree,
target: &str,
mut yield_range: impl FnMut(std::ops::Range<usize>, (usize, usize), (usize, usize)),
) {
let mut cursor = QueryCursor::new();
let mut matches = cursor.matches(&self.query, tree.root_node(), source.as_bytes());
while let Some(m) = matches.next() {
for cap in m.captures {
let name = self
.capture_names
.get(cap.index as usize)
.map(String::as_str)
.unwrap_or("");
if name != target {
continue;
}
let (b, sp, ep) = adjust_for_node(target, source.as_bytes(), cap.node);
yield_range(b, sp, ep);
}
for pred in self.query.general_predicates(m.pattern_index) {
if pred.operator.as_ref() != "make-range!" {
continue;
}
let (name, start_idx, end_idx) = match pred.args.as_ref() {
[
QueryPredicateArg::String(n),
QueryPredicateArg::Capture(s),
QueryPredicateArg::Capture(e),
] => (n.as_ref(), *s, *e),
_ => continue,
};
if name != target {
continue;
}
let mut span_start: Option<tree_sitter::Node> = None;
let mut span_end: Option<tree_sitter::Node> = None;
for cap in m.captures {
if cap.index == start_idx {
span_start = match span_start {
None => Some(cap.node),
Some(prev) if cap.node.start_byte() < prev.start_byte() => {
Some(cap.node)
}
other => other,
};
}
if cap.index == end_idx {
span_end = match span_end {
None => Some(cap.node),
Some(prev) if cap.node.end_byte() > prev.end_byte() => Some(cap.node),
other => other,
};
}
}
if let (Some(s), Some(e)) = (span_start, span_end) {
yield_range(
s.start_byte()..e.end_byte(),
point(s.start_position()),
point(e.end_position()),
);
}
}
}
}
pub(super) fn find(
&self,
source: &str,
tree: &Tree,
target: &str,
cursor_row: usize,
cursor_col_chars: usize,
) -> Option<(usize, usize, usize, usize)> {
let cursor_pt = (
cursor_row,
char_to_byte_col(source, cursor_row, cursor_col_chars),
);
let mut best: Option<Candidate> = None;
self.for_each_range(source, tree, target, |bytes, start, end| {
consider(&mut best, bytes, start, end, cursor_pt);
});
let c = best?;
Some((
c.start.0,
byte_to_char_col(source, c.start.0, c.start.1),
c.end.0,
byte_to_char_col(source, c.end.0, c.end.1),
))
}
#[cfg(test)]
pub(super) fn all(
&self,
source: &str,
tree: &Tree,
target: &str,
) -> Vec<(usize, usize, usize, usize)> {
let mut out: Vec<(usize, usize, usize, usize)> = Vec::new();
self.for_each_range(source, tree, target, |_bytes, start, end| {
out.push((
start.0,
byte_to_char_col(source, start.0, start.1),
end.0,
byte_to_char_col(source, end.0, end.1),
));
});
out.sort_unstable();
out.dedup();
out
}
}
struct Candidate {
bytes: std::ops::Range<usize>,
start: (usize, usize),
end: (usize, usize),
}
fn point(p: tree_sitter::Point) -> (usize, usize) {
(p.row, p.column)
}
fn adjust_for_node(
target: &str,
src: &[u8],
node: tree_sitter::Node,
) -> (std::ops::Range<usize>, (usize, usize), (usize, usize)) {
match target {
"parameter.outer" => extend_separator(src, node),
"function.inner" | "class.inner" => shrink_braces(src, node),
_ => raw_range(node),
}
}
fn raw_range(node: tree_sitter::Node) -> (std::ops::Range<usize>, (usize, usize), (usize, usize)) {
(
node.start_byte()..node.end_byte(),
point(node.start_position()),
point(node.end_position()),
)
}
fn extend_separator(
src: &[u8],
node: tree_sitter::Node,
) -> (std::ops::Range<usize>, (usize, usize), (usize, usize)) {
let (s, e) = (node.start_byte(), node.end_byte());
let (start, end) = (point(node.start_position()), point(node.end_position()));
let is_sep = |n: &tree_sitter::Node| !n.is_named() && n.kind() == ",";
if let Some(comma) = node.next_sibling().filter(is_sep) {
let mut i = comma.end_byte();
while i < src.len() && matches!(src[i], b' ' | b'\t') {
i += 1;
}
return (s..i, start, move_point(src, end, e, i));
}
if let Some(comma) = node.prev_sibling().filter(is_sep) {
let mut j = comma.start_byte();
while j > 0 && matches!(src[j - 1], b' ' | b'\t') {
j -= 1;
}
return (j..e, move_point(src, start, s, j), end);
}
(s..e, start, end)
}
fn shrink_braces(
src: &[u8],
node: tree_sitter::Node,
) -> (std::ops::Range<usize>, (usize, usize), (usize, usize)) {
let mut open = None;
let mut close = None;
let mut walk = node.walk();
for child in node.children(&mut walk) {
if child.is_named() {
continue;
}
match child.kind() {
"{" if open.is_none() => open = Some(child),
"}" => close = Some(child),
_ => {}
}
}
let (Some(open), Some(close)) = (open, close) else {
return raw_range(node);
};
let mut ns = open.end_byte();
let mut ne = close.start_byte();
while ns < ne && src[ns].is_ascii_whitespace() {
ns += 1;
}
while ne > ns && src[ne - 1].is_ascii_whitespace() {
ne -= 1;
}
(
ns..ne,
move_point(src, point(open.end_position()), open.end_byte(), ns),
move_point(src, point(close.start_position()), close.start_byte(), ne),
)
}
fn move_point(src: &[u8], from_pt: (usize, usize), from: usize, to: usize) -> (usize, usize) {
if to >= from {
let (mut row, mut col) = from_pt;
for &b in &src[from..to] {
if b == b'\n' {
row += 1;
col = 0;
} else {
col += 1;
}
}
(row, col)
} else {
let mut row = from_pt.0;
for &b in &src[to..from] {
if b == b'\n' {
row -= 1;
}
}
let line_start = src[..to]
.iter()
.rposition(|&b| b == b'\n')
.map_or(0, |p| p + 1);
(row, to - line_start)
}
}
fn consider(
best: &mut Option<Candidate>,
bytes: std::ops::Range<usize>,
start: (usize, usize),
end: (usize, usize),
cursor: (usize, usize),
) {
if !(start <= cursor && cursor < end) {
return;
}
let len = bytes.end - bytes.start;
let take = match best {
None => true,
Some(c) => len < c.bytes.end - c.bytes.start,
};
if take {
*best = Some(Candidate { bytes, start, end });
}
}