use crate::transform::utils::is_inside_function_body;
use crate::{Language, Result, SkimError, TransformConfig};
use tree_sitter::{Node, Tree};
pub(crate) const MAX_AST_DEPTH: usize = 500;
pub(crate) const MAX_AST_NODES: usize = 100_000;
pub(crate) struct CommentWalkContext<'a> {
pub(crate) ranges: &'a mut Vec<(usize, usize)>,
pub(crate) node_count: &'a mut usize,
}
pub(crate) fn transform_minimal(
source: &str,
tree: &Tree,
language: Language,
_config: &TransformConfig,
) -> Result<String> {
let mut ranges_to_remove: Vec<(usize, usize)> = Vec::new();
let mut node_count: usize = 0;
let mut ctx = CommentWalkContext {
ranges: &mut ranges_to_remove,
node_count: &mut node_count,
};
collect_removable_comments(tree.root_node(), source, language, &mut ctx, 0)?;
let mut final_ranges: Vec<(usize, usize)> = ctx
.ranges
.iter()
.map(|&(start, end)| adjust_range_for_line_removal(source, start, end))
.collect();
final_ranges.sort_unstable_by_key(|&(start, _)| start);
final_ranges.dedup();
let after_removal = remove_ranges(source, &final_ranges)?;
let normalized = trim_and_normalize(&after_removal);
Ok(normalized)
}
pub(crate) fn collect_removable_comments(
node: Node,
source: &str,
language: Language,
ctx: &mut CommentWalkContext<'_>,
depth: usize,
) -> Result<()> {
if depth > MAX_AST_DEPTH {
return Err(SkimError::ParseError(format!(
"Maximum AST depth exceeded: {} (possible malicious input)",
MAX_AST_DEPTH
)));
}
*ctx.node_count += 1;
if *ctx.node_count > MAX_AST_NODES {
return Err(SkimError::ParseError(format!(
"Too many AST nodes: {} (max: {}). Possible malicious input.",
*ctx.node_count, MAX_AST_NODES
)));
}
if is_removable_comment(node, source, language) {
ctx.ranges.push((node.start_byte(), node.end_byte()));
}
let mut cursor = node.walk();
for child in node.children(&mut cursor) {
collect_removable_comments(child, source, language, ctx, depth + 1)?;
}
Ok(())
}
fn is_shebang(node: Node, source: &str) -> bool {
if node.start_byte() != 0 {
return false;
}
node.utf8_text(source.as_bytes())
.map(|text| text.starts_with("#!"))
.unwrap_or(false)
}
pub(crate) fn is_comment_node(kind: &str, language: Language) -> bool {
match language {
Language::TypeScript
| Language::JavaScript
| Language::Python
| Language::Go
| Language::C
| Language::Cpp
| Language::CSharp
| Language::Ruby
| Language::Sql => kind == "comment",
Language::Rust | Language::Java | Language::Kotlin => {
kind == "line_comment" || kind == "block_comment"
}
Language::Swift => kind == "comment" || kind == "multiline_comment",
Language::Markdown | Language::Json | Language::Yaml | Language::Toml => false,
}
}
pub(crate) fn is_removable_comment(node: Node, source: &str, language: Language) -> bool {
if !is_comment_node(node.kind(), language) {
return false;
}
let should_preserve = is_shebang(node, source)
|| is_inside_function_body(node, language)
|| is_doc_comment(node, source, language);
!should_preserve
}
fn is_doc_comment(node: Node, source: &str, language: Language) -> bool {
let text = match node.utf8_text(source.as_bytes()) {
Ok(t) => t,
Err(_) => return false,
};
match language {
Language::TypeScript | Language::JavaScript => {
text.starts_with("/**")
}
Language::Python => {
false
}
Language::Rust => {
text.starts_with("///")
|| text.starts_with("//!")
|| text.starts_with("/**")
|| text.starts_with("/*!")
}
Language::Go => {
is_go_doc_comment(node, source)
}
Language::Java => {
text.starts_with("/**")
}
Language::C | Language::Cpp => {
text.starts_with("/**") || text.starts_with("///")
}
Language::CSharp => {
text.starts_with("///") || text.starts_with("/**")
}
Language::Ruby => {
false
}
Language::Kotlin => {
text.starts_with("/**")
}
Language::Swift => {
text.starts_with("///") || text.starts_with("/**")
}
Language::Sql => {
false
}
Language::Markdown | Language::Json | Language::Yaml | Language::Toml => false,
}
}
fn is_go_doc_comment(node: Node, source: &str) -> bool {
let mut current_end = node.end_byte();
let mut sibling = node.next_named_sibling();
while let Some(sib) = sibling {
let sib_start = sib.start_byte();
if current_end <= sib_start && sib_start <= source.len() {
let between = &source[current_end..sib_start];
let newline_count = between.chars().filter(|&c| c == '\n').count();
if newline_count > 1 {
return false;
}
}
if is_comment_node(sib.kind(), Language::Go) {
current_end = sib.end_byte();
sibling = sib.next_named_sibling();
continue;
}
return is_go_declaration(sib.kind());
}
false
}
fn is_go_declaration(kind: &str) -> bool {
matches!(
kind,
"function_declaration"
| "method_declaration"
| "type_declaration"
| "var_declaration"
| "const_declaration"
| "type_spec"
)
}
pub(crate) fn adjust_range_for_line_removal(
source: &str,
start: usize,
end: usize,
) -> (usize, usize) {
let line_start = source[..start].rfind('\n').map(|pos| pos + 1).unwrap_or(0);
let line_end = source[end..]
.find('\n')
.map(|pos| end + pos + 1)
.unwrap_or(source.len());
let before_range = &source[line_start..start];
let after_range = if end < line_end {
let after_end = if line_end > 0 && source.as_bytes().get(line_end - 1) == Some(&b'\n') {
line_end - 1
} else {
line_end
};
&source[end..after_end]
} else {
""
};
let only_whitespace_before = before_range.chars().all(|c| c.is_whitespace());
let only_whitespace_after = after_range.chars().all(|c| c.is_whitespace());
if only_whitespace_before && only_whitespace_after {
(line_start, line_end)
} else if only_whitespace_after {
let trimmed_start = source[line_start..start].trim_end().len() + line_start;
(trimmed_start, end)
} else {
(start, end)
}
}
pub(crate) fn remove_ranges(source: &str, ranges: &[(usize, usize)]) -> Result<String> {
if ranges.is_empty() {
return Ok(source.to_string());
}
let mut result = String::with_capacity(source.len());
let mut last_pos = 0;
for &(start, end) in ranges {
if end < start {
return Err(SkimError::ParseError(format!(
"Invalid range: start={} end={}",
start, end
)));
}
if end > source.len() {
return Err(SkimError::ParseError(format!(
"Range exceeds source length: end={} len={}",
end,
source.len()
)));
}
if start < last_pos {
last_pos = last_pos.max(end);
continue;
}
if !source.is_char_boundary(start) || !source.is_char_boundary(end) {
return Err(SkimError::ParseError(format!(
"Invalid UTF-8 boundary at range [{}, {})",
start, end
)));
}
result.push_str(&source[last_pos..start]);
last_pos = end;
}
if !source.is_char_boundary(last_pos) {
return Err(SkimError::ParseError(format!(
"Invalid UTF-8 boundary at position {}",
last_pos
)));
}
result.push_str(&source[last_pos..]);
Ok(result)
}
pub(crate) fn trim_and_normalize(source: &str) -> String {
let mut result = String::with_capacity(source.len());
let mut consecutive_blanks: usize = 0;
for line in source.lines() {
let trimmed = line.trim_end();
if trimmed.is_empty() {
consecutive_blanks += 1;
if consecutive_blanks > 2 {
continue;
}
} else {
consecutive_blanks = 0;
}
if !result.is_empty() {
result.push('\n');
}
result.push_str(trimmed);
}
if source.ends_with('\n') {
result.push('\n');
}
result
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_trim_and_normalize_preserves_two_blanks() {
let input = "a\n\n\nb\n";
let result = trim_and_normalize(input);
assert_eq!(result, "a\n\n\nb\n");
}
#[test]
fn test_trim_and_normalize_reduces_four_blanks_to_two() {
let input = "a\n\n\n\n\nb\n";
let result = trim_and_normalize(input);
assert_eq!(result, "a\n\n\nb\n");
}
#[test]
fn test_trim_and_normalize_no_change_needed() {
let input = "a\n\nb\n";
let result = trim_and_normalize(input);
assert_eq!(result, "a\n\nb\n");
}
#[test]
fn test_trim_and_normalize_trims_trailing_whitespace() {
let input = "hello \nworld \n";
let result = trim_and_normalize(input);
assert_eq!(result, "hello\nworld\n");
}
#[test]
fn test_trim_and_normalize_combined() {
let input = "hello \n\n\n\n\nworld \n";
let result = trim_and_normalize(input);
assert_eq!(result, "hello\n\n\nworld\n");
}
#[test]
fn test_adjust_range_full_line_comment() {
let source = "code\n// comment\nmore code\n";
let (start, end) = adjust_range_for_line_removal(source, 5, 15);
assert_eq!(start, 5);
assert_eq!(end, 16); }
#[test]
fn test_adjust_range_trailing_comment() {
let source = "let x = 1; // trailing\nmore code\n";
let (start, end) = adjust_range_for_line_removal(source, 11, 22);
assert!(start <= 11, "start should be at or before comment start");
assert_eq!(end, 22);
let remaining = format!("{}{}", &source[..start], &source[end..]);
assert!(
remaining.starts_with("let x = 1;"),
"should preserve code before trailing comment, got: {:?}",
remaining
);
}
#[test]
fn test_adjust_range_inline_comment_with_code_after() {
let source = "/* comment */ let x = 1;\n";
let (start, end) = adjust_range_for_line_removal(source, 0, 13);
assert_eq!(start, 0);
assert_eq!(end, 13);
}
#[test]
fn test_remove_ranges_end_before_start() {
let source = "hello world";
let ranges = vec![(5, 3)]; let result = remove_ranges(source, &ranges);
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("Invalid range"),
"Expected 'Invalid range' error, got: {}",
err_msg
);
}
#[test]
fn test_remove_ranges_end_exceeds_source_length() {
let source = "hello";
let ranges = vec![(0, 100)]; let result = remove_ranges(source, &ranges);
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("Range exceeds source length"),
"Expected 'Range exceeds source length' error, got: {}",
err_msg
);
}
#[test]
fn test_remove_ranges_non_char_boundary() {
let source = "a\u{20AC}b"; let ranges = vec![(2, 4)];
let result = remove_ranges(source, &ranges);
assert!(result.is_err());
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("Invalid UTF-8 boundary"),
"Expected 'Invalid UTF-8 boundary' error, got: {}",
err_msg
);
}
#[test]
fn test_remove_ranges_empty_ranges() {
let source = "hello world";
let ranges = vec![];
let result = remove_ranges(source, &ranges).unwrap();
assert_eq!(result, "hello world");
}
#[test]
fn test_remove_ranges_valid_removal() {
let source = "hello beautiful world";
let ranges = vec![(5, 15)]; let result = remove_ranges(source, &ranges).unwrap();
assert_eq!(result, "hello world");
}
#[test]
fn test_max_ast_nodes_limit() {
let mut source = String::new();
for i in 0..4500 {
source.push_str("x = ");
for j in 0..20 {
if j > 0 {
source.push_str(" + ");
}
source.push_str(&(i * 20 + j).to_string());
}
source.push('\n');
}
let mut parser = crate::Parser::new(Language::Python).unwrap();
let tree = parser.parse(&source).unwrap();
let config = TransformConfig::default();
let result = transform_minimal(&source, &tree, Language::Python, &config);
assert!(
result.is_err(),
"Expected error when exceeding MAX_AST_NODES"
);
let err_msg = result.unwrap_err().to_string();
assert!(
err_msg.contains("Too many AST nodes"),
"Expected 'Too many AST nodes' error, got: {}",
err_msg
);
}
}