use std::collections::HashMap;
use crate::analyses::coupling::{CouplingRow, run_coupling};
use crate::facts::FactsDb;
use crate::{Options, Result};
const PAGERANK_DAMPING: f64 = 0.85;
const POWER_ITER_MAX: usize = 30;
const POWER_ITER_EPSILON: f64 = 1e-6;
#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)]
pub struct CentralityRow {
pub path: String,
pub degree: u32,
pub weighted_degree: f64,
pub pagerank: f64,
pub eigenvector: f64,
}
#[tracing::instrument(name = "centrality", skip_all, fields(min_revs = opts.min_revs))]
pub fn run_centrality(db: &FactsDb, opts: &Options) -> Result<Vec<CentralityRow>> {
let pairs = run_coupling(db, opts)?;
Ok(compute_centrality(&pairs))
}
#[must_use]
pub fn compute_centrality(pairs: &[CouplingRow]) -> Vec<CentralityRow> {
if pairs.is_empty() {
return Vec::new();
}
let mut path_to_id: HashMap<String, usize> = HashMap::new();
let mut id_to_path: Vec<String> = Vec::new();
for pair in pairs {
for path in [&pair.entity_a, &pair.entity_b] {
if !path_to_id.contains_key(path) {
path_to_id.insert(path.clone(), id_to_path.len());
id_to_path.push(path.clone());
}
}
}
let n = id_to_path.len();
let mut adj: Vec<Vec<(usize, f64)>> = vec![Vec::new(); n];
for pair in pairs {
let a = path_to_id[&pair.entity_a];
let b = path_to_id[&pair.entity_b];
let w = pair.degree; adj[a].push((b, w));
adj[b].push((a, w));
}
let degree: Vec<u32> = adj
.iter()
.map(|edges| u32::try_from(edges.len()).unwrap_or(u32::MAX))
.collect();
let weighted_degree: Vec<f64> = adj
.iter()
.map(|edges| edges.iter().map(|(_, w)| *w).sum())
.collect();
let pagerank = pagerank_weighted(&adj, &weighted_degree);
let eigenvector = eigenvector_weighted(&adj);
let mut out: Vec<CentralityRow> = (0..n)
.map(|i| CentralityRow {
path: id_to_path[i].clone(),
degree: degree[i],
weighted_degree: weighted_degree[i],
pagerank: pagerank[i],
eigenvector: eigenvector[i],
})
.collect();
out.sort_by(|a, b| {
b.pagerank
.partial_cmp(&a.pagerank)
.unwrap_or(std::cmp::Ordering::Equal)
.then(a.path.cmp(&b.path))
});
out
}
fn pagerank_weighted(adj: &[Vec<(usize, f64)>], weighted_degree: &[f64]) -> Vec<f64> {
let n = adj.len();
if n == 0 {
return Vec::new();
}
#[allow(clippy::cast_precision_loss)]
let n_f = n as f64;
let teleport = (1.0 - PAGERANK_DAMPING) / n_f;
let mut rank = vec![1.0 / n_f; n];
let mut next = vec![0.0_f64; n];
for _ in 0..POWER_ITER_MAX {
let dangling: f64 = rank
.iter()
.enumerate()
.filter(|(i, _)| weighted_degree[*i] == 0.0)
.map(|(_, r)| *r)
.sum();
let dangling_share = PAGERANK_DAMPING * dangling / n_f;
next.fill(teleport + dangling_share);
for u in 0..n {
if weighted_degree[u] == 0.0 {
continue;
}
let share = PAGERANK_DAMPING * rank[u] / weighted_degree[u];
for &(v, w) in &adj[u] {
next[v] += share * w;
}
}
let delta: f64 = next
.iter()
.zip(rank.iter())
.map(|(n_, r)| (n_ - r).abs())
.sum();
std::mem::swap(&mut rank, &mut next);
if delta < POWER_ITER_EPSILON {
break;
}
}
rank
}
fn eigenvector_weighted(adj: &[Vec<(usize, f64)>]) -> Vec<f64> {
let n = adj.len();
if n == 0 {
return Vec::new();
}
#[allow(clippy::cast_precision_loss)]
let n_f = n as f64;
let init = 1.0 / n_f.sqrt();
let mut vec_curr = vec![init; n];
let mut vec_next = vec![0.0_f64; n];
for _ in 0..POWER_ITER_MAX {
vec_next.copy_from_slice(&vec_curr);
for u in 0..n {
for &(v, w) in &adj[u] {
vec_next[v] += w * vec_curr[u];
}
}
let norm: f64 = vec_next.iter().map(|x| x * x).sum::<f64>().sqrt();
if norm == 0.0 {
return vec_curr;
}
for cell in &mut vec_next {
*cell /= norm;
}
let delta: f64 = vec_next
.iter()
.zip(vec_curr.iter())
.map(|(n_, c)| (n_ - c).abs())
.sum::<f64>();
std::mem::swap(&mut vec_curr, &mut vec_next);
if delta < POWER_ITER_EPSILON {
break;
}
}
vec_curr
}
pub fn from_coupling_pairs(pairs: &[CouplingRow]) -> Result<Vec<CentralityRow>> {
Ok(compute_centrality(pairs))
}
#[cfg(test)]
mod tests {
use super::*;
fn pair(a: &str, b: &str, degree: f64) -> CouplingRow {
CouplingRow {
entity_a: a.into(),
entity_b: b.into(),
shared: 1,
revs_a: 10,
revs_b: 10,
average_revs: 10,
degree,
fisher_p: 0.01,
}
}
#[test]
fn empty_input_yields_empty_output() {
assert!(compute_centrality(&[]).is_empty());
}
#[test]
fn star_graph_centrality_picks_hub() {
let pairs = vec![
pair("hub", "leaf1", 50.0),
pair("hub", "leaf2", 50.0),
pair("hub", "leaf3", 50.0),
];
let result = compute_centrality(&pairs);
assert_eq!(result.len(), 4);
let hub = result.iter().find(|r| r.path == "hub").unwrap();
assert_eq!(hub.degree, 3);
assert!((hub.weighted_degree - 150.0).abs() < 1e-9);
let leaf_pr = result
.iter()
.filter(|r| r.path.starts_with("leaf"))
.map(|r| r.pagerank)
.next()
.unwrap();
assert!(hub.pagerank > leaf_pr);
let leaf_ev = result
.iter()
.filter(|r| r.path.starts_with("leaf"))
.map(|r| r.eigenvector)
.next()
.unwrap();
assert!(hub.eigenvector > leaf_ev);
}
#[test]
fn pagerank_sums_to_one_within_epsilon() {
let pairs = vec![
pair("a", "b", 30.0),
pair("b", "c", 40.0),
pair("c", "a", 50.0),
];
let result = compute_centrality(&pairs);
let sum: f64 = result.iter().map(|r| r.pagerank).sum();
assert!((sum - 1.0).abs() < 1e-6, "pagerank sum = {sum}");
}
#[test]
fn eigenvector_is_l2_normalised() {
let pairs = vec![
pair("a", "b", 30.0),
pair("b", "c", 40.0),
pair("c", "a", 50.0),
pair("a", "d", 20.0),
];
let result = compute_centrality(&pairs);
let norm_sq: f64 = result.iter().map(|r| r.eigenvector.powi(2)).sum();
assert!(
(norm_sq - 1.0).abs() < 1e-6,
"eigenvector L2 norm² = {norm_sq}"
);
}
#[test]
fn weighted_degree_respects_edge_weights() {
let pairs = vec![pair("a", "b", 10.0), pair("a", "c", 90.0)];
let result = compute_centrality(&pairs);
let a = result.iter().find(|r| r.path == "a").unwrap();
assert!((a.weighted_degree - 100.0).abs() < 1e-9);
assert_eq!(a.degree, 2);
}
#[test]
fn output_sorted_by_pagerank_descending() {
let pairs = vec![
pair("center", "a", 100.0),
pair("center", "b", 100.0),
pair("center", "c", 100.0),
pair("center", "d", 100.0),
pair("a", "b", 10.0),
];
let result = compute_centrality(&pairs);
for w in result.windows(2) {
assert!(
w[0].pagerank >= w[1].pagerank,
"{} ({}) >= {} ({})",
w[0].path,
w[0].pagerank,
w[1].path,
w[1].pagerank,
);
}
assert_eq!(result[0].path, "center");
}
}