Skip to main content

sz_orm_graph/
algorithm.rs

1//! 图算法(Graph Algorithms)
2//!
3//! 提供常用图算法:BFS/DFS 遍历、Dijkstra 最短路径、拓扑排序、
4//! 环检测、连通分量等。基于邻接表表示的图。
5
6use std::collections::{HashMap, HashSet, VecDeque};
7
8/// 图节点 ID 类型
9pub type NodeId = u64;
10
11/// 边权重类型
12pub type Weight = f64;
13
14/// 有向图
15///
16/// 使用邻接表表示,支持带权边。
17#[derive(Debug, Clone, Default)]
18pub struct DirectedGraph {
19    /// 邻接表:节点 -> [(邻居, 权重)]
20    adjacency: HashMap<NodeId, Vec<(NodeId, Weight)>>,
21    /// 节点数(含孤立节点)
22    node_count: usize,
23}
24
25impl DirectedGraph {
26    /// 创建空图
27    pub fn new() -> Self {
28        Self::default()
29    }
30
31    /// 添加节点
32    pub fn add_node(&mut self, node: NodeId) {
33        if let std::collections::hash_map::Entry::Vacant(e) = self.adjacency.entry(node) {
34            e.insert(Vec::new());
35            self.node_count += 1;
36        }
37    }
38
39    /// 添加带权有向边
40    pub fn add_edge(&mut self, from: NodeId, to: NodeId, weight: Weight) {
41        self.add_node(from);
42        self.add_node(to);
43        self.adjacency.get_mut(&from).unwrap().push((to, weight));
44    }
45
46    /// 添加无权有向边(权重=1.0)
47    pub fn add_edge_unweighted(&mut self, from: NodeId, to: NodeId) {
48        self.add_edge(from, to, 1.0);
49    }
50
51    /// 获取节点的邻居
52    pub fn neighbors(&self, node: NodeId) -> Option<&Vec<(NodeId, Weight)>> {
53        self.adjacency.get(&node)
54    }
55
56    /// 所有节点
57    pub fn nodes(&self) -> impl Iterator<Item = NodeId> + '_ {
58        self.adjacency.keys().copied()
59    }
60
61    /// 节点数
62    pub fn node_count(&self) -> usize {
63        self.node_count
64    }
65
66    /// 边数
67    pub fn edge_count(&self) -> usize {
68        self.adjacency.values().map(|v| v.len()).sum()
69    }
70
71    /// 是否包含节点
72    pub fn has_node(&self, node: NodeId) -> bool {
73        self.adjacency.contains_key(&node)
74    }
75
76    /// 是否包含边
77    pub fn has_edge(&self, from: NodeId, to: NodeId) -> bool {
78        self.adjacency
79            .get(&from)
80            .map(|v| v.iter().any(|(n, _)| *n == to))
81            .unwrap_or(false)
82    }
83
84    /// 获取边权重
85    pub fn edge_weight(&self, from: NodeId, to: NodeId) -> Option<Weight> {
86        self.adjacency
87            .get(&from)
88            .and_then(|v| v.iter().find(|(n, _)| *n == to).map(|(_, w)| *w))
89    }
90
91    /// 广度优先搜索(BFS)
92    ///
93    /// 从 `start` 出发,返回按 BFS 顺序访问的节点列表。
94    pub fn bfs(&self, start: NodeId) -> Vec<NodeId> {
95        if !self.has_node(start) {
96            return Vec::new();
97        }
98        let mut visited = HashSet::new();
99        let mut queue = VecDeque::new();
100        let mut result = Vec::new();
101        visited.insert(start);
102        queue.push_back(start);
103        while let Some(node) = queue.pop_front() {
104            result.push(node);
105            if let Some(neighbors) = self.neighbors(node) {
106                for &(neighbor, _) in neighbors {
107                    if visited.insert(neighbor) {
108                        queue.push_back(neighbor);
109                    }
110                }
111            }
112        }
113        result
114    }
115
116    /// 深度优先搜索(DFS)
117    ///
118    /// 从 `start` 出发,返回按 DFS 顺序访问的节点列表。
119    pub fn dfs(&self, start: NodeId) -> Vec<NodeId> {
120        if !self.has_node(start) {
121            return Vec::new();
122        }
123        let mut visited = HashSet::new();
124        let mut result = Vec::new();
125        self.dfs_visit(start, &mut visited, &mut result);
126        result
127    }
128
129    fn dfs_visit(&self, node: NodeId, visited: &mut HashSet<NodeId>, result: &mut Vec<NodeId>) {
130        if !visited.insert(node) {
131            return;
132        }
133        result.push(node);
134        if let Some(neighbors) = self.neighbors(node) {
135            for &(neighbor, _) in neighbors {
136                self.dfs_visit(neighbor, visited, result);
137            }
138        }
139    }
140
141    /// Dijkstra 最短路径算法
142    ///
143    /// 返回从 `start` 到 `end` 的最短路径和总距离。
144    /// 如果不可达返回 None。
145    pub fn dijkstra(&self, start: NodeId, end: NodeId) -> Option<(Vec<NodeId>, Weight)> {
146        if !self.has_node(start) || !self.has_node(end) {
147            return None;
148        }
149        let mut dist: HashMap<NodeId, Weight> = HashMap::new();
150        let mut prev: HashMap<NodeId, NodeId> = HashMap::new();
151        let mut visited = HashSet::new();
152        for &node in self.adjacency.keys() {
153            dist.insert(node, Weight::INFINITY);
154        }
155        dist.insert(start, 0.0);
156        while visited.len() < self.node_count {
157            let current = {
158                let mut best: Option<(NodeId, Weight)> = None;
159                for (&node, &d) in dist.iter() {
160                    if !visited.contains(&node) && (best.is_none() || d < best.unwrap().1) {
161                        best = Some((node, d));
162                    }
163                }
164                best
165            };
166            match current {
167                None => break,
168                Some((node, d)) => {
169                    if d == Weight::INFINITY {
170                        break;
171                    }
172                    if node == end {
173                        let mut path = vec![end];
174                        let mut current = end;
175                        while let Some(&p) = prev.get(&current) {
176                            path.push(p);
177                            current = p;
178                        }
179                        path.reverse();
180                        return Some((path, d));
181                    }
182                    visited.insert(node);
183                    if let Some(neighbors) = self.neighbors(node) {
184                        for &(neighbor, weight) in neighbors {
185                            if visited.contains(&neighbor) {
186                                continue;
187                            }
188                            let alt = d + weight;
189                            if alt < dist[&neighbor] {
190                                dist.insert(neighbor, alt);
191                                prev.insert(neighbor, node);
192                            }
193                        }
194                    }
195                }
196            }
197        }
198        None
199    }
200
201    /// 拓扑排序(Kahn 算法)
202    ///
203    /// 返回拓扑排序结果。如果图有环返回 None。
204    pub fn topological_sort(&self) -> Option<Vec<NodeId>> {
205        let mut in_degree: HashMap<NodeId, usize> = HashMap::new();
206        for &node in self.adjacency.keys() {
207            in_degree.entry(node).or_insert(0);
208        }
209        for neighbors in self.adjacency.values() {
210            for &(neighbor, _) in neighbors {
211                *in_degree.entry(neighbor).or_insert(0) += 1;
212            }
213        }
214        let mut queue: VecDeque<NodeId> = in_degree
215            .iter()
216            .filter(|(_, &deg)| deg == 0)
217            .map(|(&n, _)| n)
218            .collect();
219        let mut result = Vec::new();
220        while let Some(node) = queue.pop_front() {
221            result.push(node);
222            if let Some(neighbors) = self.neighbors(node) {
223                for &(neighbor, _) in neighbors {
224                    if let Some(deg) = in_degree.get_mut(&neighbor) {
225                        *deg -= 1;
226                        if *deg == 0 {
227                            queue.push_back(neighbor);
228                        }
229                    }
230                }
231            }
232        }
233        if result.len() == self.node_count {
234            Some(result)
235        } else {
236            None
237        }
238    }
239
240    /// 环检测(DFS)
241    ///
242    /// 检测图中是否存在环。
243    pub fn has_cycle(&self) -> bool {
244        let mut visited = HashSet::new();
245        let mut rec_stack = HashSet::new();
246        for &node in self.adjacency.keys() {
247            if !visited.contains(&node) && self.has_cycle_dfs(node, &mut visited, &mut rec_stack) {
248                return true;
249            }
250        }
251        false
252    }
253
254    fn has_cycle_dfs(
255        &self,
256        node: NodeId,
257        visited: &mut HashSet<NodeId>,
258        rec_stack: &mut HashSet<NodeId>,
259    ) -> bool {
260        visited.insert(node);
261        rec_stack.insert(node);
262        if let Some(neighbors) = self.neighbors(node) {
263            for &(neighbor, _) in neighbors {
264                if !visited.contains(&neighbor) {
265                    if self.has_cycle_dfs(neighbor, visited, rec_stack) {
266                        return true;
267                    }
268                } else if rec_stack.contains(&neighbor) {
269                    return true;
270                }
271            }
272        }
273        rec_stack.remove(&node);
274        false
275    }
276
277    /// 连通分量(弱连通)
278    ///
279    /// 返回每个连通分量的节点列表。
280    pub fn connected_components(&self) -> Vec<Vec<NodeId>> {
281        let undirected = self.to_undirected();
282        let mut visited = HashSet::new();
283        let mut components = Vec::new();
284        for &node in undirected.adjacency.keys() {
285            if !visited.contains(&node) {
286                let component = undirected.bfs(node);
287                visited.extend(component.iter().copied());
288                components.push(component);
289            }
290        }
291        components
292    }
293
294    /// 转为无向图(忽略方向)
295    fn to_undirected(&self) -> DirectedGraph {
296        let mut undirected = DirectedGraph::new();
297        for (&node, neighbors) in &self.adjacency {
298            undirected.add_node(node);
299            for &(neighbor, weight) in neighbors {
300                undirected.add_edge(node, neighbor, weight);
301                undirected.add_edge(neighbor, node, weight);
302            }
303        }
304        undirected
305    }
306
307    /// 反转图(所有边方向取反)
308    pub fn reverse(&self) -> DirectedGraph {
309        let mut reversed = DirectedGraph::new();
310        for (&node, neighbors) in &self.adjacency {
311            reversed.add_node(node);
312            for &(neighbor, weight) in neighbors {
313                reversed.add_edge(neighbor, node, weight);
314            }
315        }
316        reversed
317    }
318
319    /// 节点的入度
320    pub fn in_degree(&self, node: NodeId) -> usize {
321        self.adjacency
322            .values()
323            .map(|v| v.iter().filter(|(n, _)| *n == node).count())
324            .sum()
325    }
326
327    /// 节点的出度
328    pub fn out_degree(&self, node: NodeId) -> usize {
329        self.adjacency.get(&node).map(|v| v.len()).unwrap_or(0)
330    }
331}
332
333/// 无向图
334#[derive(Debug, Clone, Default)]
335pub struct UndirectedGraph {
336    inner: DirectedGraph,
337}
338
339impl UndirectedGraph {
340    pub fn new() -> Self {
341        Self::default()
342    }
343
344    pub fn add_node(&mut self, node: NodeId) {
345        self.inner.add_node(node);
346    }
347
348    pub fn add_edge(&mut self, from: NodeId, to: NodeId, weight: Weight) {
349        self.inner.add_edge(from, to, weight);
350        self.inner.add_edge(to, from, weight);
351    }
352
353    pub fn add_edge_unweighted(&mut self, from: NodeId, to: NodeId) {
354        self.add_edge(from, to, 1.0);
355    }
356
357    pub fn node_count(&self) -> usize {
358        self.inner.node_count()
359    }
360
361    pub fn edge_count(&self) -> usize {
362        self.inner.edge_count() / 2
363    }
364
365    pub fn bfs(&self, start: NodeId) -> Vec<NodeId> {
366        self.inner.bfs(start)
367    }
368
369    pub fn dfs(&self, start: NodeId) -> Vec<NodeId> {
370        self.inner.dfs(start)
371    }
372
373    pub fn connected_components(&self) -> Vec<Vec<NodeId>> {
374        self.inner.connected_components()
375    }
376
377    pub fn has_node(&self, node: NodeId) -> bool {
378        self.inner.has_node(node)
379    }
380
381    pub fn has_edge(&self, from: NodeId, to: NodeId) -> bool {
382        self.inner.has_edge(from, to)
383    }
384
385    pub fn degree(&self, node: NodeId) -> usize {
386        self.inner.out_degree(node)
387    }
388}
389
390#[cfg(test)]
391mod tests {
392    use super::*;
393
394    #[test]
395    fn test_directed_graph_new() {
396        let g = DirectedGraph::new();
397        assert_eq!(g.node_count(), 0);
398        assert_eq!(g.edge_count(), 0);
399    }
400
401    #[test]
402    fn test_add_node() {
403        let mut g = DirectedGraph::new();
404        g.add_node(1);
405        assert_eq!(g.node_count(), 1);
406        assert!(g.has_node(1));
407    }
408
409    #[test]
410    fn test_add_edge() {
411        let mut g = DirectedGraph::new();
412        g.add_edge(1, 2, 3.15);
413        assert_eq!(g.node_count(), 2);
414        assert_eq!(g.edge_count(), 1);
415        assert!(g.has_edge(1, 2));
416        assert!(!g.has_edge(2, 1));
417        assert_eq!(g.edge_weight(1, 2), Some(3.15));
418    }
419
420    #[test]
421    fn test_add_edge_unweighted() {
422        let mut g = DirectedGraph::new();
423        g.add_edge_unweighted(1, 2);
424        assert_eq!(g.edge_weight(1, 2), Some(1.0));
425    }
426
427    #[test]
428    fn test_bfs_simple() {
429        let mut g = DirectedGraph::new();
430        g.add_edge_unweighted(1, 2);
431        g.add_edge_unweighted(1, 3);
432        g.add_edge_unweighted(2, 4);
433        g.add_edge_unweighted(3, 4);
434        let bfs = g.bfs(1);
435        assert_eq!(bfs[0], 1);
436        assert_eq!(bfs.len(), 4);
437        assert!(bfs.contains(&4));
438    }
439
440    #[test]
441    fn test_bfs_disconnected() {
442        let mut g = DirectedGraph::new();
443        g.add_edge_unweighted(1, 2);
444        g.add_node(3);
445        let bfs = g.bfs(1);
446        assert_eq!(bfs.len(), 2);
447        assert!(!bfs.contains(&3));
448    }
449
450    #[test]
451    fn test_bfs_nonexistent_start() {
452        let g = DirectedGraph::new();
453        assert!(g.bfs(1).is_empty());
454    }
455
456    #[test]
457    fn test_dfs_simple() {
458        let mut g = DirectedGraph::new();
459        g.add_edge_unweighted(1, 2);
460        g.add_edge_unweighted(2, 3);
461        g.add_edge_unweighted(3, 4);
462        let dfs = g.dfs(1);
463        assert_eq!(dfs.len(), 4);
464        assert_eq!(dfs[0], 1);
465    }
466
467    #[test]
468    fn test_dfs_nonexistent_start() {
469        let g = DirectedGraph::new();
470        assert!(g.dfs(1).is_empty());
471    }
472
473    #[test]
474    fn test_dijkstra_shortest_path() {
475        let mut g = DirectedGraph::new();
476        g.add_edge(1, 2, 1.0);
477        g.add_edge(2, 3, 2.0);
478        g.add_edge(1, 3, 5.0);
479        let (path, dist) = g.dijkstra(1, 3).unwrap();
480        assert_eq!(path, vec![1, 2, 3]);
481        assert!((dist - 3.0).abs() < 0.001);
482    }
483
484    #[test]
485    fn test_dijkstra_direct_edge() {
486        let mut g = DirectedGraph::new();
487        g.add_edge(1, 2, 5.0);
488        let (path, dist) = g.dijkstra(1, 2).unwrap();
489        assert_eq!(path, vec![1, 2]);
490        assert!((dist - 5.0).abs() < 0.001);
491    }
492
493    #[test]
494    fn test_dijkstra_unreachable() {
495        let mut g = DirectedGraph::new();
496        g.add_edge(1, 2, 1.0);
497        g.add_node(3);
498        assert!(g.dijkstra(1, 3).is_none());
499    }
500
501    #[test]
502    fn test_dijkstra_same_node() {
503        let mut g = DirectedGraph::new();
504        g.add_node(1);
505        let (path, dist) = g.dijkstra(1, 1).unwrap();
506        assert_eq!(path, vec![1]);
507        assert!((dist - 0.0).abs() < 0.001);
508    }
509
510    #[test]
511    fn test_topological_sort_dag() {
512        let mut g = DirectedGraph::new();
513        g.add_edge_unweighted(1, 2);
514        g.add_edge_unweighted(1, 3);
515        g.add_edge_unweighted(2, 4);
516        g.add_edge_unweighted(3, 4);
517        let topo = g.topological_sort().unwrap();
518        assert_eq!(topo.len(), 4);
519        let pos: HashMap<NodeId, usize> = topo.iter().enumerate().map(|(i, &n)| (n, i)).collect();
520        assert!(pos[&1] < pos[&2]);
521        assert!(pos[&1] < pos[&3]);
522        assert!(pos[&2] < pos[&4]);
523        assert!(pos[&3] < pos[&4]);
524    }
525
526    #[test]
527    fn test_topological_sort_with_cycle() {
528        let mut g = DirectedGraph::new();
529        g.add_edge_unweighted(1, 2);
530        g.add_edge_unweighted(2, 3);
531        g.add_edge_unweighted(3, 1);
532        assert!(g.topological_sort().is_none());
533    }
534
535    #[test]
536    fn test_has_cycle_no_cycle() {
537        let mut g = DirectedGraph::new();
538        g.add_edge_unweighted(1, 2);
539        g.add_edge_unweighted(2, 3);
540        assert!(!g.has_cycle());
541    }
542
543    #[test]
544    fn test_has_cycle_with_cycle() {
545        let mut g = DirectedGraph::new();
546        g.add_edge_unweighted(1, 2);
547        g.add_edge_unweighted(2, 3);
548        g.add_edge_unweighted(3, 1);
549        assert!(g.has_cycle());
550    }
551
552    #[test]
553    fn test_has_cycle_self_loop() {
554        let mut g = DirectedGraph::new();
555        g.add_edge_unweighted(1, 1);
556        assert!(g.has_cycle());
557    }
558
559    #[test]
560    fn test_connected_components() {
561        let mut g = DirectedGraph::new();
562        g.add_edge_unweighted(1, 2);
563        g.add_edge_unweighted(3, 4);
564        g.add_node(5);
565        let components = g.connected_components();
566        assert_eq!(components.len(), 3);
567    }
568
569    #[test]
570    fn test_connected_components_single() {
571        let mut g = DirectedGraph::new();
572        g.add_edge_unweighted(1, 2);
573        g.add_edge_unweighted(2, 3);
574        let components = g.connected_components();
575        assert_eq!(components.len(), 1);
576    }
577
578    #[test]
579    fn test_reverse() {
580        let mut g = DirectedGraph::new();
581        g.add_edge_unweighted(1, 2);
582        g.add_edge_unweighted(2, 3);
583        let reversed = g.reverse();
584        assert!(reversed.has_edge(2, 1));
585        assert!(reversed.has_edge(3, 2));
586        assert!(!reversed.has_edge(1, 2));
587    }
588
589    #[test]
590    fn test_in_degree() {
591        let mut g = DirectedGraph::new();
592        g.add_edge_unweighted(1, 3);
593        g.add_edge_unweighted(2, 3);
594        assert_eq!(g.in_degree(3), 2);
595        assert_eq!(g.in_degree(1), 0);
596    }
597
598    #[test]
599    fn test_out_degree() {
600        let mut g = DirectedGraph::new();
601        g.add_edge_unweighted(1, 2);
602        g.add_edge_unweighted(1, 3);
603        assert_eq!(g.out_degree(1), 2);
604        assert_eq!(g.out_degree(2), 0);
605    }
606
607    #[test]
608    fn test_undirected_graph() {
609        let mut g = UndirectedGraph::new();
610        g.add_edge_unweighted(1, 2);
611        g.add_edge_unweighted(2, 3);
612        assert_eq!(g.node_count(), 3);
613        assert_eq!(g.edge_count(), 2);
614        assert!(g.has_edge(1, 2));
615        assert!(g.has_edge(2, 1));
616    }
617
618    #[test]
619    fn test_undirected_connected_components() {
620        let mut g = UndirectedGraph::new();
621        g.add_edge_unweighted(1, 2);
622        g.add_edge_unweighted(3, 4);
623        let components = g.connected_components();
624        assert_eq!(components.len(), 2);
625    }
626
627    #[test]
628    fn test_undirected_degree() {
629        let mut g = UndirectedGraph::new();
630        g.add_edge_unweighted(1, 2);
631        g.add_edge_unweighted(1, 3);
632        assert_eq!(g.degree(1), 2);
633    }
634
635    #[test]
636    fn test_undirected_bfs() {
637        let mut g = UndirectedGraph::new();
638        g.add_edge_unweighted(1, 2);
639        g.add_edge_unweighted(2, 3);
640        let bfs = g.bfs(1);
641        assert_eq!(bfs.len(), 3);
642    }
643}