use std::collections::BTreeSet;
use futures::future::join_all;
use nodedb_cluster::rpc_codec::DescriptorVersionEntry;
use nodedb_cluster::{
PartNodeEntry, RaftRpc, ShuffleAggregateConsumeRequest, ShuffleAggregateConsumeResponse,
ShuffleProduceRequest, SortKey,
};
use nodedb_physical::physical_plan::wire as plan_wire;
use nodedb_physical::physical_plan::{PhysicalPlan, QueryOp};
use crate::control::server::exchange::gather::outcome_to_response;
use crate::control::server::payload_merge::{encode_msgpack_array, extract_msgpack_elements};
use crate::control::state::SharedState;
use crate::types::{DatabaseId, Lsn, TenantId, TraceId};
use super::exchange::Resolved;
use super::peers::{
distinct_data_node_count, producer_nodes, register_peers_from_topology, send_produce,
};
pub async fn resolve_shuffle_aggregate(
state: &SharedState,
database_id: DatabaseId,
tenant_id: TenantId,
child: PhysicalPlan,
keys: Vec<String>,
num_parts: usize,
trace_id: TraceId,
) -> crate::Result<Resolved> {
let PhysicalPlan::Query(QueryOp::Aggregate {
collection,
input,
group_by,
aggregates,
filters,
having,
limit,
sort_keys,
..
}) = child
else {
return Err(crate::Error::Internal {
detail: "ExchangeMode::ShuffleAggregate must wrap a QueryOp::Aggregate".into(),
});
};
if input.is_some() {
return Err(crate::Error::Internal {
detail: "shuffle aggregate requires a bare per-shard collection source \
(no embedded sub-plan)"
.into(),
});
}
if group_by.is_empty() {
return Err(crate::Error::Internal {
detail: "shuffle aggregate requires a non-empty GROUP BY key list".into(),
});
}
let group_by_fields: Vec<String> = group_by.iter().filter_map(|s| s.field.clone()).collect();
if keys != group_by_fields {
return Err(crate::Error::Internal {
detail: format!(
"shuffle aggregate: Exchange keys {keys:?} do not match wrapped \
Aggregate group_by {group_by_fields:?}"
),
});
}
let (Some(transport), Some(routing)) = (
state.cluster_transport.as_ref(),
state.cluster_routing.as_ref(),
) else {
return Err(crate::Error::Internal {
detail: "distributed shuffle aggregate 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 aggregate 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 producers = producer_nodes(&routing_snapshot, database_id, &collection)?;
let producer_count = producers.len() as u32;
if producer_count == 0 {
return Err(crate::Error::Internal {
detail: "shuffle aggregate: source collection resolved to zero producer nodes".into(),
});
}
{
let mut targets: BTreeSet<u64> = BTreeSet::new();
targets.extend(producers.iter().copied());
targets.extend(part_node_map.iter().map(|e| e.node_id));
register_peers_from_topology(state, transport, &targets);
}
let producer_plan = PhysicalPlan::Query(QueryOp::PartialAggregateState {
collection: collection.clone(),
input: None,
group_by: group_by.clone(),
aggregates: aggregates.clone(),
filters: filters.clone(),
});
let plan_bytes = plan_wire::encode(&producer_plan).map_err(|e| crate::Error::Internal {
detail: format!("shuffle aggregate: encode producer plan: {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 aggregate: num_parts {effective_num_parts} exceeds u32"),
})?;
let mut produce_futures = Vec::with_capacity(producers.len());
for &node in &producers {
let req = ShuffleProduceRequest {
shuffle_id,
side: 0,
num_parts: num_parts_u32,
producer_count,
keys: group_by_fields.clone(),
part_node_map: part_node_map.clone(),
plan_bytes: 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(),
};
produce_futures.push(send_produce(transport, node, req));
}
let mut max_read_version_lsn: u64 = 0;
for result in join_all(produce_futures).await {
max_read_version_lsn = max_read_version_lsn.max(result?);
}
let aggregates_bytes =
zerompk::to_msgpack_vec(&aggregates).map_err(|e| crate::Error::Internal {
detail: format!("shuffle aggregate: encode aggregate specs: {e}"),
})?;
let wire_sort_keys: Vec<SortKey> = sort_keys
.iter()
.map(|(column, ascending)| SortKey {
column: column.clone(),
ascending: *ascending,
})
.collect();
let mut consume_futures = Vec::with_capacity(part_node_map.len());
for entry in &part_node_map {
let req = ShuffleAggregateConsumeRequest {
shuffle_id,
part: entry.part,
group_by: group_by_fields.clone(),
aggregates_bytes: aggregates_bytes.clone(),
having: having.clone(),
limit: u64::MAX,
sort_keys: wire_sort_keys.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 mut elements: Vec<Vec<u8>> = Vec::new();
for rows in &per_part_rows {
elements.extend(extract_msgpack_elements(rows));
}
if elements.len() > limit {
elements.truncate(limit);
}
let merged = encode_msgpack_array(&elements);
Ok(Resolved::Gathered(
outcome_to_response(merged, Lsn::ZERO, Lsn::new(max_read_version_lsn)),
Vec::new(),
Vec::new(),
))
}
async fn send_consume(
transport: &nodedb_cluster::NexarTransport,
node: u64,
req: ShuffleAggregateConsumeRequest,
) -> crate::Result<Vec<u8>> {
let part = req.part;
match transport
.send_rpc(node, RaftRpc::ShuffleAggregateConsumeRequest(req))
.await
{
Ok(RaftRpc::ShuffleAggregateConsumeResponse(ShuffleAggregateConsumeResponse {
rows,
error: None,
})) => Ok(rows),
Ok(RaftRpc::ShuffleAggregateConsumeResponse(ShuffleAggregateConsumeResponse {
error: Some(e),
..
})) => Err(crate::Error::Internal {
detail: format!(
"shuffle aggregate consume failed for part {part} on node {node}: {e:?}"
),
}),
Ok(other) => Err(crate::Error::Internal {
detail: format!(
"shuffle aggregate consume: unexpected reply for part {part} from node \
{node}: {other:?}"
),
}),
Err(e) => Err(crate::Error::Internal {
detail: format!(
"shuffle aggregate consume RPC for part {part} to node {node} failed: {e}"
),
}),
}
}