spark_storage/weight_peer/
serve.rs1use anyhow::{Context, Result, bail};
10use std::collections::HashMap;
11use std::io::Read;
12use std::net::{TcpListener, TcpStream, ToSocketAddrs};
13use std::path::{Path, PathBuf};
14use std::sync::{Arc, Mutex};
15
16use super::manifest::WeightManifest;
17use super::shard::{Mmap, build_manifest};
18use super::wire::{read_model_request, write_weight_manifest};
19use crate::expert_peer::MODE_VERBS;
20
21#[derive(Clone, Debug)]
24pub struct WeightPeerConfig {
25 pub rails: Vec<(String, u32)>,
28 pub max_blade_bytes: u64,
32 pub staged_dirs: Vec<PathBuf>,
36 pub allow_any_path: bool,
39}
40
41impl Default for WeightPeerConfig {
42 fn default() -> Self {
43 Self {
44 rails: vec![("roceP2p1s0f1".into(), 3)],
45 max_blade_bytes: 0,
46 staged_dirs: Vec::new(),
47 allow_any_path: false,
48 }
49 }
50}
51
52struct StagedModel {
57 #[cfg_attr(not(atlas_rdma_verbs), allow(dead_code))]
60 shard_mmaps: Vec<Mmap>,
61 manifest: WeightManifest,
62 _reservation: crate::blade_cap::Reservation,
63}
64
65type StagedMap = Arc<Mutex<HashMap<String, Arc<StagedModel>>>>;
66
67pub fn serve<A: ToSocketAddrs>(addr: A, cfg: WeightPeerConfig) -> Result<()> {
71 let cfg = Arc::new(cfg);
72 let ledger = Arc::new(crate::blade_cap::CommitLedger::new(cfg.max_blade_bytes));
73 let staged: StagedMap = Arc::new(Mutex::new(HashMap::new()));
74
75 for dir in &cfg.staged_dirs {
78 match stage_model(&staged, &ledger, dir) {
79 Ok(m) => tracing::info!(
80 "weight-peer pre-staged {} ({} shards, {} tensors, {:.1} GiB)",
81 m.manifest.model_id,
82 m.manifest.num_shards(),
83 m.manifest.tensors.len(),
84 m.manifest.total_shard_bytes() as f64 / (1024.0 * 1024.0 * 1024.0),
85 ),
86 Err(e) => tracing::warn!("weight-peer pre-stage {} failed: {e}", dir.display()),
87 }
88 }
89
90 let listener = TcpListener::bind(addr).context("bind weight-peer listener")?;
91 let local = listener.local_addr().ok();
92 tracing::info!(
93 "weight-peer serving on {:?} (verbs rails {:?}, cap {}, allow_any_path {})",
94 local,
95 cfg.rails,
96 if cfg.max_blade_bytes == 0 {
97 "unlimited".to_string()
98 } else {
99 format!(
100 "{:.1} GiB",
101 cfg.max_blade_bytes as f64 / (1024.0 * 1024.0 * 1024.0)
102 )
103 },
104 cfg.allow_any_path,
105 );
106
107 for conn in listener.incoming() {
108 let stream = match conn {
109 Ok(s) => s,
110 Err(e) => {
111 tracing::warn!("weight-peer accept error: {e}");
112 continue;
113 }
114 };
115 let cfg = cfg.clone();
116 let ledger = ledger.clone();
117 let staged = staged.clone();
118 std::thread::spawn(move || {
119 if let Err(e) = handle_conn(stream, &cfg, &ledger, &staged) {
120 tracing::warn!("weight-peer connection ended: {e}");
121 }
122 });
123 }
124 Ok(())
125}
126
127fn handle_conn(
128 mut stream: TcpStream,
129 cfg: &WeightPeerConfig,
130 ledger: &Arc<crate::blade_cap::CommitLedger>,
131 staged: &StagedMap,
132) -> Result<()> {
133 stream.set_nodelay(true).ok();
134
135 let request = read_model_request(&mut stream)?;
137 let dir = resolve_request(cfg, &request)?;
138
139 let model = stage_model(staged, ledger, &dir)?;
141 write_weight_manifest(&mut stream, &model.manifest).context("send manifest")?;
142
143 let mut mode = [0u8; 1];
145 stream
146 .read_exact(&mut mode)
147 .context("read transport mode")?;
148 match mode[0] {
149 MODE_VERBS => serve_verbs(stream, &model, cfg),
150 other => bail!("weight-peer only serves verbs; client asked for mode {other}"),
151 }
152}
153
154fn resolve_request(cfg: &WeightPeerConfig, request: &str) -> Result<PathBuf> {
157 let req = Path::new(request);
158 for d in &cfg.staged_dirs {
160 if d == req
161 || d.file_name().and_then(|n| n.to_str()) == Some(request)
162 || d.to_string_lossy() == request
163 {
164 return Ok(d.clone());
165 }
166 }
167 if cfg.allow_any_path && req.is_dir() {
168 return Ok(req.to_path_buf());
169 }
170 bail!(
171 "model '{request}' is not staged (and allow_any_path is off); \
172 pass it to the peer with --stage <dir>"
173 );
174}
175
176fn stage_model(
179 staged: &StagedMap,
180 ledger: &Arc<crate::blade_cap::CommitLedger>,
181 dir: &Path,
182) -> Result<Arc<StagedModel>> {
183 let key = dir.to_string_lossy().into_owned();
184 {
185 let map = staged.lock().unwrap();
186 if let Some(m) = map.get(&key) {
187 return Ok(m.clone());
188 }
189 }
190
191 let (shard_paths, manifest) = build_manifest(dir, &key)?;
194 let reservation = ledger
197 .try_reserve(manifest.total_shard_bytes())
198 .context("weight blade cap")?;
199
200 let mut shard_mmaps = Vec::with_capacity(shard_paths.len());
201 for p in &shard_paths {
202 shard_mmaps.push(Mmap::open_ro(p).with_context(|| format!("mmap {}", p.display()))?);
203 }
204
205 let model = Arc::new(StagedModel {
206 shard_mmaps,
207 manifest,
208 _reservation: reservation,
209 });
210 let mut map = staged.lock().unwrap();
211 Ok(map.entry(key).or_insert(model).clone())
214}
215
216#[cfg(not(atlas_rdma_verbs))]
220fn serve_verbs(
221 _stream: TcpStream,
222 _model: &Arc<StagedModel>,
223 _cfg: &WeightPeerConfig,
224) -> Result<()> {
225 bail!("client requested verbs transport but this peer was built without rdma-core");
226}
227
228#[cfg(atlas_rdma_verbs)]
229fn serve_verbs(
230 mut stream: TcpStream,
231 model: &Arc<StagedModel>,
232 cfg: &WeightPeerConfig,
233) -> Result<()> {
234 use crate::expert_peer::{STATUS_OK, VerbsClientParams, VerbsServerParams, write_server_rails};
235 use atlas_rdma::verbs::Verbs;
236 use std::io::Write;
237
238 let num_shards = model.shard_mmaps.len();
239
240 let mut b1 = [0u8; 1];
242 stream.read_exact(&mut b1).context("read n_rails")?;
243 let n_rails = b1[0] as usize;
244 if n_rails == 0 || n_rails > cfg.rails.len() {
245 bail!(
246 "client asked for {n_rails} rails; peer has {}",
247 cfg.rails.len()
248 );
249 }
250
251 let pid = std::process::id();
255 let mut rails: Vec<Verbs> = Vec::with_capacity(n_rails);
256 for (i, (dev, gid)) in cfg.rails.iter().take(n_rails).enumerate() {
257 let psn = (0x77_7777 ^ pid ^ ((i as u32) << 20)) & 0xff_ffff;
258 rails.push(Verbs::create(dev, *gid, psn)?);
259 }
260
261 let mut per_rail_shards: Vec<Vec<(u64, u32)>> = (0..n_rails)
265 .map(|_| Vec::with_capacity(num_shards))
266 .collect();
267 for m in &model.shard_mmaps {
268 for (ri, v) in rails.iter_mut().enumerate() {
269 let keys = unsafe { v.reg_mr(m.addr as *mut _, m.len, true)? };
272 per_rail_shards[ri].push((m.addr as u64, keys.rkey));
273 }
274 }
275
276 let sp: Vec<VerbsServerParams> = rails
279 .iter()
280 .enumerate()
281 .map(|(ri, v)| VerbsServerParams {
282 qpn: v.qpn(),
283 psn: v.psn(),
284 gid: v.gid(),
285 layers: std::mem::take(&mut per_rail_shards[ri]),
286 })
287 .collect();
288 write_server_rails(&mut stream, &sp).context("send verbs server params")?;
289
290 stream.read_exact(&mut b1).context("read client n_rails")?;
292 if b1[0] as usize != n_rails {
293 bail!("client rail count mismatch");
294 }
295 for v in rails.iter_mut() {
296 let cp = VerbsClientParams::read_from(&mut stream).context("read verbs client params")?;
297 v.connect(cp.qpn, cp.psn, &cp.gid)?;
298 }
299 stream
300 .write_all(&[STATUS_OK])
301 .context("send verbs ready ack")?;
302 tracing::info!(
303 "weight-peer verbs client connected to {} ({n_rails} rail(s), {num_shards} shard MRs/rail)",
304 model.manifest.model_id,
305 );
306
307 let mut sink = [0u8; 8];
309 loop {
310 match stream.read(&mut sink) {
311 Ok(0) => break,
312 Ok(_) => {}
313 Err(_) => break,
314 }
315 }
316 drop(rails);
320 Ok(())
321}