use tree_sitter::Tree;
const MAX_SCAN_DISTANCE: usize = 5000;
const BRACKET_PAIRS: [(u8, u8); 3] = [(b'(', b')'), (b'[', b']'), (b'{', b'}')];
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct BracketMatch {
pub cursor_bracket: usize,
pub matching_bracket: usize,
}
pub fn find_matching_bracket(buffer: &str, cursor: usize, tree: &Tree) -> Option<BracketMatch> {
let bytes = buffer.as_bytes();
if cursor < bytes.len()
&& let Some(m) = try_match_at(bytes, cursor, tree)
{
return Some(m);
}
if cursor > 0
&& cursor <= bytes.len()
&& is_closing_bracket(bytes[cursor - 1])
&& let Some(m) = try_match_at(bytes, cursor - 1, tree)
{
return Some(m);
}
None
}
fn try_match_at(bytes: &[u8], pos: usize, tree: &Tree) -> Option<BracketMatch> {
let ch = bytes[pos];
let (opening, closing, is_opening) = find_pair_for(ch)?;
if is_in_string_or_comment(tree, pos) {
return None;
}
let match_pos = if is_opening {
scan_forward(bytes, pos, opening, closing, tree)?
} else {
scan_backward(bytes, pos, opening, closing, tree)?
};
Some(BracketMatch {
cursor_bracket: pos,
matching_bracket: match_pos,
})
}
fn find_pair_for(ch: u8) -> Option<(u8, u8, bool)> {
for &(open, close) in &BRACKET_PAIRS {
if ch == open {
return Some((open, close, true));
}
if ch == close {
return Some((open, close, false));
}
}
None
}
fn is_closing_bracket(ch: u8) -> bool {
matches!(ch, b')' | b']' | b'}')
}
fn scan_forward(
bytes: &[u8],
start: usize,
opening: u8,
closing: u8,
tree: &Tree,
) -> Option<usize> {
let mut stack: i32 = 1;
let limit = bytes.len().min(start + MAX_SCAN_DISTANCE);
for (i, &ch) in bytes.iter().enumerate().take(limit).skip(start + 1) {
if ch != opening && ch != closing {
continue;
}
if is_in_string_or_comment(tree, i) {
continue;
}
if ch == opening {
stack += 1;
} else {
stack -= 1;
if stack == 0 {
return Some(i);
}
}
}
None
}
fn scan_backward(
bytes: &[u8],
start: usize,
opening: u8,
closing: u8,
tree: &Tree,
) -> Option<usize> {
let mut stack: i32 = 1;
let limit = start.saturating_sub(MAX_SCAN_DISTANCE);
for i in (limit..start).rev() {
let ch = bytes[i];
if ch != opening && ch != closing {
continue;
}
if is_in_string_or_comment(tree, i) {
continue;
}
if ch == closing {
stack += 1;
} else {
stack -= 1;
if stack == 0 {
return Some(i);
}
}
}
None
}
fn is_in_string_or_comment(tree: &Tree, byte_pos: usize) -> bool {
let root = tree.root_node();
let node = root.descendant_for_byte_range(byte_pos, byte_pos);
match node {
Some(n) => is_string_or_comment_kind(n.kind()),
None => false,
}
}
fn is_string_or_comment_kind(kind: &str) -> bool {
matches!(
kind,
"string" | "string_content" | "escape_sequence" | "comment"
)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::r_parser::parse_r;
fn match_at(input: &str, cursor: usize) -> Option<BracketMatch> {
let tree = parse_r(input)?;
find_matching_bracket(input, cursor, &tree)
}
#[test]
fn test_simple_parens() {
let m = match_at("f(x)", 1).unwrap();
assert_eq!(m.cursor_bracket, 1);
assert_eq!(m.matching_bracket, 3);
}
#[test]
fn test_simple_parens_from_close() {
let m = match_at("f(x)", 3).unwrap();
assert_eq!(m.cursor_bracket, 3);
assert_eq!(m.matching_bracket, 1);
}
#[test]
fn test_cursor_after_closing_paren() {
let m = match_at("f(x)", 4).unwrap();
assert_eq!(m.cursor_bracket, 3);
assert_eq!(m.matching_bracket, 1);
}
#[test]
fn test_nested_parens() {
let m = match_at("f(g(x))", 1).unwrap();
assert_eq!(m.cursor_bracket, 1);
assert_eq!(m.matching_bracket, 6);
let m = match_at("f(g(x))", 3).unwrap();
assert_eq!(m.cursor_bracket, 3);
assert_eq!(m.matching_bracket, 5);
}
#[test]
fn test_brackets() {
let m = match_at("x[1]", 1).unwrap();
assert_eq!(m.cursor_bracket, 1);
assert_eq!(m.matching_bracket, 3);
}
#[test]
fn test_braces() {
let m = match_at("{ x }", 0).unwrap();
assert_eq!(m.cursor_bracket, 0);
assert_eq!(m.matching_bracket, 4);
}
#[test]
fn test_unmatched_bracket() {
assert!(match_at("f(x", 1).is_none());
}
#[test]
fn test_bracket_in_string_skipped() {
let input = r#"paste("(", x)"#;
let m = match_at(input, 5).unwrap();
assert_eq!(m.cursor_bracket, 5);
assert_eq!(m.matching_bracket, 12);
}
#[test]
fn test_bracket_in_comment_skipped() {
assert!(match_at("# f(x)", 3).is_none());
}
#[test]
fn test_cursor_on_non_bracket() {
assert!(match_at("hello", 2).is_none());
}
#[test]
fn test_empty_input() {
assert!(match_at("", 0).is_none());
}
#[test]
fn test_mixed_bracket_types() {
let m = match_at("f(x[1])", 1).unwrap();
assert_eq!(m.cursor_bracket, 1);
assert_eq!(m.matching_bracket, 6);
let m = match_at("f(x[1])", 3).unwrap();
assert_eq!(m.cursor_bracket, 3);
assert_eq!(m.matching_bracket, 5);
}
#[test]
fn test_cursor_after_closing_bracket() {
let m = match_at("x[1]", 4).unwrap();
assert_eq!(m.cursor_bracket, 3);
assert_eq!(m.matching_bracket, 1);
}
#[test]
fn test_cursor_after_closing_brace() {
let m = match_at("{ x }", 5).unwrap();
assert_eq!(m.cursor_bracket, 4);
assert_eq!(m.matching_bracket, 0);
}
#[test]
fn test_multiline_braces() {
let input = "if (x) {\n y\n}";
let m = match_at(input, 7).unwrap();
assert_eq!(m.cursor_bracket, 7);
assert_eq!(m.matching_bracket, 13);
let m = match_at(input, 13).unwrap();
assert_eq!(m.cursor_bracket, 13);
assert_eq!(m.matching_bracket, 7);
}
#[test]
fn test_multiline_bracket_in_string() {
let input = "paste(\n \"(\",\n x\n)";
let m = match_at(input, 5).unwrap();
assert_eq!(m.cursor_bracket, 5);
assert_eq!(m.matching_bracket, 18);
}
#[test]
fn test_multiline_bracket_in_comment() {
let input = "f(\n# )\nx)";
let m = match_at(input, 1).unwrap();
assert_eq!(m.cursor_bracket, 1);
assert_eq!(m.matching_bracket, 8);
}
}