use std::any::TypeId;
use std::collections::{HashMap, VecDeque};
#[derive(Debug, Clone)]
pub struct ModuleEntry {
pub type_id: TypeId,
pub name: &'static str,
pub dependencies: Vec<(&'static str, TypeId)>,
}
#[derive(Debug)]
pub struct DependencyGraph {
entries: Vec<ModuleEntry>,
index: HashMap<TypeId, usize>,
}
impl DependencyGraph {
#[must_use]
pub fn new() -> Self {
DependencyGraph {
entries: Vec::new(),
index: HashMap::new(),
}
}
pub fn add(&mut self, entry: ModuleEntry) -> Result<(), &'static str> {
if self.index.contains_key(&entry.type_id) {
return Err(entry.name);
}
let idx = self.entries.len();
self.index.insert(entry.type_id, idx);
self.entries.push(entry);
Ok(())
}
pub fn validate(&self) -> Result<Vec<TypeId>, GraphError> {
for entry in &self.entries {
for (dep_name, dep_id) in &entry.dependencies {
if !self.index.contains_key(dep_id) {
return Err(GraphError::DependencyMissing {
module: entry.name,
missing: dep_name,
});
}
}
}
let n = self.entries.len();
let mut in_degree = vec![0usize; n];
let mut adj: Vec<Vec<usize>> = vec![Vec::new(); n];
for (i, entry) in self.entries.iter().enumerate() {
for (_dep_name, dep_id) in &entry.dependencies {
if let Some(&dep_idx) = self.index.get(dep_id) {
adj[dep_idx].push(i);
in_degree[i] += 1;
}
}
}
let mut queue: VecDeque<usize> = VecDeque::new();
for (i, deg) in in_degree.iter().enumerate() {
if *deg == 0 {
queue.push_back(i);
}
}
let mut sorted = Vec::with_capacity(n);
while let Some(node) = queue.pop_front() {
sorted.push(self.entries[node].type_id);
for &neighbor in &adj[node] {
in_degree[neighbor] -= 1;
if in_degree[neighbor] == 0 {
queue.push_back(neighbor);
}
}
}
if sorted.len() != n {
let cycle = self.find_cycle();
return Err(GraphError::CycleDetected { cycle });
}
Ok(sorted)
}
fn find_cycle(&self) -> Vec<&'static str> {
fn dfs(
node: usize,
entries: &[ModuleEntry],
index: &HashMap<TypeId, usize>,
visited: &mut [u8],
stack: &mut Vec<usize>,
stack_pos: &mut HashMap<usize, usize>,
cycle_names: &mut Vec<&'static str>,
) -> bool {
visited[node] = 1;
stack_pos.insert(node, stack.len());
stack.push(node);
for (_dep_name, dep_id) in &entries[node].dependencies {
if let Some(&dep_idx) = index.get(dep_id) {
if visited[dep_idx] == 1 {
let Some(&start) = stack_pos.get(&dep_idx) else {
cycle_names.push(entries[dep_idx].name);
cycle_names.push(entries[node].name);
return true;
};
for &idx in &stack[start..] {
cycle_names.push(entries[idx].name);
}
cycle_names.push(entries[dep_idx].name);
return true;
}
if visited[dep_idx] == 0
&& dfs(
dep_idx,
entries,
index,
visited,
stack,
stack_pos,
cycle_names,
)
{
return true;
}
}
}
stack.pop();
stack_pos.remove(&node);
visited[node] = 2;
false
}
let n = self.entries.len();
let mut visited = vec![0u8; n]; let mut stack = Vec::with_capacity(n);
let mut stack_pos = HashMap::with_capacity(n);
let mut cycle_names = Vec::new();
for i in 0..n {
if visited[i] == 0
&& dfs(
i,
&self.entries,
&self.index,
&mut visited,
&mut stack,
&mut stack_pos,
&mut cycle_names,
)
{
return cycle_names;
}
}
vec!["<unknown cycle>"]
}
#[must_use]
pub fn dependency_names(&self, type_id: TypeId) -> Vec<&'static str> {
if let Some(&idx) = self.index.get(&type_id) {
self.entries[idx]
.dependencies
.iter()
.map(|(name, _)| *name)
.collect()
} else {
Vec::new()
}
}
#[must_use]
pub fn entries(&self) -> &[ModuleEntry] {
&self.entries
}
#[must_use]
pub fn name_of(&self, type_id: TypeId) -> Option<&'static str> {
self.index.get(&type_id).map(|&idx| self.entries[idx].name)
}
#[must_use]
pub fn to_dot(&self) -> String {
use std::fmt::Write as _;
if self.entries.is_empty() {
return "digraph {}".to_string();
}
let mut out = String::from("digraph {\n");
for entry in &self.entries {
let _ = writeln!(out, " \"{}\";", entry.name);
}
for entry in &self.entries {
for (dep_name, _) in &entry.dependencies {
let _ = writeln!(out, " \"{}\" -> \"{}\";", dep_name, entry.name);
}
}
out.push('}');
out
}
#[must_use]
pub fn to_mermaid(&self) -> String {
use std::fmt::Write as _;
if self.entries.is_empty() {
return "graph TD".to_string();
}
let mut out = String::from("graph TD\n");
for (idx, entry) in self.entries.iter().enumerate() {
for (dep_name, _) in &entry.dependencies {
let dep_idx = self
.entries
.iter()
.position(|e| e.name == *dep_name)
.unwrap_or(idx);
let _ = writeln!(
out,
" n{dep_idx}[\"{dep_name}\"] --> n{idx}[\"{}\"]",
entry.name
);
}
}
for (idx, entry) in self.entries.iter().enumerate() {
if entry.dependencies.is_empty() {
let _ = writeln!(out, " n{idx}[\"{}\"]", entry.name);
}
}
out
}
}
impl Default for DependencyGraph {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum GraphError {
DependencyMissing {
module: &'static str,
missing: &'static str,
},
CycleDetected { cycle: Vec<&'static str> },
}
#[cfg(test)]
mod tests {
use super::*;
use std::any::TypeId;
mod types {
pub struct A;
pub struct B;
pub struct C;
}
fn typed_entry<T: 'static>(
name: &'static str,
deps: Vec<(&'static str, TypeId)>,
) -> ModuleEntry {
ModuleEntry {
type_id: TypeId::of::<T>(),
name,
dependencies: deps,
}
}
#[test]
fn graph_new_is_empty() {
let g = DependencyGraph::new();
assert!(g.entries().is_empty());
}
#[test]
fn graph_add_and_entries() {
let mut g = DependencyGraph::new();
g.add(typed_entry::<types::A>("a", vec![])).unwrap();
assert_eq!(g.entries().len(), 1);
}
#[test]
fn graph_add_duplicate_returns_err() {
let mut g = DependencyGraph::new();
g.add(typed_entry::<types::A>("a", vec![])).unwrap();
let err = g.add(typed_entry::<types::A>("a2", vec![])).unwrap_err();
assert_eq!(err, "a2");
}
#[test]
fn graph_validate_empty_succeeds() {
let g = DependencyGraph::new();
let sorted = g.validate().unwrap();
assert!(sorted.is_empty());
}
#[test]
fn graph_validate_single_node() {
let mut g = DependencyGraph::new();
g.add(typed_entry::<types::A>("a", vec![])).unwrap();
let sorted = g.validate().unwrap();
assert_eq!(sorted.len(), 1);
}
#[test]
fn graph_validate_missing_dependency() {
let mut g = DependencyGraph::new();
g.add(typed_entry::<types::A>(
"a",
vec![("b", TypeId::of::<types::B>())],
))
.unwrap();
let err = g.validate().unwrap_err();
assert!(matches!(
err,
GraphError::DependencyMissing {
module: "a",
missing: "b"
}
));
}
#[test]
fn graph_validate_cycle_two_nodes() {
let mut g = DependencyGraph::new();
g.add(typed_entry::<types::A>(
"a",
vec![("b", TypeId::of::<types::B>())],
))
.unwrap();
g.add(typed_entry::<types::B>(
"b",
vec![("a", TypeId::of::<types::A>())],
))
.unwrap();
let err = g.validate().unwrap_err();
assert!(matches!(err, GraphError::CycleDetected { .. }));
if let GraphError::CycleDetected { cycle } = err {
assert!(cycle.len() >= 2, "cycle should contain at least 2 names");
}
}
#[test]
fn graph_validate_cycle_three_nodes() {
let mut g = DependencyGraph::new();
g.add(typed_entry::<types::A>(
"a",
vec![("b", TypeId::of::<types::B>())],
))
.unwrap();
g.add(typed_entry::<types::B>(
"b",
vec![("c", TypeId::of::<types::C>())],
))
.unwrap();
g.add(typed_entry::<types::C>(
"c",
vec![("a", TypeId::of::<types::A>())],
))
.unwrap();
let err = g.validate().unwrap_err();
assert!(matches!(err, GraphError::CycleDetected { .. }));
}
#[test]
fn graph_validate_topo_order() {
let mut g = DependencyGraph::new();
g.add(typed_entry::<types::A>("a", vec![])).unwrap();
g.add(typed_entry::<types::B>(
"b",
vec![("a", TypeId::of::<types::A>())],
))
.unwrap();
let sorted = g.validate().unwrap();
let a_idx = sorted
.iter()
.position(|t| *t == TypeId::of::<types::A>())
.unwrap();
let b_idx = sorted
.iter()
.position(|t| *t == TypeId::of::<types::B>())
.unwrap();
assert!(a_idx < b_idx, "a should be sorted before b");
}
#[test]
fn graph_dependency_names() {
let mut g = DependencyGraph::new();
g.add(typed_entry::<types::A>("a", vec![])).unwrap();
g.add(typed_entry::<types::B>(
"b",
vec![("a", TypeId::of::<types::A>())],
))
.unwrap();
let names = g.dependency_names(TypeId::of::<types::B>());
assert_eq!(names, vec!["a"]);
}
#[test]
fn graph_dependency_names_unknown_returns_empty() {
let g = DependencyGraph::new();
let names = g.dependency_names(TypeId::of::<types::A>());
assert!(names.is_empty());
}
#[test]
fn graph_name_of() {
let mut g = DependencyGraph::new();
g.add(typed_entry::<types::A>("module-a", vec![])).unwrap();
assert_eq!(g.name_of(TypeId::of::<types::A>()), Some("module-a"));
assert_eq!(g.name_of(TypeId::of::<types::B>()), None);
}
#[test]
fn graph_default_is_empty() {
let g = DependencyGraph::default();
assert!(g.entries().is_empty());
}
#[test]
fn graph_to_dot_empty() {
let g = DependencyGraph::new();
assert_eq!(g.to_dot(), "digraph {}");
}
#[test]
fn graph_to_dot_with_nodes_and_edges() {
let mut g = DependencyGraph::new();
g.add(typed_entry::<types::A>("a", vec![])).unwrap();
g.add(typed_entry::<types::B>(
"b",
vec![("a", TypeId::of::<types::A>())],
))
.unwrap();
let dot = g.to_dot();
assert!(dot.starts_with("digraph {"));
assert!(dot.contains("\"a\""));
assert!(dot.contains("\"b\""));
assert!(dot.contains("\"a\" -> \"b\""));
assert!(dot.ends_with('}'));
}
#[test]
fn graph_to_mermaid_empty() {
let g = DependencyGraph::new();
assert_eq!(g.to_mermaid(), "graph TD");
}
#[test]
fn graph_to_mermaid_with_nodes_and_edges() {
let mut g = DependencyGraph::new();
g.add(typed_entry::<types::A>("a", vec![])).unwrap();
g.add(typed_entry::<types::B>(
"b",
vec![("a", TypeId::of::<types::A>())],
))
.unwrap();
let mermaid = g.to_mermaid();
assert!(mermaid.starts_with("graph TD"));
assert!(mermaid.contains("n0[\"a\"]"));
assert!(mermaid.contains("n1[\"b\"]"));
assert!(mermaid.contains("-->"));
}
#[test]
fn graph_to_mermaid_hyphen_names_no_collision() {
let mut g = DependencyGraph::new();
g.add(typed_entry::<types::A>("my-module", vec![])).unwrap();
g.add(typed_entry::<types::B>(
"my-dep",
vec![("my-module", TypeId::of::<types::A>())],
))
.unwrap();
let mermaid = g.to_mermaid();
assert!(mermaid.contains("n0[\"my-module\"]"));
assert!(mermaid.contains("n1[\"my-dep\"]"));
assert!(mermaid.contains("my-module"));
assert!(mermaid.contains("my-dep"));
}
#[test]
fn graph_error_debug() {
let err = GraphError::DependencyMissing {
module: "a",
missing: "b",
};
let debug = format!("{err:?}");
assert!(debug.contains("DependencyMissing"));
let err2 = GraphError::CycleDetected {
cycle: vec!["a", "b", "a"],
};
let debug2 = format!("{err2:?}");
assert!(debug2.contains("CycleDetected"));
}
}