use crate::INF_WEIGHT;
use crate::bundle::{CchView, INVALID_ID, MetricView};
pub struct ElimTreeQuery<'a> {
cch: &'a CchView<'a>,
forward_tentative_distance: Vec<u32>,
target_node: Vec<u32>,
target_elimination_tree_end: Vec<u32>,
in_backward_search_space: Vec<bool>,
stack: Vec<u32>,
has_active_source: bool,
active_source_node: u32,
active_source_end: u32,
}
impl<'a> ElimTreeQuery<'a> {
#[must_use]
pub fn new(cch: &'a CchView<'a>) -> Self {
let n = cch.node_count() as usize;
Self {
cch,
forward_tentative_distance: vec![INF_WEIGHT; n],
target_node: Vec::new(),
target_elimination_tree_end: Vec::new(),
in_backward_search_space: vec![false; n],
stack: Vec::new(),
has_active_source: false,
active_source_node: INVALID_ID,
active_source_end: INVALID_ID,
}
}
pub fn pin_targets(&mut self, targets: &[u32]) {
assert!(
self.target_node.is_empty(),
"pin_targets called twice on the same query: in_backward_search_space \
marks are never reset, so re-pinning would corrupt results; \
create a fresh ElimTreeQuery instead"
);
self.target_node.reserve(targets.len());
self.target_elimination_tree_end.reserve(targets.len());
for &t_ext in targets {
let t = self.cch.rank[t_ext as usize];
self.target_node.push(t);
let mut end = INVALID_ID;
let mut x = t;
while x != INVALID_ID {
if self.in_backward_search_space[x as usize] {
end = x;
break;
}
self.in_backward_search_space[x as usize] = true;
x = self.cch.elimination_tree_parent[x as usize];
}
self.target_elimination_tree_end.push(end);
}
}
pub fn reset_source(&mut self) {
if self.has_active_source {
let mut x = self.active_source_node;
while x != self.active_source_end {
self.forward_tentative_distance[x as usize] = INF_WEIGHT;
x = self.cch.elimination_tree_parent[x as usize];
}
self.has_active_source = false;
self.active_source_node = INVALID_ID;
self.active_source_end = INVALID_ID;
}
for i in (0..self.target_node.len()).rev() {
let t = self.target_node[i];
let end = self.target_elimination_tree_end[i];
let mut x = t;
while x != end {
self.forward_tentative_distance[x as usize] = INF_WEIGHT;
x = self.cch.elimination_tree_parent[x as usize];
}
}
}
pub fn add_source_and_run(&mut self, metric: &MetricView, source: u32) {
let s = self.cch.rank[source as usize];
self.forward_tentative_distance[s as usize] = 0;
self.active_source_node = s;
self.has_active_source = true;
let mut x = s;
loop {
let from = self.cch.up_first_out[x as usize] as usize;
let to = self.cch.up_first_out[x as usize + 1] as usize;
let dx = self.forward_tentative_distance[x as usize];
if dx != INF_WEIGHT {
for xy in from..to {
let y = self.cch.up_head[xy] as usize;
let candidate = dx.saturating_add(metric.forward[xy]);
if candidate < self.forward_tentative_distance[y] {
self.forward_tentative_distance[y] = candidate;
}
}
}
let parent = self.cch.elimination_tree_parent[x as usize];
if parent == INVALID_ID {
break;
}
x = parent;
}
self.active_source_end = INVALID_ID;
self.stack.clear();
for i in (0..self.target_node.len()).rev() {
let t = self.target_node[i];
let end = self.target_elimination_tree_end[i];
let mut x = t;
while x != end {
self.stack.push(x);
x = self.cch.elimination_tree_parent[x as usize];
}
}
while let Some(x) = self.stack.pop() {
let from = self.cch.up_first_out[x as usize] as usize;
let to = self.cch.up_first_out[x as usize + 1] as usize;
let mut best = self.forward_tentative_distance[x as usize];
for xy in from..to {
let y = self.cch.up_head[xy] as usize;
let dy = self.forward_tentative_distance[y];
if dy != INF_WEIGHT {
let candidate = dy.saturating_add(metric.backward[xy]);
if candidate < best {
best = candidate;
}
}
}
self.forward_tentative_distance[x as usize] = best;
}
}
pub fn get_distances_to_targets(&self, out: &mut [u32]) {
assert_eq!(
out.len(),
self.target_node.len(),
"out buffer length must equal pinned-target count"
);
for (i, &t) in self.target_node.iter().enumerate() {
out[i] = self.forward_tentative_distance[t as usize];
}
}
}
#[must_use]
pub fn distance_matrix(
cch: &CchView,
metric: &MetricView,
sources: &[u32],
targets: &[u32],
) -> Vec<u32> {
if sources.is_empty() || targets.is_empty() {
return Vec::new();
}
let mut q = ElimTreeQuery::new(cch);
q.pin_targets(targets);
let mut out = vec![0u32; sources.len() * targets.len()];
let mut row = vec![0u32; targets.len()];
for (i, &s) in sources.iter().enumerate() {
q.reset_source();
q.add_source_and_run(metric, s);
q.get_distances_to_targets(&mut row);
out[i * targets.len()..i * targets.len() + targets.len()].copy_from_slice(&row);
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn distance_matrix_matches_cpp_oracle() {
use routingkit_cch::ffi;
let n: u32 = 6;
let tail: Vec<u32> = vec![0, 1, 2, 3, 4, 3, 2, 1];
let head: Vec<u32> = vec![1, 2, 3, 4, 3, 2, 1, 0];
let weights: Vec<u32> = (1..=8u32).collect();
let order: Vec<u32> = (0..n).collect();
let cch = unsafe { ffi::cch_new(&order, &tail, &head, |_| {}, false) };
let cch_ref = cch.as_ref().expect("cch_new returned null");
let dir = tempfile::tempdir().expect("tempdir");
let struct_path = dir.path().join("dm.cch-struct");
let metric_path = dir.path().join("dm.cch-metric");
unsafe {
ffi::cch_save_struct(cch_ref, struct_path.to_str().unwrap()).expect("cch_save_struct");
}
let oracle_matrix = unsafe {
let mut metric = ffi::cch_metric_new(cch_ref, &weights);
ffi::cch_metric_customize(metric.as_mut().expect("metric pin"));
let sources: Vec<u32> = (0..n).collect();
let targets: Vec<u32> = (0..n).collect();
let matrix = ffi::cch_compute_distance_matrix(
metric.as_ref().expect("metric ref"),
&sources,
&targets,
);
ffi::cch_save_metric(
metric.as_ref().expect("metric ref"),
metric_path.to_str().unwrap(),
)
.expect("cch_save_metric");
matrix
};
let cch_bundle = crate::bundle::CchBundle::open(&struct_path).expect("CchBundle::open");
let metric_bundle =
crate::bundle::MetricBundle::open(&metric_path).expect("MetricBundle::open");
let cch_view = cch_bundle.view();
let metric_view = metric_bundle.view();
let sources: Vec<u32> = (0..n).collect();
let targets: Vec<u32> = (0..n).collect();
let rust_matrix = distance_matrix(&cch_view, &metric_view, &sources, &targets);
assert_eq!(
oracle_matrix.len(),
rust_matrix.len(),
"matrix length mismatch"
);
for k in 0..oracle_matrix.len() {
let oracle_val = oracle_matrix[k];
let rust_val = rust_matrix[k];
let oracle_unreachable = oracle_val == u32::MAX || oracle_val == INF_WEIGHT;
let rust_unreachable = rust_val == INF_WEIGHT;
let row = k / (n as usize);
let col = k % (n as usize);
assert_eq!(
oracle_unreachable, rust_unreachable,
"reachability mismatch at index {k} (i={row}, j={col}): oracle={oracle_val}, rust={rust_val}",
);
if !oracle_unreachable && !rust_unreachable {
assert_eq!(
oracle_val, rust_val,
"distance mismatch at index {k} (i={row}, j={col}): oracle={oracle_val}, rust={rust_val}",
);
}
}
}
#[test]
fn distance_matrix_empty_sources_returns_empty() {
use crate::bundle::{CchBundle, MetricBundle};
use routingkit_cch::ffi;
let n: u32 = 5;
let order: Vec<u32> = (0..n).collect();
let tail: Vec<u32> = (0..n - 1).collect();
let head: Vec<u32> = (1..n).collect();
let cch = unsafe { ffi::cch_new(&order, &tail, &head, |_| {}, false) };
let cch_ref = cch.as_ref().expect("cch_new returned null");
let dir = tempfile::tempdir().expect("tempdir");
let struct_path = dir.path().join("q.cch-struct");
let metric_path = dir.path().join("q.cch-metric");
#[allow(clippy::cast_possible_truncation)]
let weights: Vec<u32> = (0..tail.len() as u32).collect();
let mut metric = unsafe { ffi::cch_metric_new(cch_ref, &weights) };
unsafe {
ffi::cch_save_struct(cch_ref, struct_path.to_str().unwrap()).unwrap();
ffi::cch_metric_customize(metric.as_mut().unwrap());
ffi::cch_save_metric(metric.as_ref().unwrap(), metric_path.to_str().unwrap()).unwrap();
}
let cch_bundle = CchBundle::open(&struct_path).unwrap();
let metric_bundle = MetricBundle::open(&metric_path).unwrap();
let cv = cch_bundle.view();
let mv = metric_bundle.view();
let result = distance_matrix(&cv, &mv, &[], &[0, 1]);
assert!(result.is_empty());
let result = distance_matrix(&cv, &mv, &[0, 1], &[]);
assert!(result.is_empty());
}
#[test]
#[should_panic(expected = "pin_targets called twice")]
fn pin_targets_twice_panics() {
use crate::bundle::CchBundle;
use routingkit_cch::ffi;
let n: u32 = 5;
let order: Vec<u32> = (0..n).collect();
let tail: Vec<u32> = (0..n - 1).collect();
let head: Vec<u32> = (1..n).collect();
let cch = unsafe { ffi::cch_new(&order, &tail, &head, |_| {}, false) };
let cch_ref = cch.as_ref().expect("cch_new returned null");
let dir = tempfile::tempdir().expect("tempdir");
let struct_path = dir.path().join("p.cch-struct");
unsafe {
ffi::cch_save_struct(cch_ref, struct_path.to_str().unwrap()).unwrap();
}
let cch_bundle = CchBundle::open(&struct_path).unwrap();
let cv = cch_bundle.view();
let mut q = ElimTreeQuery::new(&cv);
q.pin_targets(&[0]);
q.pin_targets(&[1]); }
#[test]
#[should_panic(expected = "out buffer length must equal")]
fn get_distances_wrong_len_panics() {
use crate::bundle::{CchBundle, MetricBundle};
use routingkit_cch::ffi;
let n: u32 = 5;
let order: Vec<u32> = (0..n).collect();
let tail: Vec<u32> = (0..n - 1).collect();
let head: Vec<u32> = (1..n).collect();
let cch = unsafe { ffi::cch_new(&order, &tail, &head, |_| {}, false) };
let cch_ref = cch.as_ref().expect("cch_new returned null");
let dir = tempfile::tempdir().expect("tempdir");
let struct_path = dir.path().join("g.cch-struct");
let metric_path = dir.path().join("g.cch-metric");
#[allow(clippy::cast_possible_truncation)]
let weights: Vec<u32> = (0..tail.len() as u32).collect();
let mut metric = unsafe { ffi::cch_metric_new(cch_ref, &weights) };
unsafe {
ffi::cch_save_struct(cch_ref, struct_path.to_str().unwrap()).unwrap();
ffi::cch_metric_customize(metric.as_mut().unwrap());
ffi::cch_save_metric(metric.as_ref().unwrap(), metric_path.to_str().unwrap()).unwrap();
}
let cch_bundle = CchBundle::open(&struct_path).unwrap();
let metric_bundle = MetricBundle::open(&metric_path).unwrap();
let cv = cch_bundle.view();
let mv = metric_bundle.view();
let mut q = ElimTreeQuery::new(&cv);
q.pin_targets(&[0, 1]);
q.add_source_and_run(&mv, 0);
let mut out = vec![0u32; 3]; q.get_distances_to_targets(&mut out); }
#[test]
fn add_source_run_inf_weight_branch() {
use crate::bundle::{CchBundle, MetricBundle};
use routingkit_cch::ffi;
let n: u32 = 5;
let order: Vec<u32> = (0..n).collect();
let tail: Vec<u32> = (0..n - 1).collect();
let head: Vec<u32> = (1..n).collect();
let weights: Vec<u32> = vec![crate::INF_WEIGHT; tail.len()];
let cch = unsafe { ffi::cch_new(&order, &tail, &head, |_| {}, false) };
let cch_ref = cch.as_ref().expect("cch_new returned null");
let dir = tempfile::tempdir().expect("tempdir");
let struct_path = dir.path().join("infwt.cch-struct");
let metric_path = dir.path().join("infwt.cch-metric");
let mut metric = unsafe { ffi::cch_metric_new(cch_ref, &weights) };
unsafe {
ffi::cch_save_struct(cch_ref, struct_path.to_str().unwrap()).unwrap();
ffi::cch_metric_customize(metric.as_mut().unwrap());
ffi::cch_save_metric(metric.as_ref().unwrap(), metric_path.to_str().unwrap()).unwrap();
}
let cch_bundle = CchBundle::open(&struct_path).unwrap();
let metric_bundle = MetricBundle::open(&metric_path).unwrap();
let cv = cch_bundle.view();
let mv = metric_bundle.view();
let mut q = ElimTreeQuery::new(&cv);
q.pin_targets(&[4]);
q.add_source_and_run(&mv, 0);
let mut out = vec![0u32; 1];
q.get_distances_to_targets(&mut out);
assert_eq!(out[0], crate::INF_WEIGHT);
}
}