use actix_session::SessionExt;
use actix_web::dev::{forward_ready, Service, ServiceRequest, ServiceResponse, Transform};
use actix_web::http::Method;
use actix_web::{body::EitherBody, Error, HttpMessage, HttpResponse};
use futures_util::future::{ok, LocalBoxFuture, Ready};
use rand::Rng;
use regex::Regex;
use std::rc::Rc;
use std::sync::Arc;
#[derive(Debug, Clone)]
pub struct CsrfToken {
pub token: String,
pub header_name: String,
pub parameter_name: String,
}
impl CsrfToken {
pub fn new(token: String) -> Self {
Self {
token,
header_name: "X-CSRF-TOKEN".to_string(),
parameter_name: "_csrf".to_string(),
}
}
pub fn with_names(token: String, header_name: &str, parameter_name: &str) -> Self {
Self {
token,
header_name: header_name.to_string(),
parameter_name: parameter_name.to_string(),
}
}
pub fn value(&self) -> &str {
&self.token
}
pub fn header_name(&self) -> &str {
&self.header_name
}
pub fn parameter_name(&self) -> &str {
&self.parameter_name
}
}
pub trait CsrfTokenRepository: Send + Sync {
fn generate_token(&self) -> CsrfToken;
fn save_token(&self, req: &ServiceRequest, token: &CsrfToken) -> Result<(), CsrfError>;
fn load_token(&self, req: &ServiceRequest) -> Option<CsrfToken>;
}
#[derive(Clone)]
pub struct SessionCsrfTokenRepository {
session_key: String,
header_name: String,
parameter_name: String,
}
impl Default for SessionCsrfTokenRepository {
fn default() -> Self {
Self::new()
}
}
impl SessionCsrfTokenRepository {
pub fn new() -> Self {
Self {
session_key: "CSRF_TOKEN".to_string(),
header_name: "X-CSRF-TOKEN".to_string(),
parameter_name: "_csrf".to_string(),
}
}
pub fn session_key(mut self, key: &str) -> Self {
self.session_key = key.to_string();
self
}
pub fn header_name(mut self, name: &str) -> Self {
self.header_name = name.to_string();
self
}
pub fn parameter_name(mut self, name: &str) -> Self {
self.parameter_name = name.to_string();
self
}
fn generate_token_value(&self) -> String {
let mut rng = rand::thread_rng();
let bytes: [u8; 32] = rng.gen();
hex::encode(&bytes)
}
}
impl CsrfTokenRepository for SessionCsrfTokenRepository {
fn generate_token(&self) -> CsrfToken {
CsrfToken::with_names(
self.generate_token_value(),
&self.header_name,
&self.parameter_name,
)
}
fn save_token(&self, req: &ServiceRequest, token: &CsrfToken) -> Result<(), CsrfError> {
let session = req.get_session();
session
.insert(&self.session_key, &token.token)
.map_err(|e| CsrfError::StorageError(e.to_string()))
}
fn load_token(&self, req: &ServiceRequest) -> Option<CsrfToken> {
let session = req.get_session();
session
.get::<String>(&self.session_key)
.ok()
.flatten()
.map(|token| CsrfToken::with_names(token, &self.header_name, &self.parameter_name))
}
}
#[derive(Clone)]
pub struct CsrfConfig {
repository: Arc<dyn CsrfTokenRepository>,
protected_methods: Vec<Method>,
ignored_paths: Vec<Regex>,
header_name: String,
parameter_name: String,
}
impl Default for CsrfConfig {
fn default() -> Self {
Self::new()
}
}
impl CsrfConfig {
pub fn new() -> Self {
Self {
repository: Arc::new(SessionCsrfTokenRepository::new()),
protected_methods: vec![Method::POST, Method::PUT, Method::DELETE, Method::PATCH],
ignored_paths: Vec::new(),
header_name: "X-CSRF-TOKEN".to_string(),
parameter_name: "_csrf".to_string(),
}
}
pub fn repository<R: CsrfTokenRepository + 'static>(mut self, repository: R) -> Self {
self.repository = Arc::new(repository);
self
}
pub fn protected_methods(mut self, methods: Vec<Method>) -> Self {
self.protected_methods = methods;
self
}
pub fn ignore_path(mut self, pattern: &str) -> Self {
if let Ok(regex) = Regex::new(pattern) {
self.ignored_paths.push(regex);
}
self
}
pub fn header_name(mut self, name: &str) -> Self {
self.header_name = name.to_string();
self
}
pub fn parameter_name(mut self, name: &str) -> Self {
self.parameter_name = name.to_string();
self
}
fn is_path_ignored(&self, path: &str) -> bool {
self.ignored_paths.iter().any(|regex| regex.is_match(path))
}
fn requires_protection(&self, method: &Method) -> bool {
self.protected_methods.contains(method)
}
}
#[derive(Clone)]
pub struct CsrfProtection {
config: CsrfConfig,
}
impl CsrfProtection {
pub fn new(config: CsrfConfig) -> Self {
Self { config }
}
pub fn default_config() -> Self {
Self::new(CsrfConfig::default())
}
}
impl<S, B> Transform<S, ServiceRequest> for CsrfProtection
where
S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
S::Future: 'static,
B: 'static,
{
type Response = ServiceResponse<EitherBody<B>>;
type Error = Error;
type Transform = CsrfMiddleware<S>;
type InitError = ();
type Future = Ready<Result<Self::Transform, Self::InitError>>;
fn new_transform(&self, service: S) -> Self::Future {
ok(CsrfMiddleware {
service: Rc::new(service),
config: self.config.clone(),
})
}
}
pub struct CsrfMiddleware<S> {
service: Rc<S>,
config: CsrfConfig,
}
impl<S, B> Service<ServiceRequest> for CsrfMiddleware<S>
where
S: Service<ServiceRequest, Response = ServiceResponse<B>, Error = Error> + 'static,
S::Future: 'static,
B: 'static,
{
type Response = ServiceResponse<EitherBody<B>>;
type Error = Error;
type Future = LocalBoxFuture<'static, Result<Self::Response, Self::Error>>;
forward_ready!(service);
fn call(&self, req: ServiceRequest) -> Self::Future {
let service = self.service.clone();
let config = self.config.clone();
Box::pin(async move {
let path = req.path().to_string();
let method = req.method().clone();
if config.is_path_ignored(&path) {
let res = service.call(req).await?;
return Ok(res.map_into_left_body());
}
let token = match config.repository.load_token(&req) {
Some(token) => token,
None => {
let token = config.repository.generate_token();
let _ = config.repository.save_token(&req, &token);
token
}
};
req.extensions_mut().insert(token.clone());
if config.requires_protection(&method) {
let request_token = get_token_from_request(&req, &config);
match request_token {
Some(submitted_token) if submitted_token == token.token => {
let res = service.call(req).await?;
Ok(res.map_into_left_body())
}
Some(_) => {
let response = HttpResponse::Forbidden()
.body("CSRF token mismatch")
.map_into_right_body();
Ok(req.into_response(response))
}
None => {
let response = HttpResponse::Forbidden()
.body("CSRF token missing")
.map_into_right_body();
Ok(req.into_response(response))
}
}
} else {
let res = service.call(req).await?;
Ok(res.map_into_left_body())
}
})
}
}
fn get_token_from_request(req: &ServiceRequest, config: &CsrfConfig) -> Option<String> {
if let Some(header_value) = req.headers().get(&config.header_name) {
if let Ok(token) = header_value.to_str() {
return Some(token.to_string());
}
}
let query_string = req.query_string();
let param_prefix = format!("{}=", config.parameter_name);
for pair in query_string.split('&') {
if pair.starts_with(¶m_prefix) {
return Some(pair[param_prefix.len()..].to_string());
}
}
None
}
#[derive(Debug)]
pub enum CsrfError {
MissingToken,
InvalidToken,
TokenMismatch,
StorageError(String),
}
impl std::fmt::Display for CsrfError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
CsrfError::MissingToken => write!(f, "CSRF token missing"),
CsrfError::InvalidToken => write!(f, "Invalid CSRF token"),
CsrfError::TokenMismatch => write!(f, "CSRF token mismatch"),
CsrfError::StorageError(e) => write!(f, "CSRF storage error: {}", e),
}
}
}
impl std::error::Error for CsrfError {}
mod hex {
const HEX_CHARS: &[u8; 16] = b"0123456789abcdef";
pub fn encode(bytes: &[u8]) -> String {
let mut result = String::with_capacity(bytes.len() * 2);
for byte in bytes {
result.push(HEX_CHARS[(byte >> 4) as usize] as char);
result.push(HEX_CHARS[(byte & 0x0f) as usize] as char);
}
result
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_csrf_token() {
let token = CsrfToken::new("test-token".to_string());
assert_eq!(token.value(), "test-token");
assert_eq!(token.header_name(), "X-CSRF-TOKEN");
assert_eq!(token.parameter_name(), "_csrf");
}
#[test]
fn test_csrf_token_custom_names() {
let token = CsrfToken::with_names("test-token".to_string(), "X-Custom-CSRF", "csrf_token");
assert_eq!(token.header_name(), "X-Custom-CSRF");
assert_eq!(token.parameter_name(), "csrf_token");
}
#[test]
fn test_csrf_config_default() {
let config = CsrfConfig::default();
assert_eq!(config.header_name, "X-CSRF-TOKEN");
assert_eq!(config.parameter_name, "_csrf");
assert!(config.protected_methods.contains(&Method::POST));
assert!(config.protected_methods.contains(&Method::PUT));
assert!(config.protected_methods.contains(&Method::DELETE));
assert!(config.protected_methods.contains(&Method::PATCH));
assert!(!config.protected_methods.contains(&Method::GET));
}
#[test]
fn test_csrf_config_ignore_path() {
let config = CsrfConfig::new()
.ignore_path("/api/.*")
.ignore_path("/webhook");
assert!(config.is_path_ignored("/api/users"));
assert!(config.is_path_ignored("/api/posts/123"));
assert!(config.is_path_ignored("/webhook"));
assert!(!config.is_path_ignored("/admin"));
}
#[test]
fn test_csrf_config_protected_methods() {
let config = CsrfConfig::new().protected_methods(vec![Method::POST]);
assert!(config.requires_protection(&Method::POST));
assert!(!config.requires_protection(&Method::PUT));
assert!(!config.requires_protection(&Method::GET));
}
#[test]
fn test_session_csrf_repository() {
let repo = SessionCsrfTokenRepository::new()
.session_key("MY_CSRF")
.header_name("X-My-CSRF")
.parameter_name("my_csrf");
let token = repo.generate_token();
assert_eq!(token.header_name(), "X-My-CSRF");
assert_eq!(token.parameter_name(), "my_csrf");
assert_eq!(token.token.len(), 64); }
#[test]
fn test_hex_encode() {
assert_eq!(hex::encode(&[0x00]), "00");
assert_eq!(hex::encode(&[0xff]), "ff");
assert_eq!(hex::encode(&[0xde, 0xad, 0xbe, 0xef]), "deadbeef");
}
}