use rayon::prelude::*;
use super::csr::{Csr, CsrBuilder, InputError};
use crate::ids::SectionId;
#[derive(Clone, Debug)]
pub struct SectionGraph {
edges: Csr<SectionId>,
roots: Vec<SectionId>,
}
impl SectionGraph {
pub fn new(edges: Csr<SectionId>, roots: Vec<SectionId>) -> Result<Self, InputError> {
let len = edges.rows();
let out_of_range = |id: &SectionId| id.index() >= len;
if let Some(bad) = edges.values().par_iter().find_first(|id| out_of_range(id)) {
return Err(InputError::OutOfRange {
what: "edge target section",
index: u64::from(bad.as_u32()),
len,
});
}
if let Some(bad) = roots.iter().find(|id| out_of_range(id)) {
return Err(InputError::OutOfRange {
what: "root section",
index: u64::from(bad.as_u32()),
len,
});
}
Ok(Self { edges, roots })
}
pub fn build_parallel<C, F>(
sections: usize,
count: C,
fill: F,
roots: Vec<SectionId>,
) -> Result<Self, InputError>
where
C: Fn(SectionId) -> usize + Sync,
F: Fn(SectionId, &mut [SectionId]) -> usize + Sync,
{
if u32::try_from(sections).is_err() {
return Err(InputError::TooLarge("section count"));
}
let edges = Csr::build_parallel(
sections,
SectionId::from_u32(0),
|row| count(SectionId::new(row)),
|row, slot| fill(SectionId::new(row), slot),
)?;
Self::new(edges, roots)
}
#[inline]
#[must_use]
pub fn num_sections(&self) -> usize {
self.edges.rows()
}
#[inline]
#[must_use]
pub fn num_edges(&self) -> usize {
self.edges.num_values()
}
#[inline]
#[must_use]
pub fn edges(&self, section: SectionId) -> &[SectionId] {
self.edges.row(section.index())
}
#[inline]
#[must_use]
pub fn roots(&self) -> &[SectionId] {
&self.roots
}
#[inline]
#[must_use]
pub fn edge_table(&self) -> &Csr<SectionId> {
&self.edges
}
}
#[derive(Clone, Debug)]
pub struct GraphBuilder {
edges: CsrBuilder<SectionId>,
roots: Vec<SectionId>,
}
impl GraphBuilder {
#[must_use]
pub fn new(sections: usize) -> Self {
Self {
edges: CsrBuilder::new(sections),
roots: Vec::new(),
}
}
pub fn add_edge(&mut self, from: SectionId, to: SectionId) {
self.edges.push(from.index(), to);
}
pub fn add_root(&mut self, section: SectionId) {
self.roots.push(section);
}
pub fn keep_together(&mut self, members: &[SectionId]) {
if members.len() < 2 {
return;
}
for pair in members.windows(2) {
self.add_edge(pair[0], pair[1]);
}
if let (Some(&last), Some(&first)) = (members.last(), members.first()) {
self.add_edge(last, first);
}
}
pub fn build(self) -> Result<SectionGraph, InputError> {
if u32::try_from(self.edges.rows()).is_err() {
return Err(InputError::TooLarge("section count"));
}
SectionGraph::new(self.edges.build()?, self.roots)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn id(index: usize) -> SectionId {
SectionId::new(index)
}
#[test]
fn builder_and_parallel_agree() {
let mut builder = GraphBuilder::new(4);
builder.add_edge(id(0), id(1));
builder.add_edge(id(0), id(2));
builder.keep_together(&[id(2), id(3)]);
builder.add_root(id(0));
let built = builder.build().unwrap();
let parallel = SectionGraph::build_parallel(
4,
|_| 3,
|section, slot| match section.index() {
0 => {
slot[0] = id(1);
slot[1] = id(2);
2
}
2 => {
slot[0] = id(3);
1
}
3 => {
slot[0] = id(2);
1
}
_ => 0,
},
vec![id(0)],
)
.unwrap();
for section in 0..4 {
assert_eq!(built.edges(id(section)), parallel.edges(id(section)));
}
assert_eq!(built.num_edges(), 4);
assert_eq!(built.roots(), parallel.roots());
}
#[test]
fn rejects_out_of_range() {
let mut builder = GraphBuilder::new(2);
builder.add_edge(id(0), id(2));
assert!(matches!(
builder.build(),
Err(InputError::OutOfRange { index: 2, .. })
));
let mut builder = GraphBuilder::new(2);
builder.add_root(id(5));
assert!(builder.build().is_err());
let mut builder = GraphBuilder::new(2);
builder.add_edge(id(7), id(0));
assert!(builder.build().is_err());
}
}