1use super::message::ShuffleMessage;
8#[cfg(feature = "cluster")]
9use crate::checkpoint::CheckpointAttempt;
10
11const SHUFFLE_RECV_QUEUE: usize = 256;
14const MAX_STAGE_NAME_BYTES: usize = 4096;
15
16pub type ShufflePeerId = u64;
18
19const SCOPE_CANCELLED: &str = "shuffle assignment or recovery scope was cancelled";
20const NONCANONICAL_BARRIER: &str =
21 "shuffle checkpoint barrier must use one nonzero canonical checkpoint ID";
22
23fn validate_checkpoint_barrier(
24 barrier: crate::checkpoint::CheckpointBarrier,
25) -> std::io::Result<()> {
26 if barrier.is_canonical() {
27 Ok(())
28 } else {
29 Err(std::io::Error::new(
30 std::io::ErrorKind::InvalidInput,
31 NONCANONICAL_BARRIER,
32 ))
33 }
34}
35
36fn validate_frontier(stage: &str, watermark: Option<i64>) -> std::io::Result<()> {
37 if stage.is_empty() || stage.len() > MAX_STAGE_NAME_BYTES {
38 return Err(std::io::Error::new(
39 std::io::ErrorKind::InvalidInput,
40 "shuffle stage scope is empty or too long",
41 ));
42 }
43 if watermark == Some(i64::MIN) {
44 return Err(std::io::Error::new(
45 std::io::ErrorKind::InvalidInput,
46 "shuffle frontier uses the reserved uninitialized watermark",
47 ));
48 }
49 Ok(())
50}
51
52#[derive(Debug)]
53struct ScopeCancelled;
54
55impl std::fmt::Display for ScopeCancelled {
56 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
57 f.write_str(SCOPE_CANCELLED)
58 }
59}
60
61impl std::error::Error for ScopeCancelled {}
62
63#[cfg(feature = "cluster")]
64#[derive(Debug)]
65struct MayHaveAdmittedShuffleFrame(std::io::Error);
66
67#[cfg(feature = "cluster")]
68impl std::fmt::Display for MayHaveAdmittedShuffleFrame {
69 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
70 self.0.fmt(f)
71 }
72}
73
74#[cfg(feature = "cluster")]
75impl std::error::Error for MayHaveAdmittedShuffleFrame {
76 fn source(&self) -> Option<&(dyn std::error::Error + 'static)> {
77 Some(&self.0)
78 }
79}
80
81#[cfg(feature = "cluster")]
82fn may_have_admitted_shuffle_frame(error: std::io::Error) -> std::io::Error {
83 std::io::Error::new(error.kind(), MayHaveAdmittedShuffleFrame(error))
84}
85
86#[cfg(feature = "cluster")]
87fn scope_cancelled_io() -> std::io::Error {
88 std::io::Error::new(std::io::ErrorKind::ConnectionAborted, ScopeCancelled)
89}
90
91#[must_use]
94pub fn is_scope_cancelled(error: &std::io::Error) -> bool {
95 error
96 .get_ref()
97 .and_then(|source| source.downcast_ref::<ScopeCancelled>())
98 .is_some()
99}
100
101#[cfg(feature = "cluster")]
103#[must_use]
104pub fn shuffle_send_may_have_been_admitted(error: &std::io::Error) -> bool {
105 error
106 .get_ref()
107 .and_then(|source| source.downcast_ref::<MayHaveAdmittedShuffleFrame>())
108 .is_some()
109}
110
111#[cfg(feature = "cluster")]
114pub const SHUFFLE_ADDR_KEY: &str = "shuffle:addr";
115
116#[cfg(feature = "cluster")]
117#[allow(
118 clippy::doc_markdown,
119 clippy::default_trait_access,
120 clippy::missing_const_for_fn,
121 clippy::must_use_candidate,
122 clippy::too_many_lines,
123 missing_docs
124)]
125pub(crate) mod shuffle_v1 {
126 tonic::include_proto!("laminar.shuffle.v1");
127}
128
129struct Holdover {
132 staged: parking_lot::Mutex<rustc_hash::FxHashMap<String, Vec<ReceivedBatch>>>,
133 frontiers: parking_lot::Mutex<Vec<ReceivedFrontierCut>>,
134 barriers: parking_lot::Mutex<BarrierHoldover>,
135 items: std::sync::atomic::AtomicUsize,
138 capacity: usize,
139}
140
141#[derive(Default)]
142struct BarrierHoldover {
143 staged: Vec<ReceivedShuffle>,
144 #[cfg(feature = "cluster")]
145 retired_through: Option<RetiredCheckpoint>,
146}
147
148#[cfg(feature = "cluster")]
149#[derive(Clone, Copy)]
150struct RetiredCheckpoint {
151 attempt: CheckpointAttempt,
152 assignment_digest: [u8; 32],
153}
154
155impl BarrierHoldover {
156 #[cfg(feature = "cluster")]
157 fn is_retired(
158 &self,
159 attempt: CheckpointAttempt,
160 assignment_digest: Option<[u8; 32]>,
161 ) -> std::io::Result<bool> {
162 let Some(retired) = self.retired_through else {
163 return Ok(false);
164 };
165 match attempt.checkpoint_id.cmp(&retired.attempt.checkpoint_id) {
166 std::cmp::Ordering::Less => Ok(true),
167 std::cmp::Ordering::Greater => Ok(false),
168 std::cmp::Ordering::Equal
169 if assignment_digest == Some(retired.assignment_digest) =>
170 {
171 Ok(true)
172 }
173 std::cmp::Ordering::Equal => Err(std::io::Error::new(
174 std::io::ErrorKind::InvalidData,
175 format!(
176 "retired checkpoint barrier {attempt:?} has a different assignment digest from its durable terminal outcome"
177 ),
178 )),
179 }
180 }
181}
182
183impl Holdover {
184 fn new(capacity: usize) -> Self {
185 Self {
186 staged: parking_lot::Mutex::default(),
187 frontiers: parking_lot::Mutex::default(),
188 barriers: parking_lot::Mutex::default(),
189 items: std::sync::atomic::AtomicUsize::new(0),
190 capacity,
191 }
192 }
193
194 fn try_reserve_item(&self) -> bool {
195 self.items
196 .fetch_update(
197 std::sync::atomic::Ordering::AcqRel,
198 std::sync::atomic::Ordering::Acquire,
199 |items| (items < self.capacity).then_some(items + 1),
200 )
201 .is_ok()
202 }
203
204 fn release_items(&self, count: usize) {
205 if count == 0 {
206 return;
207 }
208 let released = self.items.fetch_update(
209 std::sync::atomic::Ordering::AcqRel,
210 std::sync::atomic::Ordering::Acquire,
211 |items| items.checked_sub(count),
212 );
213 debug_assert!(
214 released.is_ok(),
215 "shuffle holdover item accounting underflow"
216 );
217 }
218
219 #[cfg(feature = "cluster")]
220 fn barrier_attempt(received: &ReceivedShuffle) -> Option<CheckpointAttempt> {
221 let ShuffleMessage::Barrier(barrier) = received.message() else {
222 return None;
223 };
224 let attempt = CheckpointAttempt::new(barrier.epoch, barrier.checkpoint_id);
225 attempt.is_canonical().then_some(attempt)
226 }
227
228 #[cfg(feature = "cluster")]
229 fn is_retired_barrier(&self, received: &ReceivedShuffle) -> std::io::Result<bool> {
230 let Some(attempt) = Self::barrier_attempt(received) else {
231 return Ok(false);
232 };
233 self.barriers
234 .lock()
235 .is_retired(attempt, received.assignment_digest)
236 }
237
238 #[cfg(feature = "cluster")]
239 fn is_retired_checkpoint_barrier(
240 &self,
241 barrier: crate::checkpoint::CheckpointBarrier,
242 assignment_digest: [u8; 32],
243 ) -> std::io::Result<bool> {
244 let attempt = CheckpointAttempt::new(barrier.epoch, barrier.checkpoint_id);
245 if !attempt.is_canonical() {
246 return Err(std::io::Error::new(
247 std::io::ErrorKind::InvalidInput,
248 NONCANONICAL_BARRIER,
249 ));
250 }
251 self.barriers
252 .lock()
253 .is_retired(attempt, Some(assignment_digest))
254 }
255
256 fn has_staged_barriers(&self) -> bool {
257 !self.barriers.lock().staged.is_empty()
258 }
259
260 fn stage_barrier(&self, barrier: ReceivedShuffle) -> std::io::Result<bool> {
261 let ShuffleMessage::Barrier(value) = barrier.message() else {
262 return Ok(false);
263 };
264 if !value.is_canonical() {
265 return Err(std::io::Error::new(
266 std::io::ErrorKind::InvalidInput,
267 NONCANONICAL_BARRIER,
268 ));
269 }
270 let mut holdover = self.barriers.lock();
271 #[cfg(feature = "cluster")]
272 if let Some(attempt) = Self::barrier_attempt(&barrier) {
273 if holdover.is_retired(attempt, barrier.assignment_digest)? {
274 return Ok(false);
275 }
276 }
277 holdover.staged.push(barrier);
278 Ok(true)
279 }
280
281 fn take_staged_barriers(&self) -> Vec<ReceivedShuffle> {
282 std::mem::take(&mut self.barriers.lock().staged)
283 }
284
285 fn has_staged_frontiers(&self) -> bool {
286 !self.frontiers.lock().is_empty()
287 }
288
289 fn stage_frontier(&self, frontier: ReceivedFrontierCut) {
290 self.frontiers.lock().push(frontier);
291 }
292
293 fn take_staged_frontiers(&self) -> Vec<ReceivedFrontierCut> {
294 std::mem::take(&mut *self.frontiers.lock())
295 }
296
297 #[cfg(feature = "cluster")]
298 fn clear_staged_frontiers(&self) {
299 let frontiers = self.take_staged_frontiers();
300 let item_count = frontiers.iter().map(ReceivedFrontierCut::item_count).sum();
301 self.release_items(item_count);
302 }
303
304 #[cfg(feature = "cluster")]
305 fn retire_checkpoint_attempt(
306 &self,
307 attempt: CheckpointAttempt,
308 assignment_digest: [u8; 32],
309 ) -> std::io::Result<()> {
310 if !attempt.is_canonical() {
311 return Err(std::io::Error::new(
312 std::io::ErrorKind::InvalidInput,
313 "retired checkpoint attempt must use one nonzero canonical checkpoint ID",
314 ));
315 }
316
317 let removed = {
318 let mut holdover = self.barriers.lock();
319 let retired = match holdover.retired_through {
320 None => RetiredCheckpoint {
321 attempt,
322 assignment_digest,
323 },
324 Some(retired) if attempt.checkpoint_id > retired.attempt.checkpoint_id => {
325 RetiredCheckpoint {
326 attempt,
327 assignment_digest,
328 }
329 }
330 Some(retired) if attempt.checkpoint_id == retired.attempt.checkpoint_id => {
331 if assignment_digest != retired.assignment_digest {
332 return Err(std::io::Error::new(
333 std::io::ErrorKind::InvalidData,
334 format!(
335 "checkpoint retirement {attempt:?} has a different assignment digest from its high-water"
336 ),
337 ));
338 }
339 retired
340 }
341 Some(retired) => retired,
342 };
343 for barrier in &holdover.staged {
344 let Some(candidate) = Self::barrier_attempt(barrier) else {
345 continue;
346 };
347 if candidate.checkpoint_id == retired.attempt.checkpoint_id
348 && barrier.assignment_digest != Some(retired.assignment_digest)
349 {
350 return Err(std::io::Error::new(
351 std::io::ErrorKind::InvalidData,
352 format!(
353 "staged checkpoint barrier {candidate:?} has a different assignment digest from its durable terminal outcome"
354 ),
355 ));
356 }
357 }
358 holdover.retired_through = Some(retired);
359 let staged_before = holdover.staged.len();
360 holdover.staged.retain(|barrier| {
361 !Self::barrier_attempt(barrier).is_some_and(|candidate| {
362 candidate.checkpoint_id <= retired.attempt.checkpoint_id
363 })
364 });
365 staged_before - holdover.staged.len()
366 };
367 self.release_items(removed);
368 Ok(())
369 }
370}
371
372impl Default for Holdover {
373 fn default() -> Self {
374 Self::new(SHUFFLE_RECV_QUEUE)
375 }
376}
377
378pub struct ReceivedBatch {
384 batch: arrow_array::RecordBatch,
385 routed_vnodes: std::sync::Arc<[u32]>,
386 reservation: Option<std::sync::Arc<InboundReservation>>,
387 peer: ShufflePeerId,
388 sender_incarnation: uuid::Uuid,
389 receiver_incarnation: uuid::Uuid,
390 stream_id: uuid::Uuid,
391 assignment_version: u64,
392 recovery_gen: u64,
393 checkpoint_sequence: u64,
394}
395
396#[must_use]
399#[derive(Clone)]
400pub struct ShuffleBatchAdmission(Option<std::sync::Arc<InboundReservation>>);
401
402impl std::fmt::Debug for ShuffleBatchAdmission {
403 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
404 f.debug_struct("ShuffleBatchAdmission")
405 .field("admitted", &self.0.is_some())
406 .finish()
407 }
408}
409
410impl ReceivedBatch {
411 #[must_use]
413 pub const fn batch(&self) -> &arrow_array::RecordBatch {
414 &self.batch
415 }
416
417 #[must_use]
419 pub fn routed_vnodes(&self) -> &[u32] {
420 &self.routed_vnodes
421 }
422
423 #[must_use]
425 pub fn routed_vnodes_arc(&self) -> std::sync::Arc<[u32]> {
426 std::sync::Arc::clone(&self.routed_vnodes)
427 }
428
429 #[must_use]
431 pub const fn peer(&self) -> ShufflePeerId {
432 self.peer
433 }
434
435 #[must_use]
437 pub const fn sender_incarnation(&self) -> uuid::Uuid {
438 self.sender_incarnation
439 }
440
441 #[must_use]
443 pub const fn receiver_incarnation(&self) -> uuid::Uuid {
444 self.receiver_incarnation
445 }
446
447 #[must_use]
449 pub const fn stream_id(&self) -> uuid::Uuid {
450 self.stream_id
451 }
452
453 #[must_use]
455 pub const fn assignment_version(&self) -> u64 {
456 self.assignment_version
457 }
458
459 #[must_use]
461 pub const fn recovery_gen(&self) -> u64 {
462 self.recovery_gen
463 }
464
465 #[must_use]
467 pub const fn checkpoint_sequence(&self) -> u64 {
468 self.checkpoint_sequence
469 }
470
471 pub fn into_parts(self) -> (arrow_array::RecordBatch, ShuffleBatchAdmission) {
475 (self.batch, ShuffleBatchAdmission(self.reservation))
476 }
477}
478
479impl std::fmt::Debug for ReceivedBatch {
480 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
481 f.debug_struct("ReceivedBatch")
482 .field("batch", &self.batch)
483 .field("routed_vnodes", &self.routed_vnodes)
484 .field("peer", &self.peer)
485 .field("sender_incarnation", &self.sender_incarnation)
486 .field("receiver_incarnation", &self.receiver_incarnation)
487 .field("stream_id", &self.stream_id)
488 .field("assignment_version", &self.assignment_version)
489 .field("recovery_gen", &self.recovery_gen)
490 .field("checkpoint_sequence", &self.checkpoint_sequence)
491 .field("admitted", &self.reservation.is_some())
492 .finish()
493 }
494}
495
496pub struct ReceivedShuffle {
500 peer: ShufflePeerId,
501 message: ShuffleMessage,
502 reservation: Option<std::sync::Arc<InboundReservation>>,
503 sender_incarnation: uuid::Uuid,
504 receiver_incarnation: uuid::Uuid,
505 stream_id: uuid::Uuid,
506 assignment_version: u64,
507 assignment_digest: Option<[u8; 32]>,
508 recovery_gen: u64,
509 checkpoint_sequence: u64,
510}
511
512impl ReceivedShuffle {
513 #[must_use]
515 pub const fn peer(&self) -> ShufflePeerId {
516 self.peer
517 }
518
519 #[must_use]
521 pub const fn message(&self) -> &ShuffleMessage {
522 &self.message
523 }
524
525 #[must_use]
527 pub const fn sender_incarnation(&self) -> uuid::Uuid {
528 self.sender_incarnation
529 }
530
531 #[must_use]
533 pub const fn receiver_incarnation(&self) -> uuid::Uuid {
534 self.receiver_incarnation
535 }
536
537 #[must_use]
539 pub const fn stream_id(&self) -> uuid::Uuid {
540 self.stream_id
541 }
542
543 #[must_use]
545 pub const fn assignment_version(&self) -> u64 {
546 self.assignment_version
547 }
548
549 #[must_use]
551 pub const fn assignment_digest(&self) -> Option<[u8; 32]> {
552 self.assignment_digest
553 }
554
555 #[must_use]
557 pub const fn recovery_gen(&self) -> u64 {
558 self.recovery_gen
559 }
560
561 #[must_use]
564 pub const fn checkpoint_sequence(&self) -> u64 {
565 self.checkpoint_sequence
566 }
567
568 pub fn into_parts(self) -> (ShuffleMessage, ShuffleBatchAdmission) {
571 (self.message, ShuffleBatchAdmission(self.reservation))
572 }
573}
574
575impl std::fmt::Debug for ReceivedShuffle {
576 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
577 f.debug_struct("ReceivedShuffle")
578 .field("peer", &self.peer)
579 .field("message", &self.message)
580 .field("stream_id", &self.stream_id)
581 .field("assignment_version", &self.assignment_version)
582 .field("recovery_gen", &self.recovery_gen)
583 .field("checkpoint_sequence", &self.checkpoint_sequence)
584 .field("admitted", &self.reservation.is_some())
585 .finish_non_exhaustive()
586 }
587}
588
589#[derive(Debug)]
591pub struct ReceivedFrontierCut {
592 preceding: Vec<ReceivedBatch>,
593 frontier: ReceivedShuffle,
594}
595
596impl ReceivedFrontierCut {
597 #[must_use]
599 pub fn preceding(&self) -> &[ReceivedBatch] {
600 &self.preceding
601 }
602
603 #[must_use]
605 pub const fn frontier(&self) -> &ReceivedShuffle {
606 &self.frontier
607 }
608
609 #[must_use]
611 pub fn into_parts(self) -> (Vec<ReceivedBatch>, ReceivedShuffle) {
612 (self.preceding, self.frontier)
613 }
614
615 fn item_count(&self) -> usize {
616 self.preceding.len() + 1
617 }
618
619 fn stage(&self) -> &str {
620 let ShuffleMessage::Frontier { stage, .. } = self.frontier.message() else {
621 unreachable!("frontier cut contains a frontier message");
622 };
623 stage
624 }
625}
626
627fn take_frontier_prefix(
628 staged: &mut rustc_hash::FxHashMap<String, Vec<ReceivedBatch>>,
629 stage: &str,
630 frontier: &ReceivedShuffle,
631) -> Vec<ReceivedBatch> {
632 let Some(batches) = staged.remove(stage) else {
633 return Vec::new();
634 };
635 let (preceding, remaining): (Vec<_>, Vec<_>) = batches.into_iter().partition(|batch| {
636 batch.peer == frontier.peer
637 && batch.sender_incarnation == frontier.sender_incarnation
638 && batch.receiver_incarnation == frontier.receiver_incarnation
639 && batch.stream_id == frontier.stream_id
640 && batch.assignment_version == frontier.assignment_version
641 && batch.recovery_gen == frontier.recovery_gen
642 && batch.checkpoint_sequence < frontier.checkpoint_sequence
643 });
644 if !remaining.is_empty() {
645 staged.insert(stage.to_owned(), remaining);
646 }
647 preceding
648}
649
650#[cfg(feature = "cluster")]
651struct InboundReservation {
652 node: tokio::sync::OwnedSemaphorePermit,
653 peer: tokio::sync::OwnedSemaphorePermit,
654 wire_bytes: usize,
655}
656
657#[cfg(not(feature = "cluster"))]
658struct InboundReservation;
659
660#[cfg(feature = "cluster")]
661mod grpc {
662 use std::collections::hash_map::Entry;
663 use std::collections::VecDeque;
664 use std::io;
665 use std::net::SocketAddr;
666 use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
667 use std::sync::{Arc, OnceLock};
668
669 use arrow_array::RecordBatch;
670 use bytes::Bytes;
671 use crossfire::{mpsc, AsyncRx, MAsyncTx};
672 use futures::StreamExt as _;
673 use parking_lot::{Mutex, RwLock};
674 use rustc_hash::FxHashMap;
675 use tokio::sync::{OwnedSemaphorePermit, Semaphore};
676 use tokio::task::JoinHandle;
677 use tokio_util::sync::CancellationToken;
678 use tonic::transport::{Channel, Server};
679 use tonic::Request;
680 use uuid::Uuid;
681
682 use super::shuffle_v1::shuffle_frame;
683 use super::shuffle_v1::shuffle_transport_client::ShuffleTransportClient;
684 use super::shuffle_v1::shuffle_transport_server::{ShuffleTransport, ShuffleTransportServer};
685 use super::shuffle_v1::{
686 Barrier, Frontier as WireFrontier, HandshakeRequest, HandshakeResponse, Hello, RoutedData,
687 ShuffleFrame, ShuffleSummary,
688 };
689 use super::{
690 is_scope_cancelled, may_have_admitted_shuffle_frame, scope_cancelled_io,
691 take_frontier_prefix, validate_checkpoint_barrier, validate_frontier, CheckpointAttempt,
692 Holdover, InboundReservation, ReceivedBatch, ReceivedFrontierCut, ReceivedShuffle,
693 ShuffleMessage, ShufflePeerId, MAX_STAGE_NAME_BYTES, NONCANONICAL_BARRIER, SCOPE_CANCELLED,
694 SHUFFLE_ADDR_KEY, SHUFFLE_RECV_QUEUE,
695 };
696 use crate::checkpoint::{CheckpointAssignmentFence, CheckpointBarrier};
697 use crate::cluster::control::{ClusterKv, LeaseDeadline};
698 use crate::serialization::{serialize_batch_stream_bounded, BatchStreamDecoder};
699
700 const SEND_QUEUE: usize = 256;
701
702 const OUTBOUND_PEER_BUDGET_BYTES: usize = 32 * 1024 * 1024;
703 const CHECKPOINTED_CONTROL_PEER_BUDGET_BYTES: usize = 256 * 1024;
704 const CHECKPOINTED_CONTROL_NODE_BUDGET_BYTES: usize = 4 * 1024 * 1024;
705 const OUTBOUND_NODE_BUDGET_BYTES: usize = 128 * 1024 * 1024;
706 const INBOUND_PEER_BUDGET_BYTES: usize = 32 * 1024 * 1024;
707 const INBOUND_NODE_BUDGET_BYTES: usize = 128 * 1024 * 1024;
708 const MAX_SOURCE_SCHEMA_MEMORY_BYTES: usize = 256 * 1024;
709 const MAX_SCHEMA_WIRE_BYTES: usize = 512 * 1024;
712 const MAX_DECODED_SCHEMA_MEMORY_BYTES: usize = 1024 * 1024;
715 const MAX_DECODED_ARRAY_STRUCTURE_BYTES: usize = 2 * MAX_DECODED_SCHEMA_MEMORY_BYTES;
716 const MAX_ROUTE_METADATA_BYTES: usize =
717 crate::state::MAX_KEY_GROUP_COUNT as usize * std::mem::size_of::<u32>();
718 const FRAME_METADATA_BYTES: usize = 64 * 1024;
721 const RETAINED_BATCH_ENVELOPE_BYTES: usize = 4 * 1024;
722 const INBOUND_BATCH_METADATA_BYTES: usize = MAX_DECODED_SCHEMA_MEMORY_BYTES
723 + MAX_DECODED_ARRAY_STRUCTURE_BYTES
724 + (2 * MAX_ROUTE_METADATA_BYTES)
725 + (2 * MAX_STAGE_NAME_BYTES)
726 + FRAME_METADATA_BYTES;
727 const MAX_WIRE_PAYLOAD_BYTES: usize = 1024 * 1024;
728 const _: () = assert!(MAX_SCHEMA_WIRE_BYTES + 8 <= MAX_WIRE_PAYLOAD_BYTES);
729 const MAX_FRAGMENTS: usize =
730 crate::shuffle::message::MAX_PAYLOAD_BYTES / MAX_WIRE_PAYLOAD_BYTES;
731 const BLOCKING_IPC_THRESHOLD_BYTES: usize = 512 * 1024;
732 const FRONTIER_WORKSPACE_BYTES: usize = 2 * MAX_STAGE_NAME_BYTES + 1024;
733 const MAX_SHUFFLE_MESSAGE_BYTES: usize = 2 * 1024 * 1024;
736 const OUTBOUND_DATA_WORKSPACE_BYTES: usize = crate::shuffle::message::MAX_PAYLOAD_BYTES
737 + crate::shuffle::ROUTE_MAX_BATCH_BYTES
738 + MAX_SOURCE_SCHEMA_MEMORY_BYTES
739 + (2 * MAX_ROUTE_METADATA_BYTES)
740 + (2 * MAX_STAGE_NAME_BYTES)
741 + FRAME_METADATA_BYTES;
742 const UNADMITTED_FRAME_BUDGET_BYTES: usize = 256 * 1024 * 1024;
745 const MAX_ACTIVE_STREAMS: usize = UNADMITTED_FRAME_BUDGET_BYTES / MAX_SHUFFLE_MESSAGE_BYTES;
746 const _: () = assert!(MAX_ACTIVE_STREAMS >= crate::checkpoint::MAX_CHECKPOINT_PARTICIPANTS - 1);
747 const MAX_TRACKED_PEERS: usize = 4096;
748 const MAX_PENDING_HANDSHAKES: usize = 4096;
749 const HANDSHAKE_TOKEN_TTL: std::time::Duration = std::time::Duration::from_secs(30);
750 const PROCESS_LEASE_EXPIRED: &str = "shuffle process lease is no longer live";
751
752 type InboundRx = AsyncRx<mpsc::Array<Inbound>>;
753 type InboundTx = MAsyncTx<mpsc::Array<Inbound>>;
754
755 fn io_err<E: std::fmt::Display>(e: E) -> io::Error {
756 io::Error::other(e.to_string())
757 }
758
759 fn scope_cancelled_status() -> tonic::Status {
760 tonic::Status::cancelled(SCOPE_CANCELLED)
761 }
762
763 fn status_io(status: tonic::Status) -> io::Error {
764 if status.code() == tonic::Code::Cancelled && status.message() == SCOPE_CANCELLED {
765 scope_cancelled_io()
766 } else {
767 io_err(status)
768 }
769 }
770
771 fn process_lease_expired_io() -> io::Error {
772 io::Error::new(io::ErrorKind::PermissionDenied, PROCESS_LEASE_EXPIRED)
773 }
774
775 fn process_lease_expired_status() -> tonic::Status {
776 tonic::Status::failed_precondition(PROCESS_LEASE_EXPIRED)
777 }
778
779 struct ProcessLeaseGate {
780 deadline: OnceLock<Arc<LeaseDeadline>>,
781 cancelled: CancellationToken,
782 watcher: Mutex<Option<tokio::task::AbortHandle>>,
783 }
784
785 impl Default for ProcessLeaseGate {
786 fn default() -> Self {
787 Self {
788 deadline: OnceLock::new(),
789 cancelled: CancellationToken::new(),
790 watcher: Mutex::new(None),
791 }
792 }
793 }
794
795 impl ProcessLeaseGate {
796 fn is_installed_deadline(&self, deadline: &Arc<LeaseDeadline>) -> bool {
797 self.deadline
798 .get()
799 .is_some_and(|current| Arc::ptr_eq(current, deadline))
800 }
801
802 fn install(&self, deadline: Arc<LeaseDeadline>) -> io::Result<()> {
803 if !deadline.is_live() {
804 return Err(process_lease_expired_io());
805 }
806 let mut watcher_slot = self.watcher.lock();
807 if let Some(current) = self.deadline.get() {
808 return if Arc::ptr_eq(current, &deadline) {
809 Ok(())
810 } else {
811 Err(io::Error::new(
812 io::ErrorKind::AlreadyExists,
813 "shuffle process lease deadline is already installed",
814 ))
815 };
816 }
817 let runtime = tokio::runtime::Handle::try_current().map_err(|error| {
818 io::Error::other(format!(
819 "shuffle process lease requires a Tokio runtime: {error}"
820 ))
821 })?;
822 self.deadline.set(Arc::clone(&deadline)).map_err(|_| {
823 io::Error::new(
824 io::ErrorKind::AlreadyExists,
825 "shuffle process lease deadline is already installed",
826 )
827 })?;
828 let cancelled = self.cancelled.clone();
829 let watcher = runtime.spawn(async move {
830 deadline.wait_until_expired().await;
831 cancelled.cancel();
832 });
833 *watcher_slot = Some(watcher.abort_handle());
834 Ok(())
835 }
836
837 fn install_pair(
838 first: &Self,
839 second: &Self,
840 deadline: Arc<LeaseDeadline>,
841 ) -> io::Result<()> {
842 let mut first_watcher = first.watcher.lock();
843 let mut second_watcher = second.watcher.lock();
844 if !deadline.is_live() {
845 return Err(process_lease_expired_io());
846 }
847 for gate in [first, second] {
848 if gate
849 .deadline
850 .get()
851 .is_some_and(|current| !Arc::ptr_eq(current, &deadline))
852 {
853 return Err(io::Error::new(
854 io::ErrorKind::AlreadyExists,
855 "shuffle process lease deadline is already installed",
856 ));
857 }
858 }
859
860 let install_first = first.deadline.get().is_none();
861 let install_second = second.deadline.get().is_none();
862 let runtime = if install_first || install_second {
863 Some(tokio::runtime::Handle::try_current().map_err(|error| {
864 io::Error::other(format!(
865 "shuffle process lease requires a Tokio runtime: {error}"
866 ))
867 })?)
868 } else {
869 None
870 };
871
872 if install_first {
873 assert!(
874 first.deadline.set(Arc::clone(&deadline)).is_ok(),
875 "first shuffle lease gate changed while its watcher lock was held"
876 );
877 }
878 if install_second {
879 assert!(
880 second.deadline.set(Arc::clone(&deadline)).is_ok(),
881 "second shuffle lease gate changed while its watcher lock was held"
882 );
883 }
884
885 if let Some(runtime) = runtime {
886 if install_first {
887 let cancelled = first.cancelled.clone();
888 let first_deadline = Arc::clone(&deadline);
889 let watcher = runtime.spawn(async move {
890 first_deadline.wait_until_expired().await;
891 cancelled.cancel();
892 });
893 *first_watcher = Some(watcher.abort_handle());
894 }
895 if install_second {
896 let cancelled = second.cancelled.clone();
897 let watcher = runtime.spawn(async move {
898 deadline.wait_until_expired().await;
899 cancelled.cancel();
900 });
901 *second_watcher = Some(watcher.abort_handle());
902 }
903 }
904 Ok(())
905 }
906
907 fn require_live_io(&self) -> io::Result<()> {
908 match self.deadline.get() {
909 Some(deadline) if deadline.is_live() => Ok(()),
910 Some(_) => Err(process_lease_expired_io()),
911 None => Err(io::Error::new(
912 io::ErrorKind::PermissionDenied,
913 "shuffle process lease deadline is not installed",
914 )),
915 }
916 }
917
918 fn require_live_status(&self) -> Result<(), tonic::Status> {
919 match self.deadline.get() {
920 Some(deadline) if deadline.is_live() => Ok(()),
921 Some(_) => Err(process_lease_expired_status()),
922 None => Err(tonic::Status::failed_precondition(
923 "shuffle process lease deadline is not installed",
924 )),
925 }
926 }
927
928 async fn wait_until_lost(&self) {
929 let Some(deadline) = self.deadline.get() else {
930 return;
931 };
932 if !deadline.is_live() {
933 return;
934 }
935 tokio::select! {
936 biased;
937 () = self.cancelled.cancelled() => {}
938 () = deadline.wait_until_expired() => {}
939 }
940 }
941
942 fn scope_token(&self, active: bool) -> CancellationToken {
943 let token = self.cancelled.child_token();
944 if !active || self.require_live_io().is_err() {
945 token.cancel();
946 }
947 token
948 }
949
950 #[cfg(test)]
951 fn install_live_for_test(&self) {
952 self.deadline
953 .set(Arc::new(LeaseDeadline::live_for(
954 std::time::Duration::from_secs(60),
955 )))
956 .expect("test process lease is installed once");
957 }
958 }
959
960 impl Drop for ProcessLeaseGate {
961 fn drop(&mut self) {
962 if let Some(watcher) = self.watcher.get_mut().take() {
963 watcher.abort();
964 }
965 }
966 }
967
968 fn cancelled_token() -> CancellationToken {
969 let token = CancellationToken::new();
970 token.cancel();
971 token
972 }
973
974 fn rotate_scope_token(
975 slot: &RwLock<CancellationToken>,
976 process_lease: &ProcessLeaseGate,
977 active: bool,
978 ) {
979 let mut token = slot.write();
980 token.cancel();
981 *token = process_lease.scope_token(active);
982 }
983
984 enum PreparedMessage {
986 Barrier(CheckpointBarrier),
987 Frontier {
988 stage: String,
989 watermark: Option<i64>,
990 idle: bool,
991 },
992 Data {
993 stage: String,
994 routed_vnodes: Vec<u32>,
995 arrow_ipc: Bytes,
996 },
997 }
998
999 struct Outbound {
1000 gen: u64,
1001 assignment_version: u64,
1002 seq: u64,
1004 msg: PreparedMessage,
1005 assignment_digest: Option<[u8; 32]>,
1007 _budget: OutboundReservation,
1008 }
1009
1010 struct Encoded {
1011 frames: VecDeque<ShuffleFrame>,
1012 _budget: OutboundReservation,
1013 }
1014
1015 struct OutboundReservation {
1016 peer: OwnedSemaphorePermit,
1017 node: OwnedSemaphorePermit,
1018 }
1019
1020 impl OutboundReservation {
1021 fn shrink_to(&mut self, retained_bytes: usize) -> io::Result<()> {
1022 let peer_bytes = self.peer.num_permits();
1023 let node_bytes = self.node.num_permits();
1024 if peer_bytes != node_bytes || retained_bytes > peer_bytes {
1025 return Err(io::Error::other(
1026 "shuffle outbound reservation accounting mismatch",
1027 ));
1028 }
1029 let release = peer_bytes - retained_bytes;
1030 if release != 0 {
1031 let peer = self
1032 .peer
1033 .split(release)
1034 .expect("validated outbound peer reservation split");
1035 let node = self
1036 .node
1037 .split(release)
1038 .expect("validated outbound node reservation split");
1039 drop((peer, node));
1040 }
1041 Ok(())
1042 }
1043 }
1044
1045 struct Inbound {
1046 peer: ShufflePeerId,
1047 msg: ShuffleMessage,
1048 budget: Option<Arc<InboundReservation>>,
1049 fence: StreamFence,
1050 assignment_digest: Option<[u8; 32]>,
1051 checkpoint_sequence: u64,
1052 }
1053
1054 impl Inbound {
1055 fn into_received(self) -> ReceivedShuffle {
1056 ReceivedShuffle {
1057 peer: self.peer,
1058 message: self.msg,
1059 reservation: self.budget,
1060 sender_incarnation: self.fence.sender_incarnation,
1061 receiver_incarnation: self.fence.receiver_incarnation,
1062 stream_id: self.fence.stream_id,
1063 assignment_version: self.fence.assignment_version,
1064 assignment_digest: self.assignment_digest,
1065 recovery_gen: self.fence.recovery_gen,
1066 checkpoint_sequence: self.checkpoint_sequence,
1067 }
1068 }
1069 }
1070
1071 struct InboundBudget {
1072 node: Arc<Semaphore>,
1073 peers: Mutex<FxHashMap<ShufflePeerId, Arc<Semaphore>>>,
1074 }
1075
1076 impl InboundBudget {
1077 fn new(node_capacity: usize) -> Self {
1078 Self {
1079 node: Arc::new(Semaphore::new(node_capacity)),
1080 peers: Mutex::new(FxHashMap::default()),
1081 }
1082 }
1083
1084 async fn reserve_frame(
1085 &self,
1086 peer: ShufflePeerId,
1087 wire_bytes: usize,
1088 cancel: &CancellationToken,
1089 ) -> Result<InboundReservation, tonic::Status> {
1090 let bytes = wire_bytes
1091 .checked_add(crate::shuffle::ROUTE_MAX_BATCH_BYTES)
1092 .and_then(|bytes| bytes.checked_add(INBOUND_BATCH_METADATA_BYTES))
1093 .and_then(|bytes| bytes.checked_add(MAX_WIRE_PAYLOAD_BYTES))
1096 .ok_or_else(|| tonic::Status::resource_exhausted("shuffle frame is too large"))?;
1097 if wire_bytes == 0
1098 || wire_bytes > crate::shuffle::message::MAX_PAYLOAD_BYTES
1099 || bytes > INBOUND_PEER_BUDGET_BYTES
1100 {
1101 return Err(tonic::Status::resource_exhausted(
1102 "shuffle frame exceeds its inbound byte budget",
1103 ));
1104 }
1105 let permits = u32::try_from(bytes)
1106 .map_err(|_| tonic::Status::resource_exhausted("shuffle frame is too large"))?;
1107 let peer_budget = {
1108 let mut peers = self.peers.lock();
1109 peers.retain(|known_peer, budget| {
1110 *known_peer == peer
1111 || Arc::strong_count(budget) > 1
1112 || budget.available_permits() != INBOUND_PEER_BUDGET_BYTES
1113 });
1114 if let Some(budget) = peers.get(&peer) {
1115 Arc::clone(budget)
1116 } else {
1117 if peers.len() >= MAX_TRACKED_PEERS {
1118 return Err(tonic::Status::resource_exhausted("too many shuffle peers"));
1119 }
1120 let budget = Arc::new(Semaphore::new(INBOUND_PEER_BUDGET_BYTES));
1121 peers.insert(peer, Arc::clone(&budget));
1122 budget
1123 }
1124 };
1125 let peer_permit = tokio::select! {
1126 biased;
1127 () = cancel.cancelled() => return Err(scope_cancelled_status()),
1128 permit = peer_budget.acquire_many_owned(permits) => {
1129 permit.map_err(|_| tonic::Status::unavailable("shuffle peer budget closed"))?
1130 }
1131 };
1132 let node_permit = tokio::select! {
1133 biased;
1134 () = cancel.cancelled() => return Err(scope_cancelled_status()),
1135 permit = Arc::clone(&self.node).acquire_many_owned(permits) => {
1136 permit.map_err(|_| tonic::Status::unavailable("shuffle node budget closed"))?
1137 }
1138 };
1139 Ok(InboundReservation {
1140 node: node_permit,
1141 peer: peer_permit,
1142 wire_bytes,
1143 })
1144 }
1145
1146 fn validate_decoded(batches: &[RecordBatch]) -> Result<usize, tonic::Status> {
1147 let bytes = batches.iter().try_fold(0usize, |total, batch| {
1148 let batch_bytes = crate::shuffle::routing::logical_batch_bytes(batch)
1149 .map_err(|error| tonic::Status::invalid_argument(error.to_string()))?;
1150 total.checked_add(batch_bytes).ok_or_else(|| {
1151 tonic::Status::resource_exhausted("decoded shuffle payload size overflow")
1152 })
1153 })?;
1154 if bytes > crate::shuffle::ROUTE_MAX_BATCH_BYTES {
1155 return Err(tonic::Status::resource_exhausted(format!(
1156 "decoded shuffle payload is {bytes} bytes; limit is {}",
1157 crate::shuffle::ROUTE_MAX_BATCH_BYTES
1158 )));
1159 }
1160 Ok(bytes)
1161 }
1162 }
1163
1164 impl InboundReservation {
1165 fn retain_decoded(
1166 &mut self,
1167 decoded_bytes: usize,
1168 metadata_bytes: usize,
1169 ) -> Result<(), tonic::Status> {
1170 let retained_bytes = self
1171 .wire_bytes
1172 .checked_add(decoded_bytes)
1173 .and_then(|bytes| bytes.checked_add(metadata_bytes))
1174 .ok_or_else(|| tonic::Status::internal("shuffle inbound accounting overflow"))?;
1175 let node_bytes = self.node.num_permits();
1176 let peer_bytes = self.peer.num_permits();
1177 if node_bytes != peer_bytes || retained_bytes > node_bytes {
1178 return Err(tonic::Status::internal(
1179 "shuffle inbound reservation accounting mismatch",
1180 ));
1181 }
1182
1183 let excess_bytes = node_bytes - retained_bytes;
1184 if excess_bytes != 0 {
1185 let node_excess = self
1186 .node
1187 .split(excess_bytes)
1188 .expect("validated node reservation split");
1189 let peer_excess = self
1190 .peer
1191 .split(excess_bytes)
1192 .expect("validated peer reservation split");
1193 drop((node_excess, peer_excess));
1194 }
1195 Ok(())
1196 }
1197 }
1198
1199 fn decode_ipc_payload(
1200 decoder: &mut BatchStreamDecoder,
1201 payload: Vec<u8>,
1202 ) -> Result<RecordBatch, arrow_schema::ArrowError> {
1203 let mut batches = decoder.decode_chunk(payload)?;
1204 decoder.ensure_message_boundary()?;
1205 if batches.len() != 1 {
1206 return Err(arrow_schema::ArrowError::IpcError(format!(
1207 "logical shuffle payload decoded {} record batches; expected exactly one",
1208 batches.len()
1209 )));
1210 }
1211 Ok(batches.pop().expect("validated one decoded batch"))
1212 }
1213
1214 async fn decode_ipc_payload_isolated<F>(
1215 payload: Vec<u8>,
1216 budget: InboundReservation,
1217 before_blocking_decode: F,
1218 ) -> Result<(RecordBatch, InboundReservation), String>
1219 where
1220 F: FnOnce() + Send + 'static,
1221 {
1222 if payload.len() >= BLOCKING_IPC_THRESHOLD_BYTES {
1223 tokio::task::spawn_blocking(move || {
1224 before_blocking_decode();
1225 let mut decoder = BatchStreamDecoder::new();
1226 let batch =
1227 decode_ipc_payload(&mut decoder, payload).map_err(|error| error.to_string())?;
1228 Ok((batch, budget))
1229 })
1230 .await
1231 .map_err(|error| format!("shuffle decoder task: {error}"))?
1232 } else {
1233 let mut decoder = BatchStreamDecoder::new();
1234 let batch =
1235 decode_ipc_payload(&mut decoder, payload).map_err(|error| error.to_string())?;
1236 Ok((batch, budget))
1237 }
1238 }
1239
1240 fn schema_memory_size(schema: &arrow_schema::Schema) -> usize {
1241 let fields = schema
1242 .fields()
1243 .iter()
1244 .fold(0usize, |bytes, field| bytes.saturating_add(field.size()));
1245 let metadata = schema.metadata().iter().fold(
1246 schema
1247 .metadata()
1248 .capacity()
1249 .saturating_mul(std::mem::size_of::<(String, String)>()),
1250 |bytes, (key, value)| {
1251 bytes
1252 .saturating_add(key.capacity())
1253 .saturating_add(value.capacity())
1254 },
1255 );
1256 std::mem::size_of_val(schema)
1257 .saturating_add(
1258 schema
1259 .fields()
1260 .len()
1261 .saturating_mul(std::mem::size_of::<arrow_schema::FieldRef>()),
1262 )
1263 .saturating_add(fields)
1264 .saturating_add(metadata)
1265 }
1266
1267 fn outbound_workspace_bytes(msg: &ShuffleMessage) -> io::Result<usize> {
1268 match msg {
1269 ShuffleMessage::Barrier(barrier) => {
1270 validate_checkpoint_barrier(*barrier)?;
1271 Ok(1024)
1272 }
1273 ShuffleMessage::Frontier {
1274 stage, watermark, ..
1275 } => {
1276 validate_frontier(stage, *watermark)?;
1277 Ok(FRONTIER_WORKSPACE_BYTES)
1278 }
1279 ShuffleMessage::Data {
1280 stage,
1281 routed_vnodes,
1282 batch,
1283 } => {
1284 if stage.is_empty() || stage.len() > MAX_STAGE_NAME_BYTES {
1285 return Err(io::Error::new(
1286 io::ErrorKind::InvalidInput,
1287 "shuffle stage scope is empty or too long",
1288 ));
1289 }
1290 let batch_bytes =
1291 crate::shuffle::routing::logical_batch_bytes(batch).map_err(|error| {
1292 io::Error::new(io::ErrorKind::InvalidInput, error.to_string())
1293 })?;
1294 if batch_bytes > crate::shuffle::ROUTE_MAX_BATCH_BYTES {
1295 return Err(io::Error::new(
1296 io::ErrorKind::InvalidInput,
1297 format!(
1298 "shuffle batch is {batch_bytes} bytes; limit is {}",
1299 crate::shuffle::ROUTE_MAX_BATCH_BYTES
1300 ),
1301 ));
1302 }
1303 if batch.num_rows() > crate::shuffle::ROUTE_MAX_BATCH_ROWS {
1304 return Err(io::Error::new(
1305 io::ErrorKind::InvalidInput,
1306 "shuffle batch exceeds the row-count bound",
1307 ));
1308 }
1309 let schema_bytes = schema_memory_size(batch.schema().as_ref());
1310 if schema_bytes > MAX_SOURCE_SCHEMA_MEMORY_BYTES {
1311 return Err(io::Error::new(
1312 io::ErrorKind::InvalidInput,
1313 format!(
1314 "shuffle schema is {schema_bytes} bytes; limit is {MAX_SOURCE_SCHEMA_MEMORY_BYTES}"
1315 ),
1316 ));
1317 }
1318 let routes_are_canonical = routed_vnodes.windows(2).all(|pair| pair[0] < pair[1]);
1319 if !routes_are_canonical
1320 || routed_vnodes.is_empty()
1321 || routed_vnodes.len()
1322 > usize::try_from(crate::state::MAX_KEY_GROUP_COUNT).unwrap_or(usize::MAX)
1323 {
1324 return Err(io::Error::new(
1325 io::ErrorKind::InvalidInput,
1326 "shuffle route metadata is empty or non-canonical",
1327 ));
1328 }
1329 Ok(OUTBOUND_DATA_WORKSPACE_BYTES)
1330 }
1331 }
1332 }
1333
1334 fn encode_outbound_data(
1335 stage: String,
1336 routed_vnodes: Vec<u32>,
1337 batch: RecordBatch,
1338 logical_bytes: usize,
1339 schema_bytes: usize,
1340 mut budget: OutboundReservation,
1341 ) -> io::Result<(PreparedMessage, OutboundReservation)> {
1342 let initial_capacity = logical_bytes
1343 .saturating_add(schema_bytes)
1344 .saturating_add(FRAME_METADATA_BYTES)
1345 .min(crate::shuffle::message::MAX_PAYLOAD_BYTES);
1346 let payload = serialize_batch_stream_bounded(
1347 &batch,
1348 crate::shuffle::message::MAX_PAYLOAD_BYTES,
1349 initial_capacity,
1350 )
1351 .map_err(|error| {
1352 io::Error::new(
1353 io::ErrorKind::InvalidData,
1354 format!("shuffle IPC encode: {error}"),
1355 )
1356 })?;
1357 drop(batch);
1360 let payload_len = payload.len();
1361 let payload_capacity = payload.capacity();
1362 if payload_len == 0
1363 || payload_len > crate::shuffle::message::MAX_PAYLOAD_BYTES
1364 || payload_capacity > crate::shuffle::message::MAX_PAYLOAD_BYTES
1365 {
1366 return Err(io::Error::new(
1367 io::ErrorKind::InvalidInput,
1368 format!(
1369 "shuffle IPC payload uses {payload_len} bytes ({payload_capacity} allocated); limit is {}",
1370 crate::shuffle::message::MAX_PAYLOAD_BYTES
1371 ),
1372 ));
1373 }
1374 validate_ipc_schema_header(&payload, payload_len)
1375 .map_err(|error| io::Error::new(io::ErrorKind::InvalidData, error))?;
1376 let route_bytes = routed_vnodes
1377 .capacity()
1378 .checked_mul(std::mem::size_of::<u32>())
1379 .ok_or_else(|| io::Error::other("shuffle route accounting overflow"))?;
1380 let retained_bytes = payload_capacity
1381 .checked_add(stage.capacity())
1382 .and_then(|bytes| bytes.checked_add(route_bytes))
1383 .and_then(|bytes| bytes.checked_add(FRAME_METADATA_BYTES))
1384 .ok_or_else(|| io::Error::other("shuffle outbound accounting overflow"))?;
1385 budget.shrink_to(retained_bytes)?;
1386 Ok((
1387 PreparedMessage::Data {
1388 stage,
1389 routed_vnodes,
1390 arrow_ipc: Bytes::from(payload),
1391 },
1392 budget,
1393 ))
1394 }
1395
1396 async fn prepare_outbound_message_with_hook<F>(
1397 msg: &ShuffleMessage,
1398 budget: OutboundReservation,
1399 before_blocking_encode: F,
1400 ) -> io::Result<(PreparedMessage, OutboundReservation)>
1401 where
1402 F: FnOnce() + Send + 'static,
1403 {
1404 match msg {
1405 ShuffleMessage::Barrier(barrier) => Ok((PreparedMessage::Barrier(*barrier), budget)),
1406 ShuffleMessage::Frontier {
1407 stage,
1408 watermark,
1409 idle,
1410 } => Ok((
1411 PreparedMessage::Frontier {
1412 stage: stage.clone(),
1413 watermark: *watermark,
1414 idle: *idle,
1415 },
1416 budget,
1417 )),
1418 ShuffleMessage::Data {
1419 stage,
1420 routed_vnodes,
1421 batch,
1422 } => {
1423 let stage = stage.clone();
1424 let routed_vnodes = routed_vnodes.to_vec();
1425 let batch = batch.clone();
1426 let logical_bytes =
1427 crate::shuffle::routing::logical_batch_bytes(&batch).map_err(|error| {
1428 io::Error::new(io::ErrorKind::InvalidInput, error.to_string())
1429 })?;
1430 let schema_bytes = schema_memory_size(batch.schema().as_ref());
1431 let offload = logical_bytes >= BLOCKING_IPC_THRESHOLD_BYTES;
1432 if offload {
1433 tokio::task::spawn_blocking(move || {
1434 before_blocking_encode();
1435 encode_outbound_data(
1436 stage,
1437 routed_vnodes,
1438 batch,
1439 logical_bytes,
1440 schema_bytes,
1441 budget,
1442 )
1443 })
1444 .await
1445 .map_err(|error| io::Error::other(format!("shuffle encoder task: {error}")))?
1446 } else {
1447 encode_outbound_data(
1448 stage,
1449 routed_vnodes,
1450 batch,
1451 logical_bytes,
1452 schema_bytes,
1453 budget,
1454 )
1455 }
1456 }
1457 }
1458 }
1459
1460 async fn prepare_outbound_message(
1461 msg: &ShuffleMessage,
1462 budget: OutboundReservation,
1463 ) -> io::Result<(PreparedMessage, OutboundReservation)> {
1464 prepare_outbound_message_with_hook(msg, budget, || {}).await
1465 }
1466
1467 #[derive(Debug, Clone, Copy, PartialEq, Eq)]
1470 struct StreamFence {
1471 sender_node_id: ShufflePeerId,
1472 sender_incarnation: Uuid,
1473 receiver_incarnation: Uuid,
1474 stream_id: Uuid,
1475 assignment_version: u64,
1476 assignment_certificate_digest: [u8; 32],
1477 recovery_gen: u64,
1478 }
1479
1480 #[derive(Debug)]
1483 struct InstalledAssignment {
1484 fence: CheckpointAssignmentFence,
1485 digest: [u8; 32],
1486 owners: Arc<[ShufflePeerId]>,
1487 }
1488
1489 #[derive(Clone)]
1490 struct ScopeLease {
1491 assignment: Arc<InstalledAssignment>,
1492 recovery_gen: u64,
1493 cancel: CancellationToken,
1494 }
1495
1496 impl ScopeLease {
1497 fn matches_fence(&self, fence: &StreamFence) -> bool {
1498 !self.cancel.is_cancelled()
1499 && self.assignment.fence.assignment_version == fence.assignment_version
1500 && self.assignment.digest == fence.assignment_certificate_digest
1501 && self.recovery_gen == fence.recovery_gen
1502 }
1503 }
1504
1505 impl InstalledAssignment {
1506 fn for_process(
1507 fence: &CheckpointAssignmentFence,
1508 owners: &[ShufflePeerId],
1509 local_id: ShufflePeerId,
1510 local_incarnation: Uuid,
1511 ) -> io::Result<Arc<Self>> {
1512 if !fence.is_canonical()
1513 || !fence.matches_owner_map(owners)
1514 || owners.iter().any(|owner| !fence.contains(*owner))
1515 || fence.participant_incarnation(local_id) != Some(local_incarnation)
1516 {
1517 return Err(io::Error::new(
1518 io::ErrorKind::InvalidInput,
1519 "shuffle assignment scope does not bind this process and exact owner map",
1520 ));
1521 }
1522 Ok(Arc::new(Self {
1523 digest: fence.digest(),
1524 fence: fence.clone(),
1525 owners: Arc::from(owners),
1526 }))
1527 }
1528
1529 fn certifies(&self, node_id: ShufflePeerId, incarnation: Uuid) -> bool {
1530 self.fence.participant_incarnation(node_id) == Some(incarnation)
1531 }
1532
1533 fn matches_stream_sender(&self, fence: &StreamFence) -> bool {
1534 self.fence.assignment_version == fence.assignment_version
1535 && self.digest == fence.assignment_certificate_digest
1536 && self.certifies(fence.sender_node_id, fence.sender_incarnation)
1537 }
1538
1539 fn owns_vnode(&self, node_id: ShufflePeerId, vnode: u32) -> bool {
1540 usize::try_from(vnode)
1541 .ok()
1542 .and_then(|vnode| self.owners.get(vnode))
1543 .is_some_and(|owner| *owner == node_id)
1544 }
1545 }
1546
1547 struct PendingHandshake {
1548 fence: StreamFence,
1549 issued_at: std::time::Instant,
1550 }
1551
1552 #[derive(Default)]
1553 struct PendingHandshakes(Mutex<FxHashMap<ShufflePeerId, PendingHandshake>>);
1554
1555 impl PendingHandshakes {
1556 fn clear(&self) {
1557 self.0.lock().clear();
1558 }
1559 }
1560
1561 struct ActiveStreamEntry {
1562 owner: Arc<()>,
1563 cancel: CancellationToken,
1564 }
1565
1566 #[derive(Default)]
1567 struct ActiveStreamRegistry {
1568 streams: Mutex<FxHashMap<ShufflePeerId, ActiveStreamEntry>>,
1569 }
1570
1571 impl ActiveStreamRegistry {
1572 fn replace(
1573 self: &Arc<Self>,
1574 fence: &StreamFence,
1575 parent_cancel: &CancellationToken,
1576 ) -> ActiveStreamLease {
1577 let key = fence.sender_node_id;
1578 let cancel = parent_cancel.child_token();
1579 let owner = Arc::new(());
1580 let previous = self.streams.lock().insert(
1581 key,
1582 ActiveStreamEntry {
1583 owner: Arc::clone(&owner),
1584 cancel: cancel.clone(),
1585 },
1586 );
1587 if let Some(previous) = previous {
1588 previous.cancel.cancel();
1589 }
1590 ActiveStreamLease {
1591 registry: Arc::clone(self),
1592 key,
1593 owner,
1594 cancel,
1595 permit: None,
1596 }
1597 }
1598 }
1599
1600 struct ActiveStreamLease {
1601 registry: Arc<ActiveStreamRegistry>,
1602 key: ShufflePeerId,
1603 owner: Arc<()>,
1604 cancel: CancellationToken,
1605 permit: Option<OwnedSemaphorePermit>,
1606 }
1607
1608 impl ActiveStreamLease {
1609 async fn acquire_permit(&mut self, permits: &Arc<Semaphore>) -> Result<(), tonic::Status> {
1610 let permit = tokio::select! {
1611 biased;
1612 () = self.cancel.cancelled() => return Err(tonic::Status::cancelled(
1613 "shuffle stream was superseded",
1614 )),
1615 permit = Arc::clone(permits).acquire_owned() => permit.map_err(|_| {
1616 tonic::Status::unavailable("shuffle stream admission closed")
1617 })?,
1618 };
1619 self.permit = Some(permit);
1620 Ok(())
1621 }
1622 }
1623
1624 impl Drop for ActiveStreamLease {
1625 fn drop(&mut self) {
1626 self.cancel.cancel();
1627 self.permit.take();
1628 let mut streams = self.registry.streams.lock();
1629 if streams
1630 .get(&self.key)
1631 .is_some_and(|entry| Arc::ptr_eq(&entry.owner, &self.owner))
1632 {
1633 streams.remove(&self.key);
1634 }
1635 }
1636 }
1637
1638 struct FragmentAssembly {
1639 stage: String,
1640 routed_vnodes: Vec<u32>,
1641 recovery_gen: u64,
1642 seq: u64,
1643 fragment_count: u32,
1644 total_payload_bytes: usize,
1645 next_fragment: u32,
1646 payload: Vec<u8>,
1647 budget: InboundReservation,
1648 }
1649
1650 struct CompleteData {
1651 stage: String,
1652 routed_vnodes: Vec<u32>,
1653 seq: u64,
1654 arrow_ipc: Vec<u8>,
1655 budget: InboundReservation,
1656 }
1657
1658 fn validate_ipc_schema_header(
1659 payload: &[u8],
1660 total_payload_bytes: usize,
1661 ) -> Result<(), String> {
1662 const CONTINUATION_MARKER: [u8; 4] = [0xff; 4];
1663
1664 if payload.len() < 8 || payload[..4] != CONTINUATION_MARKER {
1665 return Err("shuffle IPC must start with a modern Arrow schema message".into());
1666 }
1667 let metadata_len = usize::try_from(u32::from_le_bytes(
1668 payload[4..8]
1669 .try_into()
1670 .expect("validated Arrow IPC prefix length"),
1671 ))
1672 .map_err(|_| "shuffle IPC schema length is invalid".to_string())?;
1673 if metadata_len == 0 || metadata_len > MAX_SCHEMA_WIRE_BYTES {
1674 return Err("shuffle IPC schema message exceeds its wire bound".into());
1675 }
1676 let end = 8usize
1677 .checked_add(metadata_len)
1678 .ok_or_else(|| "shuffle IPC schema length overflow".to_string())?;
1679 if end > payload.len() || end > total_payload_bytes {
1680 return Err("shuffle IPC schema message is not complete in fragment zero".into());
1681 }
1682 let message = arrow_ipc::root_as_message(&payload[8..end])
1683 .map_err(|error| format!("invalid shuffle IPC schema message: {error}"))?;
1684 if message.header_type() != arrow_ipc::MessageHeader::Schema
1685 || message.header_as_schema().is_none()
1686 || message.bodyLength() != 0
1687 {
1688 return Err("shuffle IPC leading message is not a bodyless schema".into());
1689 }
1690 Ok(())
1691 }
1692
1693 fn validate_fragment(fragment: &RoutedData) -> Result<usize, String> {
1694 let fragment_count = usize::try_from(fragment.fragment_count)
1695 .map_err(|_| "invalid shuffle fragment count".to_string())?;
1696 let total_payload_bytes = usize::try_from(fragment.total_payload_bytes)
1697 .map_err(|_| "invalid shuffle payload size".to_string())?;
1698 if fragment_count == 0
1699 || fragment_count > MAX_FRAGMENTS
1700 || total_payload_bytes == 0
1701 || total_payload_bytes > crate::shuffle::message::MAX_PAYLOAD_BYTES
1702 || fragment.arrow_ipc.is_empty()
1703 || fragment.arrow_ipc.len() > MAX_WIRE_PAYLOAD_BYTES
1704 || fragment_count != total_payload_bytes.div_ceil(MAX_WIRE_PAYLOAD_BYTES)
1705 || usize::try_from(fragment.fragment_index).unwrap_or(usize::MAX) >= fragment_count
1706 {
1707 return Err("invalid shuffle fragment bounds".into());
1708 }
1709 let index = usize::try_from(fragment.fragment_index)
1710 .map_err(|_| "invalid shuffle fragment index".to_string())?;
1711 let expected_len = if index + 1 == fragment_count {
1712 total_payload_bytes - index * MAX_WIRE_PAYLOAD_BYTES
1713 } else {
1714 MAX_WIRE_PAYLOAD_BYTES
1715 };
1716 if fragment.arrow_ipc.len() != expected_len {
1717 return Err("shuffle fragment length does not match its declared payload".into());
1718 }
1719 if index == 0 {
1720 validate_ipc_schema_header(&fragment.arrow_ipc, total_payload_bytes)?;
1721 let routes_are_canonical = fragment
1722 .routed_vnodes
1723 .windows(2)
1724 .all(|pair| pair[0] < pair[1]);
1725 if fragment.stage.is_empty()
1726 || fragment.stage.len() > MAX_STAGE_NAME_BYTES
1727 || !routes_are_canonical
1728 || fragment.routed_vnodes.len()
1729 > usize::try_from(crate::state::MAX_KEY_GROUP_COUNT).unwrap_or(usize::MAX)
1730 || fragment.routed_vnodes.is_empty()
1731 {
1732 return Err("shuffle fragment-zero metadata is empty or non-canonical".into());
1733 }
1734 } else if !fragment.stage.is_empty() || !fragment.routed_vnodes.is_empty() {
1735 return Err("shuffle continuation fragment repeated logical metadata".into());
1736 }
1737 Ok(total_payload_bytes)
1738 }
1739
1740 fn retained_batch_metadata_bytes(
1741 stage: &String,
1742 routed_vnodes: &[u32],
1743 batch: &RecordBatch,
1744 ) -> Result<usize, tonic::Status> {
1745 let schema_bytes = schema_memory_size(batch.schema().as_ref());
1746 if schema_bytes > MAX_DECODED_SCHEMA_MEMORY_BYTES {
1747 return Err(tonic::Status::resource_exhausted(
1748 "decoded shuffle schema exceeds its memory bound",
1749 ));
1750 }
1751 let structure_bytes = batch.columns().iter().try_fold(
1752 std::mem::size_of::<RecordBatch>()
1753 + batch
1754 .num_columns()
1755 .saturating_mul(std::mem::size_of::<arrow_array::ArrayRef>()),
1756 |total, column| {
1757 let array_bytes = column
1758 .get_array_memory_size()
1759 .saturating_sub(column.get_buffer_memory_size());
1760 total.checked_add(array_bytes).ok_or_else(|| {
1761 tonic::Status::internal("shuffle array structure accounting overflow")
1762 })
1763 },
1764 )?;
1765 if structure_bytes > MAX_DECODED_ARRAY_STRUCTURE_BYTES {
1766 return Err(tonic::Status::resource_exhausted(
1767 "decoded shuffle array structure exceeds its memory bound",
1768 ));
1769 }
1770 let route_bytes = routed_vnodes
1771 .len()
1772 .checked_mul(std::mem::size_of::<u32>())
1773 .ok_or_else(|| tonic::Status::internal("shuffle route accounting overflow"))?;
1774 let metadata_bytes = schema_bytes
1775 .checked_add(structure_bytes)
1776 .and_then(|bytes| bytes.checked_add(stage.capacity()))
1777 .and_then(|bytes| bytes.checked_add(route_bytes))
1778 .and_then(|bytes| bytes.checked_add(RETAINED_BATCH_ENVELOPE_BYTES))
1779 .ok_or_else(|| tonic::Status::internal("shuffle metadata accounting overflow"))?;
1780 if metadata_bytes > INBOUND_BATCH_METADATA_BYTES {
1781 return Err(tonic::Status::resource_exhausted(
1782 "decoded shuffle metadata exceeds its admission bound",
1783 ));
1784 }
1785 Ok(metadata_bytes)
1786 }
1787
1788 fn push_fragment(
1789 assembly: &mut Option<FragmentAssembly>,
1790 fragment: &RoutedData,
1791 budget: Option<InboundReservation>,
1792 ) -> Result<Option<CompleteData>, String> {
1793 let total_payload_bytes = validate_fragment(fragment)?;
1794
1795 if fragment.fragment_index == 0 {
1796 if assembly.is_some() {
1797 return Err("shuffle fragments interleaved across logical frames".into());
1798 }
1799 let budget = budget.ok_or_else(|| {
1800 "shuffle fragment zero was not admitted by byte budget".to_string()
1801 })?;
1802 *assembly = Some(FragmentAssembly {
1803 stage: fragment.stage.clone(),
1804 routed_vnodes: fragment.routed_vnodes.clone(),
1805 recovery_gen: fragment.recovery_gen,
1806 seq: fragment.seq,
1807 fragment_count: fragment.fragment_count,
1808 total_payload_bytes,
1809 next_fragment: 0,
1810 payload: Vec::with_capacity(total_payload_bytes),
1811 budget,
1812 });
1813 } else if budget.is_some() {
1814 return Err("shuffle continuation fragment carried a new byte reservation".into());
1815 }
1816 let current = assembly
1817 .as_mut()
1818 .ok_or_else(|| "shuffle fragment arrived without fragment zero".to_string())?;
1819 let first_metadata_changed = fragment.fragment_index == 0
1820 && (current.stage != fragment.stage || current.routed_vnodes != fragment.routed_vnodes);
1821 if first_metadata_changed
1822 || current.recovery_gen != fragment.recovery_gen
1823 || current.seq != fragment.seq
1824 || current.fragment_count != fragment.fragment_count
1825 || current.total_payload_bytes != total_payload_bytes
1826 || current.next_fragment != fragment.fragment_index
1827 {
1828 return Err("shuffle fragment metadata or order changed mid-frame".into());
1829 }
1830 current.payload.extend_from_slice(&fragment.arrow_ipc);
1831 current.next_fragment += 1;
1832 if current.next_fragment != current.fragment_count {
1833 return Ok(None);
1834 }
1835 let complete = assembly.take().expect("completed fragment assembly exists");
1836 if complete.payload.len() != complete.total_payload_bytes {
1837 return Err("reassembled shuffle payload length mismatch".into());
1838 }
1839 Ok(Some(CompleteData {
1840 stage: complete.stage,
1841 routed_vnodes: complete.routed_vnodes,
1842 seq: complete.seq,
1843 arrow_ipc: complete.payload,
1844 budget: complete.budget,
1845 }))
1846 }
1847
1848 fn parse_uuid(raw: &[u8], field: &str) -> Result<Uuid, tonic::Status> {
1849 let value = Uuid::from_slice(raw)
1850 .map_err(|_| tonic::Status::invalid_argument(format!("invalid {field} UUID")))?;
1851 if value.is_nil() {
1852 return Err(tonic::Status::invalid_argument(format!("nil {field} UUID")));
1853 }
1854 Ok(value)
1855 }
1856
1857 fn parse_certificate_digest(raw: &[u8], field: &str) -> Result<[u8; 32], tonic::Status> {
1858 let digest: [u8; 32] = raw
1859 .try_into()
1860 .map_err(|_| tonic::Status::invalid_argument(format!("invalid {field}")))?;
1861 if digest == [0; 32] {
1862 return Err(tonic::Status::invalid_argument(format!("zero {field}")));
1863 }
1864 Ok(digest)
1865 }
1866
1867 fn hello_for(node_id: ShufflePeerId, fence: &StreamFence) -> Hello {
1868 Hello {
1869 node_id,
1870 sender_incarnation: fence.sender_incarnation.as_bytes().to_vec(),
1871 receiver_incarnation: fence.receiver_incarnation.as_bytes().to_vec(),
1872 stream_id: fence.stream_id.as_bytes().to_vec(),
1873 assignment_version: fence.assignment_version,
1874 recovery_gen: fence.recovery_gen,
1875 assignment_certificate_digest: fence.assignment_certificate_digest.to_vec(),
1876 }
1877 }
1878
1879 fn fence_from_hello(hello: &Hello) -> Result<StreamFence, tonic::Status> {
1880 if hello.node_id == 0 || hello.assignment_version == 0 {
1881 return Err(tonic::Status::failed_precondition(
1882 "shuffle stream requires assigned nodes and a nonzero assignment version",
1883 ));
1884 }
1885 Ok(StreamFence {
1886 sender_node_id: hello.node_id,
1887 sender_incarnation: parse_uuid(&hello.sender_incarnation, "sender incarnation")?,
1888 receiver_incarnation: parse_uuid(&hello.receiver_incarnation, "receiver incarnation")?,
1889 stream_id: parse_uuid(&hello.stream_id, "stream id")?,
1890 assignment_version: hello.assignment_version,
1891 assignment_certificate_digest: parse_certificate_digest(
1892 &hello.assignment_certificate_digest,
1893 "assignment certificate digest",
1894 )?,
1895 recovery_gen: hello.recovery_gen,
1896 })
1897 }
1898
1899 fn frame_message(out: Outbound) -> Result<Encoded, tonic::Status> {
1901 let Outbound {
1902 gen,
1903 assignment_version,
1904 seq,
1905 msg,
1906 assignment_digest,
1907 _budget: budget,
1908 } = out;
1909 let frames = match msg {
1910 PreparedMessage::Barrier(b) => {
1911 validate_checkpoint_barrier(b)
1912 .map_err(|error| tonic::Status::invalid_argument(error.to_string()))?;
1913 let assignment_digest = assignment_digest.ok_or_else(|| {
1914 tonic::Status::failed_precondition(
1915 "shuffle checkpoint barrier has no assignment certificate",
1916 )
1917 })?;
1918 VecDeque::from([ShuffleFrame {
1919 kind: Some(shuffle_frame::Kind::Barrier(Barrier {
1920 checkpoint_id: b.checkpoint_id,
1921 epoch: b.epoch,
1922 flags: b.flags,
1923 last_seq: seq,
1924 assignment_version,
1925 assignment_digest: assignment_digest.to_vec(),
1926 recovery_gen: gen,
1927 })),
1928 }])
1929 }
1930 PreparedMessage::Data {
1931 mut stage,
1932 routed_vnodes,
1933 arrow_ipc,
1934 } => {
1935 let total = arrow_ipc.len();
1936 if total == 0 || total > crate::shuffle::message::MAX_PAYLOAD_BYTES {
1937 return Err(tonic::Status::resource_exhausted(format!(
1938 "shuffle IPC payload is {total} bytes; limit is {}",
1939 crate::shuffle::message::MAX_PAYLOAD_BYTES
1940 )));
1941 }
1942 let fragment_count = total.div_ceil(MAX_WIRE_PAYLOAD_BYTES);
1943 if fragment_count == 0 || fragment_count > MAX_FRAGMENTS {
1944 return Err(tonic::Status::resource_exhausted(
1945 "shuffle IPC payload needs too many fragments",
1946 ));
1947 }
1948 let fragment_count = u32::try_from(fragment_count)
1949 .map_err(|_| tonic::Status::resource_exhausted("too many fragments"))?;
1950 let total_payload_bytes = u32::try_from(total)
1951 .map_err(|_| tonic::Status::resource_exhausted("payload is too large"))?;
1952 let mut routes = routed_vnodes;
1953 let mut frames = VecDeque::with_capacity(fragment_count as usize);
1954 for index in 0..fragment_count {
1955 let index = index as usize;
1956 let start = index * MAX_WIRE_PAYLOAD_BYTES;
1957 let end = (start + MAX_WIRE_PAYLOAD_BYTES).min(total);
1958 let first = index == 0;
1959 frames.push_back(ShuffleFrame {
1960 kind: Some(shuffle_frame::Kind::Data(RoutedData {
1961 stage: if first {
1962 std::mem::take(&mut stage)
1963 } else {
1964 String::new()
1965 },
1966 routed_vnodes: if first {
1967 std::mem::take(&mut routes)
1968 } else {
1969 Vec::new()
1970 },
1971 arrow_ipc: arrow_ipc.slice(start..end),
1972 recovery_gen: gen,
1973 seq,
1974 fragment_index: u32::try_from(index)
1975 .expect("fragment count is bounded by MAX_FRAGMENTS"),
1976 fragment_count,
1977 total_payload_bytes,
1978 })),
1979 });
1980 }
1981 frames
1982 }
1983 PreparedMessage::Frontier {
1984 stage,
1985 watermark,
1986 idle,
1987 } => {
1988 validate_frontier(&stage, watermark)
1989 .map_err(|error| tonic::Status::invalid_argument(error.to_string()))?;
1990 VecDeque::from([ShuffleFrame {
1991 kind: Some(shuffle_frame::Kind::Frontier(WireFrontier {
1992 stage,
1993 watermark,
1994 idle,
1995 recovery_gen: gen,
1996 seq,
1997 })),
1998 }])
1999 }
2000 };
2001 Ok(Encoded {
2002 frames,
2003 _budget: budget,
2004 })
2005 }
2006
2007 struct PeerConn {
2012 tx: MAsyncTx<mpsc::Array<Outbound>>,
2013 byte_budget: Arc<Semaphore>,
2014 control_byte_budget: Arc<Semaphore>,
2015 send_lock: tokio::sync::Mutex<()>,
2018 alive: Arc<AtomicBool>,
2019 driver: JoinHandle<()>,
2020 fence: StreamFence,
2021 }
2022
2023 impl PeerConn {
2024 fn is_alive(&self) -> bool {
2025 self.alive.load(Ordering::Acquire) && !self.driver.is_finished()
2029 }
2030 }
2031
2032 impl Drop for PeerConn {
2033 fn drop(&mut self) {
2034 self.driver.abort();
2035 }
2036 }
2037
2038 struct OpenCall {
2039 local_id: ShufflePeerId,
2040 peer: ShufflePeerId,
2041 addr: SocketAddr,
2042 sender_incarnation: Uuid,
2043 assignment_version: u64,
2044 assignment_certificate_digest: [u8; 32],
2045 expected_receiver_incarnation: Uuid,
2046 recovery_gen: u64,
2047 current_assignment: Arc<AtomicU64>,
2048 current_recovery_gen: Arc<AtomicU64>,
2049 scope_cancel: CancellationToken,
2050 }
2051
2052 type ConnectLock = Arc<tokio::sync::Mutex<()>>;
2053
2054 pub struct ShuffleSender {
2056 local_id: ShufflePeerId,
2057 sender_incarnation: Uuid,
2058 peers: Mutex<FxHashMap<ShufflePeerId, SocketAddr>>,
2059 pool: Mutex<FxHashMap<ShufflePeerId, Arc<PeerConn>>>,
2060 connect_locks: Mutex<FxHashMap<ShufflePeerId, ConnectLock>>,
2062 kv: Option<Arc<dyn ClusterKv>>,
2063 assignment: Arc<RwLock<Option<Arc<InstalledAssignment>>>>,
2064 assignment_version: Arc<AtomicU64>,
2065 scope_cancel: Arc<RwLock<CancellationToken>>,
2066 process_lease: Arc<ProcessLeaseGate>,
2067 assignment_suspended: AtomicBool,
2070 recovery_gen: Arc<AtomicU64>,
2072 seqs: Mutex<FxHashMap<ShufflePeerId, u64>>,
2075 checkpointed_control_node_budget: Arc<Semaphore>,
2076 node_budget: Arc<Semaphore>,
2077 #[cfg(test)]
2078 post_enqueue_pause: Mutex<Option<(Arc<tokio::sync::Notify>, Arc<tokio::sync::Notify>)>>,
2079 }
2080
2081 impl std::fmt::Debug for ShuffleSender {
2082 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
2083 f.debug_struct("ShuffleSender")
2084 .field("local_id", &self.local_id)
2085 .finish_non_exhaustive()
2086 }
2087 }
2088
2089 impl ShuffleSender {
2090 #[must_use]
2096 pub fn new(local_id: ShufflePeerId, incarnation: Uuid) -> Self {
2097 assert!(local_id != 0, "shuffle sender node id must be nonzero");
2098 assert!(
2099 !incarnation.is_nil(),
2100 "shuffle sender incarnation must be non-nil"
2101 );
2102 Self {
2103 local_id,
2104 sender_incarnation: incarnation,
2105 peers: Mutex::new(FxHashMap::default()),
2106 pool: Mutex::new(FxHashMap::default()),
2107 connect_locks: Mutex::new(FxHashMap::default()),
2108 kv: None,
2109 assignment: Arc::new(RwLock::new(None)),
2110 assignment_version: Arc::new(AtomicU64::new(0)),
2111 scope_cancel: Arc::new(RwLock::new(cancelled_token())),
2112 process_lease: Arc::new(ProcessLeaseGate::default()),
2113 assignment_suspended: AtomicBool::new(false),
2114 recovery_gen: Arc::new(AtomicU64::new(0)),
2115 seqs: Mutex::new(FxHashMap::default()),
2116 checkpointed_control_node_budget: Arc::new(Semaphore::new(
2117 CHECKPOINTED_CONTROL_NODE_BUDGET_BYTES,
2118 )),
2119 node_budget: Arc::new(Semaphore::new(OUTBOUND_NODE_BUDGET_BYTES)),
2120 #[cfg(test)]
2121 post_enqueue_pause: Mutex::new(None),
2122 }
2123 }
2124
2125 #[cfg(test)]
2126 pub(crate) fn pause_next_send_after_enqueue_for_test(
2127 &self,
2128 ) -> (Arc<tokio::sync::Notify>, Arc<tokio::sync::Notify>) {
2129 let entered = Arc::new(tokio::sync::Notify::new());
2130 let release = Arc::new(tokio::sync::Notify::new());
2131 let replaced = self
2132 .post_enqueue_pause
2133 .lock()
2134 .replace((Arc::clone(&entered), Arc::clone(&release)));
2135 assert!(
2136 replaced.is_none(),
2137 "post-enqueue send pause already installed"
2138 );
2139 (entered, release)
2140 }
2141
2142 #[must_use]
2144 pub const fn local_id(&self) -> ShufflePeerId {
2145 self.local_id
2146 }
2147
2148 pub fn install_process_lease_deadline(
2156 &self,
2157 deadline: Arc<LeaseDeadline>,
2158 ) -> io::Result<()> {
2159 let _assignment = self.assignment.write();
2160 if self.assignment_version.load(Ordering::Acquire) != 0
2161 && !self.process_lease.is_installed_deadline(&deadline)
2162 {
2163 return Err(io::Error::new(
2164 io::ErrorKind::InvalidInput,
2165 "shuffle process lease must be installed before assignment activation",
2166 ));
2167 }
2168 self.process_lease.install(deadline)
2169 }
2170
2171 pub fn bind_process_lease_deadline_pair(
2181 &self,
2182 receiver: &ShuffleReceiver,
2183 deadline: Arc<LeaseDeadline>,
2184 ) -> io::Result<()> {
2185 let _sender_assignment = self.assignment.write();
2186 let _receiver_assignment = receiver.assignment.write();
2187 if self.assignment_version.load(Ordering::Acquire) != 0
2188 && !self.process_lease.is_installed_deadline(&deadline)
2189 {
2190 return Err(io::Error::new(
2191 io::ErrorKind::InvalidInput,
2192 "shuffle process lease must be installed before outbound assignment activation",
2193 ));
2194 }
2195 if receiver.assignment_version.load(Ordering::Acquire) != 0
2196 && !receiver.process_lease.is_installed_deadline(&deadline)
2197 {
2198 return Err(io::Error::new(
2199 io::ErrorKind::InvalidInput,
2200 "shuffle process lease must be installed before inbound assignment activation",
2201 ));
2202 }
2203 ProcessLeaseGate::install_pair(
2204 self.process_lease.as_ref(),
2205 receiver.process_lease.as_ref(),
2206 deadline,
2207 )
2208 }
2209
2210 #[cfg(test)]
2211 pub(crate) fn install_live_process_lease_for_test(&self) {
2212 self.process_lease.install_live_for_test();
2213 }
2214
2215 pub fn install_assignment_fence(
2222 &self,
2223 fence: &CheckpointAssignmentFence,
2224 owners: &[ShufflePeerId],
2225 ) -> io::Result<bool> {
2226 let next = InstalledAssignment::for_process(
2227 fence,
2228 owners,
2229 self.local_id,
2230 self.sender_incarnation,
2231 )?;
2232 let mut assignment = self.assignment.write();
2235 self.process_lease.require_live_io()?;
2236 if let Some(current) = assignment.as_ref() {
2237 if next.fence.assignment_version < current.fence.assignment_version {
2238 return Ok(false);
2239 }
2240 if next.fence.assignment_version == current.fence.assignment_version {
2241 if next.digest == current.digest
2242 && next.fence == current.fence
2243 && next.owners == current.owners
2244 {
2245 if self.assignment_version.load(Ordering::Acquire)
2246 == next.fence.assignment_version
2247 {
2248 return Ok(false);
2249 }
2250 if self.assignment_suspended.load(Ordering::Acquire) {
2251 rotate_scope_token(&self.scope_cancel, &self.process_lease, true);
2252 self.pool.lock().clear();
2253 self.connect_locks.lock().clear();
2254 self.assignment_suspended.store(false, Ordering::Release);
2255 self.assignment_version
2256 .store(next.fence.assignment_version, Ordering::Release);
2257 return Ok(true);
2258 }
2259 return Err(io::Error::new(
2260 io::ErrorKind::InvalidData,
2261 "an invalidated shuffle assignment requires a higher version",
2262 ));
2263 }
2264 return Err(io::Error::new(
2265 io::ErrorKind::InvalidData,
2266 "conflicting shuffle assignment certificate for an installed version",
2267 ));
2268 }
2269 }
2270 rotate_scope_token(&self.scope_cancel, &self.process_lease, false);
2271 let mut pool = self.pool.lock();
2272 let mut seqs = self.seqs.lock();
2273 pool.clear();
2274 seqs.clear();
2275 self.connect_locks.lock().clear();
2276 self.peers
2277 .lock()
2278 .retain(|peer, _| next.fence.contains(*peer));
2279 let version = next.fence.assignment_version;
2280 *assignment = Some(next);
2281 self.assignment_suspended.store(false, Ordering::Release);
2282 rotate_scope_token(&self.scope_cancel, &self.process_lease, true);
2283 self.assignment_version.store(version, Ordering::Release);
2284 Ok(true)
2285 }
2286
2287 pub fn suspend_assignment_fence(&self) {
2291 let assignment = self.assignment.write();
2292 if assignment.is_none() || self.assignment_version.load(Ordering::Acquire) == 0 {
2293 return;
2294 }
2295 rotate_scope_token(&self.scope_cancel, &self.process_lease, false);
2296 self.pool.lock().clear();
2297 self.connect_locks.lock().clear();
2298 self.assignment_suspended.store(true, Ordering::Release);
2299 self.assignment_version.store(0, Ordering::Release);
2300 }
2301
2302 pub fn invalidate_assignment_fence(&self) {
2306 let _assignment = self.assignment.write();
2307 rotate_scope_token(&self.scope_cancel, &self.process_lease, false);
2308 let mut pool = self.pool.lock();
2309 let mut seqs = self.seqs.lock();
2310 self.assignment_suspended.store(false, Ordering::Release);
2311 self.assignment_version.store(0, Ordering::Release);
2312 pool.clear();
2313 seqs.clear();
2314 self.connect_locks.lock().clear();
2315 }
2316
2317 #[must_use]
2319 pub fn assignment_version(&self) -> u64 {
2320 self.assignment_version.load(Ordering::Acquire)
2321 }
2322
2323 #[must_use]
2325 pub fn active_assignment_digest(&self) -> Option<[u8; 32]> {
2326 let assignment = self.assignment.read();
2327 assignment.as_ref().and_then(|installed| {
2328 (self.assignment_version.load(Ordering::Acquire)
2329 == installed.fence.assignment_version)
2330 .then_some(installed.digest)
2331 })
2332 }
2333
2334 pub fn set_recovery_gen(&self, gen: u64) {
2337 let _assignment = self.assignment.write();
2340 let previous = self.recovery_gen.load(Ordering::Acquire);
2341 if gen <= previous {
2342 return;
2343 }
2344 rotate_scope_token(&self.scope_cancel, &self.process_lease, false);
2345 let mut pool = self.pool.lock();
2346 let mut seqs = self.seqs.lock();
2347 self.recovery_gen.store(gen, Ordering::Release);
2348 pool.clear();
2349 seqs.clear();
2350 self.connect_locks.lock().clear();
2351 rotate_scope_token(
2352 &self.scope_cancel,
2353 &self.process_lease,
2354 self.assignment_version.load(Ordering::Acquire) != 0,
2355 );
2356 }
2357
2358 #[must_use]
2360 pub fn recovery_gen(&self) -> u64 {
2361 self.recovery_gen.load(Ordering::Acquire)
2362 }
2363
2364 #[must_use]
2366 pub const fn incarnation(&self) -> Uuid {
2367 self.sender_incarnation
2368 }
2369
2370 #[cfg(test)]
2372 pub fn burn_seq_for_test(&self, peer: ShufflePeerId) {
2373 *self.seqs.lock().entry(peer).or_insert(0) += 1;
2374 }
2375
2376 #[cfg(test)]
2377 pub(crate) fn tracked_resources_for_test(&self) -> (usize, usize, usize, usize) {
2378 (
2379 self.peers.lock().len(),
2380 self.pool.lock().len(),
2381 self.connect_locks.lock().len(),
2382 self.seqs.lock().len(),
2383 )
2384 }
2385
2386 #[must_use]
2389 pub fn with_kv(local_id: ShufflePeerId, kv: Arc<dyn ClusterKv>, incarnation: Uuid) -> Self {
2390 let mut s = Self::new(local_id, incarnation);
2391 s.kv = Some(kv);
2392 s
2393 }
2394
2395 pub fn register_peer(&self, peer: ShufflePeerId, addr: SocketAddr) {
2397 if peer == 0 || peer == self.local_id {
2398 return;
2399 }
2400 let mut peers = self.peers.lock();
2401 if !peers.contains_key(&peer) && peers.len() >= MAX_TRACKED_PEERS {
2402 tracing::warn!(peer, "shuffle peer address registry is full");
2403 return;
2404 }
2405 peers.insert(peer, addr);
2406 }
2407
2408 pub async fn send_to(&self, peer: ShufflePeerId, msg: &ShuffleMessage) -> io::Result<()> {
2414 self.send_to_inner(peer, msg, None, None).await
2415 }
2416
2417 pub async fn send_to_for_assignment(
2426 &self,
2427 peer: ShufflePeerId,
2428 expected_assignment_version: u64,
2429 msg: &ShuffleMessage,
2430 ) -> io::Result<()> {
2431 self.send_to_inner(peer, msg, Some(expected_assignment_version), None)
2432 .await
2433 }
2434
2435 async fn send_to_inner(
2436 &self,
2437 peer: ShufflePeerId,
2438 msg: &ShuffleMessage,
2439 expected_assignment_version: Option<u64>,
2440 assignment_fence: Option<&CheckpointAssignmentFence>,
2441 ) -> io::Result<()> {
2442 if peer == 0 || peer == self.local_id {
2443 return Err(io::Error::new(
2444 io::ErrorKind::InvalidInput,
2445 "shuffle peer must be a different assigned node",
2446 ));
2447 }
2448 let admission_bytes = outbound_workspace_bytes(msg)?;
2449 if matches!(msg, ShuffleMessage::Barrier(_)) && assignment_fence.is_none() {
2450 return Err(io::Error::new(
2451 io::ErrorKind::InvalidInput,
2452 "shuffle checkpoint barriers require an admitted assignment certificate",
2453 ));
2454 }
2455 let scope = self.current_scope(expected_assignment_version)?;
2456 let conn = self.connection_for(peer, &scope).await?;
2457 let _send_guard = tokio::select! {
2458 biased;
2459 () = scope.cancel.cancelled() => return Err(scope_cancelled_io()),
2460 guard = conn.send_lock.lock() => guard,
2461 };
2462 self.validate_scope(&scope, expected_assignment_version)?;
2463 let assignment = &scope.assignment;
2464 if !assignment.matches_stream_sender(&conn.fence)
2465 || !assignment.certifies(peer, conn.fence.receiver_incarnation)
2466 {
2467 return Err(io::Error::new(
2468 io::ErrorKind::ConnectionAborted,
2469 "shuffle stream no longer matches the installed assignment certificate",
2470 ));
2471 }
2472 if let Some(fence) = assignment_fence {
2473 if *fence != assignment.fence
2474 || fence.digest() != conn.fence.assignment_certificate_digest
2475 || fence.participant_incarnation(self.local_id) != Some(self.sender_incarnation)
2476 || fence.participant_incarnation(peer) != Some(conn.fence.receiver_incarnation)
2477 {
2478 return Err(io::Error::new(
2479 io::ErrorKind::ConnectionAborted,
2480 "shuffle barrier stream incarnations differ from its assignment certificate",
2481 ));
2482 }
2483 }
2484 let assignment_version = assignment.fence.assignment_version;
2485 self.validate_scope(&scope, expected_assignment_version)?;
2486 if assignment_version == 0 || conn.fence.assignment_version != assignment_version {
2487 return Err(io::Error::new(
2488 io::ErrorKind::ConnectionAborted,
2489 "shuffle assignment changed while opening the stream",
2490 ));
2491 }
2492 if conn.fence.recovery_gen != scope.recovery_gen {
2493 return Err(io::Error::new(
2494 io::ErrorKind::ConnectionAborted,
2495 "shuffle recovery generation changed while opening the stream",
2496 ));
2497 }
2498 let admission_permits = u32::try_from(admission_bytes).map_err(|_| {
2499 io::Error::new(
2500 io::ErrorKind::InvalidInput,
2501 "shuffle admission is too large",
2502 )
2503 })?;
2504 let control = matches!(msg, ShuffleMessage::Barrier(_));
2505 let (peer_byte_budget, node_budget) = if control {
2506 (
2507 &conn.control_byte_budget,
2508 &self.checkpointed_control_node_budget,
2509 )
2510 } else {
2511 (&conn.byte_budget, &self.node_budget)
2512 };
2513 let peer_budget = tokio::select! {
2514 biased;
2515 () = scope.cancel.cancelled() => return Err(scope_cancelled_io()),
2516 permit = Arc::clone(peer_byte_budget).acquire_many_owned(admission_permits) => {
2517 permit.map_err(|_| {
2518 io::Error::new(io::ErrorKind::BrokenPipe, "shuffle byte budget closed")
2519 })?
2520 }
2521 };
2522 let node_budget = tokio::select! {
2523 biased;
2524 () = scope.cancel.cancelled() => return Err(scope_cancelled_io()),
2525 permit = Arc::clone(node_budget).acquire_many_owned(admission_permits) => {
2526 permit.map_err(|_| {
2527 io::Error::new(
2528 io::ErrorKind::BrokenPipe,
2529 "shuffle node byte budget closed",
2530 )
2531 })?
2532 }
2533 };
2534 let budget = OutboundReservation {
2535 peer: peer_budget,
2536 node: node_budget,
2537 };
2538 let (prepared, budget) = prepare_outbound_message(msg, budget).await?;
2539 self.validate_scope(&scope, expected_assignment_version)?;
2540 let gen = scope.recovery_gen;
2543 let seq = {
2544 let assignment = self.assignment.read();
2547 let current = assignment.as_ref().ok_or_else(scope_cancelled_io)?;
2548 self.validate_scope_locked(&scope, expected_assignment_version, current)?;
2549 let mut seqs = self.seqs.lock();
2550 self.validate_scope_locked(&scope, expected_assignment_version, current)?;
2551 if !scope.matches_fence(&conn.fence) {
2552 return Err(io::Error::new(
2553 io::ErrorKind::ConnectionAborted,
2554 "shuffle assignment changed before sequence allocation",
2555 ));
2556 }
2557 let counter = seqs.entry(peer).or_insert(0);
2558 match msg {
2559 ShuffleMessage::Data { .. } | ShuffleMessage::Frontier { .. } => {
2560 let seq = *counter;
2561 *counter = counter.checked_add(1).ok_or_else(|| {
2562 io::Error::other("shuffle delivery sequence exhausted")
2563 })?;
2564 seq
2565 }
2566 ShuffleMessage::Barrier(_) => *counter,
2567 }
2568 };
2569 let out = Outbound {
2570 gen,
2571 assignment_version,
2572 seq,
2573 msg: prepared,
2574 assignment_digest: assignment_fence.map(CheckpointAssignmentFence::digest),
2575 _budget: budget,
2576 };
2577 self.validate_scope(&scope, expected_assignment_version)?;
2578 match conn.tx.try_send(out) {
2579 Ok(()) => {}
2580 Err(crossfire::TrySendError::Full(out)) => tokio::select! {
2581 biased;
2582 () = self.process_lease.wait_until_lost() => {
2583 return Err(process_lease_expired_io());
2584 }
2585 () = scope.cancel.cancelled() => return Err(scope_cancelled_io()),
2586 result = conn.tx.send(out) => result.map_err(|_| io::Error::new(
2587 io::ErrorKind::BrokenPipe,
2588 format!("shuffle stream to peer {peer} closed"),
2589 ))?,
2590 },
2591 Err(crossfire::TrySendError::Disconnected(_)) => {
2592 return Err(io::Error::new(
2593 io::ErrorKind::BrokenPipe,
2594 format!("shuffle stream to peer {peer} closed"),
2595 ));
2596 }
2597 }
2598
2599 #[cfg(test)]
2600 let post_enqueue_pause = { self.post_enqueue_pause.lock().take() };
2601 #[cfg(test)]
2602 if let Some((entered, release)) = post_enqueue_pause {
2603 entered.notify_one();
2604 release.notified().await;
2605 }
2606
2607 self.process_lease
2608 .require_live_io()
2609 .map_err(may_have_admitted_shuffle_frame)
2610 }
2611
2612 fn validate_expected_assignment(&self, expected: Option<u64>) -> io::Result<()> {
2613 let Some(expected) = expected else {
2614 return Ok(());
2615 };
2616 let current = self.assignment_version.load(Ordering::Acquire);
2617 if expected != 0 && expected == current {
2618 return Ok(());
2619 }
2620 Err(io::Error::new(
2621 io::ErrorKind::ConnectionAborted,
2622 format!(
2623 "shuffle assignment scope mismatch: routed at {expected}, sender at {current}"
2624 ),
2625 ))
2626 }
2627
2628 fn current_assignment(&self) -> io::Result<Arc<InstalledAssignment>> {
2629 self.process_lease.require_live_io()?;
2630 let assignment = self.assignment.read().clone().ok_or_else(|| {
2631 io::Error::new(
2632 io::ErrorKind::NotConnected,
2633 "shuffle assignment certificate is not installed",
2634 )
2635 })?;
2636 if self.assignment_version.load(Ordering::Acquire)
2637 != assignment.fence.assignment_version
2638 {
2639 return Err(io::Error::new(
2640 io::ErrorKind::NotConnected,
2641 "shuffle assignment certificate is not active",
2642 ));
2643 }
2644 Ok(assignment)
2645 }
2646
2647 fn current_scope(&self, expected: Option<u64>) -> io::Result<ScopeLease> {
2648 let assignment = self.assignment.read();
2649 let installed = assignment.as_ref().ok_or_else(|| {
2650 io::Error::new(
2651 io::ErrorKind::NotConnected,
2652 "shuffle assignment certificate is not installed",
2653 )
2654 })?;
2655 let version = self.assignment_version.load(Ordering::Acquire);
2656 let recovery_gen = self.recovery_gen.load(Ordering::Acquire);
2657 let cancel = self.scope_cancel.read().clone();
2658 self.process_lease.require_live_io()?;
2659 if version == 0
2660 || version != installed.fence.assignment_version
2661 || cancel.is_cancelled()
2662 {
2663 return Err(scope_cancelled_io());
2664 }
2665 if expected.is_some_and(|expected| expected == 0 || expected != version) {
2666 return Err(io::Error::new(
2667 io::ErrorKind::ConnectionAborted,
2668 format!(
2669 "shuffle assignment scope mismatch: routed at {}, sender at {version}",
2670 expected.unwrap_or_default()
2671 ),
2672 ));
2673 }
2674 Ok(ScopeLease {
2675 assignment: Arc::clone(installed),
2676 recovery_gen,
2677 cancel,
2678 })
2679 }
2680
2681 fn validate_scope(&self, scope: &ScopeLease, expected: Option<u64>) -> io::Result<()> {
2682 let assignment = self.assignment.read();
2683 let current = assignment.as_ref().ok_or_else(scope_cancelled_io)?;
2684 self.validate_scope_locked(scope, expected, current)
2685 }
2686
2687 fn validate_scope_locked(
2688 &self,
2689 scope: &ScopeLease,
2690 expected: Option<u64>,
2691 current: &Arc<InstalledAssignment>,
2692 ) -> io::Result<()> {
2693 self.process_lease.require_live_io()?;
2694 if scope.cancel.is_cancelled() {
2695 return Err(scope_cancelled_io());
2696 }
2697 if !Arc::ptr_eq(current, &scope.assignment)
2698 || self.assignment_version.load(Ordering::Acquire)
2699 != scope.assignment.fence.assignment_version
2700 || self.recovery_gen.load(Ordering::Acquire) != scope.recovery_gen
2701 {
2702 return Err(scope_cancelled_io());
2703 }
2704 self.validate_expected_assignment(expected)
2705 }
2706
2707 pub async fn establish_assignment_mesh(
2713 &self,
2714 assignment_fence: &CheckpointAssignmentFence,
2715 ) -> io::Result<()> {
2716 let installed = self.current_assignment()?;
2717 if *assignment_fence != installed.fence
2718 || assignment_fence.digest() != installed.digest
2719 || !assignment_fence.is_canonical()
2720 || assignment_fence.participant_incarnation(self.local_id)
2721 != Some(self.sender_incarnation)
2722 {
2723 return Err(io::Error::new(
2724 io::ErrorKind::InvalidInput,
2725 "shuffle mesh does not match the installed assignment certificate",
2726 ));
2727 }
2728 let scope = self.current_scope(Some(assignment_fence.assignment_version))?;
2729 let results = futures::future::join_all(
2730 assignment_fence
2731 .participants
2732 .iter()
2733 .map(|participant| participant.node_id)
2734 .filter(|peer| *peer != self.local_id)
2735 .map(|peer| {
2736 let peer_scope = scope.clone();
2737 async move { (peer, self.connection_for(peer, &peer_scope).await) }
2738 }),
2739 )
2740 .await;
2741 let mut first_error = None;
2742 for (peer, result) in results {
2743 if let Err(error) = result {
2744 first_error.get_or_insert_with(|| {
2745 io::Error::new(
2746 error.kind(),
2747 format!("shuffle assignment mesh peer {peer}: {error}"),
2748 )
2749 });
2750 }
2751 }
2752 if let Some(error) = first_error {
2753 return Err(error);
2754 }
2755 self.validate_scope(&scope, Some(assignment_fence.assignment_version))
2756 }
2757
2758 pub async fn fan_out_barrier(
2764 &self,
2765 peers: &[ShufflePeerId],
2766 barrier: CheckpointBarrier,
2767 assignment_fence: &CheckpointAssignmentFence,
2768 ) -> io::Result<()> {
2769 validate_checkpoint_barrier(barrier)?;
2770 let installed = self.current_assignment()?;
2771 let expected_peers: Vec<_> = assignment_fence
2772 .participants
2773 .iter()
2774 .map(|participant| participant.node_id)
2775 .filter(|peer| *peer != self.local_id)
2776 .collect();
2777 let mut actual_peers = peers.to_vec();
2778 actual_peers.sort_unstable();
2779 let has_duplicates = actual_peers.windows(2).any(|pair| pair[0] == pair[1]);
2780 if *assignment_fence != installed.fence
2781 || assignment_fence.digest() != installed.digest
2782 || !assignment_fence.is_canonical()
2783 || !assignment_fence.contains(self.local_id)
2784 || assignment_fence.participant_incarnation(self.local_id)
2785 != Some(self.sender_incarnation)
2786 || has_duplicates
2787 || actual_peers != expected_peers
2788 {
2789 return Err(io::Error::new(
2790 io::ErrorKind::InvalidInput,
2791 "shuffle barrier peers do not exactly cover the assignment roster",
2792 ));
2793 }
2794 self.validate_expected_assignment(Some(assignment_fence.assignment_version))?;
2795 let cid = barrier.checkpoint_id;
2796 let msg = ShuffleMessage::Barrier(barrier);
2797 let mut first_scope_cancel = None;
2798 let mut first_peer_error = None;
2799 let results = futures::future::join_all(peers.iter().map(|&peer| {
2800 let msg = &msg;
2801 async move {
2802 (
2803 peer,
2804 self.send_to_inner(
2805 peer,
2806 msg,
2807 Some(assignment_fence.assignment_version),
2808 Some(assignment_fence),
2809 )
2810 .await,
2811 )
2812 }
2813 }))
2814 .await;
2815 for (peer, result) in results {
2816 match result {
2817 Ok(()) => {}
2818 Err(e) => {
2819 tracing::warn!(
2820 peer,
2821 checkpoint_id = cid,
2822 error = %e,
2823 "shuffle barrier fan-out: required peer unreachable"
2824 );
2825 if is_scope_cancelled(&e) {
2826 first_scope_cancel.get_or_insert(e);
2827 } else {
2828 first_peer_error.get_or_insert(e);
2829 }
2830 }
2831 }
2832 }
2833 if let Some(error) = first_peer_error.or(first_scope_cancel) {
2834 if is_scope_cancelled(&error) {
2835 self.process_lease.require_live_io()?;
2836 }
2837 return Err(error);
2838 }
2839 Ok(())
2840 }
2841
2842 async fn discover_peer(&self, peer: ShufflePeerId) -> Option<SocketAddr> {
2844 let kv = self.kv.as_ref()?;
2845 let raw = kv
2846 .read_from(crate::cluster::discovery::NodeId(peer), SHUFFLE_ADDR_KEY)
2847 .await?;
2848 let addr = match raw.parse::<SocketAddr>() {
2849 Ok(a) => a,
2850 Err(_) => tokio::net::lookup_host(&raw).await.ok()?.next()?,
2851 };
2852 self.peers.lock().insert(peer, addr);
2853 Some(addr)
2854 }
2855
2856 async fn connection_for(
2857 &self,
2858 peer: ShufflePeerId,
2859 scope: &ScopeLease,
2860 ) -> io::Result<Arc<PeerConn>> {
2861 self.validate_scope(scope, None)?;
2862 let assignment = &scope.assignment;
2863 if !assignment.fence.contains(peer) || peer == self.local_id {
2864 return Err(io::Error::new(
2865 io::ErrorKind::InvalidInput,
2866 "shuffle peer is outside the installed assignment roster",
2867 ));
2868 }
2869 let assignment_version = assignment.fence.assignment_version;
2870 let recovery_gen = scope.recovery_gen;
2871 let key = peer;
2872 if let Some(existing) = self.pool.lock().get(&key).cloned() {
2873 if existing.is_alive()
2874 && existing.fence.assignment_version == assignment_version
2875 && existing.fence.assignment_certificate_digest == assignment.digest
2876 && existing.fence.recovery_gen == recovery_gen
2877 {
2878 return Ok(existing);
2879 }
2880 }
2881 let connect_lock = self
2882 .connect_locks
2883 .lock()
2884 .entry(key)
2885 .or_insert_with(|| Arc::new(tokio::sync::Mutex::new(())))
2886 .clone();
2887 let _connect_guard = tokio::select! {
2888 biased;
2889 () = self.process_lease.wait_until_lost() => {
2890 return Err(process_lease_expired_io());
2891 }
2892 () = scope.cancel.cancelled() => return Err(scope_cancelled_io()),
2893 guard = connect_lock.lock() => guard,
2894 };
2895 self.validate_scope(scope, None)?;
2896 if let Some(existing) = self.pool.lock().get(&key).cloned() {
2897 if existing.is_alive()
2898 && existing.fence.assignment_version == assignment_version
2899 && existing.fence.assignment_certificate_digest == assignment.digest
2900 && existing.fence.recovery_gen == recovery_gen
2901 {
2902 return Ok(existing);
2903 }
2904 }
2905 self.pool
2907 .lock()
2908 .retain(|pool_key, connection| *pool_key != key || connection.is_alive());
2909
2910 let discovered = tokio::select! {
2913 biased;
2914 () = self.process_lease.wait_until_lost() => {
2915 return Err(process_lease_expired_io());
2916 }
2917 () = scope.cancel.cancelled() => return Err(scope_cancelled_io()),
2918 discovered = self.discover_peer(peer) => discovered,
2919 };
2920 let addr = match discovered {
2921 Some(addr) => addr,
2922 None => self.peers.lock().get(&peer).copied().ok_or_else(|| {
2923 io::Error::new(
2924 io::ErrorKind::NotFound,
2925 format!("peer {peer} has no registered shuffle address"),
2926 )
2927 })?,
2928 };
2929
2930 tracing::debug!(peer, addr = %addr, "shuffle reconnecting to peer");
2931 let expected_receiver_incarnation = assignment
2932 .fence
2933 .participant_incarnation(peer)
2934 .ok_or_else(|| {
2935 io::Error::new(
2936 io::ErrorKind::PermissionDenied,
2937 format!("shuffle peer {peer} is absent from the assignment certificate"),
2938 )
2939 })?;
2940 let conn = Arc::new(
2941 open_call(OpenCall {
2942 local_id: self.local_id,
2943 peer,
2944 addr,
2945 sender_incarnation: self.sender_incarnation,
2946 assignment_version,
2947 assignment_certificate_digest: assignment.digest,
2948 expected_receiver_incarnation,
2949 recovery_gen,
2950 current_assignment: Arc::clone(&self.assignment_version),
2951 current_recovery_gen: Arc::clone(&self.recovery_gen),
2952 scope_cancel: scope.cancel.clone(),
2953 })
2954 .await?,
2955 );
2956 self.validate_scope(scope, None)?;
2957 let mut pool = self.pool.lock();
2958 if scope.cancel.is_cancelled()
2959 || self.assignment_version.load(Ordering::Acquire) != assignment_version
2960 || self.recovery_gen.load(Ordering::Acquire) != recovery_gen
2961 {
2962 return Err(scope_cancelled_io());
2963 }
2964 pool.insert(key, Arc::clone(&conn));
2965 Ok(conn)
2966 }
2967
2968 #[cfg(test)]
2969 pub(crate) fn disconnect_peer_for_test(&self, peer: ShufflePeerId) {
2970 self.pool.lock().remove(&peer);
2971 }
2972
2973 #[cfg(test)]
2974 pub(crate) async fn hold_outbound_budget_for_test(
2975 &self,
2976 peer: ShufflePeerId,
2977 ) -> io::Result<OwnedSemaphorePermit> {
2978 let scope = self.current_scope(None)?;
2979 let conn = self.connection_for(peer, &scope).await?;
2980 Arc::clone(&conn.byte_budget)
2981 .acquire_many_owned(
2982 u32::try_from(OUTBOUND_PEER_BUDGET_BYTES)
2983 .expect("outbound test budget fits u32"),
2984 )
2985 .await
2986 .map_err(|_| io::Error::new(io::ErrorKind::BrokenPipe, "test budget closed"))
2987 }
2988 }
2989
2990 async fn negotiate_identity(
2992 client: &mut ShuffleTransportClient<Channel>,
2993 call: &OpenCall,
2994 ) -> io::Result<StreamFence> {
2995 let stream_id = Uuid::new_v4();
2996 let response = tokio::select! {
2997 biased;
2998 () = call.scope_cancel.cancelled() => return Err(scope_cancelled_io()),
2999 response = client.handshake(Request::new(HandshakeRequest {
3000 sender_node_id: call.local_id,
3001 sender_incarnation: call.sender_incarnation.as_bytes().to_vec(),
3002 stream_id: stream_id.as_bytes().to_vec(),
3003 assignment_version: call.assignment_version,
3004 recovery_gen: call.recovery_gen,
3005 assignment_certificate_digest: call.assignment_certificate_digest.to_vec(),
3006 })) => response.map_err(status_io)?.into_inner(),
3007 };
3008 let receiver_incarnation = parse_uuid(
3009 &response.receiver_incarnation,
3010 "handshake receiver incarnation",
3011 )
3012 .map_err(io_err)?;
3013 if response.receiver_node_id != call.peer
3014 || receiver_incarnation != call.expected_receiver_incarnation
3015 || response.sender_incarnation.as_slice() != call.sender_incarnation.as_bytes()
3016 || response.stream_id.as_slice() != stream_id.as_bytes()
3017 || response.assignment_version != call.assignment_version
3018 || response.assignment_certificate_digest.as_slice()
3019 != call.assignment_certificate_digest.as_slice()
3020 || response.recovery_gen != call.recovery_gen
3021 {
3022 return Err(io::Error::new(
3023 io::ErrorKind::InvalidData,
3024 "shuffle handshake response did not match the requested stream identity",
3025 ));
3026 }
3027 Ok(StreamFence {
3028 sender_node_id: call.local_id,
3029 sender_incarnation: call.sender_incarnation,
3030 receiver_incarnation,
3031 stream_id,
3032 assignment_version: call.assignment_version,
3033 assignment_certificate_digest: call.assignment_certificate_digest,
3034 recovery_gen: call.recovery_gen,
3035 })
3036 }
3037
3038 async fn open_call(call: OpenCall) -> io::Result<PeerConn> {
3040 let endpoint = crate::cluster::control::tls::client_endpoint(&call.addr.to_string())
3041 .map_err(io_err)?
3042 .tcp_nodelay(true);
3043 let channel = tokio::select! {
3044 biased;
3045 () = call.scope_cancel.cancelled() => return Err(scope_cancelled_io()),
3046 channel = endpoint.connect() => channel.map_err(io_err)?,
3047 };
3048 let mut client = ShuffleTransportClient::<Channel>::new(channel)
3049 .max_decoding_message_size(MAX_SHUFFLE_MESSAGE_BYTES)
3050 .max_encoding_message_size(MAX_SHUFFLE_MESSAGE_BYTES);
3051 let fence = negotiate_identity(&mut client, &call).await?;
3052 let OpenCall {
3053 addr,
3054 current_assignment,
3055 current_recovery_gen,
3056 scope_cancel,
3057 ..
3058 } = call;
3059 let (tx, rx) = mpsc::bounded_async::<Outbound>(SEND_QUEUE);
3060 let alive = Arc::new(AtomicBool::new(true));
3061 let alive_for_driver = Arc::clone(&alive);
3062
3063 let hello = ShuffleFrame {
3066 kind: Some(shuffle_frame::Kind::Hello(hello_for(
3067 fence.sender_node_id,
3068 &fence,
3069 ))),
3070 };
3071 let stream_fence = fence;
3072 let outbound = futures::stream::once(async move { hello }).chain(futures::stream::unfold(
3073 (
3074 rx,
3075 current_assignment,
3076 current_recovery_gen,
3077 scope_cancel.clone(),
3078 stream_fence,
3079 None::<Encoded>,
3080 ),
3081 |(
3082 rx,
3083 current_assignment,
3084 current_recovery_gen,
3085 scope_cancel,
3086 stream_fence,
3087 mut encoded,
3088 )| async move {
3089 loop {
3090 if let Some(pending) = encoded.as_mut() {
3091 if let Some(frame) = pending.frames.pop_front() {
3092 return Some((
3093 frame,
3094 (
3095 rx,
3096 current_assignment,
3097 current_recovery_gen,
3098 scope_cancel,
3099 stream_fence,
3100 encoded,
3101 ),
3102 ));
3103 }
3104 let _ = encoded.take();
3105 }
3106 let out = tokio::select! {
3107 biased;
3108 () = scope_cancel.cancelled() => return None,
3109 out = rx.recv() => out.ok()?,
3110 };
3111 if current_assignment.load(Ordering::Acquire) != stream_fence.assignment_version
3112 || current_recovery_gen.load(Ordering::Acquire) != stream_fence.recovery_gen
3113 {
3114 return None;
3115 }
3116 match frame_message(out) {
3117 Ok(message) => encoded = Some(message),
3118 Err(error) => {
3119 tracing::warn!(%error, "shuffle frame construction failed; closing stream");
3120 return None;
3121 }
3122 }
3123 }
3124 },
3125 ));
3126
3127 let driver = tokio::spawn(async move {
3128 tokio::select! {
3131 biased;
3132 () = scope_cancel.cancelled() => {}
3133 result = client.shuffle(Request::new(outbound)) => {
3134 if let Err(status) = result {
3135 tracing::warn!(peer_addr = %addr, error = %status, "shuffle stream broke");
3136 }
3137 }
3138 }
3139 alive_for_driver.store(false, Ordering::Release);
3140 });
3141
3142 Ok(PeerConn {
3143 tx,
3144 byte_budget: Arc::new(Semaphore::new(OUTBOUND_PEER_BUDGET_BYTES)),
3145 control_byte_budget: Arc::new(Semaphore::new(CHECKPOINTED_CONTROL_PEER_BUDGET_BYTES)),
3146 send_lock: tokio::sync::Mutex::new(()),
3147 alive,
3148 driver,
3149 fence,
3150 })
3151 }
3152
3153 struct PendingInbound<'a> {
3154 slot: &'a Mutex<Option<Inbound>>,
3155 ready: &'a AtomicBool,
3156 inbound: Option<Inbound>,
3157 }
3158
3159 impl PendingInbound<'_> {
3160 fn take(mut self) -> Inbound {
3161 self.inbound.take().expect("pending shuffle inbound")
3162 }
3163 }
3164
3165 impl Drop for PendingInbound<'_> {
3166 fn drop(&mut self) {
3167 if let Some(inbound) = self.inbound.take() {
3168 let mut slot = self.slot.lock();
3169 let replaced = slot.replace(inbound);
3170 assert!(replaced.is_none(), "shuffle deferred receive slot occupied");
3171 self.ready.store(true, Ordering::Release);
3172 }
3173 }
3174 }
3175
3176 pub struct ShuffleReceiver {
3179 local_id: ShufflePeerId,
3180 local_addr: SocketAddr,
3181 receiver_incarnation: Uuid,
3182 rx: Mutex<Option<InboundRx>>,
3187 rx_returned: Arc<tokio::sync::Notify>,
3188 work_ready: Arc<tokio::sync::Notify>,
3189 deferred_recv: Mutex<Option<Inbound>>,
3190 deferred_recv_ready: AtomicBool,
3191 #[cfg(test)]
3192 recv_deferred_pause: Mutex<Option<Arc<tokio::sync::Notify>>>,
3193 #[cfg(test)]
3194 assignment_wait_pause: Mutex<Option<(Arc<tokio::sync::Notify>, Arc<tokio::sync::Notify>)>>,
3195 server: JoinHandle<()>,
3196 holdover: Arc<Holdover>,
3197 barrier_arrivals: Arc<AtomicU64>,
3199 barrier_reconciled: AtomicU64,
3200 recovery_gen: Arc<AtomicU64>,
3202 recovery_transition: Mutex<()>,
3203 assignment: Arc<RwLock<Option<Arc<InstalledAssignment>>>>,
3204 assignment_version: Arc<AtomicU64>,
3205 scope_cancel: Arc<RwLock<CancellationToken>>,
3206 process_lease: Arc<ProcessLeaseGate>,
3207 assignment_suspended: AtomicBool,
3210 assignment_resumed: tokio::sync::Notify,
3211 pending_handshakes: Arc<PendingHandshakes>,
3212 delivery: Arc<DeliveryTracker>,
3213 #[cfg(test)]
3214 active_streams: Arc<Semaphore>,
3215 }
3216
3217 impl Drop for ShuffleReceiver {
3218 fn drop(&mut self) {
3219 self.server.abort();
3222 }
3223 }
3224
3225 impl std::fmt::Debug for ShuffleReceiver {
3226 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
3227 f.debug_struct("ShuffleReceiver")
3228 .field("local_id", &self.local_id)
3229 .field("local_addr", &self.local_addr)
3230 .finish_non_exhaustive()
3231 }
3232 }
3233
3234 impl ShuffleReceiver {
3235 pub async fn bind(
3241 local_id: ShufflePeerId,
3242 addr: SocketAddr,
3243 receiver_incarnation: Uuid,
3244 ) -> io::Result<Self> {
3245 if local_id == 0 || receiver_incarnation.is_nil() {
3246 return Err(io::Error::new(
3247 io::ErrorKind::InvalidInput,
3248 "shuffle receiver requires a nonzero node and non-nil incarnation",
3249 ));
3250 }
3251 let listener = tokio::net::TcpListener::bind(addr).await?;
3252 let local_addr = listener.local_addr()?;
3253 let (tx, rx) = mpsc::bounded_async::<Inbound>(SHUFFLE_RECV_QUEUE);
3254 let work_ready = Arc::new(tokio::sync::Notify::new());
3255
3256 let recovery_gen = Arc::new(AtomicU64::new(0));
3257 let assignment = Arc::new(RwLock::new(None));
3258 let assignment_version = Arc::new(AtomicU64::new(0));
3259 let scope_cancel = Arc::new(RwLock::new(cancelled_token()));
3260 let process_lease = Arc::new(ProcessLeaseGate::default());
3261 let pending_handshakes = Arc::new(PendingHandshakes::default());
3262 let delivery = Arc::new(DeliveryTracker::default());
3263 let barrier_arrivals = Arc::new(AtomicU64::new(0));
3264 let holdover = Arc::new(Holdover::new(SHUFFLE_RECV_QUEUE));
3265 let inbound_budget = Arc::new(InboundBudget::new(INBOUND_NODE_BUDGET_BYTES));
3266 let active_streams = Arc::new(Semaphore::new(MAX_ACTIVE_STREAMS));
3267 let active_stream_registry = Arc::new(ActiveStreamRegistry::default());
3268 let service = ShuffleService {
3269 local_id,
3270 receiver_incarnation,
3271 assignment: Arc::clone(&assignment),
3272 assignment_version: Arc::clone(&assignment_version),
3273 scope_cancel: Arc::clone(&scope_cancel),
3274 process_lease: Arc::clone(&process_lease),
3275 pending_handshakes: Arc::clone(&pending_handshakes),
3276 tx,
3277 work_ready: Arc::clone(&work_ready),
3278 recovery_gen: Arc::clone(&recovery_gen),
3279 delivery: Arc::clone(&delivery),
3280 barrier_arrivals: Arc::clone(&barrier_arrivals),
3281 holdover: Arc::clone(&holdover),
3282 inbound_budget,
3283 active_streams: Arc::clone(&active_streams),
3284 active_stream_registry,
3285 };
3286 let incoming = futures::stream::unfold(listener, |listener| async move {
3287 let item = match listener.accept().await {
3288 Ok((stream, _)) => {
3289 let _ = stream.set_nodelay(true);
3290 Ok(stream)
3291 }
3292 Err(e) => Err(e),
3293 };
3294 Some((item, listener))
3295 });
3296 let mut builder = Server::builder();
3299 if let Some(tls) = crate::cluster::control::tls::server_tls() {
3300 builder = builder
3301 .tls_config(tls.clone())
3302 .map_err(|e| io::Error::other(format!("cluster shuffle TLS config: {e}")))?;
3303 }
3304 let router = builder.add_service(
3305 ShuffleTransportServer::new(service)
3306 .max_decoding_message_size(MAX_SHUFFLE_MESSAGE_BYTES)
3307 .max_encoding_message_size(MAX_SHUFFLE_MESSAGE_BYTES),
3308 );
3309 let server = tokio::spawn(async move {
3310 let _ = router.serve_with_incoming(incoming).await;
3311 });
3312
3313 Ok(Self {
3314 local_id,
3315 local_addr,
3316 receiver_incarnation,
3317 rx: Mutex::new(Some(rx)),
3318 rx_returned: Arc::new(tokio::sync::Notify::new()),
3319 work_ready,
3320 deferred_recv: Mutex::new(None),
3321 deferred_recv_ready: AtomicBool::new(false),
3322 #[cfg(test)]
3323 recv_deferred_pause: Mutex::new(None),
3324 #[cfg(test)]
3325 assignment_wait_pause: Mutex::new(None),
3326 server,
3327 holdover,
3328 barrier_arrivals,
3329 barrier_reconciled: AtomicU64::new(0),
3330 recovery_gen,
3331 recovery_transition: Mutex::new(()),
3332 assignment,
3333 assignment_version,
3334 scope_cancel,
3335 process_lease,
3336 assignment_suspended: AtomicBool::new(false),
3337 assignment_resumed: tokio::sync::Notify::new(),
3338 pending_handshakes,
3339 delivery,
3340 #[cfg(test)]
3341 active_streams,
3342 })
3343 }
3344
3345 pub fn install_process_lease_deadline(
3353 &self,
3354 deadline: Arc<LeaseDeadline>,
3355 ) -> io::Result<()> {
3356 let _assignment = self.assignment.write();
3357 if self.assignment_version.load(Ordering::Acquire) != 0
3358 && !self.process_lease.is_installed_deadline(&deadline)
3359 {
3360 return Err(io::Error::new(
3361 io::ErrorKind::InvalidInput,
3362 "shuffle process lease must be installed before assignment activation",
3363 ));
3364 }
3365 self.process_lease.install(deadline)
3366 }
3367
3368 pub fn install_assignment_fence(
3375 &self,
3376 fence: &CheckpointAssignmentFence,
3377 owners: &[ShufflePeerId],
3378 ) -> io::Result<bool> {
3379 let next = InstalledAssignment::for_process(
3380 fence,
3381 owners,
3382 self.local_id,
3383 self.receiver_incarnation,
3384 )?;
3385 let mut assignment = self.assignment.write();
3386 self.process_lease.require_live_io()?;
3387 if let Some(current) = assignment.as_ref() {
3388 if next.fence.assignment_version < current.fence.assignment_version {
3389 return Ok(false);
3390 }
3391 if next.fence.assignment_version == current.fence.assignment_version {
3392 if next.digest == current.digest
3393 && next.fence == current.fence
3394 && next.owners == current.owners
3395 {
3396 if self.assignment_version.load(Ordering::Acquire)
3397 == next.fence.assignment_version
3398 {
3399 return Ok(false);
3400 }
3401 if self.assignment_suspended.load(Ordering::Acquire) {
3402 rotate_scope_token(&self.scope_cancel, &self.process_lease, true);
3403 self.pending_handshakes.clear();
3404 self.assignment_suspended.store(false, Ordering::Release);
3405 self.assignment_version
3406 .store(next.fence.assignment_version, Ordering::Release);
3407 self.assignment_resumed.notify_waiters();
3408 return Ok(true);
3409 }
3410 return Err(io::Error::new(
3411 io::ErrorKind::InvalidData,
3412 "an invalidated shuffle assignment requires a higher version",
3413 ));
3414 }
3415 return Err(io::Error::new(
3416 io::ErrorKind::InvalidData,
3417 "conflicting shuffle assignment certificate for an installed version",
3418 ));
3419 }
3420 }
3421 rotate_scope_token(&self.scope_cancel, &self.process_lease, false);
3422 self.pending_handshakes.clear();
3423 self.delivery.reset_assignment();
3424 self.holdover.clear_staged_frontiers();
3425 let version = next.fence.assignment_version;
3426 *assignment = Some(next);
3427 self.assignment_suspended.store(false, Ordering::Release);
3428 rotate_scope_token(&self.scope_cancel, &self.process_lease, true);
3429 self.assignment_version.store(version, Ordering::Release);
3430 self.assignment_resumed.notify_waiters();
3431 Ok(true)
3432 }
3433
3434 pub fn suspend_assignment_fence(&self) {
3438 let assignment = self.assignment.write();
3439 if assignment.is_none() || self.assignment_version.load(Ordering::Acquire) == 0 {
3440 return;
3441 }
3442 rotate_scope_token(&self.scope_cancel, &self.process_lease, false);
3443 self.pending_handshakes.clear();
3444 self.assignment_suspended.store(true, Ordering::Release);
3445 self.assignment_version.store(0, Ordering::Release);
3446 }
3447
3448 pub fn invalidate_assignment_fence(&self) {
3451 let _assignment = self.assignment.write();
3452 rotate_scope_token(&self.scope_cancel, &self.process_lease, false);
3453 self.assignment_suspended.store(false, Ordering::Release);
3454 self.assignment_version.store(0, Ordering::Release);
3455 self.pending_handshakes.clear();
3456 self.delivery.reset_assignment();
3457 self.holdover.clear_staged_frontiers();
3458 self.assignment_resumed.notify_waiters();
3459 }
3460
3461 #[must_use]
3463 pub fn assignment_version(&self) -> u64 {
3464 self.assignment_version.load(Ordering::Acquire)
3465 }
3466
3467 #[must_use]
3469 pub fn active_assignment_digest(&self) -> Option<[u8; 32]> {
3470 let assignment = self.assignment.read();
3471 assignment.as_ref().and_then(|installed| {
3472 (self.assignment_version.load(Ordering::Acquire)
3473 == installed.fence.assignment_version)
3474 .then_some(installed.digest)
3475 })
3476 }
3477
3478 pub fn set_recovery_gen(&self, gen: u64) {
3481 let _transition = self.recovery_transition.lock();
3482 let _assignment = self.assignment.write();
3483 let previous = self.recovery_gen.load(Ordering::Acquire);
3484 if gen <= previous {
3485 return;
3486 }
3487 rotate_scope_token(&self.scope_cancel, &self.process_lease, false);
3488 self.pending_handshakes.clear();
3489 self.delivery.prepare_recovery(gen);
3490 self.holdover.clear_staged_frontiers();
3491 self.recovery_gen.store(gen, Ordering::Release);
3492 rotate_scope_token(
3493 &self.scope_cancel,
3494 &self.process_lease,
3495 self.assignment_version.load(Ordering::Acquire) != 0,
3496 );
3497 }
3500
3501 #[must_use]
3503 pub fn recovery_gen(&self) -> u64 {
3504 self.recovery_gen.load(Ordering::Acquire)
3505 }
3506
3507 #[must_use]
3509 pub const fn incarnation(&self) -> Uuid {
3510 self.receiver_incarnation
3511 }
3512
3513 #[must_use]
3515 pub const fn local_id(&self) -> ShufflePeerId {
3516 self.local_id
3517 }
3518
3519 #[cfg(test)]
3520 pub(crate) fn active_streams_for_test(&self) -> usize {
3521 MAX_ACTIVE_STREAMS - self.active_streams.available_permits()
3522 }
3523
3524 #[cfg(test)]
3525 pub(crate) fn committed_sequence_for_test(&self, peer: ShufflePeerId) -> Option<u64> {
3526 self.delivery
3527 .peers
3528 .lock()
3529 .get(&peer)
3530 .map(|state| state.expected)
3531 }
3532
3533 #[cfg(test)]
3534 pub(crate) fn barrier_arrivals_for_test(&self) -> u64 {
3535 self.barrier_arrivals.load(Ordering::Acquire)
3536 }
3537
3538 #[cfg(test)]
3539 pub(super) fn assignment_fence_for_test(&self) -> CheckpointAssignmentFence {
3540 self.assignment
3541 .read()
3542 .as_ref()
3543 .expect("test assignment")
3544 .fence
3545 .clone()
3546 }
3547
3548 #[cfg(test)]
3549 pub(crate) fn pause_next_recv_after_defer_for_test(&self) -> Arc<tokio::sync::Notify> {
3550 let entered = Arc::new(tokio::sync::Notify::new());
3551 let replaced = self
3552 .recv_deferred_pause
3553 .lock()
3554 .replace(Arc::clone(&entered));
3555 assert!(
3556 replaced.is_none(),
3557 "shuffle receive pause already installed"
3558 );
3559 entered
3560 }
3561
3562 #[cfg(test)]
3563 pub(super) fn pause_next_assignment_wait_for_test(
3564 &self,
3565 ) -> (Arc<tokio::sync::Notify>, Arc<tokio::sync::Notify>) {
3566 let entered = Arc::new(tokio::sync::Notify::new());
3567 let release = Arc::new(tokio::sync::Notify::new());
3568 let replaced = self
3569 .assignment_wait_pause
3570 .lock()
3571 .replace((Arc::clone(&entered), Arc::clone(&release)));
3572 assert!(
3573 replaced.is_none(),
3574 "shuffle assignment wait pause already installed"
3575 );
3576 (entered, release)
3577 }
3578
3579 #[cfg(test)]
3580 pub(super) async fn wait_while_assignment_suspended_for_test(&self) -> bool {
3581 self.wait_while_assignment_suspended().await
3582 }
3583
3584 async fn wait_while_assignment_suspended(&self) -> bool {
3585 loop {
3586 let resumed = self.assignment_resumed.notified();
3587 tokio::pin!(resumed);
3588 resumed.as_mut().enable();
3589 if self.process_lease.require_live_io().is_err() {
3590 return false;
3591 }
3592 if !self.assignment_suspended.load(Ordering::Acquire) {
3593 return true;
3594 }
3595 #[cfg(test)]
3596 let assignment_wait_pause = { self.assignment_wait_pause.lock().take() };
3597 #[cfg(test)]
3598 if let Some((entered, release)) = assignment_wait_pause {
3599 entered.notify_one();
3600 release.notified().await;
3601 }
3602 tokio::select! {
3603 biased;
3604 () = self.process_lease.wait_until_lost() => return false,
3605 () = &mut resumed => {}
3606 }
3607 }
3608 }
3609
3610 fn consumption_scope(
3611 &self,
3612 ) -> Option<parking_lot::RwLockReadGuard<'_, Option<Arc<InstalledAssignment>>>> {
3613 let assignment = self.assignment.read();
3614 (!self.assignment_suspended.load(Ordering::Acquire)).then_some(assignment)
3615 }
3616
3617 fn take_deferred_recv(&self) -> Option<Inbound> {
3618 if !self.deferred_recv_ready.load(Ordering::Acquire) {
3619 return None;
3620 }
3621 let mut slot = self.deferred_recv.lock();
3622 let inbound = slot.take();
3623 self.deferred_recv_ready.store(false, Ordering::Release);
3624 inbound
3625 }
3626
3627 #[must_use]
3631 pub fn delivery_loss_incidents(&self) -> Arc<AtomicU64> {
3632 Arc::clone(&self.delivery.delivery_loss_incidents)
3633 }
3634
3635 #[must_use]
3638 pub fn recovered_delivery_loss_incidents(&self) -> Arc<AtomicU64> {
3639 Arc::clone(&self.delivery.recovered_delivery_loss_incidents)
3640 }
3641
3642 #[must_use]
3644 pub fn has_unrecovered_delivery_loss(&self) -> bool {
3645 self.delivery
3646 .delivery_loss_incidents
3647 .load(Ordering::Acquire)
3648 > self
3649 .delivery
3650 .recovered_delivery_loss_incidents
3651 .load(Ordering::Acquire)
3652 }
3653
3654 pub fn complete_recovery(&self, gen: u64) -> bool {
3659 self.recovery_gen.load(Ordering::Acquire) == gen && self.delivery.complete_recovery(gen)
3660 }
3661
3662 pub async fn bind_with_kv(
3667 local_id: ShufflePeerId,
3668 addr: SocketAddr,
3669 kv: Arc<dyn ClusterKv>,
3670 receiver_incarnation: Uuid,
3671 ) -> io::Result<Self> {
3672 let recv = Self::bind(local_id, addr, receiver_incarnation).await?;
3673 kv.write(SHUFFLE_ADDR_KEY, recv.local_addr.to_string())
3674 .await;
3675 Ok(recv)
3676 }
3677
3678 #[must_use]
3680 pub fn local_addr(&self) -> SocketAddr {
3681 self.local_addr
3682 }
3683
3684 #[must_use]
3686 pub fn work_ready_notify(&self) -> Arc<tokio::sync::Notify> {
3687 Arc::clone(&self.work_ready)
3688 }
3689
3690 #[must_use]
3692 pub fn queued_work_ready(&self) -> bool {
3693 if self.deferred_recv_ready.load(Ordering::Acquire) {
3694 return true;
3695 }
3696 self.rx
3697 .lock()
3698 .as_ref()
3699 .is_some_and(|receiver| !receiver.is_empty())
3700 }
3701
3702 fn retain_queued_scope(
3706 &self,
3707 peer: ShufflePeerId,
3708 sender_incarnation: Uuid,
3709 receiver_incarnation: Uuid,
3710 assignment_version: u64,
3711 recovery_gen: u64,
3712 ) -> bool {
3713 if self.process_lease.require_live_io().is_err() {
3714 return false;
3715 }
3716 let Some(assignment) = self.consumption_scope() else {
3717 return false;
3718 };
3719 let Some(assignment) = assignment.as_ref() else {
3720 return false;
3721 };
3722 let current = self.retain_queued_scope_for_assignment(
3723 assignment,
3724 peer,
3725 sender_incarnation,
3726 receiver_incarnation,
3727 assignment_version,
3728 recovery_gen,
3729 );
3730 if self.process_lease.require_live_io().is_err() {
3731 return false;
3732 }
3733 current
3734 }
3735
3736 fn retain_queued_scope_for_assignment(
3737 &self,
3738 assignment: &InstalledAssignment,
3739 peer: ShufflePeerId,
3740 sender_incarnation: Uuid,
3741 receiver_incarnation: Uuid,
3742 assignment_version: u64,
3743 recovery_gen: u64,
3744 ) -> bool {
3745 let current_recovery = self.recovery_gen.load(Ordering::Acquire);
3746 if receiver_incarnation != self.receiver_incarnation {
3747 self.delivery
3748 .note_loss(peer, 1, "queued-receiver-incarnation");
3749 return false;
3750 }
3751 if recovery_gen < current_recovery {
3752 return false;
3755 }
3756 let current_assignment = self.assignment_version.load(Ordering::Acquire);
3757 let current = recovery_gen == current_recovery
3758 && assignment_version == current_assignment
3759 && assignment.fence.assignment_version == assignment_version
3760 && assignment.certifies(peer, sender_incarnation)
3761 && assignment.certifies(self.local_id, receiver_incarnation)
3762 && self.delivery.matches_process_scope(
3763 peer,
3764 sender_incarnation,
3765 receiver_incarnation,
3766 assignment_version,
3767 recovery_gen,
3768 );
3769 if !current {
3770 self.delivery.note_loss(peer, 1, "queued-stream-scope");
3771 }
3772 current
3773 }
3774
3775 fn retain_received(&self, received: &ReceivedShuffle) -> bool {
3776 match self.holdover.is_retired_barrier(received) {
3777 Ok(true) => false,
3778 Err(_) => {
3779 self.delivery
3780 .note_loss(received.peer, 1, "retired-barrier-assignment-digest");
3781 false
3782 }
3783 Ok(false) => self.retain_queued_scope(
3784 received.peer,
3785 received.sender_incarnation,
3786 received.receiver_incarnation,
3787 received.assignment_version,
3788 received.recovery_gen,
3789 ),
3790 }
3791 }
3792
3793 fn retain_batch(&self, received: &ReceivedBatch) -> bool {
3794 self.retain_queued_scope(
3795 received.peer,
3796 received.sender_incarnation,
3797 received.receiver_incarnation,
3798 received.assignment_version,
3799 received.recovery_gen,
3800 )
3801 }
3802
3803 fn stage_frontier_if_current(&self, cut: ReceivedFrontierCut) -> bool {
3804 let frontier = cut.frontier();
3805 let assignment = self.assignment.read();
3806 let Some(assignment) = assignment.as_ref() else {
3807 return false;
3808 };
3809 let current = self.retain_queued_scope_for_assignment(
3810 assignment,
3811 frontier.peer,
3812 frontier.sender_incarnation,
3813 frontier.receiver_incarnation,
3814 frontier.assignment_version,
3815 frontier.recovery_gen,
3816 ) && self.process_lease.require_live_io().is_ok();
3817 if current {
3818 self.holdover.stage_frontier(cut);
3819 }
3820 current
3821 }
3822
3823 pub async fn recv(&self) -> Option<ReceivedShuffle> {
3827 loop {
3828 if !self.wait_while_assignment_suspended().await {
3829 return None;
3830 }
3831
3832 let taken = { self.rx.lock().take() };
3835 let Some(rx) = taken else {
3836 tokio::select! {
3837 biased;
3838 () = self.process_lease.wait_until_lost() => return None,
3839 () = self.rx_returned.notified() => {}
3840 }
3841 continue;
3842 };
3843 let mut guard = RxReturnGuard {
3844 slot: &self.rx,
3845 notify: &self.rx_returned,
3846 rx: Some(rx),
3847 };
3848
3849 if let Some(inbound) = self.take_deferred_recv() {
3850 let pending = PendingInbound {
3851 slot: &self.deferred_recv,
3852 ready: &self.deferred_recv_ready,
3853 inbound: Some(inbound),
3854 };
3855 let received = pending.take().into_received();
3856 drop(guard);
3857 if self.retain_received(&received) {
3858 return Some(received);
3859 }
3860 continue;
3861 }
3862 let inbound = tokio::select! {
3863 biased;
3864 () = self.process_lease.wait_until_lost() => return None,
3865 result = guard.rx.as_mut()?.recv() => result.ok()?,
3866 };
3867 let pending = PendingInbound {
3868 slot: &self.deferred_recv,
3869 ready: &self.deferred_recv_ready,
3870 inbound: Some(inbound),
3871 };
3872 #[cfg(test)]
3873 {
3874 let pause = { self.recv_deferred_pause.lock().take() };
3875 if let Some(entered) = pause {
3876 entered.notify_one();
3877 std::future::pending::<()>().await;
3878 }
3879 }
3880 let received = pending.take().into_received();
3881 drop(guard);
3882 if self.retain_received(&received) {
3883 return Some(received);
3884 }
3885 }
3886 }
3887
3888 #[must_use]
3891 pub fn drain_available(&self) -> Vec<ReceivedShuffle> {
3892 let mut out = Vec::new();
3893 {
3894 let slot = self.rx.lock();
3895 if let Some(rx) = slot.as_ref() {
3896 if let Some(item) = self.take_deferred_recv() {
3897 let received = item.into_received();
3898 if self.retain_received(&received) {
3899 out.push(received);
3900 }
3901 }
3902 while let Ok(item) = rx.try_recv() {
3903 let received = item.into_received();
3904 if self.retain_received(&received) {
3905 out.push(received);
3906 }
3907 }
3908 }
3909 }
3910 out
3911 }
3912
3913 fn drain_inbound_into(&self, staged: &mut FxHashMap<String, Vec<ReceivedBatch>>) -> bool {
3918 if self.holdover.has_staged_barriers() || self.holdover.has_staged_frontiers() {
3919 return false;
3920 }
3921 let slot = self.rx.lock();
3922 let Some(rx) = slot.as_ref() else {
3923 return false;
3924 };
3925 while self.holdover.try_reserve_item() {
3926 let inbound = if let Some(deferred) = self.take_deferred_recv() {
3927 deferred
3928 } else {
3929 let Ok(inbound) = rx.try_recv() else {
3930 self.holdover.release_items(1);
3931 return true;
3932 };
3933 inbound
3934 };
3935 let received = inbound.into_received();
3936 if !self.retain_received(&received) {
3937 self.holdover.release_items(1);
3938 continue;
3939 }
3940 let ReceivedShuffle {
3941 peer,
3942 message,
3943 reservation,
3944 sender_incarnation,
3945 receiver_incarnation,
3946 stream_id,
3947 assignment_version,
3948 assignment_digest,
3949 recovery_gen,
3950 checkpoint_sequence,
3951 } = received;
3952 match message {
3953 ShuffleMessage::Data {
3954 stage: s,
3955 routed_vnodes,
3956 batch,
3957 } => {
3958 staged.entry(s).or_default().push(ReceivedBatch {
3959 batch,
3960 routed_vnodes,
3961 reservation,
3962 peer,
3963 sender_incarnation,
3964 receiver_incarnation,
3965 stream_id,
3966 assignment_version,
3967 recovery_gen,
3968 checkpoint_sequence,
3969 });
3970 }
3971 ShuffleMessage::Frontier {
3972 stage,
3973 watermark,
3974 idle,
3975 } => {
3976 let frontier = ReceivedShuffle {
3977 peer,
3978 message: ShuffleMessage::Frontier {
3979 stage: stage.clone(),
3980 watermark,
3981 idle,
3982 },
3983 reservation,
3984 sender_incarnation,
3985 receiver_incarnation,
3986 stream_id,
3987 assignment_version,
3988 assignment_digest,
3989 recovery_gen,
3990 checkpoint_sequence,
3991 };
3992 let preceding = take_frontier_prefix(staged, &stage, &frontier);
3993 let cut = ReceivedFrontierCut {
3994 preceding,
3995 frontier,
3996 };
3997 let item_count = cut.item_count();
3998 if self.stage_frontier_if_current(cut) {
3999 return false;
4000 }
4001 self.holdover.release_items(item_count);
4002 }
4003 ShuffleMessage::Barrier(b) => {
4004 let barrier = ReceivedShuffle {
4005 peer,
4006 message: ShuffleMessage::Barrier(b),
4007 reservation,
4008 sender_incarnation,
4009 receiver_incarnation,
4010 stream_id,
4011 assignment_version,
4012 assignment_digest,
4013 recovery_gen,
4014 checkpoint_sequence,
4015 };
4016 match self.holdover.stage_barrier(barrier) {
4017 Ok(true) => return false,
4018 Err(_) => self
4019 .delivery
4020 .note_loss(peer, 1, "barrier-holdover-protocol"),
4021 Ok(false) => {}
4022 }
4023 self.holdover.release_items(1);
4024 }
4025 }
4026 }
4027 false
4028 }
4029
4030 #[must_use]
4032 pub fn drain_checkpointed_data_for(&self, stage: &str) -> Vec<ReceivedBatch> {
4033 let mut staged = self.holdover.staged.lock();
4034 let _ = self.drain_inbound_into(&mut staged);
4035 let mut batches = staged.remove(stage).unwrap_or_default();
4036 self.holdover.release_items(batches.len());
4037 batches.retain(|batch| self.retain_batch(batch));
4038 batches
4039 }
4040
4041 #[must_use]
4043 pub fn drain_staged_barriers(&self) -> Vec<ReceivedShuffle> {
4044 let mut barriers = self.holdover.take_staged_barriers();
4045 self.holdover.release_items(barriers.len());
4046 barriers.retain(|barrier| self.retain_received(barrier));
4047 barriers
4048 }
4049
4050 #[must_use]
4052 pub fn drain_staged_frontiers(&self) -> Vec<ReceivedFrontierCut> {
4053 let mut frontiers = self.holdover.take_staged_frontiers();
4054 let item_count = frontiers.iter().map(ReceivedFrontierCut::item_count).sum();
4055 self.holdover.release_items(item_count);
4056 frontiers.retain(|cut| self.retain_received(cut.frontier()));
4057 frontiers
4058 }
4059
4060 #[must_use]
4062 pub fn has_staged_checkpoint_barriers(&self) -> bool {
4063 self.holdover.has_staged_barriers()
4064 }
4065
4066 #[must_use]
4069 pub fn stage_checkpointed_inbound(&self) -> bool {
4070 if self.holdover.has_staged_barriers() {
4071 return true;
4072 }
4073 if self.holdover.has_staged_frontiers() {
4074 return false;
4075 }
4076 let arrived = self.barrier_arrivals.load(Ordering::Acquire);
4077 if arrived == self.barrier_reconciled.load(Ordering::Acquire) {
4078 return false;
4079 }
4080 let mut staged = self.holdover.staged.lock();
4081 let exhausted = self.drain_inbound_into(&mut staged);
4082 let has_barrier = self.holdover.has_staged_barriers();
4083 if exhausted && !has_barrier {
4084 self.barrier_reconciled.store(arrived, Ordering::Release);
4085 }
4086 has_barrier
4087 }
4088
4089 pub fn retire_checkpoint_barriers(
4094 &self,
4095 attempt: CheckpointAttempt,
4096 assignment_digest: [u8; 32],
4097 ) -> io::Result<()> {
4098 self.holdover
4099 .retire_checkpoint_attempt(attempt, assignment_digest)
4100 }
4101
4102 pub fn stash_barrier(&self, barrier: ReceivedShuffle) {
4105 debug_assert!(matches!(barrier.message(), ShuffleMessage::Barrier(_)));
4106 if !self.retain_received(&barrier) {
4107 return;
4108 }
4109 if !self.holdover.try_reserve_item() {
4110 self.delivery
4111 .note_loss(barrier.peer, 1, "barrier-holdover-capacity");
4112 return;
4113 }
4114 let peer = barrier.peer;
4115 match self.holdover.stage_barrier(barrier) {
4116 Ok(true) => {}
4117 Err(_) => {
4118 self.holdover.release_items(1);
4119 self.delivery
4120 .note_loss(peer, 1, "barrier-holdover-protocol");
4121 }
4122 Ok(false) => self.holdover.release_items(1),
4123 }
4124 }
4125
4126 #[must_use]
4128 pub fn drain_checkpointed_staged(&self) -> Vec<(String, ReceivedBatch)> {
4129 let mut staged = self.holdover.staged.lock();
4130 let _ = self.drain_inbound_into(&mut staged);
4131 self.take_checkpointed_staged(&mut staged)
4132 }
4133
4134 pub fn drain_checkpointed_holdover(&self) -> io::Result<Vec<(String, ReceivedBatch)>> {
4143 let mut staged = self.holdover.staged.lock();
4144 self.process_lease.require_live_io()?;
4145 let Some(assignment) = self.consumption_scope() else {
4146 self.process_lease.require_live_io()?;
4147 return Err(scope_cancelled_io());
4148 };
4149 let Some(assignment) = assignment.as_ref() else {
4150 self.process_lease.require_live_io()?;
4151 return Err(scope_cancelled_io());
4152 };
4153 if self.assignment_version.load(Ordering::Acquire) == 0 {
4154 self.process_lease.require_live_io()?;
4155 return Err(scope_cancelled_io());
4156 }
4157 if self.holdover.has_staged_frontiers() {
4158 return Err(io::Error::new(
4159 io::ErrorKind::WouldBlock,
4160 "shuffle frontier cut must be consumed before checkpoint holdover transfer",
4161 ));
4162 }
4163
4164 let item_count = staged.values().map(Vec::len).sum();
4165 let drained = staged
4166 .drain()
4167 .flat_map(|(stage, batches)| {
4168 batches.into_iter().filter_map(move |batch| {
4169 let current = self.retain_queued_scope_for_assignment(
4170 assignment,
4171 batch.peer,
4172 batch.sender_incarnation,
4173 batch.receiver_incarnation,
4174 batch.assignment_version,
4175 batch.recovery_gen,
4176 );
4177 current.then(|| (stage.clone(), batch))
4178 })
4179 })
4180 .collect();
4181 self.holdover.release_items(item_count);
4182 Ok(drained)
4183 }
4184
4185 fn take_checkpointed_staged(
4186 &self,
4187 staged: &mut FxHashMap<String, Vec<ReceivedBatch>>,
4188 ) -> Vec<(String, ReceivedBatch)> {
4189 let item_count = staged.values().map(Vec::len).sum();
4190 let drained = staged
4191 .drain()
4192 .flat_map(|(stage, batches)| {
4193 batches.into_iter().filter_map(move |batch| {
4194 self.retain_batch(&batch).then(|| (stage.clone(), batch))
4195 })
4196 })
4197 .collect();
4198 self.holdover.release_items(item_count);
4199 drained
4200 }
4201
4202 #[must_use]
4204 pub fn drain_all_staged(&self) -> Vec<(String, ReceivedBatch)> {
4205 let mut staged = self.holdover.staged.lock();
4206 let frontiers = self.holdover.take_staged_frontiers();
4207 let item_count = staged.values().map(Vec::len).sum::<usize>()
4208 + frontiers
4209 .iter()
4210 .map(ReceivedFrontierCut::item_count)
4211 .sum::<usize>();
4212 let mut drained: Vec<_> = staged
4213 .drain()
4214 .flat_map(|(stage, batches)| {
4215 batches.into_iter().filter_map(move |batch| {
4216 self.retain_batch(&batch).then(|| (stage.clone(), batch))
4217 })
4218 })
4219 .collect();
4220 for cut in frontiers {
4221 let stage = cut.stage().to_owned();
4222 drained.extend(
4223 cut.preceding
4224 .into_iter()
4225 .filter(|batch| self.retain_batch(batch))
4226 .map(|batch| (stage.clone(), batch)),
4227 );
4228 }
4229 self.holdover.release_items(item_count);
4230 drained
4231 }
4232 }
4233
4234 struct RxReturnGuard<'a> {
4237 slot: &'a Mutex<Option<InboundRx>>,
4238 notify: &'a tokio::sync::Notify,
4239 rx: Option<InboundRx>,
4240 }
4241
4242 impl Drop for RxReturnGuard<'_> {
4243 fn drop(&mut self) {
4244 if let Some(rx) = self.rx.take() {
4245 *self.slot.lock() = Some(rx);
4246 self.notify.notify_one();
4248 }
4249 }
4250 }
4251
4252 struct PeerSeq {
4255 fence: StreamFence,
4256 expected: u64,
4257 }
4258
4259 #[derive(Debug, Clone, Copy)]
4260 struct DataReservation {
4261 fence: StreamFence,
4262 seq: u64,
4263 expected: u64,
4264 }
4265
4266 #[derive(Debug, Clone, Copy)]
4267 struct BarrierReservation {
4268 fence: StreamFence,
4269 last_seq: u64,
4270 expected: u64,
4271 }
4272
4273 struct DataAdmission<'a> {
4276 tracker: &'a DeliveryTracker,
4277 reservation: Option<DataReservation>,
4278 cancel: &'a CancellationToken,
4279 }
4280
4281 impl<'a> DataAdmission<'a> {
4282 fn new(
4283 tracker: &'a DeliveryTracker,
4284 reservation: DataReservation,
4285 cancel: &'a CancellationToken,
4286 ) -> Self {
4287 Self {
4288 tracker,
4289 reservation: Some(reservation),
4290 cancel,
4291 }
4292 }
4293
4294 fn commit_after_enqueue(mut self) -> Result<(), tonic::Status> {
4295 let reservation = self.reservation.take().expect("unresolved data admission");
4296 self.tracker.commit_data(reservation)
4297 }
4298
4299 fn cancel(mut self) {
4300 let _ = self.reservation.take();
4301 }
4302 }
4303
4304 impl Drop for DataAdmission<'_> {
4305 fn drop(&mut self) {
4306 if let Some(reservation) = self.reservation.take() {
4307 if !self.cancel.is_cancelled() {
4308 self.tracker.abort_data(reservation);
4309 }
4310 }
4311 }
4312 }
4313
4314 #[derive(Default)]
4316 struct DeliveryTracker {
4317 peers: Mutex<FxHashMap<ShufflePeerId, PeerSeq>>,
4318 ingress: Mutex<FxHashMap<ShufflePeerId, Arc<tokio::sync::Mutex<()>>>>,
4319 delivery_loss_incidents: Arc<AtomicU64>,
4320 recovered_delivery_loss_incidents: Arc<AtomicU64>,
4321 pending_recovery: Mutex<Option<(u64, u64)>>,
4322 completed_recovery_gen: AtomicU64,
4323 }
4324
4325 impl DeliveryTracker {
4326 fn reset_assignment(&self) {
4327 self.peers.lock().clear();
4328 self.ingress.lock().clear();
4329 }
4330
4331 fn ingress_lock(
4332 &self,
4333 peer: ShufflePeerId,
4334 ) -> Result<Arc<tokio::sync::Mutex<()>>, tonic::Status> {
4335 let mut ingress = self.ingress.lock();
4336 if let Some(lock) = ingress.get(&peer) {
4337 return Ok(Arc::clone(lock));
4338 }
4339 if ingress.len() >= MAX_TRACKED_PEERS {
4340 return Err(tonic::Status::resource_exhausted(
4341 "too many shuffle ingress peers",
4342 ));
4343 }
4344 let lock = Arc::new(tokio::sync::Mutex::new(()));
4345 ingress.insert(peer, Arc::clone(&lock));
4346 Ok(lock)
4347 }
4348
4349 fn prepare_recovery(&self, gen: u64) {
4352 self.ingress.lock().clear();
4353 let mut pending = self.pending_recovery.lock();
4354 if pending.is_some_and(|(pending_gen, _)| pending_gen >= gen) {
4355 return;
4356 }
4357 *pending = Some((gen, self.delivery_loss_incidents.load(Ordering::Acquire)));
4358 }
4359
4360 fn complete_recovery(&self, gen: u64) -> bool {
4362 let mut pending = self.pending_recovery.lock();
4363 if self.completed_recovery_gen.load(Ordering::Acquire) == gen {
4364 return true;
4365 }
4366 let Some((pending_gen, cutoff)) = *pending else {
4367 return false;
4368 };
4369 if pending_gen != gen {
4370 return false;
4371 }
4372 self.recovered_delivery_loss_incidents
4376 .fetch_max(cutoff.min(u64::MAX - 1), Ordering::AcqRel);
4377 self.completed_recovery_gen.store(gen, Ordering::Release);
4378 *pending = None;
4379 true
4380 }
4381
4382 fn observe_hello(&self, fence: StreamFence) -> Result<(), tonic::Status> {
4385 let mut peers = self.peers.lock();
4386 match peers.entry(fence.sender_node_id) {
4387 Entry::Vacant(entry) => {
4388 entry.insert(PeerSeq { fence, expected: 0 });
4389 Ok(())
4390 }
4391 Entry::Occupied(mut entry) => {
4392 let state = entry.get_mut();
4393 let same_process_assignment = state.fence.sender_incarnation
4394 == fence.sender_incarnation
4395 && state.fence.receiver_incarnation == fence.receiver_incarnation
4396 && state.fence.assignment_version == fence.assignment_version
4397 && state.fence.assignment_certificate_digest
4398 == fence.assignment_certificate_digest;
4399 let expected = if same_process_assignment
4400 && state.fence.recovery_gen == fence.recovery_gen
4401 {
4402 state.expected
4403 } else if fence.assignment_version > state.fence.assignment_version {
4404 0
4408 } else if state.fence.assignment_version == fence.assignment_version
4409 && fence.recovery_gen > state.fence.recovery_gen
4410 {
4411 0
4414 } else {
4415 return Err(tonic::Status::failed_precondition(
4416 "shuffle sender scope changed without assignment or recovery advance",
4417 ));
4418 };
4419 *state = PeerSeq { fence, expected };
4420 Ok(())
4421 }
4422 }
4423 }
4424
4425 fn validate_stream(&self, fence: &StreamFence) -> Result<(), tonic::Status> {
4426 let peers = self.peers.lock();
4427 if peers
4428 .get(&fence.sender_node_id)
4429 .is_some_and(|state| state.fence == *fence)
4430 {
4431 Ok(())
4432 } else {
4433 drop(peers);
4434 self.note_loss(fence.sender_node_id, 1, "stale-stream");
4435 Err(tonic::Status::failed_precondition(
4436 "shuffle stream identity was superseded",
4437 ))
4438 }
4439 }
4440
4441 fn matches_process_scope(
4445 &self,
4446 peer: ShufflePeerId,
4447 sender_incarnation: Uuid,
4448 receiver_incarnation: Uuid,
4449 assignment_version: u64,
4450 recovery_gen: u64,
4451 ) -> bool {
4452 self.peers.lock().get(&peer).is_some_and(|state| {
4453 state.fence.sender_incarnation == sender_incarnation
4454 && state.fence.receiver_incarnation == receiver_incarnation
4455 && state.fence.assignment_version == assignment_version
4456 && state.fence.recovery_gen == recovery_gen
4457 })
4458 }
4459
4460 fn reject_protocol(&self, peer: ShufflePeerId, reason: &str) -> tonic::Status {
4461 self.note_loss(peer, 1, "identity");
4462 tonic::Status::failed_precondition(reason.to_string())
4463 }
4464
4465 fn note_loss(&self, peer: ShufflePeerId, missing: u64, at: &str) {
4466 let exhausted = self
4467 .delivery_loss_incidents
4468 .fetch_update(Ordering::AcqRel, Ordering::Acquire, |incidents| {
4469 incidents.checked_add(1)
4470 })
4471 .is_err();
4472 tracing::error!(
4473 peer,
4474 missing,
4475 at,
4476 loss_counter_exhausted = exhausted,
4477 "shuffle frames lost in transit; fencing the epoch"
4478 );
4479 }
4480
4481 fn prepare_data(
4484 &self,
4485 fence: &StreamFence,
4486 seq: u64,
4487 ) -> Result<Option<DataReservation>, tonic::Status> {
4488 if seq == u64::MAX {
4489 return Err(
4490 self.reject_protocol(fence.sender_node_id, "shuffle sequence exhausted")
4491 );
4492 }
4493 let peers = self.peers.lock();
4494 let Some(state) = peers.get(&fence.sender_node_id) else {
4495 drop(peers);
4496 return Err(self.reject_protocol(
4497 fence.sender_node_id,
4498 "shuffle data arrived before its exact Hello",
4499 ));
4500 };
4501 if state.fence != *fence {
4502 drop(peers);
4503 return Err(self.reject_protocol(
4504 fence.sender_node_id,
4505 "shuffle data stream identity was superseded",
4506 ));
4507 }
4508 if seq < state.expected {
4509 return Ok(None);
4510 }
4511 Ok(Some(DataReservation {
4512 fence: *fence,
4513 seq,
4514 expected: state.expected,
4515 }))
4516 }
4517
4518 fn commit_data(&self, reservation: DataReservation) -> Result<(), tonic::Status> {
4520 let next = reservation.seq.checked_add(1).ok_or_else(|| {
4521 self.reject_protocol(
4522 reservation.fence.sender_node_id,
4523 "shuffle sequence exhausted",
4524 )
4525 })?;
4526 let mut peers = self.peers.lock();
4527 let Some(state) = peers.get_mut(&reservation.fence.sender_node_id) else {
4528 drop(peers);
4529 return Err(self.reject_protocol(
4530 reservation.fence.sender_node_id,
4531 "shuffle delivery state disappeared before commit",
4532 ));
4533 };
4534 if state.fence != reservation.fence || state.expected != reservation.expected {
4535 drop(peers);
4536 return Err(self.reject_protocol(
4537 reservation.fence.sender_node_id,
4538 "shuffle delivery scope changed before commit",
4539 ));
4540 }
4541 state.expected = next;
4542 let missing = reservation.seq - reservation.expected;
4543 drop(peers);
4544 if missing > 0 {
4545 self.note_loss(reservation.fence.sender_node_id, missing, "data");
4546 }
4547 Ok(())
4548 }
4549
4550 fn abort_data(&self, reservation: DataReservation) {
4553 let mut peers = self.peers.lock();
4554 let Some(state) = peers.get_mut(&reservation.fence.sender_node_id) else {
4555 return;
4556 };
4557 if state.fence != reservation.fence || state.expected != reservation.expected {
4558 return;
4559 }
4560 let Some(next) = reservation.seq.checked_add(1) else {
4561 drop(peers);
4562 self.note_loss(
4563 reservation.fence.sender_node_id,
4564 1,
4565 "data-admission-sequence-exhausted",
4566 );
4567 return;
4568 };
4569 state.expected = next;
4570 let missing = reservation.seq - reservation.expected + 1;
4571 drop(peers);
4572 self.note_loss(reservation.fence.sender_node_id, missing, "data-admission");
4573 }
4574
4575 fn prepare_barrier(
4578 &self,
4579 fence: &StreamFence,
4580 last_seq: u64,
4581 ) -> Result<BarrierReservation, tonic::Status> {
4582 let peers = self.peers.lock();
4583 let Some(state) = peers.get(&fence.sender_node_id) else {
4584 drop(peers);
4585 return Err(self.reject_protocol(
4586 fence.sender_node_id,
4587 "shuffle barrier arrived before its exact Hello",
4588 ));
4589 };
4590 if state.fence != *fence {
4591 drop(peers);
4592 return Err(self.reject_protocol(
4593 fence.sender_node_id,
4594 "shuffle barrier stream identity was superseded",
4595 ));
4596 }
4597 if last_seq < state.expected {
4598 drop(peers);
4599 return Err(self.reject_protocol(
4600 fence.sender_node_id,
4601 "shuffle barrier high-water moved backwards",
4602 ));
4603 }
4604 Ok(BarrierReservation {
4605 fence: *fence,
4606 last_seq,
4607 expected: state.expected,
4608 })
4609 }
4610
4611 fn commit_barrier(&self, reservation: BarrierReservation) -> Result<(), tonic::Status> {
4612 let mut peers = self.peers.lock();
4613 let Some(state) = peers.get_mut(&reservation.fence.sender_node_id) else {
4614 drop(peers);
4615 return Err(self.reject_protocol(
4616 reservation.fence.sender_node_id,
4617 "shuffle barrier state disappeared before commit",
4618 ));
4619 };
4620 if state.fence != reservation.fence || state.expected != reservation.expected {
4621 drop(peers);
4622 return Err(self.reject_protocol(
4623 reservation.fence.sender_node_id,
4624 "shuffle barrier scope changed before commit",
4625 ));
4626 }
4627 state.expected = reservation.last_seq;
4628 let missing = reservation.last_seq - reservation.expected;
4629 drop(peers);
4630 if missing > 0 {
4631 self.note_loss(reservation.fence.sender_node_id, missing, "barrier");
4632 }
4633 Ok(())
4634 }
4635 }
4636
4637 struct ShuffleService {
4640 local_id: ShufflePeerId,
4641 receiver_incarnation: Uuid,
4642 assignment: Arc<RwLock<Option<Arc<InstalledAssignment>>>>,
4643 assignment_version: Arc<AtomicU64>,
4644 scope_cancel: Arc<RwLock<CancellationToken>>,
4645 process_lease: Arc<ProcessLeaseGate>,
4646 pending_handshakes: Arc<PendingHandshakes>,
4647 tx: InboundTx,
4648 work_ready: Arc<tokio::sync::Notify>,
4649 recovery_gen: Arc<AtomicU64>,
4650 delivery: Arc<DeliveryTracker>,
4651 barrier_arrivals: Arc<AtomicU64>,
4652 holdover: Arc<Holdover>,
4653 inbound_budget: Arc<InboundBudget>,
4654 active_streams: Arc<Semaphore>,
4655 active_stream_registry: Arc<ActiveStreamRegistry>,
4656 }
4657
4658 fn active_receiver_scope(
4659 assignment: &RwLock<Option<Arc<InstalledAssignment>>>,
4660 assignment_version: &AtomicU64,
4661 recovery_gen: &AtomicU64,
4662 scope_cancel: &RwLock<CancellationToken>,
4663 process_lease: &ProcessLeaseGate,
4664 ) -> Result<ScopeLease, tonic::Status> {
4665 let assignment = assignment.read();
4666 let installed = assignment.as_ref().ok_or_else(|| {
4667 tonic::Status::failed_precondition("shuffle assignment certificate is not installed")
4668 })?;
4669 let version = assignment_version.load(Ordering::Acquire);
4670 let recovery_gen = recovery_gen.load(Ordering::Acquire);
4671 let cancel = scope_cancel.read().clone();
4672 process_lease.require_live_status()?;
4673 if version == 0 || version != installed.fence.assignment_version || cancel.is_cancelled() {
4674 return Err(scope_cancelled_status());
4675 }
4676 Ok(ScopeLease {
4677 assignment: Arc::clone(installed),
4678 recovery_gen,
4679 cancel,
4680 })
4681 }
4682
4683 #[tonic::async_trait]
4684 impl ShuffleTransport for ShuffleService {
4685 async fn handshake(
4686 &self,
4687 request: Request<HandshakeRequest>,
4688 ) -> Result<tonic::Response<HandshakeResponse>, tonic::Status> {
4689 let request = request.into_inner();
4690 if request.sender_node_id == 0
4691 || request.sender_node_id == self.local_id
4692 || request.assignment_version == 0
4693 {
4694 return Err(tonic::Status::failed_precondition(
4695 "shuffle peers do not share an established assignment scope",
4696 ));
4697 }
4698 let sender_incarnation = parse_uuid(&request.sender_incarnation, "sender incarnation")?;
4699 let stream_id = parse_uuid(&request.stream_id, "stream id")?;
4700 let requested_digest = parse_certificate_digest(
4701 &request.assignment_certificate_digest,
4702 "assignment certificate digest",
4703 )?;
4704 let scope = active_receiver_scope(
4705 &self.assignment,
4706 &self.assignment_version,
4707 &self.recovery_gen,
4708 &self.scope_cancel,
4709 &self.process_lease,
4710 )?;
4711 if request.recovery_gen != scope.recovery_gen {
4712 return Err(tonic::Status::failed_precondition(
4713 "shuffle peers do not share a recovery generation",
4714 ));
4715 }
4716 if request.assignment_version != scope.assignment.fence.assignment_version
4717 || requested_digest != scope.assignment.digest
4718 || !scope
4719 .assignment
4720 .certifies(request.sender_node_id, sender_incarnation)
4721 {
4722 return Err(tonic::Status::failed_precondition(
4723 "shuffle sender is not certified by the installed assignment",
4724 ));
4725 }
4726 let fence = StreamFence {
4727 sender_node_id: request.sender_node_id,
4728 sender_incarnation,
4729 receiver_incarnation: self.receiver_incarnation,
4730 stream_id,
4731 assignment_version: scope.assignment.fence.assignment_version,
4732 assignment_certificate_digest: scope.assignment.digest,
4733 recovery_gen: scope.recovery_gen,
4734 };
4735 let now = std::time::Instant::now();
4736 let mut pending = self.pending_handshakes.0.lock();
4737 pending.retain(|_, handshake| {
4738 now.saturating_duration_since(handshake.issued_at) < HANDSHAKE_TOKEN_TTL
4739 });
4740 if pending.len() >= MAX_PENDING_HANDSHAKES {
4741 return Err(tonic::Status::resource_exhausted(
4742 "too many unconsumed shuffle handshakes",
4743 ));
4744 }
4745 pending.insert(
4746 request.sender_node_id,
4747 PendingHandshake {
4748 fence,
4749 issued_at: now,
4750 },
4751 );
4752 drop(pending);
4753 if scope.cancel.is_cancelled() || self.process_lease.require_live_status().is_err() {
4754 self.pending_handshakes
4755 .0
4756 .lock()
4757 .remove(&request.sender_node_id);
4758 return if scope.cancel.is_cancelled() {
4759 Err(scope_cancelled_status())
4760 } else {
4761 Err(process_lease_expired_status())
4762 };
4763 }
4764 Ok(tonic::Response::new(HandshakeResponse {
4765 receiver_node_id: self.local_id,
4766 receiver_incarnation: self.receiver_incarnation.as_bytes().to_vec(),
4767 sender_incarnation: sender_incarnation.as_bytes().to_vec(),
4768 stream_id: stream_id.as_bytes().to_vec(),
4769 assignment_version: scope.assignment.fence.assignment_version,
4770 recovery_gen: scope.recovery_gen,
4771 assignment_certificate_digest: scope.assignment.digest.to_vec(),
4772 }))
4773 }
4774
4775 async fn shuffle(
4776 &self,
4777 request: Request<tonic::Streaming<ShuffleFrame>>,
4778 ) -> Result<tonic::Response<ShuffleSummary>, tonic::Status> {
4779 let summary = run_stream(self, request.into_inner()).await?;
4780 Ok(tonic::Response::new(summary))
4781 }
4782 }
4783
4784 fn consume_handshake_token(
4785 pending: &PendingHandshakes,
4786 fence: &StreamFence,
4787 now: std::time::Instant,
4788 ) -> bool {
4789 pending
4790 .0
4791 .lock()
4792 .remove(&fence.sender_node_id)
4793 .is_some_and(|token| {
4794 token.fence == *fence
4795 && now.saturating_duration_since(token.issued_at) < HANDSHAKE_TOKEN_TTL
4796 })
4797 }
4798
4799 async fn admit_stream(
4801 stream: &mut tonic::Streaming<ShuffleFrame>,
4802 receiver_incarnation: Uuid,
4803 assignment: &RwLock<Option<Arc<InstalledAssignment>>>,
4804 assignment_version: &AtomicU64,
4805 recovery_gen: &AtomicU64,
4806 scope_cancel: &RwLock<CancellationToken>,
4807 process_lease: &ProcessLeaseGate,
4808 pending_handshakes: &PendingHandshakes,
4809 delivery: &DeliveryTracker,
4810 active_streams: &Arc<Semaphore>,
4811 active_stream_registry: &Arc<ActiveStreamRegistry>,
4812 ) -> Result<
4813 (
4814 StreamFence,
4815 Arc<tokio::sync::Mutex<()>>,
4816 ScopeLease,
4817 ActiveStreamLease,
4818 ),
4819 tonic::Status,
4820 > {
4821 let scope = active_receiver_scope(
4822 assignment,
4823 assignment_version,
4824 recovery_gen,
4825 scope_cancel,
4826 process_lease,
4827 )?;
4828 let first = tokio::select! {
4829 biased;
4830 () = process_lease.wait_until_lost() => {
4831 return Err(process_lease_expired_status());
4832 }
4833 () = scope.cancel.cancelled() => return Err(scope_cancelled_status()),
4834 first = stream.message() => first?,
4835 }
4836 .ok_or_else(|| tonic::Status::invalid_argument("shuffle stream closed before Hello"))?;
4837 let Some(shuffle_frame::Kind::Hello(hello)) = first.kind else {
4838 return Err(tonic::Status::invalid_argument(
4839 "first shuffle frame must be Hello",
4840 ));
4841 };
4842 let fence = fence_from_hello(&hello)?;
4843 if fence.recovery_gen != scope.recovery_gen {
4844 return Err(tonic::Status::failed_precondition(
4845 "shuffle Hello targets a stale recovery generation",
4846 ));
4847 }
4848 if fence.receiver_incarnation != receiver_incarnation {
4849 return Err(tonic::Status::failed_precondition(
4850 "shuffle stream targets a stale receiver process",
4851 ));
4852 }
4853 if !scope.matches_fence(&fence) || !scope.assignment.matches_stream_sender(&fence) {
4854 return Err(tonic::Status::failed_precondition(
4855 "shuffle stream targets a stale or uncertified assignment",
4856 ));
4857 }
4858 if !consume_handshake_token(pending_handshakes, &fence, std::time::Instant::now()) {
4859 return Err(tonic::Status::failed_precondition(
4860 "shuffle stream did not consume its exact unexpired handshake",
4861 ));
4862 }
4863 let ingress = delivery.ingress_lock(fence.sender_node_id)?;
4864 let ingress_guard = tokio::select! {
4865 biased;
4866 () = process_lease.wait_until_lost() => {
4867 return Err(process_lease_expired_status());
4868 }
4869 () = scope.cancel.cancelled() => return Err(scope_cancelled_status()),
4870 guard = ingress.lock() => guard,
4871 };
4872 delivery.observe_hello(fence)?;
4873 let mut active_stream = active_stream_registry.replace(&fence, &scope.cancel);
4874 tokio::select! {
4875 biased;
4876 () = process_lease.wait_until_lost() => {
4877 return Err(process_lease_expired_status());
4878 }
4879 () = scope.cancel.cancelled() => return Err(scope_cancelled_status()),
4880 result = active_stream.acquire_permit(active_streams) => result?,
4881 }
4882 process_lease.require_live_status()?;
4883 drop(ingress_guard);
4884 Ok((fence, ingress, scope, active_stream))
4885 }
4886
4887 fn validate_active_stream_scope(
4888 assignment_version: &AtomicU64,
4889 recovery_gen: &AtomicU64,
4890 delivery: &DeliveryTracker,
4891 fence: &StreamFence,
4892 cancel: &CancellationToken,
4893 process_lease: &ProcessLeaseGate,
4894 ) -> Result<(), tonic::Status> {
4895 process_lease.require_live_status()?;
4896 if cancel.is_cancelled() {
4897 return Err(scope_cancelled_status());
4898 }
4899 if recovery_gen.load(Ordering::Acquire) != fence.recovery_gen {
4900 return Err(tonic::Status::failed_precondition(
4901 "shuffle recovery generation changed while admitting a frame",
4902 ));
4903 }
4904 if assignment_version.load(Ordering::Acquire) != fence.assignment_version {
4905 return Err(tonic::Status::failed_precondition(
4906 "shuffle assignment changed while admitting a frame",
4907 ));
4908 }
4909 delivery.validate_stream(fence)
4910 }
4911
4912 fn reject_stream_protocol(
4913 delivery: &DeliveryTracker,
4914 fence: &StreamFence,
4915 reason: &str,
4916 ) -> tonic::Status {
4917 delivery.reject_protocol(fence.sender_node_id, reason)
4918 }
4919
4920 async fn publish_barrier(
4924 tx: &InboundTx,
4925 work_ready: &tokio::sync::Notify,
4926 barrier_arrivals: &AtomicU64,
4927 holdover: &Holdover,
4928 assignment_version: &AtomicU64,
4929 recovery_gen: &AtomicU64,
4930 delivery: &DeliveryTracker,
4931 fence: StreamFence,
4932 barrier: CheckpointBarrier,
4933 assignment_digest: [u8; 32],
4934 last_seq: u64,
4935 cancel: &CancellationToken,
4936 process_lease: &ProcessLeaseGate,
4937 ) -> Result<bool, tonic::Status> {
4938 validate_active_stream_scope(
4939 assignment_version,
4940 recovery_gen,
4941 delivery,
4942 &fence,
4943 cancel,
4944 process_lease,
4945 )?;
4946 let retired = holdover
4947 .is_retired_checkpoint_barrier(barrier, assignment_digest)
4948 .map_err(|error| reject_stream_protocol(delivery, &fence, &error.to_string()))?;
4949 let reservation = delivery.prepare_barrier(&fence, last_seq)?;
4950 validate_active_stream_scope(
4951 assignment_version,
4952 recovery_gen,
4953 delivery,
4954 &fence,
4955 cancel,
4956 process_lease,
4957 )?;
4958 delivery.commit_barrier(reservation)?;
4959 if retired {
4960 return Ok(true);
4961 }
4962 tokio::select! {
4963 biased;
4964 () = cancel.cancelled() => Err(scope_cancelled_status()),
4965 result = tx.send(Inbound {
4966 peer: fence.sender_node_id,
4967 msg: ShuffleMessage::Barrier(barrier),
4968 budget: None,
4969 fence,
4970 assignment_digest: Some(assignment_digest),
4971 checkpoint_sequence: last_seq,
4972 }) => {
4973 let enqueued = result.is_ok();
4974 if enqueued {
4975 barrier_arrivals.fetch_add(1, Ordering::Release);
4976 work_ready.notify_one();
4977 }
4978 Ok(enqueued)
4979 },
4980 }
4981 }
4982
4983 async fn run_stream(
4985 service: &ShuffleService,
4986 mut stream: tonic::Streaming<ShuffleFrame>,
4987 ) -> Result<ShuffleSummary, tonic::Status> {
4988 let ShuffleService {
4989 receiver_incarnation,
4990 assignment,
4991 assignment_version,
4992 scope_cancel,
4993 process_lease,
4994 pending_handshakes,
4995 tx,
4996 work_ready,
4997 barrier_arrivals,
4998 holdover,
4999 recovery_gen,
5000 delivery,
5001 inbound_budget,
5002 active_streams,
5003 active_stream_registry,
5004 ..
5005 } = service;
5006 let (fence, ingress, scope, active_stream) = admit_stream(
5007 &mut stream,
5008 *receiver_incarnation,
5009 assignment,
5010 assignment_version,
5011 recovery_gen,
5012 scope_cancel,
5013 process_lease,
5014 pending_handshakes,
5015 delivery,
5016 active_streams,
5017 active_stream_registry,
5018 )
5019 .await?;
5020 let stream_cancel = &active_stream.cancel;
5021 let peer = fence.sender_node_id;
5022
5023 let mut assembly = None;
5024 let mut frames_received = 0u64;
5025 loop {
5026 let frame = tokio::select! {
5027 biased;
5028 () = stream_cancel.cancelled() => return Err(scope_cancelled_status()),
5029 frame = stream.message() => frame?,
5030 };
5031 let Some(frame) = frame else {
5032 break;
5033 };
5034 process_lease.require_live_status()?;
5035 let Some(kind) = frame.kind else {
5036 return Err(reject_stream_protocol(
5037 delivery,
5038 &fence,
5039 "empty shuffle frame",
5040 ));
5041 };
5042 if assignment_version.load(Ordering::Acquire) != fence.assignment_version {
5043 return Err(tonic::Status::failed_precondition(
5044 "shuffle assignment changed while the stream was active",
5045 ));
5046 }
5047 if recovery_gen.load(Ordering::Acquire) != fence.recovery_gen {
5048 return Err(tonic::Status::failed_precondition(
5049 "shuffle recovery generation changed while the stream was active",
5050 ));
5051 }
5052 match kind {
5053 shuffle_frame::Kind::Hello(_) => {
5054 let _ingress_guard = tokio::select! {
5055 biased;
5056 () = stream_cancel.cancelled() => return Err(scope_cancelled_status()),
5057 guard = ingress.lock() => guard,
5058 };
5059 return Err(reject_stream_protocol(
5060 delivery,
5061 &fence,
5062 "shuffle Hello is valid only as the leading stream frame",
5063 ));
5064 }
5065 shuffle_frame::Kind::Barrier(b) => {
5066 frames_received += 1;
5067 if assembly.is_some() {
5068 return Err(reject_stream_protocol(
5069 delivery,
5070 &fence,
5071 "shuffle barrier arrived before its preceding batch completed",
5072 ));
5073 }
5074 let barrier = CheckpointBarrier {
5075 checkpoint_id: b.checkpoint_id,
5076 epoch: b.epoch,
5077 flags: b.flags,
5078 };
5079 if !barrier.is_canonical() {
5080 return Err(reject_stream_protocol(
5081 delivery,
5082 &fence,
5083 NONCANONICAL_BARRIER,
5084 ));
5085 }
5086 let assignment_digest: [u8; 32] =
5087 b.assignment_digest.as_slice().try_into().map_err(|_| {
5088 reject_stream_protocol(
5089 delivery,
5090 &fence,
5091 "shuffle barrier assignment digest is not SHA-256 sized",
5092 )
5093 })?;
5094 if b.assignment_version != fence.assignment_version
5095 || b.recovery_gen != fence.recovery_gen
5096 || assignment_digest != fence.assignment_certificate_digest
5097 {
5098 return Err(reject_stream_protocol(
5099 delivery,
5100 &fence,
5101 "shuffle barrier differs from its stream assignment or recovery scope",
5102 ));
5103 }
5104 let _ingress_guard = tokio::select! {
5105 biased;
5106 () = stream_cancel.cancelled() => return Err(scope_cancelled_status()),
5107 guard = ingress.lock() => guard,
5108 };
5109 if !publish_barrier(
5110 tx,
5111 work_ready,
5112 barrier_arrivals,
5113 holdover,
5114 assignment_version,
5115 recovery_gen,
5116 delivery,
5117 fence,
5118 barrier,
5119 assignment_digest,
5120 b.last_seq,
5121 stream_cancel,
5122 process_lease,
5123 )
5124 .await?
5125 {
5126 break;
5127 }
5128 }
5129 shuffle_frame::Kind::Frontier(v) => {
5130 frames_received += 1;
5131 if assembly.is_some() {
5132 return Err(reject_stream_protocol(
5133 delivery,
5134 &fence,
5135 "shuffle frontier arrived before its preceding batch completed",
5136 ));
5137 }
5138 if v.recovery_gen != fence.recovery_gen {
5139 return Err(reject_stream_protocol(
5140 delivery,
5141 &fence,
5142 "shuffle frontier generation differs from its stream handshake",
5143 ));
5144 }
5145 validate_frontier(&v.stage, v.watermark).map_err(|error| {
5146 reject_stream_protocol(delivery, &fence, &error.to_string())
5147 })?;
5148 let _ingress_guard = tokio::select! {
5149 biased;
5150 () = stream_cancel.cancelled() => return Err(scope_cancelled_status()),
5151 guard = ingress.lock() => guard,
5152 };
5153 validate_active_stream_scope(
5154 assignment_version,
5155 recovery_gen,
5156 delivery,
5157 &fence,
5158 stream_cancel,
5159 process_lease,
5160 )?;
5161 let Some(reservation) = delivery.prepare_data(&fence, v.seq)? else {
5162 continue;
5163 };
5164 let mut reservation =
5165 Some(DataAdmission::new(delivery, reservation, stream_cancel));
5166 let forwarded = forward_frontier(
5167 tx,
5168 work_ready,
5169 fence,
5170 v.stage,
5171 v.watermark,
5172 v.idle,
5173 v.seq,
5174 stream_cancel,
5175 process_lease,
5176 )
5177 .await;
5178 match forwarded {
5179 Ok(true) => {
5180 if let Some(admission) = reservation.take() {
5181 admission.commit_after_enqueue()?;
5182 }
5183 }
5184 Ok(false) => break,
5185 Err(status) => {
5186 if status.code() == tonic::Code::Cancelled {
5187 if let Some(admission) = reservation.take() {
5188 admission.cancel();
5189 }
5190 }
5191 return Err(status);
5192 }
5193 }
5194 validate_active_stream_scope(
5195 assignment_version,
5196 recovery_gen,
5197 delivery,
5198 &fence,
5199 stream_cancel,
5200 process_lease,
5201 )?;
5202 }
5203 shuffle_frame::Kind::Data(v) => {
5204 frames_received += 1;
5205 if v.recovery_gen != fence.recovery_gen {
5206 return Err(reject_stream_protocol(
5207 delivery,
5208 &fence,
5209 "shuffle data generation differs from its stream handshake",
5210 ));
5211 }
5212 let total_payload_bytes = validate_fragment(&v)
5213 .map_err(|error| reject_stream_protocol(delivery, &fence, &error))?;
5214 if (v.fragment_index == 0 && assembly.is_some())
5215 || (v.fragment_index != 0 && assembly.is_none())
5216 {
5217 return Err(reject_stream_protocol(
5218 delivery,
5219 &fence,
5220 "shuffle fragments interleaved or started without fragment zero",
5221 ));
5222 }
5223 let budget = if v.fragment_index == 0 {
5224 Some(
5225 inbound_budget
5226 .reserve_frame(peer, total_payload_bytes, stream_cancel)
5227 .await?,
5228 )
5229 } else {
5230 None
5231 };
5232 process_lease.require_live_status()?;
5233 if assignment_version.load(Ordering::Acquire) != fence.assignment_version
5234 || recovery_gen.load(Ordering::Acquire) != fence.recovery_gen
5235 {
5236 return Err(tonic::Status::failed_precondition(
5237 "shuffle scope changed while the frame awaited memory admission",
5238 ));
5239 }
5240 let complete = match push_fragment(&mut assembly, &v, budget) {
5241 Ok(Some(complete)) => complete,
5242 Ok(None) => continue,
5243 Err(error) => {
5244 return Err(reject_stream_protocol(delivery, &fence, &error));
5245 }
5246 };
5247 drop(v);
5250 let CompleteData {
5251 stage,
5252 routed_vnodes,
5253 seq,
5254 arrow_ipc,
5255 budget,
5256 } = complete;
5257 let (batch, budget) = decode_ipc_payload_isolated(arrow_ipc, budget, || {})
5258 .await
5259 .map_err(|error| {
5260 reject_stream_protocol(
5261 delivery,
5262 &fence,
5263 &format!("invalid shuffle IPC: {error}"),
5264 )
5265 })?;
5266 let decoded_bytes = InboundBudget::validate_decoded(std::slice::from_ref(
5267 &batch,
5268 ))
5269 .map_err(|status| reject_stream_protocol(delivery, &fence, status.message()))?;
5270 let _ingress_guard = tokio::select! {
5271 biased;
5272 () = stream_cancel.cancelled() => return Err(scope_cancelled_status()),
5273 guard = ingress.lock() => guard,
5274 };
5275 validate_active_stream_scope(
5276 assignment_version,
5277 recovery_gen,
5278 delivery,
5279 &fence,
5280 stream_cancel,
5281 process_lease,
5282 )?;
5283 let Some(reservation) = delivery.prepare_data(&fence, seq)? else {
5284 continue; };
5286 let mut reservation =
5287 Some(DataAdmission::new(delivery, reservation, stream_cancel));
5288 let forwarded = forward_routed_batch(
5289 tx,
5290 work_ready,
5291 fence,
5292 service.local_id,
5293 &scope.assignment,
5294 stage,
5295 routed_vnodes,
5296 batch,
5297 budget,
5298 decoded_bytes,
5299 seq,
5300 stream_cancel,
5301 process_lease,
5302 )
5303 .await;
5304 match forwarded {
5305 Ok(true) => {
5306 if let Some(admission) = reservation.take() {
5307 admission.commit_after_enqueue()?;
5308 }
5309 }
5310 Ok(false) => break,
5311 Err(status) => {
5312 if status.code() == tonic::Code::Cancelled {
5313 if let Some(admission) = reservation.take() {
5314 admission.cancel();
5315 }
5316 }
5317 return Err(status);
5318 }
5319 }
5320 validate_active_stream_scope(
5321 assignment_version,
5322 recovery_gen,
5323 delivery,
5324 &fence,
5325 stream_cancel,
5326 process_lease,
5327 )?;
5328 }
5329 }
5330 }
5331 if assembly.is_some() {
5332 return Err(reject_stream_protocol(
5333 delivery,
5334 &fence,
5335 "shuffle stream ended mid-fragment",
5336 ));
5337 }
5338 tracing::debug!(peer, frames_received, "shuffle inbound stream ended");
5339 Ok(ShuffleSummary { frames_received })
5340 }
5341
5342 async fn forward_frontier(
5343 tx: &InboundTx,
5344 work_ready: &tokio::sync::Notify,
5345 fence: StreamFence,
5346 stage: String,
5347 watermark: Option<i64>,
5348 idle: bool,
5349 checkpoint_sequence: u64,
5350 cancel: &CancellationToken,
5351 process_lease: &ProcessLeaseGate,
5352 ) -> Result<bool, tonic::Status> {
5353 process_lease.require_live_status()?;
5354 validate_frontier(&stage, watermark)
5355 .map_err(|error| tonic::Status::invalid_argument(error.to_string()))?;
5356 tokio::select! {
5357 biased;
5358 () = cancel.cancelled() => Err(scope_cancelled_status()),
5359 result = tx.send(Inbound {
5360 peer: fence.sender_node_id,
5361 msg: ShuffleMessage::Frontier { stage, watermark, idle },
5362 budget: None,
5363 fence,
5364 assignment_digest: None,
5365 checkpoint_sequence,
5366 }) => {
5367 let enqueued = result.is_ok();
5368 if enqueued {
5369 work_ready.notify_one();
5370 }
5371 Ok(enqueued)
5372 },
5373 }
5374 }
5375
5376 async fn forward_routed_batch(
5379 tx: &InboundTx,
5380 work_ready: &tokio::sync::Notify,
5381 fence: StreamFence,
5382 receiver_node_id: ShufflePeerId,
5383 assignment: &InstalledAssignment,
5384 stage: String,
5385 routed_vnodes: Vec<u32>,
5386 batch: RecordBatch,
5387 mut budget: InboundReservation,
5388 decoded_bytes: usize,
5389 checkpoint_sequence: u64,
5390 cancel: &CancellationToken,
5391 process_lease: &ProcessLeaseGate,
5392 ) -> Result<bool, tonic::Status> {
5393 process_lease.require_live_status()?;
5394 if routed_vnodes.is_empty() || !routed_vnodes.windows(2).all(|pair| pair[0] < pair[1]) {
5395 return Err(tonic::Status::invalid_argument(
5396 "shuffle route set is empty or non-canonical",
5397 ));
5398 }
5399 if let Some(vnode) = routed_vnodes
5400 .iter()
5401 .find(|vnode| !assignment.owns_vnode(receiver_node_id, **vnode))
5402 {
5403 return Err(tonic::Status::failed_precondition(format!(
5404 "shuffle vnode {vnode} is not owned by receiver {receiver_node_id}"
5405 )));
5406 }
5407 let routed_vnodes: Arc<[u32]> = routed_vnodes.into();
5408 let msg = ShuffleMessage::Data {
5409 stage,
5410 routed_vnodes,
5411 batch,
5412 };
5413 let ShuffleMessage::Data {
5414 stage,
5415 routed_vnodes,
5416 batch,
5417 ..
5418 } = &msg
5419 else {
5420 unreachable!("constructed data message")
5421 };
5422 let metadata_bytes = retained_batch_metadata_bytes(stage, routed_vnodes, batch)?;
5423 budget.retain_decoded(decoded_bytes, metadata_bytes)?;
5424 let budget = Arc::new(budget);
5425 tokio::select! {
5426 biased;
5427 () = cancel.cancelled() => Err(scope_cancelled_status()),
5428 result = tx.send(Inbound {
5429 peer: fence.sender_node_id,
5430 msg,
5431 budget: Some(budget),
5432 fence,
5433 assignment_digest: None,
5434 checkpoint_sequence,
5435 }) => {
5436 let enqueued = result.is_ok();
5437 if enqueued {
5438 work_ready.notify_one();
5439 }
5440 Ok(enqueued)
5441 },
5442 }
5443 }
5444
5445 #[cfg(test)]
5446 mod delivery_tests;
5447
5448 #[cfg(test)]
5449 mod encode_tests;
5450
5451 #[cfg(test)]
5452 mod fragment_tests;
5453}
5454
5455#[cfg(feature = "cluster")]
5456pub use grpc::{ShuffleReceiver, ShuffleSender};
5457
5458#[cfg(not(feature = "cluster"))]
5459mod shim {
5460 use std::io;
5461 use std::net::SocketAddr;
5462 use std::sync::Arc;
5463
5464 use crossfire::{mpsc, AsyncRx, MAsyncTx};
5465 use parking_lot::Mutex;
5466 use rustc_hash::FxHashMap;
5467
5468 use super::{
5469 take_frontier_prefix, validate_checkpoint_barrier, validate_frontier, Holdover,
5470 ReceivedBatch, ReceivedFrontierCut, ReceivedShuffle, ShuffleMessage, ShufflePeerId,
5471 SHUFFLE_RECV_QUEUE,
5472 };
5473 use crate::checkpoint::{CheckpointAssignmentFence, CheckpointBarrier};
5474
5475 type InboundRx = AsyncRx<mpsc::Array<ReceivedShuffle>>;
5476 type InboundTx = MAsyncTx<mpsc::Array<ReceivedShuffle>>;
5477
5478 pub struct ShuffleSender {
5480 local_id: ShufflePeerId,
5481 }
5482
5483 impl std::fmt::Debug for ShuffleSender {
5484 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
5485 f.debug_struct("ShuffleSender")
5486 .field("local_id", &self.local_id)
5487 .finish_non_exhaustive()
5488 }
5489 }
5490
5491 impl ShuffleSender {
5492 #[must_use]
5494 pub fn new(local_id: ShufflePeerId, _incarnation: uuid::Uuid) -> Self {
5495 Self { local_id }
5496 }
5497
5498 #[must_use]
5500 pub const fn local_id(&self) -> ShufflePeerId {
5501 self.local_id
5502 }
5503
5504 pub fn set_recovery_gen(&self, _gen: u64) {}
5506
5507 pub fn install_assignment_fence(
5512 &self,
5513 _fence: &CheckpointAssignmentFence,
5514 _owners: &[ShufflePeerId],
5515 ) -> io::Result<bool> {
5516 Ok(false)
5517 }
5518
5519 pub fn invalidate_assignment_fence(&self) {}
5521
5522 pub fn suspend_assignment_fence(&self) {}
5524
5525 #[must_use]
5527 pub const fn assignment_version(&self) -> u64 {
5528 0
5529 }
5530
5531 #[must_use]
5533 pub const fn active_assignment_digest(&self) -> Option<[u8; 32]> {
5534 None
5535 }
5536
5537 #[must_use]
5539 pub const fn recovery_gen(&self) -> u64 {
5540 0
5541 }
5542
5543 pub fn register_peer(&self, _peer: ShufflePeerId, _addr: SocketAddr) {}
5545
5546 pub fn send_to(
5549 &self,
5550 peer: ShufflePeerId,
5551 msg: &ShuffleMessage,
5552 ) -> std::future::Ready<io::Result<()>> {
5553 let validation = match msg {
5554 ShuffleMessage::Barrier(barrier) => validate_checkpoint_barrier(*barrier),
5555 ShuffleMessage::Frontier {
5556 stage, watermark, ..
5557 } => validate_frontier(stage, *watermark),
5558 ShuffleMessage::Data { .. } => Ok(()),
5559 };
5560 if let Err(error) = validation {
5561 return std::future::ready(Err(error));
5562 }
5563 std::future::ready(Err(io::Error::new(
5564 io::ErrorKind::Unsupported,
5565 format!(
5566 "node {} cannot send shuffle to peer {peer}: cluster transport is disabled",
5567 self.local_id
5568 ),
5569 )))
5570 }
5571
5572 pub fn establish_assignment_mesh(
5577 &self,
5578 assignment_fence: &CheckpointAssignmentFence,
5579 ) -> std::future::Ready<io::Result<()>> {
5580 let result = if assignment_fence
5581 .participants
5582 .iter()
5583 .all(|participant| participant.node_id == self.local_id)
5584 {
5585 Ok(())
5586 } else {
5587 Err(io::Error::new(
5588 io::ErrorKind::Unsupported,
5589 "cluster shuffle mesh is disabled",
5590 ))
5591 };
5592 std::future::ready(result)
5593 }
5594
5595 pub async fn fan_out_barrier(
5601 &self,
5602 peers: &[ShufflePeerId],
5603 barrier: CheckpointBarrier,
5604 _assignment_fence: &CheckpointAssignmentFence,
5605 ) -> io::Result<()> {
5606 validate_checkpoint_barrier(barrier)?;
5607 let msg = ShuffleMessage::Barrier(barrier);
5608 let mut first_err = None;
5609 let results =
5610 futures::future::join_all(peers.iter().map(|&peer| self.send_to(peer, &msg))).await;
5611 for result in results {
5612 first_err = first_err.or(result.err());
5613 }
5614 match first_err {
5615 Some(error) => Err(error),
5616 None => Ok(()),
5617 }
5618 }
5619
5620 pub async fn send_to_for_assignment(
5625 &self,
5626 peer: ShufflePeerId,
5627 _expected_assignment_version: u64,
5628 msg: &ShuffleMessage,
5629 ) -> io::Result<()> {
5630 self.send_to(peer, msg).await
5631 }
5632 }
5633
5634 pub struct ShuffleReceiver {
5638 local_id: ShufflePeerId,
5639 local_addr: SocketAddr,
5640 _tx: InboundTx,
5643 rx: Mutex<Option<InboundRx>>,
5644 rx_returned: Arc<tokio::sync::Notify>,
5645 holdover: Arc<Holdover>,
5646 }
5647
5648 impl std::fmt::Debug for ShuffleReceiver {
5649 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
5650 f.debug_struct("ShuffleReceiver")
5651 .field("local_id", &self.local_id)
5652 .field("local_addr", &self.local_addr)
5653 .finish_non_exhaustive()
5654 }
5655 }
5656
5657 impl ShuffleReceiver {
5658 pub fn set_recovery_gen(&self, _gen: u64) {}
5660
5661 pub fn install_assignment_fence(
5666 &self,
5667 _fence: &CheckpointAssignmentFence,
5668 _owners: &[ShufflePeerId],
5669 ) -> io::Result<bool> {
5670 Ok(false)
5671 }
5672
5673 pub fn invalidate_assignment_fence(&self) {}
5675
5676 pub fn suspend_assignment_fence(&self) {}
5678
5679 #[must_use]
5681 pub const fn assignment_version(&self) -> u64 {
5682 0
5683 }
5684
5685 #[must_use]
5687 pub const fn active_assignment_digest(&self) -> Option<[u8; 32]> {
5688 None
5689 }
5690
5691 #[must_use]
5693 pub const fn recovery_gen(&self) -> u64 {
5694 0
5695 }
5696
5697 #[must_use]
5699 pub fn delivery_loss_incidents(&self) -> Arc<std::sync::atomic::AtomicU64> {
5700 Arc::new(std::sync::atomic::AtomicU64::new(0))
5701 }
5702
5703 #[must_use]
5705 pub fn recovered_delivery_loss_incidents(&self) -> Arc<std::sync::atomic::AtomicU64> {
5706 Arc::new(std::sync::atomic::AtomicU64::new(0))
5707 }
5708
5709 #[must_use]
5711 pub const fn has_unrecovered_delivery_loss(&self) -> bool {
5712 false
5713 }
5714
5715 pub async fn bind(
5718 local_id: ShufflePeerId,
5719 addr: SocketAddr,
5720 _incarnation: uuid::Uuid,
5721 ) -> io::Result<Self> {
5722 let listener = tokio::net::TcpListener::bind(addr).await?;
5724 let local_addr = listener.local_addr()?;
5725 drop(listener);
5726 let (tx, rx) = mpsc::bounded_async::<ReceivedShuffle>(SHUFFLE_RECV_QUEUE);
5727 Ok(Self {
5728 local_id,
5729 local_addr,
5730 _tx: tx,
5731 rx: Mutex::new(Some(rx)),
5732 rx_returned: Arc::new(tokio::sync::Notify::new()),
5733 holdover: Arc::new(Holdover::default()),
5734 })
5735 }
5736
5737 #[must_use]
5739 pub fn local_addr(&self) -> SocketAddr {
5740 self.local_addr
5741 }
5742
5743 pub async fn recv(&self) -> Option<ReceivedShuffle> {
5745 loop {
5746 let taken = { self.rx.lock().take() };
5747 let Some(rx) = taken else {
5748 self.rx_returned.notified().await;
5749 continue;
5750 };
5751 let mut guard = RxReturnGuard {
5752 slot: &self.rx,
5753 notify: &self.rx_returned,
5754 rx: Some(rx),
5755 };
5756 return guard.rx.as_mut()?.recv().await.ok();
5757 }
5758 }
5759
5760 #[must_use]
5762 pub fn drain_available(&self) -> Vec<ReceivedShuffle> {
5763 let mut out = Vec::new();
5764 let slot = self.rx.lock();
5765 if let Some(rx) = slot.as_ref() {
5766 while let Ok(item) = rx.try_recv() {
5767 out.push(item);
5768 }
5769 }
5770 out
5771 }
5772
5773 fn drain_inbound_into(&self, staged: &mut FxHashMap<String, Vec<ReceivedBatch>>) {
5774 if self.holdover.has_staged_barriers() || self.holdover.has_staged_frontiers() {
5775 return;
5776 }
5777 let slot = self.rx.lock();
5778 if let Some(rx) = slot.as_ref() {
5779 while self.holdover.try_reserve_item() {
5780 let Ok(received) = rx.try_recv() else {
5781 self.holdover.release_items(1);
5782 break;
5783 };
5784 let ReceivedShuffle {
5785 peer,
5786 message,
5787 reservation,
5788 sender_incarnation,
5789 receiver_incarnation,
5790 stream_id,
5791 assignment_version,
5792 assignment_digest: _,
5793 recovery_gen,
5794 checkpoint_sequence,
5795 } = received;
5796 match message {
5797 ShuffleMessage::Data {
5798 stage,
5799 routed_vnodes,
5800 batch,
5801 } => {
5802 staged.entry(stage).or_default().push(ReceivedBatch {
5803 batch,
5804 routed_vnodes,
5805 reservation,
5806 peer,
5807 sender_incarnation,
5808 receiver_incarnation,
5809 stream_id,
5810 assignment_version,
5811 recovery_gen,
5812 checkpoint_sequence,
5813 });
5814 }
5815 ShuffleMessage::Frontier {
5816 stage,
5817 watermark,
5818 idle,
5819 } => {
5820 let frontier = ReceivedShuffle {
5821 peer,
5822 message: ShuffleMessage::Frontier {
5823 stage: stage.clone(),
5824 watermark,
5825 idle,
5826 },
5827 reservation,
5828 sender_incarnation,
5829 receiver_incarnation,
5830 stream_id,
5831 assignment_version,
5832 assignment_digest: None,
5833 recovery_gen,
5834 checkpoint_sequence,
5835 };
5836 let preceding = take_frontier_prefix(staged, &stage, &frontier);
5837 self.holdover.stage_frontier(ReceivedFrontierCut {
5838 preceding,
5839 frontier,
5840 });
5841 break;
5842 }
5843 ShuffleMessage::Barrier(barrier) => {
5844 let barrier = ReceivedShuffle {
5845 peer,
5846 message: ShuffleMessage::Barrier(barrier),
5847 reservation,
5848 sender_incarnation,
5849 receiver_incarnation,
5850 stream_id,
5851 assignment_version,
5852 assignment_digest: None,
5853 recovery_gen,
5854 checkpoint_sequence,
5855 };
5856 match self.holdover.stage_barrier(barrier) {
5857 Ok(true) => break,
5858 Ok(false) | Err(_) => self.holdover.release_items(1),
5859 }
5860 }
5861 }
5862 }
5863 }
5864 }
5865
5866 #[must_use]
5868 pub fn drain_checkpointed_data_for(&self, stage: &str) -> Vec<ReceivedBatch> {
5869 let mut staged = self.holdover.staged.lock();
5870 self.drain_inbound_into(&mut staged);
5871 let batches = staged.remove(stage).unwrap_or_default();
5872 self.holdover.release_items(batches.len());
5873 batches
5874 }
5875
5876 #[must_use]
5878 pub fn drain_staged_barriers(&self) -> Vec<ReceivedShuffle> {
5879 let barriers = self.holdover.take_staged_barriers();
5880 self.holdover.release_items(barriers.len());
5881 barriers
5882 }
5883
5884 #[must_use]
5886 pub fn drain_staged_frontiers(&self) -> Vec<ReceivedFrontierCut> {
5887 let frontiers = self.holdover.take_staged_frontiers();
5888 let item_count = frontiers.iter().map(ReceivedFrontierCut::item_count).sum();
5889 self.holdover.release_items(item_count);
5890 frontiers
5891 }
5892
5893 #[must_use]
5895 pub fn drain_all_staged(&self) -> Vec<(String, ReceivedBatch)> {
5896 let mut staged = self.holdover.staged.lock();
5897 let frontiers = self.holdover.take_staged_frontiers();
5898 let item_count = staged.values().map(Vec::len).sum::<usize>()
5899 + frontiers
5900 .iter()
5901 .map(ReceivedFrontierCut::item_count)
5902 .sum::<usize>();
5903 let mut drained: Vec<_> = staged
5904 .drain()
5905 .flat_map(|(stage, batches)| {
5906 batches
5907 .into_iter()
5908 .map(move |staged| (stage.clone(), staged))
5909 })
5910 .collect();
5911 for cut in frontiers {
5912 let stage = cut.stage().to_owned();
5913 drained.extend(
5914 cut.preceding
5915 .into_iter()
5916 .map(|batch| (stage.clone(), batch)),
5917 );
5918 }
5919 self.holdover.release_items(item_count);
5920 drained
5921 }
5922 }
5923
5924 struct RxReturnGuard<'a> {
5927 slot: &'a Mutex<Option<InboundRx>>,
5928 notify: &'a tokio::sync::Notify,
5929 rx: Option<InboundRx>,
5930 }
5931
5932 impl Drop for RxReturnGuard<'_> {
5933 fn drop(&mut self) {
5934 if let Some(rx) = self.rx.take() {
5935 *self.slot.lock() = Some(rx);
5936 self.notify.notify_one();
5937 }
5938 }
5939 }
5940}
5941
5942#[cfg(not(feature = "cluster"))]
5943pub use shim::{ShuffleReceiver, ShuffleSender};
5944
5945#[cfg(all(test, not(feature = "cluster")))]
5946mod shim_tests;
5947
5948#[cfg(all(test, feature = "cluster"))]
5949mod tests;