use std::collections::{HashMap, HashSet};
use tree_sitter::Node;
use crate::code::metadata::metadata_of;
use crate::code::{Code, Language};
use crate::diff::PassCtx;
use crate::diff::apted::{self, Algorithm};
use crate::diff::nodes::flow_control_similarity_of_sets;
use crate::diff::{
ASTDiff, grouped_greedy_matcher, nodes, solve_greedy_anchor_blocks, solve_large_flat_subtrees,
};
pub fn solve(ctx: &PassCtx, diff: &mut ASTDiff) {
let (before, after) = (ctx.before, ctx.after);
solve_large_flat_subtrees::solve(ctx, diff);
solve_qualified_name_groups(before, after, diff);
solve_import_list_overlap(before, after, diff);
solve_import_path_similarity(before, after, diff);
solve_greedy_anchor_blocks::solve(ctx, diff);
}
fn solve_qualified_name_groups(before: &Code, after: &Code, diff: &mut ASTDiff) {
let before_metadata = metadata_of(before);
let after_metadata = metadata_of(after);
let language = before_metadata.language;
let Some(before_root) = before.ast.as_ref().map(|ast| ast.root_node()) else {
return;
};
let Some(after_root) = after.ast.as_ref().map(|ast| ast.root_node()) else {
return;
};
let before_groups = collect_qualified_name_groups(before_root, &language, before);
let after_groups = collect_qualified_name_groups(after_root, &language, after);
match_qualified_name_groups(
before_groups,
after_groups,
&before_metadata,
&after_metadata,
diff,
);
}
#[allow(clippy::too_many_arguments)]
pub(crate) fn solve_qualified_name_groups_within(
before_root: Node,
before_root_id: usize,
after_root: Node,
after_root_id: usize,
before_metadata: &crate::code::ASTMetadata,
after_metadata: &crate::code::ASTMetadata,
before_code: &Code,
after_code: &Code,
diff: &mut ASTDiff,
) {
let language = before_metadata.language;
let before_groups = collect_qualified_name_groups_excluding_root(
before_root,
before_root_id,
&language,
before_code,
);
let after_groups = collect_qualified_name_groups_excluding_root(
after_root,
after_root_id,
&language,
after_code,
);
match_qualified_name_groups(
before_groups,
after_groups,
before_metadata,
after_metadata,
diff,
);
}
fn flatten_and_sort_candidates(
groups: HashMap<(String, String), Vec<usize>>,
metadata: &crate::code::ASTMetadata,
) -> Vec<(usize, (String, String))> {
let mut candidates: Vec<(usize, (String, String))> = groups
.into_iter()
.flat_map(|(key, ids)| ids.into_iter().map(move |id| (id, key.clone())))
.collect();
candidates.sort_by_key(|(id, _)| {
metadata
.node_info
.get(id)
.map(|i| i.preorder_index)
.unwrap_or(usize::MAX)
});
candidates
}
fn match_qualified_name_groups(
before_groups: HashMap<(String, String), Vec<usize>>,
after_groups: HashMap<(String, String), Vec<usize>>,
before_metadata: &crate::code::ASTMetadata,
after_metadata: &crate::code::ASTMetadata,
diff: &mut ASTDiff,
) {
let before_candidates = flatten_and_sort_candidates(before_groups, before_metadata);
let after_candidates = flatten_and_sort_candidates(after_groups, after_metadata);
grouped_greedy_matcher::solve(
diff,
&before_candidates,
&after_candidates,
|before_id, after_id| {
solve_greedy_anchor_blocks::cost_ratio(
before_id,
after_id,
before_metadata,
after_metadata,
)
.unwrap_or(0.0)
},
None,
|before_id, after_id, diff| {
apted::prematch_identical_statement_siblings(
before_id,
after_id,
before_metadata,
after_metadata,
"qualified_name",
diff,
);
apted::prematch_unique_named_locals(
before_id,
after_id,
before_metadata,
after_metadata,
"unique_named_local",
diff,
);
apted::for_nodes(
before_metadata,
after_metadata,
vec![before_id],
vec![after_id],
Algorithm::Apted,
"qualified_name",
diff,
);
},
);
}
fn collect_qualified_name_groups(
root: Node,
language: &Language,
code: &Code,
) -> HashMap<(String, String), Vec<usize>> {
let mut out = HashMap::new();
let mut scope: Vec<String> = Vec::new();
collect_qualified_name_groups_rec(root, language, code, &mut scope, &mut out);
out
}
fn collect_qualified_name_groups_excluding_root(
root: Node,
root_id: usize,
language: &Language,
code: &Code,
) -> HashMap<(String, String), Vec<usize>> {
let mut groups = collect_qualified_name_groups(root, language, code);
for ids in groups.values_mut() {
ids.retain(|&id| id != root_id);
}
groups.retain(|_, ids| !ids.is_empty());
groups
}
fn collect_qualified_name_groups_rec(
node: Node,
language: &Language,
code: &Code,
scope: &mut Vec<String>,
out: &mut HashMap<(String, String), Vec<usize>>,
) {
let mut pushed_scope = false;
if let Some((kind, name)) = nodes::is_semantically_structural(&node, language, code) {
let full_name = if scope.is_empty() {
name.clone()
} else {
format!("{}::{}", scope.join("::"), name)
};
out.entry((kind, full_name)).or_default().push(node.id());
scope.push(name);
pushed_scope = true;
}
let mut cursor = node.walk();
for child in node.children(&mut cursor) {
collect_qualified_name_groups_rec(child, language, code, scope, out);
}
if pushed_scope {
scope.pop();
}
}
const IMPORT_LIST_SIMILARITY_THRESHOLD: f64 = 0.5;
fn solve_import_list_overlap(before: &Code, after: &Code, diff: &mut ASTDiff) {
let before_metadata = metadata_of(before);
let after_metadata = metadata_of(after);
if before_metadata.language != Language::Rust {
return;
}
let Some(before_root) = before.ast.as_ref().map(|ast| ast.root_node()) else {
return;
};
let Some(after_root) = after.ast.as_ref().map(|ast| ast.root_node()) else {
return;
};
let before_items =
collect_rust_grouped_use_declarations(before_root, before, &diff.before_node_map);
let after_items =
collect_rust_grouped_use_declarations(after_root, after, &diff.after_node_map);
if before_items.is_empty() || after_items.is_empty() {
return;
}
let before_candidates: Vec<(usize, String)> = before_items
.iter()
.map(|(id, path, _)| (*id, path.clone()))
.collect();
let after_candidates: Vec<(usize, String)> = after_items
.iter()
.map(|(id, path, _)| (*id, path.clone()))
.collect();
let before_symbols: HashMap<usize, &HashSet<&str>> = before_items
.iter()
.map(|(id, _, symbols)| (*id, symbols))
.collect();
let after_symbols: HashMap<usize, &HashSet<&str>> = after_items
.iter()
.map(|(id, _, symbols)| (*id, symbols))
.collect();
grouped_greedy_matcher::solve(
diff,
&before_candidates,
&after_candidates,
|before_id, after_id| {
1.0 - flow_control_similarity_of_sets(
before_symbols[&before_id],
after_symbols[&after_id],
)
},
Some(1.0 - IMPORT_LIST_SIMILARITY_THRESHOLD),
|before_id, after_id, diff| {
apted::for_nodes(
&before_metadata,
&after_metadata,
vec![before_id],
vec![after_id],
Algorithm::Apted,
"import_list_overlap",
diff,
);
},
);
}
const IMPORT_PATH_SIMILARITY_THRESHOLD: f64 = 0.5;
const IMPORT_PATH_MIN_SHARED_TOKENS: usize = 2;
const IMPORT_PATH_TOKEN_BUDGET_DIVISOR: usize = 6;
fn solve_import_path_similarity(before: &Code, after: &Code, diff: &mut ASTDiff) {
let before_metadata = metadata_of(before);
let after_metadata = metadata_of(after);
let language = before_metadata.language;
if language != after_metadata.language {
return;
}
let Some(before_root) = before.ast.as_ref().map(|ast| ast.root_node()) else {
return;
};
let Some(after_root) = after.ast.as_ref().map(|ast| ast.root_node()) else {
return;
};
let before_items =
collect_unmapped_imports(before_root, before, &diff.before_node_map, &language);
let after_items = collect_unmapped_imports(after_root, after, &diff.after_node_map, &language);
if before_items.is_empty() || after_items.is_empty() {
return;
}
let candidates =
|items: &[(usize, &'static str, HashSet<&str>)]| -> Vec<(usize, &'static str)> {
items.iter().map(|(id, kind, _)| (*id, *kind)).collect()
};
let before_candidates = candidates(&before_items);
let after_candidates = candidates(&after_items);
let before_tokens: HashMap<usize, &HashSet<&str>> = before_items
.iter()
.map(|(id, _, tokens)| (*id, tokens))
.collect();
let after_tokens: HashMap<usize, &HashSet<&str>> = after_items
.iter()
.map(|(id, _, tokens)| (*id, tokens))
.collect();
grouped_greedy_matcher::solve(
diff,
&before_candidates,
&after_candidates,
|before_id, after_id| {
let (before_set, after_set) = (before_tokens[&before_id], after_tokens[&after_id]);
if !import_paths_are_the_same_module(before_set, after_set)
|| !import_paths_have_no_rival(before_set, after_set, &before_items, &after_items)
{
return f64::INFINITY;
}
1.0 - flow_control_similarity_of_sets(before_set, after_set)
},
Some(1.0 - IMPORT_PATH_SIMILARITY_THRESHOLD),
|before_id, after_id, diff| {
apted::for_nodes(
&before_metadata,
&after_metadata,
vec![before_id],
vec![after_id],
Algorithm::Apted,
"import_path_similarity",
diff,
);
},
);
}
fn import_paths_are_the_same_module(before: &HashSet<&str>, after: &HashSet<&str>) -> bool {
let shared = before.intersection(after).count();
let union = before.union(after).count();
let differing = union - shared;
shared >= IMPORT_PATH_MIN_SHARED_TOKENS
&& differing <= 1 + union / IMPORT_PATH_TOKEN_BUDGET_DIVISOR
}
fn import_paths_have_no_rival(
before: &HashSet<&str>,
after: &HashSet<&str>,
before_items: &[(usize, &'static str, HashSet<&str>)],
after_items: &[(usize, &'static str, HashSet<&str>)],
) -> bool {
let rivals_for = |set: &HashSet<&str>, items: &[(usize, &'static str, HashSet<&str>)]| {
items
.iter()
.filter(|(_, _, other)| import_paths_are_the_same_module(set, other))
.count()
};
rivals_for(before, after_items) == 1 && rivals_for(after, before_items) == 1
}
const IMPORT_KEYWORD_TOKENS: &[&str] = &[
"import", "include", "use", "using", "from", "require", "package",
];
fn collect_unmapped_imports<'a>(
root: Node<'a>,
code: &'a Code,
mapped: &rustc_hash::FxHashMap<usize, usize>,
language: &Language,
) -> Vec<(usize, &'static str, HashSet<&'a str>)> {
let bytes = code.contents.as_bytes();
let mut out = Vec::new();
let mut stack = vec![root];
while let Some(node) = stack.pop() {
if nodes::is_import_kind(node.kind(), language) {
if !mapped.contains_key(&node.id())
&& let Ok(text) = node.utf8_text(bytes)
{
let tokens: HashSet<&str> = text
.split(|c: char| !c.is_alphanumeric())
.filter(|token| {
!token.is_empty()
&& !IMPORT_KEYWORD_TOKENS
.iter()
.any(|keyword| token.eq_ignore_ascii_case(keyword))
})
.collect();
if !tokens.is_empty() {
out.push((node.id(), node.kind(), tokens));
}
}
continue;
}
let mut cursor = node.walk();
for child in node.children(&mut cursor) {
stack.push(child);
}
}
out
}
fn collect_rust_grouped_use_declarations<'a>(
root: Node<'a>,
code: &'a Code,
mapped: &rustc_hash::FxHashMap<usize, usize>,
) -> Vec<(usize, String, HashSet<&'a str>)> {
let bytes = code.contents.as_bytes();
let mut out = Vec::new();
let mut stack = vec![root];
while let Some(node) = stack.pop() {
if node.kind() == "use_declaration" && !mapped.contains_key(&node.id()) {
if let Some(scoped_use_list) = node.child_by_field_name("argument")
&& scoped_use_list.kind() == "scoped_use_list"
&& let Some(path_node) = scoped_use_list.child_by_field_name("path")
&& let Some(list_node) = scoped_use_list.child_by_field_name("list")
&& list_node.kind() == "use_list"
&& let Ok(path_text) = path_node.utf8_text(bytes)
{
let mut symbols = HashSet::new();
let mut cursor = list_node.walk();
for child in list_node.named_children(&mut cursor) {
if let Ok(text) = child.utf8_text(bytes) {
symbols.insert(text);
}
}
if !symbols.is_empty() {
out.push((node.id(), path_text.to_string(), symbols));
}
}
continue;
}
let mut cursor = node.walk();
for child in node.children(&mut cursor) {
stack.push(child);
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
use crate::code::Language;
use crate::diff::ASTMappingOperation;
use crate::diff::NodeCache;
use crate::test::helper::find_first_of_kind;
#[test]
fn methods_in_different_impls_are_matched_within_their_own_impl() {
let before_src = "
struct Foo;
struct Bar;
impl Foo { fn new() -> Foo { Foo } }
impl Bar { fn new() -> Bar { Bar } }
";
let after_src = "
struct Foo;
struct Bar;
impl Foo { fn new() -> Foo { Foo } }
impl Bar { fn new() -> Bar { Bar::default() } }
";
let before = Code::from_string(before_src, &Language::Rust);
let after = Code::from_string(after_src, &Language::Rust);
let node_cache = NodeCache::build(&before, &after);
let mut diff = ASTDiff::default();
solve(
&crate::diff::PassCtx::new(&before, &after, &node_cache),
&mut diff,
);
let before_root = before.ast.as_ref().unwrap().root_node();
let after_root = after.ast.as_ref().unwrap().root_node();
let foo_new_mapping = crate::test::helper::mapping_for_path(
&["impl_item:1", "declaration_list", "function_item"],
&["impl_item:1", "declaration_list", "function_item"],
before_root,
after_root,
&diff,
)
.unwrap();
assert_eq!(
foo_new_mapping.operation,
ASTMappingOperation::Identical,
"Foo::new should be identical"
);
let bar_new_mapping = crate::test::helper::mapping_for_path(
&["impl_item:2", "declaration_list", "function_item"],
&["impl_item:2", "declaration_list", "function_item"],
before_root,
after_root,
&diff,
)
.unwrap();
assert_eq!(
bar_new_mapping.operation,
ASTMappingOperation::MatchButNotIdentical,
"Bar::new should be changed"
);
}
#[test]
fn go_subtests_named_by_literal_are_individually_matched_via_qualified_name() {
let before_src = "
package main
func TestThings(t *testing.T) {
t.Run(\"alpha\", func(t *testing.T) { old() })
t.Run(\"beta\", func(t *testing.T) { old() })
}
";
let after_src = "
package main
func TestThings(t *testing.T) {
t.Run(\"alpha\", func(t *testing.T) { newImpl() })
t.Run(\"beta\", func(t *testing.T) { newImpl() })
}
";
let before = Code::from_string(before_src, &Language::Go);
let after = Code::from_string(after_src, &Language::Go);
let node_cache = NodeCache::build(&before, &after);
let mut diff = ASTDiff::default();
solve(
&crate::diff::PassCtx::new(&before, &after, &node_cache),
&mut diff,
);
let qualified_name_count = diff
.mapping
.values()
.filter(|m| {
matches!(
&m.reason,
crate::diff::ASTMappingReason::APTED("qualified_name")
)
})
.count();
assert!(
qualified_name_count >= 2,
"expected each named subtest call to be independently matched via qualified_name, got {qualified_name_count}"
);
}
#[test]
fn overloaded_same_name_functions_are_matched_nm() {
let before_src = "
struct Foo;
impl Foo { fn a() -> i32 { 1 } }
impl Foo { fn b() -> i32 { 2 } }
";
let after_src = "
struct Foo;
impl Foo { fn a() -> i32 { 10 } }
impl Foo { fn b() -> i32 { 20 } }
";
let before = Code::from_string(before_src, &Language::Rust);
let after = Code::from_string(after_src, &Language::Rust);
let node_cache = NodeCache::build(&before, &after);
let mut diff = ASTDiff::default();
solve(
&crate::diff::PassCtx::new(&before, &after, &node_cache),
&mut diff,
);
let before_ast = before.ast.as_ref().unwrap();
let mapped_fn_count = before_ast
.root_node()
.children(&mut before_ast.root_node().walk())
.filter(|n| n.kind() == "impl_item")
.flat_map(|impl_node| {
let mut cursor = impl_node.walk();
impl_node
.children(&mut cursor)
.collect::<Vec<_>>()
.into_iter()
.flat_map(|decl_list| {
let mut c2 = decl_list.walk();
decl_list.children(&mut c2).collect::<Vec<_>>()
})
.filter(|n| n.kind() == "function_item")
.collect::<Vec<_>>()
})
.filter(|n| diff.before_node_map.contains_key(&n.id()))
.count();
assert_eq!(
mapped_fn_count, 2,
"both overloaded-name functions should be mapped"
);
}
#[test]
fn grouped_use_statement_survives_symbol_set_churn() {
let before_src = "use std::collections::{HashMap, HashSet, BTreeMap};\nfn f() {}\n";
let after_src = "use std::collections::{HashMap, HashSet, VecDeque};\nfn f() {}\n";
let before = Code::from_string(before_src, &Language::Rust);
let after = Code::from_string(after_src, &Language::Rust);
let node_cache = NodeCache::build(&before, &after);
let mut diff = ASTDiff::default();
solve(
&crate::diff::PassCtx::new(&before, &after, &node_cache),
&mut diff,
);
let before_use =
find_first_of_kind(before.ast.as_ref().unwrap().root_node(), "use_declaration")
.unwrap();
let after_use =
find_first_of_kind(after.ast.as_ref().unwrap().root_node(), "use_declaration").unwrap();
assert_eq!(
diff.before_node_map.get(&before_use.id()),
Some(&after_use.id()),
"use statements with the same base path and mostly-overlapping symbols should be matched"
);
}
#[test]
fn grouped_use_statements_with_no_symbol_overlap_are_not_matched() {
let before_src = "use std::collections::{HashMap, HashSet};\n";
let after_src = "use std::collections::{VecDeque, BinaryHeap};\n";
let before = Code::from_string(before_src, &Language::Rust);
let after = Code::from_string(after_src, &Language::Rust);
let node_cache = NodeCache::build(&before, &after);
let mut diff = ASTDiff::default();
solve(
&crate::diff::PassCtx::new(&before, &after, &node_cache),
&mut diff,
);
let before_use =
find_first_of_kind(before.ast.as_ref().unwrap().root_node(), "use_declaration")
.unwrap();
assert!(
!diff.before_node_map.contains_key(&before_use.id()),
"use statements with zero symbol overlap should not be matched by this pass"
);
}
fn tokens<'a>(words: &[&'a str]) -> HashSet<&'a str> {
words.iter().copied().collect()
}
#[test]
fn import_path_moved_into_a_subdirectory_is_the_same_module() {
assert!(import_paths_are_the_same_module(
&tokens(&["aoa", "hid", "h"]),
&tokens(&["usb", "aoa", "hid", "h"]),
));
}
#[test]
fn short_import_paths_differing_in_two_tokens_are_different_modules() {
assert!(!import_paths_are_the_same_module(
&tokens(&["foo", "bar", "h"]),
&tokens(&["foo", "baz", "h"]),
));
}
#[test]
fn long_import_paths_get_a_larger_differing_token_budget() {
let shared = ["a", "b", "c", "d", "e", "f", "g", "h", "i", "j"];
let before: HashSet<&str> = shared.iter().copied().chain(["x"]).collect();
let after: HashSet<&str> = shared.iter().copied().chain(["y", "z"]).collect();
assert!(import_paths_are_the_same_module(&before, &after));
}
#[test]
fn import_paths_sharing_one_token_are_different_modules() {
assert!(!import_paths_are_the_same_module(
&tokens(&["a", "b"]),
&tokens(&["a", "c"]),
));
}
#[test]
fn import_pair_with_a_rival_on_either_side_is_refused() {
let single = tokens(&["com", "unciv", "logic", "civ", "Civ"]);
let first = tokens(&["com", "unciv", "logic", "civ", "Civ", "A"]);
let second = tokens(&["com", "unciv", "logic", "civ", "Civ", "B"]);
let before_items = vec![(1, "import", first.clone()), (2, "import", second)];
let after_items = vec![(3, "import", single.clone())];
assert!(!import_paths_have_no_rival(
&first,
&single,
&before_items,
&after_items
));
assert!(import_paths_have_no_rival(
&first,
&single,
&before_items[..1],
&after_items
));
}
}