use std::iter::zip;
use itertools::Itertools;
use log::debug;
use regex::Regex;
use crate::{
ast::AstNode,
class_mapping::{ClassMapping, Leader, RevNode},
lang_profile::CommutativeParent,
merged_tree::{Conflict, MergedTree},
pcs::Revision,
signature::isomorphic_merged_trees,
};
impl<'a> MergedTree<'a> {
pub(crate) fn post_process_for_duplicate_signatures(
self,
class_mapping: &ClassMapping<'a>,
) -> Self {
match self {
Self::MixedTree { node, children, .. } => {
let recursively_processed = children
.into_iter()
.map(|element| element.post_process_for_duplicate_signatures(class_mapping))
.collect();
let commutative_parent = node.commutative_parent_definition();
if let Some(commutative_parent) = commutative_parent {
let highlighted = highlight_duplicate_signatures(
&node,
recursively_processed,
class_mapping,
commutative_parent,
);
Self::new_mixed(node, highlighted)
} else {
Self::new_mixed(node, recursively_processed)
}
}
Self::ExactTree { .. }
| Self::Conflict { .. }
| Self::LineBasedMerge { .. }
| Self::CommutativeChildSeparator { .. } => self,
}
}
}
fn highlight_duplicate_signatures<'a>(
parent: &Leader<'a>,
elements: Vec<MergedTree<'a>>,
class_mapping: &ClassMapping<'a>,
commutative_parent: &CommutativeParent,
) -> Vec<MergedTree<'a>> {
let sigs: Vec<_> = elements
.iter()
.map(|element| element.signature(class_mapping))
.collect();
let sig_to_indices = sigs
.iter()
.enumerate()
.filter_map(|(idx, sig)| sig.as_ref().map(|signature| (signature, idx)))
.into_group_map();
let mut conflict_found = false;
sig_to_indices
.iter()
.filter_map(|(signature, indices)| (indices.len() > 1).then_some(signature))
.for_each(|signature| {
conflict_found = true;
debug!(
"signature conflict found in {}: {}",
commutative_parent.parent_type(),
signature
);
});
if !conflict_found {
return elements;
}
let trimmed_separator = commutative_parent.trimmed_separator();
let separator_example = find_separator(parent, trimmed_separator, class_mapping);
let end_regex = Regex::new("\n[ \t]*$").unwrap();
let add_separator = {
if let Some(node) = separator_example {
let full_source = node.node.source_with_surrounding_whitespace();
if end_regex.is_match(full_source) {
AddSeparator::AtEnd
} else {
AddSeparator::OnlyInside
}
} else {
AddSeparator::OnlyInside
}
};
let mut filtered_elements = Vec::new();
let mut skip_next_separator = true;
debug_assert_eq!(
elements.len(),
sigs.len(),
"Inconsistent length of signature arrays and elements array"
);
for (idx, (element, sig)) in zip(&elements, &sigs).enumerate().rev() {
match sig {
None => {
let is_separator = is_separator(element, trimmed_separator);
if !(is_separator && skip_next_separator) {
filtered_elements.push((idx, is_separator, element));
}
skip_next_separator = false;
}
Some(signature) => {
let cluster = sig_to_indices
.get(signature)
.expect("Signature not indexed in sig_to_indices map");
skip_next_separator = Some(&idx) != cluster.iter().min();
if !skip_next_separator {
filtered_elements.push((idx, false, element));
}
}
}
}
let mut result = Vec::new();
skip_next_separator = true;
for (filtered_idx, (idx, is_separator, element)) in
filtered_elements.iter().copied().enumerate().rev()
{
let sig = sigs
.get(idx)
.expect("Inconsistent of length of signature arrays and elements array");
match sig {
None => {
if !(is_separator && skip_next_separator) {
result.push(element.clone());
}
skip_next_separator = false;
}
Some(signature) => {
let cluster = sig_to_indices
.get(signature)
.expect("Signature not indexed in sig_to_indices map");
skip_next_separator = false;
if cluster.len() == 1 {
result.push(element.clone());
} else {
if Some(&idx) == cluster.iter().min() {
let conflict_add_separator = match add_separator {
AddSeparator::OnlyInside => AddSeparator::OnlyInside,
AddSeparator::AtEnd => {
if let Some((_, true, _)) = filtered_elements.get(filtered_idx - 1)
{
AddSeparator::AtEnd
} else {
AddSeparator::OnlyInside
}
}
};
let (mut merged, happy_path) = merge_same_sigs(
&cluster
.iter()
.map(|idx| {
elements
.get(*idx)
.expect("Invalid element index in sig_to_indices")
})
.collect::<Vec<_>>(),
class_mapping,
separator_example,
conflict_add_separator,
);
if !happy_path {
match add_separator {
AddSeparator::OnlyInside => {}
AddSeparator::AtEnd => {
if let Some((_, true, _)) =
filtered_elements.get(filtered_idx - 1)
{
skip_next_separator = true;
}
}
};
}
result.append(&mut merged);
} else {
skip_next_separator = true;
}
}
}
}
}
result
}
fn is_separator(element: &MergedTree, trimmed_separator: &'static str) -> bool {
match element {
MergedTree::ExactTree { node, .. } => {
node.as_representative().node.source.trim() == trimmed_separator
}
MergedTree::MixedTree { .. } | MergedTree::Conflict { .. } => false,
MergedTree::LineBasedMerge { parsed, .. } => {
parsed
.render_conflictless()
.is_some_and(|r| r.trim() == trimmed_separator)
}
MergedTree::CommutativeChildSeparator { .. } => true,
}
}
#[derive(Debug, Clone, Copy, Eq, PartialEq)]
enum AddSeparator {
OnlyInside,
AtEnd,
}
fn merge_same_sigs<'a>(
elements: &[&MergedTree<'a>],
class_mapping: &ClassMapping<'a>,
separator: Option<RevNode<'a>>,
add_separator: AddSeparator,
) -> (Vec<MergedTree<'a>>, bool) {
if let &[first, second] = elements
&& isomorphic_merged_trees(first, second, class_mapping)
{
return (vec![first.clone()], true);
}
let base = filter_by_revision(elements, Revision::Base, class_mapping);
let left = filter_by_revision(elements, Revision::Left, class_mapping);
let right = filter_by_revision(elements, Revision::Right, class_mapping);
if left.len() == right.len()
&& zip(&left, &right).all(|(elem_left, elem_right)| elem_left.isomorphic_to(elem_right))
{
let left_revnodes = left
.iter()
.map(|ast_node| RevNode::new(Revision::Left, ast_node))
.collect();
let v = add_separators(left_revnodes, separator, add_separator)
.into_iter()
.map(|rev_node| {
let leader = class_mapping.map_to_leader(rev_node);
MergedTree::new_exact(leader, class_mapping.revision_set(&leader), class_mapping)
})
.collect();
(v, false)
} else {
let separator = separator.map(|revnode| revnode.node);
(
vec![MergedTree::Conflict(Conflict {
base: add_separators(base, separator, add_separator),
left: add_separators(left, separator, add_separator),
right: add_separators(right, separator, add_separator),
})],
false,
)
}
}
fn filter_by_revision<'a>(
elements: &[&MergedTree<'a>],
revision: Revision,
class_mapping: &ClassMapping<'a>,
) -> Vec<&'a AstNode<'a>> {
elements
.iter()
.copied()
.filter_map(|element| match element {
MergedTree::ExactTree { node, .. }
| MergedTree::MixedTree { node, .. }
| MergedTree::LineBasedMerge { node, .. } => class_mapping.node_at_rev(node, revision),
MergedTree::Conflict { .. } | MergedTree::CommutativeChildSeparator { .. } => None,
})
.collect()
}
fn add_separators<T: Clone + Copy>(
elements: Vec<T>,
separator: Option<T>,
add_separator: AddSeparator,
) -> Vec<T> {
if elements.is_empty() {
return vec![];
}
let Some(separator) = separator else {
return elements;
};
let mut result = Vec::with_capacity(elements.len() * 2);
#[allow(unstable_name_collisions)] result.extend(elements.into_iter().intersperse(separator));
if add_separator == AddSeparator::AtEnd {
result.push(separator);
}
result
}
fn find_separator<'a>(
parent: &Leader<'a>,
trimmed_separator: &'static str,
class_mapping: &ClassMapping<'a>,
) -> Option<RevNode<'a>> {
let revs = [Revision::Base, Revision::Left, Revision::Right];
revs.into_iter()
.filter_map(|rev| {
class_mapping
.node_at_rev(parent, rev)
.map(|node| (rev, node))
})
.flat_map(|(rev, node)| {
node.children
.iter()
.map(move |child| RevNode::new(rev, child))
})
.find(|revnode| revnode.node.source.trim() == trimmed_separator)
}