Skip to main content

kvbm_engine/worker/velo/
client.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use super::*;
5use crate::object::ObjectBlockOps;
6use futures::future::BoxFuture;
7use parking_lot::RwLock;
8use std::collections::HashSet;
9use std::sync::OnceLock;
10
11#[derive(Clone)]
12pub struct VeloWorkerClient {
13    messenger: Arc<Messenger>,
14    remote: InstanceId,
15    g1_handle: Arc<OnceLock<LayoutHandle>>,
16    g2_handle: Arc<OnceLock<LayoutHandle>>,
17    g3_handle: Arc<OnceLock<LayoutHandle>>,
18    /// Track which remote instances we've connected to for has_remote_metadata()
19    connected_instances: Arc<RwLock<HashSet<InstanceId>>>,
20}
21
22impl WorkerTransfers for VeloWorkerClient {
23    fn execute_local_transfer(
24        &self,
25        src: LogicalLayoutHandle,
26        dst: LogicalLayoutHandle,
27        src_block_ids: Arc<[BlockId]>,
28        dst_block_ids: Arc<[BlockId]>,
29        options: TransferOptions,
30    ) -> Result<TransferCompleteNotification> {
31        // Create a single local event for this operation
32        let event = self.messenger.events().new_event()?;
33        let awaiter = self.messenger.events().awaiter(event.handle())?;
34
35        // Convert to serializable options
36        // TODO: Extract bounce buffer handle if present in options.bounce_buffer
37        let options = SerializableTransferOptions {
38            layer_range: options.layer_range,
39            nixl_write_notification: options.nixl_write_notification,
40            bounce_buffer_handle: None,
41            bounce_buffer_block_ids: None,
42        };
43
44        // Create the message
45        let message = LocalTransferMessage {
46            src,
47            dst,
48            src_block_ids: src_block_ids.to_vec(),
49            dst_block_ids: dst_block_ids.to_vec(),
50            options,
51        };
52
53        let bytes = Bytes::from(serde_json::to_vec(&message)?);
54
55        // Spawn a task for the remote instance
56        let nova = self.messenger.clone();
57        let remote_instance = self.remote;
58
59        // Use unary (not am_sync) to wait for transfer completion
60        self.messenger.tracker().spawn_on(
61            async move {
62                let result = nova
63                    .unary("kvbm.worker.local_transfer")?
64                    .raw_payload(bytes)
65                    .instance(remote_instance)
66                    .send()
67                    .await;
68
69                match result {
70                    Ok(_) => event.trigger(),
71                    Err(e) => event.poison(e.to_string()),
72                }
73            },
74            self.messenger.runtime(),
75        );
76
77        Ok(TransferCompleteNotification::from_awaiter(awaiter))
78    }
79
80    fn execute_remote_onboard(
81        &self,
82        src: RemoteDescriptor,
83        dst: LogicalLayoutHandle,
84        dst_block_ids: Arc<[BlockId]>,
85        options: TransferOptions,
86    ) -> Result<TransferCompleteNotification> {
87        let event = self.messenger.events().new_event()?;
88        let awaiter = self.messenger.events().awaiter(event.handle())?;
89
90        let options = SerializableTransferOptions {
91            layer_range: options.layer_range,
92            nixl_write_notification: options.nixl_write_notification,
93            bounce_buffer_handle: None,
94            bounce_buffer_block_ids: None,
95        };
96
97        let message = RemoteOnboardMessage {
98            src,
99            dst,
100            dst_block_ids: dst_block_ids.to_vec(),
101            options,
102        };
103
104        let bytes = Bytes::from(serde_json::to_vec(&message)?);
105
106        let nova = self.messenger.clone();
107        let remote_instance = self.remote;
108
109        self.messenger.tracker().spawn_on(
110            async move {
111                // Use unary instead of am_sync for explicit response handling
112                let result = nova
113                    .unary("kvbm.worker.remote_onboard")?
114                    .raw_payload(bytes)
115                    .instance(remote_instance)
116                    .send()
117                    .await;
118
119                match result {
120                    Ok(_) => event.trigger(),
121                    Err(e) => event.poison(e.to_string()),
122                }
123            },
124            self.messenger.runtime(),
125        );
126
127        Ok(TransferCompleteNotification::from_awaiter(awaiter))
128    }
129
130    fn execute_remote_offload(
131        &self,
132        src: LogicalLayoutHandle,
133        src_block_ids: Arc<[BlockId]>,
134        dst: RemoteDescriptor,
135        options: TransferOptions,
136    ) -> Result<TransferCompleteNotification> {
137        let event = self.messenger.events().new_event()?;
138        let awaiter = self.messenger.events().awaiter(event.handle())?;
139
140        let options = SerializableTransferOptions {
141            layer_range: options.layer_range,
142            nixl_write_notification: options.nixl_write_notification,
143            bounce_buffer_handle: None,
144            bounce_buffer_block_ids: None,
145        };
146
147        let message = RemoteOffloadMessage {
148            src,
149            dst,
150            src_block_ids: src_block_ids.to_vec(),
151            options,
152        };
153
154        let bytes = Bytes::from(serde_json::to_vec(&message)?);
155
156        let nova = self.messenger.clone();
157        let remote_instance = self.remote;
158
159        self.messenger.tracker().spawn_on(
160            async move {
161                // Use unary instead of am_sync for explicit response handling
162                let result = nova
163                    .unary("kvbm.worker.remote_offload")?
164                    .raw_payload(bytes)
165                    .instance(remote_instance)
166                    .send()
167                    .await;
168
169                match result {
170                    Ok(_) => event.trigger(),
171                    Err(e) => event.poison(e.to_string()),
172                }
173            },
174            self.messenger.runtime(),
175        );
176
177        Ok(TransferCompleteNotification::from_awaiter(awaiter))
178    }
179
180    fn connect_remote(
181        &self,
182        instance_id: InstanceId,
183        metadata: Vec<SerializedLayout>,
184    ) -> Result<ConnectRemoteResponse> {
185        // Serialize metadata to bytes (SerializedLayout uses bincode internally)
186        let serialized_metadata: Vec<Vec<u8>> =
187            metadata.iter().map(|m| m.as_bytes().to_vec()).collect();
188
189        let message = ConnectRemoteMessage {
190            instance_id,
191            metadata: serialized_metadata,
192        };
193        let bytes = Bytes::from(serde_json::to_vec(&message)?);
194
195        // Create event for completion tracking
196        let event = self.messenger.events().new_event()?;
197        let awaiter = self.messenger.events().awaiter(event.handle())?;
198
199        let nova = self.messenger.clone();
200        let remote_instance = self.remote;
201        let connected = self.connected_instances.clone();
202        let target_instance = instance_id;
203
204        self.messenger.tracker().spawn_on(
205            async move {
206                let result = nova
207                    .unary("kvbm.worker.connect_remote")?
208                    .raw_payload(bytes)
209                    .instance(remote_instance)
210                    .send()
211                    .await;
212
213                match result {
214                    Ok(_) => {
215                        // Track that we've connected to this instance
216                        connected.write().insert(target_instance);
217                        event.trigger()
218                    }
219                    Err(e) => event.poison(e.to_string()),
220                }
221            },
222            self.messenger.runtime(),
223        );
224
225        Ok(ConnectRemoteResponse::from_awaiter(awaiter))
226    }
227
228    fn has_remote_metadata(&self, instance_id: InstanceId) -> bool {
229        // Check if we've successfully connected to this instance
230        self.connected_instances.read().contains(&instance_id)
231    }
232
233    fn execute_remote_onboard_for_instance(
234        &self,
235        instance_id: InstanceId,
236        remote_logical_type: LogicalLayoutHandle,
237        src_block_ids: Vec<BlockId>,
238        dst: LogicalLayoutHandle,
239        dst_block_ids: Arc<[BlockId]>,
240        options: TransferOptions,
241    ) -> Result<TransferCompleteNotification> {
242        let message = ExecuteRemoteOnboardForInstanceMessage {
243            instance_id,
244            remote_logical_type,
245            src_block_ids,
246            dst,
247            dst_block_ids: dst_block_ids.to_vec(),
248            options: SerializableTransferOptions::from(options),
249        };
250        let bytes = Bytes::from(serde_json::to_vec(&message)?);
251
252        // Create event for completion tracking
253        let event = self.messenger.events().new_event()?;
254        let awaiter = self.messenger.events().awaiter(event.handle())?;
255
256        let nova = self.messenger.clone();
257        let remote_instance = self.remote;
258
259        self.messenger.tracker().spawn_on(
260            async move {
261                let result = nova
262                    .unary("kvbm.worker.remote_onboard_for_instance")?
263                    .raw_payload(bytes)
264                    .instance(remote_instance)
265                    .send()
266                    .await;
267
268                match result {
269                    Ok(_) => event.trigger(),
270                    Err(e) => event.poison(e.to_string()),
271                }
272            },
273            self.messenger.runtime(),
274        );
275
276        Ok(TransferCompleteNotification::from_awaiter(awaiter))
277    }
278}
279
280impl Worker for VeloWorkerClient {
281    fn g1_handle(&self) -> Option<LayoutHandle> {
282        self.g1_handle.get().copied()
283    }
284
285    fn g2_handle(&self) -> Option<LayoutHandle> {
286        self.g2_handle.get().copied()
287    }
288
289    fn g3_handle(&self) -> Option<LayoutHandle> {
290        self.g3_handle.get().copied()
291    }
292
293    fn export_metadata(&self) -> Result<SerializedLayoutResponse> {
294        // Use unary (not typed_unary) to avoid JSON serialization of bincode data
295        let unary_result = self
296            .messenger
297            .unary("kvbm.worker.export_metadata")?
298            .instance(self.remote)
299            .send();
300
301        // Wrap UnaryResult to convert Bytes to SerializedLayout
302        let future = async move {
303            let bytes = unary_result.await?;
304            Ok(SerializedLayout::from_bytes(bytes.to_vec()))
305        };
306
307        Ok(SerializedLayoutResponse::from_boxed(Box::pin(future)))
308    }
309
310    fn import_metadata(&self, metadata: SerializedLayout) -> Result<ImportMetadataResponse> {
311        // Use raw_payload to avoid JSON serialization of bincode data
312        let unary_result = self
313            .messenger
314            .unary("kvbm.worker.import_metadata")?
315            .raw_payload(Bytes::from(metadata.as_bytes().to_vec()))
316            .instance(self.remote)
317            .send();
318
319        // Response is JSON-serialized Vec<LayoutHandle>
320        let future = async move {
321            let bytes = unary_result.await?;
322            serde_json::from_slice(&bytes).map_err(|e| {
323                anyhow::anyhow!("Failed to deserialize import_metadata response: {}", e)
324            })
325        };
326
327        Ok(ImportMetadataResponse::from_boxed(Box::pin(future)))
328    }
329}
330
331impl VeloWorkerClient {
332    /// Create a new VeloWorkerClient for communicating with a remote worker.
333    pub fn new(messenger: Arc<Messenger>, remote: InstanceId) -> Self {
334        Self {
335            messenger,
336            remote,
337            g1_handle: Arc::new(OnceLock::new()),
338            g2_handle: Arc::new(OnceLock::new()),
339            g3_handle: Arc::new(OnceLock::new()),
340            connected_instances: Arc::new(RwLock::new(HashSet::new())),
341        }
342    }
343
344    /// Configure layout handles from serialized metadata.
345    ///
346    /// Call this after worker initialization when handles are known from WorkerLayoutResponse.
347    /// This allows the VeloWorkerClient to provide layout handles like DirectWorker does.
348    ///
349    /// # Arguments
350    /// * `metadata` - SerializedLayout from WorkerLayoutResponse.metadata
351    ///
352    /// # Example
353    /// ```ignore
354    /// let response: WorkerLayoutResponse = worker.initialize(config).await?;
355    /// worker_client.configure_layout_handles(&response.metadata)?;
356    /// ```
357    pub fn configure_layout_handles(&self, metadata: &SerializedLayout) -> Result<()> {
358        let unpacked = metadata.unpack()?;
359        for desc in &unpacked.layouts {
360            match desc.logical_type {
361                LogicalLayoutHandle::G1 => {
362                    self.g1_handle.set(desc.handle).ok();
363                }
364                LogicalLayoutHandle::G2 => {
365                    self.g2_handle.set(desc.handle).ok();
366                }
367                LogicalLayoutHandle::G3 => {
368                    self.g3_handle.set(desc.handle).ok();
369                }
370                _ => {}
371            }
372        }
373        Ok(())
374    }
375
376    /// Get the layout configuration from the remote worker.
377    ///
378    /// This calls the `kvbm.worker.get_layout_config` handler on the remote worker.
379    /// Used by the leader during Phase 3 to gather G1 layout configs from all workers
380    /// and validate they match before creating G2/G3 layouts.
381    ///
382    /// # Returns
383    /// A typed unary result that resolves to the layout configuration
384    pub fn get_layout_config(&self) -> Result<::velo::TypedUnaryResult<LayoutConfig>> {
385        let instance = self.remote;
386
387        let awaiter = self
388            .messenger
389            .typed_unary::<LayoutConfig>("kvbm.worker.get_layout_config")?
390            .instance(instance)
391            .send();
392
393        Ok(awaiter)
394    }
395
396    /// Configure additional layouts (G2, G3) on the remote worker.
397    ///
398    /// This calls the `kvbm.worker.configure_layouts` handler on the remote worker.
399    /// The worker will create host/pinned cache (G2) and optionally disk cache (G3)
400    /// based on the provided configuration.
401    ///
402    /// # Arguments
403    /// * `config` - Leader-provided configuration specifying block counts and backends
404    ///
405    /// # Returns
406    /// A typed unary result that resolves to the worker's response with updated metadata
407    pub fn configure_layouts(
408        &self,
409        config: LeaderLayoutConfig,
410    ) -> Result<::velo::TypedUnaryResult<WorkerLayoutResponse>> {
411        let instance = self.remote;
412
413        let awaiter = self
414            .messenger
415            .typed_unary::<WorkerLayoutResponse>("kvbm.worker.configure_layouts")?
416            .payload(config)?
417            .instance(instance)
418            .send();
419
420        Ok(awaiter)
421    }
422}
423
424impl ObjectBlockOps for VeloWorkerClient {
425    fn has_blocks(
426        &self,
427        keys: Vec<SequenceHash>,
428    ) -> BoxFuture<'static, Vec<(SequenceHash, Option<usize>)>> {
429        let message = ObjectHasBlocksMessage { keys: keys.clone() };
430        let bytes = match serde_json::to_vec(&message) {
431            Ok(b) => Bytes::from(b),
432            Err(_) => {
433                return Box::pin(async move { keys.into_iter().map(|k| (k, None)).collect() });
434            }
435        };
436
437        let nova = self.messenger.clone();
438        let remote = self.remote;
439
440        Box::pin(async move {
441            let result = nova
442                .unary("kvbm.worker.object_has_blocks")
443                .ok()
444                .map(|u| u.raw_payload(bytes).instance(remote).send());
445
446            match result {
447                Some(unary_result) => match unary_result.await {
448                    Ok(response_bytes) => {
449                        match serde_json::from_slice::<ObjectHasBlocksResponse>(&response_bytes) {
450                            Ok(response) => response.results,
451                            Err(_) => keys.into_iter().map(|k| (k, None)).collect(),
452                        }
453                    }
454                    Err(_) => keys.into_iter().map(|k| (k, None)).collect(),
455                },
456                None => keys.into_iter().map(|k| (k, None)).collect(),
457            }
458        })
459    }
460
461    fn put_blocks(
462        &self,
463        keys: Vec<SequenceHash>,
464        src_layout: LogicalLayoutHandle,
465        block_ids: Vec<BlockId>,
466    ) -> BoxFuture<'static, Vec<Result<SequenceHash, SequenceHash>>> {
467        // For remote workers, we send the logical layout handle - they resolve it locally
468        let message = ObjectPutBlocksMessage {
469            keys: keys.clone(),
470            layout: src_layout,
471            block_ids,
472        };
473        let bytes = match serde_json::to_vec(&message) {
474            Ok(b) => Bytes::from(b),
475            Err(_) => return Box::pin(async move { keys.into_iter().map(Err).collect() }),
476        };
477
478        let nova = self.messenger.clone();
479        let remote = self.remote;
480
481        Box::pin(async move {
482            let result = nova
483                .unary("kvbm.worker.object_put_blocks")
484                .ok()
485                .map(|u| u.raw_payload(bytes).instance(remote).send());
486
487            match result {
488                Some(unary_result) => match unary_result.await {
489                    Ok(response_bytes) => {
490                        match serde_json::from_slice::<ObjectPutGetBlocksResponse>(&response_bytes)
491                        {
492                            Ok(response) => response.into_results(),
493                            Err(_) => keys.into_iter().map(Err).collect(),
494                        }
495                    }
496                    Err(_) => keys.into_iter().map(Err).collect(),
497                },
498                None => keys.into_iter().map(Err).collect(),
499            }
500        })
501    }
502
503    fn get_blocks(
504        &self,
505        keys: Vec<SequenceHash>,
506        dst_layout: LogicalLayoutHandle,
507        block_ids: Vec<BlockId>,
508    ) -> BoxFuture<'static, Vec<Result<SequenceHash, SequenceHash>>> {
509        // For remote workers, we send the logical layout handle - they resolve it locally
510        let message = ObjectGetBlocksMessage {
511            keys: keys.clone(),
512            layout: dst_layout,
513            block_ids,
514        };
515        let bytes = match serde_json::to_vec(&message) {
516            Ok(b) => Bytes::from(b),
517            Err(_) => return Box::pin(async move { keys.into_iter().map(Err).collect() }),
518        };
519
520        let nova = self.messenger.clone();
521        let remote = self.remote;
522
523        Box::pin(async move {
524            let result = nova
525                .unary("kvbm.worker.object_get_blocks")
526                .ok()
527                .map(|u| u.raw_payload(bytes).instance(remote).send());
528
529            match result {
530                Some(unary_result) => match unary_result.await {
531                    Ok(response_bytes) => {
532                        match serde_json::from_slice::<ObjectPutGetBlocksResponse>(&response_bytes)
533                        {
534                            Ok(response) => response.into_results(),
535                            Err(_) => keys.into_iter().map(Err).collect(),
536                        }
537                    }
538                    Err(_) => keys.into_iter().map(Err).collect(),
539                },
540                None => keys.into_iter().map(Err).collect(),
541            }
542        })
543    }
544}