use actix_session::Session;
use actix_web::HttpRequest;
use crate::{ActixAdmin, ActixAdminError, ActixAdminErrorType};
pub const CSRF_SESSION_KEY: &str = "_actix_admin_csrf";
pub const CSRF_HEADER: &str = "X-CSRF-Token";
pub const CSRF_QUERY_PARAM: &str = "_csrf";
pub type CsrfError = ActixAdminError;
pub fn csrf_token_for(session: &Session) -> Result<String, ActixAdminError> {
if let Some(existing) = session.get::<String>(CSRF_SESSION_KEY).unwrap_or(None) {
return Ok(existing);
}
let token = generate_token();
session
.insert(CSRF_SESSION_KEY, &token)
.map_err(|e| ActixAdminError::new(ActixAdminErrorType::InternalError, e.to_string()))?;
Ok(token)
}
pub fn verify_csrf(
actix_admin: &ActixAdmin,
session: &Session,
req: &HttpRequest,
) -> Result<(), ActixAdminError> {
if !actix_admin.configuration.enable_csrf {
return Ok(());
}
let expected = session
.get::<String>(CSRF_SESSION_KEY)
.unwrap_or(None)
.ok_or_else(|| {
ActixAdminError::new(
ActixAdminErrorType::CsrfError,
"no CSRF token in session; reload the page and try again",
)
})?;
let from_header = req
.headers()
.get(CSRF_HEADER)
.and_then(|v| v.to_str().ok())
.map(str::to_owned);
let from_query = if from_header.is_none() {
form_urlencoded::parse(req.query_string().as_bytes())
.find(|(k, _)| k == CSRF_QUERY_PARAM)
.map(|(_, v)| v.into_owned())
} else {
None
};
let received = from_header.or(from_query).ok_or_else(|| {
ActixAdminError::new(
ActixAdminErrorType::CsrfError,
"missing CSRF token (expected `X-CSRF-Token` header or `_csrf` query param)",
)
})?;
if constant_time_eq(received.as_bytes(), expected.as_bytes()) {
Ok(())
} else {
Err(ActixAdminError::new(
ActixAdminErrorType::CsrfError,
"CSRF token mismatch",
))
}
}
fn constant_time_eq(a: &[u8], b: &[u8]) -> bool {
if a.len() != b.len() {
return false;
}
let mut diff: u8 = 0;
for (x, y) in a.iter().zip(b.iter()) {
diff |= x ^ y;
}
diff == 0
}
fn generate_token() -> String {
let mut bytes = [0u8; 32];
if getrandom::getrandom(&mut bytes).is_err() {
fill_fallback(&mut bytes);
}
base64_url(&bytes)
}
fn fill_fallback(bytes: &mut [u8; 32]) {
use std::sync::atomic::{AtomicU64, Ordering};
use std::time::{SystemTime, UNIX_EPOCH};
static COUNTER: AtomicU64 = AtomicU64::new(0);
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_nanos() as u64)
.unwrap_or(0);
let ctr = COUNTER.fetch_add(1, Ordering::Relaxed);
let stack = &now as *const _ as usize as u64;
let mut state: u64 = now.wrapping_mul(0x9E3779B97F4A7C15) ^ ctr ^ stack;
for chunk in bytes.chunks_mut(8) {
state = state.wrapping_add(0x9E3779B97F4A7C15);
let mut z = state;
z = (z ^ (z >> 30)).wrapping_mul(0xBF58476D1CE4E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D049BB133111EB);
z ^= z >> 31;
chunk.copy_from_slice(&z.to_le_bytes()[..chunk.len()]);
}
}
fn base64_url(data: &[u8]) -> String {
const CHARSET: &[u8; 64] = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789-_";
let mut out = String::with_capacity(data.len().div_ceil(3) * 4);
let mut i = 0;
while i < data.len() {
let b0 = data[i] as u32;
let b1 = data.get(i + 1).copied().unwrap_or(0) as u32;
let b2 = data.get(i + 2).copied().unwrap_or(0) as u32;
let n = (b0 << 16) | (b1 << 8) | b2;
out.push(CHARSET[((n >> 18) & 63) as usize] as char);
out.push(CHARSET[((n >> 12) & 63) as usize] as char);
if i + 1 < data.len() {
out.push(CHARSET[((n >> 6) & 63) as usize] as char);
}
if i + 2 < data.len() {
out.push(CHARSET[(n & 63) as usize] as char);
}
i += 3;
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn tokens_are_reasonably_unique() {
let a = generate_token();
let b = generate_token();
assert_ne!(a, b);
assert!(a.len() >= 40);
}
#[test]
fn constant_time_eq_basic() {
assert!(constant_time_eq(b"abc", b"abc"));
assert!(!constant_time_eq(b"abc", b"abd"));
assert!(!constant_time_eq(b"abc", b"abcd"));
}
}