#![deny(missing_docs)]
pub(crate) mod address;
pub(crate) mod coalesce;
pub(crate) mod ingress;
pub mod tcp;
pub mod utils;
#[cfg(unix)]
pub mod uds;
#[cfg(all(target_os = "linux", feature = "ucx"))]
pub mod ucx;
#[cfg(feature = "nats-transport")]
pub mod nats;
#[cfg(feature = "grpc")]
pub mod grpc;
#[cfg(feature = "zmq")]
pub mod zmq;
mod transport;
use std::{collections::HashMap, sync::Arc};
use crate::observability::{Direction, TransportRejection, VeloMetrics};
use bytes::Bytes;
use dashmap::DashMap;
use parking_lot::Mutex;
use velo_ext::{InstanceId, PeerInfo, TransportKey, WorkerAddress, WorkerId};
use address::WorkerAddressBuilder;
pub use utils::interfaces::{InterfaceEndpoint, InterfaceFilter};
pub use transport::{
AdmissionError, AdmissionGate, AdmissionState, AdmitOutcome, DataStreams, HealthCheckError,
InFlightGuard, InboundMessage, MessageType, SendAdmission, SendOutcome, ShutdownPolicy,
ShutdownState, Transport, TransportAdapter, TransportError, TransportErrorHandler,
make_channels,
};
#[derive(Debug, thiserror::Error)]
pub enum VeloBackendError {
#[error("No compatible transports found")]
NoCompatibleTransports,
#[error("Transport not found for instance: {0}")]
InstanceNotRegistered(InstanceId),
#[error("Worker not found: {0}")]
WorkerNotRegistered(WorkerId),
#[error("Transport not found: {0}")]
TransportNotFound(TransportKey),
#[error("Invalid transport priority: {0}")]
InvalidTransportPriority(String),
}
pub struct VeloBackend {
instance_id: InstanceId,
address: WorkerAddress,
priorities: Mutex<Vec<TransportKey>>,
transports: HashMap<TransportKey, Arc<dyn Transport>>,
transport_metrics: HashMap<TransportKey, Arc<crate::observability::TransportMetricsHandle>>,
primary_transport: DashMap<InstanceId, Arc<dyn Transport>>,
alternative_transports: DashMap<InstanceId, Vec<TransportKey>>,
workers: DashMap<WorkerId, InstanceId>,
shutdown_state: ShutdownState,
#[allow(dead_code)]
runtime: tokio::runtime::Handle,
}
impl VeloBackend {
pub async fn new(
backend_transports: Vec<Arc<dyn Transport>>,
observability: Option<Arc<VeloMetrics>>,
) -> anyhow::Result<(Self, DataStreams)> {
let instance_id = InstanceId::new_v4();
let mut priorities = Vec::new();
let mut builder = WorkerAddressBuilder::new();
let mut transports = HashMap::new();
let mut transport_metrics = HashMap::new();
let (adapter, data_streams) = transport::make_channels();
let shutdown_state = adapter.shutdown_state.clone();
let runtime = tokio::runtime::Handle::current();
for transport in backend_transports {
let key = transport.key();
if let Some(metrics) = observability.as_ref() {
let handle = Arc::new(metrics.bind_transport(key.as_str()));
transport
.set_observability(handle.clone() as Arc<dyn velo_ext::TransportObservability>);
transport_metrics.insert(key.clone(), handle);
}
transport
.start(instance_id, adapter.clone(), runtime.clone())
.await?;
builder.merge(&transport.address())?;
priorities.push(key.clone());
transports.insert(key, transport);
}
let address = builder.build()?;
Ok((
Self {
instance_id,
address,
transports,
transport_metrics,
priorities: Mutex::new(priorities),
primary_transport: DashMap::new(),
alternative_transports: DashMap::new(),
workers: DashMap::new(),
shutdown_state,
runtime,
},
data_streams,
))
}
pub fn instance_id(&self) -> InstanceId {
self.instance_id
}
pub fn peer_info(&self) -> PeerInfo {
PeerInfo::new(self.instance_id, self.address.clone())
}
pub fn is_registered(&self, instance_id: InstanceId) -> bool {
self.primary_transport.contains_key(&instance_id)
}
pub fn try_translate_worker_id(
&self,
worker_id: WorkerId,
) -> Result<InstanceId, VeloBackendError> {
self.workers
.get(&worker_id)
.map(|entry| *entry)
.ok_or(VeloBackendError::WorkerNotRegistered(worker_id))
}
#[deprecated(since = "0.7.0", note = "Use try_translate_worker_id() instead")]
pub fn translate_worker_id(&self, worker_id: WorkerId) -> Result<InstanceId, VeloBackendError> {
self.try_translate_worker_id(worker_id)
}
pub fn has_instance(&self, instance_id: InstanceId) -> bool {
self.primary_transport.contains_key(&instance_id)
}
pub fn primary_transport_key(&self, target: InstanceId) -> Option<TransportKey> {
self.primary_transport
.get(&target)
.map(|entry| entry.value().key())
}
pub(crate) fn max_message_size(&self, target: InstanceId) -> Option<usize> {
self.primary_transport
.get(&target)
.and_then(|transport| transport.value().max_message_size(target))
}
pub fn alternative_transport_keys(&self, target: InstanceId) -> Option<Vec<TransportKey>> {
self.alternative_transports
.get(&target)
.map(|entry| entry.value().clone())
}
pub fn send_message(
&self,
target: InstanceId,
header: Bytes,
payload: Bytes,
message_type: MessageType,
on_error: Arc<dyn TransportErrorHandler>,
) -> anyhow::Result<SendOutcome> {
let transport = self
.primary_transport
.get(&target)
.ok_or(VeloBackendError::InstanceNotRegistered(target))?;
let transport_key = transport.value().key();
let transport_name = transport_key.to_string();
#[cfg(not(feature = "distributed-tracing"))]
let _ = &transport_name;
#[cfg(feature = "distributed-tracing")]
let bytes = header.len() + payload.len();
let metrics = self.transport_metrics.get(&transport_key);
let error_handler = instrument_transport_error_handler(metrics.cloned(), on_error);
let report = SendReport {
metrics: metrics.cloned(),
on_error: error_handler.clone(),
message_type,
header: header.clone(),
payload: payload.clone(),
};
#[cfg(feature = "distributed-tracing")]
let outcome = {
let span = tracing::info_span!(
"velo.transport.send",
transport = transport_name.as_str(),
message_type = message_type_label(message_type),
bytes
);
let _entered = span.enter();
transport.send_message(target, header, payload, message_type, error_handler)
};
#[cfg(not(feature = "distributed-tracing"))]
let outcome = transport.send_message(target, header, payload, message_type, error_handler);
Ok(finalize_send_outcome(outcome, report))
}
pub fn send_message_with_transport(
&self,
target: InstanceId,
header: Bytes,
payload: Bytes,
message_type: MessageType,
on_error: Arc<dyn TransportErrorHandler>,
transport_key: TransportKey,
) -> anyhow::Result<SendOutcome> {
let transport = self
.primary_transport
.get(&target)
.ok_or(VeloBackendError::InstanceNotRegistered(target))?;
if transport.value().key() == transport_key {
let _transport_name = transport_key.to_string();
let metrics = self.transport_metrics.get(&transport_key);
let error_handler = instrument_transport_error_handler(metrics.cloned(), on_error);
let report = SendReport {
metrics: metrics.cloned(),
on_error: error_handler.clone(),
message_type,
header: header.clone(),
payload: payload.clone(),
};
let outcome =
transport.send_message(target, header, payload, message_type, error_handler);
return Ok(finalize_send_outcome(outcome, report));
} else {
let alternative_transports = self
.alternative_transports
.get(&target)
.ok_or(VeloBackendError::InstanceNotRegistered(target))?;
for alternative_transport in alternative_transports.iter() {
if *alternative_transport == transport_key
&& let Some(transport) = self.transports.get(alternative_transport)
{
let _transport_name = alternative_transport.to_string();
let metrics = self.transport_metrics.get(alternative_transport);
let error_handler =
instrument_transport_error_handler(metrics.cloned(), on_error);
let report = SendReport {
metrics: metrics.cloned(),
on_error: error_handler.clone(),
message_type,
header: header.clone(),
payload: payload.clone(),
};
let outcome = transport.send_message(
target,
header,
payload,
message_type,
error_handler,
);
return Ok(finalize_send_outcome(outcome, report));
}
}
}
Err(VeloBackendError::NoCompatibleTransports)?
}
pub fn send_message_to_worker(
&self,
worker_id: WorkerId,
header: Bytes,
payload: Bytes,
message_type: MessageType,
on_error: Arc<dyn TransportErrorHandler>,
) -> anyhow::Result<SendOutcome> {
let instance_id = self.try_translate_worker_id(worker_id)?;
self.send_message(instance_id, header, payload, message_type, on_error)
}
pub fn register_peer(&self, peer: PeerInfo) -> Result<(), VeloBackendError> {
let instance_id = peer.instance_id();
let mut compatible_transports = Vec::new();
for (key, transport) in self.transports.iter() {
if transport.register(peer.clone()).is_ok() {
compatible_transports.push(key.clone());
}
}
if compatible_transports.is_empty() {
return Err(VeloBackendError::NoCompatibleTransports);
}
let sorted_transports = self
.priorities
.lock()
.iter()
.filter(|key| compatible_transports.contains(key))
.cloned()
.collect::<Vec<TransportKey>>();
assert!(
!sorted_transports.is_empty(),
"failed to properly sort compatible transports"
);
let primary_transport_key = sorted_transports[0].clone();
let alternative_transport_keys = sorted_transports[1..].to_vec();
let primary_transport = self.transports.get(&primary_transport_key).unwrap();
self.primary_transport
.insert(instance_id, primary_transport.clone());
self.alternative_transports
.insert(instance_id, alternative_transport_keys);
self.workers.insert(instance_id.worker_id(), instance_id);
Ok(())
}
pub fn available_transports(&self) -> Vec<TransportKey> {
self.transports.keys().cloned().collect()
}
pub fn set_transport_priority(
&self,
priorities: Vec<TransportKey>,
) -> Result<(), VeloBackendError> {
let required_transports = self.available_transports();
if required_transports.len() != priorities.len() {
return Err(VeloBackendError::InvalidTransportPriority(format!(
"Required transports: {:?}, provided priorities: {:?}",
required_transports, priorities
)));
}
for priority in &priorities {
if !required_transports.contains(priority) {
return Err(VeloBackendError::InvalidTransportPriority(format!(
"Priority transport not found: {:?}",
priority
)));
}
}
let mut guard = self.priorities.lock();
*guard = priorities;
Ok(())
}
pub fn shutdown_state(&self) -> &ShutdownState {
&self.shutdown_state
}
pub fn begin_drain(&self) {
self.shutdown_state.begin_drain();
for transport in self.transports.values() {
transport.begin_drain();
}
}
pub async fn graceful_shutdown(&self, policy: ShutdownPolicy) {
self.begin_drain();
match policy {
ShutdownPolicy::WaitForever => {
self.shutdown_state.wait_for_drain().await;
}
ShutdownPolicy::Timeout(duration) => {
let _ = tokio::time::timeout(duration, self.shutdown_state.wait_for_drain()).await;
}
}
self.shutdown_state.teardown_token().cancel();
for transport in self.transports.values() {
transport.shutdown();
}
}
}
pub(crate) fn message_type_label(message_type: MessageType) -> &'static str {
match message_type {
MessageType::Message => "message",
MessageType::Response => "response",
MessageType::Ack => "ack",
MessageType::Event => "event",
MessageType::ShuttingDown => "shutting_down",
}
}
struct InstrumentedTransportErrorHandler {
metrics: Option<Arc<crate::observability::TransportMetricsHandle>>,
inner: Arc<dyn TransportErrorHandler>,
}
impl TransportErrorHandler for InstrumentedTransportErrorHandler {
fn on_error(&self, header: Bytes, payload: Bytes, error: String) {
if let Some(metrics) = self.metrics.as_ref() {
metrics.record_rejection(TransportRejection::SendError);
}
self.inner.on_error(header, payload, error);
}
}
fn instrument_transport_error_handler(
metrics: Option<Arc<crate::observability::TransportMetricsHandle>>,
inner: Arc<dyn TransportErrorHandler>,
) -> Arc<dyn TransportErrorHandler> {
Arc::new(InstrumentedTransportErrorHandler { metrics, inner })
}
struct SendReport {
metrics: Option<Arc<crate::observability::TransportMetricsHandle>>,
on_error: Arc<dyn TransportErrorHandler>,
message_type: MessageType,
header: Bytes,
payload: Bytes,
}
fn finalize_send_outcome(outcome: SendOutcome, report: SendReport) -> SendOutcome {
let SendReport {
metrics,
on_error,
message_type,
header,
payload,
} = report;
let label = message_type_label(message_type);
let bytes = header.len() + payload.len();
match outcome {
SendOutcome::Admitted => {
if let Some(metrics) = metrics {
metrics.record_frame(Direction::Outbound, label, bytes);
}
SendOutcome::Admitted
}
SendOutcome::Pending(admission) => {
SendOutcome::Pending(admission.on_resolved(move |result| match result {
Ok(()) => {
if let Some(metrics) = metrics {
metrics.record_frame(Direction::Outbound, label, bytes);
}
}
Err(error) => {
on_error.on_error(header, payload, format!("Send not admitted: {error}"));
}
}))
}
}
}
#[cfg(test)]
mod tests;