use crate::r_parser::{is_atomic_node, parse_r};
pub fn token_left_position(buffer: &str, cursor_pos: usize) -> usize {
if cursor_pos == 0 || buffer.is_empty() {
return 0;
}
if let Some(tree) = parse_r(buffer) {
let root = tree.root_node();
if let Some(pos) = find_token_start_before(cursor_pos, &root, buffer.as_bytes()) {
return pos;
}
}
word_left_position_fallback(buffer, cursor_pos)
}
pub fn token_right_position(buffer: &str, cursor_pos: usize) -> usize {
if cursor_pos >= buffer.len() || buffer.is_empty() {
return buffer.len();
}
if let Some(tree) = parse_r(buffer) {
let root = tree.root_node();
if let Some(pos) = find_token_end_after(cursor_pos, &root, buffer.as_bytes()) {
return pos;
}
}
word_right_position_fallback(buffer, cursor_pos)
}
fn find_token_start_before(
cursor_pos: usize,
root: &tree_sitter::Node,
source: &[u8],
) -> Option<usize> {
let tokens = collect_tokens(root, source);
let effective_pos = skip_whitespace_left(source, cursor_pos);
let mut best_token: Option<(usize, usize)> = None;
for (start, end) in &tokens {
if *end <= effective_pos {
best_token = Some((*start, *end));
} else if *start < effective_pos && effective_pos < *end {
return Some(*start);
}
}
if let Some((start, _)) = best_token {
return Some(start);
}
None
}
fn find_token_end_after(
cursor_pos: usize,
root: &tree_sitter::Node,
source: &[u8],
) -> Option<usize> {
let tokens = collect_tokens(root, source);
let effective_pos = skip_whitespace_right(source, cursor_pos);
for (start, end) in &tokens {
if *start >= effective_pos {
return Some(*end);
} else if *start < effective_pos && effective_pos < *end {
return Some(*end);
}
}
None
}
fn collect_tokens(root: &tree_sitter::Node, source: &[u8]) -> Vec<(usize, usize)> {
let mut tokens = Vec::new();
collect_tokens_recursive(root, source, &mut tokens);
tokens.sort_by_key(|(start, _)| *start);
tokens
}
fn collect_tokens_recursive(
node: &tree_sitter::Node,
source: &[u8],
tokens: &mut Vec<(usize, usize)>,
) {
let kind = node.kind();
if is_atomic_node(kind) {
let start = node.start_byte();
let end = node.end_byte();
if start < end {
tokens.push((start, end));
}
return;
}
if node.child_count() == 0 {
let start = node.start_byte();
let end = node.end_byte();
if start < end {
if let Ok(text) = std::str::from_utf8(&source[start..end])
&& !text.chars().all(char::is_whitespace)
{
tokens.push((start, end));
}
}
return;
}
for i in 0..node.child_count() {
if let Some(child) = node.child(i) {
collect_tokens_recursive(&child, source, tokens);
}
}
}
fn skip_whitespace_left(source: &[u8], pos: usize) -> usize {
let mut p = pos;
while p > 0 {
let prev_char_start = find_char_start(source, p - 1);
if let Ok(s) = std::str::from_utf8(&source[prev_char_start..p])
&& let Some(c) = s.chars().next()
&& c.is_whitespace()
{
p = prev_char_start;
continue;
}
break;
}
p
}
fn skip_whitespace_right(source: &[u8], pos: usize) -> usize {
let mut p = pos;
while p < source.len() {
if let Ok(s) = std::str::from_utf8(&source[p..])
&& let Some(c) = s.chars().next()
&& c.is_whitespace()
{
p += c.len_utf8();
continue;
}
break;
}
p
}
fn find_char_start(source: &[u8], pos: usize) -> usize {
let mut p = pos;
while p > 0 && (source[p] & 0xC0) == 0x80 {
p -= 1;
}
p
}
fn word_left_position_fallback(buffer: &str, cursor_pos: usize) -> usize {
let before = &buffer[..cursor_pos];
let trimmed = before.trim_end();
if trimmed.is_empty() {
return 0;
}
if let Some(last_space) = trimmed.rfind(char::is_whitespace) {
last_space + 1
} else {
0
}
}
fn word_right_position_fallback(buffer: &str, cursor_pos: usize) -> usize {
let after = &buffer[cursor_pos..];
let trimmed = after.trim_start();
if trimmed.is_empty() {
return buffer.len();
}
let whitespace_len = after.len() - trimmed.len();
if let Some(first_space) = trimmed.find(char::is_whitespace) {
cursor_pos + whitespace_len + first_space
} else {
buffer.len()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_token_right_pipe_operator() {
let buffer = "x |> filter()";
let pos = token_right_position(buffer, 2);
assert_eq!(pos, 4); }
#[test]
fn test_token_right_assignment() {
let buffer = "x <- 42";
let pos = token_right_position(buffer, 2);
assert_eq!(pos, 4); }
#[test]
fn test_token_left_pipe_operator() {
let buffer = "x |> filter()";
let pos = token_left_position(buffer, 4);
assert_eq!(pos, 2); }
#[test]
fn test_token_left_assignment() {
let buffer = "x <- 42";
let pos = token_left_position(buffer, 4);
assert_eq!(pos, 2); }
#[test]
fn test_token_right_identifier() {
let buffer = "filter(data)";
let pos = token_right_position(buffer, 0);
assert_eq!(pos, 6); }
#[test]
fn test_token_left_identifier() {
let buffer = "filter(data)";
let pos = token_left_position(buffer, 6);
assert_eq!(pos, 0); }
#[test]
fn test_token_right_at_end() {
let buffer = "x <- 1";
let pos = token_right_position(buffer, buffer.len());
assert_eq!(pos, buffer.len());
}
#[test]
fn test_token_left_at_start() {
let buffer = "x <- 1";
let pos = token_left_position(buffer, 0);
assert_eq!(pos, 0);
}
#[test]
fn test_token_right_double_arrow() {
let buffer = "x <<- 42";
let pos = token_right_position(buffer, 2);
assert_eq!(pos, 5); }
#[test]
fn test_token_right_magrittr_pipe() {
let buffer = "x %>% y";
let pos = token_right_position(buffer, 2);
assert_eq!(pos, 5); }
#[test]
fn test_token_right_comparison() {
let buffer = "x >= 5";
let pos = token_right_position(buffer, 2);
assert_eq!(pos, 4); }
#[test]
fn test_token_right_logical_and() {
let buffer = "x && y";
let pos = token_right_position(buffer, 2);
assert_eq!(pos, 4); }
#[test]
fn test_token_right_namespace() {
let buffer = "dplyr::filter";
let pos = token_right_position(buffer, 5);
assert_eq!(pos, 7); }
#[test]
fn test_token_right_string() {
let buffer = r#"x <- "hello world""#;
let pos = token_right_position(buffer, 5);
assert_eq!(pos, 18); }
#[test]
fn test_empty_buffer() {
let buffer = "";
assert_eq!(token_left_position(buffer, 0), 0);
assert_eq!(token_right_position(buffer, 0), 0);
}
#[test]
fn test_whitespace_only() {
let buffer = " ";
assert_eq!(token_left_position(buffer, 3), 0);
assert_eq!(token_right_position(buffer, 0), 3);
}
#[test]
fn test_skip_whitespace_then_token() {
let buffer = " x <- 1";
let pos = token_right_position(buffer, 0);
assert_eq!(pos, 3); }
}