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::{AsyncMessageSource, ReceivedMessage};
use super::{AsyncMessagePublisher, TransportError};
use super::{Message, MessageKind};
const MESSAGE_ID_HEADER: &str = "Nats-Msg-Id";
const MESSAGE_KIND_HEADER: &str = "X-Sourced-Kind";
fn retryable(context: &str, err: impl std::fmt::Display) -> TransportError {
TransportError::retryable(format!("{context}: {err}"))
}
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 AsyncMessagePublisher for NatsPublisher {
async fn publish(&self, 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, kind_str(message.kind));
for (key, value) in &message.metadata {
headers.insert(key.as_str(), value.as_str());
}
let ack_future = self
.jetstream
.publish_with_headers(subject, headers, message.payload.clone().into())
.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 AsyncMessageSource 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 subject = raw.subject.to_string();
let name = match strip_prefix {
Some(prefix) => subject
.strip_prefix(prefix)
.map(str::to_string)
.unwrap_or(subject),
None => subject,
};
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 = kind_from_str(&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
}
}
fn kind_str(kind: MessageKind) -> &'static str {
match kind {
MessageKind::Command => "command",
MessageKind::Event => "event",
}
}
fn kind_from_str(value: &str) -> MessageKind {
match value {
"command" => MessageKind::Command,
_ => MessageKind::Event,
}
}