use std::sync::Arc;
use crate::messenger::Messenger;
use anyhow::Result;
use bytes::{Bytes, BytesMut};
use velo_ext::WorkerId;
use crate::rendezvous::RendezvousManager;
use crate::rendezvous::handle::DataHandle;
use crate::rendezvous::protocol::{
AcquireResponse, DataMetadata, RdmaOffer, RvAcquireRequest, RvDetachRequest, RvHandleWire,
RvMetadataRequest, RvPullRequest, RvRefRequest, RvReleaseRequest,
};
use crate::rendezvous::write::RendezvousWrite;
#[cfg(all(target_os = "linux", feature = "ucx"))]
use crate::observability::RdmaPathReason;
#[cfg(all(target_os = "linux", feature = "ucx"))]
use crate::rendezvous::protocol::RvLeaseRenewRequest;
pub struct Consumer;
impl Consumer {
pub async fn metadata(messenger: &Arc<Messenger>, handle: DataHandle) -> Result<DataMetadata> {
let target_worker = handle.worker_id();
let meta: DataMetadata = messenger
.typed_unary_streaming::<DataMetadata>("_rv_metadata")
.payload(&RvMetadataRequest {
handle: RvHandleWire::from_handle(handle),
})?
.worker(target_worker)
.send()
.await?;
Ok(meta)
}
pub async fn get(manager: &RendezvousManager, handle: DataHandle) -> Result<(Bytes, u64)> {
let messenger = manager.messenger();
let target_worker = handle.worker_id();
let response = acquire(manager, handle, rdma_offer(manager, target_worker)).await?;
match response {
AcquireResponse::Ready {
lease_id,
transfer_id,
total_len,
chunk_size,
chunk_count,
} => {
match pull_chunks(
messenger,
target_worker,
transfer_id,
total_len,
chunk_size,
chunk_count,
)
.await
{
Ok(data) => Ok((data.freeze(), lease_id)),
Err(e) => {
if let Err(cleanup_err) =
Consumer::detach(messenger, handle, lease_id).await
{
tracing::warn!(
"Failed to detach lease {lease_id} after pull failure: {cleanup_err}"
);
}
Err(e)
}
}
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
AcquireResponse::Rdma {
lease_id,
descriptor,
lease_timeout_ms,
} => {
let descriptor = with_descriptor_hook(manager, descriptor);
match rdma_pull(manager, handle, lease_id, &descriptor, lease_timeout_ms).await {
Ok(buf) => Ok((Bytes::copy_from_slice(&buf), lease_id)),
Err(reason) => fallback_chunked(manager, handle, lease_id, reason).await,
}
}
#[cfg(not(all(target_os = "linux", feature = "ucx")))]
AcquireResponse::Rdma { lease_id, .. } => {
unsolicited_rdma(manager, handle, lease_id).await
}
}
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
pub async fn get_pinned(
manager: &RendezvousManager,
handle: DataHandle,
) -> Result<(crate::rendezvous::rdma::PinnedBuf, u64)> {
let messenger = manager.messenger();
let target_worker = handle.worker_id();
let response = acquire(manager, handle, rdma_offer(manager, target_worker)).await?;
match response {
AcquireResponse::Rdma {
lease_id,
descriptor,
lease_timeout_ms,
} => {
let descriptor = with_descriptor_hook(manager, descriptor);
match rdma_pull(manager, handle, lease_id, &descriptor, lease_timeout_ms).await {
Ok(buf) => Ok((buf, lease_id)),
Err(reason) => {
let (data, lease_id) =
fallback_chunked(manager, handle, lease_id, reason).await?;
let lease = manager.lease_guard(handle, lease_id);
let buf = copy_into_pool(manager, &data).await?;
Ok((buf, lease.disarm()))
}
}
}
AcquireResponse::Ready {
lease_id,
transfer_id,
total_len,
chunk_size,
chunk_count,
} => {
match pull_chunks(
messenger,
target_worker,
transfer_id,
total_len,
chunk_size,
chunk_count,
)
.await
{
Ok(data) => {
let lease = manager.lease_guard(handle, lease_id);
let buf = copy_into_pool(manager, &data).await?;
Ok((buf, lease.disarm()))
}
Err(e) => {
if let Err(cleanup_err) =
Consumer::detach(messenger, handle, lease_id).await
{
tracing::warn!(
"Failed to detach lease {lease_id} after pull failure: {cleanup_err}"
);
}
Err(e)
}
}
}
}
}
pub async fn get_into(
manager: &RendezvousManager,
handle: DataHandle,
dest: &mut impl RendezvousWrite,
) -> Result<u64> {
let messenger = manager.messenger();
let target_worker = handle.worker_id();
let response = acquire(manager, handle, rdma_offer(manager, target_worker)).await?;
match response {
AcquireResponse::Ready {
lease_id,
transfer_id,
total_len,
chunk_size,
chunk_count,
} => {
match pull_chunks_into(
messenger,
target_worker,
transfer_id,
total_len,
chunk_size,
chunk_count,
dest,
)
.await
{
Ok(()) => Ok(lease_id),
Err(e) => {
if let Err(cleanup_err) =
Consumer::detach(messenger, handle, lease_id).await
{
tracing::warn!(
"Failed to detach lease {lease_id} after pull failure: {cleanup_err}"
);
}
Err(e)
}
}
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
AcquireResponse::Rdma {
lease_id,
descriptor,
lease_timeout_ms,
} => {
let descriptor = with_descriptor_hook(manager, descriptor);
match rdma_pull_into(
manager,
handle,
lease_id,
&descriptor,
lease_timeout_ms,
dest,
)
.await
{
Ok(()) => Ok(lease_id),
Err(reason) => {
let (data, lease_id) =
fallback_chunked(manager, handle, lease_id, reason).await?;
let lease = manager.lease_guard(handle, lease_id);
dest.write_chunk(0, &data)?;
Ok(lease.disarm())
}
}
}
#[cfg(not(all(target_os = "linux", feature = "ucx")))]
AcquireResponse::Rdma { lease_id, .. } => {
let (data, lease_id) = unsolicited_rdma(manager, handle, lease_id).await?;
let lease = manager.lease_guard(handle, lease_id);
dest.write_chunk(0, &data)?;
Ok(lease.disarm())
}
}
}
pub async fn ref_handle(messenger: &Arc<Messenger>, handle: DataHandle) -> Result<()> {
let target_worker = handle.worker_id();
messenger
.unary_streaming("_rv_ref")
.raw_payload(Bytes::from(serde_json::to_vec(&RvRefRequest {
handle: RvHandleWire::from_handle(handle),
})?))
.worker(target_worker)
.send()
.await?;
Ok(())
}
pub async fn detach(
messenger: &Arc<Messenger>,
handle: DataHandle,
lease_id: u64,
) -> Result<()> {
let target_worker = handle.worker_id();
messenger
.am_send_streaming("_rv_detach")?
.raw_payload(Bytes::from(serde_json::to_vec(&RvDetachRequest {
handle: RvHandleWire::from_handle(handle),
lease_id,
})?))
.worker(target_worker)
.send()
.await?;
Ok(())
}
pub async fn release(
messenger: &Arc<Messenger>,
handle: DataHandle,
lease_id: u64,
) -> Result<()> {
let target_worker = handle.worker_id();
messenger
.am_send_streaming("_rv_release")?
.raw_payload(Bytes::from(serde_json::to_vec(&RvReleaseRequest {
handle: RvHandleWire::from_handle(handle),
lease_id,
})?))
.worker(target_worker)
.send()
.await?;
Ok(())
}
}
async fn acquire(
manager: &RendezvousManager,
handle: DataHandle,
offer: Option<RdmaOffer>,
) -> Result<AcquireResponse> {
let response: AcquireResponse = manager
.messenger()
.typed_unary_streaming::<AcquireResponse>("_rv_acquire")
.payload(&RvAcquireRequest {
handle: RvHandleWire::from_handle(handle),
rdma: offer,
})?
.worker(handle.worker_id())
.send()
.await?;
Ok(response)
}
#[cfg(not(all(target_os = "linux", feature = "ucx")))]
fn rdma_offer(_manager: &RendezvousManager, _target: WorkerId) -> Option<RdmaOffer> {
None
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
fn rdma_offer(manager: &RendezvousManager, target: WorkerId) -> Option<RdmaOffer> {
let store = manager.data_store();
let Some(ctx) = store.rdma() else {
store.record_path(RdmaPathReason::NotConfigured);
return None;
};
if !ctx.config.enabled {
store.record_path(RdmaPathReason::KillSwitch);
return None;
}
let key = ctx.backend.key();
let backend = manager.messenger().backend();
let Ok(instance) = backend.try_translate_worker_id(target) else {
store.record_path(RdmaPathReason::NoOffer);
return None;
};
let serves = backend
.primary_transport_key(instance)
.is_some_and(|k| k.as_str() == key)
|| backend
.alternative_transport_keys(instance)
.is_some_and(|keys| keys.iter().any(|k| k.as_str() == key));
if !serves {
store.record_path(RdmaPathReason::NoOffer);
return None;
}
Some(RdmaOffer {
backends: vec![key.to_string()],
})
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
fn with_descriptor_hook(manager: &RendezvousManager, descriptor: Vec<u8>) -> Vec<u8> {
#[cfg(feature = "test-helpers")]
if let Some(hook) = manager.take_descriptor_hook() {
return hook.corrupt(descriptor);
}
let _ = manager;
descriptor
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
#[cfg_attr(not(feature = "test-helpers"), allow(dead_code))]
enum TransferHook {
Fail,
Delay(std::time::Duration),
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
fn take_transfer_hook(manager: &RendezvousManager) -> Option<TransferHook> {
#[cfg(feature = "test-helpers")]
{
use crate::rendezvous::RdmaTestHook;
match manager.take_get_hook() {
Some(RdmaTestHook::FailGet) => Some(TransferHook::Fail),
Some(RdmaTestHook::SlowGet(delay)) => Some(TransferHook::Delay(delay)),
_ => None,
}
}
#[cfg(not(feature = "test-helpers"))]
{
let _ = manager;
None
}
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
struct RdmaFallback {
reason: RdmaPathReason,
detail: String,
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
impl RdmaFallback {
fn new(reason: RdmaPathReason, detail: impl Into<String>) -> Self {
Self {
reason,
detail: detail.into(),
}
}
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
async fn rdma_pull(
manager: &RendezvousManager,
handle: DataHandle,
lease_id: u64,
descriptor: &[u8],
lease_timeout_ms: u64,
) -> Result<crate::rendezvous::rdma::PinnedBuf, RdmaFallback> {
use crate::rendezvous::rdma::RdmaError;
let (ctx, desc) = decode_for(manager, descriptor)?;
let len = usize::try_from(desc.len).map_err(|_| {
RdmaFallback::new(
RdmaPathReason::DecodeError,
format!(
"descriptor length {} does not fit this address space",
desc.len
),
)
})?;
let buf = ctx.registry.alloc_pinned(len).await.map_err(|e| {
let reason = match e {
RdmaError::BudgetExceeded { .. } => RdmaPathReason::Budget,
_ => RdmaPathReason::PoolExhausted,
};
RdmaFallback::new(reason, e.to_string())
})?;
let dest = crate::rendezvous::write::RdmaDestination::held(
buf.backend_region_id(),
buf.arena_offset(),
buf.len() as u64,
buf.hold(),
);
run_get(manager, handle, lease_id, &desc, dest, lease_timeout_ms).await?;
Ok(buf)
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
async fn rdma_pull_into(
manager: &RendezvousManager,
handle: DataHandle,
lease_id: u64,
descriptor: &[u8],
lease_timeout_ms: u64,
dest: &mut impl RendezvousWrite,
) -> Result<(), RdmaFallback> {
let (_ctx, desc) = decode_for(manager, descriptor)?;
let len = usize::try_from(desc.len).map_err(|_| {
RdmaFallback::new(
RdmaPathReason::DecodeError,
format!(
"descriptor length {} does not fit this address space",
desc.len
),
)
})?;
if let Some(target) = dest.rdma_destination() {
if target.capacity() < desc.len {
return Err(RdmaFallback::new(
RdmaPathReason::DecodeError,
format!(
"destination holds {} bytes, descriptor names {}",
target.capacity(),
desc.len
),
));
}
return run_get(manager, handle, lease_id, &desc, target, lease_timeout_ms).await;
}
let _ = len;
let buf = rdma_pull(manager, handle, lease_id, descriptor, lease_timeout_ms).await?;
dest.write_chunk(0, &buf)
.map_err(|e| RdmaFallback::new(RdmaPathReason::GetFailed, e.to_string()))
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
fn decode_for<'a>(
manager: &'a RendezvousManager,
descriptor: &[u8],
) -> Result<
(
&'a crate::rendezvous::RdmaContext,
crate::rendezvous::descriptor::RdmaDescriptor,
),
RdmaFallback,
> {
let ctx = manager.data_store().rdma().ok_or_else(|| {
RdmaFallback::new(
RdmaPathReason::PoolExhausted,
"no rdma registry on this instance",
)
})?;
let desc = crate::rendezvous::descriptor::RdmaDescriptor::decode(descriptor)
.map_err(|e| RdmaFallback::new(RdmaPathReason::DecodeError, e.to_string()))?;
if desc.backend != ctx.backend {
return Err(RdmaFallback::new(
RdmaPathReason::DecodeError,
format!(
"descriptor names backend {:?}, this instance serves {:?}",
desc.backend, ctx.backend
),
));
}
Ok((ctx, desc))
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
async fn run_get(
manager: &RendezvousManager,
handle: DataHandle,
lease_id: u64,
desc: &crate::rendezvous::descriptor::RdmaDescriptor,
mut dest: crate::rendezvous::write::RdmaDestination<'_>,
lease_timeout_ms: u64,
) -> Result<(), RdmaFallback> {
let store = manager.data_store();
let ctx = store.rdma().ok_or_else(|| {
RdmaFallback::new(
RdmaPathReason::PoolExhausted,
"no rdma registry on this instance",
)
})?;
let peer = manager
.messenger()
.backend()
.try_translate_worker_id(handle.worker_id())
.map_err(|e| {
RdmaFallback::new(
RdmaPathReason::GetFailed,
format!("owner is not registered: {e}"),
)
})?;
let req = crate::rendezvous::rdma::BackendGet {
peer,
remote_addr: desc.addr,
packed_key: desc.packed_key.clone(),
local_region_id: dest.region_id(),
local_offset: dest.offset(),
len: desc.len,
};
let delay = match take_transfer_hook(manager) {
Some(TransferHook::Fail) => {
return Err(RdmaFallback::new(
RdmaPathReason::GetFailed,
"injected transfer failure",
));
}
Some(TransferHook::Delay(delay)) => Some(delay),
None => None,
};
let hold = dest.take_hold();
let registry = Arc::clone(&ctx.registry);
let transfer = ctx.registry.runtime().spawn(async move {
if let Some(delay) = delay {
tokio::time::sleep(delay).await;
}
let outcome = registry.get(req).await;
drop(hold);
outcome
});
let started = std::time::Instant::now();
let outcome = with_lease_renewal(manager, handle, lease_id, lease_timeout_ms, async move {
match transfer.await {
Ok(outcome) => outcome,
Err(join) => Err(crate::rendezvous::rdma::RdmaError::Backend(format!(
"rdma transfer task: {join}"
))),
}
})
.await;
match outcome {
Ok(()) => {
if let Some(m) = store.metrics() {
m.record_rendezvous_rdma_get(started.elapsed());
}
store.record_path(RdmaPathReason::Ok);
Ok(())
}
Err(e) => Err(RdmaFallback::new(RdmaPathReason::GetFailed, e.to_string())),
}
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
async fn with_lease_renewal<F, T>(
manager: &RendezvousManager,
handle: DataHandle,
lease_id: u64,
lease_timeout_ms: u64,
transfer: F,
) -> T
where
F: std::future::Future<Output = T>,
{
if lease_timeout_ms == 0 {
return transfer.await;
}
let period = std::time::Duration::from_millis(lease_timeout_ms / 2)
.max(std::time::Duration::from_millis(5));
let mut transfer = std::pin::pin!(transfer);
loop {
tokio::select! {
out = &mut transfer => return out,
_ = tokio::time::sleep(period) => {
send_lease_renewal(manager, handle, lease_id).await;
}
}
}
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
async fn send_lease_renewal(manager: &RendezvousManager, handle: DataHandle, lease_id: u64) {
let payload = match serde_json::to_vec(&RvLeaseRenewRequest {
handle: RvHandleWire::from_handle(handle),
lease_id,
}) {
Ok(payload) => payload,
Err(e) => {
tracing::debug!(error = %e, "rendezvous: could not encode a lease renewal");
return;
}
};
let sent = async {
manager
.messenger()
.am_send_streaming("_rv_lease_renew")?
.raw_payload(Bytes::from(payload))
.worker(handle.worker_id())
.send()
.await
}
.await;
if let Err(e) = sent {
tracing::debug!(
lease = lease_id,
error = %e,
"rendezvous: lease renewal not delivered; the standing deadline still applies"
);
}
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
async fn fallback_chunked(
manager: &RendezvousManager,
handle: DataHandle,
lease_id: u64,
fallback: RdmaFallback,
) -> Result<(Bytes, u64)> {
let store = manager.data_store();
store.record_path(fallback.reason);
tracing::warn!(
%handle,
lease = lease_id,
reason = fallback.reason.as_str(),
detail = %fallback.detail,
"rendezvous: falling back to the chunked path for this transfer"
);
if let Err(e) = Consumer::detach(manager.messenger(), handle, lease_id).await {
tracing::warn!(%handle, error = %e, "rendezvous: could not detach before falling back");
}
chunked_only(manager, handle).await
}
#[cfg(not(all(target_os = "linux", feature = "ucx")))]
async fn unsolicited_rdma(
manager: &RendezvousManager,
handle: DataHandle,
lease_id: u64,
) -> Result<(Bytes, u64)> {
tracing::warn!(
%handle,
"rendezvous: owner answered with an RDMA descriptor for an acquire that offered \
nothing; falling back to the chunked path"
);
if let Err(e) = Consumer::detach(manager.messenger(), handle, lease_id).await {
tracing::warn!(%handle, error = %e, "rendezvous: could not detach before falling back");
}
chunked_only(manager, handle).await
}
async fn chunked_only(manager: &RendezvousManager, handle: DataHandle) -> Result<(Bytes, u64)> {
let messenger = manager.messenger();
let target_worker = handle.worker_id();
match acquire(manager, handle, None).await? {
AcquireResponse::Ready {
lease_id,
transfer_id,
total_len,
chunk_size,
chunk_count,
} => {
match pull_chunks(
messenger,
target_worker,
transfer_id,
total_len,
chunk_size,
chunk_count,
)
.await
{
Ok(data) => Ok((data.freeze(), lease_id)),
Err(e) => {
if let Err(cleanup_err) = Consumer::detach(messenger, handle, lease_id).await {
tracing::warn!(
"Failed to detach lease {lease_id} after pull failure: {cleanup_err}"
);
}
Err(e)
}
}
}
AcquireResponse::Rdma { lease_id, .. } => {
if let Err(e) = Consumer::detach(messenger, handle, lease_id).await {
tracing::warn!(
%handle,
error = %e,
"rendezvous: could not detach an unusable lease"
);
}
anyhow::bail!(
"rendezvous owner answered {handle} with an RDMA descriptor for an acquire that \
carried no offer; refusing to retry"
)
}
}
}
#[cfg(all(target_os = "linux", feature = "ucx"))]
async fn copy_into_pool(
manager: &RendezvousManager,
data: &[u8],
) -> Result<crate::rendezvous::rdma::PinnedBuf> {
let ctx = manager
.data_store()
.rdma()
.ok_or_else(|| anyhow::anyhow!("get_pinned needs an rdma registry on this instance"))?;
let mut buf = ctx.registry.alloc_pinned(data.len()).await?;
buf.copy_from_slice(data);
Ok(buf)
}
async fn pull_chunks(
messenger: &Arc<Messenger>,
target_worker: WorkerId,
transfer_id: u64,
total_len: u64,
chunk_size: u32,
chunk_count: u32,
) -> Result<BytesMut> {
let mut buf = BytesMut::with_capacity(total_len as usize);
buf.resize(total_len as usize, 0);
for chunk_index in 0..chunk_count {
let req = RvPullRequest {
transfer_id,
chunk_index,
};
let payload = serde_json::to_vec(&req)?;
let chunk_bytes: Bytes = messenger
.unary_streaming("_rv_pull")
.raw_payload(Bytes::from(payload))
.worker(target_worker)
.send()
.await?;
let offset = chunk_index as usize * chunk_size as usize;
let end = (offset + chunk_bytes.len()).min(total_len as usize);
buf[offset..end].copy_from_slice(&chunk_bytes[..end - offset]);
}
Ok(buf)
}
async fn pull_chunks_into(
messenger: &Arc<Messenger>,
target_worker: WorkerId,
transfer_id: u64,
_total_len: u64,
chunk_size: u32,
chunk_count: u32,
dest: &mut impl RendezvousWrite,
) -> Result<()> {
for chunk_index in 0..chunk_count {
let req = RvPullRequest {
transfer_id,
chunk_index,
};
let payload = serde_json::to_vec(&req)?;
let chunk_bytes: Bytes = messenger
.unary_streaming("_rv_pull")
.raw_payload(Bytes::from(payload))
.worker(target_worker)
.send()
.await?;
let offset = chunk_index as usize * chunk_size as usize;
dest.write_chunk(offset, &chunk_bytes)?;
}
Ok(())
}