ruccl/rank/communicator/
connect.rs1use 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}