use std::sync::Arc;
use futures::Stream;
use ruststream::{
Broker, ConnectedBroker, DefaultPublish, DescribeServer, OutgoingMessage, PairError,
PublishPolicy, Publisher, ServerSpec, Subscribe, Subscriber,
};
use tokio::sync::{Mutex, OnceCell, mpsc};
use zeromq::prelude::*;
use zeromq::{PullSocket, PushSocket};
use crate::common::{DriverHandle, Lifecycle, SharedLifecycle, send_with_retry};
use crate::endpoint::ZmqEndpoint;
use crate::error::ZmqError;
use crate::message::ZmqMessage;
use crate::wire;
#[derive(Debug, Clone)]
#[must_use]
pub struct ZmqQueue {
endpoint: ZmqEndpoint,
cell: Arc<OnceCell<SharedLifecycle>>,
}
impl ZmqQueue {
pub fn new(endpoint: ZmqEndpoint) -> Self {
Self {
endpoint,
cell: Arc::new(OnceCell::new()),
}
}
#[must_use]
pub fn publisher(&self) -> ZmqQueuePublisher {
ZmqQueuePublisher {
cell: Arc::clone(&self.cell),
push: Arc::new(Mutex::new(None)),
}
}
}
impl Broker for ZmqQueue {
type Error = ZmqError;
type Connected = ConnectedZmqQueue;
async fn connect(self) -> Result<Self::Connected, Self::Error> {
let lifecycle = self
.cell
.get_or_try_init(async || {
self.endpoint.validate()?;
Ok::<_, ZmqError>(Arc::new(Lifecycle::new(self.endpoint.clone())))
})
.await?
.clone();
Ok(ConnectedZmqQueue {
lifecycle,
cell: self.cell,
})
}
}
impl DescribeServer for ZmqQueue {
fn describe_server(&self) -> ServerSpec {
ServerSpec::new(self.endpoint.address(), "zeromq")
}
}
#[derive(Debug)]
pub struct ConnectedZmqQueue {
lifecycle: SharedLifecycle,
cell: Arc<OnceCell<SharedLifecycle>>,
}
impl ConnectedZmqQueue {
#[must_use]
pub fn bound_address(&self) -> Option<String> {
self.lifecycle.resolved.get().cloned()
}
#[must_use]
pub fn publisher(&self) -> ZmqQueuePublisher {
ZmqQueuePublisher {
cell: Arc::clone(&self.cell),
push: Arc::new(Mutex::new(None)),
}
}
}
impl ConnectedBroker for ConnectedZmqQueue {
type Error = ZmqError;
type Closed = ();
async fn shutdown(self) -> Result<(), Self::Error> {
self.lifecycle
.closed
.store(true, std::sync::atomic::Ordering::Release);
Ok(())
}
}
impl Subscribe for ConnectedZmqQueue {
type Subscriber = ZmqSubscriber;
async fn subscribe(&self, name: &str) -> Result<Self::Subscriber, Self::Error> {
self.lifecycle.ensure_open()?;
let mut socket = PullSocket::new();
self.lifecycle.attach_receiver(&mut socket).await?;
let (tx, rx) = mpsc::unbounded_channel();
let task = tokio::spawn(async move {
loop {
match socket.recv().await {
Ok(message) => {
let item =
wire::decode(message).map(|(name, headers, payload)| ZmqMessage {
name,
headers,
payload,
});
if tx.send(item).is_err() {
break;
}
}
Err(err) => {
if tx.send(Err(ZmqError::Receive(err.to_string()))).is_err() {
break;
}
}
}
}
});
Ok(ZmqSubscriber {
name: name.to_owned(),
rx,
_driver: DriverHandle { task },
})
}
}
pub struct ZmqSubscriber {
name: String,
rx: mpsc::UnboundedReceiver<Result<ZmqMessage, ZmqError>>,
_driver: DriverHandle,
}
impl std::fmt::Debug for ZmqSubscriber {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ZmqSubscriber")
.field("name", &self.name)
.finish_non_exhaustive()
}
}
impl ZmqSubscriber {
pub(crate) fn from_parts(
name: String,
rx: mpsc::UnboundedReceiver<Result<ZmqMessage, ZmqError>>,
driver: DriverHandle,
) -> Self {
Self {
name,
rx,
_driver: driver,
}
}
}
impl Subscriber for ZmqSubscriber {
type Message = ZmqMessage;
type Error = ZmqError;
fn stream(&mut self) -> impl Stream<Item = Result<ZmqMessage, ZmqError>> + Send + '_ {
futures::stream::poll_fn(move |cx| self.rx.poll_recv(cx))
}
}
#[derive(Clone)]
pub struct ZmqQueuePublisher {
cell: Arc<OnceCell<SharedLifecycle>>,
push: Arc<Mutex<Option<PushSocket>>>,
}
impl std::fmt::Debug for ZmqQueuePublisher {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ZmqQueuePublisher").finish_non_exhaustive()
}
}
impl Publisher for ZmqQueuePublisher {
type Error = ZmqError;
#[allow(clippy::significant_drop_tightening)]
async fn publish(&self, msg: OutgoingMessage<'_>) -> Result<(), Self::Error> {
let lifecycle = self.cell.get().ok_or(ZmqError::NotConnected)?;
lifecycle.ensure_open()?;
let mut push = self.push.lock().await;
if push.is_none() {
let mut socket = PushSocket::new();
lifecycle.attach_sender(&mut socket).await?;
*push = Some(socket);
}
let socket = push.as_mut().expect("just attached");
send_with_retry(
socket,
msg.name(),
wire::encode(msg.name(), msg.headers(), msg.payload()),
)
.await
}
}
#[derive(Debug, Clone, Copy, Default)]
#[must_use]
pub struct ZmqQueuePublish;
impl PublishPolicy<ConnectedZmqQueue> for ZmqQueuePublish {
type Live = ZmqQueuePublisher;
async fn pair(self, connected: &ConnectedZmqQueue) -> Result<Self::Live, PairError> {
Ok(connected.publisher())
}
}
impl DefaultPublish for ConnectedZmqQueue {
type Policy = ZmqQueuePublish;
}