use crate::Snapshot;
pub const UNREACHABLE: u64 = u64::MAX;
pub const DELTA: u32 = 4;
const BUCKETS: u64 = 1 << 16;
#[must_use]
pub fn sssp(g: &Snapshot, weights: &[u32], src: u32) -> Vec<u64> {
sssp_with(g, weights, src, DELTA)
}
#[must_use]
pub fn sssp_with(g: &Snapshot, weights: &[u32], src: u32, delta: u32) -> Vec<u64> {
assert_eq!(weights.len() as u64, g.edges(), "one weight an edge");
let n = g.nodes();
let mut far = vec![UNREACHABLE; n as usize];
if src >= n {
return far;
}
let mut shift = delta.max(1).ilog2();
let heaviest = u64::from(weights.iter().copied().max().unwrap_or(0));
while (heaviest >> shift) + 2 > BUCKETS {
shift += 1;
}
let width = 1u64 << shift;
let ring = ((heaviest >> shift) + 2).next_power_of_two();
let mask = ring - 1;
let mut bucket: Vec<Vec<u32>> = vec![Vec::new(); ring as usize];
far[src as usize] = 0;
bucket[0].push(src);
let mut waiting = 1usize;
let mut band = 0u64;
let mut here: Vec<u32> = Vec::new();
let mut done: Vec<u32> = Vec::new();
while waiting > 0 {
let at = (band & mask) as usize;
done.clear();
while !bucket[at].is_empty() {
std::mem::swap(&mut here, &mut bucket[at]);
waiting -= here.len();
for node in here.drain(..) {
let from = far[node as usize];
if from >> shift != band {
continue;
}
done.push(node);
let near = g.out(node);
let cost = &weights[g.out_at(node)..][..near.len()];
for (to, weight) in near.iter().zip(cost) {
let weight = u64::from(*weight);
if weight > width {
continue;
}
let now = from + weight;
if now < far[*to as usize] {
far[*to as usize] = now;
bucket[((now >> shift) & mask) as usize].push(*to);
waiting += 1;
}
}
}
}
for node in &done {
let from = far[*node as usize];
let near = g.out(*node);
let cost = &weights[g.out_at(*node)..][..near.len()];
for (to, weight) in near.iter().zip(cost) {
let weight = u64::from(*weight);
if weight <= width {
continue;
}
let now = from + weight;
if now < far[*to as usize] {
far[*to as usize] = now;
bucket[((now >> shift) & mask) as usize].push(*to);
waiting += 1;
}
}
}
band += 1;
}
far
}
#[cfg(test)]
mod tests {
use super::*;
use crate::graph::NO_PROPS;
use crate::{Graph, Snapshot};
use std::collections::BinaryHeap;
use yo_common::Rng;
use yo_doc::Builder;
fn reference(g: &Snapshot, weights: &[u32], src: u32) -> Vec<u64> {
let mut far = vec![UNREACHABLE; g.nodes() as usize];
let mut todo = BinaryHeap::new();
far[src as usize] = 0;
todo.push((std::cmp::Reverse(0u64), src));
while let Some((std::cmp::Reverse(at), node)) = todo.pop() {
if at > far[node as usize] {
continue;
}
let from = g.out_at(node);
for (i, to) in g.out(node).iter().enumerate() {
let now = at + u64::from(weights[from + i]);
if now < far[*to as usize] {
far[*to as usize] = now;
todo.push((std::cmp::Reverse(now), *to));
}
}
}
far
}
fn costing(cost: i64) -> Vec<u8> {
let mut b = Builder::new();
b.begin_object().expect("an object");
b.key(b"cost").expect("a key");
b.int(cost).expect("a number");
b.end_object().expect("an end");
b.finish().expect("a document").to_vec()
}
fn weighted(edges: &[(u64, u64, i64)]) -> (Snapshot, Vec<u32>) {
let mut g = Graph::new();
for (from, to, cost) in edges {
g.link(*from, *to, 1, &costing(*cost)).expect("an edge");
}
Snapshot::weighted(&g, &[1], b"cost", 1)
}
#[test]
fn a_chain_adds_up() {
let (s, w) = weighted(&[(1, 2, 5), (2, 3, 7), (3, 4, 2)]);
let far = sssp(&s, &w, s.dense(1).expect("1"));
assert_eq!(far[s.dense(4).expect("4") as usize], 14);
}
#[test]
fn the_cheap_way_round_wins() {
let (s, w) = weighted(&[(1, 4, 100), (1, 2, 1), (2, 3, 1), (3, 4, 1)]);
let far = sssp(&s, &w, s.dense(1).expect("1"));
assert_eq!(far[s.dense(4).expect("4") as usize], 3);
}
#[test]
fn what_cannot_be_reached_stays_unreachable() {
let (s, w) = weighted(&[(1, 2, 1), (3, 4, 1)]);
let far = sssp(&s, &w, s.dense(1).expect("1"));
assert_eq!(far[s.dense(3).expect("3") as usize], UNREACHABLE);
assert_eq!(far[s.dense(1).expect("1") as usize], 0);
}
#[test]
fn direction_is_respected() {
let (s, w) = weighted(&[(1, 2, 1)]);
let far = sssp(&s, &w, s.dense(2).expect("2"));
assert_eq!(far[s.dense(1).expect("1") as usize], UNREACHABLE);
}
#[test]
fn a_huge_weight_does_not_ask_for_a_huge_ring() {
let heaviest = i64::from(u32::MAX);
let (s, w) = weighted(&[
(1, 4, heaviest),
(1, 2, heaviest / 3),
(2, 3, heaviest / 3),
(3, 4, heaviest / 3),
]);
let far = sssp_with(&s, &w, s.dense(1).expect("1"), 1);
assert_eq!(
far[s.dense(4).expect("4") as usize],
heaviest as u64 / 3 * 3
);
assert_eq!(far, reference(&s, &w, s.dense(1).expect("1")));
}
#[test]
fn a_source_that_is_not_a_node() {
let (s, w) = weighted(&[(1, 2, 1)]);
assert!(sssp(&s, &w, 99).iter().all(|far| *far == UNREACHABLE));
}
#[test]
fn an_edge_that_weighs_nothing_is_free() {
let (s, w) = weighted(&[(1, 2, 0), (2, 3, 0)]);
let far = sssp(&s, &w, s.dense(1).expect("1"));
assert_eq!(far[s.dense(3).expect("3") as usize], 0);
}
#[test]
fn a_very_heavy_edge() {
let (s, w) = weighted(&[(1, 2, 1), (2, 3, 1_000_000), (1, 3, 999_999)]);
let far = sssp(&s, &w, s.dense(1).expect("1"));
assert_eq!(far[s.dense(3).expect("3") as usize], 999_999);
}
#[test]
fn all_the_same_weight_is_the_hop_count() {
let edges: Vec<(u64, u64, i64)> = (0..50u64).map(|i| (i, i + 1, 1)).collect();
let (s, w) = weighted(&edges);
let far = sssp(&s, &w, s.dense(0).expect("0"));
for id in 0..=50u64 {
assert_eq!(far[s.dense(id).expect("a node") as usize], id);
}
}
#[test]
fn the_band_width_does_not_change_the_answer() {
let mut rng = Rng::new(0x5551);
let edges: Vec<(u64, u64, i64)> = (0..600)
.map(|_| {
(
rng.next_u64() % 100,
rng.next_u64() % 100,
(rng.next_u64() % 50) as i64,
)
})
.collect();
let (s, w) = weighted(&edges);
let src = s.dense(0).expect("0");
let want = reference(&s, &w, src);
for delta in [0u32, 1, 3, 16, 1000] {
assert_eq!(sssp_with(&s, &w, src, delta), want, "delta {delta}");
}
}
#[test]
fn it_agrees_with_dijkstra() {
let mut rng = Rng::new(0x5550);
for case in 0..40 {
let nodes = 2 + rng.next_u64() % 80;
let edges: Vec<(u64, u64, i64)> = (0..nodes * 3)
.map(|_| {
(
rng.next_u64() % nodes,
rng.next_u64() % nodes,
(rng.next_u64() % 200) as i64,
)
})
.collect();
let (s, w) = weighted(&edges);
let src = rng.next_u64() as u32 % s.nodes();
assert_eq!(sssp(&s, &w, src), reference(&s, &w, src), "case {case}");
}
}
#[test]
fn a_weight_comes_off_the_edge_it_belongs_to() {
let mut g = Graph::new();
g.link(1, 2, 1, &costing(7)).expect("an edge");
g.link(1, 3, 1, &costing(9)).expect("an edge");
g.link(2, 3, 1, NO_PROPS).expect("an edge");
let (s, w) = Snapshot::weighted(&g, &[1], b"cost", 4);
let one = s.dense(1).expect("1");
let seen: Vec<(u64, u32)> = s
.out(one)
.iter()
.enumerate()
.map(|(i, to)| (s.id(*to), w[s.out_at(one) + i]))
.collect();
assert!(seen.contains(&(2, 7)), "{seen:?}");
assert!(seen.contains(&(3, 9)), "{seen:?}");
let two = s.dense(2).expect("2");
assert_eq!(w[s.out_at(two)], 4);
}
#[test]
fn a_weight_that_is_not_a_weight() {
let mut b = Builder::new();
b.begin_object().expect("an object");
b.key(b"cost").expect("a key");
b.text("free").expect("some text");
b.end_object().expect("an end");
let text = b.finish().expect("a document").to_vec();
let mut g = Graph::new();
g.link(1, 2, 1, &text).expect("an edge");
g.link(2, 3, 1, &costing(-5)).expect("an edge");
let (_, w) = Snapshot::weighted(&g, &[1], b"cost", 6);
assert_eq!(w, vec![6, 6]);
}
}