mod capability;
use capability::SeekControl;
pub use capability::{
MemoryPosition, MemoryRequester, MemorySeeker, MemoryTransaction, PARTITION_KEY_HEADER,
RequestError,
};
use std::{
collections::HashMap,
convert::Infallible,
fmt,
sync::{Arc, Mutex, OnceLock, atomic::AtomicU64},
task::Poll,
time::Duration,
};
#[cfg(feature = "testing")]
use crate::testing::coordinator::Coordinator;
use crate::{
AckError, Broker, ConnectedBroker, DefaultPublish, DescribeServer, Headers, IncomingMessage,
OutgoingMessage, PairError, PublishPolicy, Publisher, RawMessage, ServerSpec, Subscribe,
Subscriber, SubscriptionSource,
};
use bytes::Bytes;
use futures::Stream;
use thiserror::Error;
use tokio::sync::{Notify, mpsc};
use tokio::time::sleep;
type Sender = mpsc::UnboundedSender<MemoryDelivery>;
#[derive(Clone)]
struct MemoryOutbound {
name: String,
payload: Bytes,
headers: Headers,
}
#[derive(Clone)]
struct MemoryDelivery {
name: String,
payload: Bytes,
headers: Headers,
seq: usize,
}
enum Bus {
Live(HashMap<String, Vec<Sender>>),
ShutDown,
}
impl Default for Bus {
fn default() -> Self {
Self::Live(HashMap::new())
}
}
#[derive(Default)]
struct MemoryState {
subscribers: Mutex<Bus>,
published: Mutex<HashMap<String, Vec<RawMessage>>>,
notify: Notify,
inbox_seq: AtomicU64,
#[cfg(feature = "testing")]
coordinator: OnceLock<Coordinator>,
}
impl MemoryState {
fn register(&self, name: String, tx: Sender) -> Result<(), MemoryError> {
match &mut *self
.subscribers
.lock()
.expect("memory broker mutex poisoned")
{
Bus::Live(subscribers) => {
subscribers.entry(name).or_default().push(tx);
Ok(())
}
Bus::ShutDown => Err(MemoryError::ShutDown),
}
}
fn unregister(&self, name: &str) {
if let Bus::Live(subscribers) = &mut *self
.subscribers
.lock()
.expect("memory broker mutex poisoned")
{
subscribers.remove(name);
}
}
#[allow(clippy::significant_drop_tightening)]
fn fanout(&self, outbound: &MemoryOutbound) -> Result<(), MemoryError> {
{
let bus = self
.subscribers
.lock()
.expect("memory broker mutex poisoned");
let Bus::Live(subscribers) = &*bus else {
return Err(MemoryError::ShutDown);
};
let mut log = self.published.lock().expect("memory broker mutex poisoned");
let entries = log.entry(outbound.name.clone()).or_default();
let delivery = MemoryDelivery {
name: outbound.name.clone(),
payload: outbound.payload.clone(),
headers: outbound.headers.clone(),
seq: entries.len(),
};
entries.push(
RawMessage::new(outbound.name.clone(), outbound.payload.clone())
.with_headers(outbound.headers.clone()),
);
self.send_to(subscribers, &delivery);
}
self.notify.notify_waiters();
Ok(())
}
fn send_to(&self, subscribers: &HashMap<String, Vec<Sender>>, delivery: &MemoryDelivery) {
if let Some(senders) = subscribers.get(&delivery.name) {
for tx in senders {
let sent = tx.send(delivery.clone());
#[cfg(feature = "testing")]
if sent.is_ok() && !delivery.name.starts_with("_inbox.") {
if let Some(coordinator) = self.coordinator.get() {
coordinator.enqueued();
}
}
#[cfg(not(feature = "testing"))]
let _ = sent;
}
}
}
#[cfg(feature = "testing")]
fn install_coordinator(&self, coordinator: Coordinator) {
let _ = self.coordinator.set(coordinator);
}
#[cfg(feature = "testing")]
fn coordinator(&self) -> Option<Coordinator> {
self.coordinator.get().cloned()
}
}
#[derive(Clone, Default)]
pub struct MemoryBroker {
state: Arc<MemoryState>,
}
impl MemoryBroker {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn subscribe(&self, name: impl Into<String>) -> MemorySubscriber {
let (tx, rx) = mpsc::unbounded_channel();
let name = name.into();
let _ = self.state.register(name.clone(), tx.clone());
MemorySubscriber {
name,
rx,
requeue: tx,
batch_limit: DEFAULT_BATCH_LIMIT,
state: Arc::clone(&self.state),
seek: Arc::new(SeekControl::default()),
#[cfg(feature = "testing")]
coordinator: self.state.coordinator(),
}
}
#[must_use]
pub fn publisher(&self) -> MemoryPublisher {
MemoryPublisher {
state: Arc::clone(&self.state),
txn: Mutex::new(None),
}
}
#[must_use]
pub fn requester(&self) -> MemoryRequester {
MemoryRequester::new(Arc::clone(&self.state))
}
}
impl fmt::Debug for MemoryBroker {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("MemoryBroker").finish_non_exhaustive()
}
}
impl Broker for MemoryBroker {
type Error = MemoryError;
type Connected = ConnectedMemoryBroker;
async fn connect(self) -> Result<Self::Connected, Self::Error> {
{
let mut bus = self
.state
.subscribers
.lock()
.expect("memory broker mutex poisoned");
if matches!(*bus, Bus::ShutDown) {
*bus = Bus::Live(HashMap::new());
}
}
Ok(ConnectedMemoryBroker { state: self.state })
}
}
#[derive(Clone)]
pub struct ConnectedMemoryBroker {
state: Arc<MemoryState>,
}
impl fmt::Debug for ConnectedMemoryBroker {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ConnectedMemoryBroker")
.finish_non_exhaustive()
}
}
impl ConnectedMemoryBroker {
#[must_use]
pub fn publisher(&self) -> MemoryPublisher {
MemoryPublisher {
state: Arc::clone(&self.state),
txn: Mutex::new(None),
}
}
#[must_use]
pub fn requester(&self) -> MemoryRequester {
MemoryRequester::new(Arc::clone(&self.state))
}
}
impl ConnectedBroker for ConnectedMemoryBroker {
type Error = MemoryError;
type Closed = ClosedMemoryBroker;
async fn shutdown(self) -> Result<Self::Closed, Self::Error> {
let dropped = {
let mut bus = self
.state
.subscribers
.lock()
.expect("memory broker mutex poisoned");
match std::mem::replace(&mut *bus, Bus::ShutDown) {
Bus::Live(subscribers) => subscribers.values().map(Vec::len).sum(),
Bus::ShutDown => 0,
}
};
Ok(ClosedMemoryBroker {
subscribers_dropped: dropped,
})
}
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
#[must_use]
pub struct MemoryPublish;
impl PublishPolicy<ConnectedMemoryBroker> for MemoryPublish {
type Live = MemoryPublisher;
async fn pair(self, connected: &ConnectedMemoryBroker) -> Result<Self::Live, PairError> {
Ok(connected.publisher())
}
}
impl DefaultPublish for ConnectedMemoryBroker {
type Policy = MemoryPublish;
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
#[must_use]
pub struct MemoryRequest;
impl PublishPolicy<ConnectedMemoryBroker> for MemoryRequest {
type Live = MemoryRequester;
async fn pair(self, connected: &ConnectedMemoryBroker) -> Result<Self::Live, PairError> {
Ok(connected.requester())
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ClosedMemoryBroker {
subscribers_dropped: usize,
}
impl ClosedMemoryBroker {
#[must_use]
pub fn subscribers_dropped(&self) -> usize {
self.subscribers_dropped
}
}
impl DescribeServer for MemoryBroker {
fn describe_server(&self) -> ServerSpec {
ServerSpec::in_process("memory")
}
}
#[cfg(feature = "testing")]
impl crate::testing::TestableBroker for ConnectedMemoryBroker {
fn install_coordinator(&self, coordinator: Coordinator) {
self.state.install_coordinator(coordinator);
}
fn inject(&self, message: OutgoingMessage<'_>) {
self.state
.fanout(&MemoryOutbound {
name: message.name().to_owned(),
payload: Bytes::copy_from_slice(message.payload()),
headers: message.headers().clone(),
})
.expect("inject on a shut-down broker: drive the harness before shutdown");
}
fn published(&self, name: &str) -> Vec<RawMessage> {
self.state
.published
.lock()
.expect("memory broker mutex poisoned")
.get(name)
.cloned()
.unwrap_or_default()
}
}
#[cfg(feature = "testing")]
crate::register_testable_broker!(ConnectedMemoryBroker);
impl Subscribe for ConnectedMemoryBroker {
type Subscriber = MemorySubscriber;
async fn subscribe(&self, name: &str) -> Result<Self::Subscriber, Self::Error> {
let (tx, rx) = mpsc::unbounded_channel();
let name = name.to_owned();
self.state.register(name.clone(), tx.clone())?;
Ok(MemorySubscriber {
name,
rx,
requeue: tx,
batch_limit: DEFAULT_BATCH_LIMIT,
state: Arc::clone(&self.state),
seek: Arc::new(SeekControl::default()),
#[cfg(feature = "testing")]
coordinator: self.state.coordinator(),
})
}
}
#[derive(Debug, Clone)]
pub struct MemorySource {
name: String,
}
impl MemorySource {
#[must_use]
pub fn new(name: impl Into<String>) -> Self {
Self { name: name.into() }
}
}
impl SubscriptionSource<ConnectedMemoryBroker> for MemorySource {
type Subscriber = MemorySubscriber;
fn name(&self) -> &str {
&self.name
}
async fn subscribe(
self,
connected: &ConnectedMemoryBroker,
) -> Result<Self::Subscriber, MemoryError> {
Subscribe::subscribe(connected, &self.name).await
}
}
const DEFAULT_BATCH_LIMIT: usize = 64;
pub struct MemorySubscriber {
name: String,
rx: mpsc::UnboundedReceiver<MemoryDelivery>,
requeue: Sender,
batch_limit: usize,
state: Arc<MemoryState>,
seek: Arc<SeekControl>,
#[cfg(feature = "testing")]
coordinator: Option<Coordinator>,
}
impl MemorySubscriber {
pub fn set_batch_limit(&mut self, limit: usize) {
self.batch_limit = limit;
}
}
impl fmt::Debug for MemorySubscriber {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("MemorySubscriber")
.field("name", &self.name)
.finish_non_exhaustive()
}
}
impl Subscriber for MemorySubscriber {
type Message = MemoryMessage;
type Error = Infallible;
fn stream(&mut self) -> impl Stream<Item = Result<Self::Message, Self::Error>> + Send + '_ {
let requeue = self.requeue.clone();
#[cfg(feature = "testing")]
let coordinator = self.coordinator.clone();
futures::stream::poll_fn(move |cx| {
self.seek.waker.register(cx.waker());
self.apply_pending_seek();
loop {
match self.rx.poll_recv(cx) {
Poll::Ready(Some(delivery)) => {
if delivery.seq < self.seek.watermark() {
#[cfg(feature = "testing")]
if let Some(coordinator) = &coordinator {
coordinator.consumed();
}
continue;
}
return Poll::Ready(Some(Ok(MemoryMessage {
delivery: Some(delivery),
requeue: requeue.clone(),
#[cfg(feature = "testing")]
coordinator: coordinator.clone(),
})));
}
Poll::Ready(None) => return Poll::Ready(None),
Poll::Pending => return Poll::Pending,
}
}
})
}
}
pub struct MemoryPublisher {
state: Arc<MemoryState>,
txn: Mutex<Option<Vec<MemoryOutbound>>>,
}
impl Clone for MemoryPublisher {
fn clone(&self) -> Self {
Self {
state: Arc::clone(&self.state),
txn: Mutex::new(None),
}
}
}
impl fmt::Debug for MemoryPublisher {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("MemoryPublisher").finish_non_exhaustive()
}
}
#[derive(Debug, Error, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum MemoryError {
#[error("a transaction is already open on this publisher handle")]
TransactionBusy,
#[error("no transaction is open on this publisher handle")]
NoTransaction,
#[error("the memory broker is shut down")]
ShutDown,
}
impl Publisher for MemoryPublisher {
type Error = MemoryError;
async fn publish(&self, msg: OutgoingMessage<'_>) -> Result<(), Self::Error> {
let outbound = MemoryOutbound {
name: msg.name().to_owned(),
payload: Bytes::copy_from_slice(msg.payload()),
headers: msg.headers().clone(),
};
{
let mut txn = self.txn.lock().expect("memory broker mutex poisoned");
if let Some(buffered) = txn.as_mut() {
buffered.push(outbound);
return Ok(());
}
}
self.state.fanout(&outbound)
}
}
pub struct MemoryMessage {
delivery: Option<MemoryDelivery>,
requeue: Sender,
#[cfg(feature = "testing")]
coordinator: Option<Coordinator>,
}
#[cfg(feature = "testing")]
impl Drop for MemoryMessage {
fn drop(&mut self) {
if let Some(coordinator) = &self.coordinator {
coordinator.consumed();
}
}
}
impl fmt::Debug for MemoryMessage {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("MemoryMessage")
.field("name", &self.delivery.as_ref().map(|d| d.name.as_str()))
.finish_non_exhaustive()
}
}
impl MemoryMessage {
#[must_use]
pub fn name(&self) -> &str {
self.delivery
.as_ref()
.map(|d| d.name.as_str())
.unwrap_or_default()
}
#[must_use]
pub fn into_raw(mut self) -> RawMessage {
let delivery = self.delivery.take().expect("delivery already consumed");
RawMessage::new(delivery.name, delivery.payload).with_headers(delivery.headers)
}
}
impl IncomingMessage for MemoryMessage {
fn payload(&self) -> &[u8] {
self.delivery
.as_ref()
.map(|d| d.payload.as_ref())
.unwrap_or_default()
}
fn partition_key(&self) -> Option<&[u8]> {
crate::Partitioned::partition_key(self)
}
fn headers(&self) -> &Headers {
static EMPTY: OnceLock<Headers> = OnceLock::new();
self.delivery
.as_ref()
.map_or_else(|| EMPTY.get_or_init(Headers::new), |d| &d.headers)
}
async fn ack(mut self) -> Result<(), AckError> {
self.delivery.take();
Ok(())
}
async fn nack(mut self, requeue: bool) -> Result<(), AckError> {
let delivery = self.delivery.take().expect("delivery already consumed");
if requeue {
let sent = self.requeue.send(delivery);
#[cfg(feature = "testing")]
if sent.is_ok() {
if let Some(coordinator) = &self.coordinator {
coordinator.enqueued();
}
}
#[cfg(not(feature = "testing"))]
let _ = sent;
}
Ok(())
}
fn supports_nack_after(&self) -> bool {
true
}
async fn nack_after(mut self, delay: Duration) -> Result<(), AckError> {
let delivery = self.delivery.take().expect("delivery already consumed");
let requeue = self.requeue.clone();
#[cfg(feature = "testing")]
if let Some(coordinator) = self.coordinator.clone() {
let counter = coordinator.clone();
coordinator.schedule_redelivery(delay, move || {
if requeue.send(delivery).is_ok() {
counter.enqueued();
}
});
return Ok(());
}
tokio::spawn(async move {
sleep(delay).await;
let _ = requeue.send(delivery);
});
Ok(())
}
}
#[cfg(test)]
mod tests {
use futures::StreamExt;
use super::*;
#[tokio::test]
async fn debug_formats_and_message_accessors() {
let broker = MemoryBroker::new();
assert!(format!("{broker:?}").contains("MemoryBroker"));
let source = MemorySource::new("orders");
assert_eq!(source.name(), "orders");
let publisher = broker.publisher();
assert!(format!("{publisher:?}").contains("MemoryPublisher"));
let mut sub = broker.subscribe("dbg");
assert!(format!("{sub:?}").contains("MemorySubscriber"));
publisher
.publish(OutgoingMessage::new("dbg", b"payload".as_slice()))
.await
.unwrap();
let mut stream = std::pin::pin!(sub.stream());
let msg = stream.next().await.unwrap().unwrap();
assert!(format!("{msg:?}").contains("MemoryMessage"));
assert_eq!(msg.name(), "dbg");
let raw = msg.into_raw();
assert_eq!(raw.name(), "dbg");
assert_eq!(raw.payload(), b"payload");
}
#[tokio::test]
async fn shutdown_reports_dropped_registrations() {
let broker = MemoryBroker::new();
let connected = broker
.connect()
.await
.expect("memory connect is infallible");
let _first = connected.subscribe("orders").await.unwrap();
let _second = connected.subscribe("orders").await.unwrap();
let _third = connected.subscribe("billing").await.unwrap();
let closed = connected.shutdown().await.unwrap();
assert_eq!(closed.subscribers_dropped(), 3);
}
#[tokio::test]
async fn shutdown_after_a_sibling_shutdown_reports_nothing_dropped() {
let broker = MemoryBroker::new();
let first = broker.clone().connect().await.unwrap();
let second = broker.connect().await.unwrap();
let _sub = first.subscribe("orders").await.unwrap();
assert_eq!(first.shutdown().await.unwrap().subscribers_dropped(), 1);
assert_eq!(second.shutdown().await.unwrap().subscribers_dropped(), 0);
}
#[tokio::test(start_paused = true)]
async fn nack_after_redelivers_after_the_delay() {
let broker = MemoryBroker::new();
let mut sub = MemoryBroker::subscribe(&broker, "delayed");
let publisher = broker.publisher();
publisher
.publish(OutgoingMessage::new("delayed", b"later".as_slice()))
.await
.unwrap();
let mut stream = std::pin::pin!(sub.stream());
let msg = stream.next().await.unwrap().unwrap();
msg.nack_after(Duration::from_secs(5)).await.unwrap();
assert!(futures::poll!(stream.next()).is_pending());
tokio::time::advance(Duration::from_secs(5)).await;
tokio::task::yield_now().await;
let redelivered = stream.next().await.unwrap().unwrap();
assert_eq!(redelivered.payload(), b"later");
redelivered.ack().await.unwrap();
}
#[tokio::test]
async fn stream_can_be_reentered() {
let broker = MemoryBroker::new();
let mut sub = MemoryBroker::subscribe(&broker, "test");
let publisher = broker.publisher();
publisher
.publish(OutgoingMessage::new("test", b"one".as_slice()))
.await
.unwrap();
{
let mut stream = std::pin::pin!(sub.stream());
let msg = stream.next().await.unwrap().unwrap();
assert_eq!(msg.payload(), b"one");
msg.ack().await.unwrap();
}
publisher
.publish(OutgoingMessage::new("test", b"two".as_slice()))
.await
.unwrap();
let mut stream = std::pin::pin!(sub.stream());
let msg = stream.next().await.unwrap().unwrap();
assert_eq!(msg.payload(), b"two");
msg.ack().await.unwrap();
}
}