use std::time::Duration;
use async_nats::jetstream::consumer::pull::Config as PullConfig;
use async_nats::jetstream::consumer::Consumer;
use async_nats::jetstream::stream::Config as StreamConfig;
use async_nats::jetstream::{self, AckKind};
use futures::StreamExt;
use super::source::{MessageSource, ReceivedMessage};
use super::{retryable, MessagePublisher, TransportError};
use super::{strip_address_prefix, Message, MessageKind};
const MESSAGE_ID_HEADER: &str = "Nats-Msg-Id";
const MESSAGE_KIND_HEADER: &str = "X-Sourced-Kind";
pub struct NatsPublisher {
jetstream: jetstream::Context,
subject_prefix: Option<String>,
}
impl NatsPublisher {
pub fn new(jetstream: jetstream::Context) -> Self {
Self {
jetstream,
subject_prefix: None,
}
}
pub async fn connect(url: &str) -> Result<Self, TransportError> {
let client = async_nats::connect(url)
.await
.map_err(|err| retryable("nats connect", err))?;
Ok(Self::new(jetstream::new(client)))
}
pub fn with_subject_prefix(mut self, prefix: impl Into<String>) -> Self {
self.subject_prefix = Some(prefix.into());
self
}
fn subject(&self, message: &Message) -> String {
match &self.subject_prefix {
Some(prefix) => format!("{prefix}.{}", message.name()),
None => message.name().to_string(),
}
}
}
impl MessagePublisher for NatsPublisher {
async fn publish(&self, mut message: Message) -> Result<(), TransportError> {
let subject = self.subject(&message);
let mut headers = async_nats::HeaderMap::new();
if let Some(id) = message.id() {
headers.insert(MESSAGE_ID_HEADER, id);
}
headers.insert(MESSAGE_KIND_HEADER, message.kind.as_str());
for (key, value) in &message.metadata {
headers.insert(key.as_str(), value.as_str());
}
let payload = std::mem::take(&mut message.payload).into();
let ack_future = self
.jetstream
.publish_with_headers(subject, headers, payload)
.await
.map_err(|err| retryable("nats publish", err))?;
ack_future
.await
.map_err(|err| retryable("nats publish ack", err))?;
Ok(())
}
}
pub struct NatsJetStreamSource {
consumer: Consumer<PullConfig>,
fetch_timeout: Duration,
strip_prefix: Option<String>,
}
impl NatsJetStreamSource {
pub fn new(consumer: Consumer<PullConfig>) -> Self {
Self {
consumer,
fetch_timeout: Duration::from_millis(500),
strip_prefix: None,
}
}
pub fn with_fetch_timeout(mut self, timeout: Duration) -> Self {
self.fetch_timeout = timeout;
self
}
pub fn with_strip_prefix(mut self, prefix: impl Into<String>) -> Self {
self.strip_prefix = Some(prefix.into());
self
}
pub async fn connect(
url: &str,
stream_name: &str,
subjects: Vec<String>,
durable: &str,
) -> Result<Self, TransportError> {
let client = async_nats::connect(url)
.await
.map_err(|err| retryable("nats connect", err))?;
let jetstream = jetstream::new(client);
Self::from_context(&jetstream, stream_name, subjects, durable).await
}
pub async fn from_context(
jetstream: &jetstream::Context,
stream_name: &str,
subjects: Vec<String>,
durable: &str,
) -> Result<Self, TransportError> {
let stream = jetstream
.get_or_create_stream(StreamConfig {
name: stream_name.to_string(),
subjects,
..Default::default()
})
.await
.map_err(|err| retryable("nats get_or_create_stream", err))?;
let consumer = stream
.get_or_create_consumer(
durable,
PullConfig {
durable_name: Some(durable.to_string()),
..Default::default()
},
)
.await
.map_err(|err| retryable("nats get_or_create_consumer", err))?;
Ok(Self::new(consumer))
}
}
impl MessageSource for NatsJetStreamSource {
type Received = NatsReceived;
async fn recv(&mut self) -> Result<Option<Self::Received>, TransportError> {
let mut batch = self
.consumer
.batch()
.max_messages(1)
.expires(self.fetch_timeout)
.messages()
.await
.map_err(|err| retryable("nats fetch", err))?;
match batch.next().await {
Some(Ok(message)) => Ok(Some(NatsReceived::from_jetstream(
message,
self.strip_prefix.as_deref(),
))),
Some(Err(err)) => Err(retryable("nats batch message", err)),
None => Ok(None),
}
}
}
pub struct NatsReceived {
raw: jetstream::Message,
message: Message,
}
impl NatsReceived {
fn from_jetstream(raw: jetstream::Message, strip_prefix: Option<&str>) -> Self {
let name = strip_address_prefix(raw.subject.to_string(), strip_prefix);
let payload = raw.payload.to_vec();
let mut id = None;
let mut kind = MessageKind::Event;
let mut metadata = Vec::new();
if let Some(headers) = raw.headers.as_ref() {
for (key, values) in headers.iter() {
let key = key.to_string();
if let Some(value) = values.last() {
let value = value.to_string();
match key.as_str() {
MESSAGE_ID_HEADER => id = Some(value),
MESSAGE_KIND_HEADER => kind = MessageKind::from_str_lossy(&value),
_ => metadata.push((key, value)),
}
}
}
}
let mut message = Message::new(name, kind, payload);
message.id = id;
message.metadata = metadata;
Self { raw, message }
}
async fn settle(self, kind: AckKind) -> Result<(), TransportError> {
match kind {
AckKind::Ack => self
.raw
.ack()
.await
.map_err(|err| retryable("nats ack", err)),
other => self
.raw
.ack_with(other)
.await
.map_err(|err| retryable("nats ack_with", err)),
}
}
}
impl ReceivedMessage for NatsReceived {
fn message(&self) -> &Message {
&self.message
}
async fn ack(self) -> Result<(), TransportError> {
self.settle(AckKind::Ack).await
}
async fn nack(self, _reason: &str) -> Result<(), TransportError> {
self.settle(AckKind::Nak(None)).await
}
async fn dead_letter(self, _reason: &str) -> Result<(), TransportError> {
self.settle(AckKind::Term).await
}
async fn park(self, _reason: &str) -> Result<(), TransportError> {
self.settle(AckKind::Term).await
}
}