use crate::{
error::{Result, SframeError},
frame::{FrameValidation, validation::sliding_window::SlidingWindow},
header::{self, SframeHeader},
};
use std::cell::RefCell;
pub struct ReplayAttackProtection {
window: RefCell<Window>,
key_id: Option<header::KeyId>,
}
impl ReplayAttackProtection {
pub fn with_tolerance(tolerance: u64) -> Self {
assert!(tolerance > 0, "Tolerance must be greater than 0");
let size: usize = tolerance
.try_into()
.expect("Tolerance exceeds OS capabilities");
ReplayAttackProtection {
window: RefCell::new(Window::Empty(Empty {
window: SlidingWindow::new(size),
size: size as u64,
})),
key_id: None,
}
}
pub fn for_key_id(self, key_id: header::KeyId) -> Self {
ReplayAttackProtection {
key_id: Some(key_id),
..self
}
}
pub fn inspect(&self, header: &SframeHeader) -> Result<()> {
self.verify_key_id(header)?;
let counter = header.counter();
match &*self.window.borrow() {
Window::Empty(_) => Ok(()),
Window::Active(active) => active.inspect(counter),
}
}
fn verify_key_id(&self, header: &SframeHeader) -> Result<()> {
match self.key_id {
Some(expected) if header.key_id() != expected => {
Err(rejected_key_id(header.key_id(), expected))
}
_ => Ok(()),
}
}
}
impl FrameValidation for ReplayAttackProtection {
fn validate(&self, header: &SframeHeader) -> Result<()> {
self.verify_key_id(header)?;
let counter = header.counter();
let mut window = self.window.borrow_mut();
match &mut *window {
Window::Active(active) => active.commit(counter),
Window::Empty(empty) => {
let placeholder = Empty {
window: SlidingWindow::new(0),
size: 0,
};
let mut active = std::mem::replace(empty, placeholder).anchor(counter);
let result = active.commit(counter);
*window = Window::Active(active);
result
}
}
}
}
enum Window {
Empty(Empty),
Active(Active),
}
struct Empty {
window: SlidingWindow,
size: u64,
}
impl Empty {
fn anchor(self, counter: header::Counter) -> Active {
Active {
window: self.window,
size: self.size,
oldest: counter.wrapping_sub(self.size - 1),
}
}
}
struct Active {
window: SlidingWindow,
size: u64,
oldest: header::Counter,
}
impl Active {
fn commit(&mut self, counter: header::Counter) -> Result<()> {
if self.is_newer(counter) {
self.advance_to(counter);
}
match self.window_index(counter) {
None => Err(rejected(counter, REJECT_TOO_OLD)),
Some(idx) if self.window.is_set(idx) => Err(rejected(counter, REJECT_DUPLICATED)),
Some(idx) => {
self.window.set(idx);
Ok(())
}
}
}
fn inspect(&self, counter: header::Counter) -> Result<()> {
if self.is_newer(counter) {
return Ok(());
}
match self.window_index(counter) {
None => Err(rejected(counter, REJECT_TOO_OLD)),
Some(idx) if self.window.is_set(idx) => Err(rejected(counter, REJECT_DUPLICATED)),
Some(_) => Ok(()),
}
}
fn newest(&self) -> header::Counter {
self.oldest.wrapping_add(self.size - 1)
}
fn is_newer(&self, counter: header::Counter) -> bool {
let forward = counter.wrapping_sub(self.newest());
forward != 0 && forward <= header::Counter::MAX / 2
}
fn advance_to(&mut self, counter: header::Counter) {
let shift = counter.wrapping_sub(self.newest()).min(self.size);
self.window.shift_right(shift as usize);
self.oldest = counter.wrapping_sub(self.size - 1);
}
fn window_index(&self, counter: header::Counter) -> Option<usize> {
let index = counter.wrapping_sub(self.oldest);
if index < self.size {
Some(index as usize)
} else {
None
}
}
}
const REJECT_TOO_OLD: &str = "is too old";
const REJECT_DUPLICATED: &str = "is duplicated";
fn rejected(counter: header::Counter, reason: &str) -> SframeError {
SframeError::FrameValidationFailed(format!(
"Replay check failed, frame counter {counter} {reason}"
))
}
fn rejected_key_id(key_id: header::KeyId, expected: header::KeyId) -> SframeError {
SframeError::FrameValidationFailed(format!(
"Replay check failed, key id {key_id} does not match the associated {expected}"
))
}
#[cfg(test)]
mod test {
use crate::header;
use super::*;
const KID: u64 = 23456789;
const TOLERANCE: u64 = 128;
const NEWEST: u64 = 2480;
const WINDOW_OLDEST: u64 = NEWEST - (TOLERANCE - 1);
const TOO_OLD: u64 = NEWEST - TOLERANCE; const OLDER: u64 = NEWEST - 80;
const NEWER_KEEPS_OLDER: u64 = NEWEST + 20;
const NEWER_DROPS_OLDER: u64 = NEWEST + 60;
const FULL_WINDOW_JUMP: u64 = NEWEST + TOLERANCE;
fn validator() -> Fixture {
Fixture(ReplayAttackProtection::with_tolerance(TOLERANCE))
}
fn header(counter: header::Counter) -> SframeHeader {
SframeHeader::new(KID, counter)
}
struct Fixture(ReplayAttackProtection);
impl Fixture {
fn expect_accepted(&self, counter: header::Counter) -> &Self {
assert!(
self.0.validate(&header(counter)).is_ok(),
"counter {counter} should be accepted"
);
self
}
fn expect_rejected(&self, counter: header::Counter, reason: &str) -> &Self {
match self.0.validate(&header(counter)) {
Err(SframeError::FrameValidationFailed(msg)) => assert!(
msg.contains(reason),
"counter {counter}: expected reason {reason:?}, got: {msg}"
),
other => panic!("counter {counter}: expected rejection {reason:?}, got: {other:?}"),
}
self
}
fn expect_inspected(&self, counter: header::Counter) -> &Self {
assert!(
self.0.inspect(&header(counter)).is_ok(),
"counter {counter} should pass inspection"
);
self
}
fn expect_inspect_rejected(&self, counter: header::Counter, reason: &str) -> &Self {
match self.0.inspect(&header(counter)) {
Err(SframeError::FrameValidationFailed(msg)) => assert!(
msg.contains(reason),
"counter {counter}: expected reason {reason:?}, got: {msg}"
),
other => panic!("counter {counter}: expected rejection {reason:?}, got: {other:?}"),
}
self
}
}
#[test]
fn inspect_does_not_record_the_counter() {
validator()
.expect_accepted(OLDER)
.expect_inspected(NEWEST)
.expect_inspected(NEWEST)
.expect_accepted(NEWEST);
}
#[test]
fn inspect_rejects_already_recorded_counter() {
validator()
.expect_accepted(NEWEST)
.expect_inspect_rejected(NEWEST, REJECT_DUPLICATED);
}
#[test]
fn inspect_rejects_too_old_counter() {
validator()
.expect_accepted(NEWEST)
.expect_inspect_rejected(TOO_OLD, REJECT_TOO_OLD);
}
#[test]
fn inspect_accepts_future_counter_without_advancing() {
validator()
.expect_accepted(NEWEST)
.expect_inspected(NEWER_DROPS_OLDER)
.expect_accepted(OLDER);
}
#[test]
fn inspect_on_empty_window_accepts_anything() {
validator().expect_inspected(NEWEST);
}
const OTHER_KID: header::KeyId = KID + 1;
fn header_of(key_id: header::KeyId, counter: header::Counter) -> SframeHeader {
SframeHeader::new(key_id, counter)
}
#[test]
fn accepts_the_associated_key_id() {
let validator = ReplayAttackProtection::with_tolerance(TOLERANCE).for_key_id(KID);
assert!(validator.inspect(&header_of(KID, NEWEST)).is_ok());
assert!(validator.validate(&header_of(KID, NEWEST)).is_ok());
}
#[test]
fn rejects_another_key_id() {
let validator = ReplayAttackProtection::with_tolerance(TOLERANCE).for_key_id(KID);
assert!(validator.inspect(&header_of(OTHER_KID, NEWEST)).is_err());
assert!(validator.validate(&header_of(OTHER_KID, NEWEST)).is_err());
}
#[test]
fn another_key_id_does_not_record_the_counter() {
let validator = ReplayAttackProtection::with_tolerance(TOLERANCE).for_key_id(KID);
let _ = validator.validate(&header_of(OTHER_KID, NEWEST));
assert!(validator.validate(&header_of(KID, NEWEST)).is_ok());
}
#[test]
fn without_an_associated_key_id_every_sender_shares_the_window() {
let validator = ReplayAttackProtection::with_tolerance(TOLERANCE);
assert!(validator.validate(&header_of(KID, NEWEST)).is_ok());
assert!(
validator
.validate(&header_of(OTHER_KID, NEWEST + 1))
.is_ok()
);
}
#[test]
fn accept_newer_headers() {
validator().expect_accepted(OLDER).expect_accepted(NEWEST);
}
#[test]
fn accept_older_headers_in_tolerance() {
validator().expect_accepted(NEWEST).expect_accepted(OLDER);
}
#[test]
fn reject_too_old_headers() {
validator()
.expect_accepted(NEWEST)
.expect_rejected(TOO_OLD, REJECT_TOO_OLD);
}
#[test]
fn accepts_oldest_in_window_but_rejects_one_beyond() {
validator()
.expect_accepted(NEWEST)
.expect_accepted(WINDOW_OLDEST)
.expect_rejected(TOO_OLD, REJECT_TOO_OLD);
}
#[test]
fn rejects_header_with_duplicate_frame_counts() {
validator()
.expect_accepted(NEWEST)
.expect_rejected(NEWEST, REJECT_DUPLICATED);
}
#[test]
fn rejects_header_with_duplicate_frame_counts_within_tolerance() {
validator()
.expect_accepted(NEWEST)
.expect_accepted(OLDER)
.expect_rejected(OLDER, REJECT_DUPLICATED);
}
#[test]
fn rejects_header_with_duplicate_frame_counts_with_upper_wraparound() {
validator()
.expect_accepted(header::Counter::MAX)
.expect_accepted(0)
.expect_rejected(0, REJECT_DUPLICATED);
}
#[test]
fn rejects_header_with_duplicate_frame_counts_with_lower_wraparound() {
validator()
.expect_accepted(0)
.expect_accepted(header::Counter::MAX)
.expect_rejected(header::Counter::MAX, REJECT_DUPLICATED);
}
#[test]
fn detects_duplicate_after_window_advanced() {
validator()
.expect_accepted(NEWEST)
.expect_accepted(OLDER)
.expect_accepted(NEWER_KEEPS_OLDER)
.expect_rejected(OLDER, REJECT_DUPLICATED);
}
#[test]
fn dropped_counter_is_too_old_not_duplicate() {
validator()
.expect_accepted(NEWEST)
.expect_accepted(OLDER)
.expect_accepted(NEWER_DROPS_OLDER)
.expect_rejected(OLDER, REJECT_TOO_OLD);
}
#[test]
fn jump_beyond_window_clears_all_marks() {
let validator = validator();
validator.expect_accepted(NEWEST);
for counter in WINDOW_OLDEST..NEWEST {
validator.expect_accepted(counter);
}
validator.expect_accepted(FULL_WINDOW_JUMP);
for counter in (NEWEST + 1)..FULL_WINDOW_JUMP {
validator.expect_accepted(counter);
}
}
#[test]
fn handle_overflowing_counters() {
let start = header::Counter::MAX - 3;
let validator = validator();
validator.expect_accepted(start);
for step in 1..10 {
validator.expect_accepted(start.wrapping_add(step));
}
}
}