use futures::future::join_all;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::sync::Arc;
use tracing::{error, info, warn};
use crate::{
AckCallback, EncodedBatch, EncodedRecord, OffsetId, ZerobusError, ZerobusResult, ZerobusStream,
};
const STREAM_BITS: u32 = 6;
const OFFSET_MASK: i64 = (1i64 << (64 - STREAM_BITS)) - 1;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct MessageId(i64);
impl std::fmt::Display for MessageId {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"MessageId(stream={}, offset={})",
self.stream_index(),
self.sub_offset()
)
}
}
impl MessageId {
fn new(stream_index: usize, sub_offset: OffsetId) -> Self {
debug_assert!(stream_index < (1 << STREAM_BITS));
debug_assert!((0..=OFFSET_MASK).contains(&sub_offset));
Self(((stream_index as i64) << (64 - STREAM_BITS)) | (sub_offset & OFFSET_MASK))
}
pub fn stream_index(&self) -> usize {
((self.0 as u64) >> (64 - STREAM_BITS)) as usize
}
pub fn sub_offset(&self) -> OffsetId {
self.0 & OFFSET_MASK
}
pub fn raw(&self) -> i64 {
self.0
}
pub fn from_raw(raw: i64) -> Self {
Self(raw)
}
}
struct MultiplexedAckCallbackAdapter {
stream_index: usize,
callback: Arc<dyn AckCallback<MessageId>>,
}
impl AckCallback for MultiplexedAckCallbackAdapter {
fn on_ack(&self, offset_id: OffsetId) {
self.callback
.on_ack(MessageId::new(self.stream_index, offset_id));
}
fn on_error(&self, offset_id: OffsetId, error_message: &str) {
self.callback
.on_error(MessageId::new(self.stream_index, offset_id), error_message);
}
}
#[allow(dead_code)]
pub(crate) fn multiplexed_ack_callback(
stream_index: usize,
callback: Arc<dyn AckCallback<MessageId>>,
) -> Arc<dyn AckCallback> {
assert!(
stream_index < (1 << STREAM_BITS),
"MultiplexedStream supports at most {} sub-streams",
1 << STREAM_BITS
);
Arc::new(MultiplexedAckCallbackAdapter {
stream_index,
callback,
})
}
pub struct MultiplexedStream {
streams: Vec<ZerobusStream>,
round_robin_counter: AtomicUsize,
is_closed: AtomicBool,
}
impl MultiplexedStream {
pub fn new(streams: Vec<ZerobusStream>) -> Self {
assert!(
!streams.is_empty(),
"MultiplexedStream requires at least one sub-stream"
);
assert!(
streams.len() <= (1 << STREAM_BITS),
"MultiplexedStream supports at most {} sub-streams",
1 << STREAM_BITS
);
Self {
streams,
round_robin_counter: AtomicUsize::new(0),
is_closed: AtomicBool::new(false),
}
}
#[allow(clippy::result_large_err)]
fn check_closed(&self) -> ZerobusResult<()> {
if self.is_closed_fast() {
return Err(ZerobusError::InvalidStateError(
"MultiplexedStream is closed".to_string(),
));
}
Ok(())
}
fn is_closed_fast(&self) -> bool {
self.is_closed.load(Ordering::Relaxed)
}
async fn shutdown_on_failure(&self, trigger_index: usize, cause: &ZerobusError) {
if self.is_closed.swap(true, Ordering::Relaxed) {
return;
}
error!(
trigger_stream_index = trigger_index,
cause = %cause,
num_streams = self.streams.len(),
"MultiplexedStream poisoned due to sub-stream failure"
);
let flush_results = join_all(self.streams.iter().map(|s| s.flush())).await;
for (i, result) in flush_results.into_iter().enumerate() {
if let Err(e) = result {
warn!(stream_index = i, error = %e, "Failed to flush sub-stream during shutdown");
}
}
for s in &self.streams {
s.signal_shutdown();
}
}
fn pick_substream(&self) -> usize {
self.round_robin_counter.fetch_add(1, Ordering::Relaxed) % self.streams.len()
}
async fn wait_for_capacity(&self, stream: &ZerobusStream, idx: usize) -> ZerobusResult<()> {
let mut backoff_ms = 1u64;
let mut total_wait_ms = 0u64;
let mut logged_backpressure = false;
loop {
self.check_closed()?;
if stream.is_closed() {
let err = ZerobusError::InvalidStateError(format!(
"Sub-stream {} closed unexpectedly",
idx
));
self.shutdown_on_failure(idx, &err).await;
return Err(err);
}
if stream.has_capacity() {
return Ok(());
}
tokio::time::sleep(std::time::Duration::from_millis(backoff_ms)).await;
total_wait_ms += backoff_ms;
backoff_ms = (backoff_ms * 2).min(50);
if !logged_backpressure && total_wait_ms >= 1000 {
warn!(
stream_index = idx,
total_wait_ms, "Backpressure: sub-stream at capacity, waiting for drain"
);
logged_backpressure = true;
}
}
}
async fn handle_ingest_error(
&self,
e: ZerobusError,
stream: &ZerobusStream,
idx: usize,
) -> ZerobusError {
if stream.is_closed() {
self.shutdown_on_failure(idx, &e).await;
} else {
warn!(stream_index = idx, error = %e, "Ingest errored but sub-stream still alive");
}
e
}
pub async fn ingest_record(
&self,
payload: impl Into<EncodedRecord>,
) -> ZerobusResult<MessageId> {
self.check_closed()?;
let record = payload.into();
let idx = self.pick_substream();
let stream = &self.streams[idx];
self.wait_for_capacity(stream, idx).await?;
self.check_closed()?;
match stream.ingest_record_offset(record).await {
Ok(off) => Ok(MessageId::new(idx, off)),
Err(e) => Err(self.handle_ingest_error(e, stream, idx).await),
}
}
pub async fn ingest_records<I, T>(&self, payload: I) -> ZerobusResult<Option<MessageId>>
where
I: IntoIterator<Item = T>,
T: Into<EncodedRecord>,
{
self.check_closed()?;
let records: Vec<EncodedRecord> = payload.into_iter().map(Into::into).collect();
if records.is_empty() {
return Ok(None);
}
let idx = self.pick_substream();
let stream = &self.streams[idx];
self.wait_for_capacity(stream, idx).await?;
self.check_closed()?;
match stream.ingest_records_offset(records).await {
Ok(sub_offset) => Ok(sub_offset.map(|off| MessageId::new(idx, off))),
Err(e) => Err(self.handle_ingest_error(e, stream, idx).await),
}
}
pub async fn flush(&self) -> ZerobusResult<()> {
self.check_closed()?;
let results = join_all(self.streams.iter().map(|s| s.flush())).await;
let mut first_error: Option<ZerobusError> = None;
let mut first_closed: Option<(usize, ZerobusError)> = None;
for (i, result) in results.into_iter().enumerate() {
if let Err(e) = result {
if self.streams[i].is_closed() && first_closed.is_none() {
first_closed = Some((i, e.clone()));
}
if first_error.is_none() {
first_error = Some(e);
} else {
warn!(
stream_index = i,
error = %e,
"Additional sub-stream flush error (first error will be returned)"
);
}
}
}
match first_error {
Some(e) => {
if let Some((closed_idx, closed_err)) = first_closed {
self.shutdown_on_failure(closed_idx, &closed_err).await;
} else {
warn!(error = %e, "flush errored but sub-streams still alive");
}
Err(e)
}
None => Ok(()),
}
}
pub async fn wait_for_message_id(&self, message_id: MessageId) -> ZerobusResult<()> {
let idx = message_id.stream_index();
if idx >= self.streams.len() {
return Err(ZerobusError::InvalidArgument(format!(
"Invalid stream index {} in message id",
idx
)));
}
match self.streams[idx]
.wait_for_offset(message_id.sub_offset())
.await
{
Ok(()) => Ok(()),
Err(e) => {
if self.streams[idx].is_closed() {
self.shutdown_on_failure(idx, &e).await;
} else {
warn!(
stream_index = idx,
error = %e,
"wait_for_offset errored but sub-stream still alive"
);
}
Err(e)
}
}
}
pub async fn close(&mut self) -> ZerobusResult<()> {
info!("Closing MultiplexedStream");
self.is_closed.store(true, Ordering::Relaxed);
let mut first_error: Option<ZerobusError> = None;
let flush_results = join_all(self.streams.iter().map(|s| s.flush())).await;
for (i, result) in flush_results.into_iter().enumerate() {
if let Err(e) = result {
if first_error.is_none() {
first_error = Some(e);
} else {
warn!(
stream_index = i,
error = %e,
"Additional sub-stream flush error during close"
);
}
}
}
for (i, stream) in self.streams.iter_mut().enumerate() {
if let Err(e) = stream.close().await {
if first_error.is_none() {
first_error = Some(e);
} else {
warn!(
stream_index = i,
error = %e,
"Additional sub-stream close error"
);
}
}
}
match first_error {
Some(e) => Err(e),
None => Ok(()),
}
}
pub fn is_closed(&self) -> bool {
self.is_closed_fast() || self.streams.iter().any(ZerobusStream::is_closed)
}
pub async fn get_unacked_records(
&mut self,
) -> ZerobusResult<impl Iterator<Item = EncodedRecord>> {
let _ = self.close().await;
let mut all_records = Vec::new();
for stream in &self.streams {
all_records.extend(stream.get_unacked_records().await?);
}
Ok(all_records.into_iter())
}
pub async fn get_unacked_batches(&mut self) -> ZerobusResult<Vec<EncodedBatch>> {
let _ = self.close().await;
let mut all_batches = Vec::new();
for stream in &self.streams {
all_batches.extend(stream.get_unacked_batches().await?);
}
Ok(all_batches)
}
}
impl Drop for MultiplexedStream {
fn drop(&mut self) {
self.is_closed.store(true, Ordering::Relaxed);
for stream in &self.streams {
stream.signal_shutdown();
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Mutex;
#[derive(Default)]
struct RecordingMultiplexedCallback {
acks: Mutex<Vec<MessageId>>,
errors: Mutex<Vec<(MessageId, String)>>,
}
impl AckCallback<MessageId> for RecordingMultiplexedCallback {
fn on_ack(&self, message_id: MessageId) {
self.acks.lock().unwrap().push(message_id);
}
fn on_error(&self, message_id: MessageId, error_message: &str) {
self.errors
.lock()
.unwrap()
.push((message_id, error_message.to_string()));
}
}
#[test]
#[should_panic(expected = "MultiplexedStream requires at least one sub-stream")]
fn test_constructor_panics_on_empty_streams() {
MultiplexedStream::new(vec![]);
}
#[test]
fn test_message_id_roundtrip() {
for stream_idx in 0..64 {
for sub_offset in [0i64, 1, 100, 1_000_000, i64::MAX >> STREAM_BITS] {
let id = MessageId::new(stream_idx, sub_offset);
assert_eq!(id.stream_index(), stream_idx);
assert_eq!(id.sub_offset(), sub_offset);
}
}
}
#[test]
fn test_message_id_zero() {
let id = MessageId::new(0, 0);
assert_eq!(id.raw(), 0);
assert_eq!(id.stream_index(), 0);
assert_eq!(id.sub_offset(), 0);
}
#[test]
fn test_message_id_different_streams_same_offset() {
let a = MessageId::new(0, 42);
let b = MessageId::new(1, 42);
assert_ne!(a, b);
assert_eq!(a.sub_offset(), b.sub_offset());
assert_ne!(a.stream_index(), b.stream_index());
}
#[test]
fn test_multiplexed_ack_callback_routes_substream_ids() {
let callback = Arc::new(RecordingMultiplexedCallback::default());
let stream_0 = multiplexed_ack_callback(0, callback.clone());
let stream_1 = multiplexed_ack_callback(1, callback.clone());
stream_0.on_ack(42);
stream_1.on_ack(42);
stream_1.on_error(43, "test error");
assert_eq!(
callback.acks.lock().unwrap().as_slice(),
&[MessageId::new(0, 42), MessageId::new(1, 42)]
);
assert_eq!(
callback.errors.lock().unwrap().as_slice(),
&[(MessageId::new(1, 43), "test error".to_string())]
);
}
#[test]
#[should_panic(expected = "MultiplexedStream supports at most 64 sub-streams")]
fn test_multiplexed_ack_callback_rejects_invalid_stream_index() {
let callback = Arc::new(RecordingMultiplexedCallback::default());
multiplexed_ack_callback(64, callback);
}
}