1use ironfix_core::types::SeqNum;
12use std::num::NonZeroU64;
13use std::sync::atomic::{AtomicU64, Ordering};
14
15#[derive(Debug, Clone, Copy, PartialEq, Eq, thiserror::Error)]
24#[error(
25 "sequence counter exhausted: {counter} reached u64::MAX, session requires a sequence reset"
26)]
27pub struct SequenceExhausted {
28 pub counter: SequenceCounter,
30}
31
32#[derive(Debug, Clone, Copy, PartialEq, Eq)]
34pub enum SequenceCounter {
35 Sender,
37 Target,
39}
40
41impl std::fmt::Display for SequenceCounter {
42 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
43 match self {
44 Self::Sender => write!(f, "sender"),
45 Self::Target => write!(f, "target"),
46 }
47 }
48}
49
50#[derive(Debug)]
54pub struct SequenceManager {
55 next_sender_seq: AtomicU64,
57 next_target_seq: AtomicU64,
59}
60
61impl SequenceManager {
62 #[must_use]
64 pub fn new() -> Self {
65 Self {
66 next_sender_seq: AtomicU64::new(1),
67 next_target_seq: AtomicU64::new(1),
68 }
69 }
70
71 #[must_use]
83 pub const fn with_initial(sender_seq: NonZeroU64, target_seq: NonZeroU64) -> Self {
84 Self {
85 next_sender_seq: AtomicU64::new(sender_seq.get()),
86 next_target_seq: AtomicU64::new(target_seq.get()),
87 }
88 }
89
90 #[inline]
92 #[must_use]
93 pub fn next_sender_seq(&self) -> SeqNum {
94 SeqNum::new(self.next_sender_seq.load(Ordering::SeqCst))
95 }
96
97 #[inline]
99 #[must_use]
100 pub fn next_target_seq(&self) -> SeqNum {
101 SeqNum::new(self.next_target_seq.load(Ordering::SeqCst))
102 }
103
104 #[inline]
113 #[must_use = "dropping the allocated sequence number leaves a gap in the outbound stream"]
114 #[deprecated(
115 since = "0.4.0",
116 note = "wraps silently on overflow, which corrupts a live session; use try_allocate_sender_seq. Removed in the next breaking release."
117 )]
118 pub fn allocate_sender_seq(&self) -> SeqNum {
119 SeqNum::new(self.next_sender_seq.fetch_add(1, Ordering::SeqCst))
120 }
121
122 #[inline]
133 pub fn try_allocate_sender_seq(&self) -> Result<SeqNum, SequenceExhausted> {
134 self.next_sender_seq
135 .fetch_update(Ordering::SeqCst, Ordering::SeqCst, |current| {
136 current.checked_add(1)
137 })
138 .map(SeqNum::new)
139 .map_err(|_| SequenceExhausted {
140 counter: SequenceCounter::Sender,
141 })
142 }
143
144 #[inline]
152 #[deprecated(
153 since = "0.4.0",
154 note = "wraps silently on overflow, which corrupts a live session; use try_increment_target_seq. Removed in the next breaking release."
155 )]
156 pub fn increment_target_seq(&self) {
157 self.next_target_seq.fetch_add(1, Ordering::SeqCst);
158 }
159
160 #[inline]
171 pub fn try_increment_target_seq(&self) -> Result<SeqNum, SequenceExhausted> {
172 self.next_target_seq
173 .fetch_update(Ordering::SeqCst, Ordering::SeqCst, |current| {
174 current.checked_add(1)
175 })
176 .map(|previous| SeqNum::new(previous + 1))
177 .map_err(|_| SequenceExhausted {
178 counter: SequenceCounter::Target,
179 })
180 }
181
182 #[inline]
188 pub fn set_sender_seq(&self, seq: u64) {
189 self.next_sender_seq.store(seq, Ordering::SeqCst);
190 }
191
192 #[inline]
198 pub fn set_target_seq(&self, seq: u64) {
199 self.next_target_seq.store(seq, Ordering::SeqCst);
200 }
201
202 #[inline]
204 pub fn reset(&self) {
205 self.next_sender_seq.store(1, Ordering::SeqCst);
206 self.next_target_seq.store(1, Ordering::SeqCst);
207 }
208
209 #[must_use]
224 pub fn validate_incoming(&self, received: u64) -> SequenceResult {
225 let expected = self.next_target_seq.load(Ordering::SeqCst);
226
227 if received == expected {
228 SequenceResult::Ok
229 } else if received < expected {
230 SequenceResult::TooLow { expected, received }
231 } else {
232 SequenceResult::Gap { expected, received }
233 }
234 }
235}
236
237impl Default for SequenceManager {
238 fn default() -> Self {
239 Self::new()
240 }
241}
242
243#[derive(Debug, Clone, Copy, PartialEq, Eq)]
245pub enum SequenceResult {
246 Ok,
248 TooLow {
250 expected: u64,
252 received: u64,
254 },
255 Gap {
257 expected: u64,
259 received: u64,
261 },
262}
263
264impl SequenceResult {
265 #[must_use]
267 pub const fn is_ok(&self) -> bool {
268 matches!(self, Self::Ok)
269 }
270
271 #[must_use]
273 pub const fn is_gap(&self) -> bool {
274 matches!(self, Self::Gap { .. })
275 }
276
277 #[must_use]
279 pub const fn is_too_low(&self) -> bool {
280 matches!(self, Self::TooLow { .. })
281 }
282}
283
284#[cfg(test)]
285mod tests {
286 use super::*;
287
288 #[track_caller]
291 fn nz(value: u64) -> NonZeroU64 {
292 match NonZeroU64::new(value) {
293 Some(value) => value,
294 None => panic!("test seed {value} must be non-zero"),
295 }
296 }
297
298 #[test]
299 fn test_sequence_manager_new() {
300 let mgr = SequenceManager::new();
301 assert_eq!(mgr.next_sender_seq().value(), 1);
302 assert_eq!(mgr.next_target_seq().value(), 1);
303 }
304
305 #[test]
306 fn test_with_initial_seeds_the_given_nonzero_values() {
307 let mgr = SequenceManager::with_initial(nz(7), nz(9));
310 assert_eq!(mgr.next_sender_seq().value(), 7);
311 assert_eq!(mgr.next_target_seq().value(), 9);
312 }
313
314 #[test]
315 #[allow(deprecated)]
316 fn test_allocate_sender_seq() {
317 let mgr = SequenceManager::new();
318
319 let seq1 = mgr.allocate_sender_seq();
320 assert_eq!(seq1.value(), 1);
321 assert_eq!(mgr.next_sender_seq().value(), 2);
322
323 let seq2 = mgr.allocate_sender_seq();
324 assert_eq!(seq2.value(), 2);
325 assert_eq!(mgr.next_sender_seq().value(), 3);
326 }
327
328 #[test]
329 #[allow(deprecated)]
330 fn test_increment_target_seq() {
331 let mgr = SequenceManager::new();
332
333 mgr.increment_target_seq();
334 assert_eq!(mgr.next_target_seq().value(), 2);
335
336 mgr.increment_target_seq();
337 assert_eq!(mgr.next_target_seq().value(), 3);
338 }
339
340 #[test]
341 fn test_validate_incoming() {
342 let mgr = SequenceManager::new();
343
344 assert!(mgr.validate_incoming(1).is_ok());
345
346 mgr.set_target_seq(5);
347 assert!(mgr.validate_incoming(4).is_too_low());
348 assert!(mgr.validate_incoming(5).is_ok());
349 assert!(mgr.validate_incoming(10).is_gap());
350 }
351
352 #[test]
353 fn test_try_allocate_sender_seq() {
354 let mgr = SequenceManager::new();
355
356 assert_eq!(mgr.try_allocate_sender_seq().map(SeqNum::value), Ok(1));
357 assert_eq!(mgr.try_allocate_sender_seq().map(SeqNum::value), Ok(2));
358 assert_eq!(mgr.next_sender_seq().value(), 3);
359 }
360
361 #[test]
362 fn test_try_allocate_sender_seq_exhausted() {
363 let mgr = SequenceManager::with_initial(NonZeroU64::MAX, NonZeroU64::MIN);
364
365 assert_eq!(
366 mgr.try_allocate_sender_seq(),
367 Err(SequenceExhausted {
368 counter: SequenceCounter::Sender
369 })
370 );
371 assert_eq!(mgr.next_sender_seq().value(), u64::MAX);
373 assert!(mgr.try_allocate_sender_seq().is_err());
374
375 mgr.reset();
377 assert_eq!(mgr.try_allocate_sender_seq().map(SeqNum::value), Ok(1));
378 }
379
380 #[test]
381 fn test_try_increment_target_seq() {
382 let mgr = SequenceManager::new();
383
384 assert_eq!(mgr.try_increment_target_seq().map(SeqNum::value), Ok(2));
385 assert_eq!(mgr.try_increment_target_seq().map(SeqNum::value), Ok(3));
386 assert_eq!(mgr.next_target_seq().value(), 3);
387 }
388
389 #[test]
390 fn test_try_increment_target_seq_exhausted() {
391 let mgr = SequenceManager::with_initial(NonZeroU64::MIN, NonZeroU64::MAX);
392
393 assert_eq!(
394 mgr.try_increment_target_seq(),
395 Err(SequenceExhausted {
396 counter: SequenceCounter::Target
397 })
398 );
399 assert_eq!(mgr.next_target_seq().value(), u64::MAX);
400 assert!(mgr.try_increment_target_seq().is_err());
401 }
402
403 #[test]
404 fn test_reset() {
405 let mgr = SequenceManager::with_initial(nz(100), nz(200));
406 assert_eq!(mgr.next_sender_seq().value(), 100);
407 assert_eq!(mgr.next_target_seq().value(), 200);
408
409 mgr.reset();
410 assert_eq!(mgr.next_sender_seq().value(), 1);
411 assert_eq!(mgr.next_target_seq().value(), 1);
412 }
413}