Skip to main content

trilogy_parser/
graph.rs

1use std::cmp::Ordering;
2use std::collections::{BinaryHeap, HashMap, HashSet, VecDeque};
3
4const FLOAT_TOLERANCE: f64 = 1e-9;
5
6type NodeId = usize;
7
8#[derive(Clone, Default)]
9pub struct GraphCore {
10    directed: bool,
11    node_to_id: HashMap<String, NodeId>,
12    id_to_node: Vec<String>,
13    node_order: Vec<NodeId>,
14    succ: Vec<HashSet<NodeId>>,
15    pred: Vec<HashSet<NodeId>>,
16    succ_order: Vec<Vec<NodeId>>,
17    pred_order: Vec<Vec<NodeId>>,
18}
19
20#[derive(Copy, Clone)]
21struct IndexState {
22    cost: f64,
23    node: NodeId,
24}
25
26impl Eq for IndexState {}
27
28impl PartialEq for IndexState {
29    fn eq(&self, other: &Self) -> bool {
30        self.cost.total_cmp(&other.cost) == Ordering::Equal && self.node == other.node
31    }
32}
33
34impl Ord for IndexState {
35    fn cmp(&self, other: &Self) -> Ordering {
36        other
37            .cost
38            .total_cmp(&self.cost)
39            .then_with(|| other.node.cmp(&self.node))
40    }
41}
42
43impl PartialOrd for IndexState {
44    fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
45        Some(self.cmp(other))
46    }
47}
48
49impl GraphCore {
50    pub fn new(directed: bool) -> Self {
51        Self {
52            directed,
53            node_to_id: HashMap::new(),
54            id_to_node: Vec::new(),
55            node_order: Vec::new(),
56            succ: Vec::new(),
57            pred: Vec::new(),
58            succ_order: Vec::new(),
59            pred_order: Vec::new(),
60        }
61    }
62
63    pub fn directed(&self) -> bool {
64        self.directed
65    }
66
67    pub fn add_node(&mut self, node: &str) {
68        if self.has_node(node) {
69            return;
70        }
71        let owned = node.to_string();
72        let id = self.id_to_node.len();
73        self.node_to_id.insert(owned.clone(), id);
74        self.id_to_node.push(owned);
75        self.node_order.push(id);
76        self.succ.push(HashSet::new());
77        self.pred.push(HashSet::new());
78        self.succ_order.push(Vec::new());
79        self.pred_order.push(Vec::new());
80    }
81
82    pub fn has_node(&self, node: &str) -> bool {
83        self.node_to_id.contains_key(node)
84    }
85
86    pub fn add_edge(&mut self, left: &str, right: &str) {
87        self.add_node(left);
88        self.add_node(right);
89        let left_id = self
90            .node_to_id
91            .get(left)
92            .copied()
93            .expect("left node should exist");
94        let right_id = self
95            .node_to_id
96            .get(right)
97            .copied()
98            .expect("right node should exist");
99        self.insert_edge_ids(left_id, right_id);
100    }
101
102    pub fn add_nodes(&mut self, nodes: Vec<String>) {
103        for node in nodes {
104            self.add_node(&node);
105        }
106    }
107
108    pub fn add_edges(&mut self, edges: Vec<(String, String)>) {
109        for (left, right) in edges {
110            self.add_edge(&left, &right);
111        }
112    }
113
114    pub fn has_edge(&self, left: &str, right: &str) -> bool {
115        let Some(left_id) = self.node_to_id.get(left).copied() else {
116            return false;
117        };
118        let Some(right_id) = self.node_to_id.get(right).copied() else {
119            return false;
120        };
121        self.succ
122            .get(left_id)
123            .map(|neighbors| neighbors.contains(&right_id))
124            .unwrap_or(false)
125    }
126
127    pub fn remove_edge(&mut self, left: &str, right: &str) {
128        let Some(left_id) = self.node_to_id.get(left).copied() else {
129            return;
130        };
131        let Some(right_id) = self.node_to_id.get(right).copied() else {
132            return;
133        };
134        self.remove_edge_ids(left_id, right_id);
135    }
136
137    pub fn remove_edges(&mut self, edges: Vec<(String, String)>) {
138        for (left, right) in edges {
139            self.remove_edge(&left, &right);
140        }
141    }
142
143    pub fn remove_node(&mut self, node: &str) {
144        let Some(node_id) = self.node_to_id.remove(node) else {
145            return;
146        };
147        let outgoing = self.succ_order[node_id].clone();
148        let incoming = self.pred_order[node_id].clone();
149
150        for neighbor in outgoing {
151            if self.pred[neighbor].remove(&node_id) {
152                remove_ordered_neighbor(&mut self.pred_order[neighbor], node_id);
153            }
154            if !self.directed && self.succ[neighbor].remove(&node_id) {
155                remove_ordered_neighbor(&mut self.succ_order[neighbor], node_id);
156            }
157        }
158
159        for neighbor in incoming {
160            if self.succ[neighbor].remove(&node_id) {
161                remove_ordered_neighbor(&mut self.succ_order[neighbor], node_id);
162            }
163            if !self.directed && self.pred[neighbor].remove(&node_id) {
164                remove_ordered_neighbor(&mut self.pred_order[neighbor], node_id);
165            }
166        }
167
168        self.node_order.retain(|existing| *existing != node_id);
169        self.succ[node_id].clear();
170        self.pred[node_id].clear();
171        self.succ_order[node_id].clear();
172        self.pred_order[node_id].clear();
173    }
174
175    pub fn remove_nodes(&mut self, nodes: Vec<String>) {
176        let removed = nodes
177            .into_iter()
178            .filter_map(|node| self.node_to_id.get(&node).copied())
179            .collect::<HashSet<_>>();
180        if removed.is_empty() {
181            return;
182        }
183        if removed.len() <= 8 || removed.len() * 8 < self.node_order.len() {
184            let names = removed
185                .iter()
186                .map(|id| self.id_to_node[*id].clone())
187                .collect::<Vec<_>>();
188            for name in names {
189                self.remove_node(&name);
190            }
191            return;
192        }
193        for node_id in &removed {
194            self.node_to_id.remove(&self.id_to_node[*node_id]);
195            self.succ[*node_id].clear();
196            self.pred[*node_id].clear();
197            self.succ_order[*node_id].clear();
198            self.pred_order[*node_id].clear();
199        }
200        self.node_order.retain(|node_id| !removed.contains(node_id));
201        for node_id in &self.node_order {
202            self.succ[*node_id].retain(|neighbor| !removed.contains(neighbor));
203            self.pred[*node_id].retain(|neighbor| !removed.contains(neighbor));
204            self.succ_order[*node_id].retain(|neighbor| !removed.contains(neighbor));
205            self.pred_order[*node_id].retain(|neighbor| !removed.contains(neighbor));
206        }
207    }
208
209    pub fn nodes(&self) -> Vec<String> {
210        self.node_order
211            .iter()
212            .map(|node_id| self.id_to_node[*node_id].clone())
213            .collect()
214    }
215
216    pub fn edges(&self) -> Vec<(String, String)> {
217        let mut edges = Vec::with_capacity(
218            self.node_order
219                .iter()
220                .map(|node_id| self.succ_order[*node_id].len())
221                .sum(),
222        );
223        let mut seen = HashSet::new();
224        for left_id in &self.node_order {
225            for right_id in &self.succ_order[*left_id] {
226                if self.directed || seen.insert(canonical_id_edge(*left_id, *right_id)) {
227                    edges.push((
228                        self.id_to_node[*left_id].clone(),
229                        self.id_to_node[*right_id].clone(),
230                    ));
231                }
232            }
233        }
234        edges
235    }
236
237    pub fn neighbors(&self, node: &str) -> Vec<String> {
238        self.neighbor_names(self.node_to_id.get(node).copied(), false)
239    }
240
241    pub fn predecessors(&self, node: &str) -> Vec<String> {
242        self.neighbor_names(self.node_to_id.get(node).copied(), true)
243    }
244
245    pub fn successors(&self, node: &str) -> Vec<String> {
246        self.neighbors(node)
247    }
248
249    pub fn all_neighbors(&self, node: &str) -> Vec<String> {
250        let Some(node_id) = self.node_to_id.get(node).copied() else {
251            return Vec::new();
252        };
253        let mut seen = HashSet::new();
254        let mut output = Vec::new();
255        for neighbor in &self.pred_order[node_id] {
256            if seen.insert(*neighbor) {
257                output.push(self.id_to_node[*neighbor].clone());
258            }
259        }
260        for neighbor in &self.succ_order[node_id] {
261            if seen.insert(*neighbor) {
262                output.push(self.id_to_node[*neighbor].clone());
263            }
264        }
265        output
266    }
267
268    pub fn in_degree(&self, node: &str) -> usize {
269        self.node_to_id
270            .get(node)
271            .map(|node_id| self.pred[*node_id].len())
272            .unwrap_or(0)
273    }
274
275    pub fn out_degree(&self, node: &str) -> usize {
276        self.node_to_id
277            .get(node)
278            .map(|node_id| self.succ[*node_id].len())
279            .unwrap_or(0)
280    }
281
282    pub fn clone_graph(&self) -> Self {
283        self.clone()
284    }
285
286    pub fn induced_subgraph(&self, nodes: Vec<String>) -> Self {
287        let keep = nodes
288            .into_iter()
289            .filter_map(|node| self.node_to_id.get(&node).copied())
290            .collect::<HashSet<_>>();
291        let ordered_keep = self
292            .node_order
293            .iter()
294            .copied()
295            .filter(|node_id| keep.contains(node_id))
296            .collect::<Vec<_>>();
297        let mut graph = Self::new(self.directed);
298        let mut id_map = HashMap::new();
299        for old_id in &ordered_keep {
300            let name = &self.id_to_node[*old_id];
301            graph.add_node(name);
302            id_map.insert(
303                *old_id,
304                graph
305                    .node_to_id
306                    .get(name)
307                    .copied()
308                    .expect("new node should exist"),
309            );
310        }
311        for old_id in &ordered_keep {
312            let new_id = id_map[old_id];
313            graph.succ[new_id] = self.succ[*old_id]
314                .iter()
315                .filter_map(|neighbor| id_map.get(neighbor).copied())
316                .collect();
317            graph.pred[new_id] = self.pred[*old_id]
318                .iter()
319                .filter_map(|neighbor| id_map.get(neighbor).copied())
320                .collect();
321            graph.succ_order[new_id] = self.succ_order[*old_id]
322                .iter()
323                .filter_map(|neighbor| id_map.get(neighbor).copied())
324                .collect();
325            graph.pred_order[new_id] = self.pred_order[*old_id]
326                .iter()
327                .filter_map(|neighbor| id_map.get(neighbor).copied())
328                .collect();
329        }
330        graph
331    }
332
333    pub fn to_undirected_graph(&self) -> Self {
334        if !self.directed {
335            return self.clone();
336        }
337        let mut graph = Self::new(false);
338        let mut id_map = HashMap::new();
339        for node_id in &self.node_order {
340            let name = &self.id_to_node[*node_id];
341            graph.add_node(name);
342            id_map.insert(
343                *node_id,
344                graph
345                    .node_to_id
346                    .get(name)
347                    .copied()
348                    .expect("new node should exist"),
349            );
350        }
351        for left_id in &self.node_order {
352            let new_left = id_map[left_id];
353            for right_id in &self.succ_order[*left_id] {
354                let new_right = id_map[right_id];
355                graph.insert_edge_ids(new_left, new_right);
356            }
357        }
358        graph
359    }
360
361    pub fn connected_components(&self) -> Vec<Vec<String>> {
362        let mut seen = vec![false; self.id_to_node.len()];
363        let mut components = Vec::new();
364
365        for node_id in &self.node_order {
366            if seen[*node_id] {
367                continue;
368            }
369            let mut queue = VecDeque::from([*node_id]);
370            let mut component = Vec::new();
371            seen[*node_id] = true;
372            while let Some(current) = queue.pop_front() {
373                component.push(self.id_to_node[current].clone());
374                for neighbor in &self.pred_order[current] {
375                    if !seen[*neighbor] {
376                        seen[*neighbor] = true;
377                        queue.push_back(*neighbor);
378                    }
379                }
380                for neighbor in &self.succ_order[current] {
381                    if !seen[*neighbor] {
382                        seen[*neighbor] = true;
383                        queue.push_back(*neighbor);
384                    }
385                }
386            }
387            components.push(component);
388        }
389        components
390    }
391
392    pub fn is_weakly_connected(&self) -> bool {
393        let node_count = self.node_order.len();
394        if node_count <= 1 {
395            return true;
396        }
397        self.connected_components().len() == 1
398    }
399
400    pub fn topological_sort(&self) -> Result<Vec<String>, String> {
401        let mut indegree = vec![0usize; self.id_to_node.len()];
402        for node_id in &self.node_order {
403            indegree[*node_id] = self.pred[*node_id].len();
404        }
405
406        let mut ready = VecDeque::new();
407        for node_id in &self.node_order {
408            if indegree[*node_id] == 0 {
409                ready.push_back(*node_id);
410            }
411        }
412
413        let mut output = Vec::new();
414        while let Some(node_id) = ready.pop_front() {
415            output.push(self.id_to_node[node_id].clone());
416            for neighbor in &self.succ_order[node_id] {
417                indegree[*neighbor] -= 1;
418                if indegree[*neighbor] == 0 {
419                    ready.push_back(*neighbor);
420                }
421            }
422        }
423
424        if output.len() != self.node_order.len() {
425            return Err("Graph contains a cycle".to_string());
426        }
427        Ok(output)
428    }
429
430    pub fn shortest_path(&self, source: &str, target: &str) -> Option<Vec<String>> {
431        let source_id = self.node_to_id.get(source).copied()?;
432        let target_id = self.node_to_id.get(target).copied()?;
433        if source_id == target_id {
434            return Some(vec![source.to_string()]);
435        }
436
437        let mut queue = VecDeque::from([source_id]);
438        let mut visited = vec![false; self.id_to_node.len()];
439        let mut previous = vec![None; self.id_to_node.len()];
440        visited[source_id] = true;
441
442        while let Some(current) = queue.pop_front() {
443            for neighbor in &self.succ_order[current] {
444                if visited[*neighbor] {
445                    continue;
446                }
447                visited[*neighbor] = true;
448                previous[*neighbor] = Some(current);
449                if *neighbor == target_id {
450                    return Some(self.reconstruct_id_path(&previous, source_id, target_id));
451                }
452                queue.push_back(*neighbor);
453            }
454        }
455
456        None
457    }
458
459    pub fn shortest_path_length(&self, source: &str, target: &str) -> Option<usize> {
460        self.shortest_path(source, target)
461            .map(|path| path.len().saturating_sub(1))
462    }
463
464    pub fn ego_graph_nodes(&self, center: &str, radius: usize) -> Vec<String> {
465        let Some(center_id) = self.node_to_id.get(center).copied() else {
466            return Vec::new();
467        };
468
469        let mut queue = VecDeque::from([(center_id, 0usize)]);
470        let mut visited = vec![false; self.id_to_node.len()];
471        visited[center_id] = true;
472
473        while let Some((current, depth)) = queue.pop_front() {
474            if depth >= radius {
475                continue;
476            }
477            for neighbor in &self.succ_order[current] {
478                if !visited[*neighbor] {
479                    visited[*neighbor] = true;
480                    queue.push_back((*neighbor, depth + 1));
481                }
482            }
483        }
484
485        self.node_order
486            .iter()
487            .filter(|node_id| visited[**node_id])
488            .map(|node_id| self.id_to_node[*node_id].clone())
489            .collect()
490    }
491
492    pub fn multi_source_dijkstra_path(
493        &self,
494        sources: Vec<String>,
495        weights: Vec<(String, String, f64)>,
496    ) -> Result<Vec<(String, Vec<String>)>, String> {
497        let source_ids = dedupe_preserve_order(
498            sources
499                .into_iter()
500                .map(|source| {
501                    self.node_to_id
502                        .get(&source)
503                        .copied()
504                        .ok_or_else(|| format!("Node not found: {source}"))
505                })
506                .collect::<Result<Vec<_>, _>>()?,
507        );
508        if source_ids.is_empty() {
509            return Ok(Vec::new());
510        }
511
512        let weighted_neighbors = self.build_weighted_neighbors(weights)?;
513        let ranks = self.lex_ranks();
514        let mut heap = BinaryHeap::new();
515        let mut distances = vec![f64::INFINITY; self.id_to_node.len()];
516        let mut paths: Vec<Option<Vec<NodeId>>> = vec![None; self.id_to_node.len()];
517
518        for source_id in &source_ids {
519            distances[*source_id] = 0.0;
520            paths[*source_id] = Some(vec![*source_id]);
521            heap.push(IndexState {
522                cost: 0.0,
523                node: *source_id,
524            });
525        }
526
527        while let Some(IndexState { cost, node }) = heap.pop() {
528            if cost > distances[node] + FLOAT_TOLERANCE {
529                continue;
530            }
531            let Some(current_path) = paths[node].clone() else {
532                continue;
533            };
534
535            for (neighbor, weight) in &weighted_neighbors[node] {
536                let next_cost = cost + *weight;
537                let mut next_path = current_path.clone();
538                next_path.push(*neighbor);
539
540                let should_update = if next_cost + FLOAT_TOLERANCE < distances[*neighbor] {
541                    true
542                } else if (next_cost - distances[*neighbor]).abs() <= FLOAT_TOLERANCE {
543                    match &paths[*neighbor] {
544                        None => true,
545                        Some(existing_path) => path_less(&next_path, existing_path, &ranks),
546                    }
547                } else {
548                    false
549                };
550
551                if should_update {
552                    distances[*neighbor] = next_cost;
553                    paths[*neighbor] = Some(next_path);
554                    heap.push(IndexState {
555                        cost: next_cost,
556                        node: *neighbor,
557                    });
558                }
559            }
560        }
561
562        let mut output = Vec::new();
563        for node_id in &self.node_order {
564            let Some(path) = &paths[*node_id] else {
565                continue;
566            };
567            output.push((
568                self.id_to_node[*node_id].clone(),
569                path.iter()
570                    .map(|path_id| self.id_to_node[*path_id].clone())
571                    .collect(),
572            ));
573        }
574        Ok(output)
575    }
576
577    pub fn steiner_tree_nodes(
578        &self,
579        terminals: Vec<String>,
580        weights: Vec<(String, String, f64)>,
581    ) -> Result<Vec<String>, String> {
582        let terminal_ids = dedupe_preserve_order(
583            terminals
584                .into_iter()
585                .map(|terminal| {
586                    self.node_to_id
587                        .get(&terminal)
588                        .copied()
589                        .ok_or_else(|| format!("Node not found: {terminal}"))
590                })
591                .collect::<Result<Vec<_>, _>>()?,
592        );
593        if terminal_ids.is_empty() {
594            return Ok(Vec::new());
595        }
596        if terminal_ids.len() == 1 {
597            return Ok(vec![self.id_to_node[terminal_ids[0]].clone()]);
598        }
599
600        let weighted_neighbors = self.build_weighted_neighbors(weights)?;
601        let ranks = self.lex_ranks();
602        let mut metric_edges: Vec<(f64, NodeId, NodeId, Vec<NodeId>)> = Vec::new();
603
604        for index in 0..terminal_ids.len() {
605            let left = terminal_ids[index];
606            let targets = &terminal_ids[index + 1..];
607            let shortest_paths = self.weighted_shortest_paths_to_target_indices(
608                left,
609                targets,
610                &weighted_neighbors,
611                &ranks,
612            );
613            for right in targets {
614                let Some((distance, path)) = shortest_paths.get(right) else {
615                    return Err(format!(
616                        "No path between {} and {}",
617                        self.id_to_node[left], self.id_to_node[*right]
618                    ));
619                };
620                metric_edges.push((*distance, left, *right, path.clone()));
621            }
622        }
623
624        metric_edges.sort_by(|left, right| {
625            left.0
626                .total_cmp(&right.0)
627                .then_with(|| ranks[left.1].cmp(&ranks[right.1]))
628                .then_with(|| ranks[left.2].cmp(&ranks[right.2]))
629        });
630
631        let mut metric_dsu = IndexDisjointSet::new(self.id_to_node.len());
632        let mut expanded_nodes = terminal_ids.iter().copied().collect::<HashSet<_>>();
633        for (_distance, left, right, path) in metric_edges {
634            if metric_dsu.union(left, right) {
635                expanded_nodes.extend(path);
636            }
637        }
638
639        Ok(self.finalize_steiner_tree(
640            terminal_ids,
641            expanded_nodes,
642            &weighted_neighbors,
643        ))
644    }
645
646    fn insert_edge_ids(&mut self, left_id: NodeId, right_id: NodeId) {
647        if self.succ[left_id].insert(right_id) {
648            self.succ_order[left_id].push(right_id);
649        }
650        if self.pred[right_id].insert(left_id) {
651            self.pred_order[right_id].push(left_id);
652        }
653        if !self.directed {
654            if self.succ[right_id].insert(left_id) {
655                self.succ_order[right_id].push(left_id);
656            }
657            if self.pred[left_id].insert(right_id) {
658                self.pred_order[left_id].push(right_id);
659            }
660        }
661    }
662
663    fn remove_edge_ids(&mut self, left_id: NodeId, right_id: NodeId) {
664        if self.succ[left_id].remove(&right_id) {
665            remove_ordered_neighbor(&mut self.succ_order[left_id], right_id);
666        }
667        if self.pred[right_id].remove(&left_id) {
668            remove_ordered_neighbor(&mut self.pred_order[right_id], left_id);
669        }
670        if !self.directed {
671            if self.succ[right_id].remove(&left_id) {
672                remove_ordered_neighbor(&mut self.succ_order[right_id], left_id);
673            }
674            if self.pred[left_id].remove(&right_id) {
675                remove_ordered_neighbor(&mut self.pred_order[left_id], right_id);
676            }
677        }
678    }
679
680    fn neighbor_names(&self, node_id: Option<NodeId>, reverse: bool) -> Vec<String> {
681        let Some(node_id) = node_id else {
682            return Vec::new();
683        };
684        let neighbors = if reverse {
685            &self.pred_order[node_id]
686        } else {
687            &self.succ_order[node_id]
688        };
689        neighbors
690            .iter()
691            .map(|neighbor| self.id_to_node[*neighbor].clone())
692            .collect()
693    }
694
695    fn reconstruct_id_path(
696        &self,
697        previous: &[Option<NodeId>],
698        source: NodeId,
699        target: NodeId,
700    ) -> Vec<String> {
701        let mut path = vec![self.id_to_node[target].clone()];
702        let mut current = target;
703        while current != source {
704            let Some(parent) = previous[current] else {
705                break;
706            };
707            current = parent;
708            path.push(self.id_to_node[current].clone());
709        }
710        path.reverse();
711        path
712    }
713
714    fn lex_ranks(&self) -> Vec<usize> {
715        let mut ranked = self.node_order.clone();
716        ranked.sort_by(|left, right| self.id_to_node[*left].cmp(&self.id_to_node[*right]));
717        let mut ranks = vec![usize::MAX; self.id_to_node.len()];
718        for (rank, node_id) in ranked.into_iter().enumerate() {
719            ranks[node_id] = rank;
720        }
721        ranks
722    }
723
724    fn build_weighted_neighbors(
725        &self,
726        weights: Vec<(String, String, f64)>,
727    ) -> Result<Vec<Vec<(NodeId, f64)>>, String> {
728        let mut overrides =
729            HashMap::with_capacity(weights.len() * if self.directed { 1 } else { 2 });
730        for (left, right, weight) in weights {
731            let Some(left_id) = self.node_to_id.get(&left).copied() else {
732                return Err(format!("Node not found: {left}"));
733            };
734            let Some(right_id) = self.node_to_id.get(&right).copied() else {
735                return Err(format!("Node not found: {right}"));
736            };
737            overrides.insert((left_id, right_id), weight);
738            if !self.directed {
739                overrides.insert((right_id, left_id), weight);
740            }
741        }
742
743        let mut weighted_neighbors = vec![Vec::<(NodeId, f64)>::new(); self.id_to_node.len()];
744        for node_id in &self.node_order {
745            let mut neighbors = Vec::with_capacity(self.succ_order[*node_id].len());
746            for neighbor in &self.succ_order[*node_id] {
747                neighbors.push((
748                    *neighbor,
749                    overrides
750                        .get(&(*node_id, *neighbor))
751                        .copied()
752                        .unwrap_or(1.0),
753                ));
754            }
755            weighted_neighbors[*node_id] = neighbors;
756        }
757        Ok(weighted_neighbors)
758    }
759
760    fn finalize_steiner_tree(
761        &self,
762        terminals: Vec<NodeId>,
763        expanded_nodes: HashSet<NodeId>,
764        weighted_neighbors: &[Vec<(NodeId, f64)>],
765    ) -> Vec<String> {
766        let ordered_expanded_nodes = self
767            .node_order
768            .iter()
769            .copied()
770            .filter(|node_id| expanded_nodes.contains(node_id))
771            .collect::<Vec<_>>();
772        let mut position = vec![usize::MAX; self.id_to_node.len()];
773        for (index, node_id) in ordered_expanded_nodes.iter().enumerate() {
774            position[*node_id] = index;
775        }
776
777        let mut original_edges = Vec::new();
778        let mut seen_edges = HashSet::new();
779        for left_id in &ordered_expanded_nodes {
780            for (right_id, weight) in &weighted_neighbors[*left_id] {
781                if position[*right_id] == usize::MAX {
782                    continue;
783                }
784                if self.directed || seen_edges.insert(canonical_id_edge(*left_id, *right_id)) {
785                    original_edges.push((*weight, *left_id, *right_id));
786                }
787            }
788        }
789        let ranks = self.lex_ranks();
790        original_edges.sort_by(|left, right| {
791            left.0
792                .total_cmp(&right.0)
793                .then_with(|| ranks[left.1].cmp(&ranks[right.1]))
794                .then_with(|| ranks[left.2].cmp(&ranks[right.2]))
795        });
796
797        let mut tree = vec![HashSet::<NodeId>::new(); self.id_to_node.len()];
798        let mut induced_dsu = IndexDisjointSet::new(self.id_to_node.len());
799        for (_weight, left, right) in original_edges {
800            if induced_dsu.union(left, right) {
801                tree[left].insert(right);
802                tree[right].insert(left);
803            }
804        }
805
806        let terminals_set = terminals.into_iter().collect::<HashSet<_>>();
807        let mut removable = ordered_expanded_nodes
808            .iter()
809            .filter_map(|node_id| {
810                if !terminals_set.contains(node_id) && tree[*node_id].len() <= 1 {
811                    Some(*node_id)
812                } else {
813                    None
814                }
815            })
816            .collect::<VecDeque<_>>();
817        let mut removed = vec![false; self.id_to_node.len()];
818        while let Some(node_id) = removable.pop_front() {
819            if removed[node_id] || terminals_set.contains(&node_id) || tree[node_id].len() > 1 {
820                continue;
821            }
822            removed[node_id] = true;
823            let neighbors = std::mem::take(&mut tree[node_id]);
824            for neighbor in neighbors {
825                if tree[neighbor].remove(&node_id)
826                    && !terminals_set.contains(&neighbor)
827                    && tree[neighbor].len() <= 1
828                {
829                    removable.push_back(neighbor);
830                }
831            }
832        }
833
834        self.node_order
835            .iter()
836            .filter(|node_id| {
837                expanded_nodes.contains(node_id)
838                    && !removed[**node_id]
839                    && (!tree[**node_id].is_empty() || terminals_set.contains(node_id))
840            })
841            .map(|node_id| self.id_to_node[*node_id].clone())
842            .collect()
843    }
844
845    fn weighted_shortest_paths_to_target_indices(
846        &self,
847        source: NodeId,
848        targets: &[NodeId],
849        weighted_neighbors: &[Vec<(NodeId, f64)>],
850        ranks: &[usize],
851    ) -> HashMap<NodeId, (f64, Vec<NodeId>)> {
852        if targets.is_empty() {
853            return HashMap::new();
854        }
855        let mut remaining = targets.iter().copied().collect::<HashSet<_>>();
856        let mut heap = BinaryHeap::new();
857        let mut distances = vec![f64::INFINITY; self.id_to_node.len()];
858        let mut paths: Vec<Option<Vec<NodeId>>> = vec![None; self.id_to_node.len()];
859
860        distances[source] = 0.0;
861        paths[source] = Some(vec![source]);
862        heap.push(IndexState {
863            cost: 0.0,
864            node: source,
865        });
866
867        while let Some(IndexState { cost, node }) = heap.pop() {
868            if cost > distances[node] + FLOAT_TOLERANCE {
869                continue;
870            }
871            remaining.remove(&node);
872            if remaining.is_empty() {
873                break;
874            }
875            let Some(current_path) = paths[node].clone() else {
876                continue;
877            };
878            for (neighbor, weight) in &weighted_neighbors[node] {
879                let next_cost = cost + *weight;
880                let mut next_path = current_path.clone();
881                next_path.push(*neighbor);
882                let should_update = if next_cost + FLOAT_TOLERANCE < distances[*neighbor] {
883                    true
884                } else if (next_cost - distances[*neighbor]).abs() <= FLOAT_TOLERANCE {
885                    match &paths[*neighbor] {
886                        None => true,
887                        Some(existing_path) => path_less(&next_path, existing_path, ranks),
888                    }
889                } else {
890                    false
891                };
892                if should_update {
893                    distances[*neighbor] = next_cost;
894                    paths[*neighbor] = Some(next_path);
895                    heap.push(IndexState {
896                        cost: next_cost,
897                        node: *neighbor,
898                    });
899                }
900            }
901        }
902
903        targets
904            .iter()
905            .filter_map(|target| {
906                let path = paths[*target].clone()?;
907                Some((*target, (distances[*target], path)))
908            })
909            .collect()
910    }
911}
912
913struct IndexDisjointSet {
914    parent: Vec<NodeId>,
915    rank: Vec<usize>,
916}
917
918impl IndexDisjointSet {
919    fn new(size: usize) -> Self {
920        Self {
921            parent: (0..size).collect(),
922            rank: vec![0; size],
923        }
924    }
925
926    fn find(&mut self, node: NodeId) -> NodeId {
927        if self.parent[node] != node {
928            let root = self.find(self.parent[node]);
929            self.parent[node] = root;
930        }
931        self.parent[node]
932    }
933
934    fn union(&mut self, left: NodeId, right: NodeId) -> bool {
935        let left_root = self.find(left);
936        let right_root = self.find(right);
937        if left_root == right_root {
938            return false;
939        }
940        let left_rank = self.rank[left_root];
941        let right_rank = self.rank[right_root];
942        if left_rank < right_rank {
943            self.parent[left_root] = right_root;
944        } else if left_rank > right_rank {
945            self.parent[right_root] = left_root;
946        } else {
947            self.parent[right_root] = left_root;
948            self.rank[left_root] += 1;
949        }
950        true
951    }
952}
953
954fn dedupe_preserve_order<T: Eq + std::hash::Hash + Copy>(values: Vec<T>) -> Vec<T> {
955    let mut seen = HashSet::new();
956    let mut output = Vec::new();
957    for value in values {
958        if seen.insert(value) {
959            output.push(value);
960        }
961    }
962    output
963}
964
965fn canonical_id_edge(left: NodeId, right: NodeId) -> (NodeId, NodeId) {
966    if left <= right {
967        (left, right)
968    } else {
969        (right, left)
970    }
971}
972
973fn remove_ordered_neighbor(order: &mut Vec<NodeId>, neighbor: NodeId) {
974    order.retain(|entry| *entry != neighbor);
975}
976
977fn path_less(left: &[NodeId], right: &[NodeId], ranks: &[usize]) -> bool {
978    for (left_node, right_node) in left.iter().zip(right.iter()) {
979        match ranks[*left_node].cmp(&ranks[*right_node]) {
980            Ordering::Less => return true,
981            Ordering::Greater => return false,
982            Ordering::Equal => {}
983        }
984    }
985    left.len() < right.len()
986}