Skip to main content

hammerwork_web/
auth.rs

1//! Authentication middleware for the web dashboard.
2//!
3//! This module provides HTTP Basic authentication against a bcrypt password hash, a lockout
4//! after repeated failures, and a short-lived cache of successful verifications.
5//!
6//! - **Password hashes:** the configured `password_hash` is always a bcrypt hash and is only
7//!   ever *verified* with bcrypt, never compared as text. Verifying bcrypt needs the `auth`
8//!   feature (on by default); a build without it rejects every password, and
9//!   [`DashboardConfig::validate`](crate::config::DashboardConfig::validate) refuses to start
10//!   such a build with authentication enabled.
11//! - **Timing:** the username is compared in constant time through a keyed hash, and bcrypt
12//!   runs whether or not the username matched, so response times do not reveal the username.
13//! - **Blocking:** bcrypt runs on Tokio's blocking thread pool (at most one verification per
14//!   CPU at a time), never on the async worker threads. A successful verification is
15//!   remembered for [`VERIFIED_CREDENTIALS_TTL`] (or `session_timeout`, if shorter), so a
16//!   dashboard polling several endpoints does not pay for bcrypt on every request.
17//! - **Lockout:** failures are counted per client IP address (the connection's address, never
18//!   a header) and username, so an attacker cannot lock the administrator out from other
19//!   addresses. After `max_failed_attempts` failures the client is refused for
20//!   `lockout_duration`; afterwards the count starts over, and a successful login resets it.
21//!   At most [`MAX_TRACKED_CLIENTS`] clients are tracked.
22//!
23//! # Examples
24//!
25//! ## Basic Authentication Setup
26//!
27//! ```rust
28//! use hammerwork_web::auth::AuthState;
29//! use hammerwork_web::config::AuthConfig;
30//!
31//! let auth_config = AuthConfig {
32//!     enabled: true,
33//!     username: "admin".to_string(),
34//!     // A bcrypt hash, for example from `bcrypt::hash(password, bcrypt::DEFAULT_COST)`.
35//!     password_hash: "$2b$12$abcdefghijklmnopqrstuuJ0Y7gZ8z5d0FQqfJb8yX3QZpGQ0lW6e".to_string(),
36//!     ..Default::default()
37//! };
38//!
39//! let auth_state = AuthState::new(auth_config);
40//! assert!(auth_state.is_enabled());
41//! ```
42//!
43//! ## Extracting Basic Auth Credentials
44//!
45//! ```rust
46//! use hammerwork_web::auth::extract_basic_auth;
47//!
48//! // "admin:password" in base64 is "YWRtaW46cGFzc3dvcmQ="
49//! let auth_header = "Basic YWRtaW46cGFzc3dvcmQ=";
50//! let (username, password) = extract_basic_auth(auth_header).unwrap();
51//!
52//! assert_eq!(username, "admin");
53//! assert_eq!(password, "password");
54//!
55//! // Invalid format returns None
56//! let result = extract_basic_auth("Bearer token123");
57//! assert!(result.is_none());
58//! ```
59
60use crate::config::AuthConfig;
61use crate::security::RequestRefused;
62use base64::Engine;
63use hmac::{Hmac, Mac};
64use sha2::Sha256;
65use std::collections::HashMap;
66use std::net::{IpAddr, Ipv6Addr, SocketAddr};
67use std::sync::{Arc, Mutex, MutexGuard};
68use std::time::{Duration, Instant};
69use subtle::ConstantTimeEq;
70use warp::{Filter, Rejection, Reply};
71
72/// The most clients (IP address and username) whose failed attempts are tracked at once.
73///
74/// When the table is full, expired records are dropped first, then the least recently
75/// failed client, so the table cannot grow without bound.
76pub const MAX_TRACKED_CLIENTS: usize = 10_000;
77
78/// How long a successful verification is remembered, unless `session_timeout` is shorter.
79pub const VERIFIED_CREDENTIALS_TTL: Duration = Duration::from_secs(60);
80
81/// The most remembered verifications. Only correct credentials are remembered, so in
82/// practice there is one entry per configured user.
83const MAX_VERIFIED_CREDENTIALS: usize = 64;
84
85/// A keyed SHA-256 digest.
86type Digest = [u8; 32];
87
88/// Whose failures are counted together: one client address and whether it named the
89/// configured user. Every other username shares one record per address, so made-up usernames
90/// cannot add entries.
91#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
92struct AttemptKey {
93    client: Option<IpAddr>,
94    configured_user: bool,
95}
96
97/// The failures of one [`AttemptKey`].
98#[derive(Debug, Clone, Copy)]
99struct Failures {
100    count: u32,
101    last: Instant,
102}
103
104/// The outcome of checking a set of credentials.
105#[derive(Debug, Clone, Copy, PartialEq, Eq)]
106pub enum Verdict {
107    /// The credentials are correct (or authentication is disabled).
108    Accepted,
109    /// The credentials are wrong.
110    Rejected,
111    /// The client failed too often; nothing is checked until the lockout expires.
112    LockedOut,
113}
114
115/// Authentication middleware state. Cloning it shares the state.
116#[derive(Clone)]
117pub struct AuthState {
118    inner: Arc<Inner>,
119}
120
121struct Inner {
122    config: AuthConfig,
123    /// Random per-process key for the digests below.
124    key: Digest,
125    /// Keyed digest of the configured username.
126    username_digest: Digest,
127    failed_attempts: Mutex<HashMap<AttemptKey, Failures>>,
128    /// Keyed digests of recently verified credentials, with their expiry.
129    verified: Mutex<HashMap<Digest, Instant>>,
130    /// Limits concurrent bcrypt verifications to the number of CPUs.
131    #[cfg_attr(not(feature = "auth"), allow(dead_code))]
132    bcrypt_slots: tokio::sync::Semaphore,
133}
134
135impl std::fmt::Debug for AuthState {
136    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
137        f.debug_struct("AuthState")
138            .field("enabled", &self.inner.config.enabled)
139            .field("username", &self.inner.config.username)
140            .finish_non_exhaustive()
141    }
142}
143
144/// Locks `mutex`, recovering the data if another thread panicked while holding it (every
145/// update below leaves the maps consistent).
146fn lock<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
147    mutex
148        .lock()
149        .unwrap_or_else(|poisoned| poisoned.into_inner())
150}
151
152/// HMAC-SHA256 of the length-prefixed `parts` under `key`.
153fn keyed_digest(key: &Digest, parts: &[&[u8]]) -> Digest {
154    let Ok(mut mac) = Hmac::<Sha256>::new_from_slice(key) else {
155        unreachable!("HMAC accepts keys of any length");
156    };
157    for part in parts {
158        mac.update(&(part.len() as u64).to_be_bytes());
159        mac.update(part);
160    }
161    mac.finalize().into_bytes().into()
162}
163
164/// The address failures are counted under: IPv4-mapped IPv6 addresses count as IPv4, and
165/// IPv6 clients are grouped by their /64 network, the usual allocation for one host.
166fn client_key(addr: IpAddr) -> IpAddr {
167    match addr.to_canonical() {
168        IpAddr::V6(v6) => {
169            let mut segments = v6.segments();
170            segments[4..].fill(0);
171            IpAddr::V6(Ipv6Addr::from(segments))
172        }
173        v4 => v4,
174    }
175}
176
177impl AuthState {
178    pub fn new(config: AuthConfig) -> Self {
179        let key: Digest = rand::random();
180        let username_digest = keyed_digest(&key, &[b"username", config.username.as_bytes()]);
181        let slots = std::thread::available_parallelism()
182            .map(|n| n.get())
183            .unwrap_or(4);
184        Self {
185            inner: Arc::new(Inner {
186                config,
187                key,
188                username_digest,
189                failed_attempts: Mutex::new(HashMap::new()),
190                verified: Mutex::new(HashMap::new()),
191                bcrypt_slots: tokio::sync::Semaphore::new(slots),
192            }),
193        }
194    }
195
196    /// Check if authentication is enabled
197    pub fn is_enabled(&self) -> bool {
198        self.inner.config.enabled
199    }
200
201    /// Verify credentials from an unknown client address.
202    ///
203    /// Equivalent to [`check`](Self::check) with no client address, returning whether the
204    /// credentials were accepted.
205    pub async fn verify_credentials(&self, username: &str, password: &str) -> bool {
206        self.check(None, username, password).await == Verdict::Accepted
207    }
208
209    /// Check `username` and `password` for a request from `client`.
210    ///
211    /// A locked-out client gets [`Verdict::LockedOut`] without its credentials being looked
212    /// at. Otherwise a failure is counted against the client, and a success clears its
213    /// failures.
214    pub async fn check(&self, client: Option<IpAddr>, username: &str, password: &str) -> Verdict {
215        if !self.is_enabled() {
216            return Verdict::Accepted;
217        }
218
219        let key = self.attempt_key(client, username);
220        if self.key_locked_out(key) {
221            return Verdict::LockedOut;
222        }
223
224        let credentials = keyed_digest(
225            &self.inner.key,
226            &[b"credentials", username.as_bytes(), password.as_bytes()],
227        );
228        if self.recently_verified(&credentials) {
229            self.clear_failed_attempts(key);
230            return Verdict::Accepted;
231        }
232
233        // bcrypt runs whatever the username, so the time taken does not reveal it.
234        let password_ok = self.verify_password(password).await;
235        if key.configured_user & password_ok {
236            self.remember_verified(credentials);
237            self.clear_failed_attempts(key);
238            Verdict::Accepted
239        } else {
240            self.record_failed_attempt(key);
241            Verdict::Rejected
242        }
243    }
244
245    /// Whether `username` is the configured one, compared in constant time.
246    fn username_matches(&self, username: &str) -> bool {
247        let digest = keyed_digest(&self.inner.key, &[b"username", username.as_bytes()]);
248        digest.ct_eq(&self.inner.username_digest).into()
249    }
250
251    fn attempt_key(&self, client: Option<IpAddr>, username: &str) -> AttemptKey {
252        AttemptKey {
253            client: client.map(client_key),
254            configured_user: self.username_matches(username),
255        }
256    }
257
258    /// Verify `password` against the stored bcrypt hash, off the async worker threads. A
259    /// configuration without a password hash, or a build without the `auth` feature, accepts
260    /// no password at all. The stored hash is never compared to the password as text.
261    async fn verify_password(&self, password: &str) -> bool {
262        if self.inner.config.password_hash.is_empty() {
263            return false;
264        }
265        #[cfg(feature = "auth")]
266        {
267            let Ok(_slot) = self.inner.bcrypt_slots.acquire().await else {
268                return false;
269            };
270            let hash = self.inner.config.password_hash.clone();
271            let password = password.to_string();
272            tokio::task::spawn_blocking(move || bcrypt::verify(password, &hash).unwrap_or(false))
273                .await
274                .unwrap_or(false)
275        }
276        #[cfg(not(feature = "auth"))]
277        {
278            let _ = password;
279            false
280        }
281    }
282
283    /// How long a verification is remembered.
284    fn verified_ttl(&self) -> Duration {
285        VERIFIED_CREDENTIALS_TTL.min(self.inner.config.session_timeout)
286    }
287
288    fn recently_verified(&self, credentials: &Digest) -> bool {
289        let mut verified = lock(&self.inner.verified);
290        match verified.get(credentials) {
291            Some(expires) if *expires > Instant::now() => true,
292            Some(_) => {
293                verified.remove(credentials);
294                false
295            }
296            None => false,
297        }
298    }
299
300    fn remember_verified(&self, credentials: Digest) {
301        let ttl = self.verified_ttl();
302        if ttl.is_zero() {
303            return;
304        }
305        let now = Instant::now();
306        let mut verified = lock(&self.inner.verified);
307        if verified.len() >= MAX_VERIFIED_CREDENTIALS {
308            verified.retain(|_, expires| *expires > now);
309            if verified.len() >= MAX_VERIFIED_CREDENTIALS {
310                verified.clear();
311            }
312        }
313        verified.insert(credentials, now + ttl);
314    }
315
316    /// Whether `username` is locked out when connecting from an unknown address.
317    pub async fn is_locked_out(&self, username: &str) -> bool {
318        self.is_locked_out_from(None, username).await
319    }
320
321    /// Whether `username` is locked out when connecting from `client`.
322    pub async fn is_locked_out_from(&self, client: Option<IpAddr>, username: &str) -> bool {
323        self.key_locked_out(self.attempt_key(client, username))
324    }
325
326    fn key_locked_out(&self, key: AttemptKey) -> bool {
327        let attempts = lock(&self.inner.failed_attempts);
328        attempts.get(&key).is_some_and(|failures| {
329            failures.count >= self.inner.config.max_failed_attempts
330                && failures.last.elapsed() < self.inner.config.lockout_duration
331        })
332    }
333
334    /// Count a failure. A record whose last failure is older than the lockout duration
335    /// starts over, so a lockout always ends.
336    fn record_failed_attempt(&self, key: AttemptKey) {
337        let lockout = self.inner.config.lockout_duration;
338        let now = Instant::now();
339        let mut attempts = lock(&self.inner.failed_attempts);
340        if !attempts.contains_key(&key) && attempts.len() >= MAX_TRACKED_CLIENTS {
341            attempts.retain(|_, failures| now.duration_since(failures.last) < lockout);
342            if attempts.len() >= MAX_TRACKED_CLIENTS
343                && let Some(stalest) = attempts
344                    .iter()
345                    .min_by_key(|(_, failures)| failures.last)
346                    .map(|(key, _)| *key)
347            {
348                attempts.remove(&stalest);
349            }
350        }
351        let failures = attempts.entry(key).or_insert(Failures {
352            count: 0,
353            last: now,
354        });
355        if now.duration_since(failures.last) >= lockout {
356            failures.count = 0;
357        }
358        failures.count = failures.count.saturating_add(1);
359        failures.last = now;
360    }
361
362    fn clear_failed_attempts(&self, key: AttemptKey) {
363        lock(&self.inner.failed_attempts).remove(&key);
364    }
365
366    /// Drop failure records whose lockout has expired, and expired verifications.
367    pub async fn cleanup_expired_attempts(&self) {
368        let lockout = self.inner.config.lockout_duration;
369        let now = Instant::now();
370        lock(&self.inner.failed_attempts)
371            .retain(|_, failures| now.duration_since(failures.last) < lockout);
372        lock(&self.inner.verified).retain(|_, expires| *expires > now);
373    }
374
375    /// The number of clients with failure records (for tests and monitoring).
376    pub fn tracked_clients(&self) -> usize {
377        lock(&self.inner.failed_attempts).len()
378    }
379}
380
381/// Extract basic auth credentials from request.
382///
383/// Parses a Basic Authentication header and returns the username and password.
384/// The header format should be: `Basic <base64-encoded-credentials>`
385/// where credentials are in the format `username:password`.
386///
387/// # Examples
388///
389/// ```rust
390/// use hammerwork_web::auth::extract_basic_auth;
391///
392/// // Valid basic auth header
393/// let auth_header = "Basic YWRtaW46cGFzc3dvcmQ="; // admin:password
394/// let (username, password) = extract_basic_auth(auth_header).unwrap();
395/// assert_eq!(username, "admin");
396/// assert_eq!(password, "password");
397///
398/// // Invalid format returns None
399/// assert!(extract_basic_auth("Bearer token123").is_none());
400/// assert!(extract_basic_auth("Basic invalid_base64").is_none());
401/// ```
402///
403/// # Returns
404///
405/// - `Some((username, password))` if the header is valid
406/// - `None` if the header is malformed or not a Basic auth header
407pub fn extract_basic_auth(auth_header: &str) -> Option<(String, String)> {
408    if !auth_header.starts_with("Basic ") {
409        return None;
410    }
411
412    let encoded = &auth_header[6..];
413    let decoded = ::base64::prelude::BASE64_STANDARD.decode(encoded).ok()?;
414    let decoded_str = String::from_utf8(decoded).ok()?;
415
416    let mut parts = decoded_str.splitn(2, ':');
417    let username = parts.next()?.to_string();
418    let password = parts.next()?.to_string();
419
420    Some((username, password))
421}
422
423/// Authentication filter for Warp.
424///
425/// Failures are counted per client address, taken from the connection (never from a
426/// forwarding header, which a client could forge).
427pub fn auth_filter(
428    auth_state: AuthState,
429) -> impl Filter<Extract = ((),), Error = Rejection> + Clone {
430    warp::header::optional::<String>("authorization")
431        .and(warp::addr::remote())
432        .and_then(
433            move |auth_header: Option<String>, remote: Option<SocketAddr>| {
434                let auth_state = auth_state.clone();
435                async move {
436                    if !auth_state.is_enabled() {
437                        return Ok::<_, Rejection>(());
438                    }
439
440                    let auth_header = auth_header
441                        .ok_or_else(|| warp::reject::custom(AuthError::MissingCredentials))?;
442
443                    let (username, password) = extract_basic_auth(&auth_header)
444                        .ok_or_else(|| warp::reject::custom(AuthError::InvalidFormat))?;
445
446                    match auth_state
447                        .check(remote.map(|addr| addr.ip()), &username, &password)
448                        .await
449                    {
450                        Verdict::Accepted => Ok(()),
451                        Verdict::Rejected => {
452                            Err(warp::reject::custom(AuthError::InvalidCredentials))
453                        }
454                        Verdict::LockedOut => Err(warp::reject::custom(AuthError::AccountLocked)),
455                    }
456                }
457            },
458        )
459}
460
461/// Custom authentication errors
462#[derive(Debug)]
463pub enum AuthError {
464    MissingCredentials,
465    InvalidFormat,
466    InvalidCredentials,
467    AccountLocked,
468}
469
470impl warp::reject::Reject for AuthError {}
471
472/// Handle authentication rejections
473pub async fn handle_auth_rejection(
474    err: Rejection,
475) -> Result<Box<dyn Reply>, std::convert::Infallible> {
476    if let Some(auth_error) = err.find::<AuthError>() {
477        match auth_error {
478            AuthError::MissingCredentials => {
479                let response = warp::reply::with_header(
480                    warp::reply::with_status(
481                        "Authentication required",
482                        warp::http::StatusCode::UNAUTHORIZED,
483                    ),
484                    "WWW-Authenticate",
485                    "Basic realm=\"Hammerwork Dashboard\"",
486                );
487                Ok(Box::new(response))
488            }
489            AuthError::InvalidFormat => {
490                let error_response = serde_json::json!({"error": "Invalid authentication format"});
491                Ok(Box::new(warp::reply::with_status(
492                    warp::reply::json(&error_response),
493                    warp::http::StatusCode::BAD_REQUEST,
494                )))
495            }
496            AuthError::InvalidCredentials => {
497                let error_response = serde_json::json!({"error": "Invalid credentials"});
498                Ok(Box::new(warp::reply::with_status(
499                    warp::reply::json(&error_response),
500                    warp::http::StatusCode::UNAUTHORIZED,
501                )))
502            }
503            AuthError::AccountLocked => {
504                let error_response = serde_json::json!({"error": "Account temporarily locked"});
505                Ok(Box::new(warp::reply::with_status(
506                    warp::reply::json(&error_response),
507                    warp::http::StatusCode::TOO_MANY_REQUESTS,
508                )))
509            }
510        }
511    } else {
512        // Not an auth error: 404, 400 and 405 keep their meaning, anything else is a 500.
513        let (message, status) = if let Some(refused) = err.find::<RequestRefused>() {
514            (refused.message(), refused.status())
515        } else if err.is_not_found() {
516            ("Resource not found", warp::http::StatusCode::NOT_FOUND)
517        } else if err
518            .find::<warp::filters::body::BodyDeserializeError>()
519            .is_some()
520        {
521            ("Invalid request body", warp::http::StatusCode::BAD_REQUEST)
522        } else if err.find::<warp::reject::InvalidQuery>().is_some() {
523            (
524                "Invalid query parameters",
525                warp::http::StatusCode::BAD_REQUEST,
526            )
527        } else if err.find::<warp::reject::PayloadTooLarge>().is_some() {
528            (
529                "Request body too large",
530                warp::http::StatusCode::PAYLOAD_TOO_LARGE,
531            )
532        } else if err.find::<warp::reject::LengthRequired>().is_some() {
533            (
534                "A Content-Length header is required",
535                warp::http::StatusCode::LENGTH_REQUIRED,
536            )
537        } else if err.find::<warp::reject::MethodNotAllowed>().is_some() {
538            // Checked last: a request that reached its route and failed there also carries
539            // the 405s of the sibling routes it did not match.
540            (
541                "Method not allowed",
542                warp::http::StatusCode::METHOD_NOT_ALLOWED,
543            )
544        } else {
545            (
546                "Internal server error",
547                warp::http::StatusCode::INTERNAL_SERVER_ERROR,
548            )
549        };
550        let error_response = serde_json::json!({"error": message});
551        Ok(Box::new(warp::reply::with_status(
552            warp::reply::json(&error_response),
553            status,
554        )))
555    }
556}
557
558#[cfg(test)]
559mod tests {
560    use super::*;
561    use std::net::Ipv4Addr;
562
563    fn ip(last: u8) -> Option<IpAddr> {
564        Some(IpAddr::V4(Ipv4Addr::new(192, 0, 2, last)))
565    }
566
567    #[tokio::test]
568    async fn test_auth_state_creation() {
569        let config = AuthConfig {
570            enabled: true,
571            username: "testuser".to_string(),
572            password_hash: "testhash".to_string(),
573            ..Default::default()
574        };
575
576        let auth_state = AuthState::new(config);
577        assert!(auth_state.is_enabled());
578        let debug = format!("{auth_state:?}");
579        assert!(debug.contains("testuser"), "{debug}");
580        assert!(
581            !debug.contains("testhash"),
582            "the hash is not printed: {debug}"
583        );
584    }
585
586    #[tokio::test]
587    async fn test_disabled_auth() {
588        let config = AuthConfig {
589            enabled: false,
590            ..Default::default()
591        };
592
593        let auth_state = AuthState::new(config);
594        assert!(!auth_state.is_enabled());
595        assert!(auth_state.verify_credentials("anyone", "anything").await);
596    }
597
598    #[tokio::test]
599    async fn test_failed_attempts_tracking() {
600        let config = AuthConfig {
601            enabled: true,
602            username: "admin".to_string(),
603            password_hash: "wronghash".to_string(),
604            max_failed_attempts: 3,
605            lockout_duration: Duration::from_secs(60),
606            ..Default::default()
607        };
608
609        let auth_state = AuthState::new(config);
610
611        // Verify multiple failed attempts
612        for _ in 0..3 {
613            assert!(!auth_state.verify_credentials("admin", "wrongpass").await);
614        }
615
616        // Should be locked out now
617        assert!(auth_state.is_locked_out("admin").await);
618        assert_eq!(
619            auth_state.check(None, "admin", "wrongpass").await,
620            Verdict::LockedOut
621        );
622    }
623
624    #[test]
625    fn test_extract_basic_auth() {
626        // "admin:password" in base64 is "YWRtaW46cGFzc3dvcmQ="
627        let auth_header = "Basic YWRtaW46cGFzc3dvcmQ=";
628        let (username, password) = extract_basic_auth(auth_header).unwrap();
629        assert_eq!(username, "admin");
630        assert_eq!(password, "password");
631    }
632
633    #[test]
634    fn test_extract_basic_auth_invalid() {
635        assert!(extract_basic_auth("Bearer token").is_none());
636        assert!(extract_basic_auth("Basic invalid").is_none());
637    }
638
639    /// The stored form of `password`: always a bcrypt hash.
640    #[cfg(feature = "auth")]
641    fn stored_password(password: &str) -> String {
642        bcrypt::hash(password, 4).unwrap()
643    }
644
645    #[cfg(feature = "auth")]
646    fn auth_config(max_failed_attempts: u32, lockout: Duration) -> AuthConfig {
647        AuthConfig {
648            enabled: true,
649            username: "admin".to_string(),
650            password_hash: stored_password("s3cret"),
651            max_failed_attempts,
652            lockout_duration: lockout,
653            ..Default::default()
654        }
655    }
656
657    fn basic(user: &str, password: &str) -> String {
658        format!(
659            "Basic {}",
660            base64::prelude::BASE64_STANDARD.encode(format!("{user}:{password}"))
661        )
662    }
663
664    #[cfg(feature = "auth")]
665    #[tokio::test]
666    async fn correct_credentials_pass_and_wrong_ones_do_not() {
667        let state = AuthState::new(auth_config(50, Duration::from_secs(60)));
668        assert!(state.verify_credentials("admin", "s3cret").await);
669        assert!(!state.verify_credentials("admin", "wrong").await);
670        assert!(!state.verify_credentials("root", "s3cret").await);
671        assert!(!state.verify_credentials("admi", "s3cret").await);
672        assert!(!state.verify_credentials("admin2", "s3cret").await);
673        assert!(!state.verify_credentials("", "s3cret").await);
674        assert!(!state.verify_credentials("admin", "").await);
675        assert!(!state.verify_credentials("admin", "S3CRET").await);
676        // A remembered verification is for these exact credentials only.
677        assert!(state.verify_credentials("admin", "s3cret").await);
678        assert!(!state.verify_credentials("root", "s3cret").await);
679    }
680
681    #[tokio::test]
682    async fn an_unset_password_accepts_nothing() {
683        // The default configuration has authentication enabled but no password; an empty
684        // password must not match the empty hash.
685        let state = AuthState::new(AuthConfig::default());
686        assert!(state.is_enabled());
687        assert!(!state.verify_credentials("admin", "").await);
688        assert!(!state.verify_credentials("admin", "anything").await);
689    }
690
691    /// H5: the stored hash is never compared to the password as text, in any build.
692    #[tokio::test]
693    async fn the_stored_hash_is_never_a_password() {
694        let state = AuthState::new(AuthConfig {
695            enabled: true,
696            username: "admin".into(),
697            password_hash: "$2b$12$abcdefghijklmnopqrstuuJ0Y7gZ8z5d0FQqfJb8yX3QZpGQ0lW6e".into(),
698            ..Default::default()
699        });
700        let hash = state.inner.config.password_hash.clone();
701        assert!(!state.verify_credentials("admin", &hash).await);
702        // Nor is a plaintext "hash" accepted as its own password.
703        let plain = AuthState::new(AuthConfig {
704            enabled: true,
705            username: "admin".into(),
706            password_hash: "s3cret".into(),
707            ..Default::default()
708        });
709        assert!(!plain.verify_credentials("admin", "s3cret").await);
710    }
711
712    #[cfg(feature = "auth")]
713    #[tokio::test]
714    async fn lockout_blocks_even_the_right_password_until_it_expires() {
715        let state = AuthState::new(auth_config(2, Duration::from_millis(150)));
716        assert!(!state.verify_credentials("admin", "bad").await);
717        assert!(
718            !state.is_locked_out("admin").await,
719            "one failure is not enough"
720        );
721        assert!(!state.verify_credentials("admin", "bad").await);
722        assert!(state.is_locked_out("admin").await);
723        assert!(
724            !state.verify_credentials("admin", "s3cret").await,
725            "locked accounts reject the right password"
726        );
727        assert!(!state.is_locked_out("other").await, "lockout is per user");
728
729        tokio::time::sleep(Duration::from_millis(200)).await;
730        assert!(!state.is_locked_out("admin").await);
731        assert!(state.verify_credentials("admin", "s3cret").await);
732        assert!(
733            !state.is_locked_out("admin").await,
734            "a successful login clears the failures"
735        );
736        assert!(!state.verify_credentials("admin", "bad").await);
737        assert!(
738            !state.is_locked_out("admin").await,
739            "the count started over"
740        );
741    }
742
743    /// H6: a lockout ends, and one failure after it does not lock again.
744    #[tokio::test]
745    async fn an_expired_lockout_starts_the_count_over() {
746        let state = AuthState::new(AuthConfig {
747            enabled: true,
748            username: "admin".into(),
749            password_hash: "unusable".into(),
750            max_failed_attempts: 3,
751            lockout_duration: Duration::from_millis(100),
752            ..Default::default()
753        });
754        for _ in 0..3 {
755            assert_eq!(state.check(ip(1), "admin", "x").await, Verdict::Rejected);
756        }
757        assert_eq!(state.check(ip(1), "admin", "x").await, Verdict::LockedOut);
758        tokio::time::sleep(Duration::from_millis(150)).await;
759        // Before the fix the count stayed at 3, so this single failure re-locked the account.
760        assert_eq!(state.check(ip(1), "admin", "x").await, Verdict::Rejected);
761        assert!(!state.is_locked_out_from(ip(1), "admin").await);
762        assert_eq!(state.check(ip(1), "admin", "x").await, Verdict::Rejected);
763        assert!(!state.is_locked_out_from(ip(1), "admin").await);
764    }
765
766    /// H6: an attacker's failures lock out the attacker's address, not the administrator.
767    #[cfg(feature = "auth")]
768    #[tokio::test]
769    async fn lockout_is_per_client_address() {
770        let state = AuthState::new(auth_config(2, Duration::from_secs(60)));
771        for _ in 0..2 {
772            assert_eq!(
773                state.check(ip(66), "admin", "guess").await,
774                Verdict::Rejected
775            );
776        }
777        assert_eq!(
778            state.check(ip(66), "admin", "s3cret").await,
779            Verdict::LockedOut
780        );
781        assert_eq!(
782            state.check(ip(7), "admin", "s3cret").await,
783            Verdict::Accepted
784        );
785        assert!(state.is_locked_out_from(ip(66), "admin").await);
786        assert!(!state.is_locked_out_from(ip(7), "admin").await);
787    }
788
789    /// H6: made-up usernames share one record per address, and the table has a cap.
790    #[tokio::test]
791    async fn the_failure_table_is_bounded() {
792        let state = AuthState::new(AuthConfig {
793            enabled: true,
794            username: "admin".into(),
795            password_hash: "unusable".into(),
796            max_failed_attempts: 1_000_000,
797            lockout_duration: Duration::from_secs(60),
798            ..Default::default()
799        });
800        for n in 0..200 {
801            state
802                .check(ip(1), &format!("user-{n}-{}", "x".repeat(n)), "x")
803                .await;
804        }
805        assert_eq!(state.tracked_clients(), 1, "one record for unknown users");
806        state.check(ip(1), "admin", "x").await;
807        assert_eq!(state.tracked_clients(), 2, "and one for the real user");
808
809        // Many addresses: never more than the cap; the stalest record is evicted.
810        for n in 0..(MAX_TRACKED_CLIENTS as u32 + 50) {
811            let addr = IpAddr::V4(Ipv4Addr::from(0x0a00_0000 + n));
812            state.record_failed_attempt(state.attempt_key(Some(addr), "admin"));
813        }
814        assert_eq!(state.tracked_clients(), MAX_TRACKED_CLIENTS);
815    }
816
817    #[test]
818    fn ipv6_clients_are_grouped_by_network() {
819        let a: IpAddr = "2001:db8:1:2:aaaa:bbbb:cccc:dddd".parse().unwrap();
820        let b: IpAddr = "2001:db8:1:2::1".parse().unwrap();
821        let c: IpAddr = "2001:db8:1:3::1".parse().unwrap();
822        assert_eq!(client_key(a), client_key(b));
823        assert_ne!(client_key(a), client_key(c));
824        let mapped: IpAddr = "::ffff:192.0.2.9".parse().unwrap();
825        assert_eq!(client_key(mapped), ip(9).unwrap());
826        assert_eq!(client_key(ip(9).unwrap()), ip(9).unwrap());
827    }
828
829    #[tokio::test]
830    async fn expired_failure_records_are_cleaned_up() {
831        let state = AuthState::new(AuthConfig {
832            enabled: true,
833            username: "admin".into(),
834            password_hash: "unusable".into(),
835            lockout_duration: Duration::from_millis(20),
836            ..Default::default()
837        });
838        state.verify_credentials("admin", "bad").await;
839        state.verify_credentials("ghost", "bad").await;
840        assert_eq!(state.tracked_clients(), 2);
841        state.cleanup_expired_attempts().await;
842        assert_eq!(state.tracked_clients(), 2, "still recent");
843        tokio::time::sleep(Duration::from_millis(60)).await;
844        state.cleanup_expired_attempts().await;
845        assert_eq!(state.tracked_clients(), 0);
846    }
847
848    /// H7: a successful verification is remembered briefly, so later requests skip bcrypt.
849    #[cfg(feature = "auth")]
850    #[tokio::test]
851    async fn successful_verifications_are_remembered_briefly() {
852        let state = AuthState::new(auth_config(5, Duration::from_secs(60)));
853        assert!(state.verify_credentials("admin", "s3cret").await);
854        assert_eq!(lock(&state.inner.verified).len(), 1);
855        // Failures are never remembered.
856        assert!(!state.verify_credentials("admin", "nope").await);
857        assert_eq!(lock(&state.inner.verified).len(), 1);
858
859        // Expired entries are not used, and are cleaned up.
860        let credentials = *lock(&state.inner.verified).keys().next().unwrap();
861        lock(&state.inner.verified).insert(credentials, Instant::now());
862        assert!(!state.recently_verified(&credentials));
863        assert!(lock(&state.inner.verified).is_empty());
864        assert!(state.verify_credentials("admin", "s3cret").await);
865        lock(&state.inner.verified).insert(credentials, Instant::now());
866        state.cleanup_expired_attempts().await;
867        assert!(lock(&state.inner.verified).is_empty());
868
869        // The cache is bounded.
870        for n in 0..(MAX_VERIFIED_CREDENTIALS + 5) {
871            state.remember_verified([n as u8; 32]);
872        }
873        assert!(lock(&state.inner.verified).len() <= MAX_VERIFIED_CREDENTIALS);
874
875        // A zero session timeout turns the cache off.
876        let uncached = AuthState::new(AuthConfig {
877            session_timeout: Duration::ZERO,
878            ..auth_config(5, Duration::from_secs(60))
879        });
880        assert!(uncached.verify_credentials("admin", "s3cret").await);
881        assert!(lock(&uncached.inner.verified).is_empty());
882    }
883
884    /// H7: bcrypt does not block the async runtime: on a single-threaded runtime another task
885    /// keeps running while a password is verified.
886    #[cfg(feature = "auth")]
887    #[tokio::test(flavor = "current_thread")]
888    async fn bcrypt_runs_off_the_async_threads() {
889        let state = AuthState::new(AuthConfig {
890            password_hash: bcrypt::hash("s3cret", 10).unwrap(),
891            session_timeout: Duration::ZERO,
892            ..auth_config(5, Duration::from_secs(60))
893        });
894        let ticks = Arc::new(std::sync::atomic::AtomicUsize::new(0));
895        let counter = ticks.clone();
896        let ticker = tokio::spawn(async move {
897            loop {
898                counter.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
899                tokio::time::sleep(Duration::from_millis(1)).await;
900            }
901        });
902        tokio::task::yield_now().await;
903        let before = ticks.load(std::sync::atomic::Ordering::Relaxed);
904        let started = Instant::now();
905        assert!(state.verify_credentials("admin", "s3cret").await);
906        let elapsed = started.elapsed();
907        let during = ticks.load(std::sync::atomic::Ordering::Relaxed) - before;
908        ticker.abort();
909        assert!(
910            during >= 3,
911            "the runtime made no progress during a {elapsed:?} verification ({during} ticks)"
912        );
913    }
914
915    #[test]
916    fn basic_auth_parsing_handles_odd_input() {
917        // The password may contain ':'.
918        let header = basic("admin", "pa:ss:word");
919        assert_eq!(
920            extract_basic_auth(&header),
921            Some(("admin".to_string(), "pa:ss:word".to_string()))
922        );
923        assert_eq!(
924            extract_basic_auth(&basic("", "")),
925            Some((String::new(), String::new()))
926        );
927        // No colon at all, bad base64, non-UTF-8 payload, wrong scheme, wrong case.
928        let no_colon = format!(
929            "Basic {}",
930            base64::prelude::BASE64_STANDARD.encode("nocolon")
931        );
932        assert!(extract_basic_auth(&no_colon).is_none());
933        assert!(extract_basic_auth("Basic !!!").is_none());
934        let binary = format!(
935            "Basic {}",
936            base64::prelude::BASE64_STANDARD.encode([0xff, 0xfe, b':'])
937        );
938        assert!(extract_basic_auth(&binary).is_none());
939        assert!(extract_basic_auth("Digest abc").is_none());
940        assert!(extract_basic_auth("basic YWRtaW46cA==").is_none());
941        assert!(extract_basic_auth("").is_none());
942    }
943
944    #[cfg(feature = "auth")]
945    #[tokio::test]
946    async fn bcrypt_hashes_are_verified_not_compared() {
947        let state = AuthState::new(auth_config(5, Duration::from_secs(60)));
948        let hash = state.inner.config.password_hash.clone();
949        assert!(hash.starts_with("$2"));
950        assert!(state.verify_credentials("admin", "s3cret").await);
951        // Presenting the stored hash itself is not the password.
952        assert!(!state.verify_credentials("admin", &hash).await);
953        // A malformed stored hash accepts nothing.
954        let broken = AuthState::new(AuthConfig {
955            password_hash: "not-a-bcrypt-hash".into(),
956            ..auth_config(5, Duration::from_secs(60))
957        });
958        assert!(
959            !broken
960                .verify_credentials("admin", "not-a-bcrypt-hash")
961                .await
962        );
963    }
964
965    fn ping_route(
966        state: AuthState,
967    ) -> impl Filter<Extract = (impl Reply,), Error = std::convert::Infallible> + Clone {
968        warp::path("api")
969            .and(auth_filter(state))
970            .untuple_one()
971            .and(warp::path("ping"))
972            .map(|| "pong")
973            .recover(handle_auth_rejection)
974    }
975
976    async fn request_from(
977        state: &AuthState,
978        header: Option<&str>,
979        from: Option<IpAddr>,
980    ) -> (u16, String) {
981        let mut request = warp::test::request().path("/api/ping");
982        if let Some(header) = header {
983            request = request.header("authorization", header);
984        }
985        if let Some(from) = from {
986            request = request.remote_addr(SocketAddr::new(from, 40000));
987        }
988        let response = request.reply(&ping_route(state.clone())).await;
989        (
990            response.status().as_u16(),
991            String::from_utf8_lossy(response.body()).to_string(),
992        )
993    }
994
995    async fn request_status(state: &AuthState, header: Option<&str>) -> (u16, String) {
996        request_from(state, header, None).await
997    }
998
999    #[cfg(feature = "auth")]
1000    #[tokio::test]
1001    async fn the_filter_maps_each_failure_to_its_status() {
1002        let state = AuthState::new(auth_config(2, Duration::from_secs(60)));
1003
1004        let missing = warp::test::request()
1005            .path("/api/ping")
1006            .reply(&ping_route(state.clone()))
1007            .await;
1008        assert_eq!(missing.status(), 401);
1009        assert!(
1010            missing
1011                .headers()
1012                .get("www-authenticate")
1013                .unwrap()
1014                .to_str()
1015                .unwrap()
1016                .starts_with("Basic realm=")
1017        );
1018
1019        let (status, body) = request_status(&state, Some("Bearer token")).await;
1020        assert_eq!(status, 400);
1021        assert!(body.contains("Invalid authentication format"), "{body}");
1022
1023        let (status, body) = request_status(&state, Some(&basic("admin", "s3cret"))).await;
1024        assert_eq!((status, body.as_str()), (200, "pong"));
1025
1026        let (status, body) = request_status(&state, Some(&basic("admin", "wrong"))).await;
1027        assert_eq!(status, 401);
1028        assert!(body.contains("Invalid credentials"), "{body}");
1029
1030        // The second failure locks the account: now even the right password is refused.
1031        let (status, _) = request_status(&state, Some(&basic("admin", "wrong"))).await;
1032        assert_eq!(status, 401);
1033        let (status, body) = request_status(&state, Some(&basic("admin", "s3cret"))).await;
1034        assert_eq!(status, 429);
1035        assert!(body.contains("temporarily locked"), "{body}");
1036    }
1037
1038    /// H6: the filter counts failures per connection address.
1039    #[cfg(feature = "auth")]
1040    #[tokio::test]
1041    async fn the_filter_locks_out_the_failing_address_only() {
1042        let state = AuthState::new(auth_config(2, Duration::from_secs(60)));
1043        let wrong = basic("admin", "wrong");
1044        let right = basic("admin", "s3cret");
1045        for _ in 0..2 {
1046            assert_eq!(request_from(&state, Some(&wrong), ip(66)).await.0, 401);
1047        }
1048        assert_eq!(request_from(&state, Some(&right), ip(66)).await.0, 429);
1049        assert_eq!(request_from(&state, Some(&right), ip(7)).await.0, 200);
1050    }
1051
1052    #[tokio::test]
1053    async fn disabled_auth_lets_everything_through() {
1054        let state = AuthState::new(AuthConfig {
1055            enabled: false,
1056            ..Default::default()
1057        });
1058        assert_eq!(request_status(&state, None).await.0, 200);
1059        assert_eq!(request_status(&state, Some("garbage")).await.0, 200);
1060    }
1061
1062    #[tokio::test]
1063    async fn other_rejections_keep_their_status() {
1064        let json_route = warp::path("json")
1065            .and(warp::post())
1066            .and(warp::body::json::<serde_json::Value>())
1067            .map(|_| "ok");
1068        let query_route = warp::path("query")
1069            .and(warp::query::<std::collections::HashMap<String, u32>>())
1070            .map(|_| "ok");
1071        let sized = warp::path("sized")
1072            .and(warp::post())
1073            .and(warp::body::content_length_limit(4))
1074            .and(warp::body::bytes())
1075            .map(|_| "ok");
1076        let failing = warp::path("fail").and_then(|| async {
1077            Err::<String, _>(warp::reject::custom(AuthError::InvalidFormat))
1078        });
1079        let filter = json_route
1080            .or(query_route)
1081            .or(sized)
1082            .or(failing)
1083            .recover(handle_auth_rejection);
1084
1085        let not_found = warp::test::request().path("/nothing").reply(&filter).await;
1086        assert_eq!(not_found.status(), 404);
1087        assert!(String::from_utf8_lossy(not_found.body()).contains("Resource not found"));
1088
1089        let bad_json = warp::test::request()
1090            .method("POST")
1091            .path("/json")
1092            .body("{nope")
1093            .reply(&filter)
1094            .await;
1095        assert_eq!(bad_json.status(), 400);
1096
1097        let bad_query = warp::test::request()
1098            .path("/query?n=abc")
1099            .reply(&filter)
1100            .await;
1101        assert_eq!(bad_query.status(), 400);
1102
1103        let wrong_method = warp::test::request().path("/json").reply(&filter).await;
1104        assert_eq!(wrong_method.status(), 405);
1105
1106        let too_big = warp::test::request()
1107            .method("POST")
1108            .path("/sized")
1109            .header("content-length", "100")
1110            .body("x".repeat(100))
1111            .reply(&filter)
1112            .await;
1113        assert_eq!(too_big.status(), 413);
1114
1115        let custom = warp::test::request().path("/fail").reply(&filter).await;
1116        assert_eq!(custom.status(), 400);
1117    }
1118}