1use std::collections::HashMap;
35use std::net::IpAddr;
36use std::sync::Mutex;
37use std::time::Duration;
38
39use axum::extract::FromRequestParts;
40use axum::http::request::Parts;
41use axum::http::{HeaderMap, header};
42use base64::Engine as _;
43use base64::engine::general_purpose::URL_SAFE_NO_PAD as BASE64_URL_SAFE_NO_PAD;
44use ring::rand::{SecureRandom, SystemRandom};
45use subtle::ConstantTimeEq;
46use tracing::{info, warn};
47
48use crate::webadmin::AdminState;
49use crate::webadmin::error::AdminError;
50use acme_proxy_store::admin_session::AdminSession;
51use acme_proxy_store::admin_user::AdminRole;
52use acme_proxy_store::admin_user::AdminUser;
53use acme_proxy_store::nonce::fingerprint;
54use acme_proxy_store::nonce::now_secs;
55
56pub const COOKIE_NAME: &str = "__Host-acme_admin_session";
63
64pub const CSRF_HEADER: &str = "x-csrf-token";
66
67const TOKEN_LEN: usize = 32;
70
71const SESSION_TOUCH_INTERVAL: i64 = 60;
76
77pub const PENDING_MFA_TTL: Duration = Duration::from_secs(300);
90
91pub struct MintedToken {
93 pub token: String,
95 pub token_hash: String,
97}
98
99#[must_use]
102pub fn mint_token() -> MintedToken {
103 let mut bytes = [0u8; TOKEN_LEN];
104 SystemRandom::new()
105 .fill(&mut bytes)
106 .expect("system RNG unavailable");
107 let token = BASE64_URL_SAFE_NO_PAD.encode(bytes);
108 let token_hash = hash_token(&token);
109 MintedToken { token, token_hash }
110}
111
112#[must_use]
115pub fn mint_csrf_token() -> String {
116 let mut bytes = [0u8; TOKEN_LEN];
117 SystemRandom::new()
118 .fill(&mut bytes)
119 .expect("system RNG unavailable");
120 BASE64_URL_SAFE_NO_PAD.encode(bytes)
121}
122
123#[must_use]
125pub fn hash_token(token: &str) -> String {
126 let digest = ring::digest::digest(&ring::digest::SHA256, token.as_bytes());
127 hex::encode(digest.as_ref())
128}
129
130#[must_use]
132pub fn session_cookie(token: &str, ttl: Duration) -> String {
133 format!(
134 "{COOKIE_NAME}={token}; HttpOnly; Secure; SameSite=Strict; Path=/; Max-Age={}",
135 ttl.as_secs()
136 )
137}
138
139#[must_use]
141pub fn clearing_cookie() -> String {
142 format!("{COOKIE_NAME}=; HttpOnly; Secure; SameSite=Strict; Path=/; Max-Age=0")
143}
144
145#[must_use]
152pub fn cookie_value(headers: &HeaderMap) -> Option<String> {
153 for header in headers.get_all(header::COOKIE) {
154 let Ok(raw) = header.to_str() else { continue };
155 for pair in raw.split(';') {
156 let Some((name, value)) = pair.split_once('=') else {
157 continue;
158 };
159 if name.trim() == COOKIE_NAME {
160 let value = value.trim();
162 let value = value
163 .strip_prefix('"')
164 .and_then(|v| v.strip_suffix('"'))
165 .unwrap_or(value);
166 return Some(value.to_string());
167 }
168 }
169 }
170 None
171}
172
173#[derive(Debug, Clone, Copy)]
185pub struct AdminClientIp(pub Option<IpAddr>);
186
187impl<S: Sync> FromRequestParts<S> for AdminClientIp {
188 type Rejection = std::convert::Infallible;
189
190 async fn from_request_parts(parts: &mut Parts, _state: &S) -> Result<Self, Self::Rejection> {
191 if let Some(acme_proxy_core::client::ClientIp(resolved)) = parts
192 .extensions
193 .get::<acme_proxy_core::client::ClientIp>()
194 .copied()
195 {
196 return Ok(AdminClientIp(resolved));
197 }
198 let address = parts
199 .extensions
200 .get::<axum::extract::ConnectInfo<std::net::SocketAddr>>()
201 .map(|info| info.0.ip().to_canonical());
205 Ok(AdminClientIp(address))
206 }
207}
208
209#[derive(Debug, Clone, Copy, PartialEq, Eq)]
216pub enum MfaStep {
217 Verify,
219 Enrol,
221}
222
223impl MfaStep {
224 #[must_use]
226 pub fn as_str(self) -> &'static str {
227 match self {
228 MfaStep::Verify => "verify",
229 MfaStep::Enrol => "enrol",
230 }
231 }
232}
233
234#[derive(Debug)]
236pub struct Authenticated {
237 pub session: AdminSession,
238 pub user: AdminUser,
239}
240
241#[derive(Debug)]
250pub struct AuthenticatedWrite(pub Authenticated);
251
252#[derive(Debug)]
257pub struct AdminWrite(pub Authenticated);
258
259#[derive(Debug)]
266pub struct SelfServiceWrite(pub Authenticated);
267
268#[derive(Debug)]
282pub struct AdminRead(pub Authenticated);
283
284impl FromRequestParts<AdminState> for Authenticated {
285 type Rejection = AdminError;
286
287 async fn from_request_parts(
288 parts: &mut Parts,
289 state: &AdminState,
290 ) -> Result<Self, Self::Rejection> {
291 resolve_session(parts, state).await
292 }
293}
294
295async fn resolve_write(parts: &Parts, state: &AdminState) -> Result<Authenticated, AdminError> {
298 check_origin(&parts.headers, &state.config.admin.base_url)?;
301 let authenticated = resolve_session(parts, state).await?;
302 check_csrf(&parts.headers, &authenticated.session.csrf_token)?;
303 Ok(authenticated)
304}
305
306fn require_role(user: &AdminUser, minimum: AdminRole) -> Result<(), AdminError> {
310 if user.role() >= minimum {
311 Ok(())
312 } else {
313 Err(AdminError::insufficient_role())
314 }
315}
316
317impl FromRequestParts<AdminState> for AdminRead {
318 type Rejection = AdminError;
319
320 async fn from_request_parts(
321 parts: &mut Parts,
322 state: &AdminState,
323 ) -> Result<Self, Self::Rejection> {
324 let authenticated = resolve_session(parts, state).await?;
328 require_role(&authenticated.user, AdminRole::Admin)?;
329 Ok(AdminRead(authenticated))
330 }
331}
332
333impl FromRequestParts<AdminState> for AuthenticatedWrite {
334 type Rejection = AdminError;
335
336 async fn from_request_parts(
337 parts: &mut Parts,
338 state: &AdminState,
339 ) -> Result<Self, Self::Rejection> {
340 let authenticated = resolve_write(parts, state).await?;
341 require_role(&authenticated.user, AdminRole::Operator)?;
342 Ok(AuthenticatedWrite(authenticated))
343 }
344}
345
346impl FromRequestParts<AdminState> for AdminWrite {
347 type Rejection = AdminError;
348
349 async fn from_request_parts(
350 parts: &mut Parts,
351 state: &AdminState,
352 ) -> Result<Self, Self::Rejection> {
353 let authenticated = resolve_write(parts, state).await?;
354 require_role(&authenticated.user, AdminRole::Admin)?;
355 Ok(AdminWrite(authenticated))
356 }
357}
358
359impl FromRequestParts<AdminState> for SelfServiceWrite {
360 type Rejection = AdminError;
361
362 async fn from_request_parts(
363 parts: &mut Parts,
364 state: &AdminState,
365 ) -> Result<Self, Self::Rejection> {
366 Ok(SelfServiceWrite(resolve_write(parts, state).await?))
367 }
368}
369
370#[derive(Debug)]
377pub struct PendingMfa {
378 pub session: AdminSession,
379 pub user: AdminUser,
380 pub step: MfaStep,
381}
382
383#[derive(Debug)]
396pub struct PendingMfaSubmit(pub PendingMfa);
397
398#[derive(Debug)]
413pub struct EnrolWrite {
414 pub session: AdminSession,
415 pub user: AdminUser,
416 pub pending: bool,
419}
420
421impl FromRequestParts<AdminState> for PendingMfa {
422 type Rejection = AdminError;
423
424 async fn from_request_parts(
425 parts: &mut Parts,
426 state: &AdminState,
427 ) -> Result<Self, Self::Rejection> {
428 resolve_pending(parts, state).await
429 }
430}
431
432impl FromRequestParts<AdminState> for PendingMfaSubmit {
433 type Rejection = AdminError;
434
435 async fn from_request_parts(
436 parts: &mut Parts,
437 state: &AdminState,
438 ) -> Result<Self, Self::Rejection> {
439 check_origin(&parts.headers, &state.config.admin.base_url)?;
440 Ok(PendingMfaSubmit(resolve_pending(parts, state).await?))
441 }
442}
443
444impl FromRequestParts<AdminState> for EnrolWrite {
445 type Rejection = AdminError;
446
447 async fn from_request_parts(
448 parts: &mut Parts,
449 state: &AdminState,
450 ) -> Result<Self, Self::Rejection> {
451 check_origin(&parts.headers, &state.config.admin.base_url)?;
452 let (_, session, user) = resolve_live(parts, state).await?;
453
454 if !session.is_active() && user.has_totp() {
458 return Err(AdminError::session_invalid());
459 }
460
461 check_csrf(&parts.headers, &session.csrf_token)?;
462 let pending = !session.is_active();
463 Ok(EnrolWrite {
464 session,
465 user,
466 pending,
467 })
468 }
469}
470
471async fn resolve_live(
478 parts: &Parts,
479 state: &AdminState,
480) -> Result<(String, AdminSession, AdminUser), AdminError> {
481 let token = cookie_value(&parts.headers).ok_or_else(AdminError::session_invalid)?;
482 let token_hash = hash_token(&token);
483
484 let Some(session) = AdminSession::find_by_token_hash(&token_hash, &state.database).await?
485 else {
486 return Err(AdminError::session_invalid());
487 };
488
489 let now = now_secs();
490 if session.is_expired(now) {
491 AdminSession::delete(&token_hash, &state.database).await?;
492 return Err(AdminError::session_expired());
493 }
494 let idle_timeout = Duration::from_secs(state.config.admin.session_idle_timeout_seconds);
495 if session.is_idle(now, idle_timeout) {
496 AdminSession::delete(&token_hash, &state.database).await?;
497 return Err(AdminError::session_idle());
498 }
499
500 let Some(user) = AdminUser::find_by_id(session.user_id, &state.database).await? else {
501 warn!(event = "admin_session_orphaned", outcome = "failure", session_fp = %fingerprint(&token_hash));
504 AdminSession::delete(&token_hash, &state.database).await?;
505 return Err(AdminError::session_invalid());
506 };
507 if !user.is_active() {
508 return Err(AdminError::session_invalid());
509 }
510
511 Ok((token_hash, session, user))
512}
513
514async fn resolve_session(parts: &Parts, state: &AdminState) -> Result<Authenticated, AdminError> {
519 let (_, mut session, user) = resolve_live(parts, state).await?;
520
521 if !session.is_active() {
525 return Err(AdminError::session_invalid());
526 }
527
528 if now_secs() - session.last_seen_at >= SESSION_TOUCH_INTERVAL {
529 session.touch(&state.database).await?;
530 }
531
532 Ok(Authenticated { session, user })
533}
534
535async fn resolve_pending(parts: &Parts, state: &AdminState) -> Result<PendingMfa, AdminError> {
542 let (_, session, user) = resolve_live(parts, state).await?;
543
544 if session.is_active() {
545 return Err(AdminError::session_invalid());
546 }
547
548 let step = if user.has_totp() {
549 MfaStep::Verify
550 } else {
551 MfaStep::Enrol
552 };
553 Ok(PendingMfa {
554 session,
555 user,
556 step,
557 })
558}
559
560pub fn check_csrf(headers: &HeaderMap, expected: &str) -> Result<(), AdminError> {
563 let Some(supplied) = headers.get(CSRF_HEADER).and_then(|v| v.to_str().ok()) else {
564 return Err(AdminError::csrf_failed(format!(
565 "this request needs an {CSRF_HEADER} header carrying the session's csrfToken"
566 )));
567 };
568
569 let matches = supplied.len() == expected.len()
572 && bool::from(supplied.as_bytes().ct_eq(expected.as_bytes()));
573 if !matches {
574 return Err(AdminError::csrf_failed(
575 "the CSRF token does not match this session",
576 ));
577 }
578 Ok(())
579}
580
581pub fn check_origin(headers: &HeaderMap, base_url: &str) -> Result<(), AdminError> {
589 if let Some(site) = headers.get("sec-fetch-site").and_then(|v| v.to_str().ok())
590 && site != "same-origin"
591 && site != "none"
592 {
593 return Err(AdminError::csrf_failed(format!(
594 "cross-origin request refused (Sec-Fetch-Site: {site})"
595 )));
596 }
597
598 if let Some(origin) = headers.get(header::ORIGIN).and_then(|v| v.to_str().ok()) {
599 let expected = url::Url::parse(base_url)
600 .map(|u| u.origin().ascii_serialization())
601 .unwrap_or_default();
602 if origin != expected {
603 return Err(AdminError::csrf_failed(format!(
604 "cross-origin request refused (Origin: {origin}, expected {expected})"
605 )));
606 }
607 }
608 Ok(())
609}
610
611#[derive(Debug)]
628pub struct LoginLimiter {
629 max_attempts: u32,
630 window: Duration,
631 buckets: Mutex<HashMap<IpAddr, Bucket>>,
632}
633
634#[derive(Debug, Clone, Copy)]
635struct Bucket {
636 failures: u32,
637 in_flight: u32,
640 window_started: i64,
641}
642
643fn bucket_key(client: IpAddr) -> IpAddr {
649 match client.to_canonical() {
650 IpAddr::V4(v4) => IpAddr::V4(v4),
651 IpAddr::V6(v6) => {
652 let prefix = u128::from(v6) & (u128::MAX << 64);
653 IpAddr::V6(std::net::Ipv6Addr::from(prefix))
654 }
655 }
656}
657
658impl LoginLimiter {
659 #[must_use]
662 pub fn new(max_attempts: u32, window_seconds: u64) -> Self {
663 Self {
664 max_attempts,
665 window: Duration::from_secs(window_seconds),
666 buckets: Mutex::new(HashMap::new()),
667 }
668 }
669
670 #[must_use]
687 pub fn rebuilt(&self, max_attempts: u32, window_seconds: u64) -> Self {
688 let mut buckets =
689 std::mem::take(&mut *self.buckets.lock().unwrap_or_else(|e| e.into_inner()));
690 for bucket in buckets.values_mut() {
691 bucket.in_flight = 0;
692 }
693 Self {
694 max_attempts,
695 window: Duration::from_secs(window_seconds),
696 buckets: Mutex::new(buckets),
697 }
698 }
699
700 pub fn begin(&self, client: Option<IpAddr>) -> Result<LoginAttempt<'_>, u64> {
708 let Some(key) = client.map(bucket_key) else {
714 return Ok(LoginAttempt {
715 limiter: self,
716 key: None,
717 });
718 };
719 let now = now_secs();
720 let window = self.window.as_secs() as i64;
721
722 let mut buckets = self.buckets.lock().unwrap_or_else(|e| e.into_inner());
723 buckets.retain(|_, bucket| now - bucket.window_started < window || bucket.in_flight > 0);
726
727 let bucket = buckets.entry(key).or_insert(Bucket {
728 failures: 0,
729 in_flight: 0,
730 window_started: now,
731 });
732 if now - bucket.window_started >= window {
733 bucket.failures = 0;
734 bucket.window_started = now;
735 }
736 if bucket.failures.saturating_add(bucket.in_flight) >= self.max_attempts {
737 return Err((window - (now - bucket.window_started)).max(1) as u64);
738 }
739 bucket.in_flight += 1;
740 Ok(LoginAttempt {
741 limiter: self,
742 key: Some(key),
743 })
744 }
745
746 pub fn record_success(&self, client: Option<IpAddr>) {
749 let Some(key) = client.map(bucket_key) else {
750 return;
751 };
752 let mut buckets = self.buckets.lock().unwrap_or_else(|e| e.into_inner());
753 if let Some(bucket) = buckets.get_mut(&key) {
754 bucket.failures = 0;
757 if bucket.in_flight == 0 {
758 buckets.remove(&key);
759 }
760 }
761 }
762
763 fn settle(&self, key: IpAddr, failed: bool) {
765 let mut buckets = self.buckets.lock().unwrap_or_else(|e| e.into_inner());
766 let Some(bucket) = buckets.get_mut(&key) else {
768 return;
769 };
770 bucket.in_flight = bucket.in_flight.saturating_sub(1);
771 if failed {
772 bucket.failures += 1;
773 } else if bucket.failures == 0 && bucket.in_flight == 0 {
774 buckets.remove(&key);
775 }
776 }
777}
778
779#[derive(Debug)]
785#[must_use = "dropping the attempt releases its slot at once"]
786pub struct LoginAttempt<'a> {
787 limiter: &'a LoginLimiter,
788 key: Option<IpAddr>,
789}
790
791impl LoginAttempt<'_> {
792 pub fn failed(mut self) {
794 if let Some(key) = self.key.take() {
795 self.limiter.settle(key, true);
796 }
797 }
798}
799
800impl Drop for LoginAttempt<'_> {
801 fn drop(&mut self) {
802 if let Some(key) = self.key.take() {
803 self.limiter.settle(key, false);
804 }
805 }
806}
807
808pub fn log_login(succeeded: bool, username: &str, client: Option<IpAddr>, reason: &'static str) {
810 if succeeded {
811 info!(event = "admin_login_succeeded",
812 outcome = "success",
813 username = %username,
814 client_ip = ?client);
815 } else {
816 warn!(event = "admin_login_failed",
817 outcome = "failure",
818 username = %username,
819 client_ip = ?client,
820 reason = reason);
821 }
822}
823
824#[cfg(test)]
825mod tests {
826 use super::*;
827 use axum::http::HeaderValue;
828
829 fn headers(pairs: &[(&str, &str)]) -> HeaderMap {
830 let mut map = HeaderMap::new();
831 for (name, value) in pairs {
832 map.append(
833 header::HeaderName::from_bytes(name.as_bytes()).unwrap(),
834 HeaderValue::from_str(value).unwrap(),
835 );
836 }
837 map
838 }
839
840 async fn client_ip_of(request: axum::http::Request<()>) -> Option<IpAddr> {
841 let (mut parts, ()) = request.into_parts();
842 let AdminClientIp(ip) = AdminClientIp::from_request_parts(&mut parts, &())
843 .await
844 .unwrap();
845 ip
846 }
847
848 #[tokio::test]
851 async fn the_resolved_client_is_preferred_to_the_peer() {
852 let mut request = axum::http::Request::new(());
853 request
854 .extensions_mut()
855 .insert(axum::extract::ConnectInfo(std::net::SocketAddr::from((
856 [172, 18, 0, 2],
857 4711,
858 ))));
859 request
860 .extensions_mut()
861 .insert(acme_proxy_core::client::ClientIp(Some(
862 "198.51.100.9".parse().unwrap(),
863 )));
864 assert_eq!(
865 client_ip_of(request).await,
866 Some("198.51.100.9".parse().unwrap())
867 );
868 }
869
870 #[tokio::test]
871 async fn without_the_filter_layer_the_peer_is_the_client() {
872 let mut request = axum::http::Request::new(());
873 request.extensions_mut().insert(axum::extract::ConnectInfo(
874 "[::ffff:192.0.2.7]:4711"
875 .parse::<std::net::SocketAddr>()
876 .unwrap(),
877 ));
878 assert_eq!(
879 client_ip_of(request).await,
880 Some("192.0.2.7".parse().unwrap())
881 );
882 }
883
884 #[test]
885 fn a_minted_token_is_43_url_safe_characters_and_hashes_stably() {
886 let minted = mint_token();
887 assert_eq!(minted.token.len(), 43, "32 bytes, base64url unpadded");
888 assert!(!minted.token.contains('='));
889 assert!(!minted.token.contains('+'));
890 assert!(!minted.token.contains('/'));
891 assert_eq!(minted.token_hash, hash_token(&minted.token));
892 assert_eq!(minted.token_hash.len(), 64, "SHA-256 as hex");
893
894 assert_ne!(mint_token().token, minted.token);
896 assert!(!minted.token_hash.contains(&minted.token));
897 }
898
899 #[test]
900 fn hash_token_matches_a_known_vector() {
901 assert_eq!(
903 hash_token(""),
904 "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"
905 );
906 }
907
908 #[test]
909 fn csrf_tokens_are_unguessable_and_distinct() {
910 let first = mint_csrf_token();
911 assert_eq!(first.len(), 43);
912 assert_ne!(first, mint_csrf_token());
913 }
914
915 #[test]
916 fn the_session_cookie_carries_every_required_attribute() {
917 let cookie = session_cookie("the-token", Duration::from_secs(43_200));
918 assert!(cookie.starts_with("__Host-acme_admin_session=the-token;"));
919 assert!(cookie.contains("HttpOnly"));
920 assert!(cookie.contains("Secure"));
921 assert!(cookie.contains("SameSite=Strict"));
922 assert!(cookie.contains("Path=/"));
923 assert!(cookie.contains("Max-Age=43200"));
924 assert!(!cookie.contains("Domain"));
927 }
928
929 #[test]
930 fn the_clearing_cookie_expires_immediately_and_keeps_the_same_attributes() {
931 let cookie = clearing_cookie();
932 assert!(cookie.starts_with("__Host-acme_admin_session=;"));
933 assert!(cookie.contains("Max-Age=0"));
934 assert!(cookie.contains("Path=/"));
937 assert!(cookie.contains("Secure"));
938 assert!(cookie.contains("HttpOnly"));
939 }
940
941 type CookieCase = (
944 &'static str,
945 Vec<(&'static str, String)>,
946 Option<&'static str>,
947 );
948
949 #[test]
950 fn cookie_parsing_is_table_driven() {
951 let name = COOKIE_NAME;
952 let cases: Vec<CookieCase> = vec![
953 ("absent entirely", vec![], None),
954 (
955 "the only cookie",
956 vec![("cookie", format!("{name}=abc"))],
957 Some("abc"),
958 ),
959 (
960 "among others",
961 vec![("cookie", format!("theme=dark; {name}=abc; lang=en"))],
962 Some("abc"),
963 ),
964 (
965 "leading whitespace",
966 vec![("cookie", format!("theme=dark; {name}=abc"))],
967 Some("abc"),
968 ),
969 (
970 "a quoted value",
971 vec![("cookie", format!("{name}=\"abc\""))],
972 Some("abc"),
973 ),
974 (
975 "a segment with no equals sign",
976 vec![("cookie", format!("broken; {name}=abc"))],
977 Some("abc"),
978 ),
979 (
980 "present but empty",
981 vec![("cookie", format!("{name}="))],
982 Some(""),
983 ),
984 (
985 "a different cookie only",
986 vec![("cookie", "other=abc".to_string())],
987 None,
988 ),
989 (
990 "duplicated in one header",
993 vec![("cookie", format!("{name}=first; {name}=second"))],
994 Some("first"),
995 ),
996 (
997 "duplicated across two headers",
998 vec![
999 ("cookie", format!("{name}=first")),
1000 ("cookie", format!("{name}=second")),
1001 ],
1002 Some("first"),
1003 ),
1004 (
1005 "a name that merely contains ours",
1006 vec![("cookie", format!("x{name}=nope"))],
1007 None,
1008 ),
1009 ];
1010
1011 for (label, pairs, expected) in cases {
1012 let owned: Vec<(&str, &str)> = pairs.iter().map(|(n, v)| (*n, v.as_str())).collect();
1013 assert_eq!(
1014 cookie_value(&headers(&owned)).as_deref(),
1015 expected,
1016 "case `{label}`"
1017 );
1018 }
1019 }
1020
1021 #[test]
1022 fn the_csrf_check_accepts_only_an_exact_match() {
1023 let expected = "the-expected-token";
1024 assert!(check_csrf(&headers(&[(CSRF_HEADER, expected)]), expected).is_ok());
1025
1026 let error = check_csrf(&HeaderMap::new(), expected).unwrap_err();
1028 assert_eq!(error.code, "csrf_failed");
1029 assert!(error.message.contains(CSRF_HEADER));
1030
1031 for supplied in [
1033 "",
1034 "wrong",
1035 "the-expected-token-but-longer",
1036 "the-expected-toke",
1037 ] {
1038 let error = check_csrf(&headers(&[(CSRF_HEADER, supplied)]), expected).unwrap_err();
1039 assert_eq!(error.code, "csrf_failed", "for `{supplied}`");
1040 }
1041 }
1042
1043 #[test]
1044 fn the_origin_gate_covers_the_cases_a_browser_produces() {
1045 let base = "http://localhost:3001";
1046
1047 assert!(check_origin(&HeaderMap::new(), base).is_ok());
1049
1050 assert!(check_origin(&headers(&[("sec-fetch-site", "same-origin")]), base).is_ok());
1052 assert!(check_origin(&headers(&[("sec-fetch-site", "none")]), base).is_ok());
1053 assert!(check_origin(&headers(&[("origin", base)]), base).is_ok());
1054
1055 for site in ["cross-site", "same-site"] {
1056 let error = check_origin(&headers(&[("sec-fetch-site", site)]), base).unwrap_err();
1057 assert_eq!(error.code, "csrf_failed", "for {site}");
1058 assert!(error.message.contains(site));
1061 }
1062
1063 let error = check_origin(&headers(&[("origin", "http://evil.example")]), base).unwrap_err();
1064 assert!(error.message.contains("evil.example"));
1065
1066 let error =
1068 check_origin(&headers(&[("origin", "http://localhost:8080")]), base).unwrap_err();
1069 assert_eq!(error.code, "csrf_failed");
1070 }
1071
1072 fn ip(last: u8) -> Option<IpAddr> {
1073 Some(IpAddr::from([192, 0, 2, last]))
1074 }
1075
1076 fn fail(limiter: &LoginLimiter, client: Option<IpAddr>) {
1078 limiter.begin(client).unwrap().failed();
1079 }
1080
1081 #[test]
1082 fn the_limiter_permits_up_to_the_limit_then_refuses() {
1083 let limiter = LoginLimiter::new(3, 300);
1084
1085 for attempt in 0..3 {
1086 let slot = limiter.begin(ip(1));
1087 assert!(slot.is_ok(), "attempt {attempt} must pass");
1088 slot.unwrap().failed();
1089 }
1090
1091 let retry_after = limiter.begin(ip(1)).unwrap_err();
1092 assert!(retry_after > 0 && retry_after <= 300, "got {retry_after}");
1093
1094 assert!(limiter.begin(ip(2)).is_ok());
1096 }
1097
1098 #[test]
1101 fn attempts_in_flight_count_against_the_limit() {
1102 let limiter = LoginLimiter::new(3, 300);
1103 let held: Vec<_> = (0..3).map(|_| limiter.begin(ip(1)).unwrap()).collect();
1104 assert!(
1105 limiter.begin(ip(1)).is_err(),
1106 "a fourth concurrent attempt must be refused before any has failed"
1107 );
1108
1109 drop(held);
1112 assert!(limiter.begin(ip(1)).is_ok());
1113 assert!(
1114 limiter.buckets.lock().unwrap().is_empty(),
1115 "a bucket with nothing to remember must not linger"
1116 );
1117 }
1118
1119 #[test]
1120 fn a_success_clears_the_counter() {
1121 let limiter = LoginLimiter::new(2, 300);
1122 fail(&limiter, ip(1));
1123 limiter.record_success(ip(1));
1124 fail(&limiter, ip(1));
1125 assert!(
1126 limiter.begin(ip(1)).is_ok(),
1127 "the pre-success failure must not still count"
1128 );
1129 }
1130
1131 #[test]
1134 fn a_success_beside_an_attempt_in_flight_keeps_its_slot() {
1135 let limiter = LoginLimiter::new(1, 300);
1136 let other = limiter.begin(ip(1)).unwrap();
1137 limiter.record_success(ip(1));
1138 assert!(limiter.begin(ip(1)).is_err(), "the slot is still taken");
1139 drop(other);
1140 assert!(limiter.begin(ip(1)).is_ok());
1141 }
1142
1143 #[test]
1144 fn the_window_rolls_over_and_prunes() {
1145 let limiter = LoginLimiter::new(1, 1);
1148 fail(&limiter, ip(1));
1149 assert!(limiter.begin(ip(1)).is_err());
1150
1151 {
1153 let mut buckets = limiter.buckets.lock().unwrap();
1154 buckets.get_mut(&ip(1).unwrap()).unwrap().window_started -= 5;
1155 }
1156 assert!(limiter.begin(ip(1)).is_ok(), "the window must roll over");
1157 assert!(
1158 limiter.buckets.lock().unwrap().is_empty(),
1159 "a stale bucket must be pruned, or the map grows without bound"
1160 );
1161 }
1162
1163 #[test]
1164 fn a_missing_client_address_is_not_limited() {
1165 let limiter = LoginLimiter::new(1, 300);
1166 fail(&limiter, None);
1167 limiter.record_success(None);
1168 assert!(
1169 limiter.begin(None).is_ok(),
1170 "failing closed here would lock out every request, not every attacker"
1171 );
1172 }
1173
1174 #[test]
1177 fn an_ipv6_client_is_limited_by_its_slash_64() {
1178 let limiter = LoginLimiter::new(1, 300);
1179 let v6 = |s: &str| Some(s.parse::<IpAddr>().unwrap());
1180
1181 fail(&limiter, v6("2001:db8:1:2::1"));
1182 assert!(limiter.begin(v6("2001:db8:1:2:ffff::9")).is_err());
1183 assert!(limiter.begin(v6("2001:db8:1:3::1")).is_ok());
1184
1185 fail(&limiter, ip(7));
1186 assert!(limiter.begin(v6("::ffff:192.0.2.7")).is_err());
1187 }
1188
1189 #[test]
1192 fn a_rebuilt_limiter_keeps_failures_and_drops_slots_in_flight() {
1193 let old = LoginLimiter::new(2, 300);
1194 fail(&old, ip(1));
1195 let straddling = old.begin(ip(1)).unwrap();
1196
1197 let new = old.rebuilt(2, 300);
1198 drop(straddling);
1199 assert!(new.begin(ip(1)).is_ok(), "one failure of two is spent");
1200 fail(&new, ip(1));
1201 assert!(new.begin(ip(1)).is_err());
1202 }
1203
1204 #[test]
1205 fn log_login_renders_both_outcomes() {
1206 log_login(true, "alice", ip(1), "");
1210 log_login(false, "alice", ip(1), "wrong_password");
1211 log_login(false, "alice", None, "unknown_user");
1212 }
1213}