use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct Contact {
pub start: f64,
pub end: f64,
pub a: usize,
pub b: usize,
}
#[derive(Debug, Clone)]
pub struct ContactTrace {
pub n: usize,
pub duration: f64,
pub contacts: Vec<Contact>,
}
impl ContactTrace {
pub fn load_csv(text: &str) -> Result<Self, String> {
let mut contacts = Vec::new();
let mut max_node = 0usize;
let mut duration = 0.0f64;
for (i, line) in text.lines().enumerate() {
let line = line.trim();
if line.is_empty() {
continue;
}
let cols: Vec<&str> = line.split(',').map(str::trim).collect();
if cols.len() < 4 {
return Err(format!("line {}: expected 4 columns, got {}", i + 1, cols.len()));
}
if i == 0 && cols[0].parse::<f64>().is_err() {
continue;
}
let start: f64 = cols[0].parse().map_err(|_| format!("line {}: bad start", i + 1))?;
let end: f64 = cols[1].parse().map_err(|_| format!("line {}: bad end", i + 1))?;
let a: usize = cols[2].parse().map_err(|_| format!("line {}: bad node a", i + 1))?;
let b: usize = cols[3].parse().map_err(|_| format!("line {}: bad node b", i + 1))?;
if a == b {
continue; }
let (a, b) = if a < b { (a, b) } else { (b, a) };
max_node = max_node.max(b);
duration = duration.max(end);
contacts.push(Contact { start, end, a, b });
}
if contacts.is_empty() {
return Err("no contacts parsed".into());
}
contacts.sort_by(|x, y| x.start.partial_cmp(&y.start).unwrap());
Ok(Self { n: max_node + 1, duration, contacts })
}
pub fn synthetic(
n: usize,
duration: f64,
alpha: f64,
mean_contact: f64,
pair_active_prob: f64,
seed: u64,
) -> Self {
let mut rng = StdRng::seed_from_u64(seed);
let mut contacts = Vec::new();
let gap_min = mean_contact.max(1.0);
let gap_max = duration;
for a in 0..n {
for b in (a + 1)..n {
if rng.gen::<f64>() >= pair_active_prob {
continue; }
let mut t = rng.gen::<f64>() * gap_min; loop {
let gap = sample_truncated_power_law(&mut rng, alpha, gap_min, gap_max);
t += gap;
if t >= duration {
break;
}
let len = -mean_contact * (1.0 - rng.gen::<f64>()).ln();
let end = (t + len).min(duration);
contacts.push(Contact { start: t, end, a, b });
t = end;
}
}
}
contacts.sort_by(|x, y| x.start.partial_cmp(&y.start).unwrap());
Self { n, duration, contacts }
}
pub fn earliest_arrival(&self, src: usize, t0: f64) -> Vec<f64> {
let mut arrival = vec![f64::INFINITY; self.n];
arrival[src] = t0;
for c in &self.contacts {
if c.end < t0 {
continue;
}
let (ra, rb) = (arrival[c.a], arrival[c.b]);
if ra <= c.end {
let arrive = c.start.max(ra);
if arrive <= c.end && arrive < arrival[c.b] {
arrival[c.b] = arrive;
}
}
if rb <= c.end {
let arrive = c.start.max(rb);
if arrive <= c.end && arrive < arrival[c.a] {
arrival[c.a] = arrive;
}
}
}
arrival
}
pub fn mean_inter_contact(&self) -> f64 {
use std::collections::HashMap;
let mut last: HashMap<(usize, usize), f64> = HashMap::new();
let mut gaps = Vec::new();
for c in &self.contacts {
if let Some(&prev_end) = last.get(&(c.a, c.b)) {
gaps.push(c.start - prev_end);
}
last.insert((c.a, c.b), c.end);
}
if gaps.is_empty() {
0.0
} else {
gaps.iter().sum::<f64>() / gaps.len() as f64
}
}
}
fn sample_truncated_power_law(rng: &mut StdRng, alpha: f64, lo: f64, hi: f64) -> f64 {
let u: f64 = rng.gen();
let lo_a = lo.powf(-alpha);
let hi_a = hi.powf(-alpha);
(lo_a - u * (lo_a - hi_a)).powf(-1.0 / alpha)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn load_csv_parses_and_orders() {
let csv = "start,end,a,b\n30,40,1,0\n0,10,0,2\n";
let t = ContactTrace::load_csv(csv).unwrap();
assert_eq!(t.n, 3);
assert_eq!(t.duration, 40.0);
assert_eq!(t.contacts[0].start, 0.0);
assert_eq!((t.contacts[1].a, t.contacts[1].b), (0, 1));
}
#[test]
fn load_csv_rejects_empty() {
assert!(ContactTrace::load_csv("\n\n").is_err());
}
#[test]
fn synthetic_is_reproducible() {
let a = ContactTrace::synthetic(20, 10_000.0, 0.5, 30.0, 0.4, 7);
let b = ContactTrace::synthetic(20, 10_000.0, 0.5, 30.0, 0.4, 7);
assert_eq!(a.contacts.len(), b.contacts.len());
assert_eq!(a.contacts.first(), b.contacts.first());
}
#[test]
fn synthetic_has_heavy_tailed_gaps() {
let t = ContactTrace::synthetic(30, 200_000.0, 0.5, 30.0, 0.5, 1);
let mean_gap = t.mean_inter_contact();
assert!(
mean_gap > 30.0 * 5.0,
"expected heavy tail to lift mean gap well above the floor, got {mean_gap}"
);
}
#[test]
fn earliest_arrival_follows_temporal_path() {
let csv = "0,10,0,1\n20,30,1,2\n5,8,1,2\n";
let tr = ContactTrace::load_csv(csv).unwrap();
let arr = tr.earliest_arrival(0, 0.0);
assert_eq!(arr[0], 0.0);
assert_eq!(arr[1], 0.0);
assert_eq!(arr[2], 5.0, "should take the earliest usable 1—2 contact");
}
#[test]
fn earliest_arrival_respects_time_order() {
let csv = "5,8,1,2\n20,30,0,1\n";
let tr = ContactTrace::load_csv(csv).unwrap();
let arr = tr.earliest_arrival(0, 0.0);
assert_eq!(arr[1], 20.0);
assert!(arr[2].is_infinite(), "no forward-time path 0→2 exists");
}
#[test]
fn synthetic_thins_pairs() {
let t = ContactTrace::synthetic(40, 50_000.0, 0.5, 30.0, 0.1, 3);
let pairs: std::collections::HashSet<(usize, usize)> =
t.contacts.iter().map(|c| (c.a, c.b)).collect();
let all_pairs = 40 * 39 / 2;
assert!(pairs.len() < all_pairs / 2, "expected sparse pair set, got {}", pairs.len());
}
}