Skip to main content

kvbm_engine/worker/group/
spmd.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use super::*;
5
6use crate::object::ObjectBlockOps;
7use anyhow::Result;
8// velo event types used via fully-qualified paths (::velo::Event, ::velo::EventManager)
9use futures::future::BoxFuture;
10
11use std::collections::HashMap;
12use std::sync::{Arc, RwLock};
13
14/// SPMD (Single Program, Multiple Data) parallel worker group.
15///
16/// Wraps a set of rank-indexed [`Worker`]s and executes every operation on
17/// all of them in parallel. Each worker has its own rank, physical layout
18/// handles, and `TransferManager`, but they all receive the same logical
19/// commands (transfer, connect, import/export metadata).
20///
21/// Transfer completion notifications from individual workers are aggregated
22/// into a single notification via the event system, so callers see one
23/// completion event per logical operation regardless of worker count.
24///
25/// Remote handle mappings are stored per `(InstanceId, worker_idx,
26/// LogicalLayoutHandle)` so that each rank resolves to its own peer handle
27/// during RDMA transfers.
28pub struct SpmdParallelWorkers {
29    workers: Vec<Arc<dyn Worker>>,
30    events: Arc<::velo::EventManager>,
31    runtime: tokio::runtime::Handle,
32
33    /// Remote handle mappings: (InstanceId, worker_idx, LogicalLayoutHandle) -> remote LayoutHandle.
34    /// Populated by `connect_remote` for later use by `execute_remote_onboard_for_instance`.
35    remote_handles: RwLock<HashMap<(InstanceId, usize, LogicalLayoutHandle), LayoutHandle>>,
36}
37
38impl SpmdParallelWorkers {
39    /// Create a new SpmdParallelWorkers.
40    ///
41    /// # Arguments
42    /// * `workers` - The underlying workers (one per rank)
43    /// * `events` - The event system for aggregating completion notifications
44    /// * `runtime` - The tokio runtime handle for spawning aggregation tasks
45    pub fn new(
46        workers: Vec<Arc<dyn Worker>>,
47        events: Arc<::velo::EventManager>,
48        runtime: tokio::runtime::Handle,
49    ) -> Self {
50        Self {
51            workers,
52            events,
53            runtime,
54            remote_handles: RwLock::new(HashMap::new()),
55        }
56    }
57
58    /// Get the number of workers.
59    pub fn worker_count(&self) -> usize {
60        self.workers.len()
61    }
62}
63
64impl WorkerTransfers for SpmdParallelWorkers {
65    fn execute_local_transfer(
66        &self,
67        src: LogicalLayoutHandle,
68        dst: LogicalLayoutHandle,
69        src_block_ids: Arc<[BlockId]>,
70        dst_block_ids: Arc<[BlockId]>,
71        options: kvbm_physical::transfer::TransferOptions,
72    ) -> Result<TransferCompleteNotification> {
73        let notifications = self
74            .workers
75            .iter()
76            .map(|worker| {
77                worker.execute_local_transfer(
78                    src,
79                    dst,
80                    src_block_ids.clone(),
81                    dst_block_ids.clone(),
82                    options.clone(),
83                )
84            })
85            .collect::<Result<Vec<_>>>()?;
86
87        TransferCompleteNotification::aggregate(notifications, &self.events, &self.runtime)
88    }
89
90    fn execute_remote_onboard(
91        &self,
92        src: RemoteDescriptor,
93        dst: LogicalLayoutHandle,
94        dst_block_ids: Arc<[BlockId]>,
95        options: kvbm_physical::transfer::TransferOptions,
96    ) -> Result<TransferCompleteNotification> {
97        let notifications = self
98            .workers
99            .iter()
100            .map(|worker| {
101                worker.execute_remote_onboard(
102                    src.clone(),
103                    dst,
104                    dst_block_ids.clone(),
105                    options.clone(),
106                )
107            })
108            .collect::<Result<Vec<_>>>()?;
109
110        TransferCompleteNotification::aggregate(notifications, &self.events, &self.runtime)
111    }
112
113    fn execute_remote_offload(
114        &self,
115        src: LogicalLayoutHandle,
116        src_block_ids: Arc<[BlockId]>,
117        dst: RemoteDescriptor,
118        options: kvbm_physical::transfer::TransferOptions,
119    ) -> Result<TransferCompleteNotification> {
120        let notifications = self
121            .workers
122            .iter()
123            .map(|worker| {
124                worker.execute_remote_offload(
125                    src,
126                    src_block_ids.clone(),
127                    dst.clone(),
128                    options.clone(),
129                )
130            })
131            .collect::<Result<Vec<_>>>()?;
132
133        TransferCompleteNotification::aggregate(notifications, &self.events, &self.runtime)
134    }
135
136    fn connect_remote(
137        &self,
138        instance_id: InstanceId,
139        metadata: Vec<SerializedLayout>,
140    ) -> Result<ConnectRemoteResponse> {
141        // Validate metadata count matches worker count
142        if metadata.len() != self.workers.len() {
143            anyhow::bail!(
144                "Metadata count ({}) doesn't match worker count ({})",
145                metadata.len(),
146                self.workers.len()
147            );
148        }
149
150        // Collect handles to store and responses to await
151        let mut new_handles = Vec::new();
152        let mut import_responses = Vec::new();
153
154        for (worker_idx, (worker, meta)) in
155            self.workers.iter().zip(metadata.into_iter()).enumerate()
156        {
157            // Unpack to extract logical type info
158            let unpacked = meta.unpack()?;
159
160            // Collect handle mappings
161            for descriptor in &unpacked.layouts {
162                new_handles.push((
163                    (instance_id, worker_idx, descriptor.logical_type),
164                    descriptor.handle,
165                ));
166            }
167
168            // Repack for the underlying worker's import_metadata
169            let repacked = SerializedLayout::pack(
170                unpacked.worker_address,
171                unpacked.nixl_metadata,
172                unpacked.layouts,
173            )?;
174
175            // Call underlying worker's import_metadata
176            import_responses.push(worker.import_metadata(repacked)?);
177        }
178
179        // Store all handle mappings
180        {
181            let mut handles = self.remote_handles.write().unwrap();
182            for (key, value) in new_handles {
183                handles.insert(key, value);
184            }
185        }
186
187        // If all responses are ready (synchronous), return immediately
188        if import_responses.iter().all(|r| !r.could_yield()) {
189            return Ok(ConnectRemoteResponse::ready());
190        }
191
192        // Create an event to aggregate all import completions
193        let event = self.events.new_event()?;
194        let awaiter = self.events.awaiter(event.handle())?;
195
196        // Spawn task to await all import responses and signal completion
197        self.runtime
198            .spawn(await_import_responses(import_responses, event));
199
200        Ok(ConnectRemoteResponse::from_awaiter(awaiter))
201    }
202
203    fn has_remote_metadata(&self, instance_id: InstanceId) -> bool {
204        let handles = self.remote_handles.read().unwrap();
205        handles.keys().any(|(id, _, _)| *id == instance_id)
206    }
207
208    fn execute_remote_onboard_for_instance(
209        &self,
210        instance_id: InstanceId,
211        remote_logical_type: LogicalLayoutHandle,
212        src_block_ids: Vec<BlockId>,
213        dst: LogicalLayoutHandle,
214        dst_block_ids: Arc<[BlockId]>,
215        options: kvbm_physical::transfer::TransferOptions,
216    ) -> Result<TransferCompleteNotification> {
217        let handles = self.remote_handles.read().unwrap();
218        let mut notifications = Vec::with_capacity(self.workers.len());
219
220        // SPMD: Execute SAME transfer on EVERY worker, each with its own remote handle
221        for (worker_idx, worker) in self.workers.iter().enumerate() {
222            let remote_handle = handles
223                .get(&(instance_id, worker_idx, remote_logical_type))
224                .ok_or_else(|| {
225                    anyhow::anyhow!(
226                        "No remote {:?} handle for instance {} worker {}",
227                        remote_logical_type,
228                        instance_id,
229                        worker_idx
230                    )
231                })?;
232
233            let descriptor = RemoteDescriptor::Layout {
234                handle: *remote_handle,
235                block_ids: src_block_ids.clone(),
236            };
237
238            notifications.push(worker.execute_remote_onboard(
239                descriptor,
240                dst,
241                dst_block_ids.clone(),
242                options.clone(),
243            )?);
244        }
245
246        TransferCompleteNotification::aggregate(notifications, &self.events, &self.runtime)
247    }
248}
249
250/// Helper to await all import metadata responses and signal completion via an event.
251/// Helper to await all import metadata responses and signal completion via an event.
252async fn await_import_responses(responses: Vec<ImportMetadataResponse>, event: ::velo::Event) {
253    let results: Vec<Result<Vec<LayoutHandle>>> =
254        futures::future::join_all(responses.into_iter().map(|r| r.into_future())).await;
255
256    // Check for any failures
257    let errors: Vec<_> = results.into_iter().filter_map(|r| r.err()).collect();
258
259    if errors.is_empty() {
260        let _ = event.trigger();
261    } else {
262        let error_msg = errors
263            .iter()
264            .map(|e| e.to_string())
265            .collect::<Vec<_>>()
266            .join("; ");
267        let _ = event.poison(error_msg);
268    }
269}
270
271impl ParallelWorkers for SpmdParallelWorkers {
272    fn export_metadata(&self) -> Result<Vec<SerializedLayoutResponse>> {
273        let metadata = self
274            .workers
275            .iter()
276            .map(|worker| worker.export_metadata())
277            .collect::<Result<Vec<_>>>()?;
278
279        Ok(metadata)
280    }
281
282    fn import_metadata(
283        &self,
284        metadata: Vec<SerializedLayout>,
285    ) -> Result<Vec<ImportMetadataResponse>> {
286        // validate the size of the metadata is the same as the number of workers
287        if metadata.len() != self.workers.len() {
288            return Err(anyhow::anyhow!(
289                "Metadata size does not match number of workers"
290            ));
291        }
292
293        let results = self
294            .workers
295            .iter()
296            .zip(metadata.iter())
297            .map(|(worker, metadata)| worker.import_metadata(metadata.clone()))
298            .collect::<Result<Vec<_>>>()?;
299
300        Ok(results)
301    }
302
303    fn worker_count(&self) -> usize {
304        self.workers.len()
305    }
306
307    fn workers(&self) -> &[Arc<dyn Worker>] {
308        &self.workers
309    }
310}
311
312impl ObjectBlockOps for SpmdParallelWorkers {
313    fn has_blocks(
314        &self,
315        keys: Vec<SequenceHash>,
316    ) -> BoxFuture<'static, Vec<(SequenceHash, Option<usize>)>> {
317        // For has_blocks, we query all workers and verify consistency.
318        // All workers should agree on block presence for SPMD semantics.
319        // We return the results from worker 0 but verify all workers agree.
320        let workers = self.workers.clone();
321        let _runtime = self.runtime.clone();
322
323        Box::pin(async move {
324            if workers.is_empty() {
325                return keys.into_iter().map(|k| (k, None)).collect();
326            }
327
328            // Query all workers in parallel
329            let futures: Vec<_> = workers
330                .iter()
331                .map(|worker| worker.has_blocks(keys.clone()))
332                .collect();
333
334            let results: Vec<Vec<(SequenceHash, Option<usize>)>> =
335                futures::future::join_all(futures).await;
336
337            // Return results from first worker (all should agree in SPMD)
338            // In debug mode, we could verify consistency across workers
339            results.into_iter().next().unwrap_or_default()
340        })
341    }
342
343    fn put_blocks(
344        &self,
345        keys: Vec<SequenceHash>,
346        src_layout: LogicalLayoutHandle,
347        block_ids: Vec<BlockId>,
348    ) -> BoxFuture<'static, Vec<Result<SequenceHash, SequenceHash>>> {
349        // For put_blocks, each worker writes with its own rank-prefixed key.
350        // Each worker resolves the logical handle to its own physical layout.
351        // All workers must succeed for the operation to be considered successful.
352        let workers = self.workers.clone();
353
354        Box::pin(async move {
355            if workers.is_empty() {
356                return keys.into_iter().map(Err).collect();
357            }
358
359            // Execute put on all workers in parallel
360            // Each worker resolves src_layout to its own physical layout
361            let futures: Vec<_> = workers
362                .iter()
363                .map(|worker| worker.put_blocks(keys.clone(), src_layout, block_ids.clone()))
364                .collect();
365
366            let results: Vec<Vec<Result<SequenceHash, SequenceHash>>> =
367                futures::future::join_all(futures).await;
368
369            // Aggregate: a key succeeded only if ALL workers succeeded
370            let num_keys = keys.len();
371            let mut aggregated = Vec::with_capacity(num_keys);
372
373            for (key_idx, key) in keys.iter().enumerate() {
374                let all_succeeded = results.iter().all(|worker_results| {
375                    worker_results
376                        .get(key_idx)
377                        .map(|r| r.is_ok())
378                        .unwrap_or(false)
379                });
380
381                if all_succeeded {
382                    aggregated.push(Ok(*key));
383                } else {
384                    aggregated.push(Err(*key));
385                }
386            }
387
388            aggregated
389        })
390    }
391
392    fn get_blocks(
393        &self,
394        keys: Vec<SequenceHash>,
395        dst_layout: LogicalLayoutHandle,
396        block_ids: Vec<BlockId>,
397    ) -> BoxFuture<'static, Vec<Result<SequenceHash, SequenceHash>>> {
398        // For get_blocks, each worker reads from its own rank-prefixed key.
399        // Each worker resolves the logical handle to its own physical layout.
400        // All workers must succeed for the operation to be considered successful.
401        let workers = self.workers.clone();
402
403        Box::pin(async move {
404            if workers.is_empty() {
405                return keys.into_iter().map(Err).collect();
406            }
407
408            // Execute get on all workers in parallel
409            // Each worker resolves dst_layout to its own physical layout
410            let futures: Vec<_> = workers
411                .iter()
412                .map(|worker| worker.get_blocks(keys.clone(), dst_layout, block_ids.clone()))
413                .collect();
414
415            let results: Vec<Vec<Result<SequenceHash, SequenceHash>>> =
416                futures::future::join_all(futures).await;
417
418            // Aggregate: a key succeeded only if ALL workers succeeded
419            let num_keys = keys.len();
420            let mut aggregated = Vec::with_capacity(num_keys);
421
422            for (key_idx, key) in keys.iter().enumerate() {
423                let all_succeeded = results.iter().all(|worker_results| {
424                    worker_results
425                        .get(key_idx)
426                        .map(|r| r.is_ok())
427                        .unwrap_or(false)
428                });
429
430                if all_succeeded {
431                    aggregated.push(Ok(*key));
432                } else {
433                    aggregated.push(Err(*key));
434                }
435            }
436
437            aggregated
438        })
439    }
440}