1mod handle;
7mod local;
8mod metadata;
9mod remote;
10
11pub use handle::LayoutHandle;
12pub use metadata::{LogicalLayoutDescriptor, SerializedLayout, WorkerAddress};
13
14pub(crate) use local::LocalLayout;
15pub(crate) use metadata::LocalLayoutDescriptor;
16pub(crate) use remote::RemoteLayout;
17
18use crate::layout::PhysicalLayout;
19use crate::transfer::BounceBufferInternal;
20use crate::transfer::TransferContext;
21use crate::transfer::context::TransferCompleteNotification;
22use crate::transfer::executor::TransferOptionsInternal;
23use crate::transfer::options::TransferOptions;
24use crate::{BlockId, SequenceHash};
25use anyhow::{Result, anyhow, bail};
26use dynamo_memory::StorageKind;
27use dynamo_memory::nixl::NixlAgent;
28use kvbm_common::LogicalLayoutHandle;
29use std::collections::{HashMap, HashSet};
30use std::sync::atomic::{AtomicU16, Ordering};
31use std::sync::{Arc, RwLock};
32
33#[derive(Clone)]
42pub struct TransferManager {
43 registry: Arc<RwLock<LayoutRegistry>>,
44 context: Arc<TransferContext>,
45}
46
47impl TransferManager {
48 pub fn builder() -> crate::transfer::context::TransferConfigBuilder {
68 TransferContext::builder()
69 }
70
71 pub(crate) fn from_context(context: TransferContext) -> Self {
76 let worker_id = context.worker_id();
77 let nixl_agent = context.nixl_agent().clone();
78 let registry = Arc::new(RwLock::new(LayoutRegistry::new(nixl_agent, worker_id)));
79
80 Self {
81 registry,
82 context: Arc::new(context),
83 }
84 }
85
86 pub fn register_layout(&self, layout: PhysicalLayout) -> Result<LayoutHandle> {
102 self.registry.write().unwrap().register_local(layout)
103 }
104
105 pub fn export_metadata(&self) -> Result<SerializedLayout> {
113 self.registry.read().unwrap().export_metadata()
114 }
115
116 pub fn import_metadata(&self, metadata: SerializedLayout) -> Result<Vec<LayoutHandle>> {
131 self.registry.write().unwrap().import_metadata(metadata)
132 }
133
134 pub fn build_logical_descriptor(
151 &self,
152 handle: LayoutHandle,
153 logical_type: LogicalLayoutHandle,
154 ) -> Result<LogicalLayoutDescriptor> {
155 self.registry
156 .read()
157 .unwrap()
158 .build_logical_descriptor(handle, logical_type)
159 }
160
161 pub fn get_nixl_metadata(&self) -> Result<Vec<u8>> {
165 self.registry.read().unwrap().get_nixl_metadata()
166 }
167
168 pub fn worker_address(&self) -> WorkerAddress {
170 self.registry.read().unwrap().worker_address()
171 }
172
173 pub fn nixl_agent(&self) -> &NixlAgent {
178 self.context.nixl_agent()
179 }
180
181 pub fn get_layout_config(&self, handle: LayoutHandle) -> Result<crate::layout::LayoutConfig> {
195 let registry = self.registry.read().unwrap();
196 let physical_layout = registry
197 .get_layout(handle)
198 .ok_or_else(|| anyhow!("invalid handle: {}", handle))?;
199 Ok(physical_layout.layout().config().clone())
200 }
201
202 pub fn execute_transfer(
228 &self,
229 src_handle: LayoutHandle,
230 src_blocks: &[BlockId],
231 dst_handle: LayoutHandle,
232 dst_blocks: &[BlockId],
233 options: TransferOptions,
234 ) -> Result<TransferCompleteNotification> {
235 let (src_layout, dst_layout) = {
237 let registry = self.registry.read().unwrap();
238 let src = registry
239 .get_layout(src_handle)
240 .ok_or_else(|| anyhow!("invalid source handle: {}", src_handle))?
241 .clone(); let dst = registry
243 .get_layout(dst_handle)
244 .ok_or_else(|| anyhow!("invalid destination handle: {}", dst_handle))?
245 .clone();
246 (src, dst)
247 }; let (
250 layer_range,
251 nixl_write_notification,
252 bounce_buffer,
253 cuda_stream,
254 src_kv_layout,
255 dst_kv_layout,
256 ) = options.dissolve();
257
258 let mut internal_options = TransferOptionsInternal::builder();
259
260 if let Some(range) = layer_range {
261 internal_options = internal_options.layer_range(range);
262 }
263
264 if let Some(notification) = nixl_write_notification {
265 internal_options = internal_options.nixl_write_notification(notification);
266 }
267
268 if let Some(bounce) = bounce_buffer {
269 let (handle, block_ids) = bounce.into_parts();
270 let bounce_buffer = self.create_bounce_buffer(handle, block_ids)?;
271 internal_options = internal_options.bounce_buffer(bounce_buffer);
272 }
273
274 if let Some(stream) = cuda_stream {
275 internal_options = internal_options.cuda_stream(stream);
276 }
277
278 if let Some(layout) = src_kv_layout {
279 internal_options = internal_options.src_kv_layout(layout);
280 }
281
282 if let Some(layout) = dst_kv_layout {
283 internal_options = internal_options.dst_kv_layout(layout);
284 }
285
286 let options = internal_options.build()?;
287
288 tracing::debug!(
289 src_handle = src_handle.to_string(),
290 dst_handle = dst_handle.to_string(),
291 "Executing transfer; src_blocks = {:?}; dst_blocks = {:?}",
292 src_blocks,
293 dst_blocks,
294 );
295
296 super::transfer::executor::execute_transfer(
298 &src_layout,
299 &dst_layout,
300 src_blocks,
301 dst_blocks,
302 options,
303 &self.context,
304 )
305 }
306
307 pub fn execute_g4_offload(
315 _src_handle: LayoutHandle,
316 _src_blocks: &[BlockId],
317 _dst_object: &[SequenceHash],
318 _options: TransferOptions, ) -> Result<TransferCompleteNotification> {
320 todo!("implement remote offload")
325 }
326
327 pub fn execute_g4_onboard() {
328 todo!("implement remote onboard")
329 }
330
331 pub fn worker_id(&self) -> u64 {
335 self.context.worker_id()
336 }
337
338 pub fn get_local_handles(&self) -> Vec<LayoutHandle> {
340 self.registry.read().unwrap().local_handles()
341 }
342
343 pub fn get_remote_handles(&self) -> Vec<LayoutHandle> {
345 self.registry.read().unwrap().remote_handles()
346 }
347
348 pub fn get_physical_layout(&self, handle: LayoutHandle) -> Option<PhysicalLayout> {
356 self.registry.read().unwrap().get_layout(handle).cloned()
357 }
358
359 pub(crate) fn create_bounce_buffer(
364 &self,
365 handle: LayoutHandle,
366 block_ids: Vec<BlockId>,
367 ) -> Result<BounceBufferInternal> {
368 let layout = {
369 let registry = self.registry.read().unwrap();
370 registry
371 .get_layout(handle)
372 .ok_or_else(|| anyhow!("invalid bounce buffer handle: {}", handle))?
373 .clone()
374 };
375
376 Ok(BounceBufferInternal::from_layout(layout, block_ids))
377 }
378
379 #[doc(hidden)]
383 pub fn context(&self) -> &TransferContext {
384 &self.context
385 }
386
387 #[doc(hidden)]
392 pub fn registry(&self) -> &RwLock<LayoutRegistry> {
393 &self.registry
394 }
395
396 #[cfg(test)]
398 #[allow(dead_code)]
399 pub(crate) fn h2d_stream(&self) -> &std::sync::Arc<cudarc::driver::CudaStream> {
400 self.context.h2d_stream()
401 }
402
403 #[cfg(test)]
405 #[allow(dead_code)]
406 pub(crate) fn d2h_stream(&self) -> &std::sync::Arc<cudarc::driver::CudaStream> {
407 self.context.d2h_stream()
408 }
409
410 #[cfg(test)]
412 #[allow(dead_code)]
413 pub(crate) fn cuda_context(&self) -> &std::sync::Arc<cudarc::driver::CudaContext> {
414 self.context.cuda_context()
415 }
416
417 #[cfg(test)]
419 #[allow(dead_code)]
420 pub(crate) fn register_cuda_event(
421 &self,
422 event: cudarc::driver::CudaEvent,
423 ) -> TransferCompleteNotification {
424 self.context.register_cuda_event(event)
425 }
426
427 #[cfg(test)]
429 #[expect(dead_code)]
430 pub(crate) fn cuda_pool(&self) -> &std::sync::Arc<dynamo_memory::CudaMemPool> {
431 self.context.cuda_pool()
432 }
433}
434
435#[derive(Debug)]
443#[doc(hidden)]
444pub struct LayoutRegistry {
445 nixl_agent: NixlAgent,
447 worker_id: u64,
449 next_layout_id: AtomicU16,
451 local_layouts: HashMap<LayoutHandle, LocalLayout>,
453 remote_layouts: HashMap<LayoutHandle, RemoteLayout>,
455 loaded_remotes: HashSet<(String, u64)>,
457}
458
459#[expect(dead_code)]
460impl LayoutRegistry {
461 pub(crate) fn new(nixl_agent: NixlAgent, worker_id: u64) -> Self {
467 Self {
468 nixl_agent,
469 worker_id,
470 next_layout_id: AtomicU16::new(0),
471 local_layouts: HashMap::new(),
472 remote_layouts: HashMap::new(),
473 loaded_remotes: HashSet::new(),
474 }
475 }
476
477 pub(crate) fn register_local(&mut self, layout: PhysicalLayout) -> Result<LayoutHandle> {
488 let current = self.next_layout_id.load(Ordering::SeqCst);
490 if current == u16::MAX {
491 bail!(
492 "Layout ID overflow: maximum number of layouts ({}) reached",
493 u16::MAX
494 );
495 }
496 let layout_id = self.next_layout_id.fetch_add(1, Ordering::SeqCst);
497
498 let handle = LayoutHandle::new(self.worker_id, layout_id);
500
501 let local_layout = LocalLayout::new(handle, layout);
503
504 self.local_layouts.insert(handle, local_layout);
506
507 Ok(handle)
508 }
509
510 pub(crate) fn export_metadata(&self) -> Result<SerializedLayout> {
520 let nixl_metadata = self
522 .nixl_agent
523 .get_local_md()
524 .map_err(|e| anyhow!("failed to get NIXL local metadata: {:?}", e))?;
525
526 let worker_address = WorkerAddress::new(self.worker_id, self.nixl_agent.name().to_string());
528
529 let mut serialized_layouts = Vec::new();
531 for (handle, local_layout) in &self.local_layouts {
532 let location = local_layout.layout().location();
533
534 if matches!(
536 location,
537 StorageKind::System | StorageKind::Device(_) | StorageKind::Pinned
538 ) {
539 let serialized = local_layout
540 .layout()
541 .to_descriptor()
542 .map_err(|e| anyhow!("failed to serialize layout {}: {}", handle, e))?;
543
544 serialized_layouts.push(LocalLayoutDescriptor::new_with_default_type(
545 *handle, serialized,
546 ));
547 }
548 }
549
550 SerializedLayout::pack(worker_address, nixl_metadata, serialized_layouts)
552 }
553
554 pub(crate) fn import_metadata(
575 &mut self,
576 metadata: SerializedLayout,
577 ) -> Result<Vec<LayoutHandle>> {
578 let inner = metadata.unpack()?;
580
581 let remote_key = (
583 inner.worker_address.nixl_agent_name.clone(),
584 inner.worker_address.worker_id,
585 );
586 if self.loaded_remotes.contains(&remote_key) {
587 bail!(
588 "Remote worker already loaded: {} (worker_id={})",
589 remote_key.0,
590 remote_key.1
591 );
592 }
593
594 let returned_agent_name = self
596 .nixl_agent
597 .load_remote_md(&inner.nixl_metadata)
598 .map_err(|e| anyhow!("failed to load remote NIXL metadata: {:?}", e))?;
599
600 if returned_agent_name != inner.worker_address.nixl_agent_name {
602 bail!(
603 "Agent name mismatch: expected '{}', got '{}'",
604 inner.worker_address.nixl_agent_name,
605 returned_agent_name
606 );
607 }
608
609 let mut imported_handles = Vec::new();
611 for serialized_with_handle in inner.layouts {
612 let handle = serialized_with_handle.handle;
613 let layout = PhysicalLayout::from_descriptor(serialized_with_handle.layout)
614 .map_err(|e| anyhow!("failed to reconstruct layout {}: {}", handle, e))?;
615
616 let remote_layout = RemoteLayout::new(handle, layout);
617 self.remote_layouts.insert(handle, remote_layout);
618 imported_handles.push(handle);
619 }
620
621 self.loaded_remotes.insert(remote_key);
623
624 Ok(imported_handles)
625 }
626
627 pub(crate) fn build_logical_descriptor(
636 &self,
637 handle: LayoutHandle,
638 logical_type: LogicalLayoutHandle,
639 ) -> Result<LogicalLayoutDescriptor> {
640 let local_layout = self
641 .local_layouts
642 .get(&handle)
643 .ok_or_else(|| anyhow!("Layout handle not found: {:?}", handle))?;
644
645 let layout_descriptor = local_layout
646 .layout()
647 .to_descriptor()
648 .map_err(|e| anyhow!("failed to serialize layout {}: {}", handle, e))?;
649
650 Ok(LogicalLayoutDescriptor::new(
651 handle,
652 logical_type,
653 layout_descriptor,
654 ))
655 }
656
657 pub(crate) fn get_nixl_metadata(&self) -> Result<Vec<u8>> {
659 self.nixl_agent
660 .get_local_md()
661 .map_err(|e| anyhow!("failed to get NIXL local metadata: {:?}", e))
662 }
663
664 pub(crate) fn worker_address(&self) -> WorkerAddress {
666 WorkerAddress::new(self.worker_id, self.nixl_agent.name().to_string())
667 }
668
669 pub(crate) fn get_local(&self, handle: LayoutHandle) -> Option<&LocalLayout> {
671 self.local_layouts.get(&handle)
672 }
673
674 pub(crate) fn get_remote(&self, handle: LayoutHandle) -> Option<&RemoteLayout> {
676 self.remote_layouts.get(&handle)
677 }
678
679 pub fn get_layout(&self, handle: LayoutHandle) -> Option<&PhysicalLayout> {
684 self.local_layouts
685 .get(&handle)
686 .map(|l| l.layout())
687 .or_else(|| self.remote_layouts.get(&handle).map(|r| r.layout()))
688 }
689
690 pub(crate) fn is_local(&self, handle: LayoutHandle) -> bool {
692 self.local_layouts.contains_key(&handle)
693 }
694
695 pub(crate) fn is_remote(&self, handle: LayoutHandle) -> bool {
697 self.remote_layouts.contains_key(&handle)
698 }
699
700 pub(crate) fn local_count(&self) -> usize {
702 self.local_layouts.len()
703 }
704
705 pub(crate) fn remote_count(&self) -> usize {
707 self.remote_layouts.len()
708 }
709
710 pub(crate) fn worker_id(&self) -> u64 {
712 self.worker_id
713 }
714
715 pub(crate) fn local_handles(&self) -> Vec<LayoutHandle> {
717 self.local_layouts.keys().copied().collect()
718 }
719
720 pub(crate) fn remote_handles(&self) -> Vec<LayoutHandle> {
722 self.remote_layouts.keys().copied().collect()
723 }
724}
725
726#[cfg(all(test, feature = "testing-kvbm"))]
727mod tests {
728 use super::*;
729 use crate::layout::LayoutConfig;
730 use dynamo_memory::nixl::NixlAgent;
731
732 fn make_test_agent(name: &str) -> NixlAgent {
733 NixlAgent::new(name).expect("failed to create agent")
734 }
735
736 fn make_test_layout(agent: &NixlAgent) -> PhysicalLayout {
737 let config = LayoutConfig::builder()
738 .num_blocks(2)
739 .num_layers(2)
740 .outer_dim(2)
741 .page_size(4)
742 .inner_dim(8)
743 .dtype_width_bytes(2)
744 .build()
745 .unwrap();
746
747 PhysicalLayout::builder(agent.clone())
748 .with_config(config)
749 .fully_contiguous()
750 .allocate_system()
751 .build()
752 .unwrap()
753 }
754
755 #[test]
756 fn test_manager_creation() {
757 let agent = make_test_agent("test-manager");
758 let manager = LayoutRegistry::new(agent, 42);
759
760 assert_eq!(manager.worker_id(), 42);
761 assert_eq!(manager.local_count(), 0);
762 assert_eq!(manager.remote_count(), 0);
763 }
764
765 #[test]
766 fn test_register_local() {
767 let agent = make_test_agent("test-register");
768 let mut manager = LayoutRegistry::new(agent.clone(), 100);
769
770 let layout = make_test_layout(&agent);
771 let handle = manager.register_local(layout).unwrap();
772
773 assert_eq!(handle.worker_id(), 100);
774 assert_eq!(handle.layout_id(), 0);
775 assert_eq!(manager.local_count(), 1);
776 assert!(manager.is_local(handle));
777 assert!(!manager.is_remote(handle));
778 }
779
780 #[test]
781 fn test_register_multiple_locals() {
782 let agent = make_test_agent("test-multiple");
783 let mut manager = LayoutRegistry::new(agent.clone(), 1);
784
785 let handle1 = manager.register_local(make_test_layout(&agent)).unwrap();
786 let handle2 = manager.register_local(make_test_layout(&agent)).unwrap();
787 let handle3 = manager.register_local(make_test_layout(&agent)).unwrap();
788
789 assert_eq!(handle1.layout_id(), 0);
790 assert_eq!(handle2.layout_id(), 1);
791 assert_eq!(handle3.layout_id(), 2);
792 assert_eq!(manager.local_count(), 3);
793 }
794
795 #[test]
796 #[ignore] fn test_export_import_roundtrip() {
798 let source_agent = make_test_agent("source");
800 let mut source_manager = LayoutRegistry::new(source_agent.clone(), 1);
801
802 let handle1 = source_manager
803 .register_local(make_test_layout(&source_agent))
804 .unwrap();
805 let handle2 = source_manager
806 .register_local(make_test_layout(&source_agent))
807 .unwrap();
808
809 let metadata = source_manager.export_metadata().unwrap();
811 assert!(!metadata.is_empty());
812
813 let dest_agent = make_test_agent("dest");
815 let mut dest_manager = LayoutRegistry::new(dest_agent, 2);
816
817 let imported_handles = dest_manager.import_metadata(metadata).unwrap();
818
819 assert_eq!(imported_handles.len(), 2);
821 assert_eq!(dest_manager.remote_count(), 2);
822 assert!(dest_manager.is_remote(handle1));
823 assert!(dest_manager.is_remote(handle2));
824
825 assert!(dest_manager.get_remote(handle1).is_some());
827 assert!(dest_manager.get_remote(handle2).is_some());
828 assert!(dest_manager.get_layout(handle1).is_some());
829 }
830
831 #[test]
832 #[ignore] fn test_import_duplicate_remote_fails() {
834 let source_agent = make_test_agent("source2");
835 let mut source_manager = LayoutRegistry::new(source_agent.clone(), 10);
836
837 source_manager
838 .register_local(make_test_layout(&source_agent))
839 .unwrap();
840
841 let metadata = source_manager.export_metadata().unwrap();
842
843 let dest_agent = make_test_agent("dest2");
844 let mut dest_manager = LayoutRegistry::new(dest_agent, 20);
845
846 let metadata_clone = SerializedLayout::from_bytes(metadata.as_bytes().to_vec());
848 dest_manager.import_metadata(metadata).unwrap();
849
850 let result = dest_manager.import_metadata(metadata_clone);
852 assert!(result.is_err());
853 assert!(result.unwrap_err().to_string().contains("already loaded"));
854 }
855
856 #[test]
857 fn test_get_layout_handles() {
858 let agent = make_test_agent("test-handles");
859 let mut manager = LayoutRegistry::new(agent.clone(), 5);
860
861 let h1 = manager.register_local(make_test_layout(&agent)).unwrap();
862 let h2 = manager.register_local(make_test_layout(&agent)).unwrap();
863
864 let handles = manager.local_handles();
865 assert_eq!(handles.len(), 2);
866 assert!(handles.contains(&h1));
867 assert!(handles.contains(&h2));
868 }
869}