use std::collections::{HashMap, HashSet};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct TraversalSpec {
pub start_token_id: u64,
pub end_token_id: u64,
pub include_reverse: bool,
pub min_depth: usize,
}
#[must_use]
pub fn prepare_traversal_plan(
start_token_ids: &[u64],
end_token_ids: &[u64],
min_depth: usize,
filter_len: Option<usize>,
) -> Vec<TraversalSpec> {
let effective_min_depth = filter_len.map_or(min_depth, |n| min_depth.max(n));
let starts = dedup_ordered(start_token_ids);
let ends = dedup_ordered(end_token_ids);
let mut entries: Vec<Option<(u64, u64, bool)>> = Vec::with_capacity(starts.len() * ends.len());
let mut positions: HashMap<(u64, u64), usize> = HashMap::new();
for &start in &starts {
for &end in &ends {
positions.insert((start, end), entries.len());
entries.push(Some((start, end, false)));
}
}
let end_set: HashSet<u64> = ends.iter().copied().collect();
let mut common: Vec<u64> = Vec::new();
let mut seen: HashSet<u64> = HashSet::new();
for &token in &starts {
if end_set.contains(&token) && seen.insert(token) {
common.push(token);
}
}
if common.len() > 1 {
for i in 0..common.len() {
for j in (i + 1)..common.len() {
let (a, b) = (common[i], common[j]);
if let Some(&index) = positions.get(&(a, b)) {
if let Some(entry) = entries.get_mut(index).and_then(Option::as_mut) {
entry.2 = true;
}
}
if let Some(&index) = positions.get(&(b, a)) {
entries[index] = None;
}
}
}
}
entries
.into_iter()
.flatten()
.map(
|(start_token_id, end_token_id, include_reverse)| TraversalSpec {
start_token_id,
end_token_id,
include_reverse,
min_depth: effective_min_depth,
},
)
.collect()
}
fn dedup_ordered(ids: &[u64]) -> Vec<u64> {
let mut seen = HashSet::with_capacity(ids.len());
let mut out = Vec::with_capacity(ids.len());
for &id in ids {
if seen.insert(id) {
out.push(id);
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
fn spec(start: u64, end: u64, reverse: bool, min_depth: usize) -> TraversalSpec {
TraversalSpec {
start_token_id: start,
end_token_id: end,
include_reverse: reverse,
min_depth,
}
}
#[test]
fn cartesian_product_is_forward() {
let plan = prepare_traversal_plan(&[1, 2], &[3, 4], 2, None);
assert_eq!(
plan,
vec![
spec(1, 3, false, 2),
spec(1, 4, false, 2),
spec(2, 3, false, 2),
spec(2, 4, false, 2),
]
);
}
#[test]
fn common_tokens_consolidate_into_one_reverse_traversal() {
let plan = prepare_traversal_plan(&[1, 2], &[1, 2], 2, None);
assert_eq!(
plan,
vec![
spec(1, 1, false, 2),
spec(1, 2, true, 2),
spec(2, 2, false, 2)
]
);
}
#[test]
fn single_common_token_is_not_consolidated() {
let plan = prepare_traversal_plan(&[1], &[1], 2, None);
assert_eq!(plan, vec![spec(1, 1, false, 2)]);
}
#[test]
fn filter_length_floors_the_min_depth() {
let plan = prepare_traversal_plan(&[1], &[1], 2, Some(3));
assert_eq!(plan, vec![spec(1, 1, false, 3)]);
}
#[test]
fn filter_length_below_min_depth_keeps_min_depth() {
let plan = prepare_traversal_plan(&[1], &[1], 4, Some(2));
assert_eq!(plan, vec![spec(1, 1, false, 4)]);
}
#[test]
fn duplicate_boundary_ids_are_deduplicated() {
let plan = prepare_traversal_plan(&[1, 1, 2], &[2, 2], 2, None);
assert_eq!(plan, vec![spec(1, 2, false, 2), spec(2, 2, false, 2)]);
}
}