use std::collections::BTreeSet;
use futures::future::{join, join_all};
use nodedb_cluster::rpc_codec::DescriptorVersionEntry;
use nodedb_cluster::{
JoinKeyPair, PartNodeEntry, RaftRpc, ShuffleConsumeRequest, ShuffleConsumeResponse,
ShuffleProduceRequest,
};
use nodedb_physical::physical_plan::wire as plan_wire;
use nodedb_physical::physical_plan::{PhysicalPlan, QueryOp};
use crate::control::server::exchange::full_scan::full_scan_plan_for_collection;
use crate::control::server::exchange::gather::outcome_to_response;
use crate::control::server::payload_merge::merge_msgpack_arrays;
use crate::control::state::SharedState;
use crate::types::{DatabaseId, Lsn, TenantId, TraceId};
use super::capture::DistributedReadCapture;
use super::exchange::Resolved;
use super::peers::{
distinct_data_node_count, producer_nodes, register_peers_from_topology, send_produce,
};
pub async fn resolve_shuffle_join(
state: &SharedState,
database_id: DatabaseId,
tenant_id: TenantId,
child: PhysicalPlan,
_keys: Vec<(String, String)>,
num_parts: usize,
trace_id: TraceId,
) -> crate::Result<Resolved> {
let PhysicalPlan::Query(QueryOp::HashJoin {
left_collection,
right_collection,
left_alias,
right_alias,
on,
join_type,
limit,
left_input,
right_input,
..
}) = child
else {
return Err(crate::Error::Internal {
detail: "ExchangeMode::Shuffle must wrap a QueryOp::HashJoin".into(),
});
};
if left_input.is_some() || right_input.is_some() {
return Err(crate::Error::Internal {
detail: "shuffle join requires both inputs as bare collection scans \
(no embedded sub-plan)"
.into(),
});
}
if on.is_empty() {
return Err(crate::Error::Internal {
detail: "shuffle join requires a non-empty equi-join key list".into(),
});
}
let (Some(transport), Some(routing)) = (
state.cluster_transport.as_ref(),
state.cluster_routing.as_ref(),
) else {
return Err(crate::Error::Internal {
detail: "distributed shuffle join requires cluster mode \
(no transport / routing table on this node)"
.into(),
});
};
let routing_snapshot = {
let guard = routing.read().unwrap_or_else(|p| p.into_inner());
guard.clone()
};
let shuffle_id = state.next_shuffle_id();
let data_node_count = distinct_data_node_count(&routing_snapshot);
let effective_num_parts = if num_parts == 0 {
data_node_count.max(1)
} else {
num_parts
};
let part_map = nodedb_cluster::distributed_join::plan_shuffle_partitions(
&routing_snapshot,
effective_num_parts,
);
if part_map.len() != effective_num_parts {
return Err(crate::Error::Internal {
detail: format!(
"shuffle partition plan incomplete: expected {effective_num_parts} parts, \
got {} (no data groups?)",
part_map.len()
),
});
}
let part_node_map: Vec<PartNodeEntry> = {
let mut entries: Vec<PartNodeEntry> = part_map
.iter()
.map(|(&part, &node_id)| PartNodeEntry { part, node_id })
.collect();
entries.sort_by_key(|e| e.part);
entries
};
let probe_keys: Vec<String> = on.iter().map(|(l, _)| l.clone()).collect();
let build_keys: Vec<String> = on.iter().map(|(_, r)| r.clone()).collect();
let build_nodes = producer_nodes(&routing_snapshot, database_id, &right_collection)?;
let probe_nodes = producer_nodes(&routing_snapshot, database_id, &left_collection)?;
let build_producer_count = build_nodes.len() as u32;
let probe_producer_count = probe_nodes.len() as u32;
if build_producer_count == 0 || probe_producer_count == 0 {
return Err(crate::Error::Internal {
detail: "shuffle join: a join side resolved to zero producer nodes".into(),
});
}
{
let mut targets: BTreeSet<u64> = BTreeSet::new();
targets.extend(build_nodes.iter().copied());
targets.extend(probe_nodes.iter().copied());
targets.extend(part_node_map.iter().map(|e| e.node_id));
register_peers_from_topology(state, transport, &targets);
}
let build_scan = require_scan_plan(state, database_id, tenant_id, &right_collection)?;
let probe_scan = require_scan_plan(state, database_id, tenant_id, &left_collection)?;
let build_plan_bytes = plan_wire::encode(&build_scan).map_err(|e| crate::Error::Internal {
detail: format!("shuffle join: encode build scan: {e}"),
})?;
let probe_plan_bytes = plan_wire::encode(&probe_scan).map_err(|e| crate::Error::Internal {
detail: format!("shuffle join: encode probe scan: {e}"),
})?;
let deadline_remaining_ms = state
.tuning
.network
.default_deadline_secs
.saturating_mul(1000)
.max(1);
let num_parts_u32 = u32::try_from(effective_num_parts).map_err(|_| crate::Error::Internal {
detail: format!("shuffle join: num_parts {effective_num_parts} exceeds u32"),
})?;
let mut build_produce_futures = Vec::with_capacity(build_nodes.len());
for &node in &build_nodes {
let req = ShuffleProduceRequest {
shuffle_id,
side: 0,
num_parts: num_parts_u32,
producer_count: build_producer_count,
keys: build_keys.clone(),
part_node_map: part_node_map.clone(),
plan_bytes: build_plan_bytes.clone(),
tenant_id: tenant_id.as_u64(),
database_id: database_id.as_u64(),
deadline_remaining_ms,
trace_id: trace_id.0,
descriptor_versions: Vec::<DescriptorVersionEntry>::new(),
};
build_produce_futures.push(send_produce(transport, node, req));
}
let mut probe_produce_futures = Vec::with_capacity(probe_nodes.len());
for &node in &probe_nodes {
let req = ShuffleProduceRequest {
shuffle_id,
side: 1,
num_parts: num_parts_u32,
producer_count: probe_producer_count,
keys: probe_keys.clone(),
part_node_map: part_node_map.clone(),
plan_bytes: probe_plan_bytes.clone(),
tenant_id: tenant_id.as_u64(),
database_id: database_id.as_u64(),
deadline_remaining_ms,
trace_id: trace_id.0,
descriptor_versions: Vec::<DescriptorVersionEntry>::new(),
};
probe_produce_futures.push(send_produce(transport, node, req));
}
let (build_results, probe_results) = join(
join_all(build_produce_futures),
join_all(probe_produce_futures),
)
.await;
let mut build_rv: u64 = 0;
for result in build_results {
build_rv = build_rv.max(result?);
}
let mut probe_rv: u64 = 0;
for result in probe_results {
probe_rv = probe_rv.max(result?);
}
let on_pairs: Vec<JoinKeyPair> = on
.iter()
.map(|(l, r)| JoinKeyPair {
left: l.clone(),
right: r.clone(),
})
.collect();
let limit_u64 = u64::try_from(limit).unwrap_or(u64::MAX);
let probe_qualifier = left_alias.unwrap_or_default();
let index_qualifier = right_alias.unwrap_or_default();
let mut consume_futures = Vec::with_capacity(part_node_map.len());
for entry in &part_node_map {
let req = ShuffleConsumeRequest {
shuffle_id,
part: entry.part,
on: on_pairs.clone(),
join_type: join_type.clone(),
limit: limit_u64,
probe_qualifier: probe_qualifier.clone(),
index_qualifier: index_qualifier.clone(),
tenant_id: tenant_id.as_u64(),
database_id: database_id.as_u64(),
deadline_remaining_ms,
trace_id: trace_id.0,
};
consume_futures.push(send_consume(transport, entry.node_id, req));
}
let mut per_part_rows: Vec<Vec<u8>> = Vec::with_capacity(consume_futures.len());
for result in join_all(consume_futures).await {
per_part_rows.push(result?);
}
let merged = merge_msgpack_arrays(&per_part_rows);
let captures = vec![
DistributedReadCapture {
scan_plan: probe_scan,
read_version_lsn: Lsn::new(probe_rv),
},
DistributedReadCapture {
scan_plan: build_scan,
read_version_lsn: Lsn::new(build_rv),
},
];
Ok(Resolved::Gathered(
outcome_to_response(merged, Lsn::ZERO, Lsn::ZERO),
Vec::new(),
captures,
))
}
async fn send_consume(
transport: &nodedb_cluster::NexarTransport,
node: u64,
req: ShuffleConsumeRequest,
) -> crate::Result<Vec<u8>> {
let part = req.part;
match transport
.send_rpc(node, RaftRpc::ShuffleConsumeRequest(req))
.await
{
Ok(RaftRpc::ShuffleConsumeResponse(ShuffleConsumeResponse { rows, error: None })) => {
Ok(rows)
}
Ok(RaftRpc::ShuffleConsumeResponse(ShuffleConsumeResponse { error: Some(e), .. })) => {
Err(crate::Error::Internal {
detail: format!("shuffle consume failed for part {part} on node {node}: {e:?}"),
})
}
Ok(other) => Err(crate::Error::Internal {
detail: format!(
"shuffle consume: unexpected reply for part {part} from node {node}: {other:?}"
),
}),
Err(e) => Err(crate::Error::Internal {
detail: format!("shuffle consume RPC for part {part} to node {node} failed: {e}"),
}),
}
}
fn require_scan_plan(
state: &SharedState,
database_id: DatabaseId,
tenant_id: TenantId,
collection: &str,
) -> crate::Result<PhysicalPlan> {
full_scan_plan_for_collection(state, database_id, tenant_id, collection)?.ok_or_else(|| {
crate::Error::Internal {
detail: format!(
"shuffle join: no catalog entry for collection '{collection}' on coordinator"
),
}
})
}