use super::query::node_range;
use crate::config::query::AnnotationQuery;
use crate::diff::changes::{ChangeKind, ChangeMap};
use crate::hash::DftHashMap;
use crate::lines::{SourcePosition, SourceRange};
use crate::parse::syntax::{FoldMetadata, Syntax, SyntaxId};
use hashbrown::hash_map::Entry;
use std::collections::BTreeSet;
use streaming_iterator::StreamingIterator as _;
use tree_sitter::{QueryCursor, Tree};
#[derive(Debug)]
pub(crate) struct Fold {
pub(crate) tags: Vec<String>,
pub(crate) range: SourceRange,
pub(crate) syntax_id: SyntaxId,
pub(crate) match_kind: FoldMatch,
pub(crate) placeholder: String,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum FoldMatch {
Matched {
opposite: SyntaxId,
},
Novel,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct Conflict {
pub(crate) line: usize,
pub(crate) kind: String,
pub(crate) sources: (String, String),
}
pub(crate) fn classify(
tree: &Tree,
src: &str,
compiled: Option<&AnnotationQuery>,
) -> Result<DftHashMap<usize, FoldMetadata>, Conflict> {
let mut kinds: DftHashMap<usize, (FoldMetadata, usize)> = DftHashMap::default();
let Some(compiled) = compiled else {
return Ok(DftHashMap::default());
};
let query = &compiled.query;
let mut cursor = QueryCursor::new();
let mut matches = cursor.matches(query, tree.root_node(), src.as_bytes());
while let Some(matched) = matches.next() {
let pattern = &compiled.patterns[matched.pattern_index];
let named = |name: &'static str| {
matched
.captures
.iter()
.filter(move |capture| query.capture_names()[capture.index as usize] == name)
};
let folds: Vec<_> = named("fold").collect();
let Some(first) = folds.iter().min_by_key(|capture| capture.node.start_byte()) else {
continue;
};
let region = match (named("fold.open").last(), named("fold.close").last()) {
(Some(open), Some(close))
if folds.len() == 1
&& open.node.start_byte() >= first.node.start_byte()
&& close.node.end_byte() <= first.node.end_byte()
&& open.node.end_byte() <= close.node.start_byte() =>
{
SourceRange {
start: node_range(open.node).end,
end: node_range(close.node).start,
}
}
(Some(open), None)
if folds.len() == 1 && open.node.end_byte() <= first.node.end_byte() =>
{
SourceRange {
start: node_range(open.node).end,
end: node_range(first.node).end,
}
}
(None, None) => {
let last = folds
.iter()
.max_by_key(|capture| capture.node.end_byte())
.expect("a fold capture");
SourceRange {
start: node_range(first.node).start,
end: node_range(last.node).end,
}
}
_ => continue,
};
if region.start == region.end {
continue;
}
match kinds.entry(first.node.id()) {
Entry::Vacant(entry) => {
let mut tags = pattern.tags.clone();
tags.sort();
tags.dedup();
entry.insert((
FoldMetadata {
tags,
range_override: Some(region),
},
pattern.source,
));
}
Entry::Occupied(mut entry) => {
let (metadata, source) = entry.get_mut();
if metadata.range_override != Some(region) {
let mut sources = [
compiled.sources[*source].clone(),
compiled.sources[pattern.source].clone(),
];
sources.sort();
let [first_source, second_source] = sources;
return Err(Conflict {
line: first.node.start_position().row,
kind: first.node.kind().to_owned(),
sources: (first_source, second_source),
});
}
metadata.tags.extend(pattern.tags.iter().cloned());
metadata.tags.sort();
metadata.tags.dedup();
}
}
}
Ok(kinds
.into_iter()
.map(|(id, (metadata, _))| (id, metadata))
.collect())
}
pub(crate) fn interior_range(
open: &[line_numbers::SingleLineSpan],
close: &[line_numbers::SingleLineSpan],
) -> SourceRange {
let open = open.last().expect("list opening position");
let close = close.first().expect("list closing position");
SourceRange {
start: SourcePosition {
line: open.line,
byte_column: open.end_col as usize,
},
end: SourcePosition {
line: close.line,
byte_column: close.start_col as usize,
},
}
}
fn range(node: &Syntax<'_>, metadata: &FoldMetadata) -> Option<SourceRange> {
let region = match (metadata.range_override, node) {
(Some(region), _) => region,
(
None,
Syntax::List {
open_position,
close_position,
..
},
) => interior_range(open_position, close_position),
(None, Syntax::Atom { position, .. }) => {
let first = position.first().expect("atom start");
let last = position.last().expect("atom end");
SourceRange {
start: SourcePosition {
line: first.line,
byte_column: first.start_col as usize,
},
end: SourcePosition {
line: last.line,
byte_column: last.end_col as usize,
},
}
}
};
if (region.start.line, region.start.byte_column) >= (region.end.line, region.end.byte_column) {
return None;
}
Some(region)
}
pub(crate) fn line_span(fold: &Fold, lines: &[&str]) -> (usize, usize) {
let line_count = lines.len();
let range = &fold.range;
let first = range.start.line.as_usize();
let leading = lines.get(first).is_some_and(|line| {
let column = range.start.byte_column.min(line.len());
!line[..column].trim().is_empty()
});
let start = if leading { first + 1 } else { first };
let last = range.end.line.as_usize();
let trailing = lines.get(last).is_some_and(|line| {
let column = range.end.byte_column.min(line.len());
!line[column..].trim().is_empty()
});
let end = if trailing { last } else { last + 1 };
(
start.min(line_count),
end.min(line_count).max(start.min(line_count)),
)
}
pub(crate) fn nested_spans(spans: &[(usize, usize)]) -> Vec<Option<(usize, usize)>> {
let mut order: Vec<usize> = (0..spans.len()).collect();
order.sort_by_key(|&index| (spans[index].0, std::cmp::Reverse(spans[index].1)));
let mut out: Vec<Option<(usize, usize)>> = spans.iter().map(|&span| Some(span)).collect();
let mut open: Vec<usize> = Vec::new();
let closed = |out: &[Option<(usize, usize)>], top: usize, start: usize| {
out[top].is_none_or(|(_, top_end)| top_end <= start)
};
for index in order {
let (start, end) = spans[index];
while open.last().is_some_and(|&top| closed(&out, top, start)) {
open.pop();
}
for &top in open.iter().rev() {
let Some((top_start, top_end)) = out[top] else {
continue;
};
if top_end >= end {
break;
}
out[top] = (start > top_start).then_some((top_start, start));
}
while open.last().is_some_and(|&top| closed(&out, top, start)) {
open.pop();
}
open.push(index);
}
for span in &mut out {
if span.is_some_and(|(start, end)| end - start < 2) {
*span = None;
}
}
out
}
pub(crate) fn split_lines(spans: impl IntoIterator<Item = (usize, usize)>) -> BTreeSet<usize> {
spans
.into_iter()
.flat_map(|(start, end)| [start, end])
.collect()
}
pub(crate) fn unmatched(nodes: &[&Syntax<'_>], folds: &mut Vec<Fold>) {
for node in nodes {
folds.extend(project(node, None));
if let Syntax::List { children, .. } = node {
unmatched(children, folds);
}
}
}
pub(crate) fn partner<'a>(
node: &Syntax<'_>,
change: ChangeKind<'a>,
change_map: &ChangeMap<'a>,
) -> Option<&'a Syntax<'a>> {
let other = match change {
ChangeKind::Unchanged(other)
| ChangeKind::ReplacedComment(_, other)
| ChangeKind::ReplacedString(_, other) => other,
ChangeKind::IgnoredPunctuation | ChangeKind::Novel => return None,
};
let back = match change_map
.get(other)
.expect("the matcher records a change on every node")
{
ChangeKind::Unchanged(back)
| ChangeKind::ReplacedComment(_, back)
| ChangeKind::ReplacedString(_, back) => back,
ChangeKind::IgnoredPunctuation | ChangeKind::Novel => return None,
};
(back.id() == node.id()).then_some(other)
}
pub(crate) fn project(node: &Syntax<'_>, partner: Option<&Syntax<'_>>) -> Option<Fold> {
let own = node.info().fold.borrow();
let metadata = own.as_ref()?;
let own_range = range(node, metadata)?;
let opposite = partner.filter(|partner| {
partner
.info()
.fold
.borrow()
.as_ref()
.is_some_and(|metadata| range(partner, metadata).is_some())
});
Some(Fold {
tags: metadata.tags.clone(),
range: own_range,
syntax_id: node.id(),
match_kind: match opposite {
Some(partner) => FoldMatch::Matched {
opposite: partner.id(),
},
None => FoldMatch::Novel,
},
placeholder: String::new(),
})
}
pub(crate) fn merge_spans(folds: &mut Vec<Fold>, lines: &[&str]) {
let inside = |inner: &SourceRange, outer: &SourceRange| {
let at = |position: &SourcePosition| (position.line.as_usize(), position.byte_column);
at(&inner.start) >= at(&outer.start) && at(&inner.end) <= at(&outer.end)
};
let mut kept: Vec<Fold> = Vec::with_capacity(folds.len());
let mut by_span: DftHashMap<(usize, usize), usize> = DftHashMap::default();
for fold in folds.drain(..) {
let span = line_span(&fold, lines);
if span.1 - span.0 < 2 {
kept.push(fold);
continue;
}
match by_span.entry(span) {
Entry::Vacant(entry) => {
entry.insert(kept.len());
kept.push(fold);
}
Entry::Occupied(entry) => {
let merged = &mut kept[*entry.get()];
if inside(&fold.range, &merged.range) {
merged.range = fold.range;
merged.syntax_id = fold.syntax_id;
merged.match_kind = fold.match_kind;
}
merged.tags.extend(fold.tags);
merged.tags.sort();
merged.tags.dedup();
}
}
}
*folds = kept;
}
#[cfg(test)]
mod tests {
use super::*;
use crate::diff::sliders::fix_all_sliders;
use crate::parse::guess_language::Language;
use crate::parse::syntax::{init_all_info, AtomKind};
use line_numbers::SingleLineSpan;
use typed_arena::Arena;
#[test]
fn a_pair_moved_by_a_one_sided_slider_fix_is_not_a_pair() {
let span = |line: u32| {
vec![SingleLineSpan {
line: line.into(),
start_col: 0,
end_col: 1,
}]
};
let arena = Arena::new();
let lhs_x = Syntax::new_atom(&arena, span(0), "x".to_owned(), AtomKind::Normal);
let lhs_inner = Syntax::new_list(&arena, "(", span(0), vec![lhs_x], ")", span(0));
let lhs_outer = Syntax::new_list(&arena, "(", span(0), vec![lhs_inner], ")", span(0));
let rhs_x = Syntax::new_atom(&arena, span(1), "x".to_owned(), AtomKind::Normal);
let rhs_outer = Syntax::new_list(&arena, "(", span(1), vec![rhs_x], ")", span(1));
init_all_info(&[lhs_outer], &[rhs_outer]);
let mut change_map = ChangeMap::default();
change_map.insert(lhs_outer, ChangeKind::Unchanged(rhs_outer));
change_map.insert(rhs_outer, ChangeKind::Unchanged(lhs_outer));
change_map.insert(lhs_inner, ChangeKind::Novel);
change_map.insert(lhs_x, ChangeKind::Unchanged(rhs_x));
change_map.insert(rhs_x, ChangeKind::Unchanged(lhs_x));
fn partner_of<'a>(node: &'a Syntax<'a>, change_map: &ChangeMap<'a>) -> Option<SyntaxId> {
partner(node, change_map.get(node).unwrap(), change_map).map(Syntax::id)
}
assert_eq!(partner_of(lhs_outer, &change_map), Some(rhs_outer.id()));
assert_eq!(partner_of(rhs_outer, &change_map), Some(lhs_outer.id()));
fix_all_sliders(Language::EmacsLisp, &[lhs_outer], &mut change_map);
assert_eq!(
change_map.get(lhs_inner),
Some(ChangeKind::Unchanged(rhs_outer))
);
assert_eq!(
change_map.get(rhs_outer),
Some(ChangeKind::Unchanged(lhs_outer))
);
assert_eq!(partner_of(lhs_inner, &change_map), None);
assert_eq!(partner_of(rhs_outer, &change_map), None);
assert_eq!(partner_of(lhs_outer, &change_map), None);
}
#[test]
fn crossing_folds_give_the_shared_line_to_the_later_fold() {
let spans = [(0, 6), (5, 9), (7, 8)];
assert_eq!(
nested_spans(&spans),
vec![Some((0, 5)), Some((5, 9)), None],
"the earlier fold ends where the later starts; one-line spans go"
);
assert_eq!(
nested_spans(&[(0, 10), (2, 6), (4, 8)]),
vec![Some((0, 10)), Some((2, 4)), Some((4, 8))]
);
assert_eq!(nested_spans(&[(3, 5), (4, 9)]), vec![None, Some((4, 9))]);
assert_eq!(
nested_spans(&[(5, 9), (0, 4), (1, 3)]),
vec![Some((5, 9)), Some((0, 4)), Some((1, 3))]
);
}
}