use std::collections::HashMap;
use std::net::IpAddr;
use std::sync::Mutex;
use std::time::Duration;
use axum::extract::FromRequestParts;
use axum::http::request::Parts;
use axum::http::{HeaderMap, header};
use base64::Engine as _;
use base64::engine::general_purpose::URL_SAFE_NO_PAD as BASE64_URL_SAFE_NO_PAD;
use ring::rand::{SecureRandom, SystemRandom};
use subtle::ConstantTimeEq;
use tracing::{info, warn};
use crate::webadmin::AdminState;
use crate::webadmin::error::AdminError;
use acme_proxy_store::admin_session::AdminSession;
use acme_proxy_store::admin_user::AdminRole;
use acme_proxy_store::admin_user::AdminUser;
use acme_proxy_store::nonce::fingerprint;
use acme_proxy_store::nonce::now_secs;
pub const COOKIE_NAME: &str = "__Host-acme_admin_session";
pub const CSRF_HEADER: &str = "x-csrf-token";
const TOKEN_LEN: usize = 32;
const SESSION_TOUCH_INTERVAL: i64 = 60;
pub const PENDING_MFA_TTL: Duration = Duration::from_secs(300);
pub struct MintedToken {
pub token: String,
pub token_hash: String,
}
#[must_use]
pub fn mint_token() -> MintedToken {
let mut bytes = [0u8; TOKEN_LEN];
SystemRandom::new()
.fill(&mut bytes)
.expect("system RNG unavailable");
let token = BASE64_URL_SAFE_NO_PAD.encode(bytes);
let token_hash = hash_token(&token);
MintedToken { token, token_hash }
}
#[must_use]
pub fn mint_csrf_token() -> String {
let mut bytes = [0u8; TOKEN_LEN];
SystemRandom::new()
.fill(&mut bytes)
.expect("system RNG unavailable");
BASE64_URL_SAFE_NO_PAD.encode(bytes)
}
#[must_use]
pub fn hash_token(token: &str) -> String {
let digest = ring::digest::digest(&ring::digest::SHA256, token.as_bytes());
hex::encode(digest.as_ref())
}
#[must_use]
pub fn session_cookie(token: &str, ttl: Duration) -> String {
format!(
"{COOKIE_NAME}={token}; HttpOnly; Secure; SameSite=Strict; Path=/; Max-Age={}",
ttl.as_secs()
)
}
#[must_use]
pub fn clearing_cookie() -> String {
format!("{COOKIE_NAME}=; HttpOnly; Secure; SameSite=Strict; Path=/; Max-Age=0")
}
#[must_use]
pub fn cookie_value(headers: &HeaderMap) -> Option<String> {
for header in headers.get_all(header::COOKIE) {
let Ok(raw) = header.to_str() else { continue };
for pair in raw.split(';') {
let Some((name, value)) = pair.split_once('=') else {
continue;
};
if name.trim() == COOKIE_NAME {
let value = value.trim();
let value = value
.strip_prefix('"')
.and_then(|v| v.strip_suffix('"'))
.unwrap_or(value);
return Some(value.to_string());
}
}
}
None
}
#[derive(Debug, Clone, Copy)]
pub struct AdminClientIp(pub Option<IpAddr>);
impl<S: Sync> FromRequestParts<S> for AdminClientIp {
type Rejection = std::convert::Infallible;
async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
if let Some(acme_proxy_core::client::ClientIp(resolved)) = parts
.extensions
.get::<acme_proxy_core::client::ClientIp>()
.copied()
{
return Ok(AdminClientIp(resolved));
}
let address = parts
.extensions
.get::<axum::extract::ConnectInfo<std::net::SocketAddr>>()
.map(|info| info.0.ip().to_canonical());
Ok(AdminClientIp(address))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum MfaStep {
Verify,
Enrol,
}
impl MfaStep {
#[must_use]
pub fn as_str(self) -> &'static str {
match self {
MfaStep::Verify => "verify",
MfaStep::Enrol => "enrol",
}
}
}
#[derive(Debug)]
pub struct Authenticated {
pub session: AdminSession,
pub user: AdminUser,
}
#[derive(Debug)]
pub struct AuthenticatedWrite(pub Authenticated);
#[derive(Debug)]
pub struct AdminWrite(pub Authenticated);
#[derive(Debug)]
pub struct SelfServiceWrite(pub Authenticated);
#[derive(Debug)]
pub struct AdminRead(pub Authenticated);
impl FromRequestParts<AdminState> for Authenticated {
type Rejection = AdminError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
resolve_session(parts, state).await
}
}
async fn resolve_write(parts: &Parts, state: &AdminState) -> Result<Authenticated, AdminError> {
check_origin(&parts.headers, &state.config.admin.base_url)?;
let authenticated = resolve_session(parts, state).await?;
check_csrf(&parts.headers, &authenticated.session.csrf_token)?;
Ok(authenticated)
}
fn require_role(user: &AdminUser, minimum: AdminRole) -> Result<(), AdminError> {
if user.role() >= minimum {
Ok(())
} else {
Err(AdminError::insufficient_role())
}
}
impl FromRequestParts<AdminState> for AdminRead {
type Rejection = AdminError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
let authenticated = resolve_session(parts, state).await?;
require_role(&authenticated.user, AdminRole::Admin)?;
Ok(AdminRead(authenticated))
}
}
impl FromRequestParts<AdminState> for AuthenticatedWrite {
type Rejection = AdminError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
let authenticated = resolve_write(parts, state).await?;
require_role(&authenticated.user, AdminRole::Operator)?;
Ok(AuthenticatedWrite(authenticated))
}
}
impl FromRequestParts<AdminState> for AdminWrite {
type Rejection = AdminError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
let authenticated = resolve_write(parts, state).await?;
require_role(&authenticated.user, AdminRole::Admin)?;
Ok(AdminWrite(authenticated))
}
}
impl FromRequestParts<AdminState> for SelfServiceWrite {
type Rejection = AdminError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
Ok(SelfServiceWrite(resolve_write(parts, state).await?))
}
}
#[derive(Debug)]
pub struct PendingMfa {
pub session: AdminSession,
pub user: AdminUser,
pub step: MfaStep,
}
#[derive(Debug)]
pub struct PendingMfaSubmit(pub PendingMfa);
#[derive(Debug)]
pub struct EnrolWrite {
pub session: AdminSession,
pub user: AdminUser,
pub pending: bool,
}
impl FromRequestParts<AdminState> for PendingMfa {
type Rejection = AdminError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
resolve_pending(parts, state).await
}
}
impl FromRequestParts<AdminState> for PendingMfaSubmit {
type Rejection = AdminError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
check_origin(&parts.headers, &state.config.admin.base_url)?;
Ok(PendingMfaSubmit(resolve_pending(parts, state).await?))
}
}
impl FromRequestParts<AdminState> for EnrolWrite {
type Rejection = AdminError;
async fn from_request_parts(
parts: &mut Parts,
state: &AdminState,
) -> Result<Self, Self::Rejection> {
check_origin(&parts.headers, &state.config.admin.base_url)?;
let (_, session, user) = resolve_live(parts, state).await?;
if !session.is_active() && user.has_totp() {
return Err(AdminError::session_invalid());
}
check_csrf(&parts.headers, &session.csrf_token)?;
let pending = !session.is_active();
Ok(EnrolWrite {
session,
user,
pending,
})
}
}
async fn resolve_live(
parts: &Parts,
state: &AdminState,
) -> Result<(String, AdminSession, AdminUser), AdminError> {
let token = cookie_value(&parts.headers).ok_or_else(AdminError::session_invalid)?;
let token_hash = hash_token(&token);
let Some(session) = AdminSession::find_by_token_hash(&token_hash, &state.database).await?
else {
return Err(AdminError::session_invalid());
};
let now = now_secs();
if session.is_expired(now) {
AdminSession::delete(&token_hash, &state.database).await?;
return Err(AdminError::session_expired());
}
let idle_timeout = Duration::from_secs(state.config.admin.session_idle_timeout_seconds);
if session.is_idle(now, idle_timeout) {
AdminSession::delete(&token_hash, &state.database).await?;
return Err(AdminError::session_idle());
}
let Some(user) = AdminUser::find_by_id(session.user_id, &state.database).await? else {
warn!(event = "admin_session_orphaned", outcome = "failure", session_fp = %fingerprint(&token_hash));
AdminSession::delete(&token_hash, &state.database).await?;
return Err(AdminError::session_invalid());
};
if !user.is_active() {
return Err(AdminError::session_invalid());
}
Ok((token_hash, session, user))
}
async fn resolve_session(parts: &Parts, state: &AdminState) -> Result<Authenticated, AdminError> {
let (_, mut session, user) = resolve_live(parts, state).await?;
if !session.is_active() {
return Err(AdminError::session_invalid());
}
if now_secs() - session.last_seen_at >= SESSION_TOUCH_INTERVAL {
session.touch(&state.database).await?;
}
Ok(Authenticated { session, user })
}
async fn resolve_pending(parts: &Parts, state: &AdminState) -> Result<PendingMfa, AdminError> {
let (_, session, user) = resolve_live(parts, state).await?;
if session.is_active() {
return Err(AdminError::session_invalid());
}
let step = if user.has_totp() {
MfaStep::Verify
} else {
MfaStep::Enrol
};
Ok(PendingMfa {
session,
user,
step,
})
}
pub fn check_csrf(headers: &HeaderMap, expected: &str) -> Result<(), AdminError> {
let Some(supplied) = headers.get(CSRF_HEADER).and_then(|v| v.to_str().ok()) else {
return Err(AdminError::csrf_failed(format!(
"this request needs an {CSRF_HEADER} header carrying the session's csrfToken"
)));
};
let matches = supplied.len() == expected.len()
&& bool::from(supplied.as_bytes().ct_eq(expected.as_bytes()));
if !matches {
return Err(AdminError::csrf_failed(
"the CSRF token does not match this session",
));
}
Ok(())
}
pub fn check_origin(headers: &HeaderMap, base_url: &str) -> Result<(), AdminError> {
if let Some(site) = headers.get("sec-fetch-site").and_then(|v| v.to_str().ok())
&& site != "same-origin"
&& site != "none"
{
return Err(AdminError::csrf_failed(format!(
"cross-origin request refused (Sec-Fetch-Site: {site})"
)));
}
if let Some(origin) = headers.get(header::ORIGIN).and_then(|v| v.to_str().ok()) {
let expected = url::Url::parse(base_url)
.map(|u| u.origin().ascii_serialization())
.unwrap_or_default();
if origin != expected {
return Err(AdminError::csrf_failed(format!(
"cross-origin request refused (Origin: {origin}, expected {expected})"
)));
}
}
Ok(())
}
#[derive(Debug)]
pub struct LoginLimiter {
max_attempts: u32,
window: Duration,
buckets: Mutex<HashMap<IpAddr, Bucket>>,
}
#[derive(Debug, Clone, Copy)]
struct Bucket {
failures: u32,
in_flight: u32,
window_started: i64,
}
fn bucket_key(client: IpAddr) -> IpAddr {
match client.to_canonical() {
IpAddr::V4(v4) => IpAddr::V4(v4),
IpAddr::V6(v6) => {
let prefix = u128::from(v6) & (u128::MAX << 64);
IpAddr::V6(std::net::Ipv6Addr::from(prefix))
}
}
}
impl LoginLimiter {
#[must_use]
pub fn new(max_attempts: u32, window_seconds: u64) -> Self {
Self {
max_attempts,
window: Duration::from_secs(window_seconds),
buckets: Mutex::new(HashMap::new()),
}
}
#[must_use]
pub fn rebuilt(&self, max_attempts: u32, window_seconds: u64) -> Self {
let mut buckets =
std::mem::take(&mut *self.buckets.lock().unwrap_or_else(|e| e.into_inner()));
for bucket in buckets.values_mut() {
bucket.in_flight = 0;
}
Self {
max_attempts,
window: Duration::from_secs(window_seconds),
buckets: Mutex::new(buckets),
}
}
pub fn begin(&self, client: Option<IpAddr>) -> Result<LoginAttempt<'_>, u64> {
let Some(key) = client.map(bucket_key) else {
return Ok(LoginAttempt {
limiter: self,
key: None,
});
};
let now = now_secs();
let window = self.window.as_secs() as i64;
let mut buckets = self.buckets.lock().unwrap_or_else(|e| e.into_inner());
buckets.retain(|_, bucket| now - bucket.window_started < window || bucket.in_flight > 0);
let bucket = buckets.entry(key).or_insert(Bucket {
failures: 0,
in_flight: 0,
window_started: now,
});
if now - bucket.window_started >= window {
bucket.failures = 0;
bucket.window_started = now;
}
if bucket.failures.saturating_add(bucket.in_flight) >= self.max_attempts {
return Err((window - (now - bucket.window_started)).max(1) as u64);
}
bucket.in_flight += 1;
Ok(LoginAttempt {
limiter: self,
key: Some(key),
})
}
pub fn record_success(&self, client: Option<IpAddr>) {
let Some(key) = client.map(bucket_key) else {
return;
};
let mut buckets = self.buckets.lock().unwrap_or_else(|e| e.into_inner());
if let Some(bucket) = buckets.get_mut(&key) {
bucket.failures = 0;
if bucket.in_flight == 0 {
buckets.remove(&key);
}
}
}
fn settle(&self, key: IpAddr, failed: bool) {
let mut buckets = self.buckets.lock().unwrap_or_else(|e| e.into_inner());
let Some(bucket) = buckets.get_mut(&key) else {
return;
};
bucket.in_flight = bucket.in_flight.saturating_sub(1);
if failed {
bucket.failures += 1;
} else if bucket.failures == 0 && bucket.in_flight == 0 {
buckets.remove(&key);
}
}
}
#[derive(Debug)]
#[must_use = "dropping the attempt releases its slot at once"]
pub struct LoginAttempt<'a> {
limiter: &'a LoginLimiter,
key: Option<IpAddr>,
}
impl LoginAttempt<'_> {
pub fn failed(mut self) {
if let Some(key) = self.key.take() {
self.limiter.settle(key, true);
}
}
}
impl Drop for LoginAttempt<'_> {
fn drop(&mut self) {
if let Some(key) = self.key.take() {
self.limiter.settle(key, false);
}
}
}
pub fn log_login(succeeded: bool, username: &str, client: Option<IpAddr>, reason: &'static str) {
if succeeded {
info!(event = "admin_login_succeeded",
outcome = "success",
username = %username,
client_ip = ?client);
} else {
warn!(event = "admin_login_failed",
outcome = "failure",
username = %username,
client_ip = ?client,
reason = reason);
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::http::HeaderValue;
fn headers(pairs: &[(&str, &str)]) -> HeaderMap {
let mut map = HeaderMap::new();
for (name, value) in pairs {
map.append(
header::HeaderName::from_bytes(name.as_bytes()).unwrap(),
HeaderValue::from_str(value).unwrap(),
);
}
map
}
async fn client_ip_of(request: axum::http::Request<()>) -> Option<IpAddr> {
let (mut parts, ()) = request.into_parts();
let AdminClientIp(ip) = AdminClientIp::from_request_parts(&mut parts, &())
.await
.unwrap();
ip
}
#[tokio::test]
async fn the_resolved_client_is_preferred_to_the_peer() {
let mut request = axum::http::Request::new(());
request
.extensions_mut()
.insert(axum::extract::ConnectInfo(std::net::SocketAddr::from((
[172, 18, 0, 2],
4711,
))));
request
.extensions_mut()
.insert(acme_proxy_core::client::ClientIp(Some(
"198.51.100.9".parse().unwrap(),
)));
assert_eq!(
client_ip_of(request).await,
Some("198.51.100.9".parse().unwrap())
);
}
#[tokio::test]
async fn without_the_filter_layer_the_peer_is_the_client() {
let mut request = axum::http::Request::new(());
request.extensions_mut().insert(axum::extract::ConnectInfo(
"[::ffff:192.0.2.7]:4711"
.parse::<std::net::SocketAddr>()
.unwrap(),
));
assert_eq!(
client_ip_of(request).await,
Some("192.0.2.7".parse().unwrap())
);
}
#[test]
fn a_minted_token_is_43_url_safe_characters_and_hashes_stably() {
let minted = mint_token();
assert_eq!(minted.token.len(), 43, "32 bytes, base64url unpadded");
assert!(!minted.token.contains('='));
assert!(!minted.token.contains('+'));
assert!(!minted.token.contains('/'));
assert_eq!(minted.token_hash, hash_token(&minted.token));
assert_eq!(minted.token_hash.len(), 64, "SHA-256 as hex");
assert_ne!(mint_token().token, minted.token);
assert!(!minted.token_hash.contains(&minted.token));
}
#[test]
fn hash_token_matches_a_known_vector() {
assert_eq!(
hash_token(""),
"e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"
);
}
#[test]
fn csrf_tokens_are_unguessable_and_distinct() {
let first = mint_csrf_token();
assert_eq!(first.len(), 43);
assert_ne!(first, mint_csrf_token());
}
#[test]
fn the_session_cookie_carries_every_required_attribute() {
let cookie = session_cookie("the-token", Duration::from_secs(43_200));
assert!(cookie.starts_with("__Host-acme_admin_session=the-token;"));
assert!(cookie.contains("HttpOnly"));
assert!(cookie.contains("Secure"));
assert!(cookie.contains("SameSite=Strict"));
assert!(cookie.contains("Path=/"));
assert!(cookie.contains("Max-Age=43200"));
assert!(!cookie.contains("Domain"));
}
#[test]
fn the_clearing_cookie_expires_immediately_and_keeps_the_same_attributes() {
let cookie = clearing_cookie();
assert!(cookie.starts_with("__Host-acme_admin_session=;"));
assert!(cookie.contains("Max-Age=0"));
assert!(cookie.contains("Path=/"));
assert!(cookie.contains("Secure"));
assert!(cookie.contains("HttpOnly"));
}
type CookieCase = (
&'static str,
Vec<(&'static str, String)>,
Option<&'static str>,
);
#[test]
fn cookie_parsing_is_table_driven() {
let name = COOKIE_NAME;
let cases: Vec<CookieCase> = vec![
("absent entirely", vec![], None),
(
"the only cookie",
vec![("cookie", format!("{name}=abc"))],
Some("abc"),
),
(
"among others",
vec![("cookie", format!("theme=dark; {name}=abc; lang=en"))],
Some("abc"),
),
(
"leading whitespace",
vec![("cookie", format!("theme=dark; {name}=abc"))],
Some("abc"),
),
(
"a quoted value",
vec![("cookie", format!("{name}=\"abc\""))],
Some("abc"),
),
(
"a segment with no equals sign",
vec![("cookie", format!("broken; {name}=abc"))],
Some("abc"),
),
(
"present but empty",
vec![("cookie", format!("{name}="))],
Some(""),
),
(
"a different cookie only",
vec![("cookie", "other=abc".to_string())],
None,
),
(
"duplicated in one header",
vec![("cookie", format!("{name}=first; {name}=second"))],
Some("first"),
),
(
"duplicated across two headers",
vec![
("cookie", format!("{name}=first")),
("cookie", format!("{name}=second")),
],
Some("first"),
),
(
"a name that merely contains ours",
vec![("cookie", format!("x{name}=nope"))],
None,
),
];
for (label, pairs, expected) in cases {
let owned: Vec<(&str, &str)> = pairs.iter().map(|(n, v)| (*n, v.as_str())).collect();
assert_eq!(
cookie_value(&headers(&owned)).as_deref(),
expected,
"case `{label}`"
);
}
}
#[test]
fn the_csrf_check_accepts_only_an_exact_match() {
let expected = "the-expected-token";
assert!(check_csrf(&headers(&[(CSRF_HEADER, expected)]), expected).is_ok());
let error = check_csrf(&HeaderMap::new(), expected).unwrap_err();
assert_eq!(error.code, "csrf_failed");
assert!(error.message.contains(CSRF_HEADER));
for supplied in [
"",
"wrong",
"the-expected-token-but-longer",
"the-expected-toke",
] {
let error = check_csrf(&headers(&[(CSRF_HEADER, supplied)]), expected).unwrap_err();
assert_eq!(error.code, "csrf_failed", "for `{supplied}`");
}
}
#[test]
fn the_origin_gate_covers_the_cases_a_browser_produces() {
let base = "http://localhost:3001";
assert!(check_origin(&HeaderMap::new(), base).is_ok());
assert!(check_origin(&headers(&[("sec-fetch-site", "same-origin")]), base).is_ok());
assert!(check_origin(&headers(&[("sec-fetch-site", "none")]), base).is_ok());
assert!(check_origin(&headers(&[("origin", base)]), base).is_ok());
for site in ["cross-site", "same-site"] {
let error = check_origin(&headers(&[("sec-fetch-site", site)]), base).unwrap_err();
assert_eq!(error.code, "csrf_failed", "for {site}");
assert!(error.message.contains(site));
}
let error = check_origin(&headers(&[("origin", "http://evil.example")]), base).unwrap_err();
assert!(error.message.contains("evil.example"));
let error =
check_origin(&headers(&[("origin", "http://localhost:8080")]), base).unwrap_err();
assert_eq!(error.code, "csrf_failed");
}
fn ip(last: u8) -> Option<IpAddr> {
Some(IpAddr::from([192, 0, 2, last]))
}
fn fail(limiter: &LoginLimiter, client: Option<IpAddr>) {
limiter.begin(client).unwrap().failed();
}
#[test]
fn the_limiter_permits_up_to_the_limit_then_refuses() {
let limiter = LoginLimiter::new(3, 300);
for attempt in 0..3 {
let slot = limiter.begin(ip(1));
assert!(slot.is_ok(), "attempt {attempt} must pass");
slot.unwrap().failed();
}
let retry_after = limiter.begin(ip(1)).unwrap_err();
assert!(retry_after > 0 && retry_after <= 300, "got {retry_after}");
assert!(limiter.begin(ip(2)).is_ok());
}
#[test]
fn attempts_in_flight_count_against_the_limit() {
let limiter = LoginLimiter::new(3, 300);
let held: Vec<_> = (0..3).map(|_| limiter.begin(ip(1)).unwrap()).collect();
assert!(
limiter.begin(ip(1)).is_err(),
"a fourth concurrent attempt must be refused before any has failed"
);
drop(held);
assert!(limiter.begin(ip(1)).is_ok());
assert!(
limiter.buckets.lock().unwrap().is_empty(),
"a bucket with nothing to remember must not linger"
);
}
#[test]
fn a_success_clears_the_counter() {
let limiter = LoginLimiter::new(2, 300);
fail(&limiter, ip(1));
limiter.record_success(ip(1));
fail(&limiter, ip(1));
assert!(
limiter.begin(ip(1)).is_ok(),
"the pre-success failure must not still count"
);
}
#[test]
fn a_success_beside_an_attempt_in_flight_keeps_its_slot() {
let limiter = LoginLimiter::new(1, 300);
let other = limiter.begin(ip(1)).unwrap();
limiter.record_success(ip(1));
assert!(limiter.begin(ip(1)).is_err(), "the slot is still taken");
drop(other);
assert!(limiter.begin(ip(1)).is_ok());
}
#[test]
fn the_window_rolls_over_and_prunes() {
let limiter = LoginLimiter::new(1, 1);
fail(&limiter, ip(1));
assert!(limiter.begin(ip(1)).is_err());
{
let mut buckets = limiter.buckets.lock().unwrap();
buckets.get_mut(&ip(1).unwrap()).unwrap().window_started -= 5;
}
assert!(limiter.begin(ip(1)).is_ok(), "the window must roll over");
assert!(
limiter.buckets.lock().unwrap().is_empty(),
"a stale bucket must be pruned, or the map grows without bound"
);
}
#[test]
fn a_missing_client_address_is_not_limited() {
let limiter = LoginLimiter::new(1, 300);
fail(&limiter, None);
limiter.record_success(None);
assert!(
limiter.begin(None).is_ok(),
"failing closed here would lock out every request, not every attacker"
);
}
#[test]
fn an_ipv6_client_is_limited_by_its_slash_64() {
let limiter = LoginLimiter::new(1, 300);
let v6 = |s: &str| Some(s.parse::<IpAddr>().unwrap());
fail(&limiter, v6("2001:db8:1:2::1"));
assert!(limiter.begin(v6("2001:db8:1:2:ffff::9")).is_err());
assert!(limiter.begin(v6("2001:db8:1:3::1")).is_ok());
fail(&limiter, ip(7));
assert!(limiter.begin(v6("::ffff:192.0.2.7")).is_err());
}
#[test]
fn a_rebuilt_limiter_keeps_failures_and_drops_slots_in_flight() {
let old = LoginLimiter::new(2, 300);
fail(&old, ip(1));
let straddling = old.begin(ip(1)).unwrap();
let new = old.rebuilt(2, 300);
drop(straddling);
assert!(new.begin(ip(1)).is_ok(), "one failure of two is spent");
fail(&new, ip(1));
assert!(new.begin(ip(1)).is_err());
}
#[test]
fn log_login_renders_both_outcomes() {
log_login(true, "alice", ip(1), "");
log_login(false, "alice", ip(1), "wrong_password");
log_login(false, "alice", None, "unknown_user");
}
}