1#[cfg(unix)]
25use anyhow::{Context, Result, bail};
26
27pub use atlas_rdma::wire::{
33 MODE_TCP, MODE_VERBS, STATUS_ERR, STATUS_OK, VerbsClientParams, VerbsServerParams,
34 read_server_rails, write_server_rails,
35};
36
37pub const SHUTDOWN_MARKER: u32 = u32::MAX;
39
40pub 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
48pub 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 #[derive(Clone, Debug)]
76 pub struct RdmaConfig {
77 pub rails: Vec<(String, u32)>,
79 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 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 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 stream.write_all(&(manifest.len() as u32).to_le_bytes())?;
159 stream.write_all(manifest)?;
160
161 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 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; }
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 #[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 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 let _reservation = ledger
248 .try_reserve(index.total_bytes())
249 .context("expert blade cap")?;
250
251 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 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 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 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 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 let mut sink = [0u8; 8];
314 loop {
315 match stream.read(&mut sink) {
316 Ok(0) => break, Ok(_) => {} Err(_) => break,
319 }
320 }
321 drop(rails);
324 drop(mmaps);
325 Ok(())
326 }
327
328 #[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 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 unsafe { libc::munmap(self.addr, self.len) };
372 }
373 }
374}
375
376#[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 }