use std::fmt::{Debug, Formatter};
use std::sync::Arc;
use async_nats::jetstream::Context;
use async_nats::jetstream::message::PublishMessage;
use bytes::Bytes;
use ruststream::{OutgoingMessage, PairError, PublishPolicy, Publisher};
use crate::broker::{ConnectedNatsBroker, NatsConnection};
use crate::publisher::NatsPublishPolicy;
use crate::{convert::headers_to_nats, error::NatsError};
pub use async_nats::jetstream::publish::PublishAck;
#[derive(Debug, Clone, Default, PartialEq, Eq)]
#[must_use]
pub struct JetStreamPublish {
stream: Option<String>,
last_sequence: Option<u64>,
last_subject_sequence: Option<u64>,
last_message_id: Option<String>,
}
impl JetStreamPublish {
pub fn expect_stream(mut self, stream: impl Into<String>) -> Self {
self.stream = Some(stream.into());
self
}
pub const fn expect_last_sequence(mut self, sequence: u64) -> Self {
self.last_sequence = Some(sequence);
self
}
pub const fn expect_last_subject_sequence(mut self, sequence: u64) -> Self {
self.last_subject_sequence = Some(sequence);
self
}
pub fn expect_last_message_id(mut self, id: impl Into<String>) -> Self {
self.last_message_id = Some(id.into());
self
}
fn apply(&self, mut message: PublishMessage) -> PublishMessage {
if let Some(stream) = &self.stream {
message = message.expected_stream(stream);
}
if let Some(sequence) = self.last_sequence {
message = message.expected_last_sequence(sequence);
}
if let Some(sequence) = self.last_subject_sequence {
message = message.expected_last_subject_sequence(sequence);
}
if let Some(id) = &self.last_message_id {
message = message.expected_last_message_id(id);
}
message
}
}
impl PublishPolicy<ConnectedNatsBroker> for JetStreamPublish {
type Live = JetStreamPublisher;
async fn pair(self, connected: &ConnectedNatsBroker) -> Result<Self::Live, PairError> {
Ok(self.bind(connected))
}
}
impl NatsPublishPolicy for JetStreamPublish {
fn bind(self, connected: &ConnectedNatsBroker) -> Self::Live {
JetStreamPublisher {
connection: Arc::clone(connected.connection()),
context: connected.jetstream(),
policy: self,
}
}
}
#[derive(Clone)]
pub struct JetStreamPublisher {
connection: Arc<NatsConnection>,
context: Context,
policy: JetStreamPublish,
}
impl Debug for JetStreamPublisher {
fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
f.debug_struct("JetStreamPublisher")
.field("policy", &self.policy)
.finish_non_exhaustive()
}
}
impl JetStreamPublisher {
pub async fn publish_ack(&self, msg: OutgoingMessage<'_>) -> Result<PublishAck, NatsError> {
self.connection.live_client(msg.name())?;
let mut message = PublishMessage::build().payload(Bytes::copy_from_slice(msg.payload()));
if let Some(headers) = headers_to_nats(msg.headers()) {
message = message.headers(headers);
}
self.context
.send_publish(msg.name().to_owned(), self.policy.apply(message))
.await
.map_err(|err| NatsError::Publish(Box::new(err)))?
.await
.map_err(|err| NatsError::JetStream(Box::new(err)))
}
}
impl Publisher for JetStreamPublisher {
type Error = NatsError;
async fn publish(&self, msg: OutgoingMessage<'_>) -> Result<(), Self::Error> {
self.publish_ack(msg).await.map(|_ack| ())
}
}