spark_storage/expert_tier_rdma.rs
1// SPDX-License-Identifier: AGPL-3.0-only
2//
3// RdmaTier — the peer weight-fetch tier (Stage 4).
4//
5// Fetches expert records from an `atlas-expert-peer` straight into the pinned
6// arena, then returns residency addresses pointing INTO that arena — exactly
7// like `UmaArenaTier`, only the source is a peer instead of local NVMe. Two
8// transports share the tier and the arena machinery:
9//
10// * `Transport::Tcp` — Phase A: two-sided record streaming. The peer `pread`s
11// each record and writes it back over the socket; simple, bit-identical, but
12// the peer CPU is busy and single-stream bandwidth is ~5 GB/s.
13// * `Transport::Verbs` — Phase B: one-sided `IBV_WR_RDMA_READ`. The client
14// pulls each record directly out of the peer's registered store MR into the
15// arena slot with zero peer-CPU involvement (~14 GB/s, measured). This is a
16// pure transport swap — the bytes still land in the same pinned LPDDR that
17// the GPU reads at the same VA, and `residency_from` + the record header's
18// identity check catch any misplacement, so it cannot change a GEMM byte.
19//
20// The transport is chosen by the `--expert-backend` value: `rdma` = TCP,
21// `rdma-verbs` = one-sided verbs. Device/GID for verbs come from
22// `$ATLAS_EXPERT_RDMA_DEV` (default `roceP2p1s0f1`) / `$ATLAS_EXPERT_RDMA_GID`
23// (default 3, the RoCEv2 IPv4 GID on GB10/CX7).
24
25use std::io::{Read, Write};
26use std::net::TcpStream;
27
28use anyhow::{Context, Result, bail};
29
30use crate::expert::{ExpertKey, ExpertLayout, ExpertRecordSpec};
31use crate::expert_arena::ExpertArena;
32use crate::expert_peer::{MODE_TCP, STATUS_OK, encode_request, read_manifest};
33use crate::expert_tier::{ArenaSlot, ExpertResidency, ExpertTier, TierKind, residency_from};
34
35/// The active peer transport. `Verbs` only exists where the shim is compiled.
36/// The verbs transport is N-rail: one `Rail` per CX7 adapter, and a fetch is
37/// striped to `rail = expert % n_rails` (single-rail = the unchanged path).
38enum Transport {
39 Tcp,
40 #[cfg(atlas_rdma_verbs)]
41 Verbs(Vec<Rail>),
42}
43
44/// One-sided verbs state for a single rail: the QP, this rail's arena MR lkey,
45/// and the per-layer remote MR `{base, rkey}` table the peer published for it.
46/// The base VA is shared across rails (the peer mmaps each layer once); only the
47/// rkey (and QP/NIC) differ per rail.
48#[cfg(atlas_rdma_verbs)]
49struct Rail {
50 verbs: atlas_rdma::verbs::Verbs,
51 arena_lkey: u32,
52 /// `(remote_base_addr, rkey)` per MoE layer, layer-indexed.
53 layers: Vec<(u64, u32)>,
54}
55
56pub struct RdmaTier {
57 stream: TcpStream,
58 // `transport` is declared BEFORE `arena` so it drops first: the verbs rails
59 // hold MRs registered over the arena's pinned pages, so their `ibv_dereg_mr`
60 // must run before the arena frees those pages. (Struct fields drop in
61 // declaration order.) With N rails this is load-bearing — reverting the order
62 // would dereg N MRs over freed memory.
63 transport: Transport,
64 arena: ExpertArena,
65 spec: ExpertRecordSpec,
66 layout: ExpertLayout,
67 healthy: bool,
68}
69
70impl RdmaTier {
71 /// Connect to a peer at `addr`, receive its manifest, allocate the arena, and
72 /// bring up the chosen transport. The peer's `ExpertIndex` geometry defines
73 /// the record stride, so the arena matches the remote store exactly.
74 pub fn connect(
75 addr: &str,
76 num_slabs: u32,
77 slots_per_slab: u32,
78 use_verbs: bool,
79 ) -> Result<Self> {
80 let mut stream =
81 TcpStream::connect(addr).with_context(|| format!("connect expert peer {addr}"))?;
82 stream.set_nodelay(true).ok();
83 let index = read_manifest(&mut stream)?;
84 let spec = index.spec();
85 let layout = index.layout();
86 let arena = ExpertArena::new(num_slabs, slots_per_slab, layout.record_stride as usize)?;
87
88 let transport = if use_verbs {
89 #[cfg(atlas_rdma_verbs)]
90 {
91 connect_verbs(&mut stream, &arena, index.num_moe_layers)?
92 }
93 // Built without rdma-core (no C shim) — verbs is unavailable; the
94 // TCP `rdma` backend still works. Keeps the crate compiling under
95 // ATLAS_SKIP_BUILD / hosts without libibverbs.
96 #[cfg(not(atlas_rdma_verbs))]
97 {
98 let _ = &arena;
99 bail!(
100 "--expert-backend rdma-verbs needs a build with rdma-core \
101 (atlas_rdma_verbs cfg); use --expert-backend rdma (TCP) instead"
102 );
103 }
104 } else {
105 stream
106 .write_all(&[MODE_TCP])
107 .context("send TCP transport mode")?;
108 Transport::Tcp
109 };
110
111 let label = match &transport {
112 Transport::Tcp => "TCP".to_string(),
113 #[cfg(atlas_rdma_verbs)]
114 Transport::Verbs(rails) => {
115 format!("verbs (one-sided RDMA READ, {} rail(s))", rails.len())
116 }
117 };
118 tracing::info!(
119 "RdmaTier[{label}] connected to {addr}: {} layers, {} experts, stride {}",
120 index.num_moe_layers,
121 index.num_experts,
122 layout.record_stride
123 );
124 Ok(Self {
125 stream,
126 arena,
127 spec,
128 layout,
129 transport,
130 healthy: true,
131 })
132 }
133
134 pub fn arena(&self) -> &ExpertArena {
135 &self.arena
136 }
137
138 /// Two-sided TCP fetch: request the record, read `[status][stride bytes]`
139 /// straight into the pinned slot.
140 fn fetch_tcp(&mut self, key: ExpertKey, host: *mut u8, stride: usize) -> Result<()> {
141 if let Err(e) = self
142 .stream
143 .write_all(&encode_request(key.layer, key.expert))
144 {
145 self.healthy = false;
146 return Err(e).with_context(|| format!("peer request {:?}", key));
147 }
148 let mut status = [0u8; 1];
149 if let Err(e) = self.stream.read_exact(&mut status) {
150 self.healthy = false;
151 return Err(e).with_context(|| format!("peer status {:?}", key));
152 }
153 if status[0] != STATUS_OK {
154 bail!("peer returned error status {} for {:?}", status[0], key);
155 }
156 // Land the record bytes DIRECTLY into the pinned, GPU-addressable slot.
157 // SAFETY: `host` points at a `stride`-byte slot inside the pinned arena.
158 let dst = unsafe { std::slice::from_raw_parts_mut(host, stride) };
159 if let Err(e) = self.stream.read_exact(dst) {
160 self.healthy = false;
161 return Err(e).with_context(|| format!("peer payload {:?}", key));
162 }
163 Ok(())
164 }
165}
166
167/// Bring up the one-sided verbs transport via [`atlas_rdma::railset::RailSet`]:
168/// create N rails, register the arena on each, exchange per-rail QP params over
169/// the TCP control channel, connect INIT->RTR->RTS, await the ack. Dual-rail is
170/// env-driven (ATLAS_EXPERT_DUAL_RAIL=1): rail 0 = ATLAS_EXPERT_RDMA_DEV/GID
171/// (the existing single-rail defaults), rail 1 = ATLAS_EXPERT_RAIL2_DEV/GID
172/// (default rocep1s0f1 / 3). Single-rail is the default and is byte-for-byte
173/// the previous path.
174#[cfg(atlas_rdma_verbs)]
175fn connect_verbs(
176 stream: &mut TcpStream,
177 arena: &ExpertArena,
178 num_layers: u32,
179) -> Result<Transport> {
180 use crate::expert_peer::MODE_VERBS;
181 use atlas_rdma::env::{first_set, first_set_u32};
182 use atlas_rdma::railset::{RailSet, RailSpec};
183
184 stream
185 .write_all(&[MODE_VERBS])
186 .context("send verbs transport mode")?;
187
188 // Rail 0 from the expert env (the cabled CX7 link); rail 1 from the expert
189 // rail-2 env. Dual-rail only when ATLAS_EXPERT_DUAL_RAIL=1. PSN = fresh
190 // random 24-bit per rail (caller-supplied by RailSet design).
191 let spec = |dev: String, gid: u32| RailSpec::new(dev, gid, rand::random::<u32>() & 0xff_ffff);
192 let rail0 = spec(
193 first_set(&["ATLAS_EXPERT_RDMA_DEV"], "roceP2p1s0f1"),
194 first_set_u32(&["ATLAS_EXPERT_RDMA_GID"], 3),
195 );
196 let dual = std::env::var("ATLAS_EXPERT_DUAL_RAIL").ok().as_deref() == Some("1");
197 let specs: Vec<RailSpec> = if dual {
198 let rail1 = spec(
199 first_set(&["ATLAS_EXPERT_RAIL2_DEV"], "rocep1s0f1"),
200 first_set_u32(&["ATLAS_EXPERT_RAIL2_GID"], 3),
201 );
202 vec![rail0, rail1]
203 } else {
204 vec![rail0]
205 };
206
207 // [u8 n_rails] + one QP per rail, then register the WHOLE arena as each
208 // rail's READ landing MR (LOCAL_WRITE only — `remote_read == false` is the
209 // access-flag invariant for every client landing buffer). The N MRs pin the
210 // SAME arena pages (one lkey per rail).
211 let mut rs = RailSet::begin(stream, &specs)?;
212 let mut arena_lkeys: Vec<u32> = Vec::with_capacity(rs.n_rails());
213 for rail in &mut rs.rails {
214 // SAFETY: the arena's pinned buffer lives as long as the tier (and thus
215 // every MR); base_ptr()/total_bytes() describe exactly that allocation.
216 let keys = unsafe {
217 rail.verbs
218 .reg_mr(arena.base_ptr(), arena.total_bytes(), false)?
219 };
220 arena_lkeys.push(keys.lkey);
221 }
222
223 // Peer publishes N per-rail QP + per-layer MR tables; validate each rail's
224 // layer count against the manifest BEFORE replying (a mismatch must bail
225 // with no client params written — the pre-RailSet behavior).
226 let server = rs
227 .read_server_ro(stream)
228 .context("read verbs server params")?;
229 for sp in &server {
230 if sp.layers.len() != num_layers as usize {
231 bail!(
232 "verbs peer published {} layer MRs but manifest has {num_layers} MoE layers",
233 sp.layers.len()
234 );
235 }
236 }
237
238 // Reply with each rail's client QP, connect each rail, await the ack.
239 rs.complete(stream, &server, "verbs peer")?;
240 let rails: Vec<Rail> = rs
241 .into_verbs()
242 .into_iter()
243 .zip(arena_lkeys)
244 .zip(server)
245 .map(|((verbs, arena_lkey), sp)| Rail {
246 verbs,
247 arena_lkey,
248 layers: sp.layers,
249 })
250 .collect();
251 Ok(Transport::Verbs(rails))
252}
253
254impl ExpertTier for RdmaTier {
255 fn fetch(&mut self, key: ExpertKey, slot: ArenaSlot, _stream: u64) -> Result<ExpertResidency> {
256 let stride = self.layout.record_stride as usize;
257 let host = self.arena.slot_host_ptr(slot.slab, slot.slot)?;
258 let dev_va = self.arena.slot_dev_va(slot.slab, slot.slot)?;
259 let spec = self.spec; // Copy — release the field borrow before matching.
260
261 match &mut self.transport {
262 Transport::Tcp => {
263 self.fetch_tcp(key, host, stride)?;
264 }
265 #[cfg(atlas_rdma_verbs)]
266 Transport::Verbs(rails) => {
267 // Stripe the fetch onto rail = expert % n_rails. Single-rail
268 // (n == 1) => always rail 0, the unchanged path.
269 let ri = (key.expert as usize) % rails.len();
270 let rail = &mut rails[ri];
271 let (base, rkey) = *rail.layers.get(key.layer as usize).with_context(|| {
272 format!("verbs: no layer MR for layer {} ({:?})", key.layer, key)
273 })?;
274 let remote_addr = base + (key.expert as u64) * (stride as u64);
275 let wr_id = ((key.layer as u64) << 32) | (key.expert as u64);
276 // SAFETY: `host` is a `stride`-byte slot inside this rail's arena
277 // MR (arena_lkey); remote_addr/rkey address the peer's layer MR on
278 // the SAME rail.
279 let post = unsafe {
280 rail.verbs.post_read(
281 host as *mut std::ffi::c_void,
282 rail.arena_lkey,
283 remote_addr,
284 rkey,
285 stride as u32,
286 wr_id,
287 )
288 };
289 if let Err(e) = post {
290 self.healthy = false;
291 return Err(e).with_context(|| format!("verbs post_read {:?}", key));
292 }
293 match rail.verbs.poll() {
294 Ok(got) if got == wr_id => {}
295 Ok(got) => {
296 self.healthy = false;
297 bail!("verbs completion wr_id {got:#x} != expected {wr_id:#x} ({key:?})");
298 }
299 Err(e) => {
300 self.healthy = false;
301 return Err(e).with_context(|| format!("verbs poll {:?}", key));
302 }
303 }
304 }
305 }
306
307 // SAFETY: the slot now holds `stride` valid bytes (landed by TCP or RDMA).
308 let record = unsafe { std::slice::from_raw_parts(host, stride) };
309 residency_from(&spec, record, dev_va, key)
310 }
311
312 fn kind(&self) -> TierKind {
313 TierKind::Rdma
314 }
315
316 /// Link health: false after any transport error so the streamer can fall
317 /// back to the local NVMe UMA tier (graceful degradation on CX7 flap).
318 fn healthy(&self) -> bool {
319 self.healthy
320 }
321}