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 self.workers.iter().zip(metadata).enumerate() {
155            // Unpack to extract logical type info
156            let unpacked = meta.unpack()?;
157
158            // Collect handle mappings
159            for descriptor in &unpacked.layouts {
160                new_handles.push((
161                    (instance_id, worker_idx, descriptor.logical_type),
162                    descriptor.handle,
163                ));
164            }
165
166            // Repack for the underlying worker's import_metadata
167            let repacked = SerializedLayout::pack(
168                unpacked.worker_address,
169                unpacked.nixl_metadata,
170                unpacked.layouts,
171            )?;
172
173            // Call underlying worker's import_metadata
174            import_responses.push(worker.import_metadata(repacked)?);
175        }
176
177        // Store all handle mappings
178        {
179            let mut handles = self.remote_handles.write().unwrap();
180            for (key, value) in new_handles {
181                handles.insert(key, value);
182            }
183        }
184
185        // If all responses are ready (synchronous), return immediately
186        if import_responses.iter().all(|r| !r.could_yield()) {
187            return Ok(ConnectRemoteResponse::ready());
188        }
189
190        // Create an event to aggregate all import completions
191        let event = self.events.new_event()?;
192        let awaiter = self.events.awaiter(event.handle())?;
193
194        // Spawn task to await all import responses and signal completion
195        self.runtime
196            .spawn(await_import_responses(import_responses, event));
197
198        Ok(ConnectRemoteResponse::from_awaiter(awaiter))
199    }
200
201    fn has_remote_metadata(&self, instance_id: InstanceId) -> bool {
202        let handles = self.remote_handles.read().unwrap();
203        handles.keys().any(|(id, _, _)| *id == instance_id)
204    }
205
206    fn execute_remote_onboard_for_instance(
207        &self,
208        instance_id: InstanceId,
209        remote_logical_type: LogicalLayoutHandle,
210        src_block_ids: Vec<BlockId>,
211        dst: LogicalLayoutHandle,
212        dst_block_ids: Arc<[BlockId]>,
213        options: kvbm_physical::transfer::TransferOptions,
214    ) -> Result<TransferCompleteNotification> {
215        let handles = self.remote_handles.read().unwrap();
216        let mut notifications = Vec::with_capacity(self.workers.len());
217
218        // SPMD: Execute SAME transfer on EVERY worker, each with its own remote handle
219        for (worker_idx, worker) in self.workers.iter().enumerate() {
220            let remote_handle = handles
221                .get(&(instance_id, worker_idx, remote_logical_type))
222                .ok_or_else(|| {
223                    anyhow::anyhow!(
224                        "No remote {:?} handle for instance {} worker {}",
225                        remote_logical_type,
226                        instance_id,
227                        worker_idx
228                    )
229                })?;
230
231            let descriptor = RemoteDescriptor::Layout {
232                handle: *remote_handle,
233                block_ids: src_block_ids.clone(),
234            };
235
236            notifications.push(worker.execute_remote_onboard(
237                descriptor,
238                dst,
239                dst_block_ids.clone(),
240                options.clone(),
241            )?);
242        }
243
244        TransferCompleteNotification::aggregate(notifications, &self.events, &self.runtime)
245    }
246}
247
248/// Helper to await all import metadata responses and signal completion via an event.
249/// Helper to await all import metadata responses and signal completion via an event.
250async fn await_import_responses(responses: Vec<ImportMetadataResponse>, event: ::velo::Event) {
251    let results: Vec<Result<Vec<LayoutHandle>>> =
252        futures::future::join_all(responses.into_iter().map(|r| r.into_future())).await;
253
254    // Check for any failures
255    let errors: Vec<_> = results.into_iter().filter_map(|r| r.err()).collect();
256
257    if errors.is_empty() {
258        let _ = event.trigger();
259    } else {
260        let error_msg = errors
261            .iter()
262            .map(|e| e.to_string())
263            .collect::<Vec<_>>()
264            .join("; ");
265        let _ = event.poison(error_msg);
266    }
267}
268
269impl ParallelWorkers for SpmdParallelWorkers {
270    fn export_metadata(&self) -> Result<Vec<SerializedLayoutResponse>> {
271        let metadata = self
272            .workers
273            .iter()
274            .map(|worker| worker.export_metadata())
275            .collect::<Result<Vec<_>>>()?;
276
277        Ok(metadata)
278    }
279
280    fn import_metadata(
281        &self,
282        metadata: Vec<SerializedLayout>,
283    ) -> Result<Vec<ImportMetadataResponse>> {
284        // validate the size of the metadata is the same as the number of workers
285        if metadata.len() != self.workers.len() {
286            return Err(anyhow::anyhow!(
287                "Metadata size does not match number of workers"
288            ));
289        }
290
291        let results = self
292            .workers
293            .iter()
294            .zip(metadata.iter())
295            .map(|(worker, metadata)| worker.import_metadata(metadata.clone()))
296            .collect::<Result<Vec<_>>>()?;
297
298        Ok(results)
299    }
300
301    fn worker_count(&self) -> usize {
302        self.workers.len()
303    }
304
305    fn workers(&self) -> &[Arc<dyn Worker>] {
306        &self.workers
307    }
308}
309
310impl ObjectBlockOps for SpmdParallelWorkers {
311    fn has_blocks(
312        &self,
313        keys: Vec<SequenceHash>,
314    ) -> BoxFuture<'static, Vec<(SequenceHash, Option<usize>)>> {
315        // For has_blocks, we query all workers and verify consistency.
316        // All workers should agree on block presence for SPMD semantics.
317        // We return the results from worker 0 but verify all workers agree.
318        let workers = self.workers.clone();
319        let _runtime = self.runtime.clone();
320
321        Box::pin(async move {
322            if workers.is_empty() {
323                return keys.into_iter().map(|k| (k, None)).collect();
324            }
325
326            // Query all workers in parallel
327            let futures: Vec<_> = workers
328                .iter()
329                .map(|worker| worker.has_blocks(keys.clone()))
330                .collect();
331
332            let results: Vec<Vec<(SequenceHash, Option<usize>)>> =
333                futures::future::join_all(futures).await;
334
335            // Return results from first worker (all should agree in SPMD)
336            // In debug mode, we could verify consistency across workers
337            results.into_iter().next().unwrap_or_default()
338        })
339    }
340
341    fn put_blocks(
342        &self,
343        keys: Vec<SequenceHash>,
344        src_layout: LogicalLayoutHandle,
345        block_ids: Vec<BlockId>,
346    ) -> BoxFuture<'static, Vec<Result<SequenceHash, SequenceHash>>> {
347        // For put_blocks, each worker writes with its own rank-prefixed key.
348        // Each worker resolves the logical handle to its own physical layout.
349        // All workers must succeed for the operation to be considered successful.
350        let workers = self.workers.clone();
351
352        Box::pin(async move {
353            if workers.is_empty() {
354                return keys.into_iter().map(Err).collect();
355            }
356
357            // Execute put on all workers in parallel
358            // Each worker resolves src_layout to its own physical layout
359            let futures: Vec<_> = workers
360                .iter()
361                .map(|worker| worker.put_blocks(keys.clone(), src_layout, block_ids.clone()))
362                .collect();
363
364            let results: Vec<Vec<Result<SequenceHash, SequenceHash>>> =
365                futures::future::join_all(futures).await;
366
367            // Aggregate: a key succeeded only if ALL workers succeeded
368            let num_keys = keys.len();
369            let mut aggregated = Vec::with_capacity(num_keys);
370
371            for (key_idx, key) in keys.iter().enumerate() {
372                let all_succeeded = results.iter().all(|worker_results| {
373                    worker_results
374                        .get(key_idx)
375                        .map(|r| r.is_ok())
376                        .unwrap_or(false)
377                });
378
379                if all_succeeded {
380                    aggregated.push(Ok(*key));
381                } else {
382                    aggregated.push(Err(*key));
383                }
384            }
385
386            aggregated
387        })
388    }
389
390    fn get_blocks(
391        &self,
392        keys: Vec<SequenceHash>,
393        dst_layout: LogicalLayoutHandle,
394        block_ids: Vec<BlockId>,
395    ) -> BoxFuture<'static, Vec<Result<SequenceHash, SequenceHash>>> {
396        // For get_blocks, each worker reads from its own rank-prefixed key.
397        // Each worker resolves the logical handle to its own physical layout.
398        // All workers must succeed for the operation to be considered successful.
399        let workers = self.workers.clone();
400
401        Box::pin(async move {
402            if workers.is_empty() {
403                return keys.into_iter().map(Err).collect();
404            }
405
406            // Execute get on all workers in parallel
407            // Each worker resolves dst_layout to its own physical layout
408            let futures: Vec<_> = workers
409                .iter()
410                .map(|worker| worker.get_blocks(keys.clone(), dst_layout, block_ids.clone()))
411                .collect();
412
413            let results: Vec<Vec<Result<SequenceHash, SequenceHash>>> =
414                futures::future::join_all(futures).await;
415
416            // Aggregate: a key succeeded only if ALL workers succeeded
417            let num_keys = keys.len();
418            let mut aggregated = Vec::with_capacity(num_keys);
419
420            for (key_idx, key) in keys.iter().enumerate() {
421                let all_succeeded = results.iter().all(|worker_results| {
422                    worker_results
423                        .get(key_idx)
424                        .map(|r| r.is_ok())
425                        .unwrap_or(false)
426                });
427
428                if all_succeeded {
429                    aggregated.push(Ok(*key));
430                } else {
431                    aggregated.push(Err(*key));
432                }
433            }
434
435            aggregated
436        })
437    }
438}