kvbm_engine/worker/group/
spmd.rs1use super::*;
5
6use crate::object::ObjectBlockOps;
7use anyhow::Result;
8use futures::future::BoxFuture;
10
11use std::collections::HashMap;
12use std::sync::{Arc, RwLock};
13
14pub struct SpmdParallelWorkers {
29 workers: Vec<Arc<dyn Worker>>,
30 events: Arc<::velo::EventManager>,
31 runtime: tokio::runtime::Handle,
32
33 remote_handles: RwLock<HashMap<(InstanceId, usize, LogicalLayoutHandle), LayoutHandle>>,
36}
37
38impl SpmdParallelWorkers {
39 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 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 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 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 let unpacked = meta.unpack()?;
157
158 for descriptor in &unpacked.layouts {
160 new_handles.push((
161 (instance_id, worker_idx, descriptor.logical_type),
162 descriptor.handle,
163 ));
164 }
165
166 let repacked = SerializedLayout::pack(
168 unpacked.worker_address,
169 unpacked.nixl_metadata,
170 unpacked.layouts,
171 )?;
172
173 import_responses.push(worker.import_metadata(repacked)?);
175 }
176
177 {
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 import_responses.iter().all(|r| !r.could_yield()) {
187 return Ok(ConnectRemoteResponse::ready());
188 }
189
190 let event = self.events.new_event()?;
192 let awaiter = self.events.awaiter(event.handle())?;
193
194 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 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
248async 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 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 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 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 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 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 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 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 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 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 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 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}