use super::path::{PathCompletionOptions, complete_path};
use reedline::{Span, Suggestion};
use std::cell::RefCell;
use tree_sitter::{Parser, Tree};
thread_local! {
static R_PARSER: RefCell<Parser> = RefCell::new({
let mut parser = Parser::new();
parser
.set_language(&tree_sitter_r::LANGUAGE.into())
.expect("Failed to set tree-sitter-r language");
parser
});
}
#[derive(Debug, Clone, PartialEq)]
pub struct StringContext {
pub content: String,
pub start: usize,
pub quote: char,
}
fn parse_r_code(code: &str) -> Option<Tree> {
R_PARSER.with(|parser| parser.borrow_mut().parse(code.as_bytes(), None))
}
fn find_node_at_position<'a>(tree: &'a Tree, pos: usize) -> Option<tree_sitter::Node<'a>> {
let root = tree.root_node();
let mut cursor = root.walk();
let mut best_node = None;
loop {
let node = cursor.node();
if pos >= node.start_byte() && pos <= node.end_byte() {
best_node = Some(node);
if cursor.goto_first_child() {
loop {
let child = cursor.node();
if pos >= child.start_byte() && pos <= child.end_byte() {
break; }
if !cursor.goto_next_sibling() {
cursor.goto_parent();
return best_node;
}
}
} else {
return best_node;
}
} else {
return best_node;
}
}
}
fn find_string_ancestor<'a>(node: tree_sitter::Node<'a>) -> Option<tree_sitter::Node<'a>> {
let mut current = Some(node);
while let Some(n) = current {
if n.kind() == "string" {
return Some(n);
}
current = n.parent();
}
None
}
fn find_incomplete_string_in_error<'a>(
node: tree_sitter::Node<'a>,
source: &str,
) -> Option<(usize, char)> {
let mut current = Some(node);
while let Some(n) = current {
if n.kind() == "ERROR" {
let start = n.start_byte();
let end = n.end_byte().min(source.len());
let text = &source[start..end];
let mut in_double = false;
let mut in_single = false;
let mut last_double_pos = None;
let mut last_single_pos = None;
let mut skip_next = false;
for (i, c) in text.char_indices() {
if skip_next {
skip_next = false;
continue;
}
match c {
'\\' => {
skip_next = true;
}
'"' if !in_single => {
if in_double {
in_double = false;
last_double_pos = None;
} else {
in_double = true;
last_double_pos = Some(start + i);
}
}
'\'' if !in_double => {
if in_single {
in_single = false;
last_single_pos = None;
} else {
in_single = true;
last_single_pos = Some(start + i);
}
}
_ => {}
}
}
if in_double && let Some(pos) = last_double_pos {
return Some((pos, '"'));
}
if in_single && let Some(pos) = last_single_pos {
return Some((pos, '\''));
}
}
current = n.parent();
}
None
}
pub fn detect_string_context(line: &str, cursor_pos: usize) -> Option<StringContext> {
let tree = parse_r_code(line)?;
let node = find_node_at_position(&tree, cursor_pos)?;
if let Some(string_node) = find_string_ancestor(node) {
let string_start = string_node.start_byte();
let string_end = string_node.end_byte();
if cursor_pos < string_start || cursor_pos > string_end {
return None;
}
let string_text = &line[string_start..string_end.min(line.len())];
let (quote_char, content_start_offset) = if string_text.starts_with(r#"r""#)
|| string_text.starts_with(r#"R""#)
|| string_text.starts_with("r'")
|| string_text.starts_with("R'")
{
let quote = if string_text.contains('"') { '"' } else { '\'' };
let delim_end = string_text.find('(').map(|p| p + 1).unwrap_or(2);
(quote, delim_end)
} else if string_text.starts_with('"') {
('"', 1)
} else if string_text.starts_with('\'') {
('\'', 1)
} else {
return None;
};
let content_start = string_start + content_start_offset;
if cursor_pos < content_start {
return Some(StringContext {
content: String::new(),
start: content_start,
quote: quote_char,
});
}
let content = if cursor_pos <= line.len() && content_start <= cursor_pos {
line[content_start..cursor_pos].to_string()
} else {
String::new()
};
return Some(StringContext {
content,
start: content_start,
quote: quote_char,
});
}
if let Some((quote_pos, quote_char)) = find_incomplete_string_in_error(node, line) {
let content_start = quote_pos + 1;
if cursor_pos <= quote_pos {
return None;
}
let content = if cursor_pos <= line.len() && content_start <= cursor_pos {
line[content_start..cursor_pos].to_string()
} else {
String::new()
};
return Some(StringContext {
content,
start: content_start,
quote: quote_char,
});
}
None
}
pub(super) fn path_to_suggestions(
partial: &str,
pos: usize,
span_start: usize,
options: &PathCompletionOptions,
) -> Vec<Suggestion> {
let cwd = std::env::current_dir().unwrap_or_else(|_| std::path::PathBuf::from("."));
complete_path(partial, &cwd, options)
.into_iter()
.map(|c| Suggestion {
value: c.path,
display_override: None,
description: if c.is_dir {
Some("directory".to_string())
} else {
None
},
extra: None,
span: Span {
start: span_start,
end: pos,
},
append_whitespace: false,
style: None,
match_indices: c.match_indices,
})
.collect()
}
pub fn complete_path_in_string(_line: &str, pos: usize, ctx: &StringContext) -> Vec<Suggestion> {
path_to_suggestions(
&ctx.content,
pos,
ctx.start,
&PathCompletionOptions::default(),
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_detect_string_context_double_quote() {
let ctx = detect_string_context(r#"read.csv("data/"#, 15);
assert!(ctx.is_some());
let ctx = ctx.unwrap();
assert_eq!(ctx.content, "data/");
assert_eq!(ctx.quote, '"');
}
#[test]
fn test_detect_string_context_single_quote() {
let ctx = detect_string_context("source('script", 14);
assert!(ctx.is_some());
let ctx = ctx.unwrap();
assert_eq!(ctx.content, "script");
assert_eq!(ctx.quote, '\'');
}
#[test]
fn test_detect_string_context_not_in_string() {
let ctx = detect_string_context("print(x)", 7);
assert!(ctx.is_none());
let ctx = detect_string_context(r#"read.csv("data.csv")"#, 20);
assert!(ctx.is_none());
}
#[test]
fn test_detect_string_context_with_escaped_quotes() {
let ctx = detect_string_context(r#"paste("hello \"world"#, 20);
assert!(ctx.is_some());
let ctx = ctx.unwrap();
assert_eq!(ctx.content, r#"hello \"world"#);
}
#[test]
fn test_detect_string_context_empty_string() {
let ctx = detect_string_context(r#"read.csv(""#, 10);
assert!(ctx.is_some());
let ctx = ctx.unwrap();
assert_eq!(ctx.content, "");
assert_eq!(ctx.start, 10);
}
#[test]
fn test_detect_string_context_tilde_path() {
let ctx = detect_string_context(r#"setwd("~/"#, 9);
assert!(ctx.is_some());
let ctx = ctx.unwrap();
assert_eq!(ctx.content, "~/");
}
#[test]
fn test_detect_string_context_absolute_path() {
let ctx = detect_string_context(r#"source("/usr/local/lib/"#, 23);
assert!(ctx.is_some());
let ctx = ctx.unwrap();
assert_eq!(ctx.content, "/usr/local/lib/");
}
#[test]
fn test_detect_string_context_complete_string_cursor_inside() {
let ctx = detect_string_context(r#"read.csv("data.csv")"#, 14);
assert!(ctx.is_some());
let ctx = ctx.unwrap();
assert_eq!(ctx.content, "data");
}
#[test]
fn test_detect_string_context_in_comment() {
let ctx = detect_string_context(r#"# "data/"#, 8);
assert!(ctx.is_none(), "Should not detect string inside comment");
}
#[test]
fn test_detect_string_context_raw_string() {
let ctx = detect_string_context(r#"x <- r"(hello)""#, 11);
assert!(ctx.is_some(), "Should detect raw string");
let ctx = ctx.unwrap();
assert_eq!(ctx.content, "hel");
assert_eq!(ctx.start, 8); }
}