use std::sync::Arc;
use nodedb_cluster::TypedClusterError;
use super::inbox::ShuffleReceiverRegistry;
pub struct RegistryShuffleReceiver {
pub registry: Arc<ShuffleReceiverRegistry>,
}
impl RegistryShuffleReceiver {
pub fn new(registry: Arc<ShuffleReceiverRegistry>) -> Self {
Self { registry }
}
}
#[async_trait::async_trait]
impl nodedb_cluster::ShuffleReceiver for RegistryShuffleReceiver {
async fn on_shuffle_request(&self, shuffle_id: u64, part: u32, side: u8, producer_count: u32) {
self.registry
.get_or_create(shuffle_id, part, side, producer_count as usize);
}
async fn on_shuffle_chunk(
&self,
shuffle_id: u64,
part: u32,
side: u8,
payload: Vec<u8>,
) -> nodedb_cluster::Result<()> {
let inbox = self
.registry
.get((shuffle_id, part, side))
.unwrap_or_else(|| self.registry.get_or_create(shuffle_id, part, side, 1));
match inbox.append_chunk(&payload).await {
Ok(()) => Ok(()),
Err(e) => {
let detail = format!("shuffle stage append ({shuffle_id},{part},{side}): {e}");
inbox.set_error(TypedClusterError::Internal {
code: 0,
message: detail.clone(),
});
Err(nodedb_cluster::ClusterError::Storage { detail })
}
}
}
async fn on_shuffle_end(
&self,
shuffle_id: u64,
part: u32,
side: u8,
error: Option<TypedClusterError>,
) {
let inbox = self
.registry
.get((shuffle_id, part, side))
.unwrap_or_else(|| self.registry.get_or_create(shuffle_id, part, side, 1));
if let Some(e) = error {
inbox.set_error(e);
}
if inbox.record_end()
&& let Err(e) = inbox.finalize().await
{
inbox.set_error(TypedClusterError::Internal {
code: 0,
message: format!("shuffle stage finalize ({shuffle_id},{part},{side}): {e}"),
});
}
}
}