use crate::config::RColorConfig;
use crate::r_parser::{is_atomic_node, parse_r};
use nu_ansi_term::{Color, Style};
use once_cell::sync::Lazy;
use reedline::{Highlighter, StyledText};
use std::collections::HashSet;
use tree_sitter::Node;
use super::bracket_match::find_matching_bracket;
use super::r_regex::TokenType;
static KEYWORDS: Lazy<HashSet<&'static str>> = Lazy::new(|| {
[
"if", "else", "for", "while", "repeat", "in", "next", "break", "return", "function",
]
.into_iter()
.collect()
});
static CONSTANTS: Lazy<HashSet<&'static str>> = Lazy::new(|| {
[
"TRUE",
"FALSE",
"NULL",
"Inf",
"NaN",
"NA",
"NA_integer_",
"NA_real_",
"NA_complex_",
"NA_character_",
]
.into_iter()
.collect()
});
#[derive(Debug, Clone)]
pub struct Token {
pub start: usize,
pub end: usize,
pub token_type: TokenType,
}
fn node_to_token_type(node: &Node, source: &[u8]) -> TokenType {
match node.kind() {
"integer" | "float" | "complex" => TokenType::Number,
"string" | "string_content" => TokenType::String,
"escape_sequence" => TokenType::String,
"comment" => TokenType::Comment,
"true" | "false" => TokenType::Constant,
"null" | "inf" | "nan" | "na" => TokenType::Constant,
"dots" | "dot_dot_i" => TokenType::Constant,
"function" | "if" | "else" | "for" | "while" | "repeat" | "in" | "next" | "break"
| "return" => TokenType::Keyword,
"?" | ":=" | "=" | "<-" | "<<-" | "->" | "->>" | "~" | "|>" | "||" | "|" | "&&" | "&"
| "<" | "<=" | ">" | ">=" | "==" | "!=" | "+" | "-" | "*" | "/" | "::" | ":::" | "**"
| "^" | "$" | "@" | ":" | "!" | r#"\"# | "special" => TokenType::Operator,
"(" | ")" | "{" | "}" | "[" | "]" | "[[" | "]]" => TokenType::Punctuation,
"comma" | ";" => TokenType::Punctuation,
"identifier" => {
let text = node.utf8_text(source).unwrap_or("");
if KEYWORDS.contains(text) {
TokenType::Keyword
} else if CONSTANTS.contains(text) {
TokenType::Constant
} else {
TokenType::Identifier
}
}
_ => TokenType::Other,
}
}
fn visit_node(cursor: &mut tree_sitter::TreeCursor, source: &[u8], tokens: &mut Vec<Token>) {
let node = cursor.node();
let kind = node.kind();
if is_atomic_node(kind) {
tokens.push(Token {
start: node.start_byte(),
end: node.end_byte(),
token_type: node_to_token_type(&node, source),
});
return;
}
if node.child_count() == 0 {
let token_type = node_to_token_type(&node, source);
if token_type != TokenType::Other || node.start_byte() < node.end_byte() {
tokens.push(Token {
start: node.start_byte(),
end: node.end_byte(),
token_type,
});
}
} else {
if cursor.goto_first_child() {
loop {
visit_node(cursor, source, tokens);
if !cursor.goto_next_sibling() {
break;
}
}
cursor.goto_parent();
}
}
}
fn fill_gaps(tokens: &[Token], total_len: usize) -> Vec<Token> {
let mut result = Vec::new();
let mut pos = 0;
for token in tokens.iter() {
if token.start > pos {
result.push(Token {
start: pos,
end: token.start,
token_type: TokenType::Whitespace,
});
}
result.push(token.clone());
pos = token.end;
}
if pos < total_len {
result.push(Token {
start: pos,
end: total_len,
token_type: TokenType::Whitespace,
});
}
result
}
pub fn tokenize_r(source: &str) -> Vec<Token> {
let tree = match parse_r(source) {
Some(t) => t,
None => {
return vec![Token {
start: 0,
end: source.len(),
token_type: TokenType::Other,
}];
}
};
let source_bytes = source.as_bytes();
let mut tokens = Vec::new();
let mut cursor = tree.walk();
visit_node(&mut cursor, source_bytes, &mut tokens);
tokens.sort_by_key(|t| t.start);
fill_gaps(&tokens, source.len())
}
pub struct RTreeSitterHighlighter {
config: RColorConfig,
highlight_matching_bracket: bool,
}
impl RTreeSitterHighlighter {
pub fn new(config: RColorConfig, highlight_matching_bracket: bool) -> Self {
RTreeSitterHighlighter {
config,
highlight_matching_bracket,
}
}
}
impl Default for RTreeSitterHighlighter {
fn default() -> Self {
Self::new(RColorConfig::default(), true)
}
}
impl Highlighter for RTreeSitterHighlighter {
fn highlight(&self, line: &str, cursor: usize) -> StyledText {
let mut styled = StyledText::new();
let tree = parse_r(line);
if let Some(ref tree) = tree {
let source = line.as_bytes();
let mut tokens = Vec::new();
let mut tree_cursor = tree.walk();
visit_node(&mut tree_cursor, source, &mut tokens);
tokens.sort_by_key(|t| t.start);
let tokens = fill_gaps(&tokens, source.len());
for token in tokens {
if token.start < line.len() && token.end <= line.len() {
let text = &line[token.start..token.end];
let style = token.token_type.style(&self.config);
styled.push((style, text.to_string()));
}
}
} else {
styled.push((Style::new(), line.to_string()));
}
if styled.buffer.is_empty() {
styled.push((Style::new(), String::new()));
}
if self.highlight_matching_bracket
&& self.config.matching_bracket != Color::Default
&& let Some(ref tree) = tree
&& let Some(bracket) = find_matching_bracket(line, cursor, tree)
{
let bg = self.config.matching_bracket;
apply_bracket_highlight(&mut styled, bracket.cursor_bracket, bg);
apply_bracket_highlight(&mut styled, bracket.matching_bracket, bg);
}
styled
}
}
fn apply_bracket_highlight(styled: &mut StyledText, pos: usize, highlight_bg: Color) {
let existing_style = {
let mut offset = 0;
let mut found = None;
for (style, text) in &styled.buffer {
let end = offset + text.len();
if pos >= offset && pos < end {
found = Some(*style);
break;
}
offset = end;
}
found.unwrap_or_default()
};
let merged = Style {
background: Some(highlight_bg),
..existing_style
};
styled.style_range(pos, pos + 1, merged);
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::RColorConfig;
use nu_ansi_term::Color;
fn get_token_types(input: &str) -> Vec<(String, TokenType)> {
tokenize_r(input)
.into_iter()
.filter(|t| t.token_type != TokenType::Whitespace)
.map(|t| {
let text = input[t.start..t.end].to_string();
(text, t.token_type)
})
.collect()
}
#[test]
fn test_comment() {
let tokens = get_token_types("# this is a comment");
assert_eq!(tokens.len(), 1);
assert_eq!(tokens[0].1, TokenType::Comment);
}
#[test]
fn test_string() {
let tokens = get_token_types(r#""hello world""#);
assert_eq!(tokens.len(), 1);
assert_eq!(tokens[0].1, TokenType::String);
}
#[test]
fn test_raw_string() {
let tokens = get_token_types(r#"r"(hello "world")""#);
assert_eq!(tokens.len(), 1);
assert_eq!(tokens[0].1, TokenType::String);
}
#[test]
fn test_numbers() {
let cases = vec!["42", "3.14", "1e-5", "0xFF", "1L", "2i"];
for case in cases {
let tokens = get_token_types(case);
assert_eq!(tokens[0].1, TokenType::Number, "Failed for: {}", case);
}
}
#[test]
fn test_keywords() {
let tokens = get_token_types("if (TRUE) else FALSE");
let keywords: Vec<_> = tokens
.iter()
.filter(|(_, t)| *t == TokenType::Keyword)
.collect();
assert_eq!(keywords.len(), 2); }
#[test]
fn test_constants() {
let tokens = get_token_types("TRUE FALSE NULL NA Inf NaN");
let constants: Vec<_> = tokens
.iter()
.filter(|(_, t)| *t == TokenType::Constant)
.collect();
assert_eq!(constants.len(), 6);
}
#[test]
fn test_operators() {
let tokens = get_token_types("x <- 1 + 2");
let operators: Vec<_> = tokens
.iter()
.filter(|(_, t)| *t == TokenType::Operator)
.collect();
assert_eq!(operators.len(), 2); }
#[test]
fn test_assignment() {
let tokens = get_token_types("x <- 42");
assert_eq!(tokens[0], ("x".to_string(), TokenType::Identifier));
assert_eq!(tokens[1], ("<-".to_string(), TokenType::Operator));
assert_eq!(tokens[2], ("42".to_string(), TokenType::Number));
}
#[test]
fn test_highlight_preserves_text() {
let highlighter = RTreeSitterHighlighter::default();
let input = "x <- c(1, 2, 3)";
let styled = highlighter.highlight(input, 0);
assert_eq!(styled.raw_string(), input);
}
#[test]
fn test_highlight_empty() {
let highlighter = RTreeSitterHighlighter::default();
let styled = highlighter.highlight("", 0);
assert_eq!(styled.raw_string(), "");
}
#[test]
fn test_custom_keyword_color() {
let config = RColorConfig {
keyword: Color::Red,
..Default::default()
};
let highlighter = RTreeSitterHighlighter::new(config, false);
let styled = highlighter.highlight("if", 0);
let if_segment = styled
.buffer
.iter()
.find(|(_, text)| text == "if")
.expect("Should find 'if' segment");
assert_eq!(
if_segment.0.foreground,
Some(Color::Red),
"Keyword 'if' should be styled with Red"
);
}
#[test]
fn test_custom_string_color() {
let config = RColorConfig {
string: Color::Yellow,
..Default::default()
};
let highlighter = RTreeSitterHighlighter::new(config, false);
let styled = highlighter.highlight(r#""hello""#, 0);
let string_segment = styled
.buffer
.iter()
.find(|(_, text)| text.contains("hello"))
.expect("Should find string segment");
assert_eq!(
string_segment.0.foreground,
Some(Color::Yellow),
"String should be styled with Yellow"
);
}
#[test]
fn test_custom_number_color() {
let config = RColorConfig {
number: Color::Blue,
..Default::default()
};
let highlighter = RTreeSitterHighlighter::new(config, false);
let styled = highlighter.highlight("42", 0);
let number_segment = styled
.buffer
.iter()
.find(|(_, text)| text == "42")
.expect("Should find number segment");
assert_eq!(
number_segment.0.foreground,
Some(Color::Blue),
"Number should be styled with Blue"
);
}
#[test]
fn test_custom_comment_color() {
let config = RColorConfig {
comment: Color::Cyan,
..Default::default()
};
let highlighter = RTreeSitterHighlighter::new(config, false);
let styled = highlighter.highlight("# comment", 0);
let comment_segment = styled
.buffer
.iter()
.find(|(_, text)| text.contains("comment"))
.expect("Should find comment segment");
assert_eq!(
comment_segment.0.foreground,
Some(Color::Cyan),
"Comment should be styled with Cyan"
);
}
#[test]
fn test_custom_constant_color() {
let config = RColorConfig {
constant: Color::Magenta,
..Default::default()
};
let highlighter = RTreeSitterHighlighter::new(config, false);
let styled = highlighter.highlight("TRUE", 0);
let constant_segment = styled
.buffer
.iter()
.find(|(_, text)| text == "TRUE")
.expect("Should find constant segment");
assert_eq!(
constant_segment.0.foreground,
Some(Color::Magenta),
"Constant TRUE should be styled with Magenta"
);
}
#[test]
fn test_custom_operator_color() {
let config = RColorConfig {
operator: Color::Green,
..Default::default()
};
let highlighter = RTreeSitterHighlighter::new(config, false);
let styled = highlighter.highlight("x <- 1", 0);
let operator_segment = styled
.buffer
.iter()
.find(|(_, text)| text == "<-")
.expect("Should find operator segment");
assert_eq!(
operator_segment.0.foreground,
Some(Color::Green),
"Operator <- should be styled with Green"
);
}
#[test]
fn test_custom_identifier_color() {
let config = RColorConfig {
identifier: Color::White,
..Default::default()
};
let highlighter = RTreeSitterHighlighter::new(config, false);
let styled = highlighter.highlight("myvar", 0);
let identifier_segment = styled
.buffer
.iter()
.find(|(_, text)| text == "myvar")
.expect("Should find identifier segment");
assert_eq!(
identifier_segment.0.foreground,
Some(Color::White),
"Identifier should be styled with White"
);
}
#[test]
fn test_default_color_no_styling() {
let config = RColorConfig {
identifier: Color::Default,
punctuation: Color::Default,
..Default::default()
};
let highlighter = RTreeSitterHighlighter::new(config, false);
let styled = highlighter.highlight("x()", 0);
let identifier_segment = styled
.buffer
.iter()
.find(|(_, text)| text == "x")
.expect("Should find identifier segment");
assert_eq!(
identifier_segment.0.foreground, None,
"Color::Default should result in no foreground color"
);
let paren_segment = styled
.buffer
.iter()
.find(|(_, text)| text == "(")
.expect("Should find parenthesis segment");
assert_eq!(
paren_segment.0.foreground, None,
"Punctuation with Color::Default should have no foreground color"
);
}
#[test]
fn test_multiple_custom_colors() {
let config = RColorConfig {
keyword: Color::Red,
constant: Color::Blue,
operator: Color::Green,
number: Color::Yellow,
identifier: Color::White,
comment: Color::DarkGray,
string: Color::Cyan,
punctuation: Color::Magenta,
matching_bracket: Color::LightYellow,
};
let highlighter = RTreeSitterHighlighter::new(config, false);
let styled = highlighter.highlight("if (x <- 1) TRUE", 0);
let if_seg = styled.buffer.iter().find(|(_, t)| t == "if");
assert_eq!(if_seg.map(|(s, _)| s.foreground), Some(Some(Color::Red)));
let op_seg = styled.buffer.iter().find(|(_, t)| t == "<-");
assert_eq!(op_seg.map(|(s, _)| s.foreground), Some(Some(Color::Green)));
let num_seg = styled.buffer.iter().find(|(_, t)| t == "1");
assert_eq!(
num_seg.map(|(s, _)| s.foreground),
Some(Some(Color::Yellow))
);
let const_seg = styled.buffer.iter().find(|(_, t)| t == "TRUE");
assert_eq!(
const_seg.map(|(s, _)| s.foreground),
Some(Some(Color::Blue))
);
let id_seg = styled.buffer.iter().find(|(_, t)| t == "x");
assert_eq!(id_seg.map(|(s, _)| s.foreground), Some(Some(Color::White)));
let paren_seg = styled.buffer.iter().find(|(_, t)| t == "(");
assert_eq!(
paren_seg.map(|(s, _)| s.foreground),
Some(Some(Color::Magenta))
);
}
#[test]
fn test_bracket_highlight_applied() {
let config = RColorConfig {
punctuation: Color::Magenta,
matching_bracket: Color::LightYellow,
..Default::default()
};
let highlighter = RTreeSitterHighlighter::new(config, true);
let styled = highlighter.highlight("f(x)", 1);
let mut offset = 0;
let mut found = false;
for (style, text) in &styled.buffer {
if offset <= 1 && offset + text.len() > 1 {
found = true;
assert_eq!(
style.background,
Some(Color::LightYellow),
"Opening bracket should have highlight background"
);
assert_eq!(
style.foreground,
Some(Color::Magenta),
"Opening bracket should preserve foreground color"
);
break;
}
offset += text.len();
}
assert!(found, "Should have found a segment containing the bracket");
}
#[test]
fn test_bracket_highlight_disabled() {
let config = RColorConfig {
matching_bracket: Color::LightYellow,
..Default::default()
};
let highlighter = RTreeSitterHighlighter::new(config, false);
let styled = highlighter.highlight("f(x)", 1);
for (style, _) in &styled.buffer {
assert_ne!(
style.background,
Some(Color::LightYellow),
"Bracket highlight should not be applied when disabled"
);
}
}
}