pub use crate::errors::{ConsumerError, HandlerError, PublisherError};
pub use crate::outcomes::{Handled, Received, ReceivedBatch, Sent, SentBatch};
use crate::CanonicalMessage;
use anyhow::anyhow;
use async_trait::async_trait;
pub use futures::future::BoxFuture;
use std::any::Any;
use std::sync::Arc;
use tracing::warn;
#[derive(Default, Debug, Clone)]
#[allow(clippy::large_enum_variant)]
pub enum MessageDisposition {
#[default]
Ack,
Reply(CanonicalMessage),
Nack,
}
impl From<Option<CanonicalMessage>> for MessageDisposition {
fn from(opt: Option<CanonicalMessage>) -> Self {
match opt {
Some(msg) => MessageDisposition::Reply(msg),
None => MessageDisposition::Ack,
}
}
}
impl From<Handled> for MessageDisposition {
fn from(handled: Handled) -> Self {
match handled {
Handled::Ack => MessageDisposition::Ack,
Handled::Publish(msg) => MessageDisposition::Reply(msg),
}
}
}
#[async_trait]
pub trait Handler: Send + Sync + 'static {
async fn handle(&self, msg: CanonicalMessage) -> Result<Handled, HandlerError>;
async fn handle_many(&self, msgs: Vec<CanonicalMessage>) -> Vec<Result<Handled, HandlerError>> {
let mut results = Vec::with_capacity(msgs.len());
let mut remaining = msgs.len();
for msg in msgs {
remaining -= 1;
let result = self.handle(msg).await;
let aborted = match &result {
Err(HandlerError::Retryable(_)) => Some("retryable"),
Err(HandlerError::Connection(_)) => Some("connection"),
Err(HandlerError::NonRetryable(_)) => Some("non-retryable"),
Ok(_) => None,
};
results.push(result);
if let Some(kind) = aborted {
for _ in 0..remaining {
results.push(Err(match kind {
"retryable" => HandlerError::Retryable(anyhow!(
"batch aborted after earlier retryable handler failure"
)),
"connection" => HandlerError::Connection(anyhow!(
"batch aborted after earlier handler connection failure"
)),
_ => HandlerError::NonRetryable(anyhow!(
"batch aborted after earlier non-retryable handler failure"
)),
}));
}
break;
}
}
results
}
fn register_handler(
&self,
_type_name: &str,
_handler: Arc<dyn Handler>,
) -> Option<Arc<dyn Handler>> {
None
}
}
#[async_trait]
impl<T: Handler + ?Sized> Handler for Arc<T> {
async fn handle(&self, msg: CanonicalMessage) -> Result<Handled, HandlerError> {
(**self).handle(msg).await
}
async fn handle_many(&self, msgs: Vec<CanonicalMessage>) -> Vec<Result<Handled, HandlerError>> {
(**self).handle_many(msgs).await
}
fn register_handler(
&self,
type_name: &str,
handler: Arc<dyn Handler>,
) -> Option<Arc<dyn Handler>> {
(**self).register_handler(type_name, handler)
}
}
pub trait AsyncHandler: Send + Sync + 'static {
fn handle<'a>(&'a self, msg: CanonicalMessage) -> BoxFuture<'a, Result<Handled, HandlerError>>;
}
pub struct SimpleHandler<T>(pub T);
#[async_trait]
impl<T: AsyncHandler> Handler for SimpleHandler<T> {
async fn handle(&self, msg: CanonicalMessage) -> Result<Handled, HandlerError> {
self.0.handle(msg).await
}
}
pub type CommitFunc =
Box<dyn FnOnce(MessageDisposition) -> BoxFuture<'static, anyhow::Result<()>> + Send + 'static>;
pub type BatchCommitFunc = Box<
dyn FnOnce(Vec<MessageDisposition>) -> BoxFuture<'static, anyhow::Result<()>> + Send + 'static,
>;
#[derive(Debug, Clone, serde::Serialize)]
pub struct EndpointStatus {
pub healthy: bool,
pub target: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub pending: Option<usize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub capacity: Option<usize>,
#[serde(skip_serializing_if = "Option::is_none")]
pub error: Option<String>,
pub details: serde_json::Value,
}
impl Default for EndpointStatus {
fn default() -> Self {
Self {
healthy: true,
target: String::new(),
pending: None,
capacity: None,
error: None,
details: serde_json::Value::Null,
}
}
}
pub(crate) fn drain_idle_timeout() -> std::time::Duration {
static V: std::sync::OnceLock<std::time::Duration> = std::sync::OnceLock::new();
*V.get_or_init(|| {
std::env::var("MQ_BRIDGE_DRAIN_IDLE_TIMEOUT_MS")
.ok()
.and_then(|s| s.parse::<u64>().ok())
.map(std::time::Duration::from_millis)
.unwrap_or(std::time::Duration::from_millis(1000))
})
}
pub(crate) async fn drain_gated<F: std::future::Future>(
exit_on_empty: bool,
fut: F,
) -> Option<F::Output> {
if exit_on_empty {
tokio::time::timeout(drain_idle_timeout(), fut).await.ok()
} else {
Some(fut.await)
}
}
#[async_trait]
pub trait MessageConsumer: Send + Sync {
fn on_connect_hook(&self) -> Option<BoxFuture<'_, anyhow::Result<()>>> {
None
}
fn on_disconnect_hook(&self) -> Option<BoxFuture<'_, anyhow::Result<()>>> {
None
}
async fn receive_batch(&mut self, _max_messages: usize)
-> Result<ReceivedBatch, ConsumerError>;
async fn receive(&mut self) -> Result<Received, ConsumerError> {
loop {
let mut batch = self.receive_batch(1).await?;
if let Some(msg) = batch.messages.pop() {
debug_assert!(batch.messages.is_empty());
if !batch.messages.is_empty() {
tracing::error!(
"receive_batch(1) returned {} extra messages; dropping them (implementation bug)",
batch.messages.len()
);
}
return Ok(Received {
message: msg,
commit: into_commit_func(batch.commit),
});
}
tokio::time::sleep(std::time::Duration::from_millis(1)).await;
tokio::task::yield_now().await;
}
}
async fn receive_batch_helper(
&mut self,
_max_messages: usize,
) -> Result<ReceivedBatch, ConsumerError> {
let received = self.receive().await?; let batch_commit = Box::new(move |dispositions: Vec<MessageDisposition>| {
let single_disposition = dispositions
.into_iter()
.next()
.unwrap_or(MessageDisposition::Ack);
(received.commit)(single_disposition)
}) as BatchCommitFunc;
Ok(ReceivedBatch {
messages: vec![received.message],
commit: batch_commit,
})
}
fn set_exit_on_empty(&mut self, _exit_on_empty: bool) {}
fn commit_requires_order(&self) -> bool {
true
}
async fn status(&self) -> EndpointStatus {
EndpointStatus {
healthy: true,
..Default::default()
}
}
async fn close(&mut self) -> anyhow::Result<()> {
if let Some(hook) = self.on_disconnect_hook() {
hook.await?;
}
Ok(())
}
fn as_any(&self) -> &dyn Any;
}
#[async_trait]
pub trait MessagePublisher: Send + Sync + 'static {
fn on_connect_hook(&self) -> Option<BoxFuture<'_, anyhow::Result<()>>> {
None
}
fn on_disconnect_hook(&self) -> Option<BoxFuture<'_, anyhow::Result<()>>> {
None
}
async fn send_batch(
&self,
messages: Vec<CanonicalMessage>,
) -> Result<SentBatch, PublisherError>;
async fn send(&self, message: CanonicalMessage) -> Result<Sent, PublisherError> {
let message_id = message.message_id;
let expects_reply = message.metadata.contains_key("reply_to");
match self.send_batch(vec![message]).await {
Ok(SentBatch::Ack) => {
if expects_reply {
warn!("Message {:032x} expected a reply (reply_to set), but publisher returned Ack. Response loop might be broken.", message_id);
}
Ok(Sent::Ack)
}
Ok(SentBatch::Partial {
mut responses,
mut failed,
}) => {
if let Some((_, err)) = failed.pop() {
Err(err)
} else if let Some(res) = responses.as_mut().and_then(|r| r.pop()) {
Ok(Sent::Response(res))
} else {
if expects_reply {
warn!("Message {:032x} expected a reply (reply_to set), but publisher returned Ack. Response loop might be broken.", message_id);
}
Ok(Sent::Ack)
}
}
Err(e) => Err(e),
}
}
async fn flush(&self) -> anyhow::Result<()> {
Ok(())
}
async fn status(&self) -> EndpointStatus {
EndpointStatus {
healthy: true,
..Default::default()
}
}
fn as_any(&self) -> &dyn Any;
}
#[async_trait]
impl<T: MessagePublisher + ?Sized> MessagePublisher for Arc<T> {
fn on_connect_hook(&self) -> Option<BoxFuture<'_, anyhow::Result<()>>> {
(**self).on_connect_hook()
}
fn on_disconnect_hook(&self) -> Option<BoxFuture<'_, anyhow::Result<()>>> {
(**self).on_disconnect_hook()
}
async fn send(&self, message: CanonicalMessage) -> Result<Sent, PublisherError> {
(**self).send(message).await
}
async fn send_batch(
&self,
messages: Vec<CanonicalMessage>,
) -> Result<SentBatch, PublisherError> {
(**self).send_batch(messages).await
}
async fn flush(&self) -> anyhow::Result<()> {
(**self).flush().await
}
async fn status(&self) -> EndpointStatus {
(**self).status().await
}
fn as_any(&self) -> &dyn Any {
(**self).as_any()
}
}
#[async_trait]
impl<T: MessagePublisher + ?Sized> MessagePublisher for Box<T> {
fn on_connect_hook(&self) -> Option<BoxFuture<'_, anyhow::Result<()>>> {
(**self).on_connect_hook()
}
fn on_disconnect_hook(&self) -> Option<BoxFuture<'_, anyhow::Result<()>>> {
(**self).on_disconnect_hook()
}
async fn send(&self, message: CanonicalMessage) -> Result<Sent, PublisherError> {
(**self).send(message).await
}
async fn send_batch(
&self,
messages: Vec<CanonicalMessage>,
) -> Result<SentBatch, PublisherError> {
(**self).send_batch(messages).await
}
async fn flush(&self) -> anyhow::Result<()> {
(**self).flush().await
}
async fn status(&self) -> EndpointStatus {
(**self).status().await
}
fn as_any(&self) -> &dyn Any {
(**self).as_any()
}
}
#[async_trait]
pub trait CustomEndpointFactory: Send + Sync + std::fmt::Debug {
async fn create_consumer(
&self,
_route_name: &str,
_config: &serde_json::Value,
) -> anyhow::Result<Box<dyn MessageConsumer>> {
Err(anyhow::anyhow!(
"This custom endpoint does not support creating consumers"
))
}
async fn create_publisher(
&self,
_route_name: &str,
_config: &serde_json::Value,
) -> anyhow::Result<Box<dyn MessagePublisher>> {
Err(anyhow::anyhow!(
"This custom endpoint does not support creating publishers"
))
}
}
#[async_trait]
pub trait CustomMiddlewareFactory: Send + Sync + std::fmt::Debug {
async fn apply_consumer(
&self,
consumer: Box<dyn MessageConsumer>,
_route_name: &str,
_config: &serde_json::Value,
) -> anyhow::Result<Box<dyn MessageConsumer>> {
Ok(consumer)
}
async fn apply_publisher(
&self,
publisher: Box<dyn MessagePublisher>,
_route_name: &str,
_config: &serde_json::Value,
) -> anyhow::Result<Box<dyn MessagePublisher>> {
Ok(publisher)
}
}
pub const SEND_BATCH_CONCURRENCY: usize = 128;
pub async fn send_batch_helper<P: MessagePublisher + ?Sized>(
publisher: &P,
messages: Vec<CanonicalMessage>,
callback: impl for<'a> Fn(&'a P, CanonicalMessage) -> BoxFuture<'a, Result<Sent, PublisherError>>
+ Send
+ Sync,
) -> Result<SentBatch, PublisherError> {
use futures::stream::StreamExt;
let mut responses = Vec::new();
let mut failed_messages = Vec::new();
let callback = &callback;
let mut results = futures::stream::iter(messages.into_iter().enumerate().map(
|(idx, msg)| async move {
let result = callback(publisher, msg.clone()).await;
(idx, msg, result)
},
))
.buffer_unordered(SEND_BATCH_CONCURRENCY);
while let Some((idx, msg, result)) = results.next().await {
match result {
Ok(Sent::Response(resp)) => responses.push((idx, resp)),
Ok(Sent::Ack) => {}
Err(e) => failed_messages.push((idx, msg, e)),
}
}
responses.sort_by_key(|(idx, _)| *idx);
let responses: Vec<_> = responses.into_iter().map(|(_, resp)| resp).collect();
failed_messages.sort_by_key(|(idx, _, _)| *idx);
let failed_messages: Vec<_> = failed_messages
.into_iter()
.map(|(_, msg, err)| (msg, err))
.collect();
if failed_messages.is_empty() && responses.is_empty() {
Ok(SentBatch::Ack)
} else {
Ok(SentBatch::Partial {
responses: if responses.is_empty() {
None
} else {
Some(responses)
},
failed: failed_messages,
})
}
}
pub fn into_commit_func(batch_commit: BatchCommitFunc) -> CommitFunc {
Box::new(move |disposition: MessageDisposition| {
let batch_disposition = vec![disposition];
batch_commit(batch_disposition)
})
}
pub fn into_batch_commit_func(commit: CommitFunc) -> BatchCommitFunc {
Box::new(move |mut dispositions: Vec<MessageDisposition>| {
let single_disposition = if dispositions.len() > 1 {
warn!(
"into_batch_commit_func called with batch of {} messages; dropping all responses to avoid partial commit (incorrect usage)",
dispositions.len()
);
MessageDisposition::Ack
} else {
dispositions.pop().unwrap_or(MessageDisposition::Ack)
};
commit(single_disposition)
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::CanonicalMessage;
use anyhow::anyhow;
use std::sync::{
atomic::{AtomicUsize, Ordering},
Arc,
};
struct MockPublisher;
#[async_trait]
impl MessagePublisher for MockPublisher {
async fn send_batch(
&self,
_msgs: Vec<CanonicalMessage>,
) -> Result<SentBatch, PublisherError> {
Ok(SentBatch::Ack)
}
fn as_any(&self) -> &dyn Any {
self
}
}
#[tokio::test]
async fn test_send_batch_helper_partial_failure() {
let publisher = MockPublisher;
let msgs = vec![
CanonicalMessage::from("1"),
CanonicalMessage::from("2"),
CanonicalMessage::from("3"),
];
let result = send_batch_helper(&publisher, msgs.clone(), |_pub, msg| {
Box::pin(async move {
let payload = msg.get_payload_str();
if payload == "1" {
Ok(Sent::Response(CanonicalMessage::from("resp1")))
} else if payload == "2" {
Err(PublisherError::Retryable(anyhow!("fail")))
} else {
Ok(Sent::Ack)
}
})
})
.await;
match result {
Ok(SentBatch::Partial { responses, failed }) => {
assert!(responses.is_some());
let resps = responses.unwrap();
assert_eq!(resps.len(), 1);
assert_eq!(resps[0].get_payload_str(), "resp1");
assert_eq!(failed.len(), 1);
assert_eq!(failed[0].0.get_payload_str(), "2");
assert!(matches!(failed[0].1, PublisherError::Retryable(_)));
}
_ => panic!("Expected Partial result"),
}
}
#[tokio::test]
async fn test_send_batch_helper_preserves_response_order() {
let publisher = MockPublisher;
let count = 16u64;
let msgs: Vec<CanonicalMessage> = (0..count)
.map(|i| CanonicalMessage::from(i.to_string()))
.collect();
let result = send_batch_helper(&publisher, msgs, |_pub, msg| {
Box::pin(async move {
let i: u64 = msg.get_payload_str().parse().unwrap();
tokio::time::sleep(std::time::Duration::from_millis((count - i) * 2)).await;
let mut resp = CanonicalMessage::from(msg.get_payload_str().to_string());
resp.message_id = msg.message_id;
Ok(Sent::Response(resp))
})
})
.await
.unwrap();
match result {
SentBatch::Partial { responses, failed } => {
assert!(failed.is_empty());
let responses = responses.expect("expected responses");
let order: Vec<u64> = responses
.iter()
.map(|r| r.get_payload_str().parse().unwrap())
.collect();
assert_eq!(
order,
(0..count).collect::<Vec<u64>>(),
"send_batch_helper must preserve input order",
);
}
SentBatch::Ack => panic!("expected per-message responses"),
}
}
#[tokio::test]
async fn test_send_batch_helper_keeps_pipeline_full_when_early_send_is_slow() {
let publisher = Arc::new(MockPublisher);
let total = SEND_BATCH_CONCURRENCY + 1;
let msgs: Vec<CanonicalMessage> = (0..total)
.map(|i| CanonicalMessage::from(i.to_string()))
.collect();
let started = Arc::new(AtomicUsize::new(0));
let all_started = Arc::new(tokio::sync::Notify::new());
let release_first = Arc::new(tokio::sync::Notify::new());
let helper = tokio::spawn({
let publisher = Arc::clone(&publisher);
let started = Arc::clone(&started);
let all_started = Arc::clone(&all_started);
let release_first = Arc::clone(&release_first);
async move {
send_batch_helper(&publisher, msgs, |_pub, msg| {
let started = Arc::clone(&started);
let all_started = Arc::clone(&all_started);
let release_first = Arc::clone(&release_first);
Box::pin(async move {
let idx: usize = msg.get_payload_str().parse().unwrap();
if started.fetch_add(1, Ordering::SeqCst) + 1 == total {
all_started.notify_waiters();
}
if idx == 0 {
release_first.notified().await;
}
let mut resp = CanonicalMessage::from(idx.to_string());
resp.message_id = msg.message_id;
Ok(Sent::Response(resp))
})
})
.await
}
});
tokio::time::timeout(std::time::Duration::from_millis(200), async {
loop {
let notified = all_started.notified();
tokio::pin!(notified);
notified.as_mut().enable();
if started.load(Ordering::SeqCst) == total {
break;
}
notified.await;
}
})
.await
.expect("a completed later send should free a slot even while the first send is blocked");
release_first.notify_waiters();
let result = helper.await.unwrap().unwrap();
match result {
SentBatch::Partial { responses, failed } => {
assert!(failed.is_empty());
let order: Vec<usize> = responses
.expect("expected responses")
.iter()
.map(|r| r.get_payload_str().parse().unwrap())
.collect();
assert_eq!(order, (0..total).collect::<Vec<_>>());
}
SentBatch::Ack => panic!("expected per-message responses"),
}
}
#[tokio::test]
async fn test_send_propagates_single_error() {
struct FailPublisher;
#[async_trait]
impl MessagePublisher for FailPublisher {
async fn send_batch(
&self,
msgs: Vec<CanonicalMessage>,
) -> Result<SentBatch, PublisherError> {
Ok(SentBatch::Partial {
responses: None,
failed: vec![(
msgs[0].clone(),
PublisherError::NonRetryable(anyhow!("inner")),
)],
})
}
fn as_any(&self) -> &dyn Any {
self
}
}
let publ = FailPublisher;
let res = publ.send(CanonicalMessage::from("test")).await;
assert!(res.is_err());
match res.unwrap_err() {
PublisherError::NonRetryable(e) => assert_eq!(e.to_string(), "inner"),
_ => panic!("Expected NonRetryable error"),
}
}
#[tokio::test]
async fn test_simple_handler_wrapper() {
struct MyLogic;
impl AsyncHandler for MyLogic {
fn handle<'a>(
&'a self,
_msg: CanonicalMessage,
) -> BoxFuture<'a, Result<Handled, HandlerError>> {
Box::pin(async { Ok(Handled::Ack) })
}
}
let handler = SimpleHandler(MyLogic);
let res = handler.handle(CanonicalMessage::from("test")).await;
assert!(matches!(res, Ok(Handled::Ack)));
}
}