Skip to main content

trait_kit/kit/
graph.rs

1// Copyright (c) 2026 Kirky.X
2// SPDX-License-Identifier: MIT
3//! Dependency graph with topological sort and cycle detection.
4
5use std::any::TypeId;
6use std::collections::{HashMap, VecDeque};
7
8/// A node in the dependency graph.
9#[derive(Debug, Clone)]
10pub struct ModuleEntry {
11    /// The module's `TypeId`.
12    pub type_id: TypeId,
13    /// The module's diagnostic name.
14    pub name: &'static str,
15    /// (name, `TypeId`) pairs of modules this module depends on.
16    pub dependencies: Vec<(&'static str, TypeId)>,
17}
18
19/// Dependency graph for topological sort and cycle detection.
20#[derive(Debug)]
21pub struct DependencyGraph {
22    entries: Vec<ModuleEntry>,
23    index: HashMap<TypeId, usize>,
24}
25
26impl DependencyGraph {
27    /// Create an empty graph.
28    #[must_use]
29    pub fn new() -> Self {
30        DependencyGraph {
31            entries: Vec::new(),
32            index: HashMap::new(),
33        }
34    }
35
36    /// Add a module to the graph.
37    ///
38    /// # Errors
39    ///
40    /// Returns the module's name if it is already registered.
41    pub fn add(&mut self, entry: ModuleEntry) -> Result<(), &'static str> {
42        if self.index.contains_key(&entry.type_id) {
43            return Err(entry.name);
44        }
45        let idx = self.entries.len();
46        self.index.insert(entry.type_id, idx);
47        self.entries.push(entry);
48        Ok(())
49    }
50
51    /// Validate the graph: check for missing dependencies and cycles.
52    /// Returns the topologically sorted `TypeIds` on success.
53    ///
54    /// # Errors
55    ///
56    /// Returns `GraphError::DependencyMissing` if a module depends on an unregistered module.
57    /// Returns `GraphError::CycleDetected` if a dependency cycle is found.
58    pub fn validate(&self) -> Result<Vec<TypeId>, GraphError> {
59        // Check for missing dependencies
60        for entry in &self.entries {
61            for (dep_name, dep_id) in &entry.dependencies {
62                if !self.index.contains_key(dep_id) {
63                    return Err(GraphError::DependencyMissing {
64                        module: entry.name,
65                        missing: dep_name,
66                    });
67                }
68            }
69        }
70
71        // Kahn's algorithm for topological sort + cycle detection
72        let n = self.entries.len();
73        let mut in_degree = vec![0usize; n];
74        let mut adj: Vec<Vec<usize>> = vec![Vec::new(); n];
75
76        for (i, entry) in self.entries.iter().enumerate() {
77            for (_dep_name, dep_id) in &entry.dependencies {
78                if let Some(&dep_idx) = self.index.get(dep_id) {
79                    adj[dep_idx].push(i);
80                    in_degree[i] += 1;
81                }
82            }
83        }
84
85        let mut queue: VecDeque<usize> = VecDeque::new();
86        for (i, deg) in in_degree.iter().enumerate() {
87            if *deg == 0 {
88                queue.push_back(i);
89            }
90        }
91
92        let mut sorted = Vec::with_capacity(n);
93        while let Some(node) = queue.pop_front() {
94            sorted.push(self.entries[node].type_id);
95            for &neighbor in &adj[node] {
96                in_degree[neighbor] -= 1;
97                if in_degree[neighbor] == 0 {
98                    queue.push_back(neighbor);
99                }
100            }
101        }
102
103        if sorted.len() != n {
104            // Cycle detected — find the cycle for a useful error message
105            let cycle = self.find_cycle();
106            return Err(GraphError::CycleDetected { cycle });
107        }
108
109        Ok(sorted)
110    }
111
112    /// Find a cycle in the graph using DFS (for error reporting).
113    fn find_cycle(&self) -> Vec<&'static str> {
114        fn dfs(
115            node: usize,
116            entries: &[ModuleEntry],
117            index: &HashMap<TypeId, usize>,
118            visited: &mut [u8],
119            stack: &mut Vec<usize>,
120            stack_pos: &mut HashMap<usize, usize>,
121            cycle_names: &mut Vec<&'static str>,
122        ) -> bool {
123            visited[node] = 1;
124            stack_pos.insert(node, stack.len());
125            stack.push(node);
126
127            for (_dep_name, dep_id) in &entries[node].dependencies {
128                if let Some(&dep_idx) = index.get(dep_id) {
129                    if visited[dep_idx] == 1 {
130                        // Found cycle — O(1) lookup via stack_pos map
131                        let Some(&start) = stack_pos.get(&dep_idx) else {
132                            // Invariant violation: dep_idx should be in the
133                            // stack when visited[dep_idx] == 1. Fall back to
134                            // a generic cycle report instead of panicking.
135                            cycle_names.push(entries[dep_idx].name);
136                            cycle_names.push(entries[node].name);
137                            return true;
138                        };
139                        for &idx in &stack[start..] {
140                            cycle_names.push(entries[idx].name);
141                        }
142                        cycle_names.push(entries[dep_idx].name);
143                        return true;
144                    }
145                    if visited[dep_idx] == 0
146                        && dfs(
147                            dep_idx,
148                            entries,
149                            index,
150                            visited,
151                            stack,
152                            stack_pos,
153                            cycle_names,
154                        )
155                    {
156                        return true;
157                    }
158                }
159            }
160
161            stack.pop();
162            stack_pos.remove(&node);
163            visited[node] = 2;
164            false
165        }
166
167        let n = self.entries.len();
168        let mut visited = vec![0u8; n]; // 0=unvisited, 1=in-stack, 2=done
169        let mut stack = Vec::with_capacity(n);
170        let mut stack_pos = HashMap::with_capacity(n);
171        let mut cycle_names = Vec::new();
172
173        for i in 0..n {
174            if visited[i] == 0
175                && dfs(
176                    i,
177                    &self.entries,
178                    &self.index,
179                    &mut visited,
180                    &mut stack,
181                    &mut stack_pos,
182                    &mut cycle_names,
183                )
184            {
185                return cycle_names;
186            }
187        }
188
189        vec!["<unknown cycle>"]
190    }
191
192    /// Get the registered names of all dependencies for a module.
193    #[must_use]
194    pub fn dependency_names(&self, type_id: TypeId) -> Vec<&'static str> {
195        if let Some(&idx) = self.index.get(&type_id) {
196            self.entries[idx]
197                .dependencies
198                .iter()
199                .map(|(name, _)| *name)
200                .collect()
201        } else {
202            Vec::new()
203        }
204    }
205
206    /// Get all entries in registration order.
207    #[must_use]
208    pub fn entries(&self) -> &[ModuleEntry] {
209        &self.entries
210    }
211
212    /// Look up a module's diagnostic name by `TypeId` in O(1).
213    #[must_use]
214    pub fn name_of(&self, type_id: TypeId) -> Option<&'static str> {
215        self.index.get(&type_id).map(|&idx| self.entries[idx].name)
216    }
217
218    /// Export the dependency graph as a Graphviz DOT format string.
219    ///
220    /// Nodes are module names; directed edges represent dependencies
221    /// (dependency → dependent).
222    #[must_use]
223    pub fn to_dot(&self) -> String {
224        use std::fmt::Write as _;
225        if self.entries.is_empty() {
226            return "digraph {}".to_string();
227        }
228        let mut out = String::from("digraph {\n");
229        // Nodes
230        for entry in &self.entries {
231            let _ = writeln!(out, "    \"{}\";", entry.name);
232        }
233        // Edges: dependency → dependent
234        for entry in &self.entries {
235            for (dep_name, _) in &entry.dependencies {
236                let _ = writeln!(out, "    \"{}\" -> \"{}\";", dep_name, entry.name);
237            }
238        }
239        out.push('}');
240        out
241    }
242
243    /// Export the dependency graph as a Mermaid flowchart format string.
244    ///
245    /// Uses `graph TD` (top-down) layout. Edges: dependency --> dependent.
246    #[must_use]
247    pub fn to_mermaid(&self) -> String {
248        use std::fmt::Write as _;
249        if self.entries.is_empty() {
250            return "graph TD".to_string();
251        }
252        let mut out = String::from("graph TD\n");
253        // Use index-based node IDs to avoid collisions when names contain
254        // hyphens or other special characters (e.g. 'my-module' vs 'my_module').
255        for (idx, entry) in self.entries.iter().enumerate() {
256            for (dep_name, _) in &entry.dependencies {
257                // Find the index of the dependency entry for its node ID
258                let dep_idx = self
259                    .entries
260                    .iter()
261                    .position(|e| e.name == *dep_name)
262                    .unwrap_or(idx);
263                let _ = writeln!(
264                    out,
265                    "    n{dep_idx}[\"{dep_name}\"] --> n{idx}[\"{}\"]",
266                    entry.name
267                );
268            }
269        }
270        // Ensure nodes with no dependencies still appear
271        for (idx, entry) in self.entries.iter().enumerate() {
272            if entry.dependencies.is_empty() {
273                let _ = writeln!(out, "    n{idx}[\"{}\"]", entry.name);
274            }
275        }
276        out
277    }
278}
279
280impl Default for DependencyGraph {
281    fn default() -> Self {
282        Self::new()
283    }
284}
285
286/// Errors from graph validation.
287#[derive(Debug, Clone, PartialEq, Eq)]
288pub enum GraphError {
289    /// A module depends on an unregistered module.
290    DependencyMissing {
291        module: &'static str,
292        missing: &'static str,
293    },
294    /// A dependency cycle was detected.
295    CycleDetected { cycle: Vec<&'static str> },
296}
297
298#[cfg(test)]
299mod tests {
300    use super::*;
301    use std::any::TypeId;
302
303    /// Each module needs a unique `TypeId`, so we use distinct zero-sized types.
304    mod types {
305        pub struct A;
306        pub struct B;
307        pub struct C;
308    }
309
310    fn typed_entry<T: 'static>(
311        name: &'static str,
312        deps: Vec<(&'static str, TypeId)>,
313    ) -> ModuleEntry {
314        ModuleEntry {
315            type_id: TypeId::of::<T>(),
316            name,
317            dependencies: deps,
318        }
319    }
320
321    #[test]
322    fn graph_new_is_empty() {
323        let g = DependencyGraph::new();
324        assert!(g.entries().is_empty());
325    }
326
327    #[test]
328    fn graph_add_and_entries() {
329        let mut g = DependencyGraph::new();
330        g.add(typed_entry::<types::A>("a", vec![])).unwrap();
331        assert_eq!(g.entries().len(), 1);
332    }
333
334    #[test]
335    fn graph_add_duplicate_returns_err() {
336        let mut g = DependencyGraph::new();
337        g.add(typed_entry::<types::A>("a", vec![])).unwrap();
338        let err = g.add(typed_entry::<types::A>("a2", vec![])).unwrap_err();
339        assert_eq!(err, "a2");
340    }
341
342    #[test]
343    fn graph_validate_empty_succeeds() {
344        let g = DependencyGraph::new();
345        let sorted = g.validate().unwrap();
346        assert!(sorted.is_empty());
347    }
348
349    #[test]
350    fn graph_validate_single_node() {
351        let mut g = DependencyGraph::new();
352        g.add(typed_entry::<types::A>("a", vec![])).unwrap();
353        let sorted = g.validate().unwrap();
354        assert_eq!(sorted.len(), 1);
355    }
356
357    #[test]
358    fn graph_validate_missing_dependency() {
359        let mut g = DependencyGraph::new();
360        g.add(typed_entry::<types::A>(
361            "a",
362            vec![("b", TypeId::of::<types::B>())],
363        ))
364        .unwrap();
365        let err = g.validate().unwrap_err();
366        assert!(matches!(
367            err,
368            GraphError::DependencyMissing {
369                module: "a",
370                missing: "b"
371            }
372        ));
373    }
374
375    #[test]
376    fn graph_validate_cycle_two_nodes() {
377        let mut g = DependencyGraph::new();
378        g.add(typed_entry::<types::A>(
379            "a",
380            vec![("b", TypeId::of::<types::B>())],
381        ))
382        .unwrap();
383        g.add(typed_entry::<types::B>(
384            "b",
385            vec![("a", TypeId::of::<types::A>())],
386        ))
387        .unwrap();
388        let err = g.validate().unwrap_err();
389        assert!(matches!(err, GraphError::CycleDetected { .. }));
390        if let GraphError::CycleDetected { cycle } = err {
391            assert!(cycle.len() >= 2, "cycle should contain at least 2 names");
392        }
393    }
394
395    #[test]
396    fn graph_validate_cycle_three_nodes() {
397        let mut g = DependencyGraph::new();
398        g.add(typed_entry::<types::A>(
399            "a",
400            vec![("b", TypeId::of::<types::B>())],
401        ))
402        .unwrap();
403        g.add(typed_entry::<types::B>(
404            "b",
405            vec![("c", TypeId::of::<types::C>())],
406        ))
407        .unwrap();
408        g.add(typed_entry::<types::C>(
409            "c",
410            vec![("a", TypeId::of::<types::A>())],
411        ))
412        .unwrap();
413        let err = g.validate().unwrap_err();
414        assert!(matches!(err, GraphError::CycleDetected { .. }));
415    }
416
417    #[test]
418    fn graph_validate_topo_order() {
419        let mut g = DependencyGraph::new();
420        g.add(typed_entry::<types::A>("a", vec![])).unwrap();
421        g.add(typed_entry::<types::B>(
422            "b",
423            vec![("a", TypeId::of::<types::A>())],
424        ))
425        .unwrap();
426        let sorted = g.validate().unwrap();
427        let a_idx = sorted
428            .iter()
429            .position(|t| *t == TypeId::of::<types::A>())
430            .unwrap();
431        let b_idx = sorted
432            .iter()
433            .position(|t| *t == TypeId::of::<types::B>())
434            .unwrap();
435        assert!(a_idx < b_idx, "a should be sorted before b");
436    }
437
438    #[test]
439    fn graph_dependency_names() {
440        let mut g = DependencyGraph::new();
441        g.add(typed_entry::<types::A>("a", vec![])).unwrap();
442        g.add(typed_entry::<types::B>(
443            "b",
444            vec![("a", TypeId::of::<types::A>())],
445        ))
446        .unwrap();
447        let names = g.dependency_names(TypeId::of::<types::B>());
448        assert_eq!(names, vec!["a"]);
449    }
450
451    #[test]
452    fn graph_dependency_names_unknown_returns_empty() {
453        let g = DependencyGraph::new();
454        let names = g.dependency_names(TypeId::of::<types::A>());
455        assert!(names.is_empty());
456    }
457
458    #[test]
459    fn graph_name_of() {
460        let mut g = DependencyGraph::new();
461        g.add(typed_entry::<types::A>("module-a", vec![])).unwrap();
462        assert_eq!(g.name_of(TypeId::of::<types::A>()), Some("module-a"));
463        assert_eq!(g.name_of(TypeId::of::<types::B>()), None);
464    }
465
466    #[test]
467    fn graph_default_is_empty() {
468        let g = DependencyGraph::default();
469        assert!(g.entries().is_empty());
470    }
471
472    #[test]
473    fn graph_to_dot_empty() {
474        let g = DependencyGraph::new();
475        assert_eq!(g.to_dot(), "digraph {}");
476    }
477
478    #[test]
479    fn graph_to_dot_with_nodes_and_edges() {
480        let mut g = DependencyGraph::new();
481        g.add(typed_entry::<types::A>("a", vec![])).unwrap();
482        g.add(typed_entry::<types::B>(
483            "b",
484            vec![("a", TypeId::of::<types::A>())],
485        ))
486        .unwrap();
487        let dot = g.to_dot();
488        assert!(dot.starts_with("digraph {"));
489        assert!(dot.contains("\"a\""));
490        assert!(dot.contains("\"b\""));
491        assert!(dot.contains("\"a\" -> \"b\""));
492        assert!(dot.ends_with('}'));
493    }
494
495    #[test]
496    fn graph_to_mermaid_empty() {
497        let g = DependencyGraph::new();
498        assert_eq!(g.to_mermaid(), "graph TD");
499    }
500
501    #[test]
502    fn graph_to_mermaid_with_nodes_and_edges() {
503        let mut g = DependencyGraph::new();
504        g.add(typed_entry::<types::A>("a", vec![])).unwrap();
505        g.add(typed_entry::<types::B>(
506            "b",
507            vec![("a", TypeId::of::<types::A>())],
508        ))
509        .unwrap();
510        let mermaid = g.to_mermaid();
511        assert!(mermaid.starts_with("graph TD"));
512        // Index-based node IDs: n0 for "a", n1 for "b"
513        assert!(mermaid.contains("n0[\"a\"]"));
514        assert!(mermaid.contains("n1[\"b\"]"));
515        assert!(mermaid.contains("-->"));
516    }
517
518    #[test]
519    fn graph_to_mermaid_hyphen_names_no_collision() {
520        let mut g = DependencyGraph::new();
521        g.add(typed_entry::<types::A>("my-module", vec![])).unwrap();
522        g.add(typed_entry::<types::B>(
523            "my-dep",
524            vec![("my-module", TypeId::of::<types::A>())],
525        ))
526        .unwrap();
527        let mermaid = g.to_mermaid();
528        // Index-based IDs avoid collision between hyphens and underscores
529        assert!(mermaid.contains("n0[\"my-module\"]"));
530        assert!(mermaid.contains("n1[\"my-dep\"]"));
531        // Original names (with hyphens) are preserved in display labels
532        assert!(mermaid.contains("my-module"));
533        assert!(mermaid.contains("my-dep"));
534    }
535
536    #[test]
537    fn graph_error_debug() {
538        let err = GraphError::DependencyMissing {
539            module: "a",
540            missing: "b",
541        };
542        let debug = format!("{err:?}");
543        assert!(debug.contains("DependencyMissing"));
544
545        let err2 = GraphError::CycleDetected {
546            cycle: vec!["a", "b", "a"],
547        };
548        let debug2 = format!("{err2:?}");
549        assert!(debug2.contains("CycleDetected"));
550    }
551}