use std::{collections::HashMap, net::IpAddr, time::Instant};
#[cfg(feature = "locktick")]
use locktick::parking_lot::RwLock;
#[cfg(not(feature = "locktick"))]
use parking_lot::RwLock;
#[derive(Clone)]
pub struct BanDetails {
banned_at: Instant,
}
impl BanDetails {
pub fn new() -> Self {
Self { banned_at: Instant::now() }
}
}
impl Default for BanDetails {
fn default() -> Self {
Self::new()
}
}
#[derive(Default)]
pub struct BannedPeers(RwLock<HashMap<IpAddr, BanDetails>>);
impl BannedPeers {
pub fn is_ip_banned(&self, ip: &IpAddr) -> bool {
self.0.read().contains_key(ip)
}
pub fn get_banned_ips(&self) -> Vec<IpAddr> {
self.0.read().keys().cloned().collect()
}
pub fn get_ban_config(&self, ip: IpAddr) -> Option<BanDetails> {
self.0.read().get(&ip).cloned()
}
pub fn update_ip_ban(&self, ip: IpAddr) {
self.0.write().insert(ip, BanDetails::default());
}
pub fn remove_old_bans(&self, ban_time_in_secs: u64) {
self.0.write().retain(|_, ban_config| ban_config.banned_at.elapsed().as_secs() < ban_time_in_secs);
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::{
net::{Ipv4Addr, Ipv6Addr},
str::FromStr,
thread::sleep,
time::Duration,
};
const NEVER_EXPIRES: u64 = u64::MAX;
fn ipv4(last_octet: u8) -> IpAddr {
IpAddr::V4(Ipv4Addr::new(1, 2, 3, last_octet))
}
#[test]
fn an_ip_starts_out_unbanned() {
let banned_peers = BannedPeers::default();
assert!(!banned_peers.is_ip_banned(&ipv4(1)));
assert!(banned_peers.get_ban_config(ipv4(1)).is_none());
assert!(banned_peers.get_banned_ips().is_empty());
}
#[test]
fn banning_an_ip_is_visible_through_every_accessor() {
let banned_peers = BannedPeers::default();
banned_peers.update_ip_ban(ipv4(1));
assert!(banned_peers.is_ip_banned(&ipv4(1)));
assert!(banned_peers.get_ban_config(ipv4(1)).is_some());
assert_eq!(banned_peers.get_banned_ips(), vec![ipv4(1)]);
}
#[test]
fn bans_apply_only_to_the_banned_ip() {
let banned_peers = BannedPeers::default();
banned_peers.update_ip_ban(ipv4(1));
assert!(!banned_peers.is_ip_banned(&ipv4(2)));
assert!(!banned_peers.is_ip_banned(&IpAddr::V6(Ipv6Addr::LOCALHOST)));
}
#[test]
fn bans_are_keyed_by_the_literal_ip_not_its_canonical_form() {
let banned_peers = BannedPeers::default();
let native = ipv4(4);
let mapped = IpAddr::from_str("::ffff:1.2.3.4").unwrap();
banned_peers.update_ip_ban(native);
assert!(banned_peers.is_ip_banned(&native));
assert!(!banned_peers.is_ip_banned(&mapped));
assert_eq!(banned_peers.get_banned_ips(), vec![native]);
}
#[test]
fn re_banning_an_ip_refreshes_the_ban_rather_than_duplicating_it() {
let banned_peers = BannedPeers::default();
banned_peers.update_ip_ban(ipv4(1));
let first_ban = banned_peers.get_ban_config(ipv4(1)).unwrap().banned_at;
sleep(Duration::from_millis(10));
banned_peers.update_ip_ban(ipv4(1));
let second_ban = banned_peers.get_ban_config(ipv4(1)).unwrap().banned_at;
assert!(second_ban > first_ban);
assert_eq!(banned_peers.get_banned_ips().len(), 1);
}
#[test]
fn remove_old_bans_retains_unexpired_bans() {
let banned_peers = BannedPeers::default();
banned_peers.update_ip_ban(ipv4(1));
banned_peers.update_ip_ban(ipv4(2));
banned_peers.remove_old_bans(NEVER_EXPIRES);
assert!(banned_peers.is_ip_banned(&ipv4(1)));
assert!(banned_peers.is_ip_banned(&ipv4(2)));
}
#[test]
fn remove_old_bans_evicts_expired_bans() {
let banned_peers = BannedPeers::default();
banned_peers.update_ip_ban(ipv4(1));
banned_peers.update_ip_ban(ipv4(2));
banned_peers.remove_old_bans(0);
assert!(!banned_peers.is_ip_banned(&ipv4(1)));
assert!(!banned_peers.is_ip_banned(&ipv4(2)));
assert!(banned_peers.get_banned_ips().is_empty());
}
#[test]
fn remove_old_bans_evicts_only_the_expired_entries() {
let banned_peers = BannedPeers::default();
banned_peers
.0
.write()
.insert(ipv4(1), BanDetails { banned_at: Instant::now().checked_sub(Duration::from_secs(60)).unwrap() });
banned_peers.update_ip_ban(ipv4(2));
banned_peers.remove_old_bans(1);
assert!(!banned_peers.is_ip_banned(&ipv4(1)));
assert!(banned_peers.is_ip_banned(&ipv4(2)));
}
#[test]
fn remove_old_bans_is_a_no_op_on_an_empty_list() {
let banned_peers = BannedPeers::default();
banned_peers.remove_old_bans(1);
assert!(banned_peers.get_banned_ips().is_empty());
}
}