use crate::edge::{Directedness, Edge};
use crate::error::GraphError;
#[derive(Debug, Clone, PartialEq)]
pub struct Graph<N, W> {
pub nodes: Vec<N>,
pub edges: Vec<Edge<W>>,
pub directedness: Directedness,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Neighbor<'a, W> {
pub edge_id: usize,
pub node: usize,
pub weight: &'a W,
}
impl<N, W> Graph<N, W> {
pub fn new(directedness: Directedness) -> Self {
Graph {
nodes: Vec::new(),
edges: Vec::new(),
directedness,
}
}
pub fn with_nodes(nodes: Vec<N>, directedness: Directedness) -> Self {
Graph {
nodes,
edges: Vec::new(),
directedness,
}
}
pub fn node_count(&self) -> usize {
self.nodes.len()
}
pub fn edge_count(&self) -> usize {
self.edges.len()
}
pub fn is_directed(&self) -> bool {
matches!(self.directedness, Directedness::Directed)
}
pub fn add_node(&mut self, label: N) -> usize {
self.nodes.push(label);
self.nodes.len() - 1
}
pub fn add_edge(
&mut self,
source: usize,
target: usize,
weight: W,
) -> Result<usize, GraphError> {
let n = self.nodes.len();
let id = self.edges.len();
for node in [source, target] {
if node >= n {
return Err(GraphError::InvalidEndpoint {
edge: id,
node,
len: n,
});
}
}
self.edges.push(Edge {
id,
source,
target,
weight,
});
Ok(id)
}
pub fn validate(&self) -> Result<(), GraphError> {
let n = self.nodes.len();
for e in &self.edges {
for node in [e.source, e.target] {
if node >= n {
return Err(GraphError::InvalidEndpoint {
edge: e.id,
node,
len: n,
});
}
}
}
Ok(())
}
pub fn neighbors(&self, node: usize) -> Result<Vec<Neighbor<'_, W>>, GraphError> {
if node >= self.nodes.len() {
return Err(GraphError::NodeOutOfRange {
node,
count: self.nodes.len(),
});
}
let directed = self.is_directed();
let mut out = Vec::new();
for e in &self.edges {
if e.source == node {
out.push(Neighbor {
edge_id: e.id,
node: e.target,
weight: &e.weight,
});
} else if !directed && e.target == node {
out.push(Neighbor {
edge_id: e.id,
node: e.source,
weight: &e.weight,
});
}
}
out.sort_by_key(|adj| (adj.node, adj.edge_id));
Ok(out)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn empty_graph_is_valid() {
let g: Graph<(), ()> = Graph::new(Directedness::Undirected);
assert_eq!(g.node_count(), 0);
assert_eq!(g.edge_count(), 0);
assert!(g.validate().is_ok());
}
#[test]
fn add_edge_rejects_invalid_endpoint() {
let mut g: Graph<&str, u64> = Graph::with_nodes(vec!["a", "b"], Directedness::Directed);
let r = g.add_edge(0, 5, 1);
assert!(matches!(
r,
Err(GraphError::InvalidEndpoint { node: 5, .. })
));
}
#[test]
fn self_loop_and_multiedge_preserved() {
let mut g: Graph<u8, u64> = Graph::with_nodes(vec![0, 1], Directedness::Undirected);
g.add_edge(0, 0, 1).unwrap(); g.add_edge(0, 1, 2).unwrap();
g.add_edge(0, 1, 3).unwrap(); assert_eq!(g.edge_count(), 3);
assert!(g.edges[0].is_self_loop());
}
#[test]
fn undirected_neighbors_expand_both_ways() {
let mut g: Graph<u8, u64> = Graph::with_nodes(vec![0, 1, 2], Directedness::Undirected);
g.add_edge(0, 1, 10).unwrap();
g.add_edge(2, 1, 20).unwrap();
let n1: Vec<usize> = g.neighbors(1).unwrap().iter().map(|a| a.node).collect();
assert_eq!(n1, vec![0, 2]); }
#[test]
fn directed_neighbors_are_outgoing_only() {
let mut g: Graph<u8, u64> = Graph::with_nodes(vec![0, 1], Directedness::Directed);
g.add_edge(0, 1, 1).unwrap();
assert_eq!(g.neighbors(0).unwrap().len(), 1);
assert_eq!(g.neighbors(1).unwrap().len(), 0);
}
}