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
155 self.workers.iter().zip(metadata.into_iter()).enumerate()
156 {
157 let unpacked = meta.unpack()?;
159
160 for descriptor in &unpacked.layouts {
162 new_handles.push((
163 (instance_id, worker_idx, descriptor.logical_type),
164 descriptor.handle,
165 ));
166 }
167
168 let repacked = SerializedLayout::pack(
170 unpacked.worker_address,
171 unpacked.nixl_metadata,
172 unpacked.layouts,
173 )?;
174
175 import_responses.push(worker.import_metadata(repacked)?);
177 }
178
179 {
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 import_responses.iter().all(|r| !r.could_yield()) {
189 return Ok(ConnectRemoteResponse::ready());
190 }
191
192 let event = self.events.new_event()?;
194 let awaiter = self.events.awaiter(event.handle())?;
195
196 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 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
250async 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 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 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 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 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 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 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 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 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 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 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 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}