spark_storage/
expert_peer.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2//
3// Expert weight-serving peer + wire protocol (Stage 4, Phase A: TCP).
4//
5// The RDMA weight-tier's first incarnation: a peer that holds the resident
6// expert store and serves records to a streaming client over a socket, into the
7// client's pinned arena. This proves the residency-tier abstraction (a peer as
8// a fetch tier, distinct from EP sharding) with zero verbs risk, over the RoCE
9// Ethernet netdev. Phase B swaps the transport for one-sided RDMA_READ into the
10// SAME arena (see RESEARCH-RDMA-TIER.md) — the protocol geometry is identical.
11//
12// Wire protocol (little-endian), connection-oriented:
13//   1. On accept the server sends the manifest: [u32 len][len bytes of JSON].
14//      The client parses it to size its arena (same ExpertIndex geometry).
15//   2. Request loop: client sends [u32 layer][u32 expert]; server replies
16//      [u8 status][record_stride bytes] (status 0 = OK, nonzero = error, no
17//      payload). A layer/expert of u32::MAX/u32::MAX is a graceful shutdown.
18//
19// The peer is pure I/O (no CUDA); the client half lives in the cuda-gated
20// `expert_tier_rdma` module because it lands bytes in the pinned arena.
21
22// Consumed only by the `#[cfg(unix)]` peer server below; same gate, so a
23// Windows build does not trip `unused_imports` under `warnings = "deny"`.
24#[cfg(unix)]
25use anyhow::{Context, Result, bail};
26
27// The handshake wire codecs moved verbatim to the CUDA-free `atlas-rdma`
28// crate (extracted to atlas-rdma); re-exported here at their old paths so
29// the server below and every external user are zero-diff. The byte layouts
30// are golden-pinned in `tests/rdma_wire_golden.rs` and frozen vs the live
31// gx10 peer.
32pub use atlas_rdma::wire::{
33    MODE_TCP, MODE_VERBS, STATUS_ERR, STATUS_OK, VerbsClientParams, VerbsServerParams,
34    read_server_rails, write_server_rails,
35};
36
37/// Sentinel request that asks the server to close the connection.
38pub const SHUTDOWN_MARKER: u32 = u32::MAX;
39
40/// Serialize a request: `(layer, expert)`.
41pub fn encode_request(layer: u32, expert: u32) -> [u8; 8] {
42    let mut b = [0u8; 8];
43    b[0..4].copy_from_slice(&layer.to_le_bytes());
44    b[4..8].copy_from_slice(&expert.to_le_bytes());
45    b
46}
47
48/// Parse a request buffer.
49pub fn decode_request(b: &[u8; 8]) -> (u32, u32) {
50    let layer = u32::from_le_bytes([b[0], b[1], b[2], b[3]]);
51    let expert = u32::from_le_bytes([b[4], b[5], b[6], b[7]]);
52    (layer, expert)
53}
54
55#[cfg(unix)]
56pub use server_impl::{RdmaConfig, serve};
57
58#[cfg(unix)]
59mod server_impl {
60    use super::*;
61    use crate::expert::ExpertKey;
62    use crate::expert_pack::ExpertFileReader;
63    use std::io::{Read, Write};
64    use std::net::{TcpListener, TcpStream, ToSocketAddrs};
65    use std::path::{Path, PathBuf};
66    use std::sync::Arc;
67
68    /// RDMA device selection for the verbs (`MODE_VERBS`) transport. Ignored for
69    /// TCP clients. One `(device, gid_idx)` per CX7 adapter: a dual-rail client
70    /// requests N rails and the peer registers each layer mmap on every rail so
71    /// the client can stripe expert fetches across both adapters. The default is
72    /// a SINGLE rail (`roceP2p1s0f1`, RoCEv2 GID index 3) — the pre-dual-rail
73    /// behavior, unchanged for `verify_verbs.sh`; add a second `--rail` to serve
74    /// dual-rail clients.
75    #[derive(Clone, Debug)]
76    pub struct RdmaConfig {
77        /// `(device, gid_idx)` per rail, in link order (rail 0 = the cabled link).
78        pub rails: Vec<(String, u32)>,
79        /// Ceiling on total registered store RAM across concurrent verbs
80        /// connections, in bytes. `0` = unlimited (the default). Each verbs
81        /// connection registers the whole store (`index.total_bytes()`) once (the
82        /// N per-rail MRs share the same mmap pages), so this bounds the number of
83        /// concurrent store registrations.
84        pub max_blade_bytes: u64,
85    }
86
87    impl Default for RdmaConfig {
88        fn default() -> Self {
89            Self {
90                rails: vec![("roceP2p1s0f1".into(), 3)],
91                max_blade_bytes: 0,
92            }
93        }
94    }
95
96    /// Serve records from `dir` on `addr` until interrupted. One thread per
97    /// connection. Blocking; intended to run as its own process (`atlas-expert-peer`).
98    pub fn serve<A: ToSocketAddrs>(dir: &Path, addr: A, rdma: RdmaConfig) -> Result<()> {
99        let reader = Arc::new(ExpertFileReader::open(dir)?);
100        let manifest = serde_json::to_vec(reader.index())?;
101        let dir: Arc<PathBuf> = Arc::new(dir.to_path_buf());
102        let rdma = Arc::new(rdma);
103        // One process-global ledger; each verbs connection reserves the store
104        // size before it mmaps/registers the layer files.
105        let ledger = Arc::new(crate::blade_cap::CommitLedger::new(rdma.max_blade_bytes));
106        let listener = TcpListener::bind(addr).context("bind expert-peer listener")?;
107        let local = listener.local_addr().ok();
108        tracing::info!(
109            "expert-peer serving {} ({} layers, {} experts, stride {}) on {:?} \
110             (verbs rails {:?}, cap {})",
111            dir.display(),
112            reader.index().num_moe_layers,
113            reader.index().num_experts,
114            reader.index().record_stride,
115            local,
116            rdma.rails,
117            if rdma.max_blade_bytes == 0 {
118                "unlimited".to_string()
119            } else {
120                format!(
121                    "{:.1} GiB",
122                    rdma.max_blade_bytes as f64 / (1024.0 * 1024.0 * 1024.0)
123                )
124            },
125        );
126        for conn in listener.incoming() {
127            let stream = match conn {
128                Ok(s) => s,
129                Err(e) => {
130                    tracing::warn!("expert-peer accept error: {e}");
131                    continue;
132                }
133            };
134            let reader = reader.clone();
135            let manifest = manifest.clone();
136            let dir = dir.clone();
137            let rdma = rdma.clone();
138            let ledger = ledger.clone();
139            std::thread::spawn(move || {
140                if let Err(e) = handle_conn(stream, &reader, &manifest, &dir, &rdma, &ledger) {
141                    tracing::warn!("expert-peer connection ended: {e}");
142                }
143            });
144        }
145        Ok(())
146    }
147
148    fn handle_conn(
149        mut stream: TcpStream,
150        reader: &ExpertFileReader,
151        manifest: &[u8],
152        dir: &Path,
153        rdma: &RdmaConfig,
154        ledger: &Arc<crate::blade_cap::CommitLedger>,
155    ) -> Result<()> {
156        stream.set_nodelay(true).ok();
157        // 1. Send the manifest.
158        stream.write_all(&(manifest.len() as u32).to_le_bytes())?;
159        stream.write_all(manifest)?;
160
161        // 2. The client picks a transport. TCP `pread`s on demand and pins no
162        // RAM, so only the verbs path (which mmaps + registers the store) is
163        // charged against the ledger.
164        let mut mode = [0u8; 1];
165        stream
166            .read_exact(&mut mode)
167            .context("read transport mode")?;
168        match mode[0] {
169            MODE_TCP => serve_tcp(stream, reader),
170            MODE_VERBS => serve_verbs(stream, reader, dir, rdma, ledger),
171            other => bail!("client requested unknown transport mode {other}"),
172        }
173    }
174
175    /// Two-sided record streaming (Phase A). The request loop is unchanged.
176    fn serve_tcp(mut stream: TcpStream, reader: &ExpertFileReader) -> Result<()> {
177        let stride = reader.index().record_stride as usize;
178        let mut req = [0u8; 8];
179        loop {
180            if stream.read_exact(&mut req).is_err() {
181                break; // client hung up
182            }
183            let (layer, expert) = decode_request(&req);
184            if layer == SHUTDOWN_MARKER && expert == SHUTDOWN_MARKER {
185                break;
186            }
187            match reader.read_record_raw(ExpertKey::new(layer, expert)) {
188                Ok(rec) => {
189                    debug_assert_eq!(rec.len(), stride);
190                    stream.write_all(&[STATUS_OK])?;
191                    stream.write_all(&rec)?;
192                }
193                Err(e) => {
194                    tracing::warn!("expert-peer read {layer}/{expert}: {e}");
195                    stream.write_all(&[STATUS_ERR])?;
196                }
197            }
198        }
199        Ok(())
200    }
201
202    #[cfg(not(atlas_rdma_verbs))]
203    fn serve_verbs(
204        _stream: TcpStream,
205        _reader: &ExpertFileReader,
206        _dir: &Path,
207        _rdma: &RdmaConfig,
208        _ledger: &Arc<crate::blade_cap::CommitLedger>,
209    ) -> Result<()> {
210        bail!("client requested verbs transport but this peer was built without rdma-core");
211    }
212
213    /// One-sided RDMA READ (Phase B). The server `mmap`s + registers each layer
214    /// file with REMOTE_READ, publishes the MRs' `{base, rkey}` + its QP params,
215    /// connects to the client's QP, then goes idle: the client pulls records
216    /// directly out of these MRs. The server CPU never touches a record byte.
217    #[cfg(atlas_rdma_verbs)]
218    fn serve_verbs(
219        mut stream: TcpStream,
220        reader: &ExpertFileReader,
221        dir: &Path,
222        rdma: &RdmaConfig,
223        ledger: &Arc<crate::blade_cap::CommitLedger>,
224    ) -> Result<()> {
225        use atlas_rdma::verbs::Verbs;
226
227        let index = reader.index();
228        let num_layers = index.num_moe_layers;
229
230        // The client requests how many rails it wants to stripe across (default
231        // 1). Validate against what this peer is configured to serve.
232        let mut b1 = [0u8; 1];
233        stream.read_exact(&mut b1).context("read n_rails")?;
234        let n_rails = b1[0] as usize;
235        if n_rails == 0 || n_rails > rdma.rails.len() {
236            bail!(
237                "client asked for {n_rails} rails; peer has {}",
238                rdma.rails.len()
239            );
240        }
241
242        // Admission gate: charge the deterministic store size (identical for
243        // every client) ONCE — the N per-rail MRs pin the SAME refcounted mmap
244        // pages, so the committed footprint is `total_bytes`, not total*n_rails.
245        // Reserve BEFORE any mmap/reg_mr; the RAII guard releases on every exit
246        // below (early bail, reg_mr error, hangup).
247        let _reservation = ledger
248            .try_reserve(index.total_bytes())
249            .context("expert blade cap")?;
250
251        // One QP per rail (distinct per-rail PSN so successive clients/rails
252        // don't collide).
253        let pid = std::process::id();
254        let mut rails: Vec<Verbs> = Vec::with_capacity(n_rails);
255        for (i, (dev, gid)) in rdma.rails.iter().take(n_rails).enumerate() {
256            let psn = (0x424242 ^ pid ^ ((i as u32) << 20)) & 0xff_ffff;
257            rails.push(Verbs::create(dev, *gid, psn)?);
258        }
259
260        // mmap each layer file ONCE (REMOTE_READ) and register that SAME virtual
261        // range on EVERY rail's PD — one rkey per (rail, layer), identical base
262        // VA, shared physical pages (not N× RAM). Keep the mappings alive for the
263        // whole connection — the NIC DMAs out of them.
264        let mut mmaps: Vec<Mmap> = Vec::with_capacity(num_layers as usize);
265        let mut per_rail_layers: Vec<Vec<(u64, u32)>> = (0..n_rails)
266            .map(|_| Vec::with_capacity(num_layers as usize))
267            .collect();
268        for l in 0..num_layers {
269            let path = dir.join(index.file_name(l));
270            let m = Mmap::open_ro(&path).with_context(|| format!("mmap {}", path.display()))?;
271            for (ri, v) in rails.iter_mut().enumerate() {
272                // SAFETY: the mapping covers `m.len` bytes at `m.addr` and
273                // outlives every rail's Verbs (mmaps dropped after rails below).
274                let keys = unsafe { v.reg_mr(m.addr as *mut _, m.len, true)? };
275                per_rail_layers[ri].push((m.addr as u64, keys.rkey));
276            }
277            mmaps.push(m);
278        }
279
280        // Publish one VerbsServerParams per rail (shared base, per-rail rkey).
281        let sp: Vec<VerbsServerParams> = rails
282            .iter()
283            .enumerate()
284            .map(|(ri, v)| VerbsServerParams {
285                qpn: v.qpn(),
286                psn: v.psn(),
287                gid: v.gid(),
288                layers: std::mem::take(&mut per_rail_layers[ri]),
289            })
290            .collect();
291        write_server_rails(&mut stream, &sp).context("send verbs server params")?;
292
293        // Learn each client rail's QP, connect, ack.
294        stream.read_exact(&mut b1).context("read client n_rails")?;
295        if b1[0] as usize != n_rails {
296            bail!("client rail count mismatch");
297        }
298        for v in rails.iter_mut() {
299            let cp =
300                VerbsClientParams::read_from(&mut stream).context("read verbs client params")?;
301            v.connect(cp.qpn, cp.psn, &cp.gid)?;
302        }
303        stream
304            .write_all(&[STATUS_OK])
305            .context("send verbs ready ack")?;
306        tracing::info!(
307            "expert-peer verbs client connected ({n_rails} rail(s), {} layer MRs/rail)",
308            num_layers,
309        );
310
311        // Idle until the client hangs up. All record movement is one-sided RDMA
312        // READ initiated by the client; the server just holds the MRs open.
313        let mut sink = [0u8; 8];
314        loop {
315            match stream.read(&mut sink) {
316                Ok(0) => break, // client closed
317                Ok(_) => {}     // ignore (shutdown marker or stray bytes)
318                Err(_) => break,
319            }
320        }
321        // Drop order: `rails` (which dereg's every MR) must fall before `mmaps`
322        // are unmapped, so dereg happens over live mappings. Drop rails first.
323        drop(rails);
324        drop(mmaps);
325        Ok(())
326    }
327
328    /// A read-only `mmap` of a whole file, unmapped on drop.
329    #[cfg(atlas_rdma_verbs)]
330    struct Mmap {
331        addr: *mut libc::c_void,
332        len: usize,
333    }
334
335    #[cfg(atlas_rdma_verbs)]
336    impl Mmap {
337        fn open_ro(path: &Path) -> Result<Self> {
338            use std::os::fd::AsRawFd;
339            let f = std::fs::File::open(path)?;
340            let len = f.metadata()?.len() as usize;
341            if len == 0 {
342                bail!("empty layer file {}", path.display());
343            }
344            // SAFETY: fd is a valid open RO file; MAP_SHARED read mapping of `len`
345            // bytes. The kernel keeps the mapping valid after the fd closes.
346            let addr = unsafe {
347                libc::mmap(
348                    std::ptr::null_mut(),
349                    len,
350                    libc::PROT_READ,
351                    libc::MAP_SHARED,
352                    f.as_raw_fd(),
353                    0,
354                )
355            };
356            if addr == libc::MAP_FAILED {
357                bail!(
358                    "mmap {} failed: {}",
359                    path.display(),
360                    std::io::Error::last_os_error()
361                );
362            }
363            Ok(Self { addr, len })
364        }
365    }
366
367    #[cfg(atlas_rdma_verbs)]
368    impl Drop for Mmap {
369        fn drop(&mut self) {
370            // SAFETY: addr/len came from a successful mmap and are unmapped once.
371            unsafe { libc::munmap(self.addr, self.len) };
372        }
373    }
374}
375
376/// Read the length-prefixed manifest from a freshly-connected stream and parse
377/// it. Shared by the client (`expert_tier_rdma`).
378#[cfg(unix)]
379pub fn read_manifest<R: std::io::Read>(stream: &mut R) -> Result<crate::expert_pack::ExpertIndex> {
380    let mut lenb = [0u8; 4];
381    stream
382        .read_exact(&mut lenb)
383        .context("read manifest length")?;
384    let len = u32::from_le_bytes(lenb) as usize;
385    if len == 0 || len > 16 * 1024 * 1024 {
386        bail!("implausible peer manifest length: {len}");
387    }
388    let mut buf = vec![0u8; len];
389    stream.read_exact(&mut buf).context("read manifest json")?;
390    let index: crate::expert_pack::ExpertIndex =
391        serde_json::from_slice(&buf).context("parse peer manifest")?;
392    Ok(index)
393}
394
395#[cfg(test)]
396mod tests {
397    use super::*;
398
399    #[test]
400    fn request_round_trips() {
401        let b = encode_request(7, 42);
402        assert_eq!(decode_request(&b), (7, 42));
403        let s = encode_request(SHUTDOWN_MARKER, SHUTDOWN_MARKER);
404        assert_eq!(decode_request(&s), (SHUTDOWN_MARKER, SHUTDOWN_MARKER));
405    }
406
407    // The codec round-trip / validation tests moved WITH the codecs to
408    // `crates/atlas-rdma/tests/wire_roundtrip.rs` (extracted to atlas-rdma);
409    // the exact byte layouts stay pinned by `tests/rdma_wire_golden.rs` here.
410}