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}