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            cycle_names: &mut Vec<&'static str>,
121        ) -> bool {
122            visited[node] = 1;
123            stack.push(node);
124
125            for (_dep_name, dep_id) in &entries[node].dependencies {
126                if let Some(&dep_idx) = index.get(dep_id) {
127                    if visited[dep_idx] == 1 {
128                        // Found cycle — extract it
129                        let Some(start) = stack.iter().position(|&x| x == dep_idx) else {
130                            // Invariant violation: dep_idx should be in the
131                            // stack when visited[dep_idx] == 1. Fall back to
132                            // a generic cycle report instead of panicking.
133                            cycle_names.push(entries[dep_idx].name);
134                            cycle_names.push(entries[node].name);
135                            return true;
136                        };
137                        for &idx in &stack[start..] {
138                            cycle_names.push(entries[idx].name);
139                        }
140                        cycle_names.push(entries[dep_idx].name);
141                        return true;
142                    }
143                    if visited[dep_idx] == 0
144                        && dfs(dep_idx, entries, index, visited, stack, cycle_names)
145                    {
146                        return true;
147                    }
148                }
149            }
150
151            stack.pop();
152            visited[node] = 2;
153            false
154        }
155
156        let n = self.entries.len();
157        let mut visited = vec![0u8; n]; // 0=unvisited, 1=in-stack, 2=done
158        let mut stack = Vec::new();
159        let mut cycle_names = Vec::new();
160
161        for i in 0..n {
162            if visited[i] == 0
163                && dfs(
164                    i,
165                    &self.entries,
166                    &self.index,
167                    &mut visited,
168                    &mut stack,
169                    &mut cycle_names,
170                )
171            {
172                return cycle_names;
173            }
174        }
175
176        vec!["<unknown cycle>"]
177    }
178
179    /// Get the registered names of all dependencies for a module.
180    #[must_use]
181    pub fn dependency_names(&self, type_id: TypeId) -> Vec<&'static str> {
182        if let Some(&idx) = self.index.get(&type_id) {
183            self.entries[idx]
184                .dependencies
185                .iter()
186                .map(|(name, _)| *name)
187                .collect()
188        } else {
189            Vec::new()
190        }
191    }
192
193    /// Get all entries in registration order.
194    #[must_use]
195    pub fn entries(&self) -> &[ModuleEntry] {
196        &self.entries
197    }
198
199    /// Look up a module's diagnostic name by `TypeId` in O(1).
200    #[must_use]
201    pub fn name_of(&self, type_id: TypeId) -> Option<&'static str> {
202        self.index.get(&type_id).map(|&idx| self.entries[idx].name)
203    }
204
205    /// Export the dependency graph as a Graphviz DOT format string.
206    ///
207    /// Nodes are module names; directed edges represent dependencies
208    /// (dependency → dependent).
209    #[must_use]
210    pub fn to_dot(&self) -> String {
211        use std::fmt::Write as _;
212        if self.entries.is_empty() {
213            return "digraph {}".to_string();
214        }
215        let mut out = String::from("digraph {\n");
216        // Nodes
217        for entry in &self.entries {
218            let _ = writeln!(out, "    \"{}\";", entry.name);
219        }
220        // Edges: dependency → dependent
221        for entry in &self.entries {
222            for (dep_name, _) in &entry.dependencies {
223                let _ = writeln!(out, "    \"{}\" -> \"{}\";", dep_name, entry.name);
224            }
225        }
226        out.push('}');
227        out
228    }
229
230    /// Export the dependency graph as a Mermaid flowchart format string.
231    ///
232    /// Uses `graph TD` (top-down) layout. Edges: dependency --> dependent.
233    #[must_use]
234    pub fn to_mermaid(&self) -> String {
235        use std::fmt::Write as _;
236        if self.entries.is_empty() {
237            return "graph TD".to_string();
238        }
239        let mut out = String::from("graph TD\n");
240        for entry in &self.entries {
241            for (dep_name, _) in &entry.dependencies {
242                // Mermaid node IDs: replace hyphens with underscores
243                let from_id = dep_name.replace('-', "_");
244                let to_id = entry.name.replace('-', "_");
245                let _ = writeln!(
246                    out,
247                    "    {}[\"{}\"] --> {}[\"{}\"]",
248                    from_id, dep_name, to_id, entry.name
249                );
250            }
251        }
252        // Ensure nodes with no dependencies still appear
253        for entry in &self.entries {
254            if entry.dependencies.is_empty() {
255                let id = entry.name.replace('-', "_");
256                let _ = writeln!(out, "    {}[\"{}\"]", id, entry.name);
257            }
258        }
259        out
260    }
261}
262
263impl Default for DependencyGraph {
264    fn default() -> Self {
265        Self::new()
266    }
267}
268
269/// Errors from graph validation.
270#[derive(Debug, Clone, PartialEq, Eq)]
271pub enum GraphError {
272    /// A module depends on an unregistered module.
273    DependencyMissing {
274        module: &'static str,
275        missing: &'static str,
276    },
277    /// A dependency cycle was detected.
278    CycleDetected { cycle: Vec<&'static str> },
279}
280
281#[cfg(test)]
282mod tests {
283    use super::*;
284    use std::any::TypeId;
285
286    /// Each module needs a unique TypeId, so we use distinct zero-sized types.
287    mod types {
288        pub struct A;
289        pub struct B;
290        pub struct C;
291        pub struct D;
292    }
293
294    fn typed_entry<T: 'static>(
295        name: &'static str,
296        deps: Vec<(&'static str, TypeId)>,
297    ) -> ModuleEntry {
298        ModuleEntry {
299            type_id: TypeId::of::<T>(),
300            name,
301            dependencies: deps,
302        }
303    }
304
305    #[test]
306    fn graph_new_is_empty() {
307        let g = DependencyGraph::new();
308        assert!(g.entries().is_empty());
309    }
310
311    #[test]
312    fn graph_add_and_entries() {
313        let mut g = DependencyGraph::new();
314        g.add(typed_entry::<types::A>("a", vec![])).unwrap();
315        assert_eq!(g.entries().len(), 1);
316    }
317
318    #[test]
319    fn graph_add_duplicate_returns_err() {
320        let mut g = DependencyGraph::new();
321        g.add(typed_entry::<types::A>("a", vec![])).unwrap();
322        let err = g.add(typed_entry::<types::A>("a2", vec![])).unwrap_err();
323        assert_eq!(err, "a2");
324    }
325
326    #[test]
327    fn graph_validate_empty_succeeds() {
328        let g = DependencyGraph::new();
329        let sorted = g.validate().unwrap();
330        assert!(sorted.is_empty());
331    }
332
333    #[test]
334    fn graph_validate_single_node() {
335        let mut g = DependencyGraph::new();
336        g.add(typed_entry::<types::A>("a", vec![])).unwrap();
337        let sorted = g.validate().unwrap();
338        assert_eq!(sorted.len(), 1);
339    }
340
341    #[test]
342    fn graph_validate_missing_dependency() {
343        let mut g = DependencyGraph::new();
344        g.add(typed_entry::<types::A>(
345            "a",
346            vec![("b", TypeId::of::<types::B>())],
347        ))
348        .unwrap();
349        let err = g.validate().unwrap_err();
350        assert!(matches!(
351            err,
352            GraphError::DependencyMissing {
353                module: "a",
354                missing: "b"
355            }
356        ));
357    }
358
359    #[test]
360    fn graph_validate_cycle_two_nodes() {
361        let mut g = DependencyGraph::new();
362        g.add(typed_entry::<types::A>(
363            "a",
364            vec![("b", TypeId::of::<types::B>())],
365        ))
366        .unwrap();
367        g.add(typed_entry::<types::B>(
368            "b",
369            vec![("a", TypeId::of::<types::A>())],
370        ))
371        .unwrap();
372        let err = g.validate().unwrap_err();
373        assert!(matches!(err, GraphError::CycleDetected { .. }));
374        if let GraphError::CycleDetected { cycle } = err {
375            assert!(cycle.len() >= 2, "cycle should contain at least 2 names");
376        }
377    }
378
379    #[test]
380    fn graph_validate_cycle_three_nodes() {
381        let mut g = DependencyGraph::new();
382        g.add(typed_entry::<types::A>(
383            "a",
384            vec![("b", TypeId::of::<types::B>())],
385        ))
386        .unwrap();
387        g.add(typed_entry::<types::B>(
388            "b",
389            vec![("c", TypeId::of::<types::C>())],
390        ))
391        .unwrap();
392        g.add(typed_entry::<types::C>(
393            "c",
394            vec![("a", TypeId::of::<types::A>())],
395        ))
396        .unwrap();
397        let err = g.validate().unwrap_err();
398        assert!(matches!(err, GraphError::CycleDetected { .. }));
399    }
400
401    #[test]
402    fn graph_validate_topo_order() {
403        let mut g = DependencyGraph::new();
404        g.add(typed_entry::<types::A>("a", vec![])).unwrap();
405        g.add(typed_entry::<types::B>(
406            "b",
407            vec![("a", TypeId::of::<types::A>())],
408        ))
409        .unwrap();
410        let sorted = g.validate().unwrap();
411        let a_idx = sorted
412            .iter()
413            .position(|t| *t == TypeId::of::<types::A>())
414            .unwrap();
415        let b_idx = sorted
416            .iter()
417            .position(|t| *t == TypeId::of::<types::B>())
418            .unwrap();
419        assert!(a_idx < b_idx, "a should be sorted before b");
420    }
421
422    #[test]
423    fn graph_dependency_names() {
424        let mut g = DependencyGraph::new();
425        g.add(typed_entry::<types::A>("a", vec![])).unwrap();
426        g.add(typed_entry::<types::B>(
427            "b",
428            vec![("a", TypeId::of::<types::A>())],
429        ))
430        .unwrap();
431        let names = g.dependency_names(TypeId::of::<types::B>());
432        assert_eq!(names, vec!["a"]);
433    }
434
435    #[test]
436    fn graph_dependency_names_unknown_returns_empty() {
437        let g = DependencyGraph::new();
438        let names = g.dependency_names(TypeId::of::<types::A>());
439        assert!(names.is_empty());
440    }
441
442    #[test]
443    fn graph_name_of() {
444        let mut g = DependencyGraph::new();
445        g.add(typed_entry::<types::A>("module-a", vec![])).unwrap();
446        assert_eq!(g.name_of(TypeId::of::<types::A>()), Some("module-a"));
447        assert_eq!(g.name_of(TypeId::of::<types::B>()), None);
448    }
449
450    #[test]
451    fn graph_default_is_empty() {
452        let g = DependencyGraph::default();
453        assert!(g.entries().is_empty());
454    }
455
456    #[test]
457    fn graph_to_dot_empty() {
458        let g = DependencyGraph::new();
459        assert_eq!(g.to_dot(), "digraph {}");
460    }
461
462    #[test]
463    fn graph_to_dot_with_nodes_and_edges() {
464        let mut g = DependencyGraph::new();
465        g.add(typed_entry::<types::A>("a", vec![])).unwrap();
466        g.add(typed_entry::<types::B>(
467            "b",
468            vec![("a", TypeId::of::<types::A>())],
469        ))
470        .unwrap();
471        let dot = g.to_dot();
472        assert!(dot.starts_with("digraph {"));
473        assert!(dot.contains("\"a\""));
474        assert!(dot.contains("\"b\""));
475        assert!(dot.contains("\"a\" -> \"b\""));
476        assert!(dot.ends_with('}'));
477    }
478
479    #[test]
480    fn graph_to_mermaid_empty() {
481        let g = DependencyGraph::new();
482        assert_eq!(g.to_mermaid(), "graph TD");
483    }
484
485    #[test]
486    fn graph_to_mermaid_with_nodes_and_edges() {
487        let mut g = DependencyGraph::new();
488        g.add(typed_entry::<types::A>("a", vec![])).unwrap();
489        g.add(typed_entry::<types::B>(
490            "b",
491            vec![("a", TypeId::of::<types::A>())],
492        ))
493        .unwrap();
494        let mermaid = g.to_mermaid();
495        assert!(mermaid.starts_with("graph TD"));
496        assert!(mermaid.contains("a[\"a\"]"));
497        assert!(mermaid.contains("b[\"b\"]"));
498        assert!(mermaid.contains("-->"));
499    }
500
501    #[test]
502    fn graph_to_mermaid_hyphen_replacement() {
503        let mut g = DependencyGraph::new();
504        g.add(typed_entry::<types::A>("my-module", vec![])).unwrap();
505        g.add(typed_entry::<types::B>(
506            "my-dep",
507            vec![("my-module", TypeId::of::<types::A>())],
508        ))
509        .unwrap();
510        let mermaid = g.to_mermaid();
511        // Hyphens in names should be replaced with underscores for node IDs
512        assert!(mermaid.contains("my_module"));
513        assert!(mermaid.contains("my_dep"));
514    }
515
516    #[test]
517    fn graph_error_debug() {
518        let err = GraphError::DependencyMissing {
519            module: "a",
520            missing: "b",
521        };
522        let debug = format!("{err:?}");
523        assert!(debug.contains("DependencyMissing"));
524
525        let err2 = GraphError::CycleDetected {
526            cycle: vec!["a", "b", "a"],
527        };
528        let debug2 = format!("{err2:?}");
529        assert!(debug2.contains("CycleDetected"));
530    }
531}