1use actix_session::Session;
25use actix_web::HttpRequest;
26
27use crate::{ActixAdmin, ActixAdminError, ActixAdminErrorType};
28
29pub const CSRF_SESSION_KEY: &str = "_actix_admin_csrf";
31pub const CSRF_HEADER: &str = "X-CSRF-Token";
33pub const CSRF_QUERY_PARAM: &str = "_csrf";
36
37pub type CsrfError = ActixAdminError;
41
42pub fn csrf_token_for(session: &Session) -> Result<String, ActixAdminError> {
46 if let Some(existing) = session.get::<String>(CSRF_SESSION_KEY).unwrap_or(None) {
47 return Ok(existing);
48 }
49 let token = generate_token();
50 session
51 .insert(CSRF_SESSION_KEY, &token)
52 .map_err(|e| ActixAdminError::new(ActixAdminErrorType::InternalError, e.to_string()))?;
53 Ok(token)
54}
55
56pub fn verify_csrf(
62 actix_admin: &ActixAdmin,
63 session: &Session,
64 req: &HttpRequest,
65) -> Result<(), ActixAdminError> {
66 if !actix_admin.configuration.enable_csrf {
67 return Ok(());
68 }
69
70 let expected = session
71 .get::<String>(CSRF_SESSION_KEY)
72 .unwrap_or(None)
73 .ok_or_else(|| {
74 ActixAdminError::new(
75 ActixAdminErrorType::CsrfError,
76 "no CSRF token in session; reload the page and try again",
77 )
78 })?;
79
80 let from_header = req
81 .headers()
82 .get(CSRF_HEADER)
83 .and_then(|v| v.to_str().ok())
84 .map(str::to_owned);
85
86 let from_query = if from_header.is_none() {
87 form_urlencoded::parse(req.query_string().as_bytes())
88 .find(|(k, _)| k == CSRF_QUERY_PARAM)
89 .map(|(_, v)| v.into_owned())
90 } else {
91 None
92 };
93
94 let received = from_header.or(from_query).ok_or_else(|| {
95 ActixAdminError::new(
96 ActixAdminErrorType::CsrfError,
97 "missing CSRF token (expected `X-CSRF-Token` header or `_csrf` query param)",
98 )
99 })?;
100
101 if constant_time_eq(received.as_bytes(), expected.as_bytes()) {
102 Ok(())
103 } else {
104 Err(ActixAdminError::new(
105 ActixAdminErrorType::CsrfError,
106 "CSRF token mismatch",
107 ))
108 }
109}
110
111fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
113 if a.len() != b.len() {
114 return false;
115 }
116 let mut diff: u8 = 0;
117 for (x, y) in a.iter().zip(b.iter()) {
118 diff |= x ^ y;
119 }
120 diff == 0
121}
122
123fn generate_token() -> String {
133 let mut bytes = [0u8; 32];
134 if getrandom::getrandom(&mut bytes).is_err() {
135 fill_fallback(&mut bytes);
136 }
137 base64_url(&bytes)
138}
139
140fn fill_fallback(bytes: &mut [u8; 32]) {
141 use std::sync::atomic::{AtomicU64, Ordering};
142 use std::time::{SystemTime, UNIX_EPOCH};
143 static COUNTER: AtomicU64 = AtomicU64::new(0);
144
145 let now = SystemTime::now()
146 .duration_since(UNIX_EPOCH)
147 .map(|d| d.as_nanos() as u64)
148 .unwrap_or(0);
149 let ctr = COUNTER.fetch_add(1, Ordering::Relaxed);
150 let stack = &now as *const _ as usize as u64;
151
152 let mut state: u64 = now.wrapping_mul(0x9E3779B97F4A7C15) ^ ctr ^ stack;
153 for chunk in bytes.chunks_mut(8) {
154 state = state.wrapping_add(0x9E3779B97F4A7C15);
155 let mut z = state;
156 z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
157 z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
158 z ^= z >> 31;
159 chunk.copy_from_slice(&z.to_le_bytes()[..chunk.len()]);
160 }
161}
162
163fn base64_url(data: &[u8]) -> String {
164 const CHARSET: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_";
165 let mut out = String::with_capacity(data.len().div_ceil(3) * 4);
166 let mut i = 0;
167 while i < data.len() {
168 let b0 = data[i] as u32;
169 let b1 = data.get(i + 1).copied().unwrap_or(0) as u32;
170 let b2 = data.get(i + 2).copied().unwrap_or(0) as u32;
171 let n = (b0 << 16) | (b1 << 8) | b2;
172 out.push(CHARSET[((n >> 18) & 63) as usize] as char);
173 out.push(CHARSET[((n >> 12) & 63) as usize] as char);
174 if i + 1 < data.len() {
175 out.push(CHARSET[((n >> 6) & 63) as usize] as char);
176 }
177 if i + 2 < data.len() {
178 out.push(CHARSET[(n & 63) as usize] as char);
179 }
180 i += 3;
181 }
182 out
183}
184
185#[cfg(test)]
186mod tests {
187 use super::*;
188
189 #[test]
190 fn tokens_are_reasonably_unique() {
191 let a = generate_token();
192 let b = generate_token();
193 assert_ne!(a, b);
194 assert!(a.len() >= 40);
195 }
196
197 #[test]
198 fn constant_time_eq_basic() {
199 assert!(constant_time_eq(b"abc", b"abc"));
200 assert!(!constant_time_eq(b"abc", b"abd"));
201 assert!(!constant_time_eq(b"abc", b"abcd"));
202 }
203}