use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use bytes::Bytes;
use super::{UcxConfig, UcxTransport, UcxTransportBuilder};
use crate::transports::transport::{
DataStreams, HealthCheckError, SendOutcome, Transport, TransportErrorHandler, make_channels,
};
use crate::transports::ucx::rma::{
MAX_PACKED_RKEY, MappedRegion, RdmaEndpoint, RmaError, RmaGetRequest, SYS_DEV_UNKNOWN,
preparse_packed_rkey,
};
use crate::transports::ucx::worker::Cmd;
use velo_ext::{InstanceId, MessageType, PeerInfo};
struct CountingErrors {
count: AtomicUsize,
notify: tokio::sync::Notify,
}
impl CountingErrors {
fn new() -> Arc<Self> {
Arc::new(Self {
count: AtomicUsize::new(0),
notify: tokio::sync::Notify::new(),
})
}
fn count(&self) -> usize {
self.count.load(Ordering::SeqCst)
}
async fn wait_for_error(&self, timeout: Duration) -> bool {
let notified = self.notify.notified();
tokio::pin!(notified);
notified.as_mut().enable();
if self.count() > 0 {
return true;
}
tokio::time::timeout(timeout, notified).await.is_ok()
}
}
impl TransportErrorHandler for CountingErrors {
fn on_error(&self, _header: Bytes, _payload: Bytes, _error: String) {
self.count.fetch_add(1, Ordering::SeqCst);
self.notify.notify_waiters();
}
}
struct Node {
transport: Arc<UcxTransport>,
streams: DataStreams,
instance_id: InstanceId,
}
async fn start_node() -> Node {
start_node_with(|b| b).await
}
async fn start_node_with(
configure: impl FnOnce(UcxTransportBuilder) -> UcxTransportBuilder,
) -> Node {
start_transport(Arc::new(
configure(UcxTransportBuilder::new().tls("tcp"))
.build()
.expect("build ucx transport"),
))
.await
}
async fn start_node_with_config(config: UcxConfig) -> Node {
start_transport(Arc::new(UcxTransport::new(
velo_ext::TransportKey::from("ucx"),
config,
)))
.await
}
async fn start_transport(transport: Arc<UcxTransport>) -> Node {
let instance_id = InstanceId::new_v4();
let (adapter, streams) = make_channels();
tokio::time::timeout(
T,
transport.start(instance_id, adapter, tokio::runtime::Handle::current()),
)
.await
.expect("ucx transport startup must not hang")
.expect("start ucx transport");
Node {
transport,
streams,
instance_id,
}
}
fn cross_register(a: &Node, b: &Node) {
a.transport
.register(PeerInfo::new(b.instance_id, b.transport.address()))
.expect("register b in a");
b.transport
.register(PeerInfo::new(a.instance_id, a.transport.address()))
.expect("register a in b");
}
async fn recv(rx: &flume::Receiver<(Bytes, Bytes)>, timeout: Duration) -> Option<(Bytes, Bytes)> {
tokio::time::timeout(timeout, rx.recv_async())
.await
.ok()?
.ok()
}
async fn recv_message(
rx: &flume::Receiver<crate::transports::transport::InboundMessage>,
timeout: Duration,
) -> Option<(Bytes, Bytes)> {
let msg = tokio::time::timeout(timeout, rx.recv_async())
.await
.ok()?
.ok()?;
Some((msg.header, msg.payload))
}
const T: Duration = Duration::from_secs(10);
#[tokio::test(flavor = "multi_thread")]
async fn message_round_trip_and_stream_routing() {
let a = start_node().await;
let b = start_node().await;
cross_register(&a, &b);
let errs = CountingErrors::new();
let out = a.transport.send_message(
b.instance_id,
Bytes::from_static(b"hdr"),
Bytes::from_static(b"payload"),
MessageType::Message,
errs.clone(),
);
assert!(matches!(
out,
SendOutcome::Admitted | SendOutcome::Pending(_)
));
let (h, p) = recv_message(&b.streams.message_stream, T)
.await
.expect("message arrives");
assert_eq!(&h[..], b"hdr");
assert_eq!(&p[..], b"payload");
b.transport.send_message(
a.instance_id,
Bytes::from_static(b"resp-h"),
Bytes::from_static(b"resp-p"),
MessageType::Response,
errs.clone(),
);
let (h, _) = recv(&a.streams.response_stream, T)
.await
.expect("response arrives");
assert_eq!(&h[..], b"resp-h");
a.transport.send_message(
b.instance_id,
Bytes::from_static(b"ev-h"),
Bytes::new(),
MessageType::Event,
errs.clone(),
);
let (h, p) = recv(&b.streams.event_stream, T)
.await
.expect("event arrives");
assert_eq!(&h[..], b"ev-h");
assert!(p.is_empty());
assert_eq!(errs.count(), 0, "no send errors expected");
a.transport.shutdown();
b.transport.shutdown();
}
#[tokio::test(flavor = "multi_thread")]
async fn many_messages_preserve_order() {
let a = start_node().await;
let b = start_node().await;
cross_register(&a, &b);
let errs = CountingErrors::new();
const N: u32 = 200;
for i in 0..N {
a.transport.send_message(
b.instance_id,
Bytes::from(i.to_le_bytes().to_vec()),
Bytes::from(vec![0u8; 1024]),
MessageType::Message,
errs.clone(),
);
}
for i in 0..N {
let (h, p) = recv_message(&b.streams.message_stream, T)
.await
.expect("ordered message");
assert_eq!(
u32::from_le_bytes(h[..4].try_into().unwrap()),
i,
"order preserved"
);
assert_eq!(p.len(), 1024);
}
assert_eq!(errs.count(), 0);
a.transport.shutdown();
b.transport.shutdown();
}
#[tokio::test(flavor = "multi_thread")]
async fn unregistered_peer_reports_through_on_error() {
let a = start_node().await;
let errs = CountingErrors::new();
let out = a.transport.send_message(
InstanceId::new_v4(),
Bytes::from_static(b"h"),
Bytes::from_static(b"p"),
MessageType::Message,
errs.clone(),
);
assert!(matches!(out, SendOutcome::Admitted));
assert!(
errs.wait_for_error(T).await,
"pre-wire failure must reach on_error"
);
a.transport.shutdown();
}
#[tokio::test(flavor = "multi_thread")]
async fn oversized_frame_fails_pre_wire() {
let a = start_node().await;
let b = start_node().await;
cross_register(&a, &b);
let errs = CountingErrors::new();
let limit = a
.transport
.max_message_size(b.instance_id)
.expect("limit known");
let out = a.transport.send_message(
b.instance_id,
Bytes::from_static(b"h"),
Bytes::from(vec![0u8; limit + 1]),
MessageType::Message,
errs.clone(),
);
assert!(matches!(out, SendOutcome::Admitted));
assert!(
errs.wait_for_error(T).await,
"oversized frame must reach on_error"
);
a.transport.shutdown();
b.transport.shutdown();
}
#[tokio::test(flavor = "multi_thread")]
async fn draining_receiver_echoes_shutting_down() {
let a = start_node().await;
let b = start_node().await;
cross_register(&a, &b);
let errs = CountingErrors::new();
a.transport.send_message(
b.instance_id,
Bytes::from_static(b"warm"),
Bytes::new(),
MessageType::Message,
errs.clone(),
);
recv_message(&b.streams.message_stream, T)
.await
.expect("warmup arrives");
b.streams.shutdown_state.begin_drain();
a.transport.send_message(
b.instance_id,
Bytes::from_static(b"corr-id"),
Bytes::from_static(b"ignored"),
MessageType::Message,
errs.clone(),
);
assert!(
recv_message(&b.streams.message_stream, Duration::from_millis(500))
.await
.is_none(),
"draining receiver must not deliver new messages"
);
let (h, _) = recv(&a.streams.shutdown_stream, T)
.await
.expect("ShuttingDown echo");
assert_eq!(&h[..], b"corr-id");
a.transport.shutdown();
b.transport.shutdown();
}
#[tokio::test(flavor = "multi_thread")]
async fn health_check_semantics() {
let a = start_node().await;
let b = start_node().await;
cross_register(&a, &b);
let errs = CountingErrors::new();
assert!(matches!(
a.transport.check_health(InstanceId::new_v4(), T).await,
Err(HealthCheckError::PeerNotRegistered)
));
assert!(matches!(
a.transport.check_health(b.instance_id, T).await,
Err(HealthCheckError::NeverConnected)
));
a.transport.send_message(
b.instance_id,
Bytes::from_static(b"h"),
Bytes::new(),
MessageType::Message,
errs.clone(),
);
recv_message(&b.streams.message_stream, T)
.await
.expect("message arrives");
assert!(a.transport.check_health(b.instance_id, T).await.is_ok());
a.transport.shutdown();
b.transport.shutdown();
}
#[tokio::test(flavor = "multi_thread")]
async fn shutdown_fails_queued_sends() {
let a = start_node().await;
let b = start_node().await;
cross_register(&a, &b);
let errs = CountingErrors::new();
a.transport.shutdown();
let out = a.transport.send_message(
b.instance_id,
Bytes::from_static(b"h"),
Bytes::from_static(b"p"),
MessageType::Message,
errs.clone(),
);
match out {
SendOutcome::Admitted => {
assert!(
errs.wait_for_error(T).await,
"post-shutdown send must surface an error"
);
}
SendOutcome::Pending(admission) => {
let resolved = tokio::time::timeout(T, admission)
.await
.expect("admission must resolve, not hang");
assert!(resolved.is_err());
}
}
b.transport.shutdown();
}
struct PageBuf {
ptr: *mut u8,
len: usize,
}
impl PageBuf {
const ALIGN: usize = 4096;
fn new(len: usize) -> Self {
let layout = std::alloc::Layout::from_size_align(len, Self::ALIGN).expect("valid layout");
let ptr = unsafe { std::alloc::alloc_zeroed(layout) };
assert!(!ptr.is_null(), "allocating {len} bytes failed");
Self { ptr, len }
}
fn addr(&self) -> usize {
self.ptr as usize
}
fn as_slice(&self) -> &[u8] {
unsafe { std::slice::from_raw_parts(self.ptr, self.len) }
}
fn fill_pattern(&mut self) {
let slice = unsafe { std::slice::from_raw_parts_mut(self.ptr, self.len) };
for (i, byte) in slice.iter_mut().enumerate() {
*byte = (i % 251) as u8;
}
}
}
impl Drop for PageBuf {
fn drop(&mut self) {
let layout =
std::alloc::Layout::from_size_align(self.len, Self::ALIGN).expect("valid layout");
unsafe { std::alloc::dealloc(self.ptr, layout) };
}
}
struct RmaPair {
owner: Node,
puller: Node,
owner_rma: RdmaEndpoint,
puller_rma: RdmaEndpoint,
}
async fn start_rma_pair() -> RmaPair {
let owner = start_node().await;
let puller = start_node().await;
cross_register(&owner, &puller);
let owner_rma = owner.transport.rdma_endpoint();
let puller_rma = puller.transport.rdma_endpoint();
RmaPair {
owner,
puller,
owner_rma,
puller_rma,
}
}
fn get_request(
pair: &RmaPair,
src: &PageBuf,
remote: &MappedRegion,
local: &MappedRegion,
) -> RmaGetRequest {
RmaGetRequest {
peer: pair.owner.instance_id,
remote_addr: src.addr() as u64,
packed_rkey: remote.packed_rkey.clone(),
local_region: local.region_id,
local_offset: 0,
len: src.len as u64,
}
}
fn assert_rma_balanced(node: &Node) {
assert_eq!(
node.transport.shared.live_regions.load(Ordering::SeqCst),
0,
"a registered region outlived the transport"
);
assert_eq!(
node.transport.shared.live_rkeys.load(Ordering::SeqCst),
0,
"an unpacked rkey outlived the transport"
);
assert_eq!(
node.transport.shared.eps_open.load(Ordering::SeqCst),
0,
"an endpoint outlived the transport, or its close went uncounted"
);
}
fn assert_pair_balanced(pair: &RmaPair) {
assert_rma_balanced(&pair.owner);
assert_rma_balanced(&pair.puller);
}
async fn wait_until(budget: Duration, mut cond: impl FnMut() -> bool) -> bool {
let deadline = tokio::time::Instant::now() + budget;
loop {
if cond() {
return true;
}
if tokio::time::Instant::now() >= deadline {
return false;
}
tokio::time::sleep(Duration::from_micros(200)).await;
}
}
async fn ring_get(transport: &UcxTransport, req: RmaGetRequest) -> Result<(), RmaError> {
let (tx, rx) = tokio::sync::oneshot::channel();
transport
.shared
.ring_tx
.send_async(Cmd::RmaGet { req, reply: tx })
.await
.expect("ring accepts the command");
transport.shared.doorbell.ring();
tokio::time::timeout(T, rx)
.await
.expect("worker must answer")
.expect("worker must not drop the reply")
}
async fn ring_unmap_enqueue(
transport: &UcxTransport,
region_id: u64,
) -> tokio::sync::oneshot::Receiver<Result<(), RmaError>> {
let (tx, rx) = tokio::sync::oneshot::channel();
transport
.shared
.ring_tx
.send_async(Cmd::UnmapRegion {
region_id,
reply: tx,
})
.await
.expect("ring accepts the command");
transport.shared.doorbell.ring();
rx
}
#[tokio::test(flavor = "multi_thread")]
async fn map_get_roundtrip() {
const LEN: usize = 256 * 1024;
let mut src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
src.fill_pattern();
let pair = start_rma_pair().await;
let remote = pair
.owner_rma
.map_region(src.addr(), LEN)
.await
.expect("map source region");
let local = pair
.puller_rma
.map_region(dst.addr(), LEN)
.await
.expect("map destination region");
assert!(remote.effective_addr <= src.addr() as u64);
assert!(remote.effective_addr + remote.effective_len >= (src.addr() + LEN) as u64);
tokio::time::timeout(
T,
pair.puller_rma
.get(get_request(&pair, &src, &remote, &local)),
)
.await
.expect("get must not hang")
.expect("get succeeds");
assert_eq!(dst.as_slice(), src.as_slice(), "GET must copy the pattern");
pair.puller_rma
.unmap_region(local.region_id)
.await
.expect("unmap destination");
pair.owner_rma
.unmap_region(remote.region_id)
.await
.expect("unmap source");
pair.owner.transport.shutdown();
pair.puller.transport.shutdown();
assert_pair_balanced(&pair);
}
#[tokio::test(flavor = "multi_thread")]
async fn get_zero_length() {
const LEN: usize = 4096;
let mut src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
src.fill_pattern();
let pair = start_rma_pair().await;
let remote = pair.owner_rma.map_region(src.addr(), LEN).await.unwrap();
let local = pair.puller_rma.map_region(dst.addr(), LEN).await.unwrap();
let mut req = get_request(&pair, &src, &remote, &local);
req.len = 0;
tokio::time::timeout(T, pair.puller_rma.get(req))
.await
.expect("zero-length get must not hang")
.expect("zero-length get succeeds");
assert_eq!(dst.as_slice(), &[0u8; LEN][..], "no bytes may be written");
pair.owner.transport.shutdown();
pair.puller.transport.shutdown();
assert_pair_balanced(&pair);
}
#[tokio::test(flavor = "multi_thread")]
async fn get_out_of_range() {
const LEN: usize = 4096;
let src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
let pair = start_rma_pair().await;
let remote = pair.owner_rma.map_region(src.addr(), LEN).await.unwrap();
let local = pair.puller_rma.map_region(dst.addr(), LEN).await.unwrap();
let mut req = get_request(&pair, &src, &remote, &local);
req.local_offset = 1;
assert!(matches!(
pair.puller_rma.get(req).await,
Err(RmaError::OutOfRange)
));
let mut req = get_request(&pair, &src, &remote, &local);
req.local_offset = LEN as u64 * 2;
req.len = 1;
assert!(matches!(
pair.puller_rma.get(req).await,
Err(RmaError::OutOfRange)
));
let mut req = get_request(&pair, &src, &remote, &local);
req.local_offset = u64::MAX;
req.len = 2;
assert!(matches!(
pair.puller_rma.get(req).await,
Err(RmaError::OutOfRange)
));
let mut req = get_request(&pair, &src, &remote, &local);
req.local_region = local.region_id + 1_000;
assert!(matches!(
pair.puller_rma.get(req).await,
Err(RmaError::RegionNotFound)
));
pair.owner.transport.shutdown();
pair.puller.transport.shutdown();
assert_pair_balanced(&pair);
}
#[tokio::test(flavor = "multi_thread")]
async fn unmap_waits_for_inflight() {
const CHUNK: usize = 8 * 1024 * 1024;
const CHUNKS: usize = 8;
const LEN: usize = CHUNK * CHUNKS;
let mut src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
src.fill_pattern();
let pair = start_rma_pair().await;
let remote = pair.owner_rma.map_region(src.addr(), LEN).await.unwrap();
let local = pair.puller_rma.map_region(dst.addr(), LEN).await.unwrap();
let mut gets = Vec::with_capacity(CHUNKS);
for i in 0..CHUNKS {
let endpoint = pair.puller_rma.clone();
let req = RmaGetRequest {
peer: pair.owner.instance_id,
remote_addr: (src.addr() + i * CHUNK) as u64,
packed_rkey: remote.packed_rkey.clone(),
local_region: local.region_id,
local_offset: (i * CHUNK) as u64,
len: CHUNK as u64,
};
gets.push(tokio::spawn(async move { endpoint.get(req).await }));
}
assert!(
wait_until(T, || {
pair.puller
.transport
.shared
.inflight_ops
.load(Ordering::SeqCst)
>= CHUNKS
})
.await,
"all {CHUNKS} GETs must be in flight before the unmap is issued"
);
let mut unmap = Box::pin(pair.puller_rma.unmap_region(local.region_id));
match tokio::time::timeout(Duration::from_millis(5), &mut unmap).await {
Ok(early) => early.expect("unmap resolves"),
Err(_) => tokio::time::timeout(T, &mut unmap)
.await
.expect("unmap must resolve once the GETs complete")
.expect("unmap succeeds"),
}
for (i, task) in gets.into_iter().enumerate() {
task.await
.expect("get task must not panic")
.unwrap_or_else(|e| panic!("get {i} failed: {e}"));
}
assert_eq!(
dst.as_slice(),
src.as_slice(),
"every GET must have completed before the region was unmapped"
);
let mut req = get_request(&pair, &src, &remote, &local);
req.len = 1;
assert!(matches!(
pair.puller_rma.get(req).await,
Err(RmaError::RegionNotFound)
));
pair.owner.transport.shutdown();
pair.puller.transport.shutdown();
assert_pair_balanced(&pair);
}
#[tokio::test(flavor = "multi_thread")]
async fn get_unknown_peer() {
const LEN: usize = 4096;
let src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
let pair = start_rma_pair().await;
let remote = pair.owner_rma.map_region(src.addr(), LEN).await.unwrap();
let local = pair.puller_rma.map_region(dst.addr(), LEN).await.unwrap();
let mut req = get_request(&pair, &src, &remote, &local);
req.peer = InstanceId::new_v4();
assert!(matches!(
tokio::time::timeout(T, pair.puller_rma.get(req))
.await
.expect("must not hang"),
Err(RmaError::PeerNotRegistered(_))
));
pair.owner.transport.shutdown();
pair.puller.transport.shutdown();
assert_pair_balanced(&pair);
}
#[tokio::test(flavor = "multi_thread")]
async fn rkey_pack_canary() {
const LEN: usize = 64 * 1024;
let first = PageBuf::new(LEN);
let second = PageBuf::new(LEN);
let node = start_node().await;
let rma = node.transport.rdma_endpoint();
let a = rma.map_region(first.addr(), LEN).await.expect("map first");
println!(
"ucp_rkey_pack under UCX_TLS=tcp: {} bytes",
a.packed_rkey.len()
);
assert!(
a.packed_rkey.len() >= 9,
"packed rkey is implausibly small ({} bytes)",
a.packed_rkey.len()
);
let b = rma
.map_region(second.addr(), LEN)
.await
.expect("map second");
assert!(!b.packed_rkey.is_empty(), "second pack must also succeed");
assert_ne!(a.region_id, b.region_id);
rma.unmap_region(a.region_id).await.expect("unmap first");
rma.unmap_region(b.region_id).await.expect("unmap second");
rma.unmap_region(a.region_id)
.await
.expect("repeat unmap is a no-op, not a failure");
rma.unmap_region(u64::MAX)
.await
.expect("unmapping an id that never existed is also a no-op");
node.transport.shutdown();
assert_rma_balanced(&node);
}
#[tokio::test(flavor = "multi_thread")]
async fn get_cancelled_by_endpoint_replacement() {
const LEN: usize = 64 * 1024 * 1024;
let mut src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
src.fill_pattern();
let pair = start_rma_pair().await;
let decoy = start_node().await;
let remote = pair.owner_rma.map_region(src.addr(), LEN).await.unwrap();
let local = pair.puller_rma.map_region(dst.addr(), LEN).await.unwrap();
let endpoint = pair.puller_rma.clone();
let req = get_request(&pair, &src, &remote, &local);
let get = tokio::spawn(async move { endpoint.get(req).await });
tokio::time::sleep(Duration::from_millis(2)).await;
pair.puller
.transport
.register(PeerInfo::new(
pair.owner.instance_id,
decoy.transport.address(),
))
.expect("re-register the owner under a new incarnation");
let outcome = tokio::time::timeout(T, get)
.await
.expect("get must resolve, not hang")
.expect("get task must not panic");
assert!(
outcome.is_ok() || matches!(outcome, Err(RmaError::Ucx { .. })),
"unexpected get outcome: {outcome:?}"
);
tokio::time::timeout(T, pair.puller_rma.unmap_region(local.region_id))
.await
.expect("unmap must resolve")
.expect("unmap succeeds once the cancelled op has been accounted for");
pair.owner_rma
.unmap_region(remote.region_id)
.await
.expect("unmap source");
decoy.transport.shutdown();
pair.owner.transport.shutdown();
pair.puller.transport.shutdown();
assert_pair_balanced(&pair);
assert_rma_balanced(&decoy);
}
#[tokio::test(flavor = "multi_thread")]
async fn shutdown_with_inflight_get() {
const LEN: usize = 32 * 1024 * 1024;
let mut src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
src.fill_pattern();
let pair = start_rma_pair().await;
let remote = pair.owner_rma.map_region(src.addr(), LEN).await.unwrap();
let local = pair.puller_rma.map_region(dst.addr(), LEN).await.unwrap();
let endpoint = pair.puller_rma.clone();
let req = get_request(&pair, &src, &remote, &local);
let get = tokio::spawn(async move { endpoint.get(req).await });
tokio::time::sleep(Duration::from_millis(2)).await;
pair.puller.transport.shutdown();
let outcome = tokio::time::timeout(T, get)
.await
.expect("get must resolve, not hang")
.expect("get task must not panic");
match outcome {
Ok(()) | Err(RmaError::ShuttingDown) | Err(RmaError::ChannelClosed) => {}
Err(RmaError::Ucx { status_name }) => {
println!("get completed with ucx status: {status_name}");
}
Err(other) => panic!("unexpected get outcome: {other}"),
}
assert!(matches!(
tokio::time::timeout(T, pair.puller_rma.unmap_region(local.region_id))
.await
.expect("unmap must resolve"),
Err(RmaError::ShuttingDown)
));
pair.owner.transport.shutdown();
assert_pair_balanced(&pair);
}
#[tokio::test(flavor = "multi_thread")]
async fn worker_rejects_bad_get_commands() {
const LEN: usize = 64 * 1024;
let src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
let pair = start_rma_pair().await;
let remote = pair.owner_rma.map_region(src.addr(), LEN).await.unwrap();
let local = pair.puller_rma.map_region(dst.addr(), LEN).await.unwrap();
let puller = &*pair.puller.transport;
let mut req = get_request(&pair, &src, &remote, &local);
req.local_offset = 1;
assert!(matches!(
ring_get(puller, req).await,
Err(RmaError::OutOfRange)
));
let mut req = get_request(&pair, &src, &remote, &local);
req.local_offset = u64::MAX;
req.len = 8;
assert!(matches!(
ring_get(puller, req).await,
Err(RmaError::OutOfRange)
));
let mut req = get_request(&pair, &src, &remote, &local);
req.local_region = local.region_id + 4_096;
assert!(matches!(
ring_get(puller, req).await,
Err(RmaError::RegionNotFound)
));
let mut req = get_request(&pair, &src, &remote, &local);
req.peer = InstanceId::new_v4();
assert!(matches!(
ring_get(puller, req).await,
Err(RmaError::PeerNotRegistered(_))
));
let mut req = get_request(&pair, &src, &remote, &local);
req.packed_rkey = Bytes::new();
assert!(matches!(
ring_get(puller, req).await,
Err(RmaError::InvalidRkey)
));
let mut req = get_request(&pair, &src, &remote, &local);
req.packed_rkey = Bytes::from(vec![0xABu8; 2048]);
assert!(matches!(
ring_get(puller, req).await,
Err(RmaError::InvalidRkey)
));
let mut req = get_request(&pair, &src, &remote, &local);
req.len = 0;
req.packed_rkey = Bytes::new();
ring_get(puller, req)
.await
.expect("a zero-length GET is a no-op, whatever key it carries");
assert_eq!(
pair.puller
.transport
.shared
.live_rkeys
.load(Ordering::SeqCst),
0
);
pair.puller_rma.unmap_region(local.region_id).await.unwrap();
pair.owner_rma.unmap_region(remote.region_id).await.unwrap();
pair.owner.transport.shutdown();
pair.puller.transport.shutdown();
assert_pair_balanced(&pair);
}
#[tokio::test(flavor = "multi_thread")]
async fn out_of_range_mem_type_is_refused_before_ucx() {
const LEN: usize = 64 * 1024;
let src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
let pair = start_rma_pair().await;
let remote = pair.owner_rma.map_region(src.addr(), LEN).await.unwrap();
let local = pair.puller_rma.map_region(dst.addr(), LEN).await.unwrap();
let mut blob = 1u64.to_le_bytes().to_vec();
blob.push(0xFF);
blob.push(0);
blob.push(0xFF);
let mut req = get_request(&pair, &src, &remote, &local);
req.packed_rkey = Bytes::from(blob.clone());
req.len = 64;
assert!(matches!(
pair.puller_rma.get(req).await,
Err(RmaError::InvalidRkey)
));
let mut req = get_request(&pair, &src, &remote, &local);
req.packed_rkey = Bytes::from(blob);
req.len = 64;
assert!(matches!(
ring_get(&pair.puller.transport, req).await,
Err(RmaError::InvalidRkey)
));
assert_eq!(
pair.puller
.transport
.shared
.live_rkeys
.load(Ordering::SeqCst),
0
);
pair.puller_rma.unmap_region(local.region_id).await.unwrap();
pair.owner_rma.unmap_region(remote.region_id).await.unwrap();
pair.owner.transport.shutdown();
pair.puller.transport.shutdown();
assert_pair_balanced(&pair);
}
#[tokio::test(flavor = "multi_thread")]
async fn unusable_rkey_is_refused_by_ucx_over_tcp() {
const LEN: usize = 64 * 1024;
let src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
let pair = start_rma_pair().await;
let remote = pair.owner_rma.map_region(src.addr(), LEN).await.unwrap();
let local = pair.puller_rma.map_region(dst.addr(), LEN).await.unwrap();
let mut well_formed = 1u64.to_le_bytes().to_vec();
well_formed.push(0); well_formed.push(0); well_formed.push(0xFF); preparse_packed_rkey(&well_formed).expect("this blob is self-terminating");
let mut req = get_request(&pair, &src, &remote, &local);
req.packed_rkey = Bytes::from(well_formed);
req.len = 64;
let outcome = ring_get(&pair.puller.transport, req).await;
assert!(
matches!(outcome, Err(RmaError::Ucx { .. })),
"UCX should refuse an unreachable memory domain, got {outcome:?}"
);
assert_eq!(
pair.puller
.transport
.shared
.live_rkeys
.load(Ordering::SeqCst),
0,
"a failed unpack must not be counted as a live rkey"
);
pair.puller_rma.unmap_region(local.region_id).await.unwrap();
pair.owner_rma.unmap_region(remote.region_id).await.unwrap();
pair.owner.transport.shutdown();
pair.puller.transport.shutdown();
assert_pair_balanced(&pair);
}
#[tokio::test(flavor = "multi_thread")]
async fn refused_rma_commands_answer_their_callers() {
let (map_tx, map_rx) = tokio::sync::oneshot::channel();
Cmd::MapRegion {
ptr: 0x1000,
len: 4096,
region_id: 1,
reply: map_tx,
}
.refuse_for_shutdown();
assert!(matches!(
map_rx.await.expect("MapRegion reply must be sent"),
Err(RmaError::ShuttingDown)
));
let (unmap_tx, unmap_rx) = tokio::sync::oneshot::channel();
Cmd::UnmapRegion {
region_id: 1,
reply: unmap_tx,
}
.refuse_for_shutdown();
assert!(matches!(
unmap_rx.await.expect("UnmapRegion reply must be sent"),
Err(RmaError::ShuttingDown)
));
let (get_tx, get_rx) = tokio::sync::oneshot::channel();
Cmd::RmaGet {
req: RmaGetRequest {
peer: InstanceId::new_v4(),
remote_addr: 0x2000,
packed_rkey: Bytes::from_static(&[1, 2, 3]),
local_region: 1,
local_offset: 0,
len: 16,
},
reply: get_tx,
}
.refuse_for_shutdown();
assert!(matches!(
get_rx.await.expect("RmaGet reply must be sent"),
Err(RmaError::ShuttingDown)
));
}
#[tokio::test(flavor = "multi_thread")]
async fn map_region_cancel_rolls_back() {
const LEN: usize = 64 * 1024;
let buf = PageBuf::new(LEN);
let node = start_node().await;
let rma = node.transport.rdma_endpoint();
let live = || node.transport.shared.live_regions.load(Ordering::SeqCst);
let (tx, rx) = tokio::sync::oneshot::channel();
drop(rx);
node.transport
.shared
.ring_tx
.send_async(Cmd::MapRegion {
ptr: buf.addr(),
len: LEN,
region_id: u64::MAX / 2,
reply: tx,
})
.await
.expect("ring accepts the command");
node.transport.shared.doorbell.ring();
assert!(
wait_until(T, || live() == 0).await,
"an orphaned registration must be rolled back by the worker"
);
let probe = rma.map_region(buf.addr(), LEN).await.expect("map succeeds");
assert_eq!(live(), 1);
rma.unmap_region(probe.region_id).await.expect("unmap");
assert_eq!(live(), 0);
let mut pending = Box::pin(rma.map_region(buf.addr(), LEN));
assert!(
futures::poll!(pending.as_mut()).is_pending(),
"the first poll should submit and then await the reply"
);
drop(pending);
assert!(
wait_until(T, || live() == 0).await,
"a cancelled map_region must leave no region behind"
);
node.transport.shutdown();
assert_rma_balanced(&node);
}
#[tokio::test(flavor = "multi_thread")]
async fn unmap_cancel_then_retry() {
const LEN: usize = 32 * 1024 * 1024;
let mut src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
src.fill_pattern();
let pair = start_rma_pair().await;
let remote = pair.owner_rma.map_region(src.addr(), LEN).await.unwrap();
let local = pair.puller_rma.map_region(dst.addr(), LEN).await.unwrap();
let endpoint = pair.puller_rma.clone();
let req = get_request(&pair, &src, &remote, &local);
let get = tokio::spawn(async move { endpoint.get(req).await });
assert!(
wait_until(T, || {
pair.puller
.transport
.shared
.inflight_ops
.load(Ordering::SeqCst)
>= 1
})
.await,
"the GET must be posted before the unmap is issued"
);
let orphaned = ring_unmap_enqueue(&pair.puller.transport, local.region_id).await;
drop(orphaned);
ring_unmap_enqueue(&pair.puller.transport, u64::MAX)
.await
.await
.expect("fence reply is sent")
.expect("unmapping an unknown id is a no-op");
assert_eq!(
pair.puller
.transport
.shared
.live_regions
.load(Ordering::SeqCst),
1,
"the region must stay mapped while its GET is in flight"
);
tokio::time::timeout(T, pair.puller_rma.unmap_region(local.region_id))
.await
.expect("retry must resolve")
.expect("retry must report success, not RegionNotFound");
get.await
.expect("get task")
.expect("the GET completes before the region goes");
assert_eq!(dst.as_slice(), src.as_slice());
assert_eq!(
pair.puller
.transport
.shared
.live_regions
.load(Ordering::SeqCst),
0,
"the puller's region is gone once the retry reports success"
);
pair.owner_rma.unmap_region(remote.region_id).await.unwrap();
pair.owner.transport.shutdown();
pair.puller.transport.shutdown();
assert_pair_balanced(&pair);
}
#[tokio::test(flavor = "multi_thread")]
async fn shutdown_resolves_parked_unmap_and_get() {
const LEN: usize = 32 * 1024 * 1024;
let mut src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
src.fill_pattern();
let pair = start_rma_pair().await;
let remote = pair.owner_rma.map_region(src.addr(), LEN).await.unwrap();
let local = pair.puller_rma.map_region(dst.addr(), LEN).await.unwrap();
let endpoint = pair.puller_rma.clone();
let req = get_request(&pair, &src, &remote, &local);
let get = tokio::spawn(async move { endpoint.get(req).await });
assert!(
wait_until(T, || {
pair.puller
.transport
.shared
.inflight_ops
.load(Ordering::SeqCst)
>= 1
})
.await,
"the GET must be posted before the unmap is issued"
);
let unmap = ring_unmap_enqueue(&pair.puller.transport, local.region_id).await;
pair.puller.transport.shutdown();
let unmap_outcome = tokio::time::timeout(T, unmap)
.await
.expect("the parked unmap must resolve, not hang")
.expect("the reply must be sent, not dropped");
let get_outcome = tokio::time::timeout(T, get)
.await
.expect("the GET must resolve, not hang")
.expect("get task");
println!("shutdown: unmap={unmap_outcome:?} get={get_outcome:?}");
pair.owner.transport.shutdown();
assert_pair_balanced(&pair);
}
#[tokio::test(flavor = "multi_thread")]
async fn get_cancel_still_releases_the_region() {
const LEN: usize = 32 * 1024 * 1024;
let mut src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
src.fill_pattern();
let pair = start_rma_pair().await;
let remote = pair.owner_rma.map_region(src.addr(), LEN).await.unwrap();
let local = pair.puller_rma.map_region(dst.addr(), LEN).await.unwrap();
let cancelled = tokio::time::timeout(
Duration::from_millis(1),
pair.puller_rma
.get(get_request(&pair, &src, &remote, &local)),
)
.await;
assert!(cancelled.is_err(), "the GET must be cancelled mid-transfer");
tokio::time::timeout(T, pair.puller_rma.unmap_region(local.region_id))
.await
.expect("unmap must resolve after an abandoned GET")
.expect("unmap succeeds");
assert_eq!(
dst.as_slice(),
src.as_slice(),
"the abandoned transfer still ran to completion"
);
pair.owner_rma.unmap_region(remote.region_id).await.unwrap();
pair.owner.transport.shutdown();
pair.puller.transport.shutdown();
assert_pair_balanced(&pair);
}
#[test]
fn preparse_rejects_blobs_ucx_would_walk_off_the_end_of() {
let mut truncated = vec![0xFFu8; 8];
truncated.push(0);
assert_eq!(truncated.len(), 9);
assert!(matches!(
preparse_packed_rkey(&truncated),
Err(RmaError::InvalidRkey)
));
let mut tail_declares_255 = 1u64.to_le_bytes().to_vec();
tail_declares_255.push(0); tail_declares_255.push(255); assert_eq!(tail_declares_255.len(), 10);
assert!(matches!(
preparse_packed_rkey(&tail_declares_255),
Err(RmaError::InvalidRkey)
));
let mut crafted = 0b111_1111u64.to_le_bytes().to_vec();
crafted.push(0); for _ in 0..6 {
crafted.push(168);
crafted.extend(std::iter::repeat_n(0u8, 168));
}
crafted.push(255);
assert_eq!(crafted.len(), MAX_PACKED_RKEY);
assert!(matches!(
preparse_packed_rkey(&crafted),
Err(RmaError::InvalidRkey)
));
let mut unterminated = 1u64.to_le_bytes().to_vec();
unterminated.push(0); unterminated.push(0); unterminated.push(7); unterminated.extend_from_slice(&[1, 2, 3]); assert!(matches!(
preparse_packed_rkey(&unterminated),
Err(RmaError::InvalidRkey)
));
let mut terminated = unterminated.clone();
terminated.push(0xFF);
preparse_packed_rkey(&terminated).expect("a terminated distance list is parseable");
assert!(matches!(
preparse_packed_rkey(&[0u8; 4]),
Err(RmaError::InvalidRkey)
));
assert!(matches!(
preparse_packed_rkey(&[0u8; 8]),
Err(RmaError::InvalidRkey)
));
preparse_packed_rkey(&[0, 0, 0, 0, 0, 0, 0, 0, 0]).expect("an empty md_map is well formed");
let mut bad_mem_type = 1u64.to_le_bytes().to_vec();
bad_mem_type.push(0xFF); bad_mem_type.push(0); bad_mem_type.push(SYS_DEV_UNKNOWN); assert!(matches!(
preparse_packed_rkey(&bad_mem_type),
Err(RmaError::InvalidRkey)
));
let mut good_mem_type = bad_mem_type.clone();
good_mem_type[8] = 9; preparse_packed_rkey(&good_mem_type).expect("an in-range mem_type is accepted");
good_mem_type[8] = 0; preparse_packed_rkey(&good_mem_type).expect("host memory is accepted");
}
#[tokio::test(flavor = "multi_thread")]
async fn preparse_accepts_real_packed_rkeys() {
const LEN: usize = 64 * 1024;
let buf = PageBuf::new(LEN);
let node = start_node().await;
let rma = node.transport.rdma_endpoint();
for len in [4096usize, LEN] {
let region = rma.map_region(buf.addr(), len).await.expect("map");
preparse_packed_rkey(®ion.packed_rkey).unwrap_or_else(|e| {
panic!(
"a genuine {}-byte rkey was rejected: {e}",
region.packed_rkey.len()
)
});
rma.unmap_region(region.region_id).await.expect("unmap");
}
node.transport.shutdown();
assert_rma_balanced(&node);
}
#[tokio::test(flavor = "multi_thread")]
async fn truncated_rkey_is_refused_before_ucx() {
const LEN: usize = 64 * 1024;
let src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
let pair = start_rma_pair().await;
let remote = pair.owner_rma.map_region(src.addr(), LEN).await.unwrap();
let local = pair.puller_rma.map_region(dst.addr(), LEN).await.unwrap();
let mut hostile = vec![0xFFu8; 8];
hostile.push(0);
let mut req = get_request(&pair, &src, &remote, &local);
req.packed_rkey = Bytes::from(hostile.clone());
req.len = 64;
assert!(matches!(
pair.puller_rma.get(req).await,
Err(RmaError::InvalidRkey)
));
let mut req = get_request(&pair, &src, &remote, &local);
req.packed_rkey = Bytes::from(hostile);
req.len = 64;
assert!(matches!(
ring_get(&pair.puller.transport, req).await,
Err(RmaError::InvalidRkey)
));
assert_eq!(
pair.puller
.transport
.shared
.live_rkeys
.load(Ordering::SeqCst),
0
);
pair.puller_rma.unmap_region(local.region_id).await.unwrap();
pair.owner_rma.unmap_region(remote.region_id).await.unwrap();
pair.owner.transport.shutdown();
pair.puller.transport.shutdown();
assert_pair_balanced(&pair);
}
#[tokio::test(flavor = "multi_thread")]
async fn peer_shutdown_during_get_answers_caller() {
const LEN: usize = 64 * 1024 * 1024;
let mut src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
src.fill_pattern();
let pair = start_rma_pair().await;
let remote = pair.owner_rma.map_region(src.addr(), LEN).await.unwrap();
let local = pair.puller_rma.map_region(dst.addr(), LEN).await.unwrap();
let endpoint = pair.puller_rma.clone();
let req = get_request(&pair, &src, &remote, &local);
let get = tokio::spawn(async move { endpoint.get(req).await });
assert!(
wait_until(T, || {
pair.puller
.transport
.shared
.inflight_ops
.load(Ordering::SeqCst)
>= 1
})
.await,
"the GET must be in flight before the owner goes"
);
pair.owner.transport.shutdown();
pair.puller.transport.shutdown();
let outcome = tokio::time::timeout(T, get)
.await
.expect("the GET must resolve, not hang")
.expect("get task must not panic");
println!("peer-shutdown GET resolved as {outcome:?}");
assert_pair_balanced(&pair);
}
const IDLE: Duration = Duration::from_millis(500);
fn eps_open(node: &Node) -> usize {
node.transport.shared.eps_open.load(Ordering::SeqCst)
}
fn eps_closed_idle(node: &Node) -> u64 {
node.transport.shared.eps_closed_idle.load(Ordering::SeqCst)
}
async fn ping_message(from: &Node, to: &Node, errs: &Arc<CountingErrors>) {
ping_message_to(from, to.instance_id, to, errs, "frame").await
}
async fn ping_message_to(
from: &Node,
target: InstanceId,
arrives_at: &Node,
errs: &Arc<CountingErrors>,
what: &str,
) {
let out = from.transport.send_message(
target,
Bytes::from_static(b"h"),
Bytes::from_static(b"p"),
MessageType::Message,
errs.clone(),
);
assert!(matches!(out, SendOutcome::Admitted));
assert!(
recv_message(&arrives_at.streams.message_stream, T)
.await
.is_some(),
"{what}: the frame must arrive before the endpoint is called established"
);
}
#[test]
fn a_sub_floor_ep_idle_timeout_is_clamped() {
let clamped = UcxTransportBuilder::new()
.ep_idle_timeout(Some(Duration::from_millis(1)))
.build()
.expect("build")
.config
.ep_idle_timeout;
assert_eq!(clamped, Some(super::MIN_EP_IDLE_TIMEOUT));
let honoured = UcxTransportBuilder::new()
.ep_idle_timeout(Some(Duration::from_secs(60)))
.build()
.expect("build")
.config
.ep_idle_timeout;
assert_eq!(honoured, Some(Duration::from_secs(60)));
assert_eq!(
UcxTransportBuilder::new()
.ep_idle_timeout(None)
.build()
.expect("build")
.config
.ep_idle_timeout,
None,
"explicitly disabling must stay disabled"
);
}
#[tokio::test(flavor = "multi_thread")]
async fn an_inbound_frame_refreshes_the_endpoint_it_arrived_on() {
let a = start_node_with(|b| b.ep_idle_timeout(Some(Duration::from_secs(3600)))).await;
let b = start_node().await;
cross_register(&a, &b);
let errs = CountingErrors::new();
ping_message(&a, &b, &errs).await;
assert_eq!(eps_open(&a), 1);
let before = a
.transport
.shared
.eps_stamped_inbound
.load(Ordering::SeqCst);
ping_message(&b, &a, &errs).await;
let stamped = wait_until(T, || {
a.transport
.shared
.eps_stamped_inbound
.load(Ordering::SeqCst)
> before
})
.await;
let unmatched = a
.transport
.shared
.eps_inbound_unmatched
.load(Ordering::SeqCst);
println!(
"reply_ep identity: stamped={} unmatched={unmatched}",
a.transport
.shared
.eps_stamped_inbound
.load(Ordering::SeqCst)
);
assert!(
stamped,
"an inbound frame from a peer we hold an endpoint to did not refresh it \
(unmatched sightings: {unmatched}). UCX is handing the recv callback a \
reply endpoint that is not the one `ucp_ep_create` gave us, so no inbound \
freshness stamp is possible at this layer — see the operator guidance on \
UcxTransportBuilder::ep_idle_timeout, which depends on this holding."
);
a.transport.shutdown();
b.transport.shutdown();
assert_rma_balanced(&a);
assert_rma_balanced(&b);
}
#[tokio::test(flavor = "multi_thread")]
async fn a_peer_that_keeps_sending_keeps_its_endpoint() {
let a = start_node_with(|b| b.ep_idle_timeout(Some(IDLE))).await;
let b = start_node().await;
cross_register(&a, &b);
let errs = CountingErrors::new();
ping_message(&a, &b, &errs).await;
assert_eq!(eps_open(&a), 1);
let deadline = tokio::time::Instant::now() + IDLE * 3;
while tokio::time::Instant::now() < deadline {
ping_message(&b, &a, &errs).await;
tokio::time::sleep(IDLE / 4).await;
}
assert_eq!(
eps_closed_idle(&a),
0,
"a peer's endpoint was reaped while that peer was actively sending to us"
);
assert_eq!(
eps_open(&a),
1,
"the endpoint was retired under live traffic"
);
assert_eq!(errs.count(), 0);
assert!(
wait_until(T, || eps_closed_idle(&a) >= 1).await,
"the endpoint was never reaped after the inbound traffic stopped"
);
a.transport.shutdown();
b.transport.shutdown();
assert_rma_balanced(&a);
assert_rma_balanced(&b);
}
#[tokio::test(flavor = "multi_thread")]
async fn idle_reaper_is_off_by_default() {
let a = start_node().await;
let b = start_node().await;
cross_register(&a, &b);
let errs = CountingErrors::new();
ping_message(&a, &b, &errs).await;
ping_message(&b, &a, &errs).await;
assert_eq!(eps_open(&a), 1);
tokio::time::sleep(Duration::from_millis(300)).await;
assert_eq!(
eps_closed_idle(&a),
0,
"the reaper ran without being configured"
);
assert_eq!(
eps_open(&a),
1,
"an endpoint was closed with the reaper off"
);
a.transport.shutdown();
b.transport.shutdown();
assert_rma_balanced(&a);
assert_rma_balanced(&b);
}
#[tokio::test(flavor = "multi_thread")]
async fn idle_endpoint_closes_and_the_next_send_wires_up_again() {
let a = start_node_with(|b| b.ep_idle_timeout(Some(IDLE))).await;
let b = start_node().await;
cross_register(&a, &b);
let errs = CountingErrors::new();
ping_message_to(&a, b.instance_id, &b, &errs, "first send").await;
assert_eq!(eps_open(&a), 1);
assert!(
wait_until(T, || eps_closed_idle(&a) >= 1).await,
"the idle endpoint was never reaped"
);
assert_eq!(eps_open(&a), 0, "the reaped endpoint was not retired");
ping_message_to(&a, b.instance_id, &b, &errs, "send after the reap").await;
assert!(
eps_open(&a) >= 1,
"the next send did not re-establish an endpoint"
);
assert_eq!(errs.count(), 0, "a send failed across the reap");
assert!(
a.transport.shared.failed_peers.is_empty(),
"an idle close must not mark the peer failed"
);
a.transport.shutdown();
b.transport.shutdown();
assert_rma_balanced(&a);
assert_rma_balanced(&b);
}
#[tokio::test(flavor = "multi_thread")]
async fn reaping_disrupts_the_peers_path_back() {
let a = start_node_with(|b| b.ep_idle_timeout(Some(IDLE))).await;
let b = start_node().await;
cross_register(&a, &b);
let errs = CountingErrors::new();
ping_message_to(&a, b.instance_id, &b, &errs, "first send").await;
assert!(
wait_until(T, || eps_closed_idle(&a) >= 1).await,
"the idle endpoint was never reaped"
);
let out = b.transport.send_message(
a.instance_id,
Bytes::from_static(b"h"),
Bytes::from_static(b"p"),
MessageType::Message,
errs.clone(),
);
assert!(
matches!(out, SendOutcome::Admitted),
"the peer's send is admitted; the loss is downstream of admission"
);
assert!(
recv_message(&a.streams.message_stream, Duration::from_millis(1500))
.await
.is_none(),
"the peer's first frame after our reap arrived — the UCX endpoint-matching \
interaction this test pins has been fixed, so update the caveat on \
UcxTransportBuilder::ep_idle_timeout and delete this test"
);
assert_eq!(
errs.count(),
0,
"the loss is silent at this point; keepalive is what eventually reports it"
);
a.transport.shutdown();
b.transport.shutdown();
assert_rma_balanced(&a);
assert_rma_balanced(&b);
}
#[tokio::test(flavor = "multi_thread")]
async fn an_endpoint_with_an_inflight_get_is_not_reaped() {
const CHUNK: usize = 8 * 1024 * 1024;
const CHUNKS: usize = 8;
const LEN: usize = CHUNK * CHUNKS;
let mut src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
src.fill_pattern();
let puller = start_node_with_config(UcxConfig {
tls: Some("tcp".into()),
ep_idle_timeout: Some(Duration::from_millis(10)),
eager_endpoints: true,
..UcxConfig::default()
})
.await;
let owner = start_node().await;
let idle_peer = start_node().await;
cross_register(&puller, &owner);
let owner_rma = owner.transport.rdma_endpoint();
let puller_rma = puller.transport.rdma_endpoint();
let remote = owner_rma
.map_region(src.addr(), LEN)
.await
.expect("map src");
let local = puller_rma
.map_region(dst.addr(), LEN)
.await
.expect("map dst");
let mut gets = Vec::with_capacity(CHUNKS);
for i in 0..CHUNKS {
let endpoint = puller_rma.clone();
let req = RmaGetRequest {
peer: owner.instance_id,
remote_addr: (src.addr() + i * CHUNK) as u64,
packed_rkey: remote.packed_rkey.clone(),
local_region: local.region_id,
local_offset: (i * CHUNK) as u64,
len: CHUNK as u64,
};
gets.push(tokio::spawn(async move { endpoint.get(req).await }));
}
assert!(
wait_until(T, || {
puller.transport.shared.inflight_ops.load(Ordering::SeqCst) >= CHUNKS
})
.await,
"all {CHUNKS} GETs must be posted before the idle peer is introduced"
);
let baseline = eps_closed_idle(&puller);
cross_register(&puller, &idle_peer);
let observed = Arc::new(AtomicUsize::new(usize::MAX));
let seen = {
let observed = Arc::clone(&observed);
wait_until(T, || {
let closed = eps_closed_idle(&puller);
let inflight = puller.transport.shared.inflight_ops.load(Ordering::SeqCst);
if closed > baseline && inflight > 0 {
observed.store((closed - baseline) as usize, Ordering::SeqCst);
true
} else {
false
}
})
.await
};
assert!(
seen,
"no scan closed the idle endpoint while the GETs were still outstanding; \
64 MiB over the tcp lane finished faster than a 10 ms idle window"
);
assert_eq!(
observed.load(Ordering::SeqCst),
1,
"the endpoint carrying the in-flight GETs was reaped too"
);
for (i, task) in gets.into_iter().enumerate() {
task.await
.expect("get task must not panic")
.unwrap_or_else(|e| panic!("get {i} failed: {e}"));
}
assert_eq!(dst.as_slice(), src.as_slice(), "every GET must have landed");
puller_rma
.unmap_region(local.region_id)
.await
.expect("unmap dst");
owner_rma
.unmap_region(remote.region_id)
.await
.expect("unmap src");
puller.transport.shutdown();
owner.transport.shutdown();
idle_peer.transport.shutdown();
assert_rma_balanced(&puller);
assert_rma_balanced(&owner);
assert_rma_balanced(&idle_peer);
}
#[tokio::test(flavor = "multi_thread")]
async fn shutdown_during_an_idle_close_is_clean() {
let a = start_node_with(|b| b.ep_idle_timeout(Some(IDLE))).await;
let b = start_node().await;
cross_register(&a, &b);
let errs = CountingErrors::new();
ping_message(&a, &b, &errs).await;
assert!(
wait_until(T, || eps_closed_idle(&a) >= 1).await,
"the idle endpoint was never reaped"
);
a.transport.shutdown();
b.transport.shutdown();
assert_rma_balanced(&a);
assert_rma_balanced(&b);
}
#[tokio::test(flavor = "multi_thread")]
async fn eager_endpoints_wire_up_at_registration() {
let eager = start_node_with(|b| b.eager_endpoints(true)).await;
let lazy = start_node().await;
cross_register(&eager, &lazy);
assert!(
wait_until(T, || eps_open(&eager) == 1).await,
"registration did not establish an endpoint eagerly"
);
tokio::time::sleep(Duration::from_millis(50)).await;
assert_eq!(
eps_open(&lazy),
0,
"the default wired an endpoint up without being asked to"
);
eager.transport.shutdown();
lazy.transport.shutdown();
assert_rma_balanced(&eager);
assert_rma_balanced(&lazy);
}
#[tokio::test(flavor = "multi_thread")]
async fn eager_wireup_follows_a_re_registration() {
let eager = start_node_with(|b| b.eager_endpoints(true)).await;
let first = start_node().await;
let restarted = start_node().await;
cross_register(&eager, &first);
assert!(
wait_until(T, || eps_open(&eager) == 1).await,
"the first registration did not wire up"
);
eager
.transport
.register(PeerInfo::new(
first.instance_id,
restarted.transport.address(),
))
.expect("re-register under a new incarnation");
assert!(
wait_until(T, || eps_open(&eager) == 1
&& eps_closed_idle(&eager) == 0
&& eager.transport.shared.failed_peers.is_empty())
.await,
"the re-registered peer did not end up on exactly one fresh endpoint"
);
let errs = CountingErrors::new();
ping_message_to(
&eager,
first.instance_id,
&restarted,
&errs,
"re-registered",
)
.await;
assert_eq!(errs.count(), 0);
eager.transport.shutdown();
first.transport.shutdown();
restarted.transport.shutdown();
assert_rma_balanced(&eager);
assert_rma_balanced(&first);
assert_rma_balanced(&restarted);
}
#[tokio::test(flavor = "multi_thread")]
async fn eager_wireup_and_the_reaper_compose() {
let a = start_node_with(|b| b.eager_endpoints(true).ep_idle_timeout(Some(IDLE))).await;
let b = start_node().await;
cross_register(&a, &b);
assert!(
wait_until(T, || eps_open(&a) == 1).await,
"eager wireup did not run"
);
assert!(
wait_until(T, || eps_closed_idle(&a) >= 1 && eps_open(&a) == 0).await,
"an eagerly established but unused endpoint was never reclaimed"
);
let errs = CountingErrors::new();
ping_message(&a, &b, &errs).await;
assert_eq!(errs.count(), 0, "the send after the reap failed");
a.transport.shutdown();
b.transport.shutdown();
assert_rma_balanced(&a);
assert_rma_balanced(&b);
}
#[tokio::test(flavor = "multi_thread")]
async fn rma_lifecycle_soak() {
const CYCLES: usize = 24;
const LEN: usize = 256 * 1024;
let mut src = PageBuf::new(LEN);
let dst = PageBuf::new(LEN);
src.fill_pattern();
let pair = start_rma_pair().await;
for cycle in 0..CYCLES {
let remote = pair
.owner_rma
.map_region(src.addr(), LEN)
.await
.unwrap_or_else(|e| panic!("cycle {cycle}: map src: {e}"));
let local = pair
.puller_rma
.map_region(dst.addr(), LEN)
.await
.unwrap_or_else(|e| panic!("cycle {cycle}: map dst: {e}"));
tokio::time::timeout(
T,
pair.puller_rma
.get(get_request(&pair, &src, &remote, &local)),
)
.await
.unwrap_or_else(|_| panic!("cycle {cycle}: get hung"))
.unwrap_or_else(|e| panic!("cycle {cycle}: get failed: {e}"));
assert_eq!(dst.as_slice(), src.as_slice(), "cycle {cycle}: bad data");
pair.puller_rma
.unmap_region(local.region_id)
.await
.unwrap_or_else(|e| panic!("cycle {cycle}: unmap dst: {e}"));
pair.owner_rma
.unmap_region(remote.region_id)
.await
.unwrap_or_else(|e| panic!("cycle {cycle}: unmap src: {e}"));
assert_eq!(
pair.puller
.transport
.shared
.live_regions
.load(Ordering::SeqCst),
0,
"cycle {cycle}: a destination region survived its unmap"
);
assert_eq!(
pair.owner
.transport
.shared
.live_regions
.load(Ordering::SeqCst),
0,
"cycle {cycle}: a source region survived its unmap"
);
assert_eq!(
pair.puller
.transport
.shared
.live_rkeys
.load(Ordering::SeqCst),
0,
"cycle {cycle}: an unpacked rkey survived its operation"
);
assert_eq!(
eps_open(&pair.puller),
1,
"cycle {cycle}: the puller's endpoint count drifted"
);
}
pair.owner.transport.shutdown();
pair.puller.transport.shutdown();
assert_pair_balanced(&pair);
}
#[tokio::test(flavor = "multi_thread")]
#[ignore = "benchmark: prints timings, asserts nothing"]
async fn bench_rma() {
const MAP_SIZES: [usize; 3] = [4 * 1024, 1024 * 1024, 64 * 1024 * 1024];
const GET_SIZES: [usize; 3] = [64 * 1024, 1024 * 1024, 16 * 1024 * 1024];
let largest = *MAP_SIZES.iter().max().unwrap();
let mut src = PageBuf::new(largest);
let dst = PageBuf::new(largest);
src.fill_pattern();
let pair = start_rma_pair().await;
println!("-- ucp_mem_map latency (UCX_TLS=tcp) --");
for len in MAP_SIZES {
let started = std::time::Instant::now();
let region = pair.owner_rma.map_region(src.addr(), len).await.unwrap();
let mapped = started.elapsed();
let started = std::time::Instant::now();
pair.owner_rma.unmap_region(region.region_id).await.unwrap();
println!(
" {:>9} B: map {:>10.3?} unmap {:>10.3?} rkey {} B",
len,
mapped,
started.elapsed(),
region.packed_rkey.len()
);
}
let remote = pair
.owner_rma
.map_region(src.addr(), largest)
.await
.unwrap();
let local = pair
.puller_rma
.map_region(dst.addr(), largest)
.await
.unwrap();
println!("-- ucp_get_nbx latency (UCX_TLS=tcp) --");
for len in GET_SIZES {
for round in 0..4 {
let started = std::time::Instant::now();
pair.puller_rma
.get(RmaGetRequest {
peer: pair.owner.instance_id,
remote_addr: src.addr() as u64,
packed_rkey: remote.packed_rkey.clone(),
local_region: local.region_id,
local_offset: 0,
len: len as u64,
})
.await
.expect("get succeeds");
let elapsed = started.elapsed();
if round > 0 {
let mib = len as f64 / (1024.0 * 1024.0);
println!(
" {:>9} B: {:>10.3?} ({:.0} MiB/s)",
len,
elapsed,
mib / elapsed.as_secs_f64()
);
}
}
}
pair.puller_rma.unmap_region(local.region_id).await.unwrap();
pair.owner_rma.unmap_region(remote.region_id).await.unwrap();
pair.owner.transport.shutdown();
pair.puller.transport.shutdown();
assert_pair_balanced(&pair);
}