use crate::http::security::crypto::PasswordEncoder;
use crate::http::security::User;
use async_trait::async_trait;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::RwLock;
#[derive(Debug)]
pub enum UserDetailsError {
NotFound,
AlreadyExists,
InvalidCredentials,
AccountDisabled,
AccountLocked,
AccountExpired,
CredentialsExpired,
StorageError(String),
Other(String),
}
impl std::fmt::Display for UserDetailsError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
UserDetailsError::NotFound => write!(f, "User not found"),
UserDetailsError::AlreadyExists => write!(f, "User already exists"),
UserDetailsError::InvalidCredentials => write!(f, "Invalid credentials"),
UserDetailsError::AccountDisabled => write!(f, "Account is disabled"),
UserDetailsError::AccountLocked => write!(f, "Account is locked"),
UserDetailsError::AccountExpired => write!(f, "Account is expired"),
UserDetailsError::CredentialsExpired => write!(f, "Credentials are expired"),
UserDetailsError::StorageError(e) => write!(f, "Storage error: {}", e),
UserDetailsError::Other(e) => write!(f, "Error: {}", e),
}
}
}
impl std::error::Error for UserDetailsError {}
#[async_trait]
pub trait UserDetailsService: Send + Sync {
async fn load_user_by_username(&self, username: &str)
-> Result<Option<User>, UserDetailsError>;
async fn user_exists(&self, username: &str) -> Result<bool, UserDetailsError> {
Ok(self.load_user_by_username(username).await?.is_some())
}
}
#[async_trait]
pub trait UserDetailsManager: UserDetailsService {
async fn create_user(&self, user: &User) -> Result<(), UserDetailsError>;
async fn update_user(&self, user: &User) -> Result<(), UserDetailsError>;
async fn delete_user(&self, username: &str) -> Result<(), UserDetailsError>;
async fn change_password(
&self,
username: &str,
old_password: &str,
new_password: &str,
) -> Result<(), UserDetailsError>;
}
#[derive(Clone)]
pub struct InMemoryUserDetailsService {
users: Arc<RwLock<HashMap<String, User>>>,
}
impl Default for InMemoryUserDetailsService {
fn default() -> Self {
Self::new()
}
}
impl InMemoryUserDetailsService {
pub fn new() -> Self {
Self {
users: Arc::new(RwLock::new(HashMap::new())),
}
}
pub async fn add_user(&self, user: User) {
let mut users = self.users.write().await;
users.insert(user.get_username().to_string(), user);
}
pub async fn add_users(&self, users: Vec<User>) {
let mut store = self.users.write().await;
for user in users {
store.insert(user.get_username().to_string(), user);
}
}
}
#[async_trait]
impl UserDetailsService for InMemoryUserDetailsService {
async fn load_user_by_username(
&self,
username: &str,
) -> Result<Option<User>, UserDetailsError> {
let users = self.users.read().await;
Ok(users.get(username).cloned())
}
}
#[async_trait]
impl UserDetailsManager for InMemoryUserDetailsService {
async fn create_user(&self, user: &User) -> Result<(), UserDetailsError> {
let mut users = self.users.write().await;
let username = user.get_username().to_string();
if users.contains_key(&username) {
return Err(UserDetailsError::AlreadyExists);
}
users.insert(username, user.clone());
Ok(())
}
async fn update_user(&self, user: &User) -> Result<(), UserDetailsError> {
let mut users = self.users.write().await;
let username = user.get_username().to_string();
if !users.contains_key(&username) {
return Err(UserDetailsError::NotFound);
}
users.insert(username, user.clone());
Ok(())
}
async fn delete_user(&self, username: &str) -> Result<(), UserDetailsError> {
let mut users = self.users.write().await;
if users.remove(username).is_none() {
return Err(UserDetailsError::NotFound);
}
Ok(())
}
async fn change_password(
&self,
username: &str,
_old_password: &str,
new_password: &str,
) -> Result<(), UserDetailsError> {
let mut users = self.users.write().await;
match users.get_mut(username) {
Some(user) => {
let updated = User::new(user.get_username().to_string(), new_password.to_string())
.roles(user.get_roles())
.authorities(user.get_authorities());
*user = updated;
Ok(())
}
None => Err(UserDetailsError::NotFound),
}
}
}
struct CachedUser {
user: User,
cached_at: Instant,
}
pub struct CachingUserDetailsService<S>
where
S: UserDetailsService,
{
inner: S,
cache: Arc<RwLock<HashMap<String, CachedUser>>>,
ttl: Duration,
}
impl<S> CachingUserDetailsService<S>
where
S: UserDetailsService,
{
pub fn new(inner: S) -> Self {
Self {
inner,
cache: Arc::new(RwLock::new(HashMap::new())),
ttl: Duration::from_secs(300),
}
}
pub fn ttl(mut self, ttl: Duration) -> Self {
self.ttl = ttl;
self
}
pub async fn clear_cache(&self) {
let mut cache = self.cache.write().await;
cache.clear();
}
pub async fn invalidate(&self, username: &str) {
let mut cache = self.cache.write().await;
cache.remove(username);
}
fn is_valid(&self, entry: &CachedUser) -> bool {
entry.cached_at.elapsed() < self.ttl
}
}
#[async_trait]
impl<S> UserDetailsService for CachingUserDetailsService<S>
where
S: UserDetailsService + Send + Sync,
{
async fn load_user_by_username(
&self,
username: &str,
) -> Result<Option<User>, UserDetailsError> {
{
let cache = self.cache.read().await;
if let Some(cached) = cache.get(username) {
if self.is_valid(cached) {
return Ok(Some(cached.user.clone()));
}
}
}
let result = self.inner.load_user_by_username(username).await?;
if let Some(ref user) = result {
let mut cache = self.cache.write().await;
cache.insert(
username.to_string(),
CachedUser {
user: user.clone(),
cached_at: Instant::now(),
},
);
}
Ok(result)
}
}
#[derive(Clone)]
pub struct UserDetailsAuthenticator<S, E>
where
S: UserDetailsService + Clone,
E: PasswordEncoder + Clone,
{
service: Arc<S>,
encoder: Arc<E>,
}
impl<S, E> UserDetailsAuthenticator<S, E>
where
S: UserDetailsService + Clone,
E: PasswordEncoder + Clone,
{
pub fn new(service: S, encoder: E) -> Self {
Self {
service: Arc::new(service),
encoder: Arc::new(encoder),
}
}
pub async fn authenticate(
&self,
username: &str,
password: &str,
) -> Result<User, UserDetailsError> {
let user = self
.service
.load_user_by_username(username)
.await?
.ok_or(UserDetailsError::NotFound)?;
if self.encoder.matches(password, user.get_password()) {
Ok(user)
} else {
Err(UserDetailsError::InvalidCredentials)
}
}
pub fn service(&self) -> &S {
&self.service
}
pub fn encoder(&self) -> &E {
&self.encoder
}
}
#[cfg(test)]
mod tests {
use super::*;
fn test_user() -> User {
User::new("testuser".to_string(), "password".to_string())
.roles(&["USER".into()])
.authorities(&["read".into()])
}
#[tokio::test]
async fn test_in_memory_service() {
let service = InMemoryUserDetailsService::new();
let user = test_user();
service.add_user(user.clone()).await;
let loaded = service.load_user_by_username("testuser").await.unwrap();
assert!(loaded.is_some());
assert_eq!(loaded.unwrap().get_username(), "testuser");
assert!(service.user_exists("testuser").await.unwrap());
assert!(!service.user_exists("unknown").await.unwrap());
}
#[tokio::test]
async fn test_in_memory_manager() {
let service = InMemoryUserDetailsService::new();
let user = test_user();
service.create_user(&user).await.unwrap();
assert!(service.user_exists("testuser").await.unwrap());
let result = service.create_user(&user).await;
assert!(matches!(result, Err(UserDetailsError::AlreadyExists)));
let updated =
User::new("testuser".to_string(), "newpass".to_string()).roles(&["ADMIN".into()]);
service.update_user(&updated).await.unwrap();
let loaded = service
.load_user_by_username("testuser")
.await
.unwrap()
.unwrap();
assert!(loaded.has_role("ADMIN"));
service.delete_user("testuser").await.unwrap();
assert!(!service.user_exists("testuser").await.unwrap());
let result = service.delete_user("testuser").await;
assert!(matches!(result, Err(UserDetailsError::NotFound)));
}
#[tokio::test]
async fn test_caching_service() {
let inner = InMemoryUserDetailsService::new();
inner.add_user(test_user()).await;
let cached = CachingUserDetailsService::new(inner).ttl(Duration::from_secs(60));
let user1 = cached.load_user_by_username("testuser").await.unwrap();
assert!(user1.is_some());
let user2 = cached.load_user_by_username("testuser").await.unwrap();
assert!(user2.is_some());
cached.invalidate("testuser").await;
let user3 = cached.load_user_by_username("testuser").await.unwrap();
assert!(user3.is_some());
}
}