1use super::protocol::{
6 ANY_RANK, CollectiveAgreement, ElementType, FLAG_COUNTS_PREFIX, FLAG_P2P_CHANNEL, Frame,
7 FrameHeader, Opcode, OperationDescriptor, ProtocolError, UniqueId,
8};
9use super::{
10 CollectiveTopology, CollectiveTransport, TopologyAggregateLink, TopologyLink, TopologyRailLink,
11 peer::DirectPeerMesh,
12};
13use std::collections::{HashMap, VecDeque};
14use std::env;
15use std::error::Error;
16use std::fmt::{Display, Formatter};
17use std::io::{self, IoSlice, Read, Write};
18use std::net::{
19 IpAddr, Ipv4Addr, Ipv6Addr, Shutdown, SocketAddr, TcpListener, TcpStream, ToSocketAddrs,
20};
21use std::sync::mpsc::{self, Receiver, RecvTimeoutError, Sender};
22use std::sync::{Arc, Barrier, Condvar, Mutex};
23use std::thread::{self, JoinHandle};
24use std::time::{Duration, Instant};
25
26const DEFAULT_COLLECTIVE_TIMEOUT: Duration = Duration::from_secs(300);
27const COLLECTIVE_RESPONSE_GRACE: Duration = Duration::from_millis(250);
28const MAX_P2P_RAILS: usize = 64;
29const TOPOLOGY_PROBE_TAG_PREFIX: u64 = 0x475a_0000_0000_0000;
30const SERVER_FAILURE_POLL_INTERVAL: Duration = Duration::from_millis(10);
31
32#[derive(Debug, Clone, Copy, PartialEq, Eq)]
33pub struct TopologyProbeOptions {
34 pub payload_bytes: usize,
35 pub latency_iterations: usize,
36 pub bandwidth_iterations: usize,
37 pub warmup_iterations: usize,
38}
39
40impl Default for TopologyProbeOptions {
41 fn default() -> Self {
42 Self {
43 payload_bytes: 1024 * 1024,
44 latency_iterations: 16,
45 bandwidth_iterations: 3,
46 warmup_iterations: 1,
47 }
48 }
49}
50
51impl TopologyProbeOptions {
52 pub fn from_environment() -> Result<Self, NetworkError> {
53 let defaults = Self::default();
54 let options = Self {
55 payload_bytes: parse_probe_usize("GX1_TOPOLOGY_PROBE_BYTES", defaults.payload_bytes)?,
56 latency_iterations: parse_probe_usize(
57 "GX1_TOPOLOGY_PROBE_LATENCY_ITERATIONS",
58 defaults.latency_iterations,
59 )?,
60 bandwidth_iterations: parse_probe_usize(
61 "GX1_TOPOLOGY_PROBE_BANDWIDTH_ITERATIONS",
62 defaults.bandwidth_iterations,
63 )?,
64 warmup_iterations: parse_probe_usize(
65 "GX1_TOPOLOGY_PROBE_WARMUP_ITERATIONS",
66 defaults.warmup_iterations,
67 )?,
68 };
69 options.validate()?;
70 Ok(options)
71 }
72
73 fn validate(self) -> Result<(), NetworkError> {
74 if self.payload_bytes == 0 || self.payload_bytes > super::protocol::MAX_FRAME_PAYLOAD_BYTES
75 {
76 return Err(NetworkError::InvalidConfiguration(format!(
77 "topology probe payload {} is outside 1..={}",
78 self.payload_bytes,
79 super::protocol::MAX_FRAME_PAYLOAD_BYTES
80 )));
81 }
82 if self.latency_iterations == 0 || self.bandwidth_iterations == 0 {
83 return Err(NetworkError::InvalidConfiguration(
84 "topology probe latency and bandwidth iterations must be greater than zero".into(),
85 ));
86 }
87 if self.latency_iterations > 10_000
88 || self.bandwidth_iterations > 10_000
89 || self.warmup_iterations > 10_000
90 {
91 return Err(NetworkError::InvalidConfiguration(
92 "topology probe iteration counts must not exceed 10000".into(),
93 ));
94 }
95 Ok(())
96 }
97}
98
99#[derive(Debug, Clone, Copy, PartialEq, Eq)]
100struct LinkProbe {
101 bandwidth_mbps: u64,
102 latency_ns: u64,
103}
104
105type SharedServerFailure = Arc<(Mutex<Option<String>>, Condvar)>;
106
107fn publish_server_failure(failure: &SharedServerFailure, message: &str) {
108 let (state, ready) = &**failure;
109 if let Ok(mut state) = state.lock() {
110 state.get_or_insert_with(|| message.to_owned());
111 ready.notify_all();
112 }
113}
114
115fn server_failure_text(error: &NetworkError) -> String {
116 match error {
117 NetworkError::RemoteAbort(message) => message.clone(),
118 _ => error.to_string(),
119 }
120}
121
122fn server_failure_message(failure: &SharedServerFailure) -> Result<Option<String>, NetworkError> {
123 let (state, _) = &**failure;
124 Ok(state.lock().map_err(|_| NetworkError::Poisoned)?.clone())
125}
126
127fn wait_for_server_failure(
128 failure: &SharedServerFailure,
129 timeout: Duration,
130) -> Result<Option<String>, NetworkError> {
131 let (state, ready) = &**failure;
132 let state = state.lock().map_err(|_| NetworkError::Poisoned)?;
133 if state.is_some() {
134 return Ok(state.clone());
135 }
136 let (state, _) = ready
137 .wait_timeout(state, timeout)
138 .map_err(|_| NetworkError::Poisoned)?;
139 Ok(state.clone())
140}
141
142#[derive(Debug)]
143pub enum NetworkError {
144 Io(io::Error),
145 Protocol(ProtocolError),
146 InvalidConfiguration(String),
147 RankAlreadyJoined(u32),
148 WrongSession,
149 WrongDestination { expected: u32, actual: u32 },
150 RemoteAbort(String),
151 Timeout(String),
152 ChannelClosed,
153 Poisoned,
154}
155
156impl Display for NetworkError {
157 fn fmt(&self, formatter: &mut Formatter<'_>) -> std::fmt::Result {
158 match self {
159 Self::Io(error) => Display::fmt(error, formatter),
160 Self::Protocol(error) => Display::fmt(error, formatter),
161 Self::InvalidConfiguration(message) => formatter.write_str(message),
162 Self::RankAlreadyJoined(rank) => write!(formatter, "rank {rank} joined twice"),
163 Self::WrongSession => write!(formatter, "collective frame belongs to another session"),
164 Self::WrongDestination { expected, actual } => write!(
165 formatter,
166 "collective response destination {actual} does not match rank {expected}"
167 ),
168 Self::RemoteAbort(message) => {
169 write!(formatter, "collective coordinator aborted: {message}")
170 }
171 Self::Timeout(operation) => write!(formatter, "{operation} timed out"),
172 Self::ChannelClosed => write!(formatter, "collective response channel closed"),
173 Self::Poisoned => write!(formatter, "TCP collective session is poisoned"),
174 }
175 }
176}
177
178impl Error for NetworkError {
179 fn source(&self) -> Option<&(dyn Error + 'static)> {
180 match self {
181 Self::Io(error) => Some(error),
182 Self::Protocol(error) => Some(error),
183 _ => None,
184 }
185 }
186}
187
188impl From<io::Error> for NetworkError {
189 fn from(error: io::Error) -> Self {
190 Self::Io(error)
191 }
192}
193
194impl From<ProtocolError> for NetworkError {
195 fn from(error: ProtocolError) -> Self {
196 Self::Protocol(error)
197 }
198}
199
200#[derive(Debug)]
201pub struct TcpRendezvousServer {
202 listener: TcpListener,
203 unique_id: UniqueId,
204 world_size: usize,
205 p2p_rails: usize,
206 collective_timeout: Duration,
207 heartbeat_timeout: Option<Duration>,
208 transport: CollectiveTransport,
209}
210
211impl TcpRendezvousServer {
212 pub fn bind(
213 address: impl ToSocketAddrs,
214 unique_id: UniqueId,
215 world_size: usize,
216 ) -> Result<Self, NetworkError> {
217 if world_size == 0 || world_size > u32::MAX as usize {
218 return Err(NetworkError::InvalidConfiguration(format!(
219 "TCP collective world size {world_size} is outside 1..={} ",
220 u32::MAX
221 )));
222 }
223 let listener = TcpListener::bind(address)?;
224 let p2p_rails = p2p_rails_from_environment()?;
225 let transport = tcp_transport_from_environment()?;
226 Ok(Self {
227 listener,
228 unique_id,
229 world_size,
230 p2p_rails,
231 collective_timeout: DEFAULT_COLLECTIVE_TIMEOUT,
232 heartbeat_timeout: None,
233 transport,
234 })
235 }
236
237 pub fn with_collective_timeout(mut self, timeout: Duration) -> Result<Self, NetworkError> {
238 if timeout.is_zero() {
239 return Err(NetworkError::InvalidConfiguration(
240 "TCP collective timeout must be greater than zero".into(),
241 ));
242 }
243 self.collective_timeout = timeout;
244 Ok(self)
245 }
246
247 pub fn with_heartbeat_timeout(mut self, timeout: Duration) -> Result<Self, NetworkError> {
248 if timeout.is_zero() {
249 return Err(NetworkError::InvalidConfiguration(
250 "TCP heartbeat timeout must be greater than zero".into(),
251 ));
252 }
253 self.heartbeat_timeout = Some(timeout);
254 Ok(self)
255 }
256
257 pub fn with_p2p_rails(mut self, rails: usize) -> Result<Self, NetworkError> {
258 validate_p2p_rails(rails)?;
259 self.p2p_rails = rails;
260 Ok(self)
261 }
262
263 pub fn with_transport(mut self, transport: CollectiveTransport) -> Result<Self, NetworkError> {
264 validate_tcp_transport(transport)?;
265 self.transport = transport;
266 Ok(self)
267 }
268
269 pub fn local_addr(&self) -> Result<SocketAddr, NetworkError> {
270 Ok(self.listener.local_addr()?)
271 }
272
273 pub fn run(self) -> Result<(), NetworkError> {
276 let mut collective_streams =
277 accept_rank_channels(&self.listener, self.unique_id, self.world_size, false, 0)?;
278 let unique_id = self.unique_id;
279 let world_size = self.world_size;
280 let server_failure = Arc::new((Mutex::new(None), Condvar::new()));
281 let mut p2p_workers = Vec::with_capacity(self.p2p_rails);
282 match self.transport {
283 CollectiveTransport::TcpHostStaged => {
284 for rail in 0..self.p2p_rails {
285 let p2p_streams = accept_rank_channels(
286 &self.listener,
287 self.unique_id,
288 self.world_size,
289 true,
290 rail,
291 )?;
292 let heartbeat_timeout = (rail == 0).then_some(self.heartbeat_timeout).flatten();
293 p2p_workers.push(spawn_p2p_router(
294 p2p_streams,
295 unique_id,
296 world_size,
297 heartbeat_timeout,
298 rail,
299 Arc::clone(&server_failure),
300 )?);
301 }
302 }
303 CollectiveTransport::TcpPeer => {
304 for stream in &collective_streams {
305 stream.set_read_timeout(Some(self.collective_timeout))?;
306 }
307 exchange_peer_endpoints_server(
308 &mut collective_streams,
309 unique_id,
310 world_size,
311 self.p2p_rails,
312 )?;
313 for stream in &collective_streams {
314 stream.set_read_timeout(None)?;
315 }
316 let control_streams =
317 accept_rank_channels(&self.listener, self.unique_id, self.world_size, true, 0)?;
318 p2p_workers.push(spawn_p2p_router(
319 control_streams,
320 unique_id,
321 world_size,
322 self.heartbeat_timeout,
323 0,
324 Arc::clone(&server_failure),
325 )?);
326 }
327 CollectiveTransport::HostStaged
328 | CollectiveTransport::PciePeer
329 | CollectiveTransport::Rdma
330 | CollectiveTransport::GxLink => {
331 return Err(NetworkError::InvalidConfiguration(format!(
332 "{:?} is not a TCP transport",
333 self.transport
334 )));
335 }
336 }
337 let collective_result = run_collective_loop(
338 collective_streams,
339 self.unique_id,
340 self.world_size,
341 self.collective_timeout,
342 server_failure,
343 );
344 let mut p2p_result = Ok(());
345 for (rail, worker) in p2p_workers.into_iter().enumerate() {
346 let result = match worker.join() {
347 Ok(result) => result,
348 Err(_) => Err(NetworkError::InvalidConfiguration(format!(
349 "point-to-point rail {rail} router panicked"
350 ))),
351 };
352 if p2p_result.is_ok() {
353 p2p_result = result;
354 }
355 }
356 collective_result.and(p2p_result)
357 }
358}
359
360fn spawn_p2p_router(
361 streams: Vec<TcpStream>,
362 unique_id: UniqueId,
363 world_size: usize,
364 heartbeat_timeout: Option<Duration>,
365 rail: usize,
366 server_failure: SharedServerFailure,
367) -> Result<JoinHandle<Result<(), NetworkError>>, NetworkError> {
368 thread::Builder::new()
369 .name(format!("gx1-p2p-router-{rail}"))
370 .spawn(move || {
371 run_p2p_loop(
372 streams,
373 unique_id,
374 world_size,
375 heartbeat_timeout,
376 server_failure,
377 )
378 })
379 .map_err(|error| {
380 NetworkError::InvalidConfiguration(format!(
381 "cannot start point-to-point rail {rail} router: {error}"
382 ))
383 })
384}
385
386fn accept_rank_channels(
387 listener: &TcpListener,
388 unique_id: UniqueId,
389 world_size: usize,
390 p2p: bool,
391 rail: usize,
392) -> Result<Vec<TcpStream>, NetworkError> {
393 let mut ranks = (0..world_size)
394 .map(|_| None)
395 .collect::<Vec<Option<TcpStream>>>();
396 for _ in 0..world_size {
397 let (mut stream, _) = listener.accept()?;
398 stream.set_nodelay(true)?;
399 let frame = read_frame(&mut stream)?
400 .ok_or_else(|| NetworkError::InvalidConfiguration("rank closed before JOIN".into()))?;
401 validate_join(&frame, unique_id, world_size, p2p, rail)?;
402 let rank = frame.header.source_rank as usize;
403 if ranks[rank].is_some() {
404 return Err(NetworkError::RankAlreadyJoined(rank as u32));
405 }
406 ranks[rank] = Some(stream);
407 }
408 let mut streams = ranks
409 .into_iter()
410 .map(|stream| stream.expect("every rank joined"))
411 .collect::<Vec<_>>();
412 for (rank, stream) in streams.iter_mut().enumerate() {
413 let mut ready = control_frame(unique_id, Opcode::Ready, 0, rank as u32, world_size as u32)?;
414 if p2p {
415 ready.header.flags |= FLAG_P2P_CHANNEL;
416 ready.header.tag = rail as u64;
417 }
418 write_frame(stream, &ready)?;
419 }
420 Ok(streams)
421}
422
423fn exchange_peer_endpoints_server(
424 streams: &mut [TcpStream],
425 unique_id: UniqueId,
426 world_size: usize,
427 rails: usize,
428) -> Result<(), NetworkError> {
429 let mut endpoints = Vec::with_capacity(world_size);
430 let mut all_ranks_use_shared_endpoint = true;
431 for (rank, stream) in streams.iter_mut().enumerate() {
432 let frame = read_frame(stream)?.ok_or_else(|| {
433 NetworkError::InvalidConfiguration(format!(
434 "rank {rank} closed before publishing its direct peer endpoint"
435 ))
436 })?;
437 if frame.header.opcode != Opcode::PeerEndpoint
438 || frame.header.element_type != ElementType::U8
439 || frame.header.unique_id != unique_id
440 || frame.header.source_rank as usize != rank
441 || frame.header.destination_rank != ANY_RANK
442 || frame.header.root_rank != ANY_RANK
443 || frame.header.world_size as usize != world_size
444 || frame.header.sequence != 0
445 || frame.header.tag != rails as u64
446 || frame.header.flags != super::protocol::FLAG_PAYLOAD
447 || frame.header.element_count != frame.payload.len() as u64
448 || frame.payload.is_empty()
449 {
450 return Err(NetworkError::InvalidConfiguration(format!(
451 "rank {rank} published an invalid direct peer endpoint frame: {:?}, payload bytes {}",
452 frame.header,
453 frame.payload.len()
454 )));
455 }
456 let text = std::str::from_utf8(&frame.payload).map_err(|error| {
457 NetworkError::InvalidConfiguration(format!(
458 "rank {rank} direct peer endpoints are not UTF-8: {error}"
459 ))
460 })?;
461 let entries = text.lines().collect::<Vec<_>>();
462 if entries.len() != 1 && entries.len() != rails {
463 return Err(NetworkError::InvalidConfiguration(format!(
464 "rank {rank} published {} direct peer endpoints, expected one shared endpoint or {rails} rail endpoints",
465 entries.len()
466 )));
467 }
468 all_ranks_use_shared_endpoint &= entries.len() == 1;
469 let peer_ip = stream.peer_addr()?.ip();
470 let mut rank_endpoints = Vec::with_capacity(rails);
471 for (rail, entry) in entries.iter().enumerate() {
472 let mut endpoint = entry.parse::<SocketAddr>().map_err(|error| {
473 NetworkError::InvalidConfiguration(format!(
474 "rank {rank} direct peer endpoint {entry:?} for rail {rail} is invalid: {error}"
475 ))
476 })?;
477 if endpoint.port() == 0 {
478 return Err(NetworkError::InvalidConfiguration(format!(
479 "rank {rank} direct peer endpoint for rail {rail} uses port zero"
480 )));
481 }
482 if endpoint.ip().is_unspecified() {
483 endpoint.set_ip(peer_ip);
484 }
485 rank_endpoints.push(endpoint);
486 }
487 if rank_endpoints.len() == 1 {
488 rank_endpoints.resize(rails, rank_endpoints[0]);
489 }
490 endpoints.push(rank_endpoints);
491 }
492
493 let table = endpoints
497 .iter()
498 .flat_map(|rank_endpoints| {
499 if all_ranks_use_shared_endpoint {
500 &rank_endpoints[..1]
501 } else {
502 rank_endpoints.as_slice()
503 }
504 })
505 .map(SocketAddr::to_string)
506 .collect::<Vec<_>>()
507 .join("\n")
508 .into_bytes();
509 for (rank, stream) in streams.iter_mut().enumerate() {
510 let mut header = FrameHeader::collective(
511 unique_id,
512 Opcode::PeerEndpoint,
513 ElementType::U8,
514 0,
515 ANY_RANK,
516 world_size as u32,
517 0,
518 table.len() as u64,
519 );
520 header.destination_rank = rank as u32;
521 header.tag = rails as u64;
522 write_frame(stream, &Frame::new(header, table.clone())?)?;
523 }
524 Ok(())
525}
526
527#[derive(Debug, Clone, PartialEq, Eq)]
528struct PeerEndpointConfiguration {
529 listen_addresses: Vec<SocketAddr>,
530 advertise_addresses: Vec<IpAddr>,
531}
532
533impl PeerEndpointConfiguration {
534 fn from_environment(control_stream: &TcpStream, rails: usize) -> Result<Self, NetworkError> {
535 let default_listen = match control_stream.local_addr()?.ip() {
536 IpAddr::V4(_) => SocketAddr::new(IpAddr::V4(Ipv4Addr::UNSPECIFIED), 0),
537 IpAddr::V6(_) => SocketAddr::new(IpAddr::V6(Ipv6Addr::UNSPECIFIED), 0),
538 };
539 let listen_addresses = match optional_environment("GX1_P2P_RAIL_LISTEN_ADDRS")? {
540 Some(value) => parse_rail_socket_addresses("GX1_P2P_RAIL_LISTEN_ADDRS", &value, rails)?,
541 None => match optional_environment("GX1_P2P_LISTEN_ADDR")? {
542 Some(value) => vec![parse_socket_address("GX1_P2P_LISTEN_ADDR", &value)?],
543 None => vec![default_listen],
544 },
545 };
546 let advertise_addresses = match optional_environment("GX1_P2P_RAIL_ADVERTISE_ADDRS")? {
547 Some(value) => parse_rail_ip_addresses("GX1_P2P_RAIL_ADVERTISE_ADDRS", &value, rails)?,
548 None => match optional_environment("GX1_P2P_ADVERTISE_ADDR")? {
549 Some(value) => vec![parse_ip_address("GX1_P2P_ADVERTISE_ADDR", &value)?],
550 None => Vec::new(),
551 },
552 };
553 let configuration = Self {
554 listen_addresses,
555 advertise_addresses,
556 };
557 configuration.validate(rails)?;
558 Ok(configuration)
559 }
560
561 fn validate(&self, rails: usize) -> Result<(), NetworkError> {
562 if self.listen_addresses.len() != 1 && self.listen_addresses.len() != rails {
563 return Err(NetworkError::InvalidConfiguration(format!(
564 "direct peer listen configuration contains {} addresses, expected one shared address or {rails} rail addresses",
565 self.listen_addresses.len()
566 )));
567 }
568 if !self.advertise_addresses.is_empty()
569 && self.advertise_addresses.len() != 1
570 && self.advertise_addresses.len() != rails
571 {
572 return Err(NetworkError::InvalidConfiguration(format!(
573 "direct peer advertise configuration contains {} addresses, expected zero, one shared address, or {rails} rail addresses",
574 self.advertise_addresses.len()
575 )));
576 }
577 Ok(())
578 }
579}
580
581fn optional_environment(name: &'static str) -> Result<Option<String>, NetworkError> {
582 match env::var(name) {
583 Ok(value) => Ok(Some(value)),
584 Err(env::VarError::NotPresent) => Ok(None),
585 Err(env::VarError::NotUnicode(value)) => Err(NetworkError::InvalidConfiguration(format!(
586 "{name} is not Unicode: {:?}",
587 value.to_string_lossy()
588 ))),
589 }
590}
591
592fn parse_socket_address(name: &'static str, value: &str) -> Result<SocketAddr, NetworkError> {
593 value.trim().parse::<SocketAddr>().map_err(|error| {
594 NetworkError::InvalidConfiguration(format!(
595 "{name} must be a socket address, got {value:?}: {error}"
596 ))
597 })
598}
599
600fn parse_ip_address(name: &'static str, value: &str) -> Result<IpAddr, NetworkError> {
601 value.trim().parse::<IpAddr>().map_err(|error| {
602 NetworkError::InvalidConfiguration(format!(
603 "{name} must be an IP address, got {value:?}: {error}"
604 ))
605 })
606}
607
608fn parse_rail_socket_addresses(
609 name: &'static str,
610 value: &str,
611 rails: usize,
612) -> Result<Vec<SocketAddr>, NetworkError> {
613 let entries = value.split(',').collect::<Vec<_>>();
614 if entries.len() != rails {
615 return Err(NetworkError::InvalidConfiguration(format!(
616 "{name} contains {} addresses, expected {rails}",
617 entries.len()
618 )));
619 }
620 entries
621 .into_iter()
622 .enumerate()
623 .map(|(rail, entry)| {
624 parse_socket_address(name, entry).map_err(|error| {
625 NetworkError::InvalidConfiguration(format!("{error} (rail {rail})"))
626 })
627 })
628 .collect()
629}
630
631fn parse_rail_ip_addresses(
632 name: &'static str,
633 value: &str,
634 rails: usize,
635) -> Result<Vec<IpAddr>, NetworkError> {
636 let entries = value.split(',').collect::<Vec<_>>();
637 if entries.len() != rails {
638 return Err(NetworkError::InvalidConfiguration(format!(
639 "{name} contains {} addresses, expected {rails}",
640 entries.len()
641 )));
642 }
643 entries
644 .into_iter()
645 .enumerate()
646 .map(|(rail, entry)| {
647 parse_ip_address(name, entry).map_err(|error| {
648 NetworkError::InvalidConfiguration(format!("{error} (rail {rail})"))
649 })
650 })
651 .collect()
652}
653
654fn bind_peer_listeners(
655 control_stream: &TcpStream,
656 rails: usize,
657) -> Result<(Vec<TcpListener>, Vec<SocketAddr>), NetworkError> {
658 let configuration = PeerEndpointConfiguration::from_environment(control_stream, rails)?;
659 bind_peer_listeners_with_configuration(configuration, rails)
660}
661
662fn bind_peer_listeners_with_configuration(
663 configuration: PeerEndpointConfiguration,
664 rails: usize,
665) -> Result<(Vec<TcpListener>, Vec<SocketAddr>), NetworkError> {
666 configuration.validate(rails)?;
667 let mut listeners = Vec::with_capacity(configuration.listen_addresses.len());
668 let mut local_addresses = Vec::with_capacity(configuration.listen_addresses.len());
669 for (listener_index, bind_address) in configuration.listen_addresses.iter().enumerate() {
670 let listener = TcpListener::bind(bind_address).map_err(|error| {
671 NetworkError::InvalidConfiguration(format!(
672 "cannot bind direct peer listener {listener_index} to {bind_address}: {error}"
673 ))
674 })?;
675 local_addresses.push(listener.local_addr()?);
676 listeners.push(listener);
677 }
678
679 let mut endpoints = Vec::with_capacity(rails);
680 for rail in 0..rails {
681 let listener_index = if listeners.len() == 1 { 0 } else { rail };
682 let local = local_addresses[listener_index];
683 let advertised_ip = match configuration.advertise_addresses.len() {
684 0 => local.ip(),
685 1 => configuration.advertise_addresses[0],
686 _ => configuration.advertise_addresses[rail],
687 };
688 endpoints.push(SocketAddr::new(advertised_ip, local.port()));
689 }
690 Ok((listeners, endpoints))
691}
692
693fn exchange_peer_endpoints_client(
694 stream: &mut TcpStream,
695 rank_endpoints: &[SocketAddr],
696 unique_id: UniqueId,
697 rank: u32,
698 world_size: u32,
699 rails: usize,
700) -> Result<Vec<Vec<SocketAddr>>, NetworkError> {
701 if rank_endpoints.len() != rails {
702 return Err(NetworkError::InvalidConfiguration(format!(
703 "rank {rank} has {} direct peer endpoints, expected {rails}",
704 rank_endpoints.len()
705 )));
706 }
707 let shared_endpoint = rank_endpoints
708 .windows(2)
709 .all(|endpoints| endpoints[0] == endpoints[1]);
710 let payload = rank_endpoints[..if shared_endpoint { 1 } else { rails }]
711 .iter()
712 .map(SocketAddr::to_string)
713 .collect::<Vec<_>>()
714 .join("\n")
715 .into_bytes();
716 let mut header = FrameHeader::collective(
717 unique_id,
718 Opcode::PeerEndpoint,
719 ElementType::U8,
720 rank,
721 ANY_RANK,
722 world_size,
723 0,
724 payload.len() as u64,
725 );
726 header.tag = rails as u64;
727 write_frame(stream, &Frame::new(header, payload)?)?;
728 let response = read_frame(stream)?.ok_or_else(|| {
729 NetworkError::InvalidConfiguration(
730 "coordinator closed before returning direct peer endpoints".into(),
731 )
732 })?;
733 validate_response(&response, unique_id, rank, world_size, 0)?;
734 if response.header.opcode != Opcode::PeerEndpoint
735 || response.header.element_type != ElementType::U8
736 || response.header.tag != rails as u64
737 || response.header.element_count != response.payload.len() as u64
738 {
739 return Err(NetworkError::InvalidConfiguration(
740 "coordinator returned an invalid direct peer endpoint table".into(),
741 ));
742 }
743 let table = std::str::from_utf8(&response.payload).map_err(|error| {
744 NetworkError::InvalidConfiguration(format!(
745 "direct peer endpoint table is not UTF-8: {error}"
746 ))
747 })?;
748 let flat_endpoints = table
749 .lines()
750 .map(|entry| {
751 entry.parse::<SocketAddr>().map_err(|error| {
752 NetworkError::InvalidConfiguration(format!(
753 "direct peer endpoint {entry:?} is invalid: {error}"
754 ))
755 })
756 })
757 .collect::<Result<Vec<_>, _>>()?;
758 if flat_endpoints.len() == world_size as usize {
759 return Ok(flat_endpoints
760 .into_iter()
761 .map(|endpoint| vec![endpoint; rails])
762 .collect());
763 }
764 let expected = world_size as usize * rails;
765 if flat_endpoints.len() != expected {
766 return Err(NetworkError::InvalidConfiguration(format!(
767 "coordinator returned {} direct peer endpoints, expected {world_size} shared endpoints or {expected} rail endpoints",
768 flat_endpoints.len()
769 )));
770 }
771 Ok(flat_endpoints
772 .chunks_exact(rails)
773 .map(<[SocketAddr]>::to_vec)
774 .collect())
775}
776
777#[derive(Debug)]
778enum CollectiveEvent {
779 Frame { rank: usize, frame: Frame },
780 Closed { rank: usize },
781 Failed { rank: usize, message: String },
782}
783
784fn run_collective_loop(
785 mut streams: Vec<TcpStream>,
786 unique_id: UniqueId,
787 world_size: usize,
788 mut timeout: Duration,
789 server_failure: SharedServerFailure,
790) -> Result<(), NetworkError> {
791 let (sender, receiver) = mpsc::channel::<CollectiveEvent>();
792 let mut readers = Vec::with_capacity(world_size);
793 for (rank, stream) in streams.iter().enumerate() {
794 let mut stream = stream.try_clone()?;
795 stream.set_read_timeout(None)?;
796 let sender = sender.clone();
797 readers.push(
798 thread::Builder::new()
799 .name(format!("gx1-collective-server-rank-{rank}"))
800 .spawn(move || {
801 loop {
802 let event = match read_frame(&mut stream) {
803 Ok(Some(frame)) => CollectiveEvent::Frame { rank, frame },
804 Ok(None) => CollectiveEvent::Closed { rank },
805 Err(error) => CollectiveEvent::Failed {
806 rank,
807 message: error.to_string(),
808 },
809 };
810 let terminal = !matches!(event, CollectiveEvent::Frame { .. });
811 if sender.send(event).is_err() || terminal {
812 return;
813 }
814 }
815 })
816 .map_err(|error| {
817 NetworkError::InvalidConfiguration(format!(
818 "cannot start collective rank reader: {error}"
819 ))
820 })?,
821 );
822 }
823 drop(sender);
824 let mut agreement = CollectiveAgreement::new(world_size)?;
825 let mut active_sequence = agreement.next_sequence();
826 let result = 'session: loop {
827 if let Some(message) = server_failure_message(&server_failure)? {
828 break Err(NetworkError::InvalidConfiguration(message));
829 }
830 let sequence = agreement.next_sequence();
831 let deadline = Instant::now() + timeout;
832 let mut requests = (0..world_size)
833 .map(|_| None)
834 .collect::<Vec<Option<Frame>>>();
835 while requests.iter().any(Option::is_none) {
836 if let Some(message) = server_failure_message(&server_failure)? {
837 break 'session Err(NetworkError::InvalidConfiguration(message));
838 }
839 let remaining = deadline.saturating_duration_since(Instant::now());
840 if remaining.is_zero() {
841 break 'session Err(collective_timeout_error(sequence, &requests));
842 }
843 match receiver.recv_timeout(remaining.min(SERVER_FAILURE_POLL_INTERVAL)) {
844 Ok(CollectiveEvent::Frame { rank, frame }) => {
845 if requests[rank].replace(frame).is_some() {
846 break 'session Err(NetworkError::InvalidConfiguration(format!(
847 "rank {rank} submitted collective sequence {sequence} twice"
848 )));
849 }
850 }
851 Ok(CollectiveEvent::Closed { rank }) => {
852 if requests.iter().all(Option::is_none) {
853 break 'session Ok(());
854 }
855 if let Some(message) =
856 wait_for_server_failure(&server_failure, SERVER_FAILURE_POLL_INTERVAL)?
857 {
858 break 'session Err(NetworkError::InvalidConfiguration(message));
859 }
860 break 'session Err(NetworkError::InvalidConfiguration(format!(
861 "rank {rank} disconnected during collective sequence {sequence}"
862 )));
863 }
864 Ok(CollectiveEvent::Failed { rank, message }) => {
865 if let Some(message) =
866 wait_for_server_failure(&server_failure, SERVER_FAILURE_POLL_INTERVAL)?
867 {
868 break 'session Err(NetworkError::InvalidConfiguration(message));
869 }
870 break 'session Err(NetworkError::InvalidConfiguration(format!(
871 "rank {rank} failed during collective sequence {sequence}: {message}"
872 )));
873 }
874 Err(RecvTimeoutError::Timeout) => continue,
875 Err(RecvTimeoutError::Disconnected) => {
876 break 'session Err(NetworkError::ChannelClosed);
877 }
878 }
879 }
880 let requests = requests
881 .into_iter()
882 .map(|request| request.expect("every collective rank submitted"))
883 .collect::<Vec<_>>();
884 let responses = match validate_and_route(&mut agreement, unique_id, sequence, &requests) {
885 Ok(responses) => responses,
886 Err(error) => break Err(error),
887 };
888 for (stream, response) in streams.iter_mut().zip(&responses) {
889 if let Err(error) = write_frame(stream, response) {
890 break 'session Err(error);
891 }
892 }
893 if requests[0].header.opcode == Opcode::SetTimeout {
894 timeout = Duration::from_millis(requests[0].header.tag);
895 }
896 active_sequence = agreement.next_sequence();
897 };
898 if let Err(error) = &result {
899 broadcast_collective_abort(
900 &mut streams,
901 unique_id,
902 active_sequence,
903 world_size,
904 &error.to_string(),
905 );
906 }
907 for stream in &streams {
908 let _ = stream.shutdown(Shutdown::Both);
909 }
910 for reader in readers {
911 let _ = reader.join();
912 }
913 result
914}
915
916fn collective_timeout_error(sequence: u64, requests: &[Option<Frame>]) -> NetworkError {
917 let missing = requests
918 .iter()
919 .enumerate()
920 .filter_map(|(rank, request)| request.is_none().then_some(rank.to_string()))
921 .collect::<Vec<_>>()
922 .join(",");
923 NetworkError::Timeout(format!(
924 "collective sequence {sequence}; missing ranks [{missing}]"
925 ))
926}
927
928fn broadcast_collective_abort(
929 streams: &mut [TcpStream],
930 unique_id: UniqueId,
931 sequence: u64,
932 world_size: usize,
933 message: &str,
934) {
935 for (rank, stream) in streams.iter_mut().enumerate() {
936 if let Ok(frame) = abort_frame(unique_id, sequence, rank as u32, world_size as u32, message)
937 {
938 let _ = write_frame(stream, &frame);
939 }
940 }
941}
942
943#[derive(Debug)]
944struct TcpSessionInner {
945 stream: TcpStream,
946 next_sequence: u64,
947}
948
949#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
950struct P2pResponseKey {
951 opcode: Opcode,
952 sequence: u64,
953}
954
955type PendingP2pResponse = Sender<Result<Frame, String>>;
956
957#[derive(Debug)]
958struct P2pSubmission {
959 stream: TcpStream,
960 next_send_sequence: u64,
961 next_receive_sequence: u64,
962 next_heartbeat_sequence: u64,
963}
964
965#[derive(Debug)]
966struct P2pClient {
967 submission: Mutex<P2pSubmission>,
968 pending: Arc<Mutex<HashMap<P2pResponseKey, PendingP2pResponse>>>,
969 failure: Arc<Mutex<Option<String>>>,
970 reader: Mutex<Option<JoinHandle<()>>>,
971 unique_id: UniqueId,
972 rank: u32,
973 world_size: u32,
974 timeout: Mutex<Duration>,
975}
976
977#[derive(Debug)]
978enum P2pDataPlane {
979 Coordinator(Vec<P2pClient>),
980 Peer {
981 mesh: Box<DirectPeerMesh>,
982 control: P2pClient,
983 },
984}
985
986#[derive(Debug)]
987pub struct TcpRankSession {
988 unique_id: UniqueId,
989 rank: u32,
990 world_size: u32,
991 transport: CollectiveTransport,
992 inner: Mutex<TcpSessionInner>,
993 p2p: P2pDataPlane,
994}
995
996#[derive(Debug, Clone, Copy, PartialEq, Eq)]
997pub struct ExchangeOptions {
998 pub root_rank: u32,
999 pub element_count: u64,
1000 pub flags: u16,
1001 pub tag: u64,
1002}
1003
1004impl ExchangeOptions {
1005 pub const fn new(root_rank: u32, element_count: u64) -> Self {
1006 Self {
1007 root_rank,
1008 element_count,
1009 flags: 0,
1010 tag: 0,
1011 }
1012 }
1013}
1014
1015impl P2pClient {
1016 fn new(
1017 stream: TcpStream,
1018 unique_id: UniqueId,
1019 rank: u32,
1020 world_size: u32,
1021 timeout: Duration,
1022 rail: usize,
1023 failure: Arc<Mutex<Option<String>>>,
1024 ) -> Result<Self, NetworkError> {
1025 stream.set_write_timeout(Some(timeout))?;
1026 let mut read_stream = stream.try_clone()?;
1027 read_stream.set_read_timeout(None)?;
1028 let pending = Arc::new(Mutex::new(HashMap::new()));
1029 let reader_pending = Arc::clone(&pending);
1030 let reader_failure = Arc::clone(&failure);
1031 let reader = thread::Builder::new()
1032 .name(format!("gx1-p2p-rank-{rank}-rail-{rail}"))
1033 .spawn(move || {
1034 run_p2p_client_reader(
1035 &mut read_stream,
1036 unique_id,
1037 rank,
1038 world_size,
1039 &reader_pending,
1040 &reader_failure,
1041 );
1042 })
1043 .map_err(|error| {
1044 NetworkError::InvalidConfiguration(format!(
1045 "cannot start point-to-point response reader: {error}"
1046 ))
1047 })?;
1048 Ok(Self {
1049 submission: Mutex::new(P2pSubmission {
1050 stream,
1051 next_send_sequence: 0,
1052 next_receive_sequence: 0,
1053 next_heartbeat_sequence: 0,
1054 }),
1055 pending,
1056 failure,
1057 reader: Mutex::new(Some(reader)),
1058 unique_id,
1059 rank,
1060 world_size,
1061 timeout: Mutex::new(timeout),
1062 })
1063 }
1064
1065 #[allow(clippy::too_many_arguments)]
1066 fn request(
1067 &self,
1068 unique_id: UniqueId,
1069 rank: u32,
1070 world_size: u32,
1071 opcode: Opcode,
1072 peer: u32,
1073 tag: u64,
1074 element_type: ElementType,
1075 element_count: u64,
1076 payload: Vec<u8>,
1077 ) -> Result<Frame, NetworkError> {
1078 let timeout = *self.timeout.lock().map_err(|_| NetworkError::Poisoned)?;
1079 self.request_with_timeout(
1080 unique_id,
1081 rank,
1082 world_size,
1083 opcode,
1084 peer,
1085 tag,
1086 element_type,
1087 element_count,
1088 payload,
1089 timeout,
1090 )
1091 }
1092
1093 fn set_timeout(&self, timeout: Duration) -> Result<(), NetworkError> {
1094 if timeout.is_zero() {
1095 return Err(NetworkError::InvalidConfiguration(
1096 "point-to-point timeout must be greater than zero".into(),
1097 ));
1098 }
1099 self.submission
1100 .lock()
1101 .map_err(|_| NetworkError::Poisoned)?
1102 .stream
1103 .set_write_timeout(Some(timeout))?;
1104 *self.timeout.lock().map_err(|_| NetworkError::Poisoned)? = timeout;
1105 Ok(())
1106 }
1107
1108 fn abort(&self, message: &str) -> Result<(), NetworkError> {
1109 let mut submission = self.submission.lock().map_err(|_| NetworkError::Poisoned)?;
1110 let sequence = submission.next_send_sequence;
1111 submission.next_send_sequence = sequence
1112 .checked_add(1)
1113 .ok_or(ProtocolError::SequenceOverflow)?;
1114 let payload = message.as_bytes().to_vec();
1115 let mut header = FrameHeader::collective(
1116 self.unique_id,
1117 Opcode::Abort,
1118 ElementType::U8,
1119 self.rank,
1120 ANY_RANK,
1121 self.world_size,
1122 sequence,
1123 payload.len() as u64,
1124 );
1125 header.destination_rank = ANY_RANK;
1126 write_frame(&mut submission.stream, &Frame::new(header, payload)?)
1127 }
1128
1129 #[allow(clippy::too_many_arguments)]
1130 fn request_with_timeout(
1131 &self,
1132 unique_id: UniqueId,
1133 rank: u32,
1134 world_size: u32,
1135 opcode: Opcode,
1136 peer: u32,
1137 tag: u64,
1138 element_type: ElementType,
1139 element_count: u64,
1140 payload: Vec<u8>,
1141 response_timeout: Duration,
1142 ) -> Result<Frame, NetworkError> {
1143 let (response_sender, response_receiver) = mpsc::channel();
1144 let key;
1145 {
1146 let mut submission = self.submission.lock().map_err(|_| NetworkError::Poisoned)?;
1147 let failure = self.failure.lock().map_err(|_| NetworkError::Poisoned)?;
1148 if let Some(message) = failure.as_deref() {
1149 return Err(NetworkError::RemoteAbort(message.to_owned()));
1150 }
1151 let sequence = match opcode {
1152 Opcode::Send => &mut submission.next_send_sequence,
1153 Opcode::Receive => &mut submission.next_receive_sequence,
1154 Opcode::Heartbeat => &mut submission.next_heartbeat_sequence,
1155 _ => {
1156 return Err(NetworkError::InvalidConfiguration(format!(
1157 "{opcode:?} is not a point-to-point request"
1158 )));
1159 }
1160 };
1161 let current_sequence = *sequence;
1162 *sequence = sequence
1163 .checked_add(1)
1164 .ok_or(ProtocolError::SequenceOverflow)?;
1165 key = P2pResponseKey {
1166 opcode,
1167 sequence: current_sequence,
1168 };
1169 drop(failure);
1170 let mut header = FrameHeader::collective(
1171 unique_id,
1172 opcode,
1173 element_type,
1174 rank,
1175 ANY_RANK,
1176 world_size,
1177 current_sequence,
1178 element_count,
1179 );
1180 header.destination_rank = peer;
1181 header.tag = tag;
1182 let frame = Frame::new(header, payload)?;
1183 self.pending
1184 .lock()
1185 .map_err(|_| NetworkError::Poisoned)?
1186 .insert(key, response_sender);
1187 if let Err(error) = write_frame(&mut submission.stream, &frame) {
1188 fail_pending(&self.pending, &self.failure, &error.to_string());
1189 return Err(error);
1190 }
1191 }
1192 match response_receiver.recv_timeout(response_timeout) {
1193 Ok(Ok(frame)) => Ok(frame),
1194 Ok(Err(message)) => Err(NetworkError::RemoteAbort(message)),
1195 Err(RecvTimeoutError::Timeout) => {
1196 if let Ok(mut pending) = self.pending.lock() {
1197 pending.remove(&key);
1198 }
1199 if let Ok(submission) = self.submission.lock() {
1200 let _ = submission.stream.shutdown(Shutdown::Both);
1201 }
1202 Err(NetworkError::Timeout(format!(
1203 "point-to-point {opcode:?} with tag {tag}"
1204 )))
1205 }
1206 Err(RecvTimeoutError::Disconnected) => Err(NetworkError::ChannelClosed),
1207 }
1208 }
1209}
1210
1211impl Drop for P2pClient {
1212 fn drop(&mut self) {
1213 if let Ok(mut submission) = self.submission.lock() {
1214 let header = FrameHeader::collective(
1215 self.unique_id,
1216 Opcode::Leave,
1217 ElementType::None,
1218 self.rank,
1219 ANY_RANK,
1220 self.world_size,
1221 0,
1222 0,
1223 );
1224 if let Ok(frame) = Frame::new(header, Vec::new()) {
1225 let _ = write_frame(&mut submission.stream, &frame);
1226 }
1227 let _ = submission.stream.shutdown(Shutdown::Both);
1228 }
1229 if let Ok(reader) = self.reader.get_mut()
1230 && let Some(reader) = reader.take()
1231 {
1232 let _ = reader.join();
1233 }
1234 }
1235}
1236
1237impl TcpRankSession {
1238 pub fn connect(
1239 address: impl ToSocketAddrs,
1240 unique_id: UniqueId,
1241 rank: u32,
1242 world_size: u32,
1243 timeout: Duration,
1244 ) -> Result<Self, NetworkError> {
1245 Self::connect_with_transport(
1246 address,
1247 unique_id,
1248 rank,
1249 world_size,
1250 timeout,
1251 p2p_rails_from_environment()?,
1252 tcp_transport_from_environment()?,
1253 )
1254 }
1255
1256 pub fn connect_with_p2p_rails(
1257 address: impl ToSocketAddrs,
1258 unique_id: UniqueId,
1259 rank: u32,
1260 world_size: u32,
1261 timeout: Duration,
1262 p2p_rails: usize,
1263 ) -> Result<Self, NetworkError> {
1264 Self::connect_with_transport(
1265 address,
1266 unique_id,
1267 rank,
1268 world_size,
1269 timeout,
1270 p2p_rails,
1271 tcp_transport_from_environment()?,
1272 )
1273 }
1274
1275 #[allow(clippy::too_many_arguments)]
1276 pub fn connect_with_transport(
1277 address: impl ToSocketAddrs,
1278 unique_id: UniqueId,
1279 rank: u32,
1280 world_size: u32,
1281 timeout: Duration,
1282 p2p_rails: usize,
1283 transport: CollectiveTransport,
1284 ) -> Result<Self, NetworkError> {
1285 Self::connect_with_transport_configuration(
1286 address, unique_id, rank, world_size, timeout, p2p_rails, transport, None,
1287 )
1288 }
1289
1290 #[allow(clippy::too_many_arguments)]
1291 fn connect_with_transport_configuration(
1292 address: impl ToSocketAddrs,
1293 unique_id: UniqueId,
1294 rank: u32,
1295 world_size: u32,
1296 timeout: Duration,
1297 p2p_rails: usize,
1298 transport: CollectiveTransport,
1299 peer_endpoint_configuration: Option<PeerEndpointConfiguration>,
1300 ) -> Result<Self, NetworkError> {
1301 if world_size == 0 || rank >= world_size {
1302 return Err(NetworkError::InvalidConfiguration(format!(
1303 "rank {rank} is outside TCP collective world size {world_size}"
1304 )));
1305 }
1306 validate_p2p_rails(p2p_rails)?;
1307 validate_tcp_transport(transport)?;
1308 let addresses = address.to_socket_addrs()?.collect::<Vec<_>>();
1309 let mut stream =
1310 connect_rank_channel(&addresses, unique_id, rank, world_size, timeout, false, 0)?;
1311 let failure = Arc::new(Mutex::new(None));
1312 let p2p = match transport {
1313 CollectiveTransport::TcpHostStaged => {
1314 let mut clients = Vec::with_capacity(p2p_rails);
1315 for rail in 0..p2p_rails {
1316 let p2p_stream = connect_rank_channel(
1317 &addresses, unique_id, rank, world_size, timeout, true, rail,
1318 )?;
1319 clients.push(P2pClient::new(
1320 p2p_stream,
1321 unique_id,
1322 rank,
1323 world_size,
1324 timeout,
1325 rail,
1326 Arc::clone(&failure),
1327 )?);
1328 }
1329 P2pDataPlane::Coordinator(clients)
1330 }
1331 CollectiveTransport::TcpPeer => {
1332 let (listeners, rank_endpoints) = match peer_endpoint_configuration {
1333 Some(configuration) => {
1334 bind_peer_listeners_with_configuration(configuration, p2p_rails)?
1335 }
1336 None => bind_peer_listeners(&stream, p2p_rails)?,
1337 };
1338 let endpoints = exchange_peer_endpoints_client(
1339 &mut stream,
1340 &rank_endpoints,
1341 unique_id,
1342 rank,
1343 world_size,
1344 p2p_rails,
1345 )?;
1346 let control_stream = connect_rank_channel(
1347 &addresses, unique_id, rank, world_size, timeout, true, 0,
1348 )?;
1349 let control = P2pClient::new(
1350 control_stream,
1351 unique_id,
1352 rank,
1353 world_size,
1354 timeout,
1355 0,
1356 Arc::clone(&failure),
1357 )?;
1358 let mesh = match DirectPeerMesh::connect(
1359 listeners, &endpoints, unique_id, rank, world_size, p2p_rails, timeout, failure,
1360 ) {
1361 Ok(mesh) => Box::new(mesh),
1362 Err(error) => {
1363 let _ = control.abort(&format!(
1364 "rank {rank} direct peer mesh setup failed: {error}"
1365 ));
1366 return Err(error);
1367 }
1368 };
1369 P2pDataPlane::Peer { control, mesh }
1370 }
1371 CollectiveTransport::HostStaged
1372 | CollectiveTransport::PciePeer
1373 | CollectiveTransport::Rdma
1374 | CollectiveTransport::GxLink => unreachable!("validated TCP transport"),
1375 };
1376 stream.set_read_timeout(Some(timeout.saturating_add(COLLECTIVE_RESPONSE_GRACE)))?;
1377 stream.set_write_timeout(Some(timeout))?;
1378 Ok(Self {
1379 unique_id,
1380 rank,
1381 world_size,
1382 transport,
1383 inner: Mutex::new(TcpSessionInner {
1384 stream,
1385 next_sequence: 0,
1386 }),
1387 p2p,
1388 })
1389 }
1390
1391 pub const fn rank(&self) -> u32 {
1392 self.rank
1393 }
1394
1395 pub const fn world_size(&self) -> u32 {
1396 self.world_size
1397 }
1398
1399 pub const fn transport(&self) -> CollectiveTransport {
1400 self.transport
1401 }
1402
1403 pub fn p2p_rails(&self) -> usize {
1404 match &self.p2p {
1405 P2pDataPlane::Coordinator(rails) => rails.len(),
1406 P2pDataPlane::Peer { mesh, .. } => mesh.rails(),
1407 }
1408 }
1409
1410 fn p2p_rail(&self, rail: usize) -> Result<&P2pClient, NetworkError> {
1411 let P2pDataPlane::Coordinator(rails) = &self.p2p else {
1412 return Err(NetworkError::InvalidConfiguration(
1413 "coordinator P2P rail requested for a direct peer transport".into(),
1414 ));
1415 };
1416 rails.get(rail).ok_or_else(|| {
1417 NetworkError::InvalidConfiguration(format!(
1418 "point-to-point rail {rail} is outside 0..{}",
1419 rails.len()
1420 ))
1421 })
1422 }
1423
1424 fn control(&self) -> Result<&P2pClient, NetworkError> {
1425 match &self.p2p {
1426 P2pDataPlane::Coordinator(rails) => rails.first().ok_or_else(|| {
1427 NetworkError::InvalidConfiguration("point-to-point control rail is absent".into())
1428 }),
1429 P2pDataPlane::Peer { control, .. } => Ok(control),
1430 }
1431 }
1432
1433 pub fn probe_topology_from_environment(&self) -> Result<CollectiveTopology, NetworkError> {
1434 self.probe_topology(TopologyProbeOptions::from_environment()?)
1435 }
1436
1437 pub fn probe_topology(
1438 &self,
1439 options: TopologyProbeOptions,
1440 ) -> Result<CollectiveTopology, NetworkError> {
1441 options.validate()?;
1442 if self.transport != CollectiveTransport::TcpPeer {
1443 return Err(NetworkError::InvalidConfiguration(
1444 "automatic link probing requires GX1_COLLECTIVE_TRANSPORT=tcp_peer".into(),
1445 ));
1446 }
1447 let mut topology = CollectiveTopology::empty(self.world_size).map_err(|error| {
1448 NetworkError::InvalidConfiguration(format!(
1449 "cannot create probed collective topology: {error}"
1450 ))
1451 })?;
1452 self.barrier()?;
1453 for first_rank in 0..self.world_size {
1454 for second_rank in first_rank + 1..self.world_size {
1455 let mut pair_bandwidth_mbps = u64::MAX;
1456 let mut pair_latency_ns = 0_u64;
1457 let mut pair_rail_bandwidth_sum_mbps = 0_u64;
1458 for rail in 0..self.p2p_rails() {
1459 let forward = self.probe_direction(
1460 first_rank,
1461 second_rank,
1462 first_rank,
1463 second_rank,
1464 rail,
1465 0,
1466 options,
1467 )?;
1468 let reverse = self.probe_direction(
1469 first_rank,
1470 second_rank,
1471 second_rank,
1472 first_rank,
1473 rail,
1474 1,
1475 options,
1476 )?;
1477 let mut local = [0_u64; 4];
1478 if self.rank == first_rank {
1479 let probe = forward.expect("forward sender owns its link measurement");
1480 local[0] = probe.bandwidth_mbps;
1481 local[1] = probe.latency_ns;
1482 }
1483 if self.rank == second_rank {
1484 let probe = reverse.expect("reverse sender owns its link measurement");
1485 local[2] = probe.bandwidth_mbps;
1486 local[3] = probe.latency_ns;
1487 }
1488 let payload = local
1489 .into_iter()
1490 .flat_map(u64::to_le_bytes)
1491 .collect::<Vec<_>>();
1492 let gathered =
1493 self.exchange(Opcode::AllGather, ElementType::U64, ANY_RANK, 4, payload)?;
1494 let values = gathered
1495 .payload
1496 .chunks_exact(8)
1497 .map(|bytes| u64::from_le_bytes(bytes.try_into().unwrap()))
1498 .collect::<Vec<_>>();
1499 let expected_values = self.world_size as usize * 4;
1500 if values.len() != expected_values {
1501 return Err(NetworkError::InvalidConfiguration(format!(
1502 "topology probe gathered {} metrics, expected {expected_values}",
1503 values.len()
1504 )));
1505 }
1506 let forward_offset = first_rank as usize * 4;
1507 let reverse_offset = second_rank as usize * 4;
1508 let forward = LinkProbe {
1509 bandwidth_mbps: values[forward_offset],
1510 latency_ns: values[forward_offset + 1],
1511 };
1512 let reverse = LinkProbe {
1513 bandwidth_mbps: values[reverse_offset + 2],
1514 latency_ns: values[reverse_offset + 3],
1515 };
1516 if forward.bandwidth_mbps == 0 || reverse.bandwidth_mbps == 0 {
1517 return Err(NetworkError::InvalidConfiguration(format!(
1518 "topology probe for ranks {first_rank}-{second_rank} rail {rail} returned zero bandwidth"
1519 )));
1520 }
1521 let bandwidth_mbps = conservative_probe_bandwidth_bucket(
1522 forward.bandwidth_mbps.min(reverse.bandwidth_mbps),
1523 );
1524 let latency_ns = conservative_probe_latency_bucket(
1525 forward.latency_ns.max(reverse.latency_ns),
1526 );
1527 topology
1528 .add_rail_link(TopologyRailLink {
1529 rail,
1530 first_rank,
1531 second_rank,
1532 bandwidth_mbps,
1533 latency_ns,
1534 })
1535 .map_err(|error| {
1536 NetworkError::InvalidConfiguration(format!(
1537 "cannot add probed topology rail link {rail}@{first_rank}-{second_rank}: {error}"
1538 ))
1539 })?;
1540 pair_bandwidth_mbps = pair_bandwidth_mbps.min(bandwidth_mbps);
1541 pair_rail_bandwidth_sum_mbps =
1542 pair_rail_bandwidth_sum_mbps.saturating_add(bandwidth_mbps);
1543 pair_latency_ns = pair_latency_ns.max(latency_ns);
1544 }
1545 if self.p2p_rails() > 1 {
1546 let forward = self.probe_concurrent_rails_direction(
1547 first_rank,
1548 second_rank,
1549 first_rank,
1550 second_rank,
1551 0,
1552 options,
1553 )?;
1554 let reverse = self.probe_concurrent_rails_direction(
1555 first_rank,
1556 second_rank,
1557 second_rank,
1558 first_rank,
1559 1,
1560 options,
1561 )?;
1562 let mut local = [0_u64; 2];
1563 if self.rank == first_rank {
1564 local[0] =
1565 forward.expect("forward sender owns its aggregate link measurement");
1566 }
1567 if self.rank == second_rank {
1568 local[1] =
1569 reverse.expect("reverse sender owns its aggregate link measurement");
1570 }
1571 let payload = local
1572 .into_iter()
1573 .flat_map(u64::to_le_bytes)
1574 .collect::<Vec<_>>();
1575 let gathered =
1576 self.exchange(Opcode::AllGather, ElementType::U64, ANY_RANK, 2, payload)?;
1577 let values = gathered
1578 .payload
1579 .chunks_exact(8)
1580 .map(|bytes| u64::from_le_bytes(bytes.try_into().unwrap()))
1581 .collect::<Vec<_>>();
1582 let expected_values = self.world_size as usize * 2;
1583 if values.len() != expected_values {
1584 return Err(NetworkError::InvalidConfiguration(format!(
1585 "topology aggregate probe gathered {} metrics, expected {expected_values}",
1586 values.len()
1587 )));
1588 }
1589 let forward = values[first_rank as usize * 2];
1590 let reverse = values[second_rank as usize * 2 + 1];
1591 if forward == 0 || reverse == 0 {
1592 return Err(NetworkError::InvalidConfiguration(format!(
1593 "topology aggregate probe for ranks {first_rank}-{second_rank} returned zero bandwidth"
1594 )));
1595 }
1596 let bandwidth_mbps = conservative_probe_bandwidth_bucket(forward.min(reverse))
1597 .min(pair_rail_bandwidth_sum_mbps)
1598 .max(1);
1599 topology
1600 .add_aggregate_link(TopologyAggregateLink {
1601 first_rank,
1602 second_rank,
1603 bandwidth_mbps,
1604 })
1605 .map_err(|error| {
1606 NetworkError::InvalidConfiguration(format!(
1607 "cannot add probed topology aggregate link {first_rank}-{second_rank}: {error}"
1608 ))
1609 })?;
1610 }
1611 topology
1612 .add_link(TopologyLink {
1613 first_rank,
1614 second_rank,
1615 bandwidth_mbps: pair_bandwidth_mbps,
1616 latency_ns: pair_latency_ns,
1617 })
1618 .map_err(|error| {
1619 NetworkError::InvalidConfiguration(format!(
1620 "cannot add probed topology link {first_rank}-{second_rank}: {error}"
1621 ))
1622 })?;
1623 }
1624 }
1625 Ok(topology)
1626 }
1627
1628 #[allow(clippy::too_many_arguments)]
1629 fn probe_direction(
1630 &self,
1631 first_rank: u32,
1632 second_rank: u32,
1633 sender_rank: u32,
1634 receiver_rank: u32,
1635 rail: usize,
1636 direction: u64,
1637 options: TopologyProbeOptions,
1638 ) -> Result<Option<LinkProbe>, NetworkError> {
1639 let latency_tag = topology_probe_tag(first_rank, second_rank, rail, direction, 0);
1640 let bandwidth_tag = topology_probe_tag(first_rank, second_rank, rail, direction, 1);
1641 let mut latency_samples = Vec::with_capacity(options.latency_iterations);
1642 let mut bandwidth_samples = Vec::with_capacity(options.bandwidth_iterations);
1643 if self.rank == sender_rank {
1644 for _ in 0..options.warmup_iterations {
1645 self.send_on_rail(
1646 rail,
1647 receiver_rank,
1648 latency_tag,
1649 ElementType::U8,
1650 1,
1651 vec![0],
1652 )?;
1653 }
1654 for _ in 0..options.latency_iterations {
1655 let started = Instant::now();
1656 self.send_on_rail(
1657 rail,
1658 receiver_rank,
1659 latency_tag,
1660 ElementType::U8,
1661 1,
1662 vec![0],
1663 )?;
1664 latency_samples.push(duration_ns(started.elapsed()));
1665 }
1666 let probe_payload = vec![0xa5; options.payload_bytes];
1667 for _ in 0..options.warmup_iterations {
1668 self.send_on_rail(
1669 rail,
1670 receiver_rank,
1671 bandwidth_tag,
1672 ElementType::U8,
1673 options.payload_bytes as u64,
1674 probe_payload.clone(),
1675 )?;
1676 }
1677 for _ in 0..options.bandwidth_iterations {
1678 let payload = probe_payload.clone();
1679 let started = Instant::now();
1680 self.send_on_rail(
1681 rail,
1682 receiver_rank,
1683 bandwidth_tag,
1684 ElementType::U8,
1685 options.payload_bytes as u64,
1686 payload,
1687 )?;
1688 bandwidth_samples.push(duration_ns(started.elapsed()));
1689 }
1690 } else if self.rank == receiver_rank {
1691 for _ in 0..options.warmup_iterations + options.latency_iterations {
1692 self.receive_on_rail(rail, Some(sender_rank), latency_tag, ElementType::U8, 1)?;
1693 }
1694 for _ in 0..options.warmup_iterations + options.bandwidth_iterations {
1695 self.receive_on_rail(
1696 rail,
1697 Some(sender_rank),
1698 bandwidth_tag,
1699 ElementType::U8,
1700 options.payload_bytes as u64,
1701 )?;
1702 }
1703 }
1704 self.barrier()?;
1705 if self.rank != sender_rank {
1706 return Ok(None);
1707 }
1708 let latency_ns = median(&mut latency_samples).saturating_add(1) / 2;
1709 let elapsed_ns = median(&mut bandwidth_samples);
1710 let transfer_ns = elapsed_ns
1711 .saturating_sub(latency_ns.saturating_mul(2))
1712 .max(1);
1713 let bandwidth_mbps = u64::try_from(
1714 (options.payload_bytes as u128)
1715 .saturating_mul(1_000)
1716 .checked_div(u128::from(transfer_ns))
1717 .unwrap_or(0),
1718 )
1719 .unwrap_or(u64::MAX)
1720 .max(1);
1721 Ok(Some(LinkProbe {
1722 bandwidth_mbps,
1723 latency_ns: latency_ns.max(1),
1724 }))
1725 }
1726
1727 #[allow(clippy::too_many_arguments)]
1728 fn probe_concurrent_rails_direction(
1729 &self,
1730 first_rank: u32,
1731 second_rank: u32,
1732 sender_rank: u32,
1733 receiver_rank: u32,
1734 direction: u64,
1735 options: TopologyProbeOptions,
1736 ) -> Result<Option<u64>, NetworkError> {
1737 let rails = self.p2p_rails();
1738 if rails <= 1 {
1739 return Err(NetworkError::InvalidConfiguration(
1740 "concurrent topology probing requires at least two P2P rails".into(),
1741 ));
1742 }
1743 let samples = options
1744 .warmup_iterations
1745 .saturating_add(options.bandwidth_iterations);
1746 let mut bandwidth_samples = Vec::with_capacity(options.bandwidth_iterations);
1747 if self.rank == sender_rank {
1748 for sample in 0..samples {
1749 let elapsed_ns = self.send_concurrent_probe_sample(
1750 first_rank,
1751 second_rank,
1752 receiver_rank,
1753 direction,
1754 options.payload_bytes,
1755 )?;
1756 if sample >= options.warmup_iterations {
1757 bandwidth_samples.push(elapsed_ns);
1758 }
1759 }
1760 } else if self.rank == receiver_rank {
1761 for _ in 0..samples {
1762 for rail in 0..rails {
1763 self.receive_on_rail(
1764 rail,
1765 Some(sender_rank),
1766 topology_probe_tag(first_rank, second_rank, rail, direction, 2),
1767 ElementType::U8,
1768 options.payload_bytes as u64,
1769 )?;
1770 }
1771 }
1772 }
1773 self.barrier()?;
1774 if self.rank != sender_rank {
1775 return Ok(None);
1776 }
1777 let elapsed_ns = median(&mut bandwidth_samples).max(1);
1778 let total_bytes = (options.payload_bytes as u128).saturating_mul(rails as u128);
1779 let bandwidth_mbps = u64::try_from(
1780 total_bytes
1781 .saturating_mul(1_000)
1782 .checked_div(u128::from(elapsed_ns))
1783 .unwrap_or(0),
1784 )
1785 .unwrap_or(u64::MAX)
1786 .max(1);
1787 Ok(Some(bandwidth_mbps))
1788 }
1789
1790 fn send_concurrent_probe_sample(
1791 &self,
1792 first_rank: u32,
1793 second_rank: u32,
1794 receiver_rank: u32,
1795 direction: u64,
1796 payload_bytes: usize,
1797 ) -> Result<u64, NetworkError> {
1798 let rails = self.p2p_rails();
1799 thread::scope(|scope| {
1800 let start = Arc::new(Barrier::new(rails + 1));
1801 let mut workers = Vec::with_capacity(rails);
1802 for rail in 0..rails {
1803 let start = Arc::clone(&start);
1804 let payload = vec![0x5a; payload_bytes];
1805 workers.push(scope.spawn(move || {
1806 start.wait();
1807 self.send_on_rail(
1808 rail,
1809 receiver_rank,
1810 topology_probe_tag(first_rank, second_rank, rail, direction, 2),
1811 ElementType::U8,
1812 payload_bytes as u64,
1813 payload,
1814 )
1815 }));
1816 }
1817 let started = Instant::now();
1818 start.wait();
1819 let mut first_error = None;
1820 for worker in workers {
1821 match worker.join() {
1822 Ok(Ok(())) => {}
1823 Ok(Err(error)) if first_error.is_none() => first_error = Some(error),
1824 Ok(Err(_)) => {}
1825 Err(_) if first_error.is_none() => {
1826 first_error = Some(NetworkError::InvalidConfiguration(
1827 "concurrent topology probe worker panicked".into(),
1828 ));
1829 }
1830 Err(_) => {}
1831 }
1832 }
1833 match first_error {
1834 Some(error) => Err(error),
1835 None => Ok(duration_ns(started.elapsed())),
1836 }
1837 })
1838 }
1839
1840 pub fn heartbeat(&self, timeout: Duration) -> Result<Duration, NetworkError> {
1841 if timeout.is_zero() {
1842 return Err(NetworkError::InvalidConfiguration(
1843 "heartbeat timeout must be greater than zero".into(),
1844 ));
1845 }
1846 let started = Instant::now();
1847 let response = self.control()?.request_with_timeout(
1848 self.unique_id,
1849 self.rank,
1850 self.world_size,
1851 Opcode::Heartbeat,
1852 ANY_RANK,
1853 0,
1854 ElementType::None,
1855 0,
1856 Vec::new(),
1857 timeout,
1858 )?;
1859 if response.header.opcode != Opcode::Heartbeat || !response.payload.is_empty() {
1860 return Err(NetworkError::InvalidConfiguration(
1861 "invalid heartbeat response".into(),
1862 ));
1863 }
1864 Ok(started.elapsed())
1865 }
1866
1867 pub fn send(
1868 &self,
1869 destination: u32,
1870 tag: u64,
1871 element_type: ElementType,
1872 element_count: u64,
1873 payload: Vec<u8>,
1874 ) -> Result<(), NetworkError> {
1875 let rail = (tag % self.p2p_rails() as u64) as usize;
1876 self.send_on_rail(rail, destination, tag, element_type, element_count, payload)
1877 }
1878
1879 #[allow(clippy::too_many_arguments)]
1880 pub fn send_on_rail(
1881 &self,
1882 rail: usize,
1883 destination: u32,
1884 tag: u64,
1885 element_type: ElementType,
1886 element_count: u64,
1887 payload: Vec<u8>,
1888 ) -> Result<(), NetworkError> {
1889 if destination >= self.world_size {
1890 return Err(ProtocolError::RankOutOfRange {
1891 name: "destination",
1892 rank: destination,
1893 world_size: self.world_size,
1894 }
1895 .into());
1896 }
1897 match &self.p2p {
1898 P2pDataPlane::Coordinator(_) => {
1899 let response = self.p2p_rail(rail)?.request(
1900 self.unique_id,
1901 self.rank,
1902 self.world_size,
1903 Opcode::Send,
1904 destination,
1905 tag,
1906 element_type,
1907 element_count,
1908 payload,
1909 )?;
1910 if response.header.opcode != Opcode::Send || !response.payload.is_empty() {
1911 return Err(NetworkError::InvalidConfiguration(
1912 "invalid point-to-point send acknowledgement".into(),
1913 ));
1914 }
1915 Ok(())
1916 }
1917 P2pDataPlane::Peer { control, mesh } => {
1918 let result =
1919 mesh.send_on_rail(rail, destination, tag, element_type, element_count, payload);
1920 if let Err(error) = &result {
1921 let message = error.to_string();
1922 let _ = mesh.abort(&message);
1923 let _ = control.abort(&message);
1924 }
1925 result
1926 }
1927 }
1928 }
1929
1930 pub fn receive(
1931 &self,
1932 source: Option<u32>,
1933 tag: u64,
1934 element_type: ElementType,
1935 element_count: u64,
1936 ) -> Result<Frame, NetworkError> {
1937 let rail = (tag % self.p2p_rails() as u64) as usize;
1938 self.receive_on_rail(rail, source, tag, element_type, element_count)
1939 }
1940
1941 pub fn receive_on_rail(
1942 &self,
1943 rail: usize,
1944 source: Option<u32>,
1945 tag: u64,
1946 element_type: ElementType,
1947 element_count: u64,
1948 ) -> Result<Frame, NetworkError> {
1949 if let Some(source) = source
1950 && source >= self.world_size
1951 {
1952 return Err(ProtocolError::RankOutOfRange {
1953 name: "source",
1954 rank: source,
1955 world_size: self.world_size,
1956 }
1957 .into());
1958 }
1959 match &self.p2p {
1960 P2pDataPlane::Coordinator(_) => {
1961 let response = self.p2p_rail(rail)?.request(
1962 self.unique_id,
1963 self.rank,
1964 self.world_size,
1965 Opcode::Receive,
1966 source.unwrap_or(ANY_RANK),
1967 tag,
1968 element_type,
1969 element_count,
1970 Vec::new(),
1971 )?;
1972 if response.header.opcode != Opcode::Receive {
1973 return Err(NetworkError::InvalidConfiguration(
1974 "invalid point-to-point receive response".into(),
1975 ));
1976 }
1977 Ok(response)
1978 }
1979 P2pDataPlane::Peer { control, mesh } => {
1980 let result = mesh.receive_on_rail(rail, source, tag, element_type, element_count);
1981 if let Err(error) = &result {
1982 let message = error.to_string();
1983 let _ = mesh.abort(&message);
1984 let _ = control.abort(&message);
1985 }
1986 result
1987 }
1988 }
1989 }
1990
1991 pub fn exchange(
1992 &self,
1993 opcode: Opcode,
1994 element_type: ElementType,
1995 root_rank: u32,
1996 element_count: u64,
1997 payload: Vec<u8>,
1998 ) -> Result<Frame, NetworkError> {
1999 self.exchange_with_options(
2000 opcode,
2001 element_type,
2002 ExchangeOptions::new(root_rank, element_count),
2003 payload,
2004 )
2005 }
2006
2007 pub fn exchange_with_options(
2008 &self,
2009 opcode: Opcode,
2010 element_type: ElementType,
2011 options: ExchangeOptions,
2012 payload: Vec<u8>,
2013 ) -> Result<Frame, NetworkError> {
2014 let mut inner = self.inner.lock().map_err(|_| NetworkError::Poisoned)?;
2015 let sequence = inner.next_sequence;
2016 let mut header = FrameHeader::collective(
2017 self.unique_id,
2018 opcode,
2019 element_type,
2020 self.rank,
2021 options.root_rank,
2022 self.world_size,
2023 sequence,
2024 options.element_count,
2025 );
2026 header.flags |= options.flags;
2027 header.tag = options.tag;
2028 let request = Frame::new(header, payload)?;
2029 write_frame(&mut inner.stream, &request)?;
2030 let response = read_frame(&mut inner.stream)?.ok_or_else(|| {
2031 NetworkError::InvalidConfiguration(
2032 "coordinator closed before collective response".into(),
2033 )
2034 })?;
2035 validate_response(
2036 &response,
2037 self.unique_id,
2038 self.rank,
2039 self.world_size,
2040 sequence,
2041 )?;
2042 if response.header.opcode == Opcode::Abort {
2043 return Err(NetworkError::RemoteAbort(
2044 String::from_utf8_lossy(&response.payload).into_owned(),
2045 ));
2046 }
2047 inner.next_sequence = inner
2048 .next_sequence
2049 .checked_add(1)
2050 .ok_or(ProtocolError::SequenceOverflow)?;
2051 Ok(response)
2052 }
2053
2054 pub fn barrier(&self) -> Result<(), NetworkError> {
2055 let response =
2056 self.exchange(Opcode::Barrier, ElementType::None, ANY_RANK, 0, Vec::new())?;
2057 if response.header.opcode != Opcode::Barrier || !response.payload.is_empty() {
2058 return Err(NetworkError::InvalidConfiguration(
2059 "invalid barrier response".into(),
2060 ));
2061 }
2062 Ok(())
2063 }
2064
2065 pub fn set_timeout(&self, timeout: Duration) -> Result<(), NetworkError> {
2066 if timeout.is_zero() {
2067 return Err(NetworkError::InvalidConfiguration(
2068 "collective timeout must be greater than zero".into(),
2069 ));
2070 }
2071 let timeout_ms = u64::try_from(timeout.as_millis().max(1)).map_err(|_| {
2072 NetworkError::InvalidConfiguration("collective timeout exceeds u64 milliseconds".into())
2073 })?;
2074 let mut inner = self.inner.lock().map_err(|_| NetworkError::Poisoned)?;
2075 let sequence = inner.next_sequence;
2076 let mut header = FrameHeader::collective(
2077 self.unique_id,
2078 Opcode::SetTimeout,
2079 ElementType::None,
2080 self.rank,
2081 ANY_RANK,
2082 self.world_size,
2083 sequence,
2084 0,
2085 );
2086 header.tag = timeout_ms;
2087 write_frame(&mut inner.stream, &Frame::new(header, Vec::new())?)?;
2088 let response = read_frame(&mut inner.stream)?.ok_or_else(|| {
2089 NetworkError::InvalidConfiguration(
2090 "coordinator closed before SET_TIMEOUT response".into(),
2091 )
2092 })?;
2093 validate_response(
2094 &response,
2095 self.unique_id,
2096 self.rank,
2097 self.world_size,
2098 sequence,
2099 )?;
2100 if response.header.opcode == Opcode::Abort {
2101 return Err(NetworkError::RemoteAbort(
2102 String::from_utf8_lossy(&response.payload).into_owned(),
2103 ));
2104 }
2105 if response.header.opcode != Opcode::SetTimeout || !response.payload.is_empty() {
2106 return Err(NetworkError::InvalidConfiguration(
2107 "invalid SET_TIMEOUT response".into(),
2108 ));
2109 }
2110 inner.next_sequence = inner
2111 .next_sequence
2112 .checked_add(1)
2113 .ok_or(ProtocolError::SequenceOverflow)?;
2114 inner
2115 .stream
2116 .set_read_timeout(Some(timeout.saturating_add(COLLECTIVE_RESPONSE_GRACE)))?;
2117 inner.stream.set_write_timeout(Some(timeout))?;
2118 match &self.p2p {
2119 P2pDataPlane::Coordinator(rails) => {
2120 for rail in rails {
2121 rail.set_timeout(timeout)?;
2122 }
2123 }
2124 P2pDataPlane::Peer { control, mesh } => {
2125 control.set_timeout(timeout)?;
2126 mesh.set_timeout(timeout)?;
2127 }
2128 }
2129 Ok(())
2130 }
2131
2132 pub fn abort(&self, message: &str) -> Result<(), NetworkError> {
2136 let mut first_error = None;
2137 match &self.p2p {
2138 P2pDataPlane::Coordinator(rails) => {
2139 for rail in rails {
2140 if let Err(error) = rail.abort(message)
2141 && first_error.is_none()
2142 {
2143 first_error = Some(error);
2144 }
2145 }
2146 }
2147 P2pDataPlane::Peer { control, mesh } => {
2148 if let Err(error) = control.abort(message) {
2149 first_error = Some(error);
2150 }
2151 if let Err(error) = mesh.abort(message)
2152 && first_error.is_none()
2153 {
2154 first_error = Some(error);
2155 }
2156 }
2157 }
2158 match first_error {
2159 Some(error) => Err(error),
2160 None => Ok(()),
2161 }
2162 }
2163}
2164
2165impl super::RankTransport for TcpRankSession {
2166 fn rank(&self) -> u32 {
2167 Self::rank(self)
2168 }
2169
2170 fn world_size(&self) -> u32 {
2171 Self::world_size(self)
2172 }
2173
2174 fn transport(&self) -> CollectiveTransport {
2175 Self::transport(self)
2176 }
2177
2178 fn p2p_rails(&self) -> usize {
2179 Self::p2p_rails(self)
2180 }
2181
2182 fn heartbeat(&self, timeout: Duration) -> Result<Duration, NetworkError> {
2183 Self::heartbeat(self, timeout)
2184 }
2185
2186 fn send_on_rail(
2187 &self,
2188 rail: usize,
2189 destination: u32,
2190 tag: u64,
2191 element_type: ElementType,
2192 element_count: u64,
2193 payload: Vec<u8>,
2194 ) -> Result<(), NetworkError> {
2195 Self::send_on_rail(
2196 self,
2197 rail,
2198 destination,
2199 tag,
2200 element_type,
2201 element_count,
2202 payload,
2203 )
2204 }
2205
2206 fn receive_on_rail(
2207 &self,
2208 rail: usize,
2209 source: Option<u32>,
2210 tag: u64,
2211 element_type: ElementType,
2212 element_count: u64,
2213 ) -> Result<Frame, NetworkError> {
2214 Self::receive_on_rail(self, rail, source, tag, element_type, element_count)
2215 }
2216
2217 fn exchange_with_options(
2218 &self,
2219 opcode: Opcode,
2220 element_type: ElementType,
2221 options: ExchangeOptions,
2222 payload: Vec<u8>,
2223 ) -> Result<Frame, NetworkError> {
2224 Self::exchange_with_options(self, opcode, element_type, options, payload)
2225 }
2226
2227 fn set_timeout(&self, timeout: Duration) -> Result<(), NetworkError> {
2228 Self::set_timeout(self, timeout)
2229 }
2230
2231 fn abort(&self, message: &str) -> Result<(), NetworkError> {
2232 Self::abort(self, message)
2233 }
2234}
2235
2236fn connect_rank_channel(
2237 addresses: &[SocketAddr],
2238 unique_id: UniqueId,
2239 rank: u32,
2240 world_size: u32,
2241 timeout: Duration,
2242 p2p: bool,
2243 rail: usize,
2244) -> Result<TcpStream, NetworkError> {
2245 let mut last_error = None;
2246 let mut stream = None;
2247 for address in addresses {
2248 match TcpStream::connect_timeout(address, timeout) {
2249 Ok(connected) => {
2250 stream = Some(connected);
2251 break;
2252 }
2253 Err(error) => last_error = Some(error),
2254 }
2255 }
2256 let mut stream = stream.ok_or_else(|| {
2257 NetworkError::Io(last_error.unwrap_or_else(|| {
2258 io::Error::new(
2259 io::ErrorKind::InvalidInput,
2260 "no rendezvous address resolved",
2261 )
2262 }))
2263 })?;
2264 stream.set_nodelay(true)?;
2265 stream.set_read_timeout(Some(timeout))?;
2266 stream.set_write_timeout(Some(timeout))?;
2267 let mut header = FrameHeader::collective(
2268 unique_id,
2269 Opcode::Join,
2270 ElementType::None,
2271 rank,
2272 ANY_RANK,
2273 world_size,
2274 0,
2275 0,
2276 );
2277 if p2p {
2278 header.flags |= FLAG_P2P_CHANNEL;
2279 header.tag = rail as u64;
2280 }
2281 write_frame(&mut stream, &Frame::new(header, Vec::new())?)?;
2282 let ready = read_frame(&mut stream)?.ok_or_else(|| {
2283 NetworkError::InvalidConfiguration("coordinator closed before READY".into())
2284 })?;
2285 validate_response(&ready, unique_id, rank, world_size, 0)?;
2286 if ready.header.opcode != Opcode::Ready
2287 || (ready.header.flags & FLAG_P2P_CHANNEL != 0) != p2p
2288 || ready.header.tag != if p2p { rail as u64 } else { 0 }
2289 {
2290 return Err(NetworkError::InvalidConfiguration(format!(
2291 "coordinator answered JOIN on the wrong channel/rail with {:?} tag {}",
2292 ready.header.opcode, ready.header.tag
2293 )));
2294 }
2295 Ok(stream)
2296}
2297
2298fn run_p2p_client_reader(
2299 stream: &mut TcpStream,
2300 unique_id: UniqueId,
2301 rank: u32,
2302 world_size: u32,
2303 pending: &Arc<Mutex<HashMap<P2pResponseKey, PendingP2pResponse>>>,
2304 failure: &Arc<Mutex<Option<String>>>,
2305) {
2306 loop {
2307 let frame = match read_frame(stream) {
2308 Ok(Some(frame)) => frame,
2309 Ok(None) => {
2310 fail_pending(pending, failure, "point-to-point coordinator closed");
2311 return;
2312 }
2313 Err(error) => {
2314 fail_pending(
2315 pending,
2316 failure,
2317 &format!("point-to-point read failed: {error}"),
2318 );
2319 return;
2320 }
2321 };
2322 if frame.header.unique_id != unique_id
2323 || frame.header.world_size != world_size
2324 || frame.header.destination_rank != rank
2325 {
2326 fail_pending(pending, failure, "invalid point-to-point response envelope");
2327 return;
2328 }
2329 if frame.header.opcode == Opcode::Abort {
2330 fail_pending(pending, failure, &String::from_utf8_lossy(&frame.payload));
2331 return;
2332 }
2333 if !matches!(
2334 frame.header.opcode,
2335 Opcode::Send | Opcode::Receive | Opcode::Heartbeat
2336 ) {
2337 fail_pending(
2338 pending,
2339 failure,
2340 &format!(
2341 "unexpected {:?} response on point-to-point channel",
2342 frame.header.opcode
2343 ),
2344 );
2345 return;
2346 }
2347 let key = P2pResponseKey {
2348 opcode: frame.header.opcode,
2349 sequence: frame.header.sequence,
2350 };
2351 let sender = match pending.lock() {
2352 Ok(mut pending) => pending.remove(&key),
2353 Err(_) => return,
2354 };
2355 if let Some(sender) = sender {
2356 let _ = sender.send(Ok(frame));
2357 }
2358 }
2359}
2360
2361fn fail_pending(
2362 pending: &Arc<Mutex<HashMap<P2pResponseKey, PendingP2pResponse>>>,
2363 failure: &Arc<Mutex<Option<String>>>,
2364 message: &str,
2365) {
2366 if let Ok(mut failure) = failure.lock() {
2367 failure.get_or_insert_with(|| message.to_owned());
2368 }
2369 if let Ok(mut pending) = pending.lock() {
2370 for (_, sender) in pending.drain() {
2371 let _ = sender.send(Err(message.to_owned()));
2372 }
2373 }
2374}
2375
2376fn validate_join(
2377 frame: &Frame,
2378 unique_id: UniqueId,
2379 world_size: usize,
2380 p2p: bool,
2381 rail: usize,
2382) -> Result<(), NetworkError> {
2383 if frame.header.opcode != Opcode::Join
2384 || frame.header.element_type != ElementType::None
2385 || !frame.payload.is_empty()
2386 {
2387 return Err(NetworkError::InvalidConfiguration(
2388 "first rank frame must be an empty JOIN".into(),
2389 ));
2390 }
2391 if frame.header.unique_id != unique_id {
2392 return Err(NetworkError::WrongSession);
2393 }
2394 if (frame.header.flags & FLAG_P2P_CHANNEL != 0) != p2p {
2395 return Err(NetworkError::InvalidConfiguration(
2396 "rank joined the wrong TCP channel".into(),
2397 ));
2398 }
2399 let expected_rail = if p2p { rail as u64 } else { 0 };
2400 if frame.header.tag != expected_rail {
2401 return Err(NetworkError::InvalidConfiguration(format!(
2402 "rank joined point-to-point rail {}, coordinator expects {expected_rail}",
2403 frame.header.tag
2404 )));
2405 }
2406 if frame.header.world_size as usize != world_size {
2407 return Err(NetworkError::InvalidConfiguration(format!(
2408 "rank declares world size {}, coordinator expects {world_size}",
2409 frame.header.world_size
2410 )));
2411 }
2412 Ok(())
2413}
2414
2415fn p2p_rails_from_environment() -> Result<usize, NetworkError> {
2416 match env::var("GX1_P2P_RAILS") {
2417 Ok(value) => {
2418 let rails = value.trim().parse::<usize>().map_err(|_| {
2419 NetworkError::InvalidConfiguration(format!(
2420 "GX1_P2P_RAILS must be an integer in 1..={MAX_P2P_RAILS}, got {value:?}"
2421 ))
2422 })?;
2423 validate_p2p_rails(rails)?;
2424 Ok(rails)
2425 }
2426 Err(env::VarError::NotPresent) => Ok(1),
2427 Err(env::VarError::NotUnicode(value)) => Err(NetworkError::InvalidConfiguration(format!(
2428 "GX1_P2P_RAILS is not Unicode: {:?}",
2429 value.to_string_lossy()
2430 ))),
2431 }
2432}
2433
2434fn parse_probe_usize(name: &'static str, default: usize) -> Result<usize, NetworkError> {
2435 match env::var(name) {
2436 Ok(value) => value.trim().parse::<usize>().map_err(|_| {
2437 NetworkError::InvalidConfiguration(format!(
2438 "{name} must be a non-negative integer, got {value:?}"
2439 ))
2440 }),
2441 Err(env::VarError::NotPresent) => Ok(default),
2442 Err(env::VarError::NotUnicode(value)) => Err(NetworkError::InvalidConfiguration(format!(
2443 "{name} is not Unicode: {:?}",
2444 value.to_string_lossy()
2445 ))),
2446 }
2447}
2448
2449fn duration_ns(duration: Duration) -> u64 {
2450 u64::try_from(duration.as_nanos())
2451 .unwrap_or(u64::MAX)
2452 .max(1)
2453}
2454
2455fn median(samples: &mut [u64]) -> u64 {
2456 samples.sort_unstable();
2457 samples[samples.len() / 2]
2458}
2459
2460fn conservative_probe_bandwidth_bucket(value: u64) -> u64 {
2461 let exponent = 63_u32.saturating_sub(value.max(1).leading_zeros());
2462 1_u64 << (exponent & !1)
2463}
2464
2465fn conservative_probe_latency_bucket(value: u64) -> u64 {
2466 let exponent = 63_u32.saturating_sub(value.max(1).leading_zeros());
2467 let upper_exponent = (exponent & !1).saturating_add(2);
2468 if upper_exponent >= u64::BITS {
2469 u64::MAX
2470 } else {
2471 1_u64 << upper_exponent
2472 }
2473}
2474
2475fn topology_probe_tag(
2476 first_rank: u32,
2477 second_rank: u32,
2478 rail: usize,
2479 direction: u64,
2480 kind: u64,
2481) -> u64 {
2482 let mut hash = 0xcbf2_9ce4_8422_2325_u64;
2483 for byte in first_rank
2484 .to_le_bytes()
2485 .into_iter()
2486 .chain(second_rank.to_le_bytes())
2487 .chain((rail as u64).to_le_bytes())
2488 .chain(direction.to_le_bytes())
2489 .chain(kind.to_le_bytes())
2490 {
2491 hash ^= u64::from(byte);
2492 hash = hash.wrapping_mul(0x0000_0100_0000_01b3);
2493 }
2494 TOPOLOGY_PROBE_TAG_PREFIX | (hash & 0x0000_ffff_ffff_ffff)
2495}
2496
2497fn tcp_transport_from_environment() -> Result<CollectiveTransport, NetworkError> {
2498 match env::var("GX1_COLLECTIVE_TRANSPORT") {
2499 Ok(value) => match value.trim().to_ascii_lowercase().as_str() {
2500 "coordinator" | "tcp_coordinator" | "tcp_host_staged" => {
2501 Ok(CollectiveTransport::TcpHostStaged)
2502 }
2503 "peer" | "tcp_peer" => Ok(CollectiveTransport::TcpPeer),
2504 _ => Err(NetworkError::InvalidConfiguration(format!(
2505 "GX1_COLLECTIVE_TRANSPORT must be tcp_host_staged or tcp_peer, got {value:?}"
2506 ))),
2507 },
2508 Err(env::VarError::NotPresent) => Ok(CollectiveTransport::TcpHostStaged),
2509 Err(env::VarError::NotUnicode(value)) => Err(NetworkError::InvalidConfiguration(format!(
2510 "GX1_COLLECTIVE_TRANSPORT is not Unicode: {:?}",
2511 value.to_string_lossy()
2512 ))),
2513 }
2514}
2515
2516fn validate_tcp_transport(transport: CollectiveTransport) -> Result<(), NetworkError> {
2517 match transport {
2518 CollectiveTransport::TcpHostStaged | CollectiveTransport::TcpPeer => Ok(()),
2519 CollectiveTransport::HostStaged
2520 | CollectiveTransport::PciePeer
2521 | CollectiveTransport::Rdma
2522 | CollectiveTransport::GxLink => Err(NetworkError::InvalidConfiguration(format!(
2523 "{transport:?} is not a TCP transport"
2524 ))),
2525 }
2526}
2527
2528fn validate_p2p_rails(rails: usize) -> Result<(), NetworkError> {
2529 if (1..=MAX_P2P_RAILS).contains(&rails) {
2530 Ok(())
2531 } else {
2532 Err(NetworkError::InvalidConfiguration(format!(
2533 "point-to-point rail count {rails} is outside 1..={MAX_P2P_RAILS}"
2534 )))
2535 }
2536}
2537
2538fn validate_response(
2539 frame: &Frame,
2540 unique_id: UniqueId,
2541 rank: u32,
2542 world_size: u32,
2543 sequence: u64,
2544) -> Result<(), NetworkError> {
2545 if frame.header.unique_id != unique_id || frame.header.world_size != world_size {
2546 return Err(NetworkError::WrongSession);
2547 }
2548 if frame.header.destination_rank != rank {
2549 return Err(NetworkError::WrongDestination {
2550 expected: rank,
2551 actual: frame.header.destination_rank,
2552 });
2553 }
2554 if frame.header.sequence != sequence {
2555 return Err(ProtocolError::SequenceMismatch {
2556 expected: sequence,
2557 actual: frame.header.sequence,
2558 }
2559 .into());
2560 }
2561 Ok(())
2562}
2563
2564#[derive(Debug)]
2565enum P2pEvent {
2566 Frame { rank: usize, frame: Frame },
2567 Closed { rank: usize },
2568 Failed { rank: usize, message: String },
2569}
2570
2571fn run_p2p_loop(
2572 mut streams: Vec<TcpStream>,
2573 unique_id: UniqueId,
2574 world_size: usize,
2575 heartbeat_timeout: Option<Duration>,
2576 server_failure: SharedServerFailure,
2577) -> Result<(), NetworkError> {
2578 let (sender, receiver): (Sender<P2pEvent>, Receiver<P2pEvent>) = mpsc::channel();
2579 let mut readers = Vec::with_capacity(world_size);
2580 for (rank, stream) in streams.iter().enumerate() {
2581 let mut stream = stream.try_clone()?;
2582 stream.set_read_timeout(None)?;
2583 let sender = sender.clone();
2584 readers.push(
2585 thread::Builder::new()
2586 .name(format!("gx1-p2p-server-rank-{rank}"))
2587 .spawn(move || {
2588 loop {
2589 match read_frame(&mut stream) {
2590 Ok(Some(frame)) => {
2591 if sender.send(P2pEvent::Frame { rank, frame }).is_err() {
2592 return;
2593 }
2594 }
2595 Ok(None) => {
2596 let _ = sender.send(P2pEvent::Closed { rank });
2597 return;
2598 }
2599 Err(error) => {
2600 let _ = sender.send(P2pEvent::Failed {
2601 rank,
2602 message: error.to_string(),
2603 });
2604 return;
2605 }
2606 }
2607 }
2608 })
2609 .map_err(|error| {
2610 NetworkError::InvalidConfiguration(format!(
2611 "cannot start point-to-point rank reader: {error}"
2612 ))
2613 })?,
2614 );
2615 }
2616 drop(sender);
2617 let mut sends = VecDeque::<Frame>::new();
2618 let mut receives = VecDeque::<Frame>::new();
2619 let mut last_seen = vec![Instant::now(); world_size];
2620 let mut graceful = vec![false; world_size];
2621 let mut closed = vec![false; world_size];
2622 let poll_interval = heartbeat_timeout.map(|timeout| {
2623 (timeout / 4)
2624 .max(Duration::from_millis(1))
2625 .min(Duration::from_millis(250))
2626 });
2627 let result = loop {
2628 let event = match poll_interval {
2629 Some(interval) => match receiver.recv_timeout(interval) {
2630 Ok(event) => Some(event),
2631 Err(RecvTimeoutError::Timeout) => None,
2632 Err(RecvTimeoutError::Disconnected) => break Ok(()),
2633 },
2634 None => match receiver.recv() {
2635 Ok(event) => Some(event),
2636 Err(_) => break Ok(()),
2637 },
2638 };
2639 let routed = match event {
2640 Some(P2pEvent::Frame { rank, frame }) => {
2641 last_seen[rank] = Instant::now();
2642 if graceful[rank] {
2643 Err(NetworkError::InvalidConfiguration(format!(
2644 "point-to-point rank {rank} sent data after LEAVE"
2645 )))
2646 } else {
2647 let leaving = frame.header.opcode == Opcode::Leave;
2648 let result = route_p2p_request(
2649 &mut streams,
2650 unique_id,
2651 world_size,
2652 rank,
2653 frame,
2654 &mut sends,
2655 &mut receives,
2656 );
2657 if result.is_ok() && leaving {
2658 graceful[rank] = true;
2659 }
2660 result
2661 }
2662 }
2663 Some(P2pEvent::Closed { rank }) => {
2664 closed[rank] = true;
2665 if heartbeat_timeout.is_some() && !graceful[rank] {
2666 Err(NetworkError::InvalidConfiguration(format!(
2667 "point-to-point rank {rank} disconnected without LEAVE"
2668 )))
2669 } else {
2670 Ok(())
2671 }
2672 }
2673 Some(P2pEvent::Failed { rank, message }) => {
2674 closed[rank] = true;
2675 if heartbeat_timeout.is_some() && !graceful[rank] {
2676 Err(NetworkError::InvalidConfiguration(format!(
2677 "point-to-point rank {rank} failed without LEAVE: {message}"
2678 )))
2679 } else {
2680 Ok(())
2681 }
2682 }
2683 None => Ok(()),
2684 };
2685 if let Err(error) = routed {
2686 let message = server_failure_text(&error);
2687 publish_server_failure(&server_failure, &message);
2688 broadcast_p2p_abort(&mut streams, unique_id, world_size, &message);
2689 break Err(error);
2690 }
2691 if closed.iter().all(|closed| *closed) {
2692 break Ok(());
2693 }
2694 if let Some(timeout) = heartbeat_timeout
2695 && let Some((rank, elapsed)) = last_seen
2696 .iter()
2697 .enumerate()
2698 .filter(|(rank, _)| !closed[*rank] && !graceful[*rank])
2699 .map(|(rank, last_seen)| (rank, last_seen.elapsed()))
2700 .find(|(_, elapsed)| *elapsed >= timeout)
2701 {
2702 let error = NetworkError::Timeout(format!(
2703 "heartbeat lease for rank {rank}; last contact {} ms ago",
2704 elapsed.as_millis()
2705 ));
2706 let message = server_failure_text(&error);
2707 publish_server_failure(&server_failure, &message);
2708 broadcast_p2p_abort(&mut streams, unique_id, world_size, &message);
2709 break Err(error);
2710 }
2711 };
2712 for stream in &streams {
2713 let _ = stream.shutdown(Shutdown::Both);
2714 }
2715 for reader in readers {
2716 let _ = reader.join();
2717 }
2718 result
2719}
2720
2721#[allow(clippy::too_many_arguments)]
2722fn route_p2p_request(
2723 streams: &mut [TcpStream],
2724 unique_id: UniqueId,
2725 world_size: usize,
2726 connection_rank: usize,
2727 frame: Frame,
2728 sends: &mut VecDeque<Frame>,
2729 receives: &mut VecDeque<Frame>,
2730) -> Result<(), NetworkError> {
2731 validate_p2p_request(&frame, unique_id, world_size, connection_rank)?;
2732 match frame.header.opcode {
2733 Opcode::Send => {
2734 let sender = frame.header.source_rank as usize;
2735 let acknowledgement = p2p_send_acknowledgement(&frame, unique_id, world_size)?;
2736 write_frame(&mut streams[sender], &acknowledgement)?;
2737 if let Some(index) = receives
2738 .iter()
2739 .position(|receive| p2p_matches(&frame, receive))
2740 {
2741 let receive = receives.remove(index).expect("matched receive exists");
2742 deliver_p2p_receive(streams, unique_id, world_size, &frame, &receive)
2743 } else {
2744 sends.push_back(frame);
2745 Ok(())
2746 }
2747 }
2748 Opcode::Receive => {
2749 if let Some(index) = sends.iter().position(|send| p2p_matches(send, &frame)) {
2750 let send = sends.remove(index).expect("matched send exists");
2751 deliver_p2p_receive(streams, unique_id, world_size, &send, &frame)
2752 } else {
2753 receives.push_back(frame);
2754 Ok(())
2755 }
2756 }
2757 Opcode::Heartbeat => {
2758 let response = p2p_heartbeat_response(&frame, unique_id, world_size)?;
2759 write_frame(&mut streams[connection_rank], &response)
2760 }
2761 Opcode::Leave => Ok(()),
2762 Opcode::Abort => Err(NetworkError::RemoteAbort(format!(
2763 "rank {connection_rank} aborted: {}",
2764 String::from_utf8_lossy(&frame.payload)
2765 ))),
2766 _ => unreachable!("point-to-point request was validated"),
2767 }
2768}
2769
2770fn validate_p2p_request(
2771 frame: &Frame,
2772 unique_id: UniqueId,
2773 world_size: usize,
2774 connection_rank: usize,
2775) -> Result<(), NetworkError> {
2776 if frame.header.unique_id != unique_id {
2777 return Err(NetworkError::WrongSession);
2778 }
2779 if frame.header.world_size as usize != world_size
2780 || frame.header.source_rank as usize != connection_rank
2781 {
2782 return Err(NetworkError::InvalidConfiguration(format!(
2783 "point-to-point connection rank {connection_rank} submitted source {} and world {}",
2784 frame.header.source_rank, frame.header.world_size
2785 )));
2786 }
2787 if frame.header.root_rank != ANY_RANK {
2788 return Err(NetworkError::InvalidConfiguration(
2789 "point-to-point request cannot declare a collective root".into(),
2790 ));
2791 }
2792 match frame.header.opcode {
2793 Opcode::Send => {
2794 if frame.header.destination_rank == ANY_RANK {
2795 return Err(NetworkError::InvalidConfiguration(
2796 "point-to-point send requires a destination rank".into(),
2797 ));
2798 }
2799 }
2800 Opcode::Receive => {
2801 if !frame.payload.is_empty() {
2802 return Err(NetworkError::InvalidConfiguration(
2803 "point-to-point receive request cannot carry data".into(),
2804 ));
2805 }
2806 }
2807 Opcode::Abort => {
2808 if frame.header.destination_rank != ANY_RANK
2809 || frame.header.element_type != ElementType::U8
2810 || frame.payload.is_empty()
2811 {
2812 return Err(NetworkError::InvalidConfiguration(
2813 "point-to-point abort requires a non-empty U8 message".into(),
2814 ));
2815 }
2816 }
2817 Opcode::Heartbeat => {
2818 if frame.header.destination_rank != ANY_RANK
2819 || frame.header.element_type != ElementType::None
2820 || frame.header.element_count != 0
2821 || !frame.payload.is_empty()
2822 {
2823 return Err(NetworkError::InvalidConfiguration(
2824 "point-to-point heartbeat must be an empty control frame".into(),
2825 ));
2826 }
2827 }
2828 Opcode::Leave => {
2829 if frame.header.destination_rank != ANY_RANK
2830 || frame.header.element_type != ElementType::None
2831 || frame.header.element_count != 0
2832 || !frame.payload.is_empty()
2833 {
2834 return Err(NetworkError::InvalidConfiguration(
2835 "point-to-point LEAVE must be an empty control frame".into(),
2836 ));
2837 }
2838 }
2839 opcode => {
2840 return Err(NetworkError::InvalidConfiguration(format!(
2841 "opcode {opcode:?} is invalid on the point-to-point channel"
2842 )));
2843 }
2844 }
2845 Ok(())
2846}
2847
2848fn p2p_matches(send: &Frame, receive: &Frame) -> bool {
2849 send.header.destination_rank == receive.header.source_rank
2850 && (receive.header.destination_rank == ANY_RANK
2851 || receive.header.destination_rank == send.header.source_rank)
2852 && send.header.tag == receive.header.tag
2853}
2854
2855fn p2p_send_acknowledgement(
2856 send: &Frame,
2857 unique_id: UniqueId,
2858 world_size: usize,
2859) -> Result<Frame, NetworkError> {
2860 let mut header = FrameHeader::collective(
2861 unique_id,
2862 Opcode::Send,
2863 ElementType::None,
2864 send.header.source_rank,
2865 ANY_RANK,
2866 world_size as u32,
2867 send.header.sequence,
2868 0,
2869 );
2870 header.destination_rank = send.header.source_rank;
2871 header.tag = send.header.tag;
2872 Ok(Frame::new(header, Vec::new())?)
2873}
2874
2875fn p2p_heartbeat_response(
2876 heartbeat: &Frame,
2877 unique_id: UniqueId,
2878 world_size: usize,
2879) -> Result<Frame, NetworkError> {
2880 let mut header = FrameHeader::collective(
2881 unique_id,
2882 Opcode::Heartbeat,
2883 ElementType::None,
2884 heartbeat.header.source_rank,
2885 ANY_RANK,
2886 world_size as u32,
2887 heartbeat.header.sequence,
2888 0,
2889 );
2890 header.destination_rank = heartbeat.header.source_rank;
2891 Ok(Frame::new(header, Vec::new())?)
2892}
2893
2894fn deliver_p2p_receive(
2895 streams: &mut [TcpStream],
2896 unique_id: UniqueId,
2897 world_size: usize,
2898 send: &Frame,
2899 receive: &Frame,
2900) -> Result<(), NetworkError> {
2901 if send.header.element_type != receive.header.element_type
2902 || send.header.element_count != receive.header.element_count
2903 {
2904 return Err(NetworkError::InvalidConfiguration(format!(
2905 "point-to-point tag {} contract mismatch: send {:?}/{} receive {:?}/{}",
2906 send.header.tag,
2907 send.header.element_type,
2908 send.header.element_count,
2909 receive.header.element_type,
2910 receive.header.element_count
2911 )));
2912 }
2913 let receiver = receive.header.source_rank as usize;
2914 let mut header = FrameHeader::collective(
2915 unique_id,
2916 Opcode::Receive,
2917 send.header.element_type,
2918 send.header.source_rank,
2919 ANY_RANK,
2920 world_size as u32,
2921 receive.header.sequence,
2922 send.header.element_count,
2923 );
2924 header.destination_rank = receive.header.source_rank;
2925 header.tag = send.header.tag;
2926 write_frame(
2927 &mut streams[receiver],
2928 &Frame::new(header, send.payload.clone())?,
2929 )
2930}
2931
2932fn broadcast_p2p_abort(
2933 streams: &mut [TcpStream],
2934 unique_id: UniqueId,
2935 world_size: usize,
2936 message: &str,
2937) {
2938 for (rank, stream) in streams.iter_mut().enumerate() {
2939 if let Ok(frame) = abort_frame(unique_id, 0, rank as u32, world_size as u32, message) {
2940 let _ = write_frame(stream, &frame);
2941 }
2942 }
2943}
2944
2945fn validate_and_route(
2946 agreement: &mut CollectiveAgreement,
2947 unique_id: UniqueId,
2948 sequence: u64,
2949 requests: &[Frame],
2950) -> Result<Vec<Frame>, NetworkError> {
2951 let world_size = requests.len();
2952 for (rank, request) in requests.iter().enumerate() {
2953 if request.header.unique_id != unique_id {
2954 return Err(NetworkError::WrongSession);
2955 }
2956 if request.header.source_rank as usize != rank
2957 || request.header.world_size as usize != world_size
2958 {
2959 return Err(NetworkError::InvalidConfiguration(format!(
2960 "connection rank {rank} submitted source rank {} and world size {}",
2961 request.header.source_rank, request.header.world_size
2962 )));
2963 }
2964 let variable = request.header.opcode == Opcode::AllToAllV;
2965 let descriptor = OperationDescriptor::new(
2966 request.header.opcode,
2967 request.header.element_type,
2968 request.header.root_rank,
2969 if variable {
2970 0
2971 } else {
2972 request.header.element_count
2973 },
2974 )
2975 .with_layout_hash(if variable { 0 } else { request.header.tag });
2976 agreement.submit(rank, sequence, descriptor)?;
2977 }
2978 let first = &requests[0].header;
2979 let element_bytes = first.element_type.byte_width();
2980 let count = usize::try_from(first.element_count).map_err(|_| {
2981 NetworkError::InvalidConfiguration("collective element count exceeds usize".into())
2982 })?;
2983 let rank_payload_bytes = count.checked_mul(element_bytes).ok_or_else(|| {
2984 NetworkError::InvalidConfiguration("collective payload size overflow".into())
2985 })?;
2986 let root = if first.root_rank == ANY_RANK {
2987 None
2988 } else {
2989 Some(first.root_rank as usize)
2990 };
2991 match first.opcode {
2992 Opcode::Barrier => route_empty(requests, unique_id, sequence, Opcode::Barrier),
2993 Opcode::SetTimeout => {
2994 if first.element_type != ElementType::None
2995 || first.root_rank != ANY_RANK
2996 || first.tag == 0
2997 {
2998 return Err(NetworkError::InvalidConfiguration(
2999 "SET_TIMEOUT requires an empty control frame and non-zero millisecond tag"
3000 .into(),
3001 ));
3002 }
3003 route_empty(requests, unique_id, sequence, Opcode::SetTimeout)
3004 }
3005 Opcode::Broadcast => {
3006 let root = required_root(root, world_size)?;
3007 require_only_root_payload(requests, root, rank_payload_bytes)?;
3008 route_same_payload(
3009 requests,
3010 unique_id,
3011 sequence,
3012 first.opcode,
3013 first.element_type,
3014 count,
3015 &requests[root].payload,
3016 )
3017 }
3018 Opcode::AllGather => {
3019 require_all_payloads(requests, rank_payload_bytes)?;
3020 let payload = concatenate_payloads(requests)?;
3021 route_same_payload(
3022 requests,
3023 unique_id,
3024 sequence,
3025 first.opcode,
3026 first.element_type,
3027 count.checked_mul(world_size).ok_or_else(|| {
3028 NetworkError::InvalidConfiguration("all-gather count overflow".into())
3029 })?,
3030 &payload,
3031 )
3032 }
3033 Opcode::Gather | Opcode::Reduce => {
3034 let root = required_root(root, world_size)?;
3035 require_all_payloads(requests, rank_payload_bytes)?;
3036 let payload = concatenate_payloads(requests)?;
3037 route_root_payload(
3038 requests,
3039 unique_id,
3040 sequence,
3041 first.opcode,
3042 first.element_type,
3043 root,
3044 count.checked_mul(world_size).ok_or_else(|| {
3045 NetworkError::InvalidConfiguration("rooted collective count overflow".into())
3046 })?,
3047 &payload,
3048 )
3049 }
3050 Opcode::Scatter => {
3051 let root = required_root(root, world_size)?;
3052 require_only_root_payload(requests, root, rank_payload_bytes)?;
3053 if !count.is_multiple_of(world_size) {
3054 return Err(NetworkError::InvalidConfiguration(
3055 "scatter total element count is not divisible by world size".into(),
3056 ));
3057 }
3058 let shard_count = count / world_size;
3059 let shard_bytes = shard_count * element_bytes;
3060 let mut responses = Vec::with_capacity(world_size);
3061 for rank in 0..world_size {
3062 let start = rank * shard_bytes;
3063 responses.push(data_response(
3064 requests,
3065 unique_id,
3066 sequence,
3067 first.opcode,
3068 first.element_type,
3069 rank,
3070 shard_count,
3071 requests[root].payload[start..start + shard_bytes].to_vec(),
3072 )?);
3073 }
3074 Ok(responses)
3075 }
3076 Opcode::AllReduce => {
3077 require_all_payloads(requests, rank_payload_bytes)?;
3078 let payload = concatenate_payloads(requests)?;
3079 route_same_payload(
3080 requests,
3081 unique_id,
3082 sequence,
3083 first.opcode,
3084 first.element_type,
3085 count.checked_mul(world_size).ok_or_else(|| {
3086 NetworkError::InvalidConfiguration("all-reduce count overflow".into())
3087 })?,
3088 &payload,
3089 )
3090 }
3091 Opcode::ReduceScatter => {
3092 require_all_payloads(requests, rank_payload_bytes)?;
3093 if !count.is_multiple_of(world_size) {
3094 return Err(NetworkError::InvalidConfiguration(
3095 "reduce-scatter count is not divisible by world size".into(),
3096 ));
3097 }
3098 let shard_count = count / world_size;
3099 let shard_bytes = shard_count * element_bytes;
3100 let mut responses = Vec::with_capacity(world_size);
3101 for destination in 0..world_size {
3102 let mut payload = Vec::with_capacity(shard_bytes * world_size);
3103 for request in requests {
3104 let start = destination * shard_bytes;
3105 payload.extend_from_slice(&request.payload[start..start + shard_bytes]);
3106 }
3107 responses.push(data_response(
3108 requests,
3109 unique_id,
3110 sequence,
3111 first.opcode,
3112 first.element_type,
3113 destination,
3114 shard_count * world_size,
3115 payload,
3116 )?);
3117 }
3118 Ok(responses)
3119 }
3120 Opcode::AllToAll => {
3121 require_all_payloads(requests, rank_payload_bytes)?;
3122 if !count.is_multiple_of(world_size) {
3123 return Err(NetworkError::InvalidConfiguration(
3124 "all-to-all count is not divisible by world size".into(),
3125 ));
3126 }
3127 let shard_count = count / world_size;
3128 let shard_bytes = shard_count * element_bytes;
3129 let mut responses = Vec::with_capacity(world_size);
3130 for destination in 0..world_size {
3131 let mut payload = Vec::with_capacity(rank_payload_bytes);
3132 for request in requests {
3133 let start = destination * shard_bytes;
3134 payload.extend_from_slice(&request.payload[start..start + shard_bytes]);
3135 }
3136 responses.push(data_response(
3137 requests,
3138 unique_id,
3139 sequence,
3140 first.opcode,
3141 first.element_type,
3142 destination,
3143 count,
3144 payload,
3145 )?);
3146 }
3147 Ok(responses)
3148 }
3149 Opcode::AllToAllV => route_all_to_all_v(
3150 requests,
3151 unique_id,
3152 sequence,
3153 first.element_type,
3154 element_bytes,
3155 ),
3156 Opcode::Join
3157 | Opcode::Ready
3158 | Opcode::PeerEndpoint
3159 | Opcode::Send
3160 | Opcode::Receive
3161 | Opcode::Abort
3162 | Opcode::Heartbeat
3163 | Opcode::Leave => Err(NetworkError::InvalidConfiguration(format!(
3164 "opcode {:?} is not a coordinator collective request",
3165 first.opcode
3166 ))),
3167 }
3168}
3169
3170fn route_all_to_all_v(
3171 requests: &[Frame],
3172 unique_id: UniqueId,
3173 sequence: u64,
3174 element_type: ElementType,
3175 element_bytes: usize,
3176) -> Result<Vec<Frame>, NetworkError> {
3177 if element_type == ElementType::None || element_bytes == 0 {
3178 return Err(NetworkError::InvalidConfiguration(
3179 "all-to-all-v requires a concrete element type".into(),
3180 ));
3181 }
3182 let world_size = requests.len();
3183 let mut rows = Vec::with_capacity(world_size);
3184 for request in requests {
3185 rows.push(decode_counts_payload(request, world_size, element_bytes)?);
3186 }
3187 let mut responses = Vec::with_capacity(world_size);
3188 for destination in 0..world_size {
3189 let receive_counts = rows
3190 .iter()
3191 .map(|(counts, _)| counts[destination])
3192 .collect::<Vec<_>>();
3193 let receive_total = receive_counts
3194 .iter()
3195 .try_fold(0_usize, |total, count| total.checked_add(*count));
3196 let receive_total = receive_total.ok_or_else(|| {
3197 NetworkError::InvalidConfiguration("all-to-all-v receive count overflow".into())
3198 })?;
3199 let mut payload = encode_counts(&receive_counts)?;
3200 payload.reserve(receive_total.checked_mul(element_bytes).ok_or_else(|| {
3201 NetworkError::InvalidConfiguration("all-to-all-v response size overflow".into())
3202 })?);
3203 for (counts, data) in &rows {
3204 let start_elements = counts[..destination]
3205 .iter()
3206 .try_fold(0_usize, |total, count| total.checked_add(*count))
3207 .ok_or_else(|| {
3208 NetworkError::InvalidConfiguration("all-to-all-v source offset overflow".into())
3209 })?;
3210 let start = start_elements.checked_mul(element_bytes).ok_or_else(|| {
3211 NetworkError::InvalidConfiguration("all-to-all-v byte offset overflow".into())
3212 })?;
3213 let bytes = counts[destination]
3214 .checked_mul(element_bytes)
3215 .ok_or_else(|| {
3216 NetworkError::InvalidConfiguration("all-to-all-v shard size overflow".into())
3217 })?;
3218 payload.extend_from_slice(&data[start..start + bytes]);
3219 }
3220 let mut header = response_header(
3221 unique_id,
3222 Opcode::AllToAllV,
3223 element_type,
3224 sequence,
3225 destination,
3226 world_size,
3227 receive_total,
3228 );
3229 header.flags |= FLAG_COUNTS_PREFIX;
3230 responses.push(Frame::new(header, payload)?);
3231 }
3232 Ok(responses)
3233}
3234
3235fn decode_counts_payload(
3236 frame: &Frame,
3237 world_size: usize,
3238 element_bytes: usize,
3239) -> Result<(Vec<usize>, &[u8]), NetworkError> {
3240 if frame.header.flags & FLAG_COUNTS_PREFIX == 0 {
3241 return Err(NetworkError::InvalidConfiguration(
3242 "all-to-all-v frame is missing its counts prefix".into(),
3243 ));
3244 }
3245 let prefix_bytes = world_size.checked_mul(8).ok_or_else(|| {
3246 NetworkError::InvalidConfiguration("all-to-all-v prefix size overflow".into())
3247 })?;
3248 if frame.payload.len() < prefix_bytes {
3249 return Err(NetworkError::InvalidConfiguration(
3250 "all-to-all-v counts prefix is truncated".into(),
3251 ));
3252 }
3253 let counts = frame.payload[..prefix_bytes]
3254 .chunks_exact(8)
3255 .map(|bytes| {
3256 usize::try_from(u64::from_le_bytes(bytes.try_into().unwrap())).map_err(|_| {
3257 NetworkError::InvalidConfiguration("all-to-all-v count exceeds host usize".into())
3258 })
3259 })
3260 .collect::<Result<Vec<_>, _>>()?;
3261 let total = counts
3262 .iter()
3263 .try_fold(0_usize, |sum, count| sum.checked_add(*count));
3264 let total = total.ok_or_else(|| {
3265 NetworkError::InvalidConfiguration("all-to-all-v count sum overflow".into())
3266 })?;
3267 if total as u64 != frame.header.element_count {
3268 return Err(NetworkError::InvalidConfiguration(format!(
3269 "all-to-all-v counts sum to {total}, header declares {}",
3270 frame.header.element_count
3271 )));
3272 }
3273 let expected_data = total.checked_mul(element_bytes).ok_or_else(|| {
3274 NetworkError::InvalidConfiguration("all-to-all-v data size overflow".into())
3275 })?;
3276 let data = &frame.payload[prefix_bytes..];
3277 if data.len() != expected_data {
3278 return Err(NetworkError::InvalidConfiguration(format!(
3279 "all-to-all-v data has {} bytes, expected {expected_data}",
3280 data.len()
3281 )));
3282 }
3283 Ok((counts, data))
3284}
3285
3286fn encode_counts(counts: &[usize]) -> Result<Vec<u8>, NetworkError> {
3287 let mut encoded = Vec::with_capacity(counts.len().saturating_mul(8));
3288 for count in counts {
3289 encoded.extend_from_slice(
3290 &u64::try_from(*count)
3291 .map_err(|_| {
3292 NetworkError::InvalidConfiguration(
3293 "all-to-all-v count exceeds protocol u64".into(),
3294 )
3295 })?
3296 .to_le_bytes(),
3297 );
3298 }
3299 Ok(encoded)
3300}
3301
3302fn route_empty(
3303 requests: &[Frame],
3304 unique_id: UniqueId,
3305 sequence: u64,
3306 opcode: Opcode,
3307) -> Result<Vec<Frame>, NetworkError> {
3308 if requests.iter().any(|request| !request.payload.is_empty()) {
3309 return Err(NetworkError::InvalidConfiguration(format!(
3310 "{opcode:?} cannot carry a payload"
3311 )));
3312 }
3313 (0..requests.len())
3314 .map(|rank| {
3315 Frame::new(
3316 response_header(
3317 unique_id,
3318 opcode,
3319 ElementType::None,
3320 sequence,
3321 rank,
3322 requests.len(),
3323 0,
3324 ),
3325 Vec::new(),
3326 )
3327 .map_err(NetworkError::from)
3328 })
3329 .collect()
3330}
3331
3332fn route_same_payload(
3333 requests: &[Frame],
3334 unique_id: UniqueId,
3335 sequence: u64,
3336 opcode: Opcode,
3337 element_type: ElementType,
3338 count: usize,
3339 payload: &[u8],
3340) -> Result<Vec<Frame>, NetworkError> {
3341 (0..requests.len())
3342 .map(|rank| {
3343 data_response(
3344 requests,
3345 unique_id,
3346 sequence,
3347 opcode,
3348 element_type,
3349 rank,
3350 count,
3351 payload.to_vec(),
3352 )
3353 })
3354 .collect()
3355}
3356
3357#[allow(clippy::too_many_arguments)]
3358fn route_root_payload(
3359 requests: &[Frame],
3360 unique_id: UniqueId,
3361 sequence: u64,
3362 opcode: Opcode,
3363 element_type: ElementType,
3364 root: usize,
3365 count: usize,
3366 payload: &[u8],
3367) -> Result<Vec<Frame>, NetworkError> {
3368 (0..requests.len())
3369 .map(|rank| {
3370 if rank == root {
3371 data_response(
3372 requests,
3373 unique_id,
3374 sequence,
3375 opcode,
3376 element_type,
3377 rank,
3378 count,
3379 payload.to_vec(),
3380 )
3381 } else {
3382 Frame::new(
3383 response_header(
3384 unique_id,
3385 opcode,
3386 ElementType::None,
3387 sequence,
3388 rank,
3389 requests.len(),
3390 0,
3391 ),
3392 Vec::new(),
3393 )
3394 .map_err(NetworkError::from)
3395 }
3396 })
3397 .collect()
3398}
3399
3400#[allow(clippy::too_many_arguments)]
3401fn data_response(
3402 requests: &[Frame],
3403 unique_id: UniqueId,
3404 sequence: u64,
3405 opcode: Opcode,
3406 element_type: ElementType,
3407 destination: usize,
3408 count: usize,
3409 payload: Vec<u8>,
3410) -> Result<Frame, NetworkError> {
3411 Ok(Frame::new(
3412 response_header(
3413 unique_id,
3414 opcode,
3415 element_type,
3416 sequence,
3417 destination,
3418 requests.len(),
3419 count,
3420 ),
3421 payload,
3422 )?)
3423}
3424
3425fn response_header(
3426 unique_id: UniqueId,
3427 opcode: Opcode,
3428 element_type: ElementType,
3429 sequence: u64,
3430 destination: usize,
3431 world_size: usize,
3432 count: usize,
3433) -> FrameHeader {
3434 let mut header = FrameHeader::collective(
3435 unique_id,
3436 opcode,
3437 element_type,
3438 0,
3439 ANY_RANK,
3440 world_size as u32,
3441 sequence,
3442 count as u64,
3443 );
3444 header.destination_rank = destination as u32;
3445 header
3446}
3447
3448fn control_frame(
3449 unique_id: UniqueId,
3450 opcode: Opcode,
3451 sequence: u64,
3452 destination: u32,
3453 world_size: u32,
3454) -> Result<Frame, NetworkError> {
3455 let mut header = FrameHeader::collective(
3456 unique_id,
3457 opcode,
3458 ElementType::None,
3459 0,
3460 ANY_RANK,
3461 world_size,
3462 sequence,
3463 0,
3464 );
3465 header.destination_rank = destination;
3466 Ok(Frame::new(header, Vec::new())?)
3467}
3468
3469fn abort_frame(
3470 unique_id: UniqueId,
3471 sequence: u64,
3472 destination: u32,
3473 world_size: u32,
3474 message: &str,
3475) -> Result<Frame, NetworkError> {
3476 let payload = message.as_bytes().to_vec();
3477 let mut header = FrameHeader::collective(
3478 unique_id,
3479 Opcode::Abort,
3480 ElementType::U8,
3481 0,
3482 ANY_RANK,
3483 world_size,
3484 sequence,
3485 payload.len() as u64,
3486 );
3487 header.destination_rank = destination;
3488 Ok(Frame::new(header, payload)?)
3489}
3490
3491fn required_root(root: Option<usize>, world_size: usize) -> Result<usize, NetworkError> {
3492 match root {
3493 Some(root) if root < world_size => Ok(root),
3494 _ => Err(NetworkError::InvalidConfiguration(
3495 "rooted collective requires a valid root rank".into(),
3496 )),
3497 }
3498}
3499
3500fn require_all_payloads(requests: &[Frame], bytes: usize) -> Result<(), NetworkError> {
3501 if let Some((rank, actual)) = requests
3502 .iter()
3503 .enumerate()
3504 .map(|(rank, request)| (rank, request.payload.len()))
3505 .find(|(_, actual)| *actual != bytes)
3506 {
3507 return Err(NetworkError::InvalidConfiguration(format!(
3508 "rank {rank} payload has {actual} bytes, expected {bytes}"
3509 )));
3510 }
3511 Ok(())
3512}
3513
3514fn require_only_root_payload(
3515 requests: &[Frame],
3516 root: usize,
3517 bytes: usize,
3518) -> Result<(), NetworkError> {
3519 for (rank, request) in requests.iter().enumerate() {
3520 let expected = if rank == root { bytes } else { 0 };
3521 if request.payload.len() != expected {
3522 return Err(NetworkError::InvalidConfiguration(format!(
3523 "rank {rank} payload has {} bytes, expected {expected}",
3524 request.payload.len()
3525 )));
3526 }
3527 }
3528 Ok(())
3529}
3530
3531fn concatenate_payloads(requests: &[Frame]) -> Result<Vec<u8>, NetworkError> {
3532 let capacity = requests
3533 .iter()
3534 .try_fold(0_usize, |total, request| {
3535 total.checked_add(request.payload.len())
3536 })
3537 .ok_or_else(|| NetworkError::InvalidConfiguration("response payload overflow".into()))?;
3538 let mut payload = Vec::with_capacity(capacity);
3539 for request in requests {
3540 payload.extend_from_slice(&request.payload);
3541 }
3542 Ok(payload)
3543}
3544
3545pub(crate) fn write_frame(stream: &mut TcpStream, frame: &Frame) -> Result<(), NetworkError> {
3546 let header = frame.encode_transport_header()?;
3547 let total = header.len() + frame.payload.len();
3548 let written = stream.write_vectored(&[
3549 IoSlice::new(&header),
3550 IoSlice::new(frame.payload.as_slice()),
3551 ])?;
3552 if written < header.len() {
3553 stream.write_all(&header[written..])?;
3554 stream.write_all(&frame.payload)?;
3555 } else if written < total {
3556 stream.write_all(&frame.payload[written - header.len()..])?;
3557 }
3558 stream.flush()?;
3559 Ok(())
3560}
3561
3562pub(crate) fn read_frame(stream: &mut TcpStream) -> Result<Option<Frame>, NetworkError> {
3563 let mut header_bytes = [0_u8; super::protocol::FRAME_HEADER_BYTES];
3564 match stream.read(&mut header_bytes[..1]) {
3565 Ok(0) => return Ok(None),
3566 Ok(1) => {}
3567 Ok(_) => unreachable!("one-byte read cannot return more than one byte"),
3568 Err(error) => return Err(error.into()),
3569 }
3570 stream.read_exact(&mut header_bytes[1..])?;
3571 let header = FrameHeader::decode(&header_bytes)?;
3572 let payload_length = usize::try_from(header.payload_bytes).map_err(|_| {
3573 NetworkError::InvalidConfiguration("frame payload length exceeds usize".into())
3574 })?;
3575 let mut payload = vec![0_u8; payload_length];
3576 stream.read_exact(&mut payload)?;
3577 header.validate(Some(&payload))?;
3578 Ok(Some(Frame { header, payload }))
3579}
3580
3581#[cfg(test)]
3582mod tests {
3583 use super::*;
3584 use std::thread;
3585
3586 #[test]
3587 fn topology_probe_metrics_use_stable_conservative_buckets() {
3588 assert_eq!(conservative_probe_bandwidth_bucket(1), 1);
3589 assert_eq!(conservative_probe_bandwidth_bucket(63), 16);
3590 assert_eq!(conservative_probe_bandwidth_bucket(64), 64);
3591 assert_eq!(conservative_probe_bandwidth_bucket(255), 64);
3592 assert_eq!(conservative_probe_bandwidth_bucket(256), 256);
3593 assert_eq!(conservative_probe_latency_bucket(1), 4);
3594 assert_eq!(conservative_probe_latency_bucket(32_769), 65_536);
3595 assert_eq!(conservative_probe_latency_bucket(65_535), 65_536);
3596 assert_eq!(conservative_probe_latency_bucket(65_536), 262_144);
3597 }
3598
3599 #[test]
3600 fn peer_rail_address_configuration_requires_one_address_per_rail() {
3601 let listen =
3602 parse_rail_socket_addresses("GX1_P2P_RAIL_LISTEN_ADDRS", "127.0.0.1:0, 127.0.0.2:0", 2)
3603 .unwrap();
3604 assert_eq!(listen.len(), 2);
3605 assert_eq!(listen[0].ip(), IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)));
3606 assert_eq!(listen[1].ip(), IpAddr::V4(Ipv4Addr::new(127, 0, 0, 2)));
3607
3608 let advertise =
3609 parse_rail_ip_addresses("GX1_P2P_RAIL_ADVERTISE_ADDRS", "127.0.0.1,127.0.0.2", 2)
3610 .unwrap();
3611 assert_eq!(advertise.len(), 2);
3612 assert!(
3613 parse_rail_socket_addresses("GX1_P2P_RAIL_LISTEN_ADDRS", "127.0.0.1:0", 2,)
3614 .unwrap_err()
3615 .to_string()
3616 .contains("expected 2")
3617 );
3618 assert!(
3619 parse_rail_ip_addresses("GX1_P2P_RAIL_ADVERTISE_ADDRS", "127.0.0.1,not-an-ip", 2,)
3620 .unwrap_err()
3621 .to_string()
3622 .contains("rail 1")
3623 );
3624 }
3625
3626 #[test]
3627 fn peer_rail_address_configuration_binds_dedicated_listeners() {
3628 let configuration = PeerEndpointConfiguration {
3629 listen_addresses: vec![
3630 "127.0.0.1:0".parse().unwrap(),
3631 "127.0.0.2:0".parse().unwrap(),
3632 ],
3633 advertise_addresses: vec![
3634 IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)),
3635 IpAddr::V4(Ipv4Addr::new(127, 0, 0, 2)),
3636 ],
3637 };
3638 let (listeners, endpoints) =
3639 bind_peer_listeners_with_configuration(configuration, 2).unwrap();
3640 assert_eq!(listeners.len(), 2);
3641 assert_eq!(endpoints.len(), 2);
3642 assert_eq!(endpoints[0].ip(), IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)));
3643 assert_eq!(endpoints[1].ip(), IpAddr::V4(Ipv4Addr::new(127, 0, 0, 2)));
3644 assert_ne!(endpoints[0].port(), 0);
3645 assert_ne!(endpoints[1].port(), 0);
3646 }
3647
3648 #[test]
3649 fn tcp_rendezvous_routes_all_gather_all_to_all_and_barrier() {
3650 let unique_id = UniqueId::from_bytes([9; 16]);
3651 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 3).unwrap();
3652 let address = server.local_addr().unwrap();
3653 let coordinator = thread::spawn(move || server.run());
3654 let ranks = (0..3_u32)
3655 .map(|rank| {
3656 thread::spawn(move || {
3657 let session = TcpRankSession::connect(
3658 address,
3659 unique_id,
3660 rank,
3661 3,
3662 Duration::from_secs(5),
3663 )
3664 .unwrap();
3665 let payload = [rank * 10, rank * 10 + 1]
3666 .into_iter()
3667 .flat_map(u32::to_le_bytes)
3668 .collect::<Vec<_>>();
3669 let gathered = session
3670 .exchange(Opcode::AllGather, ElementType::U32, ANY_RANK, 2, payload)
3671 .unwrap();
3672 let values = gathered
3673 .payload
3674 .chunks_exact(4)
3675 .map(|bytes| u32::from_le_bytes(bytes.try_into().unwrap()))
3676 .collect::<Vec<_>>();
3677 assert_eq!(values, vec![0, 1, 10, 11, 20, 21]);
3678
3679 let all_to_all_input = (0..3_u32)
3680 .map(|destination| rank * 100 + destination)
3681 .flat_map(u32::to_le_bytes)
3682 .collect::<Vec<_>>();
3683 let all_to_all = session
3684 .exchange(
3685 Opcode::AllToAll,
3686 ElementType::U32,
3687 ANY_RANK,
3688 3,
3689 all_to_all_input,
3690 )
3691 .unwrap();
3692 let values = all_to_all
3693 .payload
3694 .chunks_exact(4)
3695 .map(|bytes| u32::from_le_bytes(bytes.try_into().unwrap()))
3696 .collect::<Vec<_>>();
3697 assert_eq!(values, vec![rank, 100 + rank, 200 + rank]);
3698 session.barrier().unwrap();
3699 })
3700 })
3701 .collect::<Vec<_>>();
3702 for rank in ranks {
3703 rank.join().unwrap();
3704 }
3705 coordinator.join().unwrap().unwrap();
3706 }
3707
3708 #[test]
3709 fn tcp_point_to_point_is_tagged_multiplexed_and_collective_independent() {
3710 let unique_id = UniqueId::from_bytes([7; 16]);
3711 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 3).unwrap();
3712 let address = server.local_addr().unwrap();
3713 let coordinator = thread::spawn(move || server.run());
3714 let ranks = (0..3_u32)
3715 .map(|rank| {
3716 thread::spawn(move || {
3717 let session = Arc::new(
3718 TcpRankSession::connect(
3719 address,
3720 unique_id,
3721 rank,
3722 3,
3723 Duration::from_secs(5),
3724 )
3725 .unwrap(),
3726 );
3727 match rank {
3728 0 => {
3729 session
3730 .send(
3731 1,
3732 101,
3733 ElementType::U32,
3734 2,
3735 [101_u32, 102]
3736 .into_iter()
3737 .flat_map(u32::to_le_bytes)
3738 .collect(),
3739 )
3740 .unwrap();
3741 session
3742 .send(
3743 1,
3744 100,
3745 ElementType::U32,
3746 2,
3747 [100_u32, 101]
3748 .into_iter()
3749 .flat_map(u32::to_le_bytes)
3750 .collect(),
3751 )
3752 .unwrap();
3753 }
3754 1 => {
3755 let receives = [100_u64, 101]
3756 .into_iter()
3757 .map(|tag| {
3758 let session = Arc::clone(&session);
3759 thread::spawn(move || {
3760 let frame = session
3761 .receive(Some(0), tag, ElementType::U32, 2)
3762 .unwrap();
3763 let values = frame
3764 .payload
3765 .chunks_exact(4)
3766 .map(|bytes| {
3767 u32::from_le_bytes(bytes.try_into().unwrap())
3768 })
3769 .collect::<Vec<_>>();
3770 (tag, values)
3771 })
3772 })
3773 .collect::<Vec<_>>();
3774 let mut received = receives
3775 .into_iter()
3776 .map(|receive| receive.join().unwrap())
3777 .collect::<Vec<_>>();
3778 received.sort_by_key(|(tag, _)| *tag);
3779 assert_eq!(
3780 received,
3781 vec![(100, vec![100, 101]), (101, vec![101, 102])]
3782 );
3783 }
3784 2 => {}
3785 _ => unreachable!(),
3786 }
3787 let gathered = session
3788 .exchange(
3789 Opcode::AllGather,
3790 ElementType::U32,
3791 ANY_RANK,
3792 1,
3793 rank.to_le_bytes().to_vec(),
3794 )
3795 .unwrap();
3796 assert_eq!(
3797 gathered
3798 .payload
3799 .chunks_exact(4)
3800 .map(|bytes| u32::from_le_bytes(bytes.try_into().unwrap()))
3801 .collect::<Vec<_>>(),
3802 vec![0, 1, 2]
3803 );
3804 })
3805 })
3806 .collect::<Vec<_>>();
3807 for rank in ranks {
3808 rank.join().unwrap();
3809 }
3810 coordinator.join().unwrap().unwrap();
3811 }
3812
3813 #[test]
3814 fn tcp_peer_mesh_routes_rank_to_rank_data_on_multiple_rails() {
3815 let unique_id = UniqueId::from_bytes([0x51; 16]);
3816 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 3)
3817 .unwrap()
3818 .with_p2p_rails(2)
3819 .unwrap()
3820 .with_transport(CollectiveTransport::TcpPeer)
3821 .unwrap();
3822 let address = server.local_addr().unwrap();
3823 let coordinator = thread::spawn(move || server.run());
3824 let ranks = (0..3_u32)
3825 .map(|rank| {
3826 thread::spawn(move || {
3827 let session = TcpRankSession::connect_with_transport(
3828 address,
3829 unique_id,
3830 rank,
3831 3,
3832 Duration::from_secs(5),
3833 2,
3834 CollectiveTransport::TcpPeer,
3835 )
3836 .unwrap();
3837 assert_eq!(session.transport(), CollectiveTransport::TcpPeer);
3838 assert_eq!(session.p2p_rails(), 2);
3839 match rank {
3840 0 => {
3841 session
3842 .send_on_rail(
3843 1,
3844 1,
3845 101,
3846 ElementType::U32,
3847 2,
3848 [10_u32, 11]
3849 .into_iter()
3850 .flat_map(u32::to_le_bytes)
3851 .collect(),
3852 )
3853 .unwrap();
3854 let frame = session
3855 .receive_on_rail(0, Some(2), 202, ElementType::U32, 1)
3856 .unwrap();
3857 assert_eq!(frame.payload, 22_u32.to_le_bytes());
3858 }
3859 1 => {
3860 let frame = session
3861 .receive_on_rail(1, None, 101, ElementType::U32, 2)
3862 .unwrap();
3863 assert_eq!(
3864 frame
3865 .payload
3866 .chunks_exact(4)
3867 .map(|bytes| u32::from_le_bytes(bytes.try_into().unwrap()))
3868 .collect::<Vec<_>>(),
3869 vec![10, 11]
3870 );
3871 session
3872 .send_on_rail(
3873 0,
3874 2,
3875 303,
3876 ElementType::U32,
3877 1,
3878 13_u32.to_le_bytes().to_vec(),
3879 )
3880 .unwrap();
3881 }
3882 2 => {
3883 session
3884 .send_on_rail(
3885 0,
3886 0,
3887 202,
3888 ElementType::U32,
3889 1,
3890 22_u32.to_le_bytes().to_vec(),
3891 )
3892 .unwrap();
3893 let frame = session
3894 .receive_on_rail(0, Some(1), 303, ElementType::U32, 1)
3895 .unwrap();
3896 assert_eq!(frame.payload, 13_u32.to_le_bytes());
3897 }
3898 _ => unreachable!(),
3899 }
3900 let gathered = session
3901 .exchange(
3902 Opcode::AllGather,
3903 ElementType::U32,
3904 ANY_RANK,
3905 1,
3906 rank.to_le_bytes().to_vec(),
3907 )
3908 .unwrap();
3909 assert_eq!(
3910 gathered
3911 .payload
3912 .chunks_exact(4)
3913 .map(|bytes| u32::from_le_bytes(bytes.try_into().unwrap()))
3914 .collect::<Vec<_>>(),
3915 vec![0, 1, 2]
3916 );
3917 })
3918 })
3919 .collect::<Vec<_>>();
3920 for rank in ranks {
3921 rank.join().unwrap();
3922 }
3923 coordinator.join().unwrap().unwrap();
3924 }
3925
3926 #[test]
3927 fn tcp_peer_mesh_routes_each_rail_through_its_own_listen_address() {
3928 let unique_id = UniqueId::from_bytes([0x5a; 16]);
3929 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 2)
3930 .unwrap()
3931 .with_p2p_rails(2)
3932 .unwrap()
3933 .with_transport(CollectiveTransport::TcpPeer)
3934 .unwrap();
3935 let address = server.local_addr().unwrap();
3936 let coordinator = thread::spawn(move || server.run());
3937 let ranks = (0..2_u32)
3938 .map(|rank| {
3939 thread::spawn(move || {
3940 let configuration = PeerEndpointConfiguration {
3941 listen_addresses: vec![
3942 "127.0.0.1:0".parse().unwrap(),
3943 "127.0.0.2:0".parse().unwrap(),
3944 ],
3945 advertise_addresses: vec![
3946 IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)),
3947 IpAddr::V4(Ipv4Addr::new(127, 0, 0, 2)),
3948 ],
3949 };
3950 let session = TcpRankSession::connect_with_transport_configuration(
3951 address,
3952 unique_id,
3953 rank,
3954 2,
3955 Duration::from_secs(5),
3956 2,
3957 CollectiveTransport::TcpPeer,
3958 Some(configuration),
3959 )
3960 .unwrap();
3961 if rank == 0 {
3962 let first = session
3963 .receive_on_rail(0, Some(1), 700, ElementType::U32, 1)
3964 .unwrap();
3965 let second = session
3966 .receive_on_rail(1, Some(1), 701, ElementType::U32, 1)
3967 .unwrap();
3968 assert_eq!(first.payload, 10_u32.to_le_bytes());
3969 assert_eq!(second.payload, 11_u32.to_le_bytes());
3970 } else {
3971 session
3972 .send_on_rail(
3973 0,
3974 0,
3975 700,
3976 ElementType::U32,
3977 1,
3978 10_u32.to_le_bytes().to_vec(),
3979 )
3980 .unwrap();
3981 session
3982 .send_on_rail(
3983 1,
3984 0,
3985 701,
3986 ElementType::U32,
3987 1,
3988 11_u32.to_le_bytes().to_vec(),
3989 )
3990 .unwrap();
3991 }
3992 session
3993 .exchange(Opcode::Barrier, ElementType::None, ANY_RANK, 0, Vec::new())
3994 .unwrap();
3995 })
3996 })
3997 .collect::<Vec<_>>();
3998 for rank in ranks {
3999 rank.join().unwrap();
4000 }
4001 coordinator.join().unwrap().unwrap();
4002 }
4003
4004 #[test]
4005 fn tcp_peer_probe_builds_the_same_complete_topology_on_every_rank() {
4006 let unique_id = UniqueId::from_bytes([0x57; 16]);
4007 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 4)
4008 .unwrap()
4009 .with_p2p_rails(2)
4010 .unwrap()
4011 .with_transport(CollectiveTransport::TcpPeer)
4012 .unwrap();
4013 let address = server.local_addr().unwrap();
4014 let coordinator = thread::spawn(move || server.run());
4015 let options = TopologyProbeOptions {
4016 payload_bytes: 16 * 1024,
4017 latency_iterations: 3,
4018 bandwidth_iterations: 2,
4019 warmup_iterations: 0,
4020 };
4021 let ranks = (0..4_u32)
4022 .map(|rank| {
4023 thread::spawn(move || {
4024 let session = TcpRankSession::connect_with_transport(
4025 address,
4026 unique_id,
4027 rank,
4028 4,
4029 Duration::from_secs(5),
4030 2,
4031 CollectiveTransport::TcpPeer,
4032 )
4033 .unwrap();
4034 let topology = session.probe_topology(options).unwrap();
4035 let links = topology.links();
4036 assert_eq!(links.len(), 6);
4037 assert!(
4038 links
4039 .iter()
4040 .all(|link| { link.bandwidth_mbps > 0 && link.latency_ns > 0 })
4041 );
4042 let rail_links = topology.rail_links();
4043 assert_eq!(rail_links.len(), 12);
4044 assert!(rail_links.iter().all(|link| {
4045 link.rail < 2 && link.bandwidth_mbps > 0 && link.latency_ns > 0
4046 }));
4047 let aggregate_links = topology.aggregate_links();
4048 assert_eq!(aggregate_links.len(), 6);
4049 assert!(aggregate_links.iter().all(|link| link.bandwidth_mbps > 0));
4050 (
4051 links,
4052 rail_links,
4053 aggregate_links,
4054 topology.best_ring_order().unwrap(),
4055 )
4056 })
4057 })
4058 .collect::<Vec<_>>();
4059 let results = ranks
4060 .into_iter()
4061 .map(|rank| rank.join().unwrap())
4062 .collect::<Vec<_>>();
4063 for result in &results[1..] {
4064 assert_eq!(result, &results[0]);
4065 }
4066 coordinator.join().unwrap().unwrap();
4067 }
4068
4069 #[test]
4070 fn topology_probe_rejects_coordinator_data_transport() {
4071 let unique_id = UniqueId::from_bytes([0x58; 16]);
4072 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 1).unwrap();
4073 let address = server.local_addr().unwrap();
4074 let coordinator = thread::spawn(move || server.run());
4075 let session = TcpRankSession::connect_with_transport(
4076 address,
4077 unique_id,
4078 0,
4079 1,
4080 Duration::from_secs(2),
4081 1,
4082 CollectiveTransport::TcpHostStaged,
4083 )
4084 .unwrap();
4085 let error = session
4086 .probe_topology(TopologyProbeOptions::default())
4087 .unwrap_err();
4088 assert!(
4089 error
4090 .to_string()
4091 .contains("requires GX1_COLLECTIVE_TRANSPORT=tcp_peer")
4092 );
4093 drop(session);
4094 coordinator.join().unwrap().unwrap();
4095 }
4096
4097 #[test]
4098 fn peer_abort_interrupts_topology_probe_on_all_other_ranks() {
4099 let unique_id = UniqueId::from_bytes([0x59; 16]);
4100 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 3)
4101 .unwrap()
4102 .with_transport(CollectiveTransport::TcpPeer)
4103 .unwrap();
4104 let address = server.local_addr().unwrap();
4105 let coordinator = thread::spawn(move || server.run());
4106 let options = TopologyProbeOptions {
4107 payload_bytes: 4096,
4108 latency_iterations: 3,
4109 bandwidth_iterations: 2,
4110 warmup_iterations: 0,
4111 };
4112 let connected = Arc::new(std::sync::Barrier::new(3));
4113 let ranks = (0..3_u32)
4114 .map(|rank| {
4115 let connected = Arc::clone(&connected);
4116 thread::spawn(move || {
4117 let session = TcpRankSession::connect_with_transport(
4118 address,
4119 unique_id,
4120 rank,
4121 3,
4122 Duration::from_secs(3),
4123 1,
4124 CollectiveTransport::TcpPeer,
4125 )
4126 .unwrap();
4127 connected.wait();
4128 if rank == 2 {
4129 session.abort("topology probe failure injection").unwrap();
4130 return;
4131 }
4132 let error = session.probe_topology(options).unwrap_err();
4133 assert!(
4134 matches!(error, NetworkError::RemoteAbort(_)),
4135 "topology probe returned a non-abort failure: {error:?}"
4136 );
4137 assert!(
4138 error
4139 .to_string()
4140 .contains("topology probe failure injection"),
4141 "unexpected probe abort reason: {error}"
4142 );
4143 })
4144 })
4145 .collect::<Vec<_>>();
4146 for rank in ranks {
4147 rank.join().unwrap();
4148 }
4149 assert!(coordinator.join().unwrap().is_err());
4150 }
4151
4152 #[test]
4153 fn tcp_peer_abort_interrupts_pending_receive_on_another_rail() {
4154 let unique_id = UniqueId::from_bytes([0x52; 16]);
4155 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 2)
4156 .unwrap()
4157 .with_p2p_rails(2)
4158 .unwrap()
4159 .with_transport(CollectiveTransport::TcpPeer)
4160 .unwrap();
4161 let address = server.local_addr().unwrap();
4162 let coordinator = thread::spawn(move || server.run());
4163 let (ready_sender, ready_receiver) = mpsc::channel();
4164 let receiver = thread::spawn(move || {
4165 let session = TcpRankSession::connect_with_transport(
4166 address,
4167 unique_id,
4168 0,
4169 2,
4170 Duration::from_secs(3),
4171 2,
4172 CollectiveTransport::TcpPeer,
4173 )
4174 .unwrap();
4175 ready_sender.send(()).unwrap();
4176 let started = Instant::now();
4177 let error = session
4178 .receive_on_rail(1, Some(1), 404, ElementType::U32, 1)
4179 .unwrap_err();
4180 assert!(
4181 matches!(error, NetworkError::RemoteAbort(_)),
4182 "unexpected disconnect error: {error:?}"
4183 );
4184 assert!(error.to_string().contains("injected direct peer failure"));
4185 assert!(started.elapsed() < Duration::from_secs(1));
4186 });
4187 let aborter = thread::spawn(move || {
4188 let session = TcpRankSession::connect_with_transport(
4189 address,
4190 unique_id,
4191 1,
4192 2,
4193 Duration::from_secs(3),
4194 2,
4195 CollectiveTransport::TcpPeer,
4196 )
4197 .unwrap();
4198 ready_receiver.recv().unwrap();
4199 session.abort("injected direct peer failure").unwrap();
4200 });
4201 receiver.join().unwrap();
4202 aborter.join().unwrap();
4203 assert!(coordinator.join().unwrap().is_err());
4204 }
4205
4206 #[test]
4207 fn eight_rank_peer_abort_fans_out_across_both_data_rails() {
4208 const WORLD_SIZE: u32 = 8;
4209
4210 let unique_id = UniqueId::from_bytes([0x74; 16]);
4211 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, WORLD_SIZE as usize)
4212 .unwrap()
4213 .with_p2p_rails(2)
4214 .unwrap()
4215 .with_transport(CollectiveTransport::TcpPeer)
4216 .unwrap();
4217 let address = server.local_addr().unwrap();
4218 let coordinator = thread::spawn(move || server.run());
4219 let connected = Arc::new(std::sync::Barrier::new(WORLD_SIZE as usize));
4220 let ranks = (0..WORLD_SIZE)
4221 .map(|rank| {
4222 let connected = Arc::clone(&connected);
4223 thread::spawn(move || {
4224 let session = TcpRankSession::connect_with_transport(
4225 address,
4226 unique_id,
4227 rank,
4228 WORLD_SIZE,
4229 Duration::from_secs(5),
4230 2,
4231 CollectiveTransport::TcpPeer,
4232 )
4233 .unwrap();
4234 connected.wait();
4235 if rank == WORLD_SIZE - 1 {
4236 session.abort("injected eight-rank peer failure").unwrap();
4237 return;
4238 }
4239
4240 let rail = rank as usize % 2;
4241 let source = (rank + 1) % (WORLD_SIZE - 1);
4242 let started = Instant::now();
4243 let error = session
4244 .receive_on_rail(
4245 rail,
4246 Some(source),
4247 0x8000 + u64::from(rank),
4248 ElementType::U32,
4249 1,
4250 )
4251 .unwrap_err();
4252 assert!(
4253 matches!(error, NetworkError::RemoteAbort(_)),
4254 "rank {rank} returned a non-abort failure: {error:?}"
4255 );
4256 assert!(
4257 error
4258 .to_string()
4259 .contains("injected eight-rank peer failure")
4260 );
4261 assert!(started.elapsed() < Duration::from_secs(1));
4262 })
4263 })
4264 .collect::<Vec<_>>();
4265 for rank in ranks {
4266 rank.join().unwrap();
4267 }
4268 assert!(coordinator.join().unwrap().is_err());
4269 }
4270
4271 #[test]
4272 fn tcp_peer_endpoint_exchange_times_out_when_a_rank_never_publishes() {
4273 let unique_id = UniqueId::from_bytes([0x54; 16]);
4274 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 2)
4275 .unwrap()
4276 .with_collective_timeout(Duration::from_millis(120))
4277 .unwrap()
4278 .with_transport(CollectiveTransport::TcpPeer)
4279 .unwrap();
4280 let address = server.local_addr().unwrap();
4281 let coordinator = thread::spawn(move || server.run());
4282 let publishing = thread::spawn(move || {
4283 let started = Instant::now();
4284 let result = TcpRankSession::connect_with_transport(
4285 address,
4286 unique_id,
4287 0,
4288 2,
4289 Duration::from_secs(2),
4290 1,
4291 CollectiveTransport::TcpPeer,
4292 );
4293 assert!(result.is_err());
4294 assert!(started.elapsed() < Duration::from_secs(1));
4295 });
4296 let stalled = thread::spawn(move || {
4297 let addresses = [address];
4298 let _stream = connect_rank_channel(
4299 &addresses,
4300 unique_id,
4301 1,
4302 2,
4303 Duration::from_secs(2),
4304 false,
4305 0,
4306 )
4307 .unwrap();
4308 thread::sleep(Duration::from_millis(250));
4309 });
4310 publishing.join().unwrap();
4311 stalled.join().unwrap();
4312 assert!(coordinator.join().unwrap().is_err());
4313 }
4314
4315 #[test]
4316 fn tcp_peer_mesh_setup_failure_is_broadcast_on_control_plane() {
4317 let unique_id = UniqueId::from_bytes([0x56; 16]);
4318 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 2)
4319 .unwrap()
4320 .with_transport(CollectiveTransport::TcpPeer)
4321 .unwrap();
4322 let address = server.local_addr().unwrap();
4323 let coordinator = thread::spawn(move || server.run());
4324 let fake_rank = thread::spawn(move || {
4325 let addresses = [address];
4326 let mut collective = connect_rank_channel(
4327 &addresses,
4328 unique_id,
4329 0,
4330 2,
4331 Duration::from_secs(2),
4332 false,
4333 0,
4334 )
4335 .unwrap();
4336 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
4337 let unreachable = listener.local_addr().unwrap();
4338 drop(listener);
4339 exchange_peer_endpoints_client(&mut collective, &[unreachable], unique_id, 0, 2, 1)
4340 .unwrap();
4341 let mut control =
4342 connect_rank_channel(&addresses, unique_id, 0, 2, Duration::from_secs(2), true, 0)
4343 .unwrap();
4344 let abort = read_frame(&mut control)
4345 .unwrap()
4346 .expect("control plane returns setup abort");
4347 assert_eq!(abort.header.opcode, Opcode::Abort);
4348 assert!(
4349 String::from_utf8_lossy(&abort.payload)
4350 .contains("rank 1 direct peer mesh setup failed")
4351 );
4352 });
4353 let failing_rank = thread::spawn(move || {
4354 let started = Instant::now();
4355 let result = TcpRankSession::connect_with_transport(
4356 address,
4357 unique_id,
4358 1,
4359 2,
4360 Duration::from_millis(300),
4361 1,
4362 CollectiveTransport::TcpPeer,
4363 );
4364 assert!(result.is_err());
4365 assert!(started.elapsed() < Duration::from_secs(1));
4366 });
4367 fake_rank.join().unwrap();
4368 failing_rank.join().unwrap();
4369 assert!(coordinator.join().unwrap().is_err());
4370 }
4371
4372 #[test]
4373 fn tcp_peer_disconnect_wakes_pending_receive_without_waiting_for_timeout() {
4374 let unique_id = UniqueId::from_bytes([0x55; 16]);
4375 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 2)
4376 .unwrap()
4377 .with_transport(CollectiveTransport::TcpPeer)
4378 .unwrap();
4379 let address = server.local_addr().unwrap();
4380 let coordinator = thread::spawn(move || server.run());
4381 let (ready_sender, ready_receiver) = mpsc::channel();
4382 let receiver = thread::spawn(move || {
4383 let session = TcpRankSession::connect_with_transport(
4384 address,
4385 unique_id,
4386 0,
4387 2,
4388 Duration::from_secs(3),
4389 1,
4390 CollectiveTransport::TcpPeer,
4391 )
4392 .unwrap();
4393 ready_sender.send(()).unwrap();
4394 let started = Instant::now();
4395 let error = session
4396 .receive(Some(1), 505, ElementType::U32, 1)
4397 .unwrap_err();
4398 assert!(
4399 matches!(error, NetworkError::RemoteAbort(_)),
4400 "unexpected disconnect error: {error:?}"
4401 );
4402 assert!(error.to_string().contains("direct peer rank 1 left"));
4403 assert!(started.elapsed() < Duration::from_secs(1));
4404 });
4405 let disconnecting = thread::spawn(move || {
4406 let session = TcpRankSession::connect_with_transport(
4407 address,
4408 unique_id,
4409 1,
4410 2,
4411 Duration::from_secs(3),
4412 1,
4413 CollectiveTransport::TcpPeer,
4414 )
4415 .unwrap();
4416 ready_receiver.recv().unwrap();
4417 drop(session);
4418 });
4419 receiver.join().unwrap();
4420 disconnecting.join().unwrap();
4421 assert!(coordinator.join().unwrap().is_err());
4422 }
4423
4424 #[test]
4425 fn tcp_rendezvous_aborts_mismatched_collective_order() {
4426 let unique_id = UniqueId::from_bytes([5; 16]);
4427 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 2).unwrap();
4428 let address = server.local_addr().unwrap();
4429 let coordinator = thread::spawn(move || server.run());
4430 let ranks = (0..2_u32)
4431 .map(|rank| {
4432 thread::spawn(move || {
4433 let session = TcpRankSession::connect(
4434 address,
4435 unique_id,
4436 rank,
4437 2,
4438 Duration::from_secs(5),
4439 )
4440 .unwrap();
4441 let opcode = if rank == 0 {
4442 Opcode::AllGather
4443 } else {
4444 Opcode::AllReduce
4445 };
4446 assert!(matches!(
4447 session.exchange(
4448 opcode,
4449 ElementType::U32,
4450 ANY_RANK,
4451 1,
4452 rank.to_le_bytes().to_vec()
4453 ),
4454 Err(NetworkError::RemoteAbort(_))
4455 ));
4456 })
4457 })
4458 .collect::<Vec<_>>();
4459 for rank in ranks {
4460 rank.join().unwrap();
4461 }
4462 assert!(coordinator.join().unwrap().is_err());
4463 }
4464
4465 #[test]
4466 fn collective_timeout_names_missing_rank_and_aborts_waiter() {
4467 let unique_id = UniqueId::from_bytes([0x31; 16]);
4468 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 2)
4469 .unwrap()
4470 .with_collective_timeout(Duration::from_millis(150))
4471 .unwrap();
4472 let address = server.local_addr().unwrap();
4473 let coordinator = thread::spawn(move || server.run());
4474 let waiting = thread::spawn(move || {
4475 let session =
4476 TcpRankSession::connect(address, unique_id, 0, 2, Duration::from_secs(2)).unwrap();
4477 session.barrier()
4478 });
4479 let missing = thread::spawn(move || {
4480 let _session =
4481 TcpRankSession::connect(address, unique_id, 1, 2, Duration::from_secs(2)).unwrap();
4482 thread::sleep(Duration::from_millis(300));
4483 });
4484 let error = waiting.join().unwrap().unwrap_err();
4485 assert!(matches!(error, NetworkError::RemoteAbort(_)));
4486 assert!(error.to_string().contains("missing ranks [1]"));
4487 missing.join().unwrap();
4488 assert!(matches!(
4489 coordinator.join().unwrap(),
4490 Err(NetworkError::Timeout(_))
4491 ));
4492 }
4493
4494 #[test]
4495 fn explicit_rank_abort_interrupts_pending_point_to_point_work() {
4496 let unique_id = UniqueId::from_bytes([0x32; 16]);
4497 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 2).unwrap();
4498 let address = server.local_addr().unwrap();
4499 let coordinator = thread::spawn(move || server.run());
4500 let waiting = thread::spawn(move || {
4501 let session =
4502 TcpRankSession::connect(address, unique_id, 0, 2, Duration::from_secs(2)).unwrap();
4503 session.receive(Some(1), 99, ElementType::U32, 1)
4504 });
4505 let aborting = thread::spawn(move || {
4506 let session =
4507 TcpRankSession::connect(address, unique_id, 1, 2, Duration::from_secs(2)).unwrap();
4508 thread::sleep(Duration::from_millis(25));
4509 session.abort("injected rank failure").unwrap();
4510 thread::sleep(Duration::from_millis(25));
4511 });
4512 let error = waiting.join().unwrap().unwrap_err();
4513 assert!(matches!(error, NetworkError::RemoteAbort(_)));
4514 assert!(error.to_string().contains("injected rank failure"));
4515 aborting.join().unwrap();
4516 assert!(coordinator.join().unwrap().is_err());
4517 }
4518
4519 #[test]
4520 fn explicit_rank_abort_interrupts_collective_without_closing_its_session() {
4521 for (index, transport) in [
4522 CollectiveTransport::TcpHostStaged,
4523 CollectiveTransport::TcpPeer,
4524 ]
4525 .into_iter()
4526 .enumerate()
4527 {
4528 let unique_id = UniqueId::from_bytes([0x5a + index as u8; 16]);
4529 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 2)
4530 .unwrap()
4531 .with_transport(transport)
4532 .unwrap();
4533 let address = server.local_addr().unwrap();
4534 let coordinator = thread::spawn(move || server.run());
4535 let (ready_sender, ready_receiver) = mpsc::channel();
4536 let (release_sender, release_receiver) = mpsc::channel();
4537 let waiting = thread::spawn(move || {
4538 let session = TcpRankSession::connect_with_transport(
4539 address,
4540 unique_id,
4541 0,
4542 2,
4543 Duration::from_secs(2),
4544 1,
4545 transport,
4546 )
4547 .unwrap();
4548 ready_sender.send(()).unwrap();
4549 let error = session.barrier().unwrap_err();
4550 assert!(matches!(error, NetworkError::RemoteAbort(_)));
4551 assert!(
4552 error.to_string().contains("collective failure reason"),
4553 "unexpected collective abort reason: {error}"
4554 );
4555 release_sender.send(()).unwrap();
4556 });
4557 let aborting = thread::spawn(move || {
4558 let session = TcpRankSession::connect_with_transport(
4559 address,
4560 unique_id,
4561 1,
4562 2,
4563 Duration::from_secs(2),
4564 1,
4565 transport,
4566 )
4567 .unwrap();
4568 ready_receiver.recv().unwrap();
4569 session.abort("collective failure reason").unwrap();
4570 release_receiver.recv().unwrap();
4571 drop(session);
4572 });
4573 waiting.join().unwrap();
4574 aborting.join().unwrap();
4575 let error = coordinator.join().unwrap().unwrap_err();
4576 assert!(error.to_string().contains("collective failure reason"));
4577 }
4578 }
4579
4580 #[test]
4581 fn abort_is_broadcast_to_pending_work_on_every_point_to_point_rail() {
4582 let unique_id = UniqueId::from_bytes([0x46; 16]);
4583 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 2)
4584 .unwrap()
4585 .with_p2p_rails(2)
4586 .unwrap();
4587 let address = server.local_addr().unwrap();
4588 let coordinator = thread::spawn(move || server.run());
4589 let (ready_sender, ready_receiver) = mpsc::channel();
4590 let receiver = thread::spawn(move || {
4591 let session = TcpRankSession::connect_with_p2p_rails(
4592 address,
4593 unique_id,
4594 0,
4595 2,
4596 Duration::from_secs(2),
4597 2,
4598 )
4599 .unwrap();
4600 ready_sender.send(()).unwrap();
4601 let error = session
4602 .receive_on_rail(1, Some(1), 77, ElementType::U32, 1)
4603 .unwrap_err();
4604 assert!(matches!(error, NetworkError::RemoteAbort(_)));
4605 });
4606 let aborter = thread::spawn(move || {
4607 let session = TcpRankSession::connect_with_p2p_rails(
4608 address,
4609 unique_id,
4610 1,
4611 2,
4612 Duration::from_secs(2),
4613 2,
4614 )
4615 .unwrap();
4616 ready_receiver.recv().unwrap();
4617 session.abort("multi-rail failure").unwrap();
4618 });
4619 receiver.join().unwrap();
4620 aborter.join().unwrap();
4621 assert!(coordinator.join().unwrap().is_err());
4622 }
4623
4624 #[test]
4625 fn heartbeat_round_trip_is_independent_of_collective_sequence() {
4626 let unique_id = UniqueId::from_bytes([0x33; 16]);
4627 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 2).unwrap();
4628 let address = server.local_addr().unwrap();
4629 let coordinator = thread::spawn(move || server.run());
4630 let ranks = (0..2_u32)
4631 .map(|rank| {
4632 thread::spawn(move || {
4633 let session = TcpRankSession::connect(
4634 address,
4635 unique_id,
4636 rank,
4637 2,
4638 Duration::from_secs(2),
4639 )
4640 .unwrap();
4641 for _ in 0..3 {
4642 assert!(session.heartbeat(Duration::from_secs(1)).unwrap().as_secs() < 1);
4643 }
4644 session.barrier().unwrap();
4645 })
4646 })
4647 .collect::<Vec<_>>();
4648 for rank in ranks {
4649 rank.join().unwrap();
4650 }
4651 coordinator.join().unwrap().unwrap();
4652 }
4653
4654 #[test]
4655 fn heartbeat_enabled_session_closes_cleanly_with_leave() {
4656 let unique_id = UniqueId::from_bytes([0x35; 16]);
4657 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 2)
4658 .unwrap()
4659 .with_heartbeat_timeout(Duration::from_secs(1))
4660 .unwrap();
4661 let address = server.local_addr().unwrap();
4662 let coordinator = thread::spawn(move || server.run());
4663 let ranks = (0..2_u32)
4664 .map(|rank| {
4665 thread::spawn(move || {
4666 let session = TcpRankSession::connect(
4667 address,
4668 unique_id,
4669 rank,
4670 2,
4671 Duration::from_secs(2),
4672 )
4673 .unwrap();
4674 session.heartbeat(Duration::from_secs(1)).unwrap();
4675 session.barrier().unwrap();
4676 })
4677 })
4678 .collect::<Vec<_>>();
4679 for rank in ranks {
4680 rank.join().unwrap();
4681 }
4682 coordinator.join().unwrap().unwrap();
4683 }
4684
4685 #[test]
4686 fn collective_timeout_is_reconfigured_for_all_future_sequences() {
4687 let unique_id = UniqueId::from_bytes([0x36; 16]);
4688 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 2)
4689 .unwrap()
4690 .with_collective_timeout(Duration::from_secs(2))
4691 .unwrap();
4692 let address = server.local_addr().unwrap();
4693 let coordinator = thread::spawn(move || server.run());
4694 let waiting = thread::spawn(move || {
4695 let session =
4696 TcpRankSession::connect(address, unique_id, 0, 2, Duration::from_secs(2)).unwrap();
4697 session.set_timeout(Duration::from_millis(120)).unwrap();
4698 session.barrier()
4699 });
4700 let missing = thread::spawn(move || {
4701 let session =
4702 TcpRankSession::connect(address, unique_id, 1, 2, Duration::from_secs(2)).unwrap();
4703 session.set_timeout(Duration::from_millis(120)).unwrap();
4704 thread::sleep(Duration::from_millis(300));
4705 });
4706 let error = waiting.join().unwrap().unwrap_err();
4707 assert!(matches!(error, NetworkError::RemoteAbort(_)));
4708 assert!(error.to_string().contains("missing ranks [1]"));
4709 missing.join().unwrap();
4710 assert!(matches!(
4711 coordinator.join().unwrap(),
4712 Err(NetworkError::Timeout(_))
4713 ));
4714 }
4715
4716 #[test]
4717 fn heartbeat_lease_aborts_a_stale_idle_rank() {
4718 let unique_id = UniqueId::from_bytes([0x34; 16]);
4719 let server = TcpRendezvousServer::bind("127.0.0.1:0", unique_id, 2)
4720 .unwrap()
4721 .with_heartbeat_timeout(Duration::from_millis(150))
4722 .unwrap();
4723 let address = server.local_addr().unwrap();
4724 let coordinator = thread::spawn(move || server.run());
4725 let active = thread::spawn(move || {
4726 let session =
4727 TcpRankSession::connect(address, unique_id, 0, 2, Duration::from_secs(2)).unwrap();
4728 for _ in 0..20 {
4729 match session.heartbeat(Duration::from_secs(1)) {
4730 Ok(_) => thread::sleep(Duration::from_millis(30)),
4731 Err(error) => return error,
4732 }
4733 }
4734 panic!("stale rank did not expire")
4735 });
4736 let stale = thread::spawn(move || {
4737 let _session =
4738 TcpRankSession::connect(address, unique_id, 1, 2, Duration::from_secs(2)).unwrap();
4739 thread::sleep(Duration::from_millis(300));
4740 });
4741 let error = active.join().unwrap();
4742 assert!(
4743 matches!(error, NetworkError::RemoteAbort(_)),
4744 "unexpected heartbeat failure: {error:?}"
4745 );
4746 assert!(error.to_string().contains("heartbeat lease for rank 1"));
4747 stale.join().unwrap();
4748 assert!(coordinator.join().unwrap().is_err());
4749 }
4750}