use std::collections::HashMap;
use std::sync::Arc;
use nodedb_cluster::{NexarTransport, ShufflePushRequest, ShufflePushStream, TypedClusterError};
use super::frame_explode::explode_row_array;
use super::inbox::ShuffleReceiverRegistry;
use crate::data::executor::response_codec::{encode_binary_rows, flatten_to_relational_rows};
const FLUSH_ROWS: usize = 1024;
enum PartTarget {
Remote {
node_id: u64,
stream: Option<ShufflePushStream>,
},
Loopback,
}
struct PartState {
target: PartTarget,
buffer: Vec<Vec<u8>>,
}
pub struct ShuffleFanoutSink {
transport: Arc<NexarTransport>,
registry: Arc<ShuffleReceiverRegistry>,
shuffle_id: u64,
side: u8,
num_parts: u32,
producer_count: u32,
keys: Vec<String>,
parts: HashMap<u32, PartState>,
max_read_version_lsn: u64,
}
pub struct ShuffleFanoutSinkParams<'a> {
pub self_node_id: u64,
pub shuffle_id: u64,
pub side: u8,
pub num_parts: u32,
pub producer_count: u32,
pub keys: Vec<String>,
pub part_node_map: &'a [(u32, u64)],
}
impl ShuffleFanoutSink {
pub fn new(
transport: Arc<NexarTransport>,
registry: Arc<ShuffleReceiverRegistry>,
params: ShuffleFanoutSinkParams<'_>,
) -> Self {
let ShuffleFanoutSinkParams {
self_node_id,
shuffle_id,
side,
num_parts,
producer_count,
keys,
part_node_map,
} = params;
let owner: HashMap<u32, u64> = part_node_map.iter().copied().collect();
let mut parts = HashMap::with_capacity(num_parts as usize);
for part in 0..num_parts {
let target = match owner.get(&part) {
Some(&node_id) if node_id != self_node_id => PartTarget::Remote {
node_id,
stream: None,
},
_ => PartTarget::Loopback,
};
parts.insert(
part,
PartState {
target,
buffer: Vec::new(),
},
);
}
Self {
transport,
registry,
shuffle_id,
side,
num_parts,
producer_count,
keys,
parts,
max_read_version_lsn: 0,
}
}
pub fn observed_read_version_lsn(&self) -> u64 {
self.max_read_version_lsn
}
fn push_request(&self, part: u32) -> ShufflePushRequest {
ShufflePushRequest {
shuffle_id: self.shuffle_id,
part,
side: self.side,
num_parts: self.num_parts,
producer_count: self.producer_count,
}
}
async fn route_chunk(&mut self, payload: Vec<u8>) -> crate::Result<()> {
let flat = flatten_to_relational_rows(&payload);
let rows = explode_row_array(&flat)?;
for row in rows {
let part =
(nodedb_query::partition_hash(row, &self.keys) % self.num_parts as u64) as u32;
let needs_flush = {
let state = self
.parts
.get_mut(&part)
.ok_or_else(|| crate::Error::Internal {
detail: format!(
"shuffle fanout: row hashed to part {part} outside 0..{}",
self.num_parts
),
})?;
state.buffer.push(row.to_vec());
state.buffer.len() >= FLUSH_ROWS
};
if needs_flush {
self.flush_part(part).await?;
}
}
Ok(())
}
async fn flush_part(&mut self, part: u32) -> crate::Result<()> {
let rows = {
let state = self
.parts
.get_mut(&part)
.ok_or_else(|| crate::Error::Internal {
detail: format!("shuffle fanout: flush of unknown part {part}"),
})?;
if state.buffer.is_empty() {
return Ok(());
}
std::mem::take(&mut state.buffer)
};
let chunk = encode_binary_rows(&rows);
let is_loopback = matches!(
self.parts.get(&part).map(|s| &s.target),
Some(PartTarget::Loopback)
);
if is_loopback {
let inbox = self.registry.get_or_create(
self.shuffle_id,
part,
self.side,
self.producer_count as usize,
);
inbox.append_chunk(&chunk).await?;
return Ok(());
}
let (node_id, opener) = {
let state = self
.parts
.get_mut(&part)
.ok_or_else(|| crate::Error::Internal {
detail: format!("shuffle fanout: flush of unknown part {part}"),
})?;
match &mut state.target {
PartTarget::Remote { node_id, stream } => (*node_id, stream.is_none()),
PartTarget::Loopback => {
return Err(crate::Error::Internal {
detail: format!(
"shuffle fanout: part {part} reclassified to loopback mid-flush"
),
});
}
}
};
if opener {
let req = self.push_request(part);
let opened = self
.transport
.open_shuffle_push_stream(node_id, req)
.await
.map_err(|e| crate::Error::Internal {
detail: format!(
"shuffle fanout: open push stream to node {node_id} for part {part}: {e}"
),
})?;
if let Some(PartState {
target: PartTarget::Remote { stream, .. },
..
}) = self.parts.get_mut(&part)
{
*stream = Some(opened);
}
}
let state = self
.parts
.get_mut(&part)
.ok_or_else(|| crate::Error::Internal {
detail: format!("shuffle fanout: flush of unknown part {part}"),
})?;
let PartTarget::Remote {
stream: Some(s), ..
} = &mut state.target
else {
return Err(crate::Error::Internal {
detail: format!("shuffle fanout: push stream absent for part {part}"),
});
};
s.push_chunk(chunk)
.await
.map_err(|e| crate::Error::Internal {
detail: format!(
"shuffle fanout: push chunk to node {node_id} for part {part}: {e}"
),
})?;
Ok(())
}
pub async fn finish(mut self, error: Option<TypedClusterError>) -> crate::Result<()> {
if error.is_none() {
for part in 0..self.num_parts {
self.flush_part(part).await?;
}
}
for part in 0..self.num_parts {
let Some(state) = self.parts.remove(&part) else {
continue;
};
match state.target {
PartTarget::Loopback => {
let inbox = self.registry.get_or_create(
self.shuffle_id,
part,
self.side,
self.producer_count as usize,
);
if let Some(e) = error.clone() {
inbox.set_error(e);
}
if inbox.record_end() {
inbox.finalize().await?;
}
}
PartTarget::Remote { node_id, stream } => {
match stream {
Some(s) => {
s.finish(error.clone())
.await
.map_err(|e| crate::Error::Internal {
detail: format!(
"shuffle fanout: finish push stream to node {node_id} \
for part {part}: {e}"
),
})?;
}
None => {
let req = self.push_request(part);
let s = self
.transport
.open_shuffle_push_stream(node_id, req)
.await
.map_err(|e| crate::Error::Internal {
detail: format!(
"shuffle fanout: open empty-part stream to node \
{node_id} for part {part}: {e}"
),
})?;
s.finish(error.clone())
.await
.map_err(|e| crate::Error::Internal {
detail: format!(
"shuffle fanout: finish empty-part stream to node \
{node_id} for part {part}: {e}"
),
})?;
}
}
}
}
}
Ok(())
}
}
impl nodedb_cluster::ChunkSink for &mut ShuffleFanoutSink {
async fn send_chunk(
&mut self,
payload: Vec<u8>,
_watermark_lsn: u64,
read_version_lsn: u64,
) -> nodedb_cluster::Result<()> {
self.max_read_version_lsn = self.max_read_version_lsn.max(read_version_lsn);
self.route_chunk(payload)
.await
.map_err(|e| nodedb_cluster::ClusterError::Storage {
detail: format!("shuffle fanout route: {e}"),
})
}
}