use crate::observability::{Direction, TransportRejection};
use bytes::{BufMut, Bytes, BytesMut};
use dashmap::DashMap;
use futures::StreamExt;
use futures::future::BoxFuture;
use std::sync::{Arc, OnceLock};
use std::time::Duration;
use tokio_util::sync::CancellationToken;
pub(crate) const HEADER_VELO_TYPE: &str = "Velo-Type";
pub(crate) const HEADER_VELO_HLEN: &str = "Velo-HLen";
use super::subjects;
use velo_ext::{
AdmissionGate, AdmitOutcome, InstanceId, MessageType, PeerInfo, SendOutcome, TransportKey,
WorkerAddress,
transport::{
HealthCheckError, ShutdownState, Transport, TransportAdapter, TransportError,
TransportErrorHandler,
},
};
const VELO_TYPE_STRINGS: [&str; 5] = ["0", "1", "2", "3", "4"];
const DEFAULT_SENDER_CAPACITY: usize = 1024;
const NATS_HEADER_OVERHEAD: usize = 64;
struct NatsSendTask {
subject: String,
message_type: MessageType,
header: Bytes,
payload: Bytes,
on_error: Arc<dyn TransportErrorHandler>,
}
pub struct NatsTransport {
key: TransportKey,
client: Arc<async_nats::Client>,
cluster_id: String,
local_address: OnceLock<WorkerAddress>,
peers: Arc<DashMap<InstanceId, String>>,
sender_tx: OnceLock<flume::Sender<NatsSendTask>>,
gates: DashMap<InstanceId, AdmissionGate<NatsSendTask>>,
sender_capacity: usize,
runtime: OnceLock<tokio::runtime::Handle>,
cancel_token: CancellationToken,
begin_shutdown_token: CancellationToken,
shutdown_state: OnceLock<ShutdownState>,
metrics: OnceLock<std::sync::Arc<dyn velo_ext::TransportObservability>>,
}
impl NatsTransport {
fn frame_capacity(&self) -> usize {
self.client
.max_payload()
.saturating_sub(NATS_HEADER_OVERHEAD)
}
fn update_peer_gauge(&self) {
if let Some(metrics) = self.metrics.get() {
metrics.set_registered_peers(self.peers.len());
}
}
fn gate_for(
&self,
target: InstanceId,
tx: &flume::Sender<NatsSendTask>,
rt: &tokio::runtime::Handle,
) -> AdmissionGate<NatsSendTask> {
self.gates
.entry(target)
.or_insert_with(|| AdmissionGate::new(tx.clone(), rt.clone()))
.clone()
}
}
impl Transport for NatsTransport {
fn key(&self) -> TransportKey {
self.key.clone()
}
fn address(&self) -> WorkerAddress {
self.local_address
.get()
.cloned()
.unwrap_or_else(|| WorkerAddress::from_encoded(Bytes::from_static(&[])))
}
fn max_message_size(&self, _target: InstanceId) -> Option<usize> {
Some(self.frame_capacity())
}
fn register(&self, peer_info: PeerInfo) -> Result<(), TransportError> {
let instance_id = peer_info.instance_id();
let entry = peer_info
.worker_address()
.get_entry("nats")
.map_err(|_| TransportError::InvalidEndpoint)?
.ok_or(TransportError::NoEndpoint)?;
let subject =
String::from_utf8(entry.to_vec()).map_err(|_| TransportError::InvalidEndpoint)?;
tracing::debug!(
instance_id = %instance_id,
subject = %subject,
"Registered NATS peer"
);
self.peers.insert(instance_id, subject);
self.update_peer_gauge();
Ok(())
}
#[inline]
fn send_message(
&self,
instance_id: InstanceId,
header: Bytes,
payload: Bytes,
message_type: MessageType,
on_error: Arc<dyn TransportErrorHandler>,
) -> SendOutcome {
let subject = match self.peers.get(&instance_id) {
Some(entry) => entry.value().clone(),
None => {
on_error.on_error(
header,
payload,
format!("Peer not registered: {}", instance_id),
);
return SendOutcome::Admitted;
}
};
let frame_size = header.len() + payload.len();
if frame_size > self.frame_capacity() {
let max = self.client.max_payload();
on_error.on_error(
header,
payload,
format!(
"Frame size {} exceeds NATS max_payload {} for peer {}",
frame_size + NATS_HEADER_OVERHEAD,
max,
instance_id
),
);
return SendOutcome::Admitted;
}
let task = NatsSendTask {
subject,
message_type,
header,
payload,
on_error,
};
let (Some(tx), Some(rt)) = (self.sender_tx.get(), self.runtime.get()) else {
task.on_error.on_error(
task.header,
task.payload,
"NATS transport not started".into(),
);
return SendOutcome::Admitted;
};
let outcome = self.gate_for(instance_id, tx, rt).send(task);
if let Some(m) = self.metrics.get()
&& !outcome.is_admitted()
{
m.record_send_backpressure();
}
outcome
}
fn start(
&self,
instance_id: InstanceId,
channels: TransportAdapter,
rt: tokio::runtime::Handle,
) -> BoxFuture<'_, anyhow::Result<()>> {
let _ = self.runtime.set(rt.clone());
let _ = self.shutdown_state.set(channels.shutdown_state.clone());
Box::pin(async move {
tracing::info!(
max_payload = self.client.max_payload(),
frame_capacity = self.frame_capacity(),
"NATS max_payload for this connection"
);
let subject = subjects::inbound_subject(&self.cluster_id, instance_id);
let health_subj = subjects::health_subject(&self.cluster_id, instance_id);
let mut addr_builder = crate::transports::address::WorkerAddressBuilder::new();
addr_builder.add_entry("nats", Bytes::from(subject.as_bytes().to_vec()))?;
let _ = self.local_address.set(addr_builder.build()?);
let data_sub = self.client.subscribe(subject.clone()).await.map_err(|e| {
anyhow::anyhow!("Failed to subscribe to inbound subject {}: {}", subject, e)
})?;
let health_sub = self
.client
.subscribe(health_subj.clone())
.await
.map_err(|e| {
anyhow::anyhow!(
"Failed to subscribe to health subject {}: {}",
health_subj,
e
)
})?;
tracing::info!(
data_subject = %subject,
health_subject = %health_subj,
"NATS transport started, subscriptions live"
);
let (sender_tx, sender_rx) = flume::bounded(self.sender_capacity);
let _ = self.sender_tx.set(sender_tx);
let sender_cancel = self.cancel_token.clone();
let sender_client = self.client.clone();
let sender_metrics = self.metrics.get().cloned();
rt.spawn(run_sender_task(
sender_rx,
sender_client,
sender_cancel,
sender_metrics,
));
let cancel = self.cancel_token.clone();
let begin_shutdown = self.begin_shutdown_token.clone();
let client = self.client.clone();
let transport_key = self.key.to_string();
let metrics = self.metrics.get().cloned();
rt.spawn(async move {
run_receive_loop(
data_sub,
health_sub,
channels,
cancel,
begin_shutdown,
client,
transport_key,
metrics,
)
.await;
});
Ok(())
})
}
fn begin_drain(&self) {
}
fn shutdown(&self) {
self.begin_shutdown_token.cancel();
self.cancel_token.cancel();
}
fn set_observability(
&self,
observability: std::sync::Arc<dyn velo_ext::TransportObservability>,
) {
let _ = self.metrics.set(observability);
self.update_peer_gauge();
}
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 {
let _rt = self
.runtime
.get()
.ok_or(HealthCheckError::TransportNotStarted)?;
let subject = self
.peers
.get(&instance_id)
.ok_or(HealthCheckError::PeerNotRegistered)?
.clone();
let health_subj = format!("{}.health", subject);
let client = self.client.clone();
match tokio::time::timeout(timeout, client.request(health_subj, Bytes::new())).await {
Ok(Ok(_response)) => Ok(()),
Ok(Err(_e)) => Err(HealthCheckError::ConnectionFailed),
Err(_elapsed) => Err(HealthCheckError::Timeout),
}
})
}
}
fn build_nats_frame(
message_type: MessageType,
header: &Bytes,
payload: &Bytes,
) -> (async_nats::HeaderMap, Bytes) {
let mut nats_headers = async_nats::HeaderMap::new();
nats_headers.insert(
HEADER_VELO_TYPE,
VELO_TYPE_STRINGS[message_type as u8 as usize],
);
let hlen_str = header.len().to_string();
nats_headers.insert(HEADER_VELO_HLEN, hlen_str);
let nats_payload: Bytes = if header.is_empty() {
payload.clone()
} else if payload.is_empty() {
header.clone()
} else {
let mut buf = BytesMut::with_capacity(header.len() + payload.len());
buf.put(header.as_ref());
buf.put(payload.as_ref());
buf.freeze()
};
(nats_headers, nats_payload)
}
async fn run_sender_task(
rx: flume::Receiver<NatsSendTask>,
client: Arc<async_nats::Client>,
cancel: CancellationToken,
metrics: Option<std::sync::Arc<dyn velo_ext::TransportObservability>>,
) {
loop {
tokio::select! {
biased;
_ = cancel.cancelled() => {
tracing::debug!("NATS sender task cancelled");
break;
}
result = rx.recv_async() => {
match result {
Ok(task) => {
let (nats_headers, nats_payload) = build_nats_frame(
task.message_type,
&task.header,
&task.payload,
);
if let Some(metrics) = metrics.as_ref() {
metrics.record_frame(
Direction::Outbound,
crate::transports::message_type_label(task.message_type),
nats_payload.len(),
);
}
if let Err(e) = client
.publish_with_headers(task.subject, nats_headers, nats_payload)
.await
{
task.on_error.on_error(
task.header,
task.payload,
format!("NATS publish failed: {}", e),
);
}
}
Err(_) => {
tracing::debug!("NATS sender channel closed, exiting");
break;
}
}
}
}
}
}
#[allow(clippy::too_many_arguments)]
async fn run_receive_loop(
mut data_sub: async_nats::Subscriber,
mut health_sub: async_nats::Subscriber,
adapter: TransportAdapter,
cancel: CancellationToken,
begin_shutdown: CancellationToken,
client: Arc<async_nats::Client>,
transport_key: String,
metrics: Option<std::sync::Arc<dyn velo_ext::TransportObservability>>,
) {
loop {
tokio::select! {
biased;
_ = begin_shutdown.cancelled() => {
tracing::debug!("NATS shutdown signaled, unsubscribing before cancel");
let _ = data_sub.unsubscribe().await;
let _ = health_sub.unsubscribe().await;
cancel.cancel(); break;
}
_ = cancel.cancelled() => {
tracing::debug!("NATS receive loop cancelled directly");
break;
}
msg = data_sub.next() => {
match msg {
Some(msg) => {
match route_frame(&msg, &adapter, &transport_key, metrics.as_ref()) {
NatsRouted::Done => {}
NatsRouted::DrainRejected { header } => {
if let Some(reply) = &msg.reply {
let mut nats_headers = async_nats::HeaderMap::new();
nats_headers.insert(
HEADER_VELO_TYPE,
(MessageType::ShuttingDown as u8).to_string().as_str(),
);
nats_headers.insert(
HEADER_VELO_HLEN,
header.len().to_string().as_str(),
);
if let Err(e) = client.publish_with_headers(
reply.clone(),
nats_headers,
header,
).await {
tracing::warn!(error = %e, "Failed to send ShuttingDown response");
}
} else {
tracing::debug!(
"Discarding fire-and-forget Message during drain (no reply inbox)"
);
}
}
}
}
None => {
tracing::warn!("NATS data subscription stream ended");
break;
}
}
}
msg = health_sub.next() => {
match msg {
Some(msg) => {
if let Some(reply) = msg.reply
&& let Err(e) = client.publish(reply, Bytes::new()).await
{
tracing::warn!(error = %e, "Failed to reply to health check");
}
}
None => {
tracing::warn!("NATS health subscription stream ended");
break;
}
}
}
}
}
let _ = data_sub.unsubscribe().await;
let _ = health_sub.unsubscribe().await;
tracing::debug!("NATS receive loop exited, subscriptions unsubscribed");
}
enum NatsRouted {
Done,
DrainRejected { header: Bytes },
}
fn route_frame(
msg: &async_nats::Message,
adapter: &TransportAdapter,
transport_key: &str,
metrics: Option<&std::sync::Arc<dyn velo_ext::TransportObservability>>,
) -> NatsRouted {
#[cfg(not(feature = "distributed-tracing"))]
let _ = transport_key;
let headers = match &msg.headers {
Some(h) => h,
None => {
if let Some(metrics) = metrics {
metrics.record_rejection(TransportRejection::MissingHeaders);
}
tracing::trace!("Dropping NATS message with no headers");
return NatsRouted::Done;
}
};
let type_str = match headers.get(HEADER_VELO_TYPE) {
Some(v) => v.as_str(),
None => {
if let Some(metrics) = metrics {
metrics.record_rejection(TransportRejection::MissingType);
}
tracing::trace!("Dropping NATS message missing Velo-Type header");
return NatsRouted::Done;
}
};
let msg_type = match type_str.parse::<u8>() {
Ok(0) => MessageType::Message,
Ok(1) => MessageType::Response,
Ok(2) => MessageType::Ack,
Ok(3) => MessageType::Event,
Ok(4) => MessageType::ShuttingDown,
_ => {
if let Some(metrics) = metrics {
metrics.record_rejection(TransportRejection::InvalidType);
}
tracing::trace!(
velo_type = type_str,
"Dropping NATS message with invalid Velo-Type"
);
return NatsRouted::Done;
}
};
let hlen: usize = match headers
.get(HEADER_VELO_HLEN)
.and_then(|v| v.as_str().parse().ok())
{
Some(n) => n,
None => {
if let Some(metrics) = metrics {
metrics.record_rejection(TransportRejection::InvalidHeaderLength);
}
tracing::trace!("Dropping NATS message missing or invalid Velo-HLen header");
return NatsRouted::Done;
}
};
if msg.payload.len() < hlen {
if let Some(metrics) = metrics {
metrics.record_rejection(TransportRejection::TruncatedFrame);
}
tracing::trace!(
expected_min = hlen,
actual = msg.payload.len(),
"Dropping truncated NATS frame"
);
return NatsRouted::Done;
}
let header = msg.payload.slice(..hlen);
let body = msg.payload.slice(hlen..);
let frame_bytes = header.len() + body.len();
let result = match msg_type {
MessageType::Message => match adapter.admit_message(header, body) {
AdmitOutcome::Admitted => Ok(()),
AdmitOutcome::Draining { header, .. } => {
if let Some(metrics) = metrics {
metrics.record_rejection(TransportRejection::DrainRejected);
}
return NatsRouted::DrainRejected { header };
}
AdmitOutcome::Disconnected { .. } => Err(()),
},
MessageType::Response => adapter
.response_stream
.try_send((header, body))
.map_err(|_| ()),
MessageType::Ack | MessageType::Event => adapter
.event_stream
.try_send((header, body))
.map_err(|_| ()),
MessageType::ShuttingDown => adapter
.shutdown_stream
.try_send((header, body))
.map_err(|_| ()),
};
match result {
Ok(()) => {
if let Some(metrics) = metrics {
#[cfg(feature = "distributed-tracing")]
let span = tracing::debug_span!(
"velo.transport.receive",
transport = transport_key,
message_type = crate::transports::message_type_label(msg_type),
bytes = frame_bytes
);
#[cfg(feature = "distributed-tracing")]
let _entered = span.enter();
metrics.record_frame(
Direction::Inbound,
crate::transports::message_type_label(msg_type),
frame_bytes,
);
}
}
Err(()) => {
if let Some(metrics) = metrics {
metrics.record_rejection(TransportRejection::RouteFailed);
}
}
}
NatsRouted::Done
}
pub struct NatsTransportBuilder {
client: Arc<async_nats::Client>,
cluster_id: String,
key: TransportKey,
sender_capacity: usize,
}
impl NatsTransportBuilder {
pub fn new(client: Arc<async_nats::Client>, cluster_id: impl Into<String>) -> Self {
Self {
client,
cluster_id: cluster_id.into(),
key: TransportKey::from("nats"),
sender_capacity: DEFAULT_SENDER_CAPACITY,
}
}
pub fn with_key(mut self, key: impl Into<TransportKey>) -> Self {
self.key = key.into();
self
}
pub fn with_sender_capacity(mut self, capacity: usize) -> Self {
self.sender_capacity = capacity;
self
}
pub fn build(self) -> NatsTransport {
NatsTransport {
key: self.key,
client: self.client,
cluster_id: self.cluster_id,
local_address: OnceLock::new(),
peers: Arc::new(DashMap::new()),
sender_tx: OnceLock::new(),
gates: DashMap::new(),
sender_capacity: self.sender_capacity,
runtime: OnceLock::new(),
cancel_token: CancellationToken::new(),
begin_shutdown_token: CancellationToken::new(),
shutdown_state: OnceLock::new(),
metrics: OnceLock::new(),
}
}
}
#[cfg(test)]
#[path = "tests.rs"]
mod tests;