use std::collections::HashMap;
use std::sync::Arc;
use actix_web::dev::ServiceRequest;
use rand::distributions::Alphanumeric;
use rand::{thread_rng, Rng};
use crate::http::security::config::Authenticator;
use crate::http::security::crypto::{NoOpPasswordEncoder, PasswordEncoder};
use crate::http::security::user::User;
#[cfg(feature = "http-basic")]
use crate::http::security::http_basic::extract_basic_auth;
pub struct MemoryAuthenticator {
users: HashMap<String, User>,
logged_users: HashMap<String, String>,
password_encoder: Arc<dyn PasswordEncoder>,
}
impl MemoryAuthenticator {
pub fn new() -> Self {
MemoryAuthenticator {
users: HashMap::new(),
logged_users: HashMap::new(),
password_encoder: Arc::new(NoOpPasswordEncoder),
}
}
pub fn password_encoder<E: PasswordEncoder + 'static>(mut self, encoder: E) -> Self {
self.password_encoder = Arc::new(encoder);
self
}
pub fn with_user(mut self, user: User) -> Self {
use std::collections::hash_map::Entry;
let user_name = user.get_username().to_string();
match self.users.entry(user_name) {
Entry::Occupied(e) => {
eprintln!("Warning: User {} already exists, skipping", e.key());
}
Entry::Vacant(e) => {
e.insert(user);
}
}
self
}
pub fn login(&mut self, user_name: String, password: String) -> Option<String> {
self.users.get(&user_name).and_then(|u| {
if self.password_encoder.matches(&password, u.get_password()) {
let id: String = thread_rng()
.sample_iter(&Alphanumeric)
.take(30)
.map(char::from)
.collect();
self.logged_users.insert(id.clone(), user_name);
Some(id)
} else {
None
}
})
}
pub fn logout(&mut self, id: &str) {
self.logged_users.remove(id);
}
pub(crate) fn verify_credentials(&self, username: &str, password: &str) -> Option<User> {
self.users.get(username).and_then(|user| {
if self.password_encoder.matches(password, user.get_password()) {
Some(user.clone())
} else {
None
}
})
}
}
impl Default for MemoryAuthenticator {
fn default() -> Self {
Self::new()
}
}
impl Clone for MemoryAuthenticator {
fn clone(&self) -> Self {
MemoryAuthenticator {
logged_users: self
.logged_users
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect(),
users: self
.users
.iter()
.map(|(k, v)| (k.clone(), v.clone()))
.collect(),
password_encoder: Arc::clone(&self.password_encoder),
}
}
}
impl Authenticator for MemoryAuthenticator {
fn get_user(&self, req: &ServiceRequest) -> Option<User> {
#[cfg(feature = "http-basic")]
if let Some(user) = extract_basic_auth(req, |username, password| {
self.verify_credentials(username, password)
}) {
return Some(user);
}
let user_name = req.headers().get("user_name")?.to_str().ok()?;
let password = req.headers().get("password")?.to_str().ok()?;
self.verify_credentials(user_name, password)
}
}