use anyhow::{Context, Result};
use bytes::Bytes;
use dashmap::DashMap;
use std::net::SocketAddr;
use std::sync::{Arc, Mutex, OnceLock};
use std::time::Duration;
use tokio::net::TcpStream;
use tokio_util::sync::CancellationToken;
use tracing::{debug, error, info, warn};
use crate::transports::transport::{
AdmissionError, AdmissionGate, HealthCheckError, SendOutcome, ShutdownState, TransportError,
TransportErrorHandler,
};
use crate::transports::utils::interfaces::{
InterfaceEndpoint, InterfaceFilter, parse_endpoints, resolve_advertise_endpoints,
select_best_endpoint,
};
use velo_ext::{MessageType, PeerInfo, Transport, TransportAdapter, TransportKey, WorkerAddress};
use super::framing::DEFAULT_MAX_FRAME_SIZE;
use super::listener::TcpListener;
use crate::transports::coalesce::{
Coalescable, WriterFailure, WriterObserver, run_coalescing_writer,
};
use crate::transports::ingress::{DialedReaderContext, run_dialed_reader};
pub struct TcpTransport {
key: TransportKey,
bind_addr: SocketAddr,
local_address: WorkerAddress,
peers: Arc<DashMap<crate::InstanceId, SocketAddr>>,
connections: Arc<DashMap<crate::InstanceId, ConnectionHandle>>,
runtime: OnceLock<tokio::runtime::Handle>,
cancel_token: CancellationToken,
shutdown_state: OnceLock<ShutdownState>,
channel_capacity: usize,
connect_timeout: Duration,
listener: Mutex<Option<std::net::TcpListener>>,
local_interfaces: OnceLock<Vec<InterfaceEndpoint>>,
numa_hint: Option<u32>,
metrics: OnceLock<std::sync::Arc<dyn velo_ext::TransportObservability>>,
shrink_threshold: usize,
dialed_ctx: OnceLock<DialedReaderContext>,
}
#[derive(Clone)]
struct ConnectionHandle {
tx: flume::Sender<SendTask>,
gate: AdmissionGate<SendTask>,
}
impl ConnectionHandle {
fn retire(&self) {
self.gate.fail_all(AdmissionError::ConnectionReplaced);
}
}
struct SendTask {
msg_type: MessageType,
header: Bytes,
payload: Bytes,
on_error: Arc<dyn TransportErrorHandler>,
}
impl SendTask {
fn on_error(self, error: impl Into<String>) {
self.on_error
.on_error(self.header, self.payload, error.into());
}
}
impl TcpTransport {
pub fn new(
bind_addr: SocketAddr,
key: TransportKey,
local_address: WorkerAddress,
channel_capacity: usize,
connect_timeout: Duration,
listener: Option<std::net::TcpListener>,
numa_hint: Option<u32>,
) -> Self {
Self {
key,
bind_addr,
local_address,
peers: Arc::new(DashMap::new()),
connections: Arc::new(DashMap::new()),
runtime: OnceLock::new(),
cancel_token: CancellationToken::new(),
shutdown_state: OnceLock::new(),
channel_capacity,
connect_timeout,
listener: Mutex::new(listener),
local_interfaces: OnceLock::new(),
numa_hint,
metrics: OnceLock::new(),
shrink_threshold: super::listener::default_shrink_threshold(),
dialed_ctx: OnceLock::new(),
}
}
pub fn ensure_connected(&self, instance_id: crate::InstanceId) -> Result<()> {
self.get_or_create_connection(instance_id)?;
Ok(())
}
fn reap_stale_connection(&self, instance_id: crate::InstanceId) {
if let Some((_, stale)) = self
.connections
.remove_if(&instance_id, |_, h| h.tx.is_disconnected())
{
stale.retire();
self.update_connection_gauge();
}
}
fn get_or_create_connection(&self, instance_id: crate::InstanceId) -> Result<ConnectionHandle> {
if let Some(handle) = self.connections.get(&instance_id) {
if !handle.tx.is_disconnected() {
return Ok(handle.clone());
}
drop(handle);
self.reap_stale_connection(instance_id);
}
let rt = self.runtime.get().ok_or(TransportError::NotStarted)?;
let handle = match self.connections.entry(instance_id) {
dashmap::mapref::entry::Entry::Occupied(mut entry) => {
if !entry.get().tx.is_disconnected() {
entry.get().clone()
} else {
entry.get().retire();
let handle = self.create_connection(instance_id, rt)?;
entry.insert(handle.clone());
self.update_connection_gauge();
handle
}
}
dashmap::mapref::entry::Entry::Vacant(entry) => {
let handle = self.create_connection(instance_id, rt)?;
entry.insert(handle.clone());
self.update_connection_gauge();
handle
}
};
Ok(handle)
}
fn create_connection(
&self,
instance_id: crate::InstanceId,
rt: &tokio::runtime::Handle,
) -> Result<ConnectionHandle> {
let addr = *self
.peers
.get(&instance_id)
.ok_or(TransportError::PeerNotRegistered(instance_id))?
.value();
let (tx, rx) = flume::bounded(self.channel_capacity);
let handle = ConnectionHandle {
gate: AdmissionGate::new(tx.clone(), rt.clone()),
tx,
};
rt.spawn(connection_writer_task(
addr,
instance_id,
rx,
WriterTaskContext {
connections: Arc::clone(&self.connections),
cancel_token: self.cancel_token.clone(),
connect_timeout: self.connect_timeout,
reader_ctx: self.dialed_ctx.get().cloned(),
metrics: self.metrics.get().cloned(),
},
));
debug!("Created new connection to {} ({})", instance_id, addr);
Ok(handle)
}
fn update_peer_gauge(&self) {
if let Some(metrics) = self.metrics.get() {
metrics.set_registered_peers(self.peers.len());
}
}
fn update_connection_gauge(&self) {
if let Some(metrics) = self.metrics.get() {
metrics.set_active_connections(self.connections.len());
}
}
fn slow_path_send(&self, instance_id: crate::InstanceId, send_msg: SendTask) -> SendOutcome {
if self.runtime.get().is_none() {
send_msg.on_error("Transport not started");
return SendOutcome::Admitted;
}
let handle = match self.get_or_create_connection(instance_id) {
Ok(h) => h,
Err(e) => {
send_msg.on_error(format!("Failed to create connection: {}", e));
return SendOutcome::Admitted;
}
};
self.admit(&handle, send_msg)
}
fn admit(&self, handle: &ConnectionHandle, send_msg: SendTask) -> SendOutcome {
let outcome = handle.gate.send(send_msg);
if let Some(m) = self.metrics.get()
&& !outcome.is_admitted()
{
m.record_send_backpressure();
}
outcome
}
}
impl Transport for TcpTransport {
fn key(&self) -> TransportKey {
self.key.clone()
}
fn address(&self) -> WorkerAddress {
self.local_address.clone()
}
fn max_message_size(&self, _target: crate::InstanceId) -> Option<usize> {
Some(DEFAULT_MAX_FRAME_SIZE as usize)
}
fn register(&self, peer_info: PeerInfo) -> Result<(), TransportError> {
let endpoint = peer_info
.worker_address()
.get_entry(&self.key)
.map_err(|_| TransportError::NoEndpoint)?
.ok_or(TransportError::NoEndpoint)?;
let remote_endpoints = parse_endpoints(&endpoint).map_err(|e| {
error!("Failed to parse TCP endpoint: {}", e);
TransportError::InvalidEndpoint
})?;
let local = self.local_interfaces.get_or_init(|| {
resolve_advertise_endpoints(self.bind_addr, &InterfaceFilter::All).unwrap_or_default()
});
let addr = select_best_endpoint(&remote_endpoints, local, self.numa_hint)
.ok_or(TransportError::InvalidEndpoint)?;
self.peers.insert(peer_info.instance_id(), addr);
self.update_peer_gauge();
debug!("Registered peer {} at {}", peer_info.instance_id(), addr);
Ok(())
}
#[inline]
fn send_message(
&self,
instance_id: crate::InstanceId,
header: Bytes,
payload: Bytes,
message_type: MessageType,
on_error: std::sync::Arc<dyn TransportErrorHandler>,
) -> SendOutcome {
let send_msg = SendTask {
msg_type: message_type,
header,
payload,
on_error,
};
if let Some(handle) = self.connections.get(&instance_id) {
let live = (!handle.tx.is_disconnected()).then(|| handle.clone());
drop(handle);
match live {
Some(handle) => return self.admit(&handle, send_msg),
None => self.reap_stale_connection(instance_id),
}
}
self.slow_path_send(instance_id, send_msg)
}
fn start(
&self,
_instance_id: crate::InstanceId,
channels: TransportAdapter,
rt: tokio::runtime::Handle,
) -> futures::future::BoxFuture<'_, anyhow::Result<()>> {
self.dialed_ctx
.set(DialedReaderContext {
adapter: channels.clone(),
error_handler: std::sync::Arc::new(DialedReaderErrorHandler),
transport_key: self.key.as_str().to_string(),
shrink_threshold: self.shrink_threshold,
})
.ok();
self.runtime.set(rt.clone()).ok();
self.shutdown_state
.set(channels.shutdown_state.clone())
.ok();
let bind_addr = self.bind_addr;
let shutdown_state = channels.shutdown_state.clone();
let listener = self
.listener
.lock()
.expect("Listener mutex poisoned")
.take();
Box::pin(async move {
struct DefaultErrorHandler;
impl TransportErrorHandler for DefaultErrorHandler {
fn on_error(&self, _header: Bytes, _payload: Bytes, error: String) {
warn!("Transport error: {}", error);
}
}
let tcp_listener = TcpListener::builder()
.bind_addr(bind_addr)
.adapter(channels)
.error_handler(std::sync::Arc::new(DefaultErrorHandler))
.shutdown_state(shutdown_state)
.listener(listener)
.transport_key(self.key.as_str())
.metrics(self.metrics.get().cloned())
.shrink_threshold(self.shrink_threshold)
.build()?;
rt.spawn(async move {
if let Err(e) = tcp_listener.serve().await {
error!("TCP listener error: {}", e);
}
});
info!("TCP transport started on {}", bind_addr);
Ok(())
})
}
fn begin_drain(&self) {
}
fn shutdown(&self) {
info!("Shutting down TCP transport");
if let Some(state) = self.shutdown_state.get() {
state.teardown_token().cancel();
}
self.cancel_token.cancel();
self.connections.clear();
self.update_connection_gauge();
}
fn set_observability(
&self,
observability: std::sync::Arc<dyn velo_ext::TransportObservability>,
) {
let _ = self.metrics.set(observability);
self.update_peer_gauge();
self.update_connection_gauge();
}
fn check_health(
&self,
instance_id: crate::InstanceId,
timeout: Duration,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = Result<(), HealthCheckError>> + Send + '_>,
> {
Box::pin(async move {
let connection_exists = self.connections.contains_key(&instance_id);
if let Some(handle) = self.connections.get(&instance_id) {
if !handle.tx.is_disconnected() {
return Ok(()); }
drop(handle);
self.reap_stale_connection(instance_id);
}
let addr = *self
.peers
.get(&instance_id)
.ok_or(HealthCheckError::PeerNotRegistered)?
.value();
match tokio::time::timeout(timeout, TcpStream::connect(addr)).await {
Ok(Ok(_stream)) => {
if connection_exists {
Ok(())
} else {
Err(HealthCheckError::NeverConnected)
}
}
Ok(Err(_)) => Err(HealthCheckError::ConnectionFailed),
Err(_) => Err(HealthCheckError::Timeout),
}
})
}
}
struct WriterTaskContext {
connections: Arc<DashMap<crate::InstanceId, ConnectionHandle>>,
cancel_token: CancellationToken,
connect_timeout: Duration,
reader_ctx: Option<DialedReaderContext>,
metrics: Option<std::sync::Arc<dyn velo_ext::TransportObservability>>,
}
async fn connection_writer_task(
addr: SocketAddr,
instance_id: crate::InstanceId,
rx: flume::Receiver<SendTask>,
ctx: WriterTaskContext,
) -> Result<()> {
let WriterTaskContext {
connections,
cancel_token,
connect_timeout,
reader_ctx,
metrics,
} = ctx;
let result = connection_writer_inner(
addr,
instance_id,
&rx,
&cancel_token,
connect_timeout,
reader_ctx,
metrics.clone(),
)
.await;
while let Ok(msg) = rx.try_recv() {
msg.on_error("Connection closed");
}
drop(rx);
if let Some((_, stale)) = connections.remove_if(&instance_id, |_, h| h.tx.is_disconnected()) {
stale.retire();
}
if let Some(metrics) = metrics.as_ref() {
metrics.set_active_connections(connections.len());
}
debug!("Connection to {} ({}) closed", instance_id, addr);
result
}
async fn connection_writer_inner(
addr: SocketAddr,
instance_id: crate::InstanceId,
rx: &flume::Receiver<SendTask>,
cancel_token: &CancellationToken,
connect_timeout: Duration,
reader_ctx: Option<DialedReaderContext>,
metrics: Option<std::sync::Arc<dyn velo_ext::TransportObservability>>,
) -> Result<()> {
debug!("Connecting to {}", addr);
let stream = tokio::select! {
_ = cancel_token.cancelled() => return Ok(()),
res = tokio::time::timeout(connect_timeout, TcpStream::connect(addr)) => {
res.context("connect timeout")?.context("connect failed")?
},
};
if let Err(e) = stream.set_nodelay(true) {
warn!("Failed to set TCP_NODELAY: {}", e);
}
let sock = socket2::SockRef::from(&stream);
if let Err(e) = sock.set_tcp_keepalive(
&socket2::TcpKeepalive::new()
.with_time(Duration::from_secs(60))
.with_interval(Duration::from_secs(10)),
) {
warn!("Failed to set keepalive: {}", e);
}
if let Err(e) = sock.set_send_buffer_size(2_097_152) {
warn!("Failed to set send buffer size: {}", e);
}
if let Err(e) = sock.set_recv_buffer_size(2_097_152) {
warn!("Failed to set recv buffer size: {}", e);
}
debug!("Connected to {}", addr);
let (read_half, mut write_half) = stream.into_split();
let conn_cancel = cancel_token.child_token();
let reader = reader_ctx.map(|ctx| {
tokio::spawn(run_dialed_reader(
read_half,
ctx,
metrics,
conn_cancel.clone(),
format!("{} ({})", instance_id, addr),
))
});
run_coalescing_writer(
&mut write_half,
rx,
std::convert::identity,
Some(&conn_cancel),
&TcpWriterObserver { instance_id, addr },
)
.await;
if let Some(reader) = reader {
reader.abort();
}
Ok(())
}
struct DialedReaderErrorHandler;
impl TransportErrorHandler for DialedReaderErrorHandler {
fn on_error(&self, _header: Bytes, _payload: Bytes, error: String) {
warn!("Transport error: {}", error);
}
}
impl Coalescable for SendTask {
type FailureToken = Self;
fn msg_type(&self) -> MessageType {
self.msg_type
}
fn header(&self) -> &[u8] {
&self.header
}
fn payload(&self) -> &[u8] {
&self.payload
}
fn into_failure_token(self) -> Self {
self
}
fn fail(token: Self, reason: &str) {
token.on_error(format!("Failed to write to stream: {}", reason));
}
}
struct TcpWriterObserver {
instance_id: crate::InstanceId,
addr: SocketAddr,
}
impl WriterObserver for TcpWriterObserver {
fn on_failure(&self, kind: WriterFailure, err: &std::io::Error, frames: usize) {
match kind {
WriterFailure::Write => error!(
"Write error to {} ({}): {} ({} message(s) in batch)",
self.instance_id, self.addr, err, frames
),
WriterFailure::Encode => error!(
"Encode error to {} ({}): {}",
self.instance_id, self.addr, err
),
}
}
}
#[cfg(test)]
fn parse_tcp_endpoint(endpoint: &[u8]) -> Result<SocketAddr> {
use std::net::ToSocketAddrs;
let endpoint_str = std::str::from_utf8(endpoint).context("endpoint is not valid UTF-8")?;
let addr_str = endpoint_str.strip_prefix("tcp://").unwrap_or(endpoint_str);
let mut addrs = addr_str
.to_socket_addrs()
.context("failed to parse socket address")?;
addrs
.next()
.ok_or_else(|| anyhow::anyhow!("no addresses resolved"))
}
pub struct TcpTransportBuilder {
bind_addr: Option<SocketAddr>,
key: Option<TransportKey>,
channel_capacity: usize,
connect_timeout: Duration,
listener: Option<std::net::TcpListener>,
interface_filter: InterfaceFilter,
numa_hint: Option<u32>,
shrink_threshold: Option<usize>,
}
impl TcpTransportBuilder {
pub fn new() -> Self {
Self {
bind_addr: None,
key: None,
channel_capacity: 256,
connect_timeout: Duration::from_secs(5),
listener: None,
interface_filter: InterfaceFilter::default(),
numa_hint: None,
shrink_threshold: None,
}
}
pub fn bind_addr(mut self, addr: SocketAddr) -> Self {
self.bind_addr = Some(addr);
self
}
pub fn key(mut self, key: TransportKey) -> Self {
self.key = Some(key);
self
}
pub fn channel_capacity(mut self, capacity: usize) -> Self {
self.channel_capacity = capacity;
self
}
pub fn connect_timeout(mut self, timeout: Duration) -> Self {
self.connect_timeout = timeout;
self
}
pub fn interface_filter(mut self, filter: InterfaceFilter) -> Self {
self.interface_filter = filter;
self
}
pub fn numa_hint(mut self, node: u32) -> Self {
self.numa_hint = Some(node);
self
}
pub fn shrink_threshold(mut self, bytes: usize) -> Self {
self.shrink_threshold = Some(bytes);
self
}
pub fn from_listener(mut self, listener: std::net::TcpListener) -> Result<Self> {
if self.bind_addr.is_some() {
anyhow::bail!(
"Cannot use both bind_addr() and from_listener() - they are mutually exclusive"
);
}
let addr = listener
.local_addr()
.context("Failed to get local address from listener")?;
self.bind_addr = Some(addr);
self.listener = Some(listener);
Ok(self)
}
pub fn build(self) -> Result<TcpTransport> {
let key = self.key.unwrap_or_else(|| TransportKey::from("tcp"));
let (bind_addr, listener) = if let Some(listener) = self.listener {
super::listener::size_listener_buffers(&listener);
let addr = listener.local_addr()?;
(addr, Some(listener))
} else {
let requested = self
.bind_addr
.unwrap_or_else(|| "0.0.0.0:0".parse().unwrap());
let domain = if requested.is_ipv4() {
socket2::Domain::IPV4
} else {
socket2::Domain::IPV6
};
let socket =
socket2::Socket::new(domain, socket2::Type::STREAM, Some(socket2::Protocol::TCP))
.context("Failed to create TCP listener socket")?;
socket
.set_reuse_address(true)
.context("Failed to set SO_REUSEADDR")?;
super::listener::size_listener_buffers(&socket);
socket
.bind(&requested.into())
.context("Failed to pre-bind TCP listener")?;
socket.listen(128).context("Failed to listen")?;
let std_listener: std::net::TcpListener = socket.into();
let actual = std_listener.local_addr()?;
(actual, Some(std_listener))
};
let endpoints = resolve_advertise_endpoints(bind_addr, &self.interface_filter)?;
if let (Some(numa), InterfaceFilter::ByName(name)) =
(self.numa_hint, &self.interface_filter)
{
for ep in &endpoints {
if let Some(ep_numa) = ep.numa_node
&& ep_numa != numa as i32
{
warn!(
"NIC {} is on NUMA node {} but GPU NUMA hint is {}",
name, ep_numa, numa
);
}
}
}
let encoded =
rmp_serde::to_vec(&endpoints).context("Failed to encode interface endpoints")?;
let mut addr_builder = crate::transports::address::WorkerAddressBuilder::new();
addr_builder.add_entry(key.clone(), encoded)?;
let local_address = addr_builder.build()?;
let mut transport = TcpTransport::new(
bind_addr,
key,
local_address,
self.channel_capacity,
self.connect_timeout,
listener,
self.numa_hint,
);
if let Some(t) = self.shrink_threshold {
transport.shrink_threshold = t;
}
Ok(transport)
}
}
impl Default for TcpTransportBuilder {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
#[path = "tests.rs"]
mod tests;