use anyhow::{Context, Result};
use bytes::Bytes;
use dashmap::DashMap;
use std::collections::HashMap;
use std::sync::{Arc, OnceLock};
use std::time::Duration;
use tracing::{debug, error, info};
use crate::transports::transport::{
HealthCheckError, SendBackpressure, ShutdownState, TransportError, TransportErrorHandler,
try_send_or_backpressure,
};
use velo_ext::{MessageType, PeerInfo, Transport, TransportAdapter, TransportKey, WorkerAddress};
use super::listener;
pub struct ZmqTransport {
key: TransportKey,
bind_endpoint: String,
local_address: WorkerAddress,
peers: Arc<DashMap<crate::InstanceId, String>>,
zmq_context: Arc<zmq::Context>,
sender_tx: OnceLock<flume::Sender<SenderCommand>>,
runtime: OnceLock<tokio::runtime::Handle>,
shutdown_state: OnceLock<ShutdownState>,
channel_capacity: usize,
sndhwm: i32,
rcvhwm: i32,
linger_ms: i32,
metrics: OnceLock<std::sync::Arc<dyn velo_ext::TransportObservability>>,
listener_handle: std::sync::Mutex<Option<std::thread::JoinHandle<()>>>,
sender_handle: std::sync::Mutex<Option<std::thread::JoinHandle<()>>>,
listener_control_endpoint: String,
router_socket: std::sync::Mutex<Option<zmq::Socket>>,
}
pub(crate) struct OutboundTask {
pub target: crate::InstanceId,
pub msg_type: MessageType,
pub header: Bytes,
pub payload: Bytes,
pub on_error: Arc<dyn TransportErrorHandler>,
}
impl OutboundTask {
fn on_error(self, error: impl Into<String>) {
self.on_error
.on_error(self.header, self.payload, error.into());
}
}
pub(crate) enum SenderCommand {
Send(OutboundTask),
Shutdown,
}
impl ZmqTransport {
fn update_peer_gauge(&self) {
if let Some(metrics) = self.metrics.get() {
metrics.set_registered_peers(self.peers.len());
}
}
}
impl Transport for ZmqTransport {
fn key(&self) -> TransportKey {
self.key.clone()
}
fn address(&self) -> WorkerAddress {
self.local_address.clone()
}
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 endpoint_str = std::str::from_utf8(&endpoint).map_err(|_| {
error!("ZMQ endpoint is not valid UTF-8");
TransportError::InvalidEndpoint
})?;
if !endpoint_str.starts_with("tcp://") && !endpoint_str.starts_with("ipc://") {
error!(
"Invalid ZMQ peer endpoint (only tcp:// and ipc:// supported): {}",
endpoint_str
);
return Err(TransportError::InvalidEndpoint);
}
self.peers
.insert(peer_info.instance_id(), endpoint_str.to_string());
self.update_peer_gauge();
debug!(
"Registered ZMQ peer {} at {}",
peer_info.instance_id(),
endpoint_str
);
Ok(())
}
#[inline]
fn send_message(
&self,
instance_id: crate::InstanceId,
header: Bytes,
payload: Bytes,
message_type: MessageType,
on_error: Arc<dyn TransportErrorHandler>,
) -> Result<(), SendBackpressure> {
let task = OutboundTask {
target: instance_id,
msg_type: message_type,
header,
payload,
on_error,
};
let tx = match self.sender_tx.get() {
Some(tx) => tx,
None => {
task.on_error("Transport not started");
return Ok(());
}
};
let r = try_send_or_backpressure(
tx,
SenderCommand::Send(task),
|cmd| match cmd {
SenderCommand::Send(task) => task.on_error("Sender thread exited"),
SenderCommand::Shutdown => {}
},
|cmd| {
if let SenderCommand::Send(task) = cmd {
task.on_error("Sender channel closed");
}
},
);
if let Some(m) = self.metrics.get()
&& r.is_err()
{
m.record_send_backpressure();
}
r
}
fn start(
&self,
instance_id: crate::InstanceId,
channels: TransportAdapter,
rt: tokio::runtime::Handle,
) -> futures::future::BoxFuture<'_, Result<()>> {
self.runtime.set(rt.clone()).ok();
self.shutdown_state
.set(channels.shutdown_state.clone())
.ok();
let ctx = self.zmq_context.clone();
let bind_endpoint = self.bind_endpoint.clone();
let listener_control_ep = self.listener_control_endpoint.clone();
let peers = self.peers.clone();
let channel_capacity = self.channel_capacity;
let sndhwm = self.sndhwm;
let rcvhwm = self.rcvhwm;
let linger_ms = self.linger_ms;
let metrics = self.metrics.get().cloned();
let shutdown_state = channels.shutdown_state.clone();
let instance_id_bytes = instance_id.as_bytes().to_vec();
let router_socket = self
.router_socket
.lock()
.expect("router_socket mutex poisoned")
.take();
Box::pin(async move {
let (sender_tx, sender_rx) = flume::bounded(channel_capacity);
let _ = self.sender_tx.set(sender_tx);
let (ready_tx, ready_rx) = std::sync::mpsc::sync_channel(1);
let listener_cfg = listener::ListenerConfig {
ctx: ctx.clone(),
bind_endpoint,
control_endpoint: listener_control_ep,
adapter: channels,
shutdown_state: shutdown_state.clone(),
rcvhwm,
linger_ms,
metrics: metrics.clone(),
router_socket,
ready_tx,
};
let listener_handle = std::thread::Builder::new()
.name("zmq-listener".to_string())
.spawn(move || {
listener::run_listener(listener_cfg);
})
.context("Failed to spawn ZMQ listener thread")?;
match ready_rx.recv() {
Ok(Ok(())) => {}
Ok(Err(e)) => anyhow::bail!("ZMQ listener failed to start: {}", e),
Err(_) => anyhow::bail!("ZMQ listener thread exited before signaling ready"),
}
*self
.listener_handle
.lock()
.expect("listener_handle mutex poisoned") = Some(listener_handle);
let (sender_ready_tx, sender_ready_rx) = std::sync::mpsc::sync_channel(1);
let sender_cfg = SenderConfig {
ctx,
rx: sender_rx,
peers,
identity: instance_id_bytes,
sndhwm,
linger_ms,
metrics,
ready_tx: sender_ready_tx,
};
let sender_handle = std::thread::Builder::new()
.name("zmq-sender".to_string())
.spawn(move || {
run_sender(sender_cfg);
})
.context("Failed to spawn ZMQ sender thread")?;
match sender_ready_rx.recv() {
Ok(Ok(())) => {}
Ok(Err(e)) => anyhow::bail!("ZMQ sender failed to start: {}", e),
Err(_) => anyhow::bail!("ZMQ sender thread exited before signaling ready"),
}
*self
.sender_handle
.lock()
.expect("sender_handle mutex poisoned") = Some(sender_handle);
info!("ZMQ transport started on {}", self.bind_endpoint);
Ok(())
})
}
fn begin_drain(&self) {
}
fn shutdown(&self) {
info!("Shutting down ZMQ transport");
if let Ok(ctrl) = self.zmq_context.socket(zmq::PAIR)
&& ctrl.connect(&self.listener_control_endpoint).is_ok()
{
let _ = ctrl.send("shutdown", 0);
}
if let Some(tx) = self.sender_tx.get()
&& let Err(e) = tx.try_send(SenderCommand::Shutdown)
{
debug!("ZMQ shutdown signal not sent (channel full or disconnected): {e}");
}
if let Some(handle) = self.listener_handle.lock().expect("mutex poisoned").take() {
let _ = handle.join();
}
if let Some(handle) = self.sender_handle.lock().expect("mutex poisoned").take() {
let _ = handle.join();
}
}
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: crate::InstanceId,
timeout: Duration,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = Result<(), HealthCheckError>> + Send + '_>,
> {
Box::pin(async move {
let endpoint = self
.peers
.get(&instance_id)
.map(|e| e.value().clone())
.ok_or(HealthCheckError::PeerNotRegistered)?;
let ctx = self.zmq_context.clone();
tokio::task::spawn_blocking(move || -> Result<(), HealthCheckError> {
let timeout_ms = timeout.as_millis() as i32;
let sock = ctx
.socket(zmq::DEALER)
.map_err(|_| HealthCheckError::ConnectionFailed)?;
sock.set_linger(0).ok();
sock.set_connect_timeout(timeout_ms).ok();
let monitor_endpoint = format!("inproc://zmq-healthcheck-monitor-{:p}", &sock);
let events = (zmq::SocketEvent::CONNECTED as i32)
| (zmq::SocketEvent::CONNECT_RETRIED as i32)
| (zmq::SocketEvent::DISCONNECTED as i32);
sock.monitor(&monitor_endpoint, events)
.map_err(|_| HealthCheckError::ConnectionFailed)?;
let monitor_sock = ctx
.socket(zmq::PAIR)
.map_err(|_| HealthCheckError::ConnectionFailed)?;
monitor_sock.set_rcvtimeo(timeout_ms).ok();
monitor_sock
.connect(&monitor_endpoint)
.map_err(|_| HealthCheckError::ConnectionFailed)?;
sock.connect(&endpoint)
.map_err(|_| HealthCheckError::ConnectionFailed)?;
const ZMQ_EVENT_CONNECTED: u16 = 0x0001;
const ZMQ_EVENT_CONNECT_RETRIED: u16 = 0x0040;
const ZMQ_EVENT_DISCONNECTED: u16 = 0x0200;
loop {
let data = monitor_sock
.recv_bytes(0)
.map_err(|_| HealthCheckError::Timeout)?;
let _ = monitor_sock.recv_bytes(0);
if data.len() >= 2 {
let event_id = u16::from_le_bytes([data[0], data[1]]);
match event_id {
ZMQ_EVENT_CONNECTED => return Ok(()),
ZMQ_EVENT_CONNECT_RETRIED | ZMQ_EVENT_DISCONNECTED => {
return Err(HealthCheckError::ConnectionFailed);
}
_ => { }
}
}
}
})
.await
.map_err(|_| HealthCheckError::Timeout)?
})
}
}
struct SenderConfig {
ctx: Arc<zmq::Context>,
rx: flume::Receiver<SenderCommand>,
peers: Arc<DashMap<crate::InstanceId, String>>,
identity: Vec<u8>,
sndhwm: i32,
linger_ms: i32,
metrics: Option<std::sync::Arc<dyn velo_ext::TransportObservability>>,
ready_tx: std::sync::mpsc::SyncSender<Result<(), String>>,
}
fn run_sender(cfg: SenderConfig) {
let _ = cfg.ready_tx.send(Ok(()));
let mut dealer_sockets: HashMap<crate::InstanceId, zmq::Socket> = HashMap::new();
while let Ok(cmd) = cfg.rx.recv() {
let task = match cmd {
SenderCommand::Send(task) => task,
SenderCommand::Shutdown => {
debug!("ZMQ sender received shutdown signal");
break;
}
};
let target = task.target;
let sock = match dealer_sockets.get(&target) {
Some(s) => s,
None => {
let endpoint = match cfg.peers.get(&target) {
Some(ep) => ep.value().clone(),
None => {
task.on_error(format!("Peer not registered: {}", target));
continue;
}
};
match create_dealer_socket(
&cfg.ctx,
&cfg.identity,
&endpoint,
cfg.sndhwm,
cfg.linger_ms,
) {
Ok(sock) => {
dealer_sockets.insert(target, sock);
dealer_sockets.get(&target).unwrap()
}
Err(e) => {
task.on_error(format!("Failed to create DEALER socket: {}", e));
continue;
}
}
}
};
let type_byte: &[u8] = &[task.msg_type.as_u8()];
let send_result = sock
.send(type_byte, zmq::SNDMORE)
.and_then(|_| sock.send(task.header.as_ref(), zmq::SNDMORE))
.and_then(|_| sock.send(task.payload.as_ref(), 0));
match send_result {
Ok(()) => {
if let Some(ref m) = cfg.metrics {
m.record_frame(
crate::observability::Direction::Outbound,
crate::transports::message_type_label(task.msg_type),
task.header.len() + task.payload.len(),
);
}
}
Err(e) => {
error!("ZMQ send error to {}: {}", target, e);
dealer_sockets.remove(&target);
task.on_error(format!("ZMQ send failed: {}", e));
}
}
}
while let Ok(cmd) = cfg.rx.try_recv() {
if let SenderCommand::Send(task) = cmd {
task.on_error("Transport shutting down");
}
}
drop(dealer_sockets);
debug!("ZMQ sender thread exited");
}
fn create_dealer_socket(
ctx: &zmq::Context,
identity: &[u8],
endpoint: &str,
sndhwm: i32,
linger_ms: i32,
) -> Result<zmq::Socket> {
let sock = ctx
.socket(zmq::DEALER)
.context("Failed to create DEALER socket")?;
sock.set_identity(identity)
.context("Failed to set DEALER identity")?;
sock.set_sndhwm(sndhwm)
.context("Failed to set ZMQ_SNDHWM")?;
sock.set_linger(linger_ms)
.context("Failed to set ZMQ_LINGER")?;
sock.set_sndtimeo(5000)
.context("Failed to set ZMQ_SNDTIMEO")?;
sock.set_immediate(true)
.context("Failed to set ZMQ_IMMEDIATE")?;
sock.connect(endpoint)
.context(format!("Failed to connect DEALER to {}", endpoint))?;
debug!("Created DEALER socket connected to {}", endpoint);
Ok(sock)
}
pub struct ZmqTransportBuilder {
bind_endpoint: Option<String>,
key: Option<TransportKey>,
channel_capacity: usize,
zmq_io_threads: usize,
sndhwm: i32,
rcvhwm: i32,
linger_ms: i32,
}
impl ZmqTransportBuilder {
pub fn new() -> Self {
Self {
bind_endpoint: None,
key: None,
channel_capacity: 256,
zmq_io_threads: 1,
sndhwm: 1000,
rcvhwm: 1000,
linger_ms: 1000,
}
}
pub fn bind_endpoint(mut self, endpoint: impl Into<String>) -> Self {
self.bind_endpoint = Some(endpoint.into());
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 zmq_io_threads(mut self, threads: usize) -> Self {
self.zmq_io_threads = threads;
self
}
pub fn sndhwm(mut self, hwm: i32) -> Self {
self.sndhwm = hwm;
self
}
pub fn rcvhwm(mut self, hwm: i32) -> Self {
self.rcvhwm = hwm;
self
}
pub fn linger_ms(mut self, ms: i32) -> Self {
self.linger_ms = ms;
self
}
pub fn build(self) -> Result<ZmqTransport> {
let key = self.key.unwrap_or_else(|| TransportKey::from("zmq"));
let requested_endpoint = self
.bind_endpoint
.unwrap_or_else(|| "tcp://127.0.0.1:0".to_string());
let ctx = zmq::Context::new();
ctx.set_io_threads(self.zmq_io_threads as i32)
.context("Failed to set ZMQ IO threads")?;
let router = ctx
.socket(zmq::ROUTER)
.context("Failed to create ROUTER socket")?;
router
.set_rcvhwm(self.rcvhwm)
.context("Failed to set ZMQ_RCVHWM")?;
router
.set_linger(self.linger_ms)
.context("Failed to set ZMQ_LINGER")?;
router
.set_router_mandatory(true)
.context("Failed to set ZMQ_ROUTER_MANDATORY")?;
router
.set_immediate(true)
.context("Failed to set ZMQ_IMMEDIATE")?;
router.bind(&requested_endpoint).context(format!(
"Failed to bind ROUTER socket to {}",
requested_endpoint
))?;
let resolved_endpoint = router
.get_last_endpoint()
.context("Failed to get last endpoint")?
.map_err(|_| anyhow::anyhow!("Failed to get resolved endpoint"))?;
let mut addr_builder = crate::transports::address::WorkerAddressBuilder::new();
addr_builder.add_entry(key.clone(), resolved_endpoint.as_bytes().to_vec())?;
let local_address = addr_builder.build()?;
let unique_id = crate::InstanceId::new_v4();
let listener_control_endpoint = format!("inproc://zmq-listener-ctrl-{}", unique_id);
Ok(ZmqTransport {
key,
bind_endpoint: resolved_endpoint,
local_address,
peers: Arc::new(DashMap::new()),
zmq_context: Arc::new(ctx),
sender_tx: OnceLock::new(),
runtime: OnceLock::new(),
shutdown_state: OnceLock::new(),
channel_capacity: self.channel_capacity,
sndhwm: self.sndhwm,
rcvhwm: self.rcvhwm,
linger_ms: self.linger_ms,
metrics: OnceLock::new(),
listener_handle: std::sync::Mutex::new(None),
sender_handle: std::sync::Mutex::new(None),
listener_control_endpoint,
router_socket: std::sync::Mutex::new(Some(router)),
})
}
}
impl Default for ZmqTransportBuilder {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::transports::address::WorkerAddressBuilder;
use velo_ext::PeerInfo;
fn make_zmq_peer(endpoint: &str) -> PeerInfo {
let instance_id = crate::InstanceId::new_v4();
let mut builder = WorkerAddressBuilder::new();
builder
.add_entry("zmq", endpoint.as_bytes().to_vec())
.unwrap();
PeerInfo::new(instance_id, builder.build().unwrap())
}
#[test]
fn test_builder_default() {
let transport = ZmqTransportBuilder::new().build();
assert!(transport.is_ok());
}
#[test]
fn test_builder_with_endpoint() {
let transport = ZmqTransportBuilder::new()
.bind_endpoint("tcp://127.0.0.1:0")
.build();
assert!(transport.is_ok());
let t = transport.unwrap();
assert!(t.bind_endpoint.starts_with("tcp://127.0.0.1:"));
}
#[test]
fn test_register_valid_peer() {
let transport = ZmqTransportBuilder::new().build().unwrap();
let peer = make_zmq_peer("tcp://127.0.0.1:9999");
let iid = peer.instance_id();
assert!(transport.register(peer).is_ok());
assert!(transport.peers.contains_key(&iid));
}
#[test]
fn test_register_invalid_endpoint() {
let transport = ZmqTransportBuilder::new().build().unwrap();
let peer = make_zmq_peer("invalid://foo");
assert!(transport.register(peer).is_err());
}
#[test]
fn test_register_inproc_rejected() {
let transport = ZmqTransportBuilder::new().build().unwrap();
let peer = make_zmq_peer("inproc://test");
assert!(transport.register(peer).is_err());
}
#[test]
fn test_address_contains_endpoint() {
let transport = ZmqTransportBuilder::new().build().unwrap();
let wa = transport.address();
let entry = wa.get_entry("zmq").unwrap().unwrap();
let endpoint = std::str::from_utf8(&entry).unwrap();
assert!(endpoint.starts_with("tcp://"));
}
}