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::{message_from_wire, strip_address_prefix, Message};
use super::{retryable, MessagePublisher, TransportError};
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;
fn transport_name(&self) -> &'static str {
"nats"
}
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 headers: Vec<(String, String)> = raw
.headers
.as_ref()
.into_iter()
.flat_map(|headers| headers.iter())
.filter_map(|(key, values)| {
values
.last()
.map(|value| (key.to_string(), value.to_string()))
})
.collect();
let message = message_from_wire(
name,
payload,
Some(MESSAGE_ID_HEADER),
MESSAGE_KIND_HEADER,
headers,
);
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
}
}