use crate::error::ZmqError;
#[cfg(feature = "io-uring")]
use crate::io_uring_backend::connection_handler::OutgoingMessage;
#[cfg(feature = "io-uring")]
use crate::io_uring_backend::ops::UringOpRequest;
#[cfg(feature = "io-uring")]
use crate::io_uring_backend::ops::{WAKEUP_STATE_SIGNALED, WAKEUP_STATE_SLEEPING};
use crate::message::{FrameBatch, Msg};
#[cfg(feature = "io-uring")]
use crate::uring;
#[cfg(feature = "io-uring")]
use crate::Context;
use std::any::Any;
use std::fmt;
#[cfg(feature = "io-uring")]
use std::os::{fd::AsRawFd, unix::io::RawFd};
#[cfg(feature = "io-uring")]
use std::sync::atomic::{AtomicU8, Ordering};
#[cfg(feature = "io-uring")]
use std::sync::Arc;
#[cfg(feature = "io-uring")]
use std::time::Duration;
use async_trait::async_trait;
#[cfg(feature = "io-uring")]
use fibre::mpsc;
#[cfg(feature = "io-uring")]
use fibre::oneshot::oneshot;
#[async_trait]
pub(crate) trait ISocketConnection: Send + Sync + fmt::Debug {
async fn send_message(&self, msg: Msg) -> Result<(), ZmqError> {
let mut fb = FrameBatch::new();
fb.push(msg);
self.send_multipart(fb).await
}
async fn send_multipart(&self, msgs: FrameBatch) -> Result<(), ZmqError>;
async fn send_multipart_owned(&self, msgs: FrameBatch) -> Result<(), (FrameBatch, ZmqError)> {
match self.send_multipart(msgs).await {
Ok(()) => Ok(()),
Err(e) => Err((FrameBatch::new(), e)),
}
}
fn try_send_multipart_owned_sync(&self, msgs: FrameBatch) -> Result<(), (FrameBatch, ZmqError)> {
Err((msgs, ZmqError::ResourceLimitReached)) }
async fn close_connection(&self) -> Result<(), ZmqError>;
fn as_any(&self) -> &dyn Any;
}
#[derive(Debug, Clone)]
pub(crate) struct DummyConnection;
#[async_trait]
impl ISocketConnection for DummyConnection {
async fn send_message(&self, _msg: Msg) -> Result<(), ZmqError> {
Err(ZmqError::UnsupportedFeature(
"DummyConnection cannot send".into(),
))
}
async fn send_multipart(&self, _msgs: FrameBatch) -> Result<(), ZmqError> {
Err(ZmqError::UnsupportedFeature(
"DummyConnection cannot send multipart".into(),
))
}
async fn close_connection(&self) -> Result<(), ZmqError> {
Ok(())
}
fn as_any(&self) -> &dyn Any {
self
}
}
#[cfg(feature = "io-uring")]
pub(crate) struct UringFdConnection {
fd: RawFd,
mpsc_tx: mpsc::BoundedAsyncSender<OutgoingMessage>,
event_fd: eventfd::EventFD,
worker_asleep: Arc<AtomicU8>,
context: Context,
}
#[cfg(feature = "io-uring")]
impl fmt::Debug for UringFdConnection {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("UringFdConnection")
.field("fd", &self.fd)
.field("mpsc_tx_is_closed", &self.mpsc_tx.is_closed())
.field("event_fd_raw", &self.event_fd.as_raw_fd())
.field("context_present", &true)
.finish()
}
}
#[cfg(feature = "io-uring")]
impl UringFdConnection {
pub(crate) fn new(
fd: RawFd,
mpsc_tx: mpsc::BoundedAsyncSender<OutgoingMessage>,
event_fd: eventfd::EventFD,
worker_asleep: Arc<AtomicU8>,
context: Context,
) -> Self {
Self {
fd,
mpsc_tx,
event_fd,
worker_asleep,
context,
}
}
}
#[cfg(feature = "io-uring")]
impl UringFdConnection {
async fn send_outgoing(&self, msg: OutgoingMessage) -> Result<(), ZmqError> {
match self.mpsc_tx.send(msg).await {
Ok(()) => {
if self.worker_asleep.load(Ordering::Relaxed) == WAKEUP_STATE_SLEEPING {
if self.worker_asleep.compare_exchange(
WAKEUP_STATE_SLEEPING,
WAKEUP_STATE_SIGNALED,
Ordering::AcqRel,
Ordering::Acquire,
).is_ok() {
if let Err(e) = self.event_fd.write(1) {
tracing::error!("UringFdConnection: Failed to signal eventfd: {}", e);
}
}
}
Ok(())
}
Err(_) => Err(ZmqError::ConnectionClosed),
}
}
}
#[cfg(feature = "io-uring")]
#[async_trait]
impl ISocketConnection for UringFdConnection {
async fn send_message(&self, msg: Msg) -> Result<(), ZmqError> {
self.send_outgoing(OutgoingMessage::Single(msg)).await
}
async fn send_multipart(&self, msgs: FrameBatch) -> Result<(), ZmqError> {
self.send_outgoing(OutgoingMessage::Multipart(msgs)).await
}
async fn close_connection(&self) -> Result<(), ZmqError> {
let (reply_tx, reply_rx) = oneshot();
let unique_user_data = self.context.inner().next_handle() as u64;
let req = UringOpRequest::ShutdownConnectionHandler {
user_data: unique_user_data,
fd: self.fd,
reply_tx,
};
let mut worker_op_tx = uring::global_state::get_global_uring_worker_op_tx()?;
worker_op_tx.send(req).await.map_err(|e| {
ZmqError::Internal(format!("UringWorker op channel error for close: {}", e))
})?;
match tokio::time::timeout(Duration::from_secs(5), reply_rx.recv()).await {
Ok(Ok(Ok(_))) => Ok(()),
Ok(Ok(Err(e))) => Err(e),
Ok(Err(_)) => Err(ZmqError::Internal("UringWorker reply channel error for close".into())),
Err(_) => Err(ZmqError::Timeout),
}
}
fn as_any(&self) -> &dyn Any {
self
}
}