use std::collections::HashMap;
use std::net::IpAddr;
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
pub const WINDOW: Duration = Duration::from_secs(10 * 60);
pub const MAX_GUESSES: u32 = 10;
pub const MAX_RECORDED: u32 = 20;
const MAX_SOURCES: usize = 4096;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
enum Source {
V4(u32),
V6(u64),
Overflow,
}
impl Source {
fn of(ip: IpAddr) -> Source {
match ip.to_canonical() {
IpAddr::V4(v4) => Source::V4(u32::from(v4)),
IpAddr::V6(v6) => Source::V6((u128::from(v6) >> 64) as u64),
}
}
}
#[derive(Debug)]
struct Window {
opened: Instant,
guesses: u32,
recorded: u32,
pending: u32,
}
impl Window {
fn new(now: Instant) -> Window {
Window {
opened: now,
guesses: 0,
recorded: 0,
pending: 0,
}
}
fn expired(&self, now: Instant) -> bool {
now.saturating_duration_since(self.opened) >= WINDOW
}
fn roll(&mut self, now: Instant) {
if self.expired(now) {
let pending = self.pending;
*self = Window::new(now);
self.pending = pending;
}
}
}
#[derive(Debug, Default)]
pub struct EnrolAttempts {
by_source: Mutex<HashMap<Source, Window>>,
}
#[derive(Debug, Default)]
pub struct SignInRefusals {
by_source: Mutex<HashMap<Source, Window>>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Entry {
Refusal,
Quiet(Duration),
Nothing,
}
impl SignInRefusals {
pub fn refused(&self, ip: IpAddr, now: Instant) -> Entry {
let mut map = self.by_source.lock().unwrap_or_else(|p| p.into_inner());
let source = EnrolAttempts::slot(&mut map, Source::of(ip), now);
let w = map.entry(source).or_insert_with(|| Window::new(now));
w.roll(now);
w.recorded = w.recorded.saturating_add(1);
match w.recorded {
n if n <= MAX_RECORDED => Entry::Refusal,
n if n == MAX_RECORDED + 1 => {
Entry::Quiet(WINDOW.saturating_sub(now.saturating_duration_since(w.opened)))
}
_ => Entry::Nothing,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Settled {
pub record: bool,
pub locked: Option<Duration>,
}
#[derive(Debug)]
pub struct Ticket {
attempts: Arc<EnrolAttempts>,
source: Source,
settled: bool,
}
impl EnrolAttempts {
pub fn admit(self: &Arc<Self>, ip: IpAddr, now: Instant) -> Result<Ticket, Duration> {
let mut map = self.by_source.lock().unwrap_or_else(|p| p.into_inner());
let source = Self::slot(&mut map, Source::of(ip), now);
let w = map.entry(source).or_insert_with(|| Window::new(now));
w.roll(now);
if w.guesses + w.pending >= MAX_GUESSES {
let left = WINDOW.saturating_sub(now.saturating_duration_since(w.opened));
return Err(left.max(Duration::from_secs(1)));
}
w.pending += 1;
Ok(Ticket {
attempts: Arc::clone(self),
source,
settled: false,
})
}
fn slot(map: &mut HashMap<Source, Window>, source: Source, now: Instant) -> Source {
if map.contains_key(&source) || map.len() < MAX_SOURCES {
return source;
}
map.retain(|_, w| w.pending > 0 || !w.expired(now));
if map.len() < MAX_SOURCES {
source
} else {
Source::Overflow
}
}
fn settle(&self, source: Source, refused: Option<bool>, now: Instant) -> Settled {
let mut map = self.by_source.lock().unwrap_or_else(|p| p.into_inner());
let w = map.entry(source).or_insert_with(|| Window::new(now));
w.roll(now);
w.pending = w.pending.saturating_sub(1);
let Some(guess) = refused else {
return Settled {
record: false,
locked: None,
};
};
if guess {
w.guesses += 1;
}
let record = w.recorded < MAX_RECORDED;
if record {
w.recorded += 1;
}
let locked = (guess && w.guesses == MAX_GUESSES)
.then(|| WINDOW.saturating_sub(now.saturating_duration_since(w.opened)));
Settled { record, locked }
}
}
impl Ticket {
pub fn enrolled(mut self, now: Instant) {
self.settled = true;
self.attempts.settle(self.source, None, now);
}
pub fn refused(mut self, guess: bool, now: Instant) -> Settled {
self.settled = true;
self.attempts.settle(self.source, Some(guess), now)
}
pub fn unasked(mut self, now: Instant) {
self.settled = true;
self.attempts.settle(self.source, None, now);
}
}
impl Drop for Ticket {
fn drop(&mut self) {
if !self.settled {
self.settled = true;
self.attempts
.settle(self.source, Some(true), Instant::now());
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn ip(s: &str) -> IpAddr {
s.parse().unwrap()
}
#[test]
fn ten_unknown_codes_lock_an_address_for_the_rest_of_the_window() {
let a = Arc::new(EnrolAttempts::default());
let t0 = Instant::now();
for n in 1..=MAX_GUESSES {
let s = a.admit(ip("192.168.1.50"), t0).unwrap().refused(true, t0);
assert!(s.record, "guess {n} is on record");
assert_eq!(s.locked.is_some(), n == MAX_GUESSES, "guess {n}");
}
let later = t0 + Duration::from_secs(60);
let left = a.admit(ip("192.168.1.50"), later).unwrap_err();
assert_eq!(left, WINDOW - Duration::from_secs(60));
assert!(
a.admit(ip("192.168.1.51"), later).is_ok(),
"another address is not affected"
);
assert!(
a.admit(ip("192.168.1.50"), t0 + WINDOW).is_ok(),
"the lock ends with the window"
);
}
#[test]
fn a_refusal_for_a_real_invitation_never_locks_but_is_recorded_only_so_often() {
let a = Arc::new(EnrolAttempts::default());
let t0 = Instant::now();
let recorded = (0..100)
.filter(|_| {
a.admit(ip("10.0.0.1"), t0)
.unwrap()
.refused(false, t0)
.record
})
.count();
assert_eq!(recorded, MAX_RECORDED as usize);
assert!(a.admit(ip("10.0.0.1"), t0).is_ok());
}
#[test]
fn attempts_in_flight_count_before_they_are_answered() {
let a = Arc::new(EnrolAttempts::default());
let t0 = Instant::now();
let held: Vec<Ticket> = (0..MAX_GUESSES)
.map(|_| a.admit(ip("10.0.0.2"), t0).unwrap())
.collect();
assert!(a.admit(ip("10.0.0.2"), t0).is_err());
for t in held {
t.unasked(t0);
}
assert!(a.admit(ip("10.0.0.2"), t0).is_ok(), "freed once answered");
}
#[test]
fn an_attempt_dropped_unanswered_counts_as_a_guess() {
let a = Arc::new(EnrolAttempts::default());
let t0 = Instant::now();
for _ in 0..MAX_GUESSES {
drop(a.admit(ip("10.0.0.3"), t0).unwrap());
}
assert!(a.admit(ip("10.0.0.3"), t0).is_err());
}
#[test]
fn an_ipv6_address_is_counted_by_its_64_and_a_mapped_ipv4_as_itself() {
let a = Arc::new(EnrolAttempts::default());
let t0 = Instant::now();
for n in 0..MAX_GUESSES {
a.admit(ip(&format!("2001:db8:1:2::{n:x}")), t0)
.unwrap()
.refused(true, t0);
}
assert!(a.admit(ip("2001:db8:1:2:ffff::1"), t0).is_err());
assert!(a.admit(ip("2001:db8:1:3::1"), t0).is_ok());
for _ in 0..MAX_GUESSES {
a.admit(ip("192.0.2.7"), t0).unwrap().refused(true, t0);
}
assert!(a.admit(ip("::ffff:192.0.2.7"), t0).is_err());
}
#[test]
fn refused_sign_ins_are_recorded_up_to_the_cap_then_said_once_and_never_locked() {
let s = SignInRefusals::default();
let t0 = Instant::now();
let entries: Vec<Entry> = (0..MAX_RECORDED + 5)
.map(|_| s.refused(ip("192.168.1.60"), t0))
.collect();
assert!(
entries[..MAX_RECORDED as usize]
.iter()
.all(|e| *e == Entry::Refusal)
);
assert_eq!(entries[MAX_RECORDED as usize], Entry::Quiet(WINDOW));
assert!(
entries[MAX_RECORDED as usize + 1..]
.iter()
.all(|e| *e == Entry::Nothing)
);
assert_eq!(
s.refused(ip("192.168.1.61"), t0),
Entry::Refusal,
"another address has its own count"
);
assert_eq!(
s.refused(ip("192.168.1.60"), t0 + WINDOW),
Entry::Refusal,
"a new window records again"
);
}
#[test]
fn a_full_table_shares_one_count_rather_than_growing() {
let a = Arc::new(EnrolAttempts::default());
let t0 = Instant::now();
for n in 0..MAX_SOURCES as u32 {
a.admit(IpAddr::V4(n.into()), t0)
.unwrap()
.refused(false, t0);
}
for n in 0..MAX_GUESSES {
a.admit(IpAddr::V4((u32::MAX - n).into()), t0)
.unwrap()
.refused(true, t0);
}
assert!(a.admit(ip("203.0.113.200"), t0).is_err());
assert!(a.by_source.lock().unwrap().len() <= MAX_SOURCES + 1);
assert!(
a.admit(ip("203.0.113.200"), t0 + WINDOW).is_ok(),
"windows that are over make room again"
);
}
}