use std::collections::HashMap;
use nodedb_cluster::distributed_graph::{BspCoordinator, SuperstepAck};
use nodedb_graph::{AlgoParams, GraphAlgorithm};
use crate::bridge::envelope::Payload;
use crate::control::state::SharedState;
use crate::engine::graph::algo::result::AlgoResultBatch;
use crate::types::{DatabaseId, TenantId};
use super::enumerate::enumerate_shards;
use super::scatter::{ScatterSuperstepParams, ShardDispatch, scatter_superstep};
const DEFAULT_MAX_ITERATIONS: u32 = 20;
struct ShardRankState {
is_local: bool,
owned_vshards: Vec<u32>,
route_vshard: u32,
node_names: Vec<String>,
rank_vec: Vec<f64>,
}
pub async fn run_bsp_pagerank(
state: &SharedState,
tenant_id: TenantId,
database_id: DatabaseId,
params: AlgoParams,
deadline_ms: u64,
) -> crate::Result<Payload> {
let algorithm = GraphAlgorithm::PageRank;
let enumeration = enumerate_shards(state)?;
let targets = enumeration.targets;
let vshard_owner = enumeration.vshard_owner;
if targets.is_empty() {
return empty_payload();
}
let global_seed_sum: f64 = params
.personalization_vector()
.map(|seed| seed.values().map(|&w| w.max(0.0)).sum())
.unwrap_or(0.0);
let max_iterations = params
.max_iterations
.map(|m| m.clamp(1, u32::MAX as usize) as u32)
.unwrap_or(DEFAULT_MAX_ITERATIONS);
let tolerance = params.convergence_tolerance();
let shard_ids: Vec<u32> = targets.iter().map(|t| t.node_id as u32).collect();
let mut bsp = BspCoordinator::new(
algorithm.name().to_string(),
max_iterations,
tolerance,
shard_ids,
);
let count_dispatches: Vec<ShardDispatch> = targets
.iter()
.map(|t| ShardDispatch {
node_id: t.node_id,
is_local: t.is_local,
owned_vshards: t.owned_vshards.clone(),
route_vshard: t.route_vshard(),
incoming_contributions: Vec::new(),
rank_seed: Vec::new(),
global_dangling: 0.0, personalization_sum: 0.0, })
.collect();
let counts = scatter_superstep(
state,
ScatterSuperstepParams {
tenant_id,
database_id,
algorithm,
params: ¶ms,
superstep: 0,
global_n: 0, dispatches: count_dispatches,
deadline_ms,
},
)
.await?;
let global_n: usize = counts.iter().map(|c| c.result.vertex_count).sum();
if global_n == 0 {
return empty_payload();
}
let global_seed_hits: usize = counts.iter().map(|c| c.result.seed_hits).sum();
let personalization_sum = if global_seed_sum > 0.0 && global_seed_hits > 0 {
global_seed_sum
} else {
0.0
};
let mut shard_state: HashMap<u64, ShardRankState> = HashMap::with_capacity(targets.len());
for (target, count) in targets.iter().zip(counts) {
shard_state.insert(
target.node_id,
ShardRankState {
is_local: target.is_local,
owned_vshards: target.owned_vshards.clone(),
route_vshard: target.route_vshard(),
node_names: count.result.node_names,
rank_vec: Vec::new(),
},
);
}
let mut incoming: HashMap<u64, Vec<(String, f64)>> = HashMap::new();
let mut global_dangling: f64 = 0.0;
let mut superstep: u32 = 0;
loop {
let mut ordered_nodes: Vec<u64> = shard_state.keys().copied().collect();
ordered_nodes.sort_unstable();
let dispatches: Vec<ShardDispatch> = ordered_nodes
.iter()
.map(|&node_id| {
let st = &shard_state[&node_id];
ShardDispatch {
node_id,
is_local: st.is_local,
owned_vshards: st.owned_vshards.clone(),
route_vshard: st.route_vshard,
incoming_contributions: incoming.remove(&node_id).unwrap_or_default(),
rank_seed: st
.node_names
.iter()
.cloned()
.zip(st.rank_vec.iter().copied())
.collect(),
global_dangling,
personalization_sum,
}
})
.collect();
let results = scatter_superstep(
state,
ScatterSuperstepParams {
tenant_id,
database_id,
algorithm,
params: ¶ms,
superstep,
global_n,
dispatches,
deadline_ms,
},
)
.await?;
incoming.clear();
global_dangling = 0.0;
for sr in results {
let node_id = sr.node_id;
let res = sr.result;
bsp.record_ack(SuperstepAck {
shard_id: node_id as u32,
iteration: superstep + 1,
local_delta: res.local_delta,
vertex_count: res.vertex_count,
contributions_sent: res.outbound.len(),
});
global_dangling += res.dangling_sum;
for (target_vshard, dst_name, contrib) in res.outbound {
let Some(&owner) = vshard_owner.get(&target_vshard) else {
return Err(crate::Error::Internal {
detail: format!(
"bsp pagerank: outbound contribution to unmapped target \
vshard={target_vshard} (dst={dst_name})"
),
});
};
if !shard_state.contains_key(&owner) {
return Err(crate::Error::Internal {
detail: format!(
"bsp pagerank: outbound contribution to unknown owner \
node={owner} for vshard={target_vshard} (dst={dst_name})"
),
});
}
incoming.entry(owner).or_default().push((dst_name, contrib));
}
if let Some(st) = shard_state.get_mut(&node_id) {
st.node_names = res.node_names;
st.rank_vec = res.rank_vec;
}
}
if !bsp.all_acked() {
return Err(crate::Error::Internal {
detail: "bsp pagerank: not all shards acked after superstep dispatch".into(),
});
}
if !bsp.advance() {
break;
}
superstep += 1;
}
assemble_result(&shard_state)
}
fn assemble_result(shard_state: &HashMap<u64, ShardRankState>) -> crate::Result<Payload> {
let mut batch = AlgoResultBatch::new(GraphAlgorithm::PageRank);
let mut ordered: Vec<u64> = shard_state.keys().copied().collect();
ordered.sort_unstable();
for node_id in ordered {
let st = &shard_state[&node_id];
for (name, rank) in st.node_names.iter().zip(st.rank_vec.iter()) {
batch.push_node_f64(name.clone(), *rank);
}
}
let bytes = batch.to_msgpack()?;
Ok(Payload::from_vec(bytes))
}
fn empty_payload() -> crate::Result<Payload> {
let bytes = AlgoResultBatch::new(GraphAlgorithm::PageRank).to_msgpack()?;
Ok(Payload::from_vec(bytes))
}