use ironfix_core::types::SeqNum;
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 fn with_initial(sender_seq: u64, target_seq: u64) -> Self {
Self {
next_sender_seq: AtomicU64::new(sender_seq),
next_target_seq: AtomicU64::new(target_seq),
}
}
#[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]
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]
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::*;
#[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_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]
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().unwrap().value(), 1);
assert_eq!(mgr.try_allocate_sender_seq().unwrap().value(), 2);
assert_eq!(mgr.next_sender_seq().value(), 3);
}
#[test]
fn test_try_allocate_sender_seq_exhausted() {
let mgr = SequenceManager::with_initial(u64::MAX, 1);
let err = mgr.try_allocate_sender_seq().unwrap_err();
assert_eq!(err.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().unwrap().value(), 1);
}
#[test]
fn test_try_increment_target_seq() {
let mgr = SequenceManager::new();
assert_eq!(mgr.try_increment_target_seq().unwrap().value(), 2);
assert_eq!(mgr.try_increment_target_seq().unwrap().value(), 3);
assert_eq!(mgr.next_target_seq().value(), 3);
}
#[test]
fn test_try_increment_target_seq_exhausted() {
let mgr = SequenceManager::with_initial(1, u64::MAX);
let err = mgr.try_increment_target_seq().unwrap_err();
assert_eq!(err.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(100, 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);
}
}