1use 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
72pub const MAX_TRACKED_CLIENTS: usize = 10_000;
77
78pub const VERIFIED_CREDENTIALS_TTL: Duration = Duration::from_secs(60);
80
81const MAX_VERIFIED_CREDENTIALS: usize = 64;
84
85type Digest = [u8; 32];
87
88#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
92struct AttemptKey {
93 client: Option<IpAddr>,
94 configured_user: bool,
95}
96
97#[derive(Debug, Clone, Copy)]
99struct Failures {
100 count: u32,
101 last: Instant,
102}
103
104#[derive(Debug, Clone, Copy, PartialEq, Eq)]
106pub enum Verdict {
107 Accepted,
109 Rejected,
111 LockedOut,
113}
114
115#[derive(Clone)]
117pub struct AuthState {
118 inner: Arc<Inner>,
119}
120
121struct Inner {
122 config: AuthConfig,
123 key: Digest,
125 username_digest: Digest,
127 failed_attempts: Mutex<HashMap<AttemptKey, Failures>>,
128 verified: Mutex<HashMap<Digest, Instant>>,
130 #[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
144fn lock<T>(mutex: &Mutex<T>) -> MutexGuard<'_, T> {
147 mutex
148 .lock()
149 .unwrap_or_else(|poisoned| poisoned.into_inner())
150}
151
152fn 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
164fn 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 pub fn is_enabled(&self) -> bool {
198 self.inner.config.enabled
199 }
200
201 pub async fn verify_credentials(&self, username: &str, password: &str) -> bool {
206 self.check(None, username, password).await == Verdict::Accepted
207 }
208
209 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 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 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 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 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 pub async fn is_locked_out(&self, username: &str) -> bool {
318 self.is_locked_out_from(None, username).await
319 }
320
321 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 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 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 pub fn tracked_clients(&self) -> usize {
377 lock(&self.inner.failed_attempts).len()
378 }
379}
380
381pub 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
423pub 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#[derive(Debug)]
463pub enum AuthError {
464 MissingCredentials,
465 InvalidFormat,
466 InvalidCredentials,
467 AccountLocked,
468}
469
470impl warp::reject::Reject for AuthError {}
471
472pub 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 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 (
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 for _ in 0..3 {
613 assert!(!auth_state.verify_credentials("admin", "wrongpass").await);
614 }
615
616 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 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 #[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 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 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 #[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 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 #[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 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 #[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 #[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 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 #[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 assert!(!state.verify_credentials("admin", "nope").await);
857 assert_eq!(lock(&state.inner.verified).len(), 1);
858
859 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 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 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 #[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 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 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 assert!(!state.verify_credentials("admin", &hash).await);
953 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 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 #[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}