Skip to main content

ruccl/rank/communicator/
connect.rs

1use super::*;
2use crate::rank::TcpRankSession;
3use crate::rank::transport::rank_transport_from_environment;
4use std::net::ToSocketAddrs;
5
6impl<D> RankCommunicator<D> {
7    pub fn connect<E>(
8        initialize: impl FnOnce() -> Result<D, E>,
9        address: impl ToSocketAddrs,
10        unique_id: UniqueId,
11        rank: u32,
12        world_size: u32,
13        timeout: Duration,
14        queue_name: &'static str,
15    ) -> Result<Self, E>
16    where
17        E: From<TopologyError> + From<NetworkError> + From<WorkError>,
18    {
19        let session = TcpRankSession::connect(address, unique_id, rank, world_size, timeout)?;
20        Self::from_session(initialize, session, queue_name)
21    }
22
23    pub fn from_session<E>(
24        initialize: impl FnOnce() -> Result<D, E>,
25        session: TcpRankSession,
26        queue_name: &'static str,
27    ) -> Result<Self, E>
28    where
29        E: From<TopologyError> + From<NetworkError> + From<WorkError>,
30    {
31        let tuning = if CollectiveTuning::topology_probe_requested()? {
32            let topology = session.probe_topology_from_environment()?;
33            CollectiveTuning::from_environment_with_topology(topology)?
34        } else {
35            CollectiveTuning::from_environment(session.world_size())?
36        };
37        Self::from_session_with_tuning(initialize, session, tuning, queue_name)
38    }
39
40    pub fn from_session_with_tuning<E>(
41        initialize: impl FnOnce() -> Result<D, E>,
42        session: TcpRankSession,
43        tuning: CollectiveTuning,
44        queue_name: &'static str,
45    ) -> Result<Self, E>
46    where
47        E: From<TopologyError> + From<NetworkError> + From<WorkError>,
48    {
49        let transport = rank_transport_from_environment(session)?;
50        Self::from_transport_arc_with_tuning(initialize, transport, tuning, queue_name)
51    }
52
53    pub fn from_transport_with_tuning<T, E>(
54        initialize: impl FnOnce() -> Result<D, E>,
55        transport: T,
56        tuning: CollectiveTuning,
57        queue_name: &'static str,
58    ) -> Result<Self, E>
59    where
60        T: RankTransport + 'static,
61        E: From<TopologyError> + From<NetworkError> + From<WorkError>,
62    {
63        let transport = rank_transport_from_environment(transport)?;
64        Self::from_transport_arc_with_tuning(initialize, transport, tuning, queue_name)
65    }
66
67    fn from_transport_arc_with_tuning<E>(
68        initialize: impl FnOnce() -> Result<D, E>,
69        session: Arc<dyn RankTransport>,
70        tuning: CollectiveTuning,
71        queue_name: &'static str,
72    ) -> Result<Self, E>
73    where
74        E: From<TopologyError> + From<NetworkError> + From<WorkError>,
75    {
76        let transport = session.transport();
77        if tuning.transport() != transport {
78            return Err(TopologyError::TransportMismatch {
79                tuning: tuning.transport(),
80                communicator: transport,
81            }
82            .into());
83        }
84        if tuning.ring_order.len() != session.world_size() as usize {
85            return Err(TopologyError::InvalidRing(format!(
86                "contains {} ranks, expected {}",
87                tuning.ring_order.len(),
88                session.world_size()
89            ))
90            .into());
91        }
92        if tuning.p2p_rails() != session.p2p_rails() {
93            return Err(TopologyError::P2pRailMismatch {
94                tuning: tuning.p2p_rails(),
95                session: session.p2p_rails(),
96            }
97            .into());
98        }
99        let execution = initialize()?;
100        let async_queue = OrderedWorkQueue::new(queue_name)?;
101        let collective_rail_order = Arc::<[usize]>::from(tuning.rail_order());
102        Ok(Self {
103            execution,
104            session,
105            async_queue,
106            tuning: Arc::new(tuning),
107            collective_rail_order,
108            internal_sequence: Arc::new(AtomicU64::new(0)),
109        })
110    }
111}