use alloc::collections::VecDeque;
use alloc::vec;
use alloc::vec::Vec;
use broadcast_common::stage::Timestamp;
use bytes::Bytes;
use core::time::Duration;
use thiserror::Error;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct SourceId(pub usize);
#[derive(Debug, Clone, PartialEq, Eq)]
#[non_exhaustive]
pub enum MergePolicy {
FirstArrival,
Failover {
primary: SourceId,
secondary: SourceId,
silence_timeout: Duration,
},
}
#[derive(Debug, Error, Clone, Copy, PartialEq, Eq)]
#[non_exhaustive]
pub enum MergeError {
#[error("source {source_id:?} is out of range for a merge with {num_sources} sources")]
UnknownSource {
source_id: SourceId,
num_sources: usize,
},
#[error("merge output queue full: {max_queued} messages already queued for poll()")]
QueueFull {
max_queued: usize,
},
}
pub struct ByteMerge {
policy: MergePolicy,
num_sources: usize,
last_seen: Vec<Option<Timestamp>>,
active: usize,
queue: VecDeque<(Bytes, Timestamp)>,
max_queued: usize,
}
impl ByteMerge {
pub fn new(policy: MergePolicy, num_sources: usize, max_queued: usize) -> Self {
assert!(num_sources > 0, "ByteMerge num_sources must be > 0");
assert!(max_queued > 0, "ByteMerge max_queued must be > 0");
let active = match &policy {
MergePolicy::Failover {
primary, secondary, ..
} => {
assert!(
primary.0 < num_sources && secondary.0 < num_sources,
"ByteMerge Failover primary/secondary must be within 0..num_sources"
);
primary.0
}
MergePolicy::FirstArrival => 0,
};
ByteMerge {
policy,
num_sources,
last_seen: vec![None; num_sources],
active,
queue: VecDeque::new(),
max_queued,
}
}
pub fn num_sources(&self) -> usize {
self.num_sources
}
pub fn len(&self) -> usize {
self.queue.len()
}
pub fn is_empty(&self) -> bool {
self.queue.is_empty()
}
pub fn feed(&mut self, source: SourceId, msg: Bytes, at: Timestamp) -> Result<(), MergeError> {
if source.0 >= self.num_sources {
return Err(MergeError::UnknownSource {
source_id: source,
num_sources: self.num_sources,
});
}
self.last_seen[source.0] = Some(at);
let forward = match &self.policy {
MergePolicy::FirstArrival => true,
MergePolicy::Failover {
primary, secondary, ..
} => {
if source.0 == primary.0 {
self.active = primary.0;
true
} else if source.0 == secondary.0 {
self.active == secondary.0
} else {
false
}
}
};
if forward {
if self.queue.len() >= self.max_queued {
return Err(MergeError::QueueFull {
max_queued: self.max_queued,
});
}
self.queue.push_back((msg, at));
}
Ok(())
}
pub fn poll(&mut self) -> Option<(Bytes, Timestamp)> {
self.queue.pop_front()
}
pub fn next_deadline(&self) -> Option<Timestamp> {
match &self.policy {
MergePolicy::FirstArrival => None,
MergePolicy::Failover {
primary,
silence_timeout,
..
} => self.last_seen[primary.0].map(|t| t.saturating_add(*silence_timeout)),
}
}
pub fn on_deadline(&mut self, now: Timestamp) {
if let MergePolicy::Failover {
primary,
secondary,
silence_timeout,
} = &self.policy
{
if let Some(last) = self.last_seen[primary.0] {
if now.saturating_sub(last) >= *silence_timeout {
self.active = secondary.0;
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn first_arrival_interleaves_two_sources_in_arrival_order() {
let mut merge = ByteMerge::new(MergePolicy::FirstArrival, 2, 8);
let a = SourceId(0);
let b = SourceId(1);
merge
.feed(a, Bytes::from_static(b"a0"), Timestamp::from_nanos(0))
.unwrap();
merge
.feed(b, Bytes::from_static(b"b0"), Timestamp::from_nanos(1))
.unwrap();
merge
.feed(a, Bytes::from_static(b"a1"), Timestamp::from_nanos(2))
.unwrap();
merge
.feed(b, Bytes::from_static(b"b1"), Timestamp::from_nanos(3))
.unwrap();
assert_eq!(
merge.poll(),
Some((Bytes::from_static(b"a0"), Timestamp::from_nanos(0)))
);
assert_eq!(
merge.poll(),
Some((Bytes::from_static(b"b0"), Timestamp::from_nanos(1)))
);
assert_eq!(
merge.poll(),
Some((Bytes::from_static(b"a1"), Timestamp::from_nanos(2)))
);
assert_eq!(
merge.poll(),
Some((Bytes::from_static(b"b1"), Timestamp::from_nanos(3)))
);
assert_eq!(merge.poll(), None);
}
#[test]
fn failover_switches_after_timeout_not_on_single_late_message_then_switches_back() {
let primary = SourceId(0);
let secondary = SourceId(1);
let silence_timeout = Duration::from_millis(100);
let mut merge = ByteMerge::new(
MergePolicy::Failover {
primary,
secondary,
silence_timeout,
},
2,
8,
);
merge
.feed(primary, Bytes::from_static(b"p0"), Timestamp::from_nanos(0))
.unwrap();
assert_eq!(
merge.poll(),
Some((Bytes::from_static(b"p0"), Timestamp::from_nanos(0)))
);
merge
.feed(
secondary,
Bytes::from_static(b"s-dropped"),
Timestamp::from_nanos(10),
)
.unwrap();
assert_eq!(merge.poll(), None);
merge.on_deadline(Timestamp::from_nanos(50_000_000));
merge
.feed(
primary,
Bytes::from_static(b"p-late"),
Timestamp::from_nanos(90_000_000),
)
.unwrap();
assert_eq!(
merge.poll(),
Some((
Bytes::from_static(b"p-late"),
Timestamp::from_nanos(90_000_000)
))
);
merge.on_deadline(Timestamp::from_nanos(150_000_000));
merge
.feed(
secondary,
Bytes::from_static(b"s-still-dropped"),
Timestamp::from_nanos(150_000_000),
)
.unwrap();
assert_eq!(
merge.poll(),
None,
"must not have flapped to secondary on a single late primary message"
);
merge.on_deadline(Timestamp::from_nanos(200_000_000));
merge
.feed(
secondary,
Bytes::from_static(b"s0"),
Timestamp::from_nanos(200_000_000),
)
.unwrap();
assert_eq!(
merge.poll(),
Some((
Bytes::from_static(b"s0"),
Timestamp::from_nanos(200_000_000)
)),
"must have switched to secondary once genuinely silent past the timeout"
);
merge
.feed(
primary,
Bytes::from_static(b"p-back"),
Timestamp::from_nanos(210_000_000),
)
.unwrap();
assert_eq!(
merge.poll(),
Some((
Bytes::from_static(b"p-back"),
Timestamp::from_nanos(210_000_000)
))
);
merge
.feed(
secondary,
Bytes::from_static(b"s-dropped-again"),
Timestamp::from_nanos(220_000_000),
)
.unwrap();
assert_eq!(merge.poll(), None);
}
#[test]
fn per_source_state_and_output_queue_are_bounded_under_flood() {
let max_queued = 4;
let mut merge = ByteMerge::new(MergePolicy::FirstArrival, 2, max_queued);
let mut full_errors = 0usize;
for i in 0..10_000u64 {
match merge.feed(
SourceId(0),
Bytes::from_static(b"x"),
Timestamp::from_nanos(i),
) {
Ok(()) => {}
Err(MergeError::QueueFull { max_queued: cap }) => {
assert_eq!(cap, max_queued);
full_errors += 1;
}
Err(other) => panic!("unexpected error: {other:?}"),
}
assert!(
merge.len() <= max_queued,
"queue exceeded its bound mid-flood"
);
}
assert_eq!(merge.num_sources(), 2);
assert_eq!(merge.len(), max_queued);
assert!(full_errors > 0, "flood should have hit the queue bound");
let mut drained = 0usize;
while merge.poll().is_some() {
drained += 1;
}
assert_eq!(drained, max_queued);
}
#[test]
fn unknown_source_is_rejected() {
let mut merge = ByteMerge::new(MergePolicy::FirstArrival, 2, 4);
let err = merge
.feed(SourceId(5), Bytes::from_static(b"x"), Timestamp::ZERO)
.unwrap_err();
assert_eq!(
err,
MergeError::UnknownSource {
source_id: SourceId(5),
num_sources: 2,
}
);
assert!(merge.is_empty());
}
#[test]
#[should_panic(expected = "num_sources must be > 0")]
fn zero_sources_panics() {
let _ = ByteMerge::new(MergePolicy::FirstArrival, 0, 4);
}
#[test]
#[should_panic(expected = "max_queued must be > 0")]
fn zero_max_queued_panics() {
let _ = ByteMerge::new(MergePolicy::FirstArrival, 2, 0);
}
#[test]
#[should_panic(expected = "within 0..num_sources")]
fn failover_out_of_range_source_panics() {
let _ = ByteMerge::new(
MergePolicy::Failover {
primary: SourceId(0),
secondary: SourceId(9),
silence_timeout: Duration::from_millis(1),
},
2,
4,
);
}
}