Skip to main content

ruccl/rank/
network.rs

1//! TCP rendezvous and host-staged collective exchange using the versioned GX
2//! collective wire protocol. Payload reduction remains a rank-side GX kernel
3//! responsibility; the coordinator only validates ordering and routes bytes.
4
5use 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    /// Accept one ordered collective channel and one or more independently
274    /// multiplexed point-to-point rails per rank, then serve until ranks disconnect.
275    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    // Keep the original one-line-per-rank table when every rank uses one
494    // shared listener. New clients accept both layouts, so existing scalar
495    // GX1_P2P_LISTEN_ADDR deployments retain their wire representation.
496    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    /// Notify every rank that this process is abandoning the current session.
2133    /// The abort uses the independent P2P channel so it can interrupt a rank
2134    /// waiting in the ordered collective channel.
2135    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}