use futures::future::join_all;
use crate::bridge::envelope::{Payload, PhysicalPlan};
use crate::control::gateway::version_set::GatewayVersionSet;
use crate::control::server::graph_dispatch::bsp_pagerank::enumerate::ShardTarget;
use crate::control::server::graph_dispatch::cluster_resolve::{
DispatchSuperstepParams, dispatch_superstep_to_node, gateway_shared,
};
use crate::types::{DatabaseId, TenantId};
use nodedb_graph::AlgoParams;
use nodedb_physical::physical_plan::{GraphOp, WccSuperstepPlan, WccSuperstepResult};
pub(super) struct ShardWccResult {
pub(super) result: WccSuperstepResult,
}
pub(super) async fn scatter_wcc_round(
state: &crate::control::state::SharedState,
tenant_id: TenantId,
database_id: DatabaseId,
params: &AlgoParams,
targets: &[ShardTarget],
deadline_ms: u64,
) -> crate::Result<Vec<ShardWccResult>> {
let shared_arc = gateway_shared(state)?;
let version_set = GatewayVersionSet::from_pairs(Vec::new());
let futs = targets.iter().map(|t| {
let plan = PhysicalPlan::Graph(GraphOp::WccSuperstep(Box::new(WccSuperstepPlan {
params: params.clone(),
owned_vshards: t.owned_vshards.clone(),
})));
let version_set = version_set.clone();
let node_id = t.node_id;
let is_local = t.is_local;
let route_vshard = t.route_vshard();
let shared_arc = shared_arc.clone();
Box::pin(async move {
let payload = dispatch_superstep_to_node(
&shared_arc,
DispatchSuperstepParams {
tenant_id,
database_id,
deadline_ms,
node_id,
is_local,
route_vshard,
plan,
version_set: &version_set,
},
)
.await?;
let result = decode_wcc_from_payload(node_id, payload)?;
Ok::<ShardWccResult, crate::Error>(ShardWccResult { result })
})
});
let results = join_all(futs).await;
let mut out = Vec::with_capacity(results.len());
for res in results {
out.push(res?);
}
Ok(out)
}
fn decode_wcc_from_payload(node_id: u64, payload: Payload) -> crate::Result<WccSuperstepResult> {
if payload.is_empty() {
return Ok(WccSuperstepResult::default());
}
zerompk::from_msgpack::<WccSuperstepResult>(payload.as_ref()).map_err(|e| crate::Error::Codec {
detail: format!("wcc: node={node_id} result decode: {e}"),
})
}