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