#[cfg(feature = "collectives")]
mod replicated;
#[cfg(feature = "collectives")]
#[allow(unused_imports)]
pub use replicated::ReplicatedDataWorker;
use std::collections::HashMap;
use std::sync::{Arc, RwLock};
#[cfg(feature = "nccl")]
use cudarc::driver::CudaEvent;
use derive_builder::Builder;
use futures::future::BoxFuture;
use crate::object::ObjectBlockOps;
use kvbm_physical::layout::PhysicalLayout;
use kvbm_physical::{
manager::{SerializedLayout, TransferManager},
transfer::{BounceBuffer, TransferOptions, context::TransferCompleteNotification},
};
use super::*;
#[derive(Builder)]
#[builder(pattern = "owned")]
pub struct PhysicalWorker {
manager: TransferManager,
#[builder(default, setter(strip_option))]
g1_handle: Option<LayoutHandle>,
#[builder(default, setter(strip_option))]
g2_handle: Option<LayoutHandle>,
#[builder(default, setter(strip_option))]
g3_handle: Option<LayoutHandle>,
#[builder(default = "RwLock::new(HashMap::new())")]
remote_handles: RwLock<HashMap<(InstanceId, LogicalLayoutHandle), LayoutHandle>>,
#[builder(default, setter(strip_option))]
rank: Option<usize>,
#[builder(default, setter(strip_option))]
object_client: Option<Arc<dyn ObjectBlockOps>>,
}
impl PhysicalWorker {
pub fn builder() -> PhysicalWorkerBuilder {
PhysicalWorkerBuilder::default()
}
pub fn rank(&self) -> Option<usize> {
self.rank
}
pub fn object_client(&self) -> Option<&Arc<dyn ObjectBlockOps>> {
self.object_client.as_ref()
}
pub fn g1_handle(&self) -> Option<LayoutHandle> {
self.g1_handle
}
pub fn g2_handle(&self) -> Option<LayoutHandle> {
self.g2_handle
}
pub fn g3_handle(&self) -> Option<LayoutHandle> {
self.g3_handle
}
pub fn transfer_manager(&self) -> &TransferManager {
&self.manager
}
pub fn resolve_layout(&self, logical: LogicalLayoutHandle) -> Result<PhysicalLayout> {
use LogicalLayoutHandle::*;
let physical_handle = match logical {
G1 => self.g1_handle(),
G2 => self.g2_handle(),
G3 => self.g3_handle(),
_ => None,
}
.ok_or_else(|| anyhow::anyhow!("No layout registered for {:?}", logical))?;
self.manager
.get_physical_layout(physical_handle)
.ok_or_else(|| {
anyhow::anyhow!(
"Layout handle {:?} not found in TransferManager",
physical_handle
)
})
}
pub fn create_bounce_buffer(
&self,
handle: LayoutHandle,
block_ids: Vec<BlockId>,
) -> Result<BounceBuffer> {
Ok(BounceBuffer::from_handle(handle, block_ids))
}
pub fn export_metadata(&self) -> Result<SerializedLayout> {
self.export_metadata_with_logical_types()
}
fn export_metadata_with_logical_types(&self) -> Result<SerializedLayout> {
let mut descriptors = Vec::new();
if let Some(handle) = self.g1_handle() {
descriptors.push(
self.manager
.build_logical_descriptor(handle, LogicalLayoutHandle::G1)?,
);
}
if let Some(handle) = self.g2_handle() {
descriptors.push(
self.manager
.build_logical_descriptor(handle, LogicalLayoutHandle::G2)?,
);
}
if let Some(handle) = self.g3_handle() {
descriptors.push(
self.manager
.build_logical_descriptor(handle, LogicalLayoutHandle::G3)?,
);
}
let worker_address = self.manager.worker_address();
let nixl_metadata = self.manager.get_nixl_metadata()?;
SerializedLayout::pack(worker_address, nixl_metadata, descriptors)
}
pub fn import_metadata(&self, metadata: SerializedLayout) -> Result<Vec<LayoutHandle>> {
self.manager.import_metadata(metadata)
}
#[cfg(feature = "nccl")]
pub fn execute_local_layerwise_onboard(
&self,
src_block_ids: &[BlockId],
dst_block_ids: &[BlockId],
layer_events: &[Arc<CudaEvent>],
) -> Result<()> {
if src_block_ids.len() != dst_block_ids.len() {
return Err(anyhow::anyhow!(
"Block ID length mismatch: src={}, dst={}",
src_block_ids.len(),
dst_block_ids.len()
));
}
let g2_handle = self
.g2_handle()
.ok_or_else(|| anyhow::anyhow!("G2 layout not registered"))?;
let g1_handle = self
.g1_handle()
.ok_or_else(|| anyhow::anyhow!("G1 layout not registered"))?;
let g2_config = self.manager.get_layout_config(g2_handle)?;
let num_layers = g2_config.num_layers;
if layer_events.len() != num_layers {
return Err(anyhow::anyhow!(
"layer_events length ({}) doesn't match num_layers ({})",
layer_events.len(),
num_layers
));
}
let stream = self.manager.context().acquire_h2d_stream();
tracing::debug!(
num_layers,
num_blocks = src_block_ids.len(),
"Starting layer-wise onboard from G2 to G1"
);
for layer in 0..num_layers {
let options = TransferOptions::builder()
.layer_range(layer..layer + 1)
.cuda_stream(stream.clone())
.build()?;
self.manager.execute_transfer(
g2_handle,
src_block_ids,
g1_handle,
dst_block_ids,
options,
)?;
layer_events[layer].record(stream.as_ref())?;
}
tracing::debug!(num_layers, "Layer-wise onboard complete - events recorded");
Ok(())
}
}
impl WorkerTransfers for PhysicalWorker {
fn execute_local_transfer(
&self,
src: LogicalLayoutHandle,
dst: LogicalLayoutHandle,
src_block_ids: Arc<[BlockId]>,
dst_block_ids: Arc<[BlockId]>,
options: TransferOptions,
) -> Result<TransferCompleteNotification> {
use LogicalLayoutHandle::*;
let src_layout = match &src {
G1 => self.g1_handle(),
G2 => self.g2_handle(),
G3 => self.g3_handle(),
G4 => return Err(anyhow::anyhow!("G4 is not supported for local transfers")),
}
.ok_or_else(|| anyhow::anyhow!("Source layout not registered: {:?}", src))?;
let dst_layout = match &dst {
G1 => self.g1_handle(),
G2 => self.g2_handle(),
G3 => self.g3_handle(),
G4 => return Err(anyhow::anyhow!("G4 is not supported for local transfers")),
}
.ok_or_else(|| anyhow::anyhow!("Destination layout not registered: {:?}", dst))?;
self.manager.execute_transfer(
src_layout,
&src_block_ids,
dst_layout,
&dst_block_ids,
options,
)
}
fn execute_remote_onboard(
&self,
src: RemoteDescriptor,
dst: LogicalLayoutHandle,
dst_block_ids: Arc<[BlockId]>,
options: TransferOptions,
) -> Result<TransferCompleteNotification> {
use LogicalLayoutHandle::*;
let dst_layout = match &dst {
G1 => self.g1_handle(),
G2 => self.g2_handle(),
G3 => self.g3_handle(),
G4 => return Err(anyhow::anyhow!("G4 is not supported for remote transfers")),
}
.ok_or_else(|| anyhow::anyhow!("Destination layout not registered: {:?}", dst))?;
match src {
RemoteDescriptor::Layout { handle, block_ids } => {
let block_ids_arc: Arc<[BlockId]> = block_ids.into();
self.manager.execute_transfer(
handle,
&block_ids_arc,
dst_layout,
&dst_block_ids,
options,
)
}
RemoteDescriptor::Object { keys } => {
let object_client = self
.object_client
.as_ref()
.ok_or_else(|| anyhow::anyhow!("Object client not configured"))?
.clone();
let dst_physical = self.resolve_layout(dst)?;
let block_ids_vec: Vec<BlockId> = dst_block_ids.to_vec();
let ctx = self.manager.context();
let event = ctx.event_system().new_event()?;
let handle = event.handle();
let awaiter = ctx.event_system().awaiter(handle)?;
ctx.tokio().spawn(async move {
let results = object_client
.get_blocks_with_layout(keys.clone(), dst_physical, block_ids_vec)
.await;
let failed: Vec<_> = results.iter().filter(|r| r.is_err()).collect();
if failed.is_empty() {
let _ = event.trigger();
} else {
let error_msg = format!(
"{} of {} blocks failed to download",
failed.len(),
results.len()
);
let _ = event.poison(error_msg);
}
});
Ok(TransferCompleteNotification::from_awaiter(awaiter))
}
}
}
fn execute_remote_offload(
&self,
src: LogicalLayoutHandle,
src_block_ids: Arc<[BlockId]>,
dst: RemoteDescriptor,
_options: TransferOptions,
) -> Result<TransferCompleteNotification> {
match dst {
RemoteDescriptor::Layout { handle, block_ids } => {
let src_layout = match &src {
LogicalLayoutHandle::G1 => self.g1_handle(),
LogicalLayoutHandle::G2 => self.g2_handle(),
LogicalLayoutHandle::G3 => self.g3_handle(),
LogicalLayoutHandle::G4 => {
return Err(anyhow::anyhow!("G4 cannot be used as source for offload"));
}
}
.ok_or_else(|| anyhow::anyhow!("Source layout not registered: {:?}", src))?;
let block_ids_arc: Arc<[BlockId]> = block_ids.into();
self.manager.execute_transfer(
src_layout,
&src_block_ids,
handle,
&block_ids_arc,
_options,
)
}
RemoteDescriptor::Object { keys } => {
let object_client = self
.object_client
.as_ref()
.ok_or_else(|| anyhow::anyhow!("Object client not configured"))?
.clone();
let src_physical = self.resolve_layout(src)?;
let block_ids_vec: Vec<BlockId> = src_block_ids.to_vec();
let ctx = self.manager.context();
let event = ctx.event_system().new_event()?;
let handle = event.handle();
let awaiter = ctx.event_system().awaiter(handle)?;
ctx.tokio().spawn(async move {
let results = object_client
.put_blocks_with_layout(keys.clone(), src_physical, block_ids_vec)
.await;
let failed: Vec<_> = results.iter().filter(|r| r.is_err()).collect();
if failed.is_empty() {
let _ = event.trigger();
} else {
let error_msg = format!(
"{} of {} blocks failed to upload",
failed.len(),
results.len()
);
let _ = event.poison(error_msg);
}
});
Ok(TransferCompleteNotification::from_awaiter(awaiter))
}
}
}
fn connect_remote(
&self,
instance_id: InstanceId,
metadata: Vec<SerializedLayout>,
) -> Result<ConnectRemoteResponse> {
if metadata.len() != 1 {
anyhow::bail!(
"PhysicalWorker expects exactly 1 metadata item, got {}",
metadata.len()
);
}
let meta = metadata.into_iter().next().unwrap();
let unpacked = meta.unpack()?;
{
let mut handles = self.remote_handles.write().unwrap();
for descriptor in &unpacked.layouts {
handles.insert((instance_id, descriptor.logical_type), descriptor.handle);
}
}
let repacked = SerializedLayout::pack(
unpacked.worker_address,
unpacked.nixl_metadata,
unpacked.layouts,
)?;
self.manager.import_metadata(repacked)?;
Ok(ConnectRemoteResponse::ready())
}
fn has_remote_metadata(&self, instance_id: InstanceId) -> bool {
let handles = self.remote_handles.read().unwrap();
handles.keys().any(|(id, _)| *id == instance_id)
}
fn execute_remote_onboard_for_instance(
&self,
instance_id: InstanceId,
remote_logical_type: LogicalLayoutHandle,
src_block_ids: Vec<BlockId>,
dst: LogicalLayoutHandle,
dst_block_ids: Arc<[BlockId]>,
options: TransferOptions,
) -> Result<TransferCompleteNotification> {
let handles = self.remote_handles.read().unwrap();
let remote_handle = handles
.get(&(instance_id, remote_logical_type))
.ok_or_else(|| {
anyhow::anyhow!(
"No remote {:?} handle for instance {}",
remote_logical_type,
instance_id
)
})?;
let descriptor = RemoteDescriptor::Layout {
handle: *remote_handle,
block_ids: src_block_ids,
};
self.execute_remote_onboard(descriptor, dst, dst_block_ids, options)
}
}
impl Worker for PhysicalWorker {
fn g1_handle(&self) -> Option<LayoutHandle> {
self.g1_handle
}
fn g2_handle(&self) -> Option<LayoutHandle> {
self.g2_handle
}
fn g3_handle(&self) -> Option<LayoutHandle> {
self.g3_handle
}
fn export_metadata(&self) -> Result<SerializedLayoutResponse> {
self.export_metadata_with_logical_types()
.map(SerializedLayoutResponse::ready)
}
fn import_metadata(&self, metadata: SerializedLayout) -> Result<ImportMetadataResponse> {
self.manager
.import_metadata(metadata)
.map(ImportMetadataResponse::ready)
}
}
impl ObjectBlockOps for PhysicalWorker {
fn has_blocks(
&self,
keys: Vec<SequenceHash>,
) -> BoxFuture<'static, Vec<(SequenceHash, Option<usize>)>> {
if let Some(client) = self.object_client.as_ref() {
client.has_blocks(keys)
} else {
Box::pin(async move { keys.into_iter().map(|k| (k, None)).collect() })
}
}
fn put_blocks(
&self,
keys: Vec<SequenceHash>,
src_layout: LogicalLayoutHandle,
block_ids: Vec<BlockId>,
) -> BoxFuture<'static, Vec<Result<SequenceHash, SequenceHash>>> {
let physical_layout = match self.resolve_layout(src_layout) {
Ok(layout) => layout,
Err(e) => {
tracing::error!(?src_layout, error = %e, "Failed to resolve layout for put_blocks");
return Box::pin(async move { keys.into_iter().map(Err).collect() });
}
};
if let Some(client) = self.object_client.as_ref() {
client.put_blocks_with_layout(keys, physical_layout, block_ids)
} else {
tracing::warn!("put_blocks called but no object client configured");
Box::pin(async move { keys.into_iter().map(Err).collect() })
}
}
fn get_blocks(
&self,
keys: Vec<SequenceHash>,
dst_layout: LogicalLayoutHandle,
block_ids: Vec<BlockId>,
) -> BoxFuture<'static, Vec<Result<SequenceHash, SequenceHash>>> {
let physical_layout = match self.resolve_layout(dst_layout) {
Ok(layout) => layout,
Err(e) => {
tracing::error!(?dst_layout, error = %e, "Failed to resolve layout for get_blocks");
return Box::pin(async move { keys.into_iter().map(Err).collect() });
}
};
if let Some(client) = self.object_client.as_ref() {
client.get_blocks_with_layout(keys, physical_layout, block_ids)
} else {
tracing::warn!("get_blocks called but no object client configured");
Box::pin(async move { keys.into_iter().map(Err).collect() })
}
}
}