use std::collections::HashMap;
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex};
use nodedb_cluster::TypedClusterError;
use tokio::io::AsyncWriteExt;
use super::frame_explode::explode_row_array;
pub type ShuffleKey = (u64, u32, u8);
pub struct ShuffleInbox {
path: PathBuf,
writer: tokio::sync::Mutex<Option<tokio::fs::File>>,
producer_count: usize,
ends_received: AtomicUsize,
error: Mutex<Option<TypedClusterError>>,
finalized: AtomicBool,
finalized_notify: tokio::sync::Notify,
}
impl ShuffleInbox {
pub fn new(path: PathBuf, producer_count: usize) -> Self {
Self {
path,
writer: tokio::sync::Mutex::new(None),
producer_count: producer_count.max(1),
ends_received: AtomicUsize::new(0),
error: Mutex::new(None),
finalized: AtomicBool::new(false),
finalized_notify: tokio::sync::Notify::new(),
}
}
pub fn producer_count(&self) -> usize {
self.producer_count
}
pub fn staged_path(&self) -> &Path {
&self.path
}
pub async fn append_chunk(&self, chunk_payload: &[u8]) -> crate::Result<()> {
let frames = explode_row_array(chunk_payload)?;
let mut guard = self.writer.lock().await;
if guard.is_none() {
*guard = Some(self.open_staging().await?);
}
let Some(file) = guard.as_mut() else {
return Err(crate::Error::Storage {
engine: "shuffle-stage".into(),
detail: format!(
"staging writer unexpectedly absent for {}",
self.path.display()
),
});
};
for row in frames {
let len = u32::try_from(row.len()).map_err(|_| crate::Error::Storage {
engine: "shuffle-stage".into(),
detail: format!(
"shuffle row exceeds u32 frame length ({} bytes) staging to {}",
row.len(),
self.path.display()
),
})?;
file.write_all(&len.to_le_bytes()).await?;
file.write_all(row).await?;
}
Ok(())
}
pub async fn finalize(&self) -> crate::Result<()> {
let mut guard = self.writer.lock().await;
if guard.is_none() {
*guard = Some(self.open_staging().await?);
}
if let Some(file) = guard.as_mut() {
file.flush().await?;
file.sync_all().await?;
}
self.finalized.store(true, Ordering::Release);
self.finalized_notify.notify_waiters();
Ok(())
}
pub fn is_finalized(&self) -> bool {
self.finalized.load(Ordering::Acquire)
}
pub async fn wait_finalized(&self) {
let notified = self.finalized_notify.notified();
tokio::pin!(notified);
notified.as_mut().enable();
if self.is_finalized() {
return;
}
notified.await;
}
async fn open_staging(&self) -> crate::Result<tokio::fs::File> {
if let Some(parent) = self.path.parent() {
tokio::fs::create_dir_all(parent)
.await
.map_err(|e| crate::Error::Storage {
engine: "shuffle-stage".into(),
detail: format!("create shuffle staging dir {}: {e}", parent.display()),
})?;
}
tokio::fs::OpenOptions::new()
.create(true)
.write(true)
.truncate(true)
.open(&self.path)
.await
.map_err(|e| crate::Error::Storage {
engine: "shuffle-stage".into(),
detail: format!("open shuffle staging file {}: {e}", self.path.display()),
})
}
pub fn record_end(&self) -> bool {
let prev = self.ends_received.fetch_add(1, Ordering::AcqRel);
prev + 1 >= self.producer_count
}
pub fn ends_received(&self) -> usize {
self.ends_received.load(Ordering::Acquire)
}
pub fn barrier_complete(&self) -> bool {
self.ends_received.load(Ordering::Acquire) >= self.producer_count
}
pub fn set_error(&self, error: TypedClusterError) {
let mut slot = self.error.lock().unwrap_or_else(|p| p.into_inner());
if slot.is_none() {
*slot = Some(error);
}
}
pub fn take_error(&self) -> Option<TypedClusterError> {
self.error.lock().unwrap_or_else(|p| p.into_inner()).take()
}
}
pub struct ShuffleReceiverRegistry {
base_dir: PathBuf,
inboxes: Mutex<HashMap<ShuffleKey, Arc<ShuffleInbox>>>,
}
impl ShuffleReceiverRegistry {
pub fn new(base_dir: PathBuf) -> Self {
Self {
base_dir,
inboxes: Mutex::new(HashMap::new()),
}
}
fn shuffle_dir(&self, shuffle_id: u64) -> PathBuf {
self.base_dir
.join("shuffle-stage")
.join(shuffle_id.to_string())
}
fn staged_path(&self, shuffle_id: u64, part: u32, side: u8) -> PathBuf {
self.shuffle_dir(shuffle_id)
.join(format!("{part}-{side}.frames"))
}
pub fn get_or_create(
&self,
shuffle_id: u64,
part: u32,
side: u8,
producer_count: usize,
) -> Arc<ShuffleInbox> {
let key = (shuffle_id, part, side);
let mut map = self.inboxes.lock().unwrap_or_else(|p| p.into_inner());
if let Some(existing) = map.get(&key) {
return Arc::clone(existing);
}
let path = self.staged_path(shuffle_id, part, side);
let inbox = Arc::new(ShuffleInbox::new(path, producer_count));
map.insert(key, Arc::clone(&inbox));
inbox
}
pub fn get(&self, key: ShuffleKey) -> Option<Arc<ShuffleInbox>> {
self.inboxes
.lock()
.unwrap_or_else(|p| p.into_inner())
.get(&key)
.map(Arc::clone)
}
pub fn unregister_shuffle(&self, shuffle_id: u64) {
self.inboxes
.lock()
.unwrap_or_else(|p| p.into_inner())
.retain(|(sid, _, _), _| *sid != shuffle_id);
let dir = self.shuffle_dir(shuffle_id);
if let Err(e) = std::fs::remove_dir_all(&dir)
&& e.kind() != std::io::ErrorKind::NotFound
{
tracing::warn!(
shuffle_id,
dir = %dir.display(),
error = %e,
"failed to remove shuffle staging dir"
);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn temp_base() -> (tempfile::TempDir, ShuffleReceiverRegistry) {
let dir = tempfile::tempdir().expect("tempdir");
let reg = ShuffleReceiverRegistry::new(dir.path().to_path_buf());
(dir, reg)
}
fn row(fields: &[(&str, serde_json::Value)]) -> Vec<u8> {
let mut map = serde_json::Map::new();
for (k, v) in fields {
map.insert((*k).to_string(), v.clone());
}
nodedb_types::json_to_msgpack(&serde_json::Value::Object(map)).expect("encode row")
}
fn encode_array(rows: &[Vec<u8>]) -> Vec<u8> {
crate::data::executor::response_codec::encode_binary_rows(rows)
}
fn read_staged(path: &Path) -> Vec<Vec<u8>> {
let bytes = std::fs::read(path).expect("read staged file");
let mut out = Vec::new();
let mut pos = 0usize;
while pos + 4 <= bytes.len() {
let len = u32::from_le_bytes(bytes[pos..pos + 4].try_into().expect("len")) as usize;
pos += 4;
assert!(pos + len <= bytes.len(), "frame body truncated");
out.push(bytes[pos..pos + len].to_vec());
pos += len;
}
assert_eq!(pos, bytes.len(), "trailing bytes after last frame");
out
}
#[tokio::test]
async fn finalize_without_any_append_creates_empty_staged_file() {
let (_d, reg) = temp_base();
let inbox = reg.get_or_create(20, 0, 0, 1);
assert!(!inbox.staged_path().exists(), "no file before finalize");
inbox.finalize().await.expect("finalize");
assert!(
inbox.staged_path().exists(),
"finalize must create the staged file even with zero rows"
);
assert!(
read_staged(inbox.staged_path()).is_empty(),
"the zero-row staged file holds no frames"
);
assert!(inbox.is_finalized());
}
#[tokio::test]
async fn append_chunk_explodes_array_into_per_row_frames() {
let (_d, reg) = temp_base();
let inbox = reg.get_or_create(1, 0, 0, 1);
let rows = vec![
row(&[("k", serde_json::json!(1))]),
row(&[("k", serde_json::json!(2))]),
row(&[("k", serde_json::json!(3))]),
];
inbox
.append_chunk(&encode_array(&rows))
.await
.expect("append");
inbox.finalize().await.expect("finalize");
let staged = read_staged(inbox.staged_path());
assert_eq!(staged, rows, "each array element becomes one frame");
}
#[tokio::test]
async fn append_chunk_is_appending_across_chunks() {
let (_d, reg) = temp_base();
let inbox = reg.get_or_create(2, 0, 1, 1);
let a = vec![row(&[("k", serde_json::json!("a"))])];
let b = vec![
row(&[("k", serde_json::json!("b"))]),
row(&[("k", serde_json::json!("c"))]),
];
inbox.append_chunk(&encode_array(&a)).await.expect("a");
inbox.append_chunk(&encode_array(&b)).await.expect("b");
inbox.finalize().await.expect("finalize");
let staged = read_staged(inbox.staged_path());
let mut want = a.clone();
want.extend(b.clone());
assert_eq!(staged, want);
}
#[tokio::test]
async fn empty_chunk_array_stages_no_frames() {
let (_d, reg) = temp_base();
let inbox = reg.get_or_create(3, 0, 0, 1);
inbox.append_chunk(&encode_array(&[])).await.expect("empty");
inbox.finalize().await.expect("finalize");
let staged = read_staged(inbox.staged_path());
assert!(staged.is_empty());
}
#[tokio::test]
async fn malformed_chunk_is_hard_error() {
let (_d, reg) = temp_base();
let inbox = reg.get_or_create(4, 0, 0, 1);
let bad = vec![0x91u8];
let res = inbox.append_chunk(&bad).await;
assert!(
matches!(res, Err(crate::Error::Storage { .. })),
"a malformed chunk must surface a Storage error, never a silent drop"
);
}
#[test]
fn barrier_fires_only_after_all_producers_end() {
let (_d, reg) = temp_base();
let inbox = reg.get_or_create(5, 0, 0, 2);
assert!(!inbox.barrier_complete());
assert!(!inbox.record_end());
assert!(!inbox.barrier_complete());
assert_eq!(inbox.ends_received(), 1);
assert!(inbox.record_end());
assert!(inbox.barrier_complete());
assert_eq!(inbox.ends_received(), 2);
}
#[test]
fn single_producer_barrier_fires_on_first_end() {
let (_d, reg) = temp_base();
let inbox = reg.get_or_create(6, 0, 0, 1);
assert!(!inbox.barrier_complete());
assert!(inbox.record_end());
assert!(inbox.barrier_complete());
}
#[test]
fn error_capture_first_writer_wins() {
let (_d, reg) = temp_base();
let inbox = reg.get_or_create(7, 0, 0, 1);
assert!(inbox.take_error().is_none());
inbox.set_error(TypedClusterError::Internal {
code: 1,
message: "first".into(),
});
inbox.set_error(TypedClusterError::Internal {
code: 2,
message: "second".into(),
});
match inbox.take_error() {
Some(TypedClusterError::Internal { code, .. }) => assert_eq!(code, 1),
other => panic!("expected first Internal error, got {other:?}"),
}
assert!(inbox.take_error().is_none());
}
#[tokio::test]
async fn wait_finalized_wakes_a_parked_waiter() {
let (_d, reg) = temp_base();
let inbox = reg.get_or_create(20, 0, 0, 1);
assert!(!inbox.is_finalized());
let waiter = Arc::clone(&inbox);
let handle = tokio::spawn(async move {
waiter.wait_finalized().await;
});
tokio::task::yield_now().await;
inbox.finalize().await.expect("finalize");
tokio::time::timeout(std::time::Duration::from_secs(5), handle)
.await
.expect("waiter must wake within 5s")
.expect("waiter task joined");
assert!(inbox.is_finalized());
}
#[tokio::test]
async fn wait_finalized_returns_immediately_when_already_finalized() {
let (_d, reg) = temp_base();
let inbox = reg.get_or_create(21, 0, 0, 1);
inbox.finalize().await.expect("finalize");
assert!(inbox.is_finalized());
tokio::time::timeout(std::time::Duration::from_secs(1), inbox.wait_finalized())
.await
.expect("already-finalized wait must return immediately");
}
#[tokio::test]
async fn failed_finalize_does_not_mark_finalized() {
let (_d, reg) = temp_base();
let inbox = reg.get_or_create(22, 0, 0, 1);
assert!(!inbox.is_finalized());
inbox.finalize().await.expect("no-op finalize");
assert!(
inbox.is_finalized(),
"a successful (no-op) finalize marks the inbox finalized"
);
}
#[test]
fn registry_get_or_create_is_idempotent() {
let (_d, reg) = temp_base();
let a = reg.get_or_create(10, 0, 0, 2);
let b = reg.get_or_create(10, 0, 0, 99);
assert!(Arc::ptr_eq(&a, &b), "same key must reuse the same inbox");
assert_eq!(a.producer_count(), 2, "first creator's producer_count wins");
let c = reg.get_or_create(10, 1, 0, 1);
assert!(!Arc::ptr_eq(&a, &c));
}
#[test]
fn registry_get_returns_none_for_missing() {
let (_d, reg) = temp_base();
assert!(reg.get((11, 0, 0)).is_none());
reg.get_or_create(11, 0, 0, 1);
assert!(reg.get((11, 0, 0)).is_some());
}
#[test]
fn staged_path_is_deterministic_and_scoped() {
let (_d, reg) = temp_base();
let inbox = reg.get_or_create(12, 3, 1, 1);
let p = inbox.staged_path();
assert!(p.ends_with("shuffle-stage/12/3-1.frames"), "path: {p:?}");
}
#[tokio::test]
async fn unregister_removes_inboxes_and_scratch_dir() {
let (_d, reg) = temp_base();
let inbox = reg.get_or_create(13, 0, 0, 1);
inbox
.append_chunk(&encode_array(&[row(&[("k", serde_json::json!(1))])]))
.await
.expect("append");
inbox.finalize().await.expect("finalize");
let dir = inbox.staged_path().parent().unwrap().to_path_buf();
assert!(dir.exists(), "staging dir created");
reg.get_or_create(13, 1, 1, 1);
reg.get_or_create(14, 0, 0, 1);
reg.unregister_shuffle(13);
assert!(reg.get((13, 0, 0)).is_none());
assert!(reg.get((13, 1, 1)).is_none());
assert!(reg.get((14, 0, 0)).is_some());
assert!(!dir.exists(), "scratch dir removed for shuffle 13");
}
}