use ironfix_core::types::SeqNum;
use std::num::NonZeroU64;
use std::sync::atomic::{AtomicU64, Ordering};
#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
#[error(
"sequence counter exhausted: {counter} reached u64::MAX, session requires a sequence reset"
)]
pub struct SequenceExhausted {
pub counter: SequenceCounter,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SequenceCounter {
Sender,
Target,
}
impl std::fmt::Display for SequenceCounter {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Sender => write!(f, "sender"),
Self::Target => write!(f, "target"),
}
}
}
#[derive(Debug)]
pub struct SequenceManager {
next_sender_seq: AtomicU64,
next_target_seq: AtomicU64,
}
impl SequenceManager {
#[must_use]
pub fn new() -> Self {
Self {
next_sender_seq: AtomicU64::new(1),
next_target_seq: AtomicU64::new(1),
}
}
#[must_use]
pub const fn with_initial(sender_seq: NonZeroU64, target_seq: NonZeroU64) -> Self {
Self {
next_sender_seq: AtomicU64::new(sender_seq.get()),
next_target_seq: AtomicU64::new(target_seq.get()),
}
}
#[inline]
#[must_use]
pub fn next_sender_seq(&self) -> SeqNum {
SeqNum::new(self.next_sender_seq.load(Ordering::SeqCst))
}
#[inline]
#[must_use]
pub fn next_target_seq(&self) -> SeqNum {
SeqNum::new(self.next_target_seq.load(Ordering::SeqCst))
}
#[inline]
#[must_use = "dropping the allocated sequence number leaves a gap in the outbound stream"]
#[deprecated(
since = "0.4.0",
note = "wraps silently on overflow, which corrupts a live session; use try_allocate_sender_seq. Removed in the next breaking release."
)]
pub fn allocate_sender_seq(&self) -> SeqNum {
SeqNum::new(self.next_sender_seq.fetch_add(1, Ordering::SeqCst))
}
#[inline]
pub fn try_allocate_sender_seq(&self) -> Result<SeqNum, SequenceExhausted> {
self.next_sender_seq
.fetch_update(Ordering::SeqCst, Ordering::SeqCst, |current| {
current.checked_add(1)
})
.map(SeqNum::new)
.map_err(|_| SequenceExhausted {
counter: SequenceCounter::Sender,
})
}
#[inline]
#[deprecated(
since = "0.4.0",
note = "wraps silently on overflow, which corrupts a live session; use try_increment_target_seq. Removed in the next breaking release."
)]
pub fn increment_target_seq(&self) {
self.next_target_seq.fetch_add(1, Ordering::SeqCst);
}
#[inline]
pub fn try_increment_target_seq(&self) -> Result<SeqNum, SequenceExhausted> {
self.next_target_seq
.fetch_update(Ordering::SeqCst, Ordering::SeqCst, |current| {
current.checked_add(1)
})
.map(|previous| SeqNum::new(previous + 1))
.map_err(|_| SequenceExhausted {
counter: SequenceCounter::Target,
})
}
#[inline]
pub fn set_sender_seq(&self, seq: u64) {
self.next_sender_seq.store(seq, Ordering::SeqCst);
}
#[inline]
pub fn set_target_seq(&self, seq: u64) {
self.next_target_seq.store(seq, Ordering::SeqCst);
}
#[inline]
pub fn reset(&self) {
self.next_sender_seq.store(1, Ordering::SeqCst);
self.next_target_seq.store(1, Ordering::SeqCst);
}
#[must_use]
pub fn validate_incoming(&self, received: u64) -> SequenceResult {
let expected = self.next_target_seq.load(Ordering::SeqCst);
if received == expected {
SequenceResult::Ok
} else if received < expected {
SequenceResult::TooLow { expected, received }
} else {
SequenceResult::Gap { expected, received }
}
}
}
impl Default for SequenceManager {
fn default() -> Self {
Self::new()
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum SequenceResult {
Ok,
TooLow {
expected: u64,
received: u64,
},
Gap {
expected: u64,
received: u64,
},
}
impl SequenceResult {
#[must_use]
pub const fn is_ok(&self) -> bool {
matches!(self, Self::Ok)
}
#[must_use]
pub const fn is_gap(&self) -> bool {
matches!(self, Self::Gap { .. })
}
#[must_use]
pub const fn is_too_low(&self) -> bool {
matches!(self, Self::TooLow { .. })
}
}
#[cfg(test)]
mod tests {
use super::*;
#[track_caller]
fn nz(value: u64) -> NonZeroU64 {
match NonZeroU64::new(value) {
Some(value) => value,
None => panic!("test seed {value} must be non-zero"),
}
}
#[test]
fn test_sequence_manager_new() {
let mgr = SequenceManager::new();
assert_eq!(mgr.next_sender_seq().value(), 1);
assert_eq!(mgr.next_target_seq().value(), 1);
}
#[test]
fn test_with_initial_seeds_the_given_nonzero_values() {
let mgr = SequenceManager::with_initial(nz(7), nz(9));
assert_eq!(mgr.next_sender_seq().value(), 7);
assert_eq!(mgr.next_target_seq().value(), 9);
}
#[test]
#[allow(deprecated)]
fn test_allocate_sender_seq() {
let mgr = SequenceManager::new();
let seq1 = mgr.allocate_sender_seq();
assert_eq!(seq1.value(), 1);
assert_eq!(mgr.next_sender_seq().value(), 2);
let seq2 = mgr.allocate_sender_seq();
assert_eq!(seq2.value(), 2);
assert_eq!(mgr.next_sender_seq().value(), 3);
}
#[test]
#[allow(deprecated)]
fn test_increment_target_seq() {
let mgr = SequenceManager::new();
mgr.increment_target_seq();
assert_eq!(mgr.next_target_seq().value(), 2);
mgr.increment_target_seq();
assert_eq!(mgr.next_target_seq().value(), 3);
}
#[test]
fn test_validate_incoming() {
let mgr = SequenceManager::new();
assert!(mgr.validate_incoming(1).is_ok());
mgr.set_target_seq(5);
assert!(mgr.validate_incoming(4).is_too_low());
assert!(mgr.validate_incoming(5).is_ok());
assert!(mgr.validate_incoming(10).is_gap());
}
#[test]
fn test_try_allocate_sender_seq() {
let mgr = SequenceManager::new();
assert_eq!(mgr.try_allocate_sender_seq().map(SeqNum::value), Ok(1));
assert_eq!(mgr.try_allocate_sender_seq().map(SeqNum::value), Ok(2));
assert_eq!(mgr.next_sender_seq().value(), 3);
}
#[test]
fn test_try_allocate_sender_seq_exhausted() {
let mgr = SequenceManager::with_initial(NonZeroU64::MAX, NonZeroU64::MIN);
assert_eq!(
mgr.try_allocate_sender_seq(),
Err(SequenceExhausted {
counter: SequenceCounter::Sender
})
);
assert_eq!(mgr.next_sender_seq().value(), u64::MAX);
assert!(mgr.try_allocate_sender_seq().is_err());
mgr.reset();
assert_eq!(mgr.try_allocate_sender_seq().map(SeqNum::value), Ok(1));
}
#[test]
fn test_try_increment_target_seq() {
let mgr = SequenceManager::new();
assert_eq!(mgr.try_increment_target_seq().map(SeqNum::value), Ok(2));
assert_eq!(mgr.try_increment_target_seq().map(SeqNum::value), Ok(3));
assert_eq!(mgr.next_target_seq().value(), 3);
}
#[test]
fn test_try_increment_target_seq_exhausted() {
let mgr = SequenceManager::with_initial(NonZeroU64::MIN, NonZeroU64::MAX);
assert_eq!(
mgr.try_increment_target_seq(),
Err(SequenceExhausted {
counter: SequenceCounter::Target
})
);
assert_eq!(mgr.next_target_seq().value(), u64::MAX);
assert!(mgr.try_increment_target_seq().is_err());
}
#[test]
fn test_reset() {
let mgr = SequenceManager::with_initial(nz(100), nz(200));
assert_eq!(mgr.next_sender_seq().value(), 100);
assert_eq!(mgr.next_target_seq().value(), 200);
mgr.reset();
assert_eq!(mgr.next_sender_seq().value(), 1);
assert_eq!(mgr.next_target_seq().value(), 1);
}
}