use std::collections::HashMap;
use std::mem::MaybeUninit;
use std::os::raw::{c_int, c_void};
use std::sync::atomic::{AtomicBool, AtomicI64, AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, Mutex, OnceLock};
use std::time::{Duration, Instant};
use bytes::Bytes;
use dashmap::DashMap;
use tracing::{debug, warn};
use ucx_rs::{decode_status_ptr, status_string, sys};
use velo_ext::{AdmitOutcome, InstanceId, MessageType, TransportAdapter, TransportErrorHandler};
use super::address::{AM_ID_BASE, AM_KIND_COUNT, AM_KIND_PING, AM_KIND_PONG, UcxEndpoint};
use super::rma::{
MAX_PACKED_RKEY, MappedRegion, RKEY_UNPACK_PAD, RmaError, RmaGetRequest, validate_packed_rkey,
};
use super::transport::UcxConfig;
pub(crate) struct SendTask {
pub peer: InstanceId,
pub msg_type: MessageType,
pub header: Bytes,
pub payload: Bytes,
pub on_error: Arc<dyn TransportErrorHandler>,
}
impl SendTask {
pub(crate) fn fail(self, why: impl Into<String>) {
self.on_error
.on_error(self.header, self.payload, why.into());
}
}
pub(crate) enum Cmd {
Send(SendTask),
Ping {
peer: InstanceId,
token: u64,
},
PongTo {
reply_ep: usize,
token: u64,
},
ShuttingDownTo {
reply_ep: usize,
header: Bytes,
},
MapRegion {
ptr: usize,
len: usize,
region_id: u64,
reply: tokio::sync::oneshot::Sender<Result<MappedRegion, RmaError>>,
},
UnmapRegion {
region_id: u64,
reply: tokio::sync::oneshot::Sender<Result<(), RmaError>>,
},
RmaGet {
req: RmaGetRequest,
reply: tokio::sync::oneshot::Sender<Result<(), RmaError>>,
},
EnsureEp {
peer: InstanceId,
},
Shutdown,
}
impl Cmd {
pub(crate) fn refuse_for_shutdown(self) {
match self {
Cmd::Send(task) => task.fail("ucx transport shutting down"),
Cmd::MapRegion { reply, .. } => {
let _ = reply.send(Err(RmaError::ShuttingDown));
}
Cmd::UnmapRegion { reply, .. } => {
let _ = reply.send(Err(RmaError::ShuttingDown));
}
Cmd::RmaGet { reply, .. } => {
let _ = reply.send(Err(RmaError::ShuttingDown));
}
Cmd::Ping { .. }
| Cmd::PongTo { .. }
| Cmd::ShuttingDownTo { .. }
| Cmd::EnsureEp { .. }
| Cmd::Shutdown => {}
}
}
}
pub(crate) struct Doorbell {
armed: AtomicBool,
worker: Mutex<usize>,
}
impl Doorbell {
pub fn new() -> Self {
Self {
armed: AtomicBool::new(false),
worker: Mutex::new(0),
}
}
pub fn ring(&self) {
if self.armed.swap(false, Ordering::AcqRel) {
let guard = self.worker.lock().unwrap_or_else(|e| e.into_inner());
let raw = *guard;
if raw != 0 {
unsafe { sys::ucp_worker_signal(raw as sys::ucp_worker_h) };
}
}
}
pub fn ring_force(&self) {
self.armed.store(false, Ordering::Release);
let guard = self.worker.lock().unwrap_or_else(|e| e.into_inner());
let raw = *guard;
if raw != 0 {
unsafe { sys::ucp_worker_signal(raw as sys::ucp_worker_h) };
}
}
fn arm(&self) {
self.armed.store(true, Ordering::Release);
}
fn disarm(&self) {
self.armed.store(false, Ordering::Release);
}
fn install(&self, worker: sys::ucp_worker_h) {
*self.worker.lock().unwrap_or_else(|e| e.into_inner()) = worker as usize;
}
fn retire(&self) {
*self.worker.lock().unwrap_or_else(|e| e.into_inner()) = 0;
}
}
pub(crate) struct WorkerShared {
pub ring_tx: flume::Sender<Cmd>,
pub doorbell: Arc<Doorbell>,
pub peers: Arc<DashMap<InstanceId, UcxEndpoint>>,
pub pending_pings: Arc<DashMap<u64, tokio::sync::oneshot::Sender<()>>>,
pub failed_peers: Arc<DashMap<InstanceId, ()>>,
pub inflight_ops: Arc<AtomicUsize>,
pub shutdown_requested: Arc<AtomicBool>,
pub reg_epoch: Arc<AtomicU64>,
pub live_regions: Arc<AtomicUsize>,
pub live_rkeys: Arc<AtomicI64>,
pub eps_open: Arc<AtomicUsize>,
pub eps_stamped_inbound: Arc<AtomicU64>,
pub eps_inbound_unmatched: Arc<AtomicU64>,
pub eps_closed_idle: Arc<AtomicU64>,
pub reply_eps: Arc<ReplyEpSightings>,
}
pub(crate) struct StartupOut {
pub worker_addr: Vec<u8>,
pub max_am_header: usize,
}
pub(crate) struct WorkerArgs {
pub config: UcxConfig,
pub ring_rx: flume::Receiver<Cmd>,
pub shared: Arc<WorkerShared>,
pub adapter: TransportAdapter,
pub startup: tokio::sync::oneshot::Sender<anyhow::Result<StartupOut>>,
}
enum OpKind {
Frame {
header: Bytes,
payload: Bytes,
on_error: Arc<dyn TransportErrorHandler>,
},
Control { _hold: Bytes },
}
struct OpState {
kind: OpKind,
inflight: Arc<AtomicUsize>,
}
impl OpState {
fn complete(self: Arc<Self>, status: sys::ucs_status_t) {
if status != sys::ucs_status_t_UCS_OK
&& let Some(state) = Arc::into_inner(self)
{
match state.kind {
OpKind::Frame {
header,
payload,
on_error,
} => {
on_error.on_error(
header,
payload,
format!("ucx send failed: {}", status_string(status)),
);
}
OpKind::Control { .. } => {
debug!("ucx control send failed: {}", status_string(status));
}
}
}
}
}
unsafe extern "C" fn send_trampoline(
request: *mut c_void,
status: sys::ucs_status_t,
user_data: *mut c_void,
) {
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let state = unsafe { Arc::from_raw(user_data as *const OpState) };
state.inflight.fetch_sub(1, Ordering::AcqRel);
state.complete(status);
}));
if !request.is_null() {
unsafe { sys::ucp_request_free(request) };
}
}
type RmaCompletion = (u64, u64);
type RmaCompletions = Vec<RmaCompletion>;
struct RmaOpState {
rkey: Mutex<Option<usize>>,
peer: InstanceId,
region_id: u64,
op_id: u64,
reply: Mutex<Option<tokio::sync::oneshot::Sender<Result<(), RmaError>>>>,
inflight: Arc<AtomicUsize>,
live_rkeys: Arc<AtomicI64>,
rma_completions: Arc<Mutex<RmaCompletions>>,
}
impl RmaOpState {
fn complete(&self, status: sys::ucs_status_t) {
if let Some(rkey) = self.rkey.lock().unwrap_or_else(|e| e.into_inner()).take() {
unsafe { sys::ucp_rkey_destroy(rkey as sys::ucp_rkey_h) };
self.live_rkeys.fetch_sub(1, Ordering::Relaxed);
}
self.resolve(if status == sys::ucs_status_t_UCS_OK {
Ok(())
} else {
Err(RmaError::Ucx {
status_name: status_string(status),
})
});
self.rma_completions
.lock()
.unwrap_or_else(|e| e.into_inner())
.push((self.region_id, self.op_id));
}
fn resolve(&self, result: Result<(), RmaError>) {
if let Some(reply) = self.reply.lock().unwrap_or_else(|e| e.into_inner()).take() {
let _ = reply.send(result);
}
}
}
unsafe extern "C" fn rma_trampoline(
request: *mut c_void,
status: sys::ucs_status_t,
user_data: *mut c_void,
) {
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let state = unsafe { Arc::from_raw(user_data as *const RmaOpState) };
state.inflight.fetch_sub(1, Ordering::AcqRel);
state.complete(status);
}));
if !request.is_null() {
unsafe { sys::ucp_request_free(request) };
}
}
struct RecvShared {
adapter: TransportAdapter,
ring_tx: flume::Sender<Cmd>,
pending_pings: Arc<DashMap<u64, tokio::sync::oneshot::Sender<()>>>,
reply_eps: Arc<ReplyEpSightings>,
stamp_inbound: bool,
}
const REPLY_EP_SLOTS: usize = 8;
pub(crate) struct ReplyEpSightings {
slots: [AtomicUsize; REPLY_EP_SLOTS],
recorded: AtomicUsize,
}
impl ReplyEpSightings {
pub(crate) fn new() -> Self {
Self {
slots: [const { AtomicUsize::new(0) }; REPLY_EP_SLOTS],
recorded: AtomicUsize::new(0),
}
}
fn record(&self, ep: usize) {
let seq = self.recorded.fetch_add(1, Ordering::Relaxed);
self.slots[seq % REPLY_EP_SLOTS].store(ep, Ordering::Release);
}
fn drain_into(&self, seen: &mut usize, out: &mut Vec<usize>) {
let recorded = self.recorded.load(Ordering::Acquire);
if recorded == *seen {
return;
}
*seen = recorded;
for slot in &self.slots {
let ep = slot.swap(0, Ordering::AcqRel);
if ep != 0 {
out.push(ep);
}
}
}
}
struct RecvArg {
shared: Arc<RecvShared>,
kind: u8,
}
unsafe extern "C" fn recv_trampoline(
arg: *mut c_void,
header: *const c_void,
header_length: usize,
data: *mut c_void,
length: usize,
param: *const sys::ucp_am_recv_param_t,
) -> sys::ucs_status_t {
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let ra = unsafe { &*(arg as *const RecvArg) };
let p = unsafe { &*param };
if ra.shared.stamp_inbound && !p.reply_ep.is_null() {
ra.shared.reply_eps.record(p.reply_ep as usize);
}
if p.recv_attr & sys::ucp_am_recv_attr_t_UCP_AM_RECV_ATTR_FLAG_RNDV as u64 != 0 {
warn!("ucx: rejecting rendezvous-mode AM (kind {})", ra.kind);
return sys::ucs_status_t_UCS_ERR_UNSUPPORTED;
}
let header = if header_length == 0 {
Bytes::new()
} else {
Bytes::copy_from_slice(unsafe {
std::slice::from_raw_parts(header as *const u8, header_length)
})
};
let payload = if length == 0 {
Bytes::new()
} else {
Bytes::copy_from_slice(unsafe { std::slice::from_raw_parts(data as *const u8, length) })
};
match ra.kind {
AM_KIND_PING => {
if header.len() >= 8 && !p.reply_ep.is_null() {
let token = u64::from_le_bytes(header[..8].try_into().unwrap());
let _ = ra.shared.ring_tx.try_send(Cmd::PongTo {
reply_ep: p.reply_ep as usize,
token,
});
}
}
AM_KIND_PONG => {
if header.len() >= 8 {
let token = u64::from_le_bytes(header[..8].try_into().unwrap());
if let Some((_, tx)) = ra.shared.pending_pings.remove(&token) {
let _ = tx.send(());
}
}
}
kind => {
let adapter = &ra.shared.adapter;
match MessageType::from_u8(kind) {
Some(MessageType::Message) => {
match adapter.admit_message(header, payload) {
AdmitOutcome::Admitted => {}
AdmitOutcome::Draining { header, .. } => {
if !p.reply_ep.is_null()
&& ra
.shared
.ring_tx
.try_send(Cmd::ShuttingDownTo {
reply_ep: p.reply_ep as usize,
header,
})
.is_err()
{
debug!("ucx: drain echo dropped (ring full)");
}
}
AdmitOutcome::Disconnected { .. } => {
debug!("ucx: inbound Message dropped (receiver gone)");
}
}
}
Some(MessageType::Response) => {
let _ = adapter.response_stream.send((header, payload));
}
Some(MessageType::ShuttingDown) => {
let _ = adapter.shutdown_stream.send((header, payload));
}
Some(MessageType::Ack) | Some(MessageType::Event) => {
let _ = adapter.event_stream.send((header, payload));
}
None => {
warn!("ucx: inbound AM with unknown kind {kind}");
}
}
}
}
sys::ucs_status_t_UCS_OK
}));
result.unwrap_or_else(|_| {
warn!("ucx: panic in AM recv handler (message dropped)");
sys::ucs_status_t_UCS_OK
})
}
struct ErrArg {
peer: InstanceId,
failed: Arc<DashMap<InstanceId, ()>>,
err_events: Arc<Mutex<Vec<InstanceId>>>,
}
unsafe extern "C" fn err_trampoline(
arg: *mut c_void,
_ep: sys::ucp_ep_h,
status: sys::ucs_status_t,
) {
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let ea = unsafe { &*(arg as *const ErrArg) };
warn!(
"ucx: endpoint to {} failed: {}",
ea.peer,
status_string(status)
);
ea.failed.insert(ea.peer, ());
ea.err_events
.lock()
.unwrap_or_else(|e| e.into_inner())
.push(ea.peer);
}));
}
struct EpEntry {
ep: sys::ucp_ep_h,
err_arg: *mut ErrArg,
incarnation: u64,
last_used: Instant,
}
struct RegionEntry {
memh: sys::ucp_mem_h,
requested_addr: u64,
requested_len: u64,
effective_addr: u64,
effective_len: u64,
inflight: usize,
pending_unmap: Vec<tokio::sync::oneshot::Sender<Result<(), RmaError>>>,
}
struct PreparedGet {
ep: sys::ucp_ep_h,
rkey: sys::ucp_rkey_h,
memh: sys::ucp_mem_h,
local_addr: u64,
remote_addr: u64,
len: usize,
region_id: u64,
peer: InstanceId,
}
struct WorkerState {
context: sys::ucp_context_h,
worker: sys::ucp_worker_h,
efd: c_int,
eps: HashMap<InstanceId, EpEntry>,
err_events: Arc<Mutex<Vec<InstanceId>>>,
regions: HashMap<u64, RegionEntry>,
rma_ops: HashMap<u64, Arc<RmaOpState>>,
next_op_id: u64,
rma_completions: Arc<Mutex<RmaCompletions>>,
seen_reg_epoch: u64,
shared: Arc<WorkerShared>,
config: UcxConfig,
_recv_args: Vec<Arc<RecvArg>>,
parked_for_close: Vec<EpEntry>,
pending_closes: Vec<sys::ucs_status_ptr_t>,
now: Instant,
next_ep_scan: Instant,
seen_reply_eps: usize,
reply_ep_scratch: Vec<usize>,
}
fn ep_scan_period(timeout: Duration) -> Duration {
(timeout / 2).clamp(Duration::from_millis(10), Duration::from_secs(1))
}
pub(crate) fn worker_main(args: WorkerArgs) {
let WorkerArgs {
config,
ring_rx,
shared,
adapter,
startup,
} = args;
let state = match unsafe { init_ucx(&config, &shared, &adapter) } {
Ok((state, out)) => {
let _ = startup.send(Ok(out));
state
}
Err(e) => {
let _ = startup.send(Err(e));
return;
}
};
run_loop(state, ring_rx);
}
unsafe fn init_ucx(
config: &UcxConfig,
shared: &Arc<WorkerShared>,
adapter: &TransportAdapter,
) -> anyhow::Result<(WorkerState, StartupOut)> {
unsafe {
let mut ucp_cfg: *mut sys::ucp_config_t = std::ptr::null_mut();
let st = sys::ucp_config_read(std::ptr::null(), std::ptr::null(), &mut ucp_cfg);
anyhow::ensure!(
st == sys::ucs_status_t_UCS_OK,
"ucp_config_read: {}",
status_string(st)
);
for (key, value) in [("RCACHE_ENABLE", "n"), ("MEM_EVENTS", "n")] {
if std::env::var_os(format!("UCX_{key}")).is_none() {
let k = std::ffi::CString::new(key).unwrap();
let v = std::ffi::CString::new(value).unwrap();
let _ = sys::ucp_config_modify(ucp_cfg, k.as_ptr(), v.as_ptr());
}
}
if let Some(tls) = &config.tls
&& std::env::var_os("UCX_TLS").is_none()
{
let k = std::ffi::CString::new("TLS").unwrap();
let v = std::ffi::CString::new(tls.as_str()).unwrap();
let st = sys::ucp_config_modify(ucp_cfg, k.as_ptr(), v.as_ptr());
anyhow::ensure!(
st == sys::ucs_status_t_UCS_OK,
"ucp_config_modify(TLS={tls}): {}",
status_string(st)
);
}
if let Some(devices) = &config.net_devices
&& std::env::var_os("UCX_NET_DEVICES").is_none()
{
let k = std::ffi::CString::new("NET_DEVICES").unwrap();
let v = std::ffi::CString::new(devices.as_str()).unwrap();
let st = sys::ucp_config_modify(ucp_cfg, k.as_ptr(), v.as_ptr());
anyhow::ensure!(
st == sys::ucs_status_t_UCS_OK,
"ucp_config_modify(NET_DEVICES={devices}): {}",
status_string(st)
);
}
let mut params: sys::ucp_params_t = MaybeUninit::zeroed().assume_init();
params.field_mask = (sys::ucp_params_field_UCP_PARAM_FIELD_FEATURES
| sys::ucp_params_field_UCP_PARAM_FIELD_MT_WORKERS_SHARED)
as u64;
params.features = (sys::ucp_feature_UCP_FEATURE_AM
| sys::ucp_feature_UCP_FEATURE_RMA
| sys::ucp_feature_UCP_FEATURE_WAKEUP) as u64;
params.mt_workers_shared = 1;
let mut context: sys::ucp_context_h = std::ptr::null_mut();
let st = sys::ucp_init_version(
sys::UCP_API_MAJOR,
sys::UCP_API_MINOR,
¶ms,
ucp_cfg,
&mut context,
);
sys::ucp_config_release(ucp_cfg);
anyhow::ensure!(
st == sys::ucs_status_t_UCS_OK,
"ucp_init: {} (if this is InvalidParam with no UCX log output, a \
constructor reference is missing — see ucx-rs)",
status_string(st)
);
let mut wparams: sys::ucp_worker_params_t = MaybeUninit::zeroed().assume_init();
wparams.field_mask = sys::ucp_worker_params_field_UCP_WORKER_PARAM_FIELD_THREAD_MODE as u64;
wparams.thread_mode = sys::ucs_thread_mode_t_UCS_THREAD_MODE_SINGLE;
let mut worker: sys::ucp_worker_h = std::ptr::null_mut();
let st = sys::ucp_worker_create(context, &wparams, &mut worker);
if st != sys::ucs_status_t_UCS_OK {
sys::ucp_cleanup(context);
anyhow::bail!("ucp_worker_create: {}", status_string(st));
}
let recv_shared = Arc::new(RecvShared {
adapter: adapter.clone(),
ring_tx: shared.ring_tx.clone(),
pending_pings: Arc::clone(&shared.pending_pings),
reply_eps: Arc::clone(&shared.reply_eps),
stamp_inbound: config.ep_idle_timeout.is_some(),
});
let mut recv_args = Vec::with_capacity(AM_KIND_COUNT as usize);
for kind in 0..AM_KIND_COUNT {
let arg = Arc::new(RecvArg {
shared: Arc::clone(&recv_shared),
kind,
});
let mut hp: sys::ucp_am_handler_param_t = MaybeUninit::zeroed().assume_init();
hp.field_mask = (sys::ucp_am_handler_param_field_UCP_AM_HANDLER_PARAM_FIELD_ID
| sys::ucp_am_handler_param_field_UCP_AM_HANDLER_PARAM_FIELD_CB
| sys::ucp_am_handler_param_field_UCP_AM_HANDLER_PARAM_FIELD_ARG)
as u64;
hp.id = (AM_ID_BASE as u32) + kind as u32;
hp.cb = Some(recv_trampoline);
hp.arg = Arc::as_ptr(&arg) as *mut c_void;
let st = sys::ucp_worker_set_am_recv_handler(worker, &hp);
if st != sys::ucs_status_t_UCS_OK {
sys::ucp_worker_destroy(worker);
sys::ucp_cleanup(context);
anyhow::bail!("set_am_recv_handler(kind {kind}): {}", status_string(st));
}
recv_args.push(arg);
}
let mut attr: sys::ucp_worker_attr_t = MaybeUninit::zeroed().assume_init();
attr.field_mask = (sys::ucp_worker_attr_field_UCP_WORKER_ATTR_FIELD_ADDRESS
| sys::ucp_worker_attr_field_UCP_WORKER_ATTR_FIELD_MAX_AM_HEADER)
as u64;
let st = sys::ucp_worker_query(worker, &mut attr);
if st != sys::ucs_status_t_UCS_OK {
sys::ucp_worker_destroy(worker);
sys::ucp_cleanup(context);
anyhow::bail!("ucp_worker_query: {}", status_string(st));
}
let worker_addr =
std::slice::from_raw_parts(attr.address as *const u8, attr.address_length).to_vec();
sys::ucp_worker_release_address(worker, attr.address);
let max_am_header = attr.max_am_header;
let mut efd: c_int = -1;
let st = sys::ucp_worker_get_efd(worker, &mut efd);
if st != sys::ucs_status_t_UCS_OK {
sys::ucp_worker_destroy(worker);
sys::ucp_cleanup(context);
anyhow::bail!("ucp_worker_get_efd: {}", status_string(st));
}
shared.doorbell.install(worker);
Ok((
WorkerState {
context,
worker,
efd,
eps: HashMap::new(),
err_events: Arc::new(Mutex::new(Vec::new())),
regions: HashMap::new(),
rma_ops: HashMap::new(),
next_op_id: 1,
rma_completions: Arc::new(Mutex::new(Vec::new())),
seen_reg_epoch: shared.reg_epoch.load(Ordering::Acquire),
shared: Arc::clone(shared),
config: config.clone(),
_recv_args: recv_args,
parked_for_close: Vec::new(),
pending_closes: Vec::new(),
now: Instant::now(),
next_ep_scan: Instant::now(),
seen_reply_eps: 0,
reply_ep_scratch: Vec::with_capacity(REPLY_EP_SLOTS),
},
StartupOut {
worker_addr,
max_am_header,
},
))
}
}
fn drain_ring(
state: &mut WorkerState,
ring_rx: &flume::Receiver<Cmd>,
budget: usize,
last_activity: &mut Instant,
) -> (bool, bool) {
let mut drained = 0;
while drained < budget {
match ring_rx.try_recv() {
Ok(cmd) => {
drained += 1;
*last_activity = Instant::now();
let cont = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
state.handle_cmd(cmd)
}))
.unwrap_or_else(|_| {
warn!("ucx: panic while handling a command (continuing)");
true
});
if !cont {
return (false, false);
}
}
Err(flume::TryRecvError::Empty) => return (true, true),
Err(flume::TryRecvError::Disconnected) => return (true, false),
}
}
(false, true)
}
fn run_loop(mut state: WorkerState, ring_rx: flume::Receiver<Cmd>) {
const DRAIN_BUDGET: usize = 64;
const PARK_MS: c_int = 100;
let spin_window = Duration::from_micros(state.config.spin_us);
let flush_budget = state.config.channel_capacity + DRAIN_BUDGET;
let mut last_activity = Instant::now();
'outer: loop {
if state.shared.shutdown_requested.load(Ordering::Acquire) {
break 'outer;
}
state.now = Instant::now();
let (_, keep_running) = drain_ring(&mut state, &ring_rx, DRAIN_BUDGET, &mut last_activity);
if !keep_running {
break 'outer;
}
while unsafe { sys::ucp_worker_progress(state.worker) } != 0 {
last_activity = Instant::now();
}
let (observed_empty, keep_running) =
drain_ring(&mut state, &ring_rx, flush_budget, &mut last_activity);
if !keep_running {
break 'outer;
}
state.drain_rma_completions();
if observed_empty {
state.close_parked();
state.revalidate_eps();
state.reap_failed_eps();
state.stamp_inbound_use();
state.reap_idle_eps();
}
state.poll_pending_closes();
if !ring_rx.is_empty() {
continue;
}
if last_activity.elapsed() < spin_window {
std::hint::spin_loop();
continue;
}
state.shared.doorbell.arm();
if !ring_rx.is_empty() {
state.shared.doorbell.disarm();
continue;
}
let st = unsafe { sys::ucp_worker_arm(state.worker) };
if st == sys::ucs_status_t_UCS_ERR_BUSY {
state.shared.doorbell.disarm();
continue;
}
if st != sys::ucs_status_t_UCS_OK {
warn!("ucx: ucp_worker_arm: {}", status_string(st));
state.shared.doorbell.disarm();
continue;
}
let mut pfd = nix::libc::pollfd {
fd: state.efd,
events: nix::libc::POLLIN,
revents: 0,
};
unsafe { nix::libc::poll(&mut pfd, 1, PARK_MS) };
state.shared.doorbell.disarm();
last_activity = Instant::now();
}
state.teardown(ring_rx);
}
impl WorkerState {
fn handle_cmd(&mut self, cmd: Cmd) -> bool {
match cmd {
Cmd::Send(task) => {
match self.ensure_ep(task.peer) {
Ok(ep) => {
let kind = task.msg_type.as_u8();
let reply = matches!(task.msg_type, MessageType::Message);
let op = Arc::new(OpState {
kind: OpKind::Frame {
header: task.header.clone(),
payload: task.payload.clone(),
on_error: task.on_error,
},
inflight: Arc::clone(&self.shared.inflight_ops),
});
self.post_am(ep, kind, task.header, task.payload, reply, op);
}
Err(e) => task.fail(format!("ucx endpoint unavailable: {e}")),
}
}
Cmd::Ping { peer, token } => {
match self.ensure_ep(peer) {
Ok(ep) => {
let header = Bytes::copy_from_slice(&token.to_le_bytes());
let op = Arc::new(OpState {
kind: OpKind::Control {
_hold: header.clone(),
},
inflight: Arc::clone(&self.shared.inflight_ops),
});
self.post_am(ep, AM_KIND_PING, header, Bytes::new(), true, op);
}
Err(_) => {
self.shared.pending_pings.remove(&token);
}
}
}
Cmd::PongTo { reply_ep, token } => {
let header = Bytes::copy_from_slice(&token.to_le_bytes());
let op = Arc::new(OpState {
kind: OpKind::Control {
_hold: header.clone(),
},
inflight: Arc::clone(&self.shared.inflight_ops),
});
self.post_am(
reply_ep as sys::ucp_ep_h,
AM_KIND_PONG,
header,
Bytes::new(),
false,
op,
);
}
Cmd::ShuttingDownTo { reply_ep, header } => {
let op = Arc::new(OpState {
kind: OpKind::Control {
_hold: header.clone(),
},
inflight: Arc::clone(&self.shared.inflight_ops),
});
self.post_am(
reply_ep as sys::ucp_ep_h,
MessageType::ShuttingDown.as_u8(),
header,
Bytes::new(),
false,
op,
);
}
Cmd::MapRegion {
ptr,
len,
region_id,
reply,
} => {
if let Err(Ok(orphan)) = reply.send(self.map_region(ptr, len, region_id)) {
debug!(
"ucx: rolling back region {} (map_region caller went away)",
orphan.region_id
);
if let Some(entry) = self.regions.remove(&orphan.region_id) {
let _ = self.unmap_entry(entry);
}
}
}
Cmd::UnmapRegion { region_id, reply } => self.unmap_region(region_id, reply),
Cmd::RmaGet { req, reply } => self.rma_get(req, reply),
Cmd::EnsureEp { peer } => {
if let Err(e) = self.ensure_ep(peer) {
debug!("ucx: eager endpoint to {peer} not established: {e}");
}
}
Cmd::Shutdown => return false,
}
true
}
fn post_am(
&mut self,
ep: sys::ucp_ep_h,
kind: u8,
header: Bytes,
payload: Bytes,
reply: bool,
op: Arc<OpState>,
) {
self.shared.inflight_ops.fetch_add(1, Ordering::AcqRel);
let user_data = Arc::into_raw(op) as *mut c_void;
let mut param: sys::ucp_request_param_t = unsafe { MaybeUninit::zeroed().assume_init() };
param.op_attr_mask = sys::ucp_op_attr_t_UCP_OP_ATTR_FIELD_CALLBACK
| sys::ucp_op_attr_t_UCP_OP_ATTR_FIELD_USER_DATA
| sys::ucp_op_attr_t_UCP_OP_ATTR_FIELD_FLAGS;
param.flags = sys::ucp_send_am_flags_UCP_AM_SEND_FLAG_EAGER
| sys::ucp_send_am_flags_UCP_AM_SEND_FLAG_COPY_HEADER
| if reply {
sys::ucp_send_am_flags_UCP_AM_SEND_FLAG_REPLY
} else {
0
};
param.cb.send = Some(send_trampoline);
param.user_data = user_data;
let ptr = unsafe {
sys::ucp_am_send_nbx(
ep,
(AM_ID_BASE as u32) + kind as u32,
header.as_ptr() as *const c_void,
header.len(),
payload.as_ptr() as *const c_void,
payload.len(),
¶m,
)
};
match decode_status_ptr(ptr) {
Ok(Some(_request)) => {
}
Ok(None) => {
let state = unsafe { Arc::from_raw(user_data as *const OpState) };
state.inflight.fetch_sub(1, Ordering::AcqRel);
drop(state);
}
Err(status) => {
let state = unsafe { Arc::from_raw(user_data as *const OpState) };
state.inflight.fetch_sub(1, Ordering::AcqRel);
let _ = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
state.complete(status)
}));
}
}
}
fn ensure_ep(&mut self, peer: InstanceId) -> anyhow::Result<sys::ucp_ep_h> {
let blob = self
.shared
.peers
.get(&peer)
.map(|e| e.value().clone())
.ok_or_else(|| anyhow::anyhow!("peer {peer} not registered"))?;
if let Some(entry) = self.eps.get_mut(&peer) {
if entry.incarnation == blob.incarnation {
entry.last_used = self.now;
return Ok(entry.ep);
}
if let Some(old) = self.eps.remove(&peer) {
debug!("ucx: peer {peer} re-registered; replacing endpoint");
self.parked_for_close.push(old);
}
}
let err_arg = Box::into_raw(Box::new(ErrArg {
peer,
failed: Arc::clone(&self.shared.failed_peers),
err_events: Arc::clone(&self.err_events),
}));
let ep = unsafe {
let mut params: sys::ucp_ep_params_t = MaybeUninit::zeroed().assume_init();
params.field_mask = (sys::ucp_ep_params_field_UCP_EP_PARAM_FIELD_REMOTE_ADDRESS
| sys::ucp_ep_params_field_UCP_EP_PARAM_FIELD_ERR_HANDLING_MODE
| sys::ucp_ep_params_field_UCP_EP_PARAM_FIELD_ERR_HANDLER)
as u64;
params.address = blob.worker_addr.as_ptr() as *const sys::ucp_address_t;
params.err_mode = sys::ucp_err_handling_mode_t_UCP_ERR_HANDLING_MODE_PEER;
params.err_handler = sys::ucp_err_handler_t {
cb: Some(err_trampoline),
arg: err_arg as *mut c_void,
};
let mut ep: sys::ucp_ep_h = std::ptr::null_mut();
let st = sys::ucp_ep_create(self.worker, ¶ms, &mut ep);
if st != sys::ucs_status_t_UCS_OK {
drop(Box::from_raw(err_arg));
anyhow::bail!("ucp_ep_create: {}", status_string(st));
}
ep
};
self.shared.failed_peers.remove(&peer);
self.shared.eps_open.fetch_add(1, Ordering::Relaxed);
self.eps.insert(
peer,
EpEntry {
ep,
err_arg,
incarnation: blob.incarnation,
last_used: self.now,
},
);
debug!("ucx: created endpoint to {peer}");
Ok(ep)
}
fn close_parked(&mut self) {
for entry in std::mem::take(&mut self.parked_for_close) {
self.close_ep(entry, true);
}
}
fn revalidate_eps(&mut self) {
let epoch = self.shared.reg_epoch.load(Ordering::Acquire);
if epoch == self.seen_reg_epoch {
return;
}
self.seen_reg_epoch = epoch;
let stale: Vec<InstanceId> = self
.eps
.iter()
.filter(|(peer, entry)| match self.shared.peers.get(peer) {
Some(blob) => blob.value().incarnation != entry.incarnation,
None => true,
})
.map(|(peer, _)| *peer)
.collect();
for peer in stale {
if let Some(entry) = self.eps.remove(&peer) {
debug!("ucx: dropping stale endpoint to re-registered peer {peer}");
self.close_ep(entry, true);
}
}
}
fn reap_failed_eps(&mut self) {
let peers: Vec<InstanceId> = {
let mut guard = self.err_events.lock().unwrap_or_else(|e| e.into_inner());
std::mem::take(&mut *guard)
};
for peer in peers {
if let Some(entry) = self.eps.remove(&peer) {
self.close_ep(entry, true);
}
}
}
fn stamp_inbound_use(&mut self) {
if self.config.ep_idle_timeout.is_none() {
return;
}
self.reply_ep_scratch.clear();
self.shared
.reply_eps
.drain_into(&mut self.seen_reply_eps, &mut self.reply_ep_scratch);
if self.reply_ep_scratch.is_empty() {
return;
}
let now = self.now;
let (mut stamped, mut unmatched) = (0u64, 0u64);
for &seen in &self.reply_ep_scratch {
let mut hit = false;
for entry in self.eps.values_mut() {
if entry.ep as usize == seen {
entry.last_used = now;
hit = true;
}
}
if hit {
stamped += 1;
} else {
unmatched += 1;
}
}
if stamped != 0 {
self.shared
.eps_stamped_inbound
.fetch_add(stamped, Ordering::Relaxed);
}
if unmatched != 0 {
self.shared
.eps_inbound_unmatched
.fetch_add(unmatched, Ordering::Relaxed);
}
}
fn reap_idle_eps(&mut self) {
let Some(timeout) = self.config.ep_idle_timeout else {
return;
};
if self.now < self.next_ep_scan {
return;
}
self.next_ep_scan = self.now + ep_scan_period(timeout);
if self.eps.is_empty() {
return;
}
let idle: Vec<InstanceId> = self
.eps
.iter()
.filter(|(peer, entry)| {
self.now.saturating_duration_since(entry.last_used) > timeout
&& !self.rma_ops.values().any(|op| op.peer == **peer)
})
.map(|(peer, _)| *peer)
.collect();
for peer in idle {
if let Some(entry) = self.eps.remove(&peer) {
debug!("ucx: closing endpoint to {peer} after {timeout:?} idle");
self.close_ep(entry, true);
self.shared.eps_closed_idle.fetch_add(1, Ordering::Relaxed);
}
}
}
fn close_ep(&mut self, entry: EpEntry, force: bool) {
self.close_ep_raw(entry, force)
}
fn close_ep_raw(&mut self, entry: EpEntry, force: bool) {
self.shared.eps_open.fetch_sub(1, Ordering::Relaxed);
unsafe {
let mut param: sys::ucp_request_param_t = MaybeUninit::zeroed().assume_init();
if force {
param.op_attr_mask = sys::ucp_op_attr_t_UCP_OP_ATTR_FIELD_FLAGS;
param.flags = sys::ucp_ep_close_flags_t_UCP_EP_CLOSE_FLAG_FORCE;
}
let ptr = sys::ucp_ep_close_nbx(entry.ep, ¶m);
if let Ok(Some(request)) = decode_status_ptr(ptr) {
self.pending_closes.push(request);
}
drop(Box::from_raw(entry.err_arg));
}
}
fn poll_pending_closes(&mut self) {
if self.pending_closes.is_empty() {
return;
}
self.pending_closes.retain(|req| {
let st = unsafe { sys::ucp_request_check_status(*req) };
if st == sys::ucs_status_t_UCS_INPROGRESS {
true
} else {
unsafe { sys::ucp_request_free(*req) };
false
}
});
}
fn map_region(
&mut self,
ptr: usize,
len: usize,
region_id: u64,
) -> Result<MappedRegion, RmaError> {
if self.shared.shutdown_requested.load(Ordering::Acquire) {
return Err(RmaError::ShuttingDown);
}
if ptr == 0 || len == 0 || ptr.checked_add(len).is_none() {
return Err(RmaError::OutOfRange);
}
let memh = unsafe {
let mut params: sys::ucp_mem_map_params_t = MaybeUninit::zeroed().assume_init();
params.field_mask = (sys::ucp_mem_map_params_field_UCP_MEM_MAP_PARAM_FIELD_ADDRESS
| sys::ucp_mem_map_params_field_UCP_MEM_MAP_PARAM_FIELD_LENGTH)
as u64;
params.address = ptr as *mut c_void;
params.length = len;
let mut memh: sys::ucp_mem_h = std::ptr::null_mut();
let st = sys::ucp_mem_map(self.context, ¶ms, &mut memh);
if st != sys::ucs_status_t_UCS_OK {
return Err(RmaError::Ucx {
status_name: status_string(st),
});
}
memh
};
match self.describe_region(memh, ptr as u64, len as u64, region_id) {
Ok(region) => Ok(region),
Err(e) => {
unsafe { sys::ucp_mem_unmap(self.context, memh) };
Err(e)
}
}
}
fn describe_region(
&mut self,
memh: sys::ucp_mem_h,
requested_addr: u64,
requested_len: u64,
region_id: u64,
) -> Result<MappedRegion, RmaError> {
let (effective_addr, effective_len) = unsafe {
let mut attr: sys::ucp_mem_attr_t = MaybeUninit::zeroed().assume_init();
attr.field_mask = (sys::ucp_mem_attr_field_UCP_MEM_ATTR_FIELD_ADDRESS
| sys::ucp_mem_attr_field_UCP_MEM_ATTR_FIELD_LENGTH)
as u64;
let st = sys::ucp_mem_query(memh, &mut attr);
if st != sys::ucs_status_t_UCS_OK {
return Err(RmaError::Ucx {
status_name: status_string(st),
});
}
(attr.address as u64, attr.length as u64)
};
let requested_end = requested_addr + requested_len;
let effective_end = effective_addr
.checked_add(effective_len)
.ok_or(RmaError::OutOfRange)?;
if effective_addr > requested_addr || effective_end < requested_end {
warn!(
"ucx: ucp_mem_query reported [{effective_addr:#x}, {effective_end:#x}) which does \
not contain the mapped range [{requested_addr:#x}, {requested_end:#x})"
);
return Err(RmaError::OutOfRange);
}
let packed_rkey = unsafe {
let mut buf: *mut c_void = std::ptr::null_mut();
let mut size: usize = 0;
let st = sys::ucp_rkey_pack(self.context, memh, &mut buf, &mut size);
if st != sys::ucs_status_t_UCS_OK {
return Err(RmaError::Ucx {
status_name: status_string(st),
});
}
let packed = if buf.is_null() || size == 0 {
Bytes::new()
} else {
Bytes::copy_from_slice(std::slice::from_raw_parts(buf as *const u8, size))
};
if !buf.is_null() {
sys::ucp_rkey_buffer_release(buf);
}
packed
};
validate_packed_rkey(&packed_rkey)?;
self.regions.insert(
region_id,
RegionEntry {
memh,
requested_addr,
requested_len,
effective_addr,
effective_len,
inflight: 0,
pending_unmap: Vec::new(),
},
);
self.shared.live_regions.fetch_add(1, Ordering::Relaxed);
Ok(MappedRegion {
region_id,
effective_addr,
effective_len,
packed_rkey,
})
}
fn unmap_region(
&mut self,
region_id: u64,
reply: tokio::sync::oneshot::Sender<Result<(), RmaError>>,
) {
let Some(entry) = self.regions.get_mut(®ion_id) else {
let _ = reply.send(Ok(()));
return;
};
entry.pending_unmap.push(reply);
if entry.inflight > 0 {
return;
}
let entry = self
.regions
.remove(®ion_id)
.expect("looked up immediately above");
self.finish_unmap(entry);
}
fn finish_unmap(&self, mut entry: RegionEntry) {
let waiters = std::mem::take(&mut entry.pending_unmap);
let result = self.unmap_entry(entry);
for waiter in waiters {
let _ = waiter.send(result.clone());
}
}
fn unmap_entry(&self, entry: RegionEntry) -> Result<(), RmaError> {
let st = unsafe { sys::ucp_mem_unmap(self.context, entry.memh) };
self.shared.live_regions.fetch_sub(1, Ordering::Relaxed);
if st == sys::ucs_status_t_UCS_OK {
Ok(())
} else {
Err(RmaError::Ucx {
status_name: status_string(st),
})
}
}
fn rma_get(
&mut self,
req: RmaGetRequest,
reply: tokio::sync::oneshot::Sender<Result<(), RmaError>>,
) {
match self.prepare_get(&req) {
Err(e) => {
let _ = reply.send(Err(e));
}
Ok(None) => {
let _ = reply.send(Ok(()));
}
Ok(Some(prepared)) => self.post_get(prepared, reply),
}
}
fn prepare_get(&mut self, req: &RmaGetRequest) -> Result<Option<PreparedGet>, RmaError> {
if self.shared.shutdown_requested.load(Ordering::Acquire) {
return Err(RmaError::ShuttingDown);
}
let (memh, requested_addr, requested_len, effective_addr, effective_len, unmapping) = {
let entry = self
.regions
.get(&req.local_region)
.ok_or(RmaError::RegionNotFound)?;
(
entry.memh,
entry.requested_addr,
entry.requested_len,
entry.effective_addr,
entry.effective_len,
!entry.pending_unmap.is_empty(),
)
};
if unmapping {
return Err(RmaError::RegionNotFound);
}
let end = req
.local_offset
.checked_add(req.len)
.ok_or(RmaError::OutOfRange)?;
if end > requested_len {
return Err(RmaError::OutOfRange);
}
let local_addr = requested_addr
.checked_add(req.local_offset)
.ok_or(RmaError::OutOfRange)?;
let local_end = local_addr
.checked_add(req.len)
.ok_or(RmaError::OutOfRange)?;
let effective_end = effective_addr
.checked_add(effective_len)
.ok_or(RmaError::OutOfRange)?;
if local_addr < effective_addr || local_end > effective_end {
return Err(RmaError::OutOfRange);
}
if !self.shared.peers.contains_key(&req.peer) {
return Err(RmaError::PeerNotRegistered(req.peer));
}
if req.len == 0 {
return Ok(None);
}
validate_packed_rkey(&req.packed_rkey)?;
let ep = self
.ensure_ep(req.peer)
.map_err(|e| RmaError::EndpointUnavailable(e.to_string()))?;
let mut unpack_buf = [0xFFu8; MAX_PACKED_RKEY + RKEY_UNPACK_PAD];
unpack_buf[..req.packed_rkey.len()].copy_from_slice(&req.packed_rkey);
let rkey = unsafe {
let mut rkey: sys::ucp_rkey_h = std::ptr::null_mut();
let st = sys::ucp_ep_rkey_unpack(ep, unpack_buf.as_ptr() as *const c_void, &mut rkey);
if st != sys::ucs_status_t_UCS_OK {
return Err(RmaError::Ucx {
status_name: status_string(st),
});
}
rkey
};
self.shared.live_rkeys.fetch_add(1, Ordering::Relaxed);
Ok(Some(PreparedGet {
ep,
rkey,
memh,
local_addr,
remote_addr: req.remote_addr,
len: req.len as usize,
region_id: req.local_region,
peer: req.peer,
}))
}
fn post_get(
&mut self,
prepared: PreparedGet,
reply: tokio::sync::oneshot::Sender<Result<(), RmaError>>,
) {
self.shared.inflight_ops.fetch_add(1, Ordering::AcqRel);
if let Some(entry) = self.regions.get_mut(&prepared.region_id) {
entry.inflight += 1;
}
let op_id = self.next_op_id;
self.next_op_id += 1;
let op = Arc::new(RmaOpState {
rkey: Mutex::new(Some(prepared.rkey as usize)),
peer: prepared.peer,
region_id: prepared.region_id,
op_id,
reply: Mutex::new(Some(reply)),
inflight: Arc::clone(&self.shared.inflight_ops),
live_rkeys: Arc::clone(&self.shared.live_rkeys),
rma_completions: Arc::clone(&self.rma_completions),
});
self.rma_ops.insert(op_id, Arc::clone(&op));
let user_data = Arc::into_raw(op) as *mut c_void;
let mut param: sys::ucp_request_param_t = unsafe { MaybeUninit::zeroed().assume_init() };
param.op_attr_mask = sys::ucp_op_attr_t_UCP_OP_ATTR_FIELD_CALLBACK
| sys::ucp_op_attr_t_UCP_OP_ATTR_FIELD_USER_DATA
| sys::ucp_op_attr_t_UCP_OP_ATTR_FIELD_MEMH
| sys::ucp_op_attr_t_UCP_OP_ATTR_FLAG_NO_IMM_CMPL;
param.cb.send = Some(rma_trampoline);
param.user_data = user_data;
param.memh = prepared.memh;
let ptr = unsafe {
sys::ucp_get_nbx(
prepared.ep,
prepared.local_addr as *mut c_void,
prepared.len,
prepared.remote_addr,
prepared.rkey,
¶m,
)
};
match decode_status_ptr(ptr) {
Ok(Some(_request)) => {
}
Ok(None) => {
let state = unsafe { Arc::from_raw(user_data as *const RmaOpState) };
state.inflight.fetch_sub(1, Ordering::AcqRel);
state.complete(sys::ucs_status_t_UCS_OK);
}
Err(status) => {
let state = unsafe { Arc::from_raw(user_data as *const RmaOpState) };
state.inflight.fetch_sub(1, Ordering::AcqRel);
state.complete(status);
}
}
}
fn drain_rma_completions(&mut self) {
let completed: RmaCompletions = {
let mut guard = self
.rma_completions
.lock()
.unwrap_or_else(|e| e.into_inner());
if guard.is_empty() {
return;
}
std::mem::take(&mut *guard)
};
for (region_id, op_id) in completed {
self.rma_ops.remove(&op_id);
let released = match self.regions.get_mut(®ion_id) {
Some(entry) => {
entry.inflight = entry.inflight.saturating_sub(1);
entry.inflight == 0 && !entry.pending_unmap.is_empty()
}
None => false,
};
if released && let Some(entry) = self.regions.remove(®ion_id) {
self.finish_unmap(entry);
}
}
}
fn unmap_idle_regions(&mut self) {
let idle: Vec<u64> = self
.regions
.iter()
.filter(|(_, entry)| entry.inflight == 0)
.map(|(id, _)| *id)
.collect();
for region_id in idle {
if let Some(entry) = self.regions.remove(®ion_id) {
self.finish_unmap(entry);
}
}
}
fn force_unmap_regions(&mut self) {
let remaining: Vec<u64> = self.regions.keys().copied().collect();
for region_id in remaining {
let Some(mut entry) = self.regions.remove(®ion_id) else {
continue;
};
let inflight = entry.inflight;
if inflight == 0 {
self.finish_unmap(entry);
continue;
}
warn!("ucx: force-unmapping region {region_id} with {inflight} rma op(s) in flight");
let waiters = std::mem::take(&mut entry.pending_unmap);
let _ = self.unmap_entry(entry);
for waiter in waiters {
let _ = waiter.send(Err(RmaError::ShuttingDown));
}
}
}
fn abandon_rma_ops(&mut self) {
for (op_id, op) in std::mem::take(&mut self.rma_ops) {
debug!("ucx: abandoning rma op {op_id} at teardown");
op.resolve(Err(RmaError::ShuttingDown));
}
}
fn teardown(mut self, ring_rx: flume::Receiver<Cmd>) {
debug!("ucx: progress thread tearing down");
for pass in 0..2 {
while let Ok(cmd) = ring_rx.try_recv() {
cmd.refuse_for_shutdown();
}
if pass == 0 {
std::thread::sleep(Duration::from_millis(1));
}
}
drop(ring_rx);
self.drain_rma_completions();
self.unmap_idle_regions();
let entries: Vec<EpEntry> = {
let eps = std::mem::take(&mut self.eps);
let mut all: Vec<EpEntry> = eps.into_values().collect();
all.append(&mut self.parked_for_close);
all
};
let mut pending: Vec<(EpEntry, sys::ucs_status_ptr_t)> = Vec::new();
for entry in entries {
let ptr = unsafe {
let param: sys::ucp_request_param_t = MaybeUninit::zeroed().assume_init();
sys::ucp_ep_close_nbx(entry.ep, ¶m)
};
match decode_status_ptr(ptr) {
Ok(Some(req)) => pending.push((entry, req)),
_ => {
self.shared.eps_open.fetch_sub(1, Ordering::Relaxed);
unsafe { drop(Box::from_raw(entry.err_arg)) }
}
}
}
let deadline = Instant::now() + Duration::from_secs(1);
while !pending.is_empty() && Instant::now() < deadline {
unsafe { sys::ucp_worker_progress(self.worker) };
self.drain_rma_completions();
pending.retain(|(entry, req)| {
let st = unsafe { sys::ucp_request_check_status(*req) };
if st == sys::ucs_status_t_UCS_INPROGRESS {
true
} else {
self.shared.eps_open.fetch_sub(1, Ordering::Relaxed);
unsafe {
sys::ucp_request_free(*req);
drop(Box::from_raw(entry.err_arg));
}
false
}
});
}
self.unmap_idle_regions();
for (entry, req) in pending {
unsafe { sys::ucp_request_free(req) };
self.close_ep_raw(entry, true);
}
let deadline = Instant::now() + Duration::from_secs(1);
while self.shared.inflight_ops.load(Ordering::Acquire) > 0 && Instant::now() < deadline {
unsafe { sys::ucp_worker_progress(self.worker) };
self.drain_rma_completions();
}
let leaked = self.shared.inflight_ops.load(Ordering::Acquire);
if leaked > 0 {
warn!("ucx: {leaked} operation(s) still in flight at teardown");
}
self.drain_rma_completions();
self.abandon_rma_ops();
self.force_unmap_regions();
let deadline = Instant::now() + Duration::from_millis(500);
while !self.pending_closes.is_empty() && Instant::now() < deadline {
unsafe { sys::ucp_worker_progress(self.worker) };
self.poll_pending_closes();
}
for req in self.pending_closes.drain(..) {
unsafe { sys::ucp_request_free(req) };
}
self.shared.doorbell.retire();
unsafe {
sys::ucp_worker_destroy(self.worker);
sys::ucp_cleanup(self.context);
}
debug!("ucx: progress thread exited");
}
}
pub(crate) type StartupSlot = OnceLock<StartupOut>;