use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::{Arc, Mutex, OnceLock};
use std::time::Duration;
use anyhow::Result;
use bytes::Bytes;
use dashmap::DashMap;
use futures::future::BoxFuture;
use tracing::{debug, info, warn};
use crate::transports::transport::{
AdmissionError, AdmissionGate, HealthCheckError, SendOutcome, ShutdownState, TransportError,
TransportErrorHandler,
};
use velo_ext::{
InstanceId, MessageType, PeerInfo, Transport, TransportAdapter, TransportKey, WorkerAddress,
};
use super::address::{AM_ID_BASE, BLOB_VERSION, UcxEndpoint};
use super::rma::{RdmaEndpoint, RmaState};
use super::worker::{Cmd, Doorbell, SendTask, StartupSlot, WorkerArgs, WorkerShared, worker_main};
const MAX_EAGER: u32 = 16 * 1024 * 1024;
#[derive(Debug, Clone)]
pub struct UcxConfig {
pub eager_max: u32,
pub spin_us: u64,
pub channel_capacity: usize,
pub tls: Option<String>,
pub net_devices: Option<String>,
pub ep_idle_timeout: Option<Duration>,
pub eager_endpoints: bool,
}
const MIN_EP_IDLE_TIMEOUT: Duration = Duration::from_millis(500);
impl Default for UcxConfig {
fn default() -> Self {
Self {
eager_max: 1 << 20,
spin_us: 20,
channel_capacity: 1024,
tls: None,
net_devices: None,
ep_idle_timeout: None,
eager_endpoints: false,
}
}
}
#[derive(Clone)]
struct ConnHandle {
gate: AdmissionGate<Cmd>,
}
pub struct UcxTransport {
key: TransportKey,
config: UcxConfig,
incarnation: u64,
ring_tx: flume::Sender<Cmd>,
ring_rx: Mutex<Option<flume::Receiver<Cmd>>>,
shared: Arc<WorkerShared>,
local_address: OnceLock<WorkerAddress>,
startup: StartupSlot,
connections: DashMap<InstanceId, ConnHandle>,
runtime: OnceLock<tokio::runtime::Handle>,
shutdown_state: OnceLock<ShutdownState>,
join: Mutex<Option<std::thread::JoinHandle<()>>>,
ping_token: AtomicU64,
metrics: OnceLock<Arc<dyn velo_ext::TransportObservability>>,
rma: Arc<RmaState>,
}
impl UcxTransport {
fn new(key: TransportKey, config: UcxConfig) -> Self {
let (ring_tx, ring_rx) = flume::bounded(config.channel_capacity);
let shared = Arc::new(WorkerShared {
ring_tx: ring_tx.clone(),
doorbell: Arc::new(Doorbell::new()),
peers: Arc::new(DashMap::new()),
pending_pings: Arc::new(DashMap::new()),
failed_peers: Arc::new(DashMap::new()),
inflight_ops: Arc::new(Default::default()),
shutdown_requested: Arc::new(Default::default()),
reg_epoch: Arc::new(Default::default()),
live_regions: Arc::new(Default::default()),
live_rkeys: Arc::new(Default::default()),
eps_open: Arc::new(Default::default()),
eps_closed_idle: Arc::new(Default::default()),
eps_stamped_inbound: Arc::new(Default::default()),
eps_inbound_unmatched: Arc::new(Default::default()),
reply_eps: Arc::new(super::worker::ReplyEpSightings::new()),
});
Self {
key,
config,
incarnation: uuid::Uuid::new_v4().as_u128() as u64,
ring_tx,
ring_rx: Mutex::new(Some(ring_rx)),
shared,
local_address: OnceLock::new(),
startup: OnceLock::new(),
connections: DashMap::new(),
runtime: OnceLock::new(),
shutdown_state: OnceLock::new(),
join: Mutex::new(None),
ping_token: AtomicU64::new(1),
metrics: OnceLock::new(),
rma: Arc::new(RmaState::new()),
}
}
pub(crate) fn rdma_endpoint(&self) -> RdmaEndpoint {
RdmaEndpoint::new(Arc::clone(&self.shared), Arc::clone(&self.rma))
}
#[allow(dead_code)]
pub(crate) fn live_regions(&self) -> usize {
self.shared.live_regions.load(Ordering::SeqCst)
}
#[allow(dead_code)]
pub(crate) fn live_rkeys(&self) -> i64 {
self.shared.live_rkeys.load(Ordering::SeqCst)
}
fn eager_limit(&self, peer: InstanceId) -> Option<usize> {
self.shared
.peers
.get(&peer)
.map(|e| self.config.eager_max.min(e.value().eager_max) as usize)
}
fn get_or_create_connection(&self, peer: InstanceId) -> Result<ConnHandle, TransportError> {
if let Some(handle) = self.connections.get(&peer) {
return Ok(handle.clone());
}
let rt = self.runtime.get().ok_or(TransportError::NotStarted)?;
if !self.shared.peers.contains_key(&peer) {
return Err(TransportError::PeerNotRegistered(peer));
}
let handle = self
.connections
.entry(peer)
.or_insert_with(|| ConnHandle {
gate: AdmissionGate::new(self.ring_tx.clone(), rt.clone()),
})
.clone();
if let Some(m) = self.metrics.get() {
m.set_active_connections(self.connections.len());
}
Ok(handle)
}
fn admit(&self, handle: &ConnHandle, task: SendTask) -> SendOutcome {
match handle.gate.send(Cmd::Send(task)) {
SendOutcome::Admitted => {
self.shared.doorbell.ring();
SendOutcome::Admitted
}
SendOutcome::Pending(admission) => {
if let Some(m) = self.metrics.get() {
m.record_send_backpressure();
}
let doorbell = Arc::clone(&self.shared.doorbell);
SendOutcome::Pending(admission.on_resolved(move |result| {
if result.is_ok() {
doorbell.ring();
}
}))
}
}
}
fn reap_failed_connection(&self, peer: InstanceId) {
if self.shared.failed_peers.remove(&peer).is_some()
&& let Some((_, stale)) = self.connections.remove(&peer)
{
stale.gate.fail_all(AdmissionError::ConnectionReplaced);
if let Some(m) = self.metrics.get() {
m.set_active_connections(self.connections.len());
}
}
}
}
impl Transport for UcxTransport {
fn key(&self) -> TransportKey {
self.key.clone()
}
fn address(&self) -> WorkerAddress {
self.local_address.get().cloned().unwrap_or_else(|| {
crate::transports::address::WorkerAddressBuilder::new()
.build()
.expect("empty WorkerAddress")
})
}
fn register(&self, peer_info: PeerInfo) -> Result<(), TransportError> {
let entry = peer_info
.worker_address()
.get_entry(&self.key)
.map_err(|_| TransportError::NoEndpoint)?
.ok_or(TransportError::NoEndpoint)?;
let endpoint = UcxEndpoint::decode(&entry).map_err(|e| {
debug!("ucx: rejecting peer blob: {e}");
TransportError::InvalidEndpoint
})?;
let peer = peer_info.instance_id();
self.shared.failed_peers.remove(&peer);
self.shared.peers.insert(peer, endpoint);
self.shared.reg_epoch.fetch_add(1, Ordering::AcqRel);
if self.config.eager_endpoints && self.ring_tx.try_send(Cmd::EnsureEp { peer }).is_err() {
debug!("ucx: eager wireup for {peer} skipped (ring full or closed)");
}
self.shared.doorbell.ring();
if let Some(m) = self.metrics.get() {
m.set_registered_peers(self.shared.peers.len());
}
debug!("ucx: registered peer {peer}");
Ok(())
}
fn send_message(
&self,
instance_id: InstanceId,
header: Bytes,
payload: Bytes,
message_type: MessageType,
on_error: Arc<dyn TransportErrorHandler>,
) -> SendOutcome {
let task = SendTask {
peer: instance_id,
msg_type: message_type,
header,
payload,
on_error,
};
if self.runtime.get().is_none() {
task.on_error.on_error(
task.header,
task.payload,
"ucx transport not started".into(),
);
return SendOutcome::Admitted;
}
if let Some(out) = self.startup.get()
&& task.header.len() > out.max_am_header
{
let why = format!(
"header {} bytes exceeds ucx max_am_header {}",
task.header.len(),
out.max_am_header
);
task.on_error.on_error(task.header, task.payload, why);
return SendOutcome::Admitted;
}
match self.eager_limit(instance_id) {
Some(limit) if task.header.len() + task.payload.len() > limit => {
let why = format!(
"frame {} bytes exceeds negotiated ucx eager limit {limit}",
task.header.len() + task.payload.len()
);
task.on_error.on_error(task.header, task.payload, why);
return SendOutcome::Admitted;
}
Some(_) => {}
None => {
task.on_error.on_error(
task.header,
task.payload,
format!("peer not registered: {instance_id}"),
);
return SendOutcome::Admitted;
}
}
self.reap_failed_connection(instance_id);
match self.get_or_create_connection(instance_id) {
Ok(handle) => self.admit(&handle, task),
Err(e) => {
task.fail(format!("ucx connection unavailable: {e}"));
SendOutcome::Admitted
}
}
}
fn max_message_size(&self, target: InstanceId) -> Option<usize> {
self.eager_limit(target)
}
fn start(
&self,
_instance_id: InstanceId,
channels: TransportAdapter,
rt: tokio::runtime::Handle,
) -> BoxFuture<'_, Result<()>> {
self.runtime.set(rt).ok();
self.shutdown_state
.set(channels.shutdown_state.clone())
.ok();
let ring_rx = self
.ring_rx
.lock()
.unwrap_or_else(|e| e.into_inner())
.take();
Box::pin(async move {
let ring_rx =
ring_rx.ok_or_else(|| anyhow::anyhow!("ucx transport already started"))?;
let (startup_tx, startup_rx) = tokio::sync::oneshot::channel();
let args = WorkerArgs {
config: self.config.clone(),
ring_rx,
shared: Arc::clone(&self.shared),
adapter: channels,
startup: startup_tx,
};
let join = std::thread::Builder::new()
.name("velo-ucx-progress".into())
.spawn(move || worker_main(args))?;
*self.join.lock().unwrap_or_else(|e| e.into_inner()) = Some(join);
let out = startup_rx
.await
.map_err(|_| anyhow::anyhow!("ucx progress thread died during startup"))??;
let blob = UcxEndpoint {
v: BLOB_VERSION,
am_id_base: AM_ID_BASE,
eager_max: self.config.eager_max.min(MAX_EAGER),
incarnation: self.incarnation,
worker_addr: out.worker_addr.clone(),
}
.encode()?;
let mut builder = crate::transports::address::WorkerAddressBuilder::new();
builder.add_entry(self.key.clone(), blob)?;
let address = builder.build()?;
info!(
"UCX transport started (worker address {} B, max_am_header {} B)",
out.worker_addr.len(),
out.max_am_header
);
self.startup.set(out).ok();
self.local_address.set(address).ok();
self.rma.mark_started(self.runtime.get().cloned());
Ok(())
})
}
fn begin_drain(&self) {
}
fn shutdown(&self) {
info!("Shutting down UCX transport");
if let Some(state) = self.shutdown_state.get() {
state.teardown_token().cancel();
}
self.shared
.shutdown_requested
.store(true, std::sync::atomic::Ordering::Release);
let _ = self.ring_tx.try_send(Cmd::Shutdown);
self.shared.doorbell.ring_force();
if let Some(join) = self.join.lock().unwrap_or_else(|e| e.into_inner()).take()
&& join.join().is_err()
{
warn!("ucx progress thread panicked during shutdown");
}
for entry in self.connections.iter() {
entry.value().gate.fail_all(AdmissionError::ChannelClosed);
}
self.connections.clear();
if let Some(m) = self.metrics.get() {
m.set_active_connections(0);
}
}
fn set_observability(&self, observability: Arc<dyn velo_ext::TransportObservability>) {
let _ = self.metrics.set(observability);
if let Some(m) = self.metrics.get() {
m.set_registered_peers(self.shared.peers.len());
m.set_active_connections(self.connections.len());
}
}
fn check_health(
&self,
instance_id: InstanceId,
timeout: Duration,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = Result<(), HealthCheckError>> + Send + '_>,
> {
Box::pin(async move {
if self.runtime.get().is_none() {
return Err(HealthCheckError::TransportNotStarted);
}
if !self.shared.peers.contains_key(&instance_id) {
return Err(HealthCheckError::PeerNotRegistered);
}
let connection_existed = self.connections.contains_key(&instance_id)
&& !self.shared.failed_peers.contains_key(&instance_id);
let token = self.ping_token.fetch_add(1, Ordering::Relaxed);
let (tx, rx) = tokio::sync::oneshot::channel();
self.shared.pending_pings.insert(token, tx);
struct PingGuard {
map: Arc<DashMap<u64, tokio::sync::oneshot::Sender<()>>>,
token: u64,
}
impl Drop for PingGuard {
fn drop(&mut self) {
self.map.remove(&self.token);
}
}
let _guard = PingGuard {
map: Arc::clone(&self.shared.pending_pings),
token,
};
let probe = async {
self.ring_tx
.send_async(Cmd::Ping {
peer: instance_id,
token,
})
.await
.map_err(|_| HealthCheckError::ConnectionFailed)?;
self.shared.doorbell.ring();
rx.await.map_err(|_| HealthCheckError::ConnectionFailed)
};
match tokio::time::timeout(timeout, probe).await {
Ok(Ok(())) => {
if connection_existed {
Ok(())
} else {
Err(HealthCheckError::NeverConnected)
}
}
Ok(Err(e)) => Err(e),
Err(_) => {
if self.shared.failed_peers.contains_key(&instance_id) {
Err(HealthCheckError::ConnectionFailed)
} else {
Err(HealthCheckError::Timeout)
}
}
}
})
}
}
pub struct UcxTransportBuilder {
key: Option<TransportKey>,
config: UcxConfig,
}
impl UcxTransportBuilder {
pub fn new() -> Self {
Self {
key: None,
config: UcxConfig::default(),
}
}
pub fn key(mut self, key: TransportKey) -> Self {
self.key = Some(key);
self
}
pub fn eager_max(mut self, bytes: u32) -> Self {
self.config.eager_max = bytes.min(MAX_EAGER);
self
}
pub fn spin_us(mut self, us: u64) -> Self {
self.config.spin_us = us;
self
}
pub fn channel_capacity(mut self, capacity: usize) -> Self {
self.config.channel_capacity = capacity.max(1);
self
}
pub fn tls(mut self, tls: impl Into<String>) -> Self {
self.config.tls = Some(tls.into());
self
}
pub fn net_devices(mut self, devices: impl Into<String>) -> Self {
self.config.net_devices = Some(devices.into());
self
}
pub fn ep_idle_timeout(mut self, timeout: Option<Duration>) -> Self {
self.config.ep_idle_timeout = timeout.map(|t| {
if t < MIN_EP_IDLE_TIMEOUT {
debug!("ucx: ep_idle_timeout {t:?} raised to the {MIN_EP_IDLE_TIMEOUT:?} floor");
MIN_EP_IDLE_TIMEOUT
} else {
t
}
});
self
}
pub fn eager_endpoints(mut self, eager: bool) -> Self {
self.config.eager_endpoints = eager;
self
}
pub fn build(self) -> Result<UcxTransport> {
let key = self.key.unwrap_or_else(|| TransportKey::from("ucx"));
Ok(UcxTransport::new(key, self.config))
}
}
impl Default for UcxTransportBuilder {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
#[path = "tests.rs"]
mod tests;