#[cfg(test)]
use std::sync::Arc;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use crate::error::{AuthError, Result};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SessionData {
pub user_id: String,
pub issued_at: u64,
pub expires_at: u64,
pub refresh_token_hash: String,
}
impl SessionData {
#[must_use]
pub fn is_expired(&self) -> bool {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap_or_default()
.as_secs();
self.expires_at <= now
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TokenPair {
pub access_token: String,
pub refresh_token: String,
pub expires_in: u64,
}
#[async_trait]
pub trait SessionStore: Send + Sync {
async fn create_session(&self, user_id: &str, expires_at: u64) -> Result<TokenPair>;
async fn get_session(&self, refresh_token_hash: &str) -> Result<SessionData>;
async fn revoke_session(&self, refresh_token_hash: &str) -> Result<()>;
async fn revoke_all_sessions(&self, user_id: &str) -> Result<()>;
}
pub fn unix_now() -> Result<u64> {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs())
.map_err(|e| AuthError::SystemTimeError {
message: format!("System clock is before Unix epoch: {e}"),
})
}
#[must_use]
pub fn hash_token(token: &str) -> String {
let mut hasher = Sha256::new();
hasher.update(token.as_bytes());
hex::encode(hasher.finalize())
}
#[must_use]
pub fn generate_refresh_token() -> String {
use base64::Engine;
use rand::Rng;
let random_bytes: Vec<u8> = (0..32).map(|_| rand::rng().random()).collect();
base64::engine::general_purpose::STANDARD.encode(&random_bytes)
}
#[cfg(test)]
pub struct InMemorySessionStore {
sessions: Arc<dashmap::DashMap<String, SessionData>>,
}
#[cfg(test)]
impl InMemorySessionStore {
#[must_use]
pub fn new() -> Self {
Self {
sessions: Arc::new(dashmap::DashMap::new()),
}
}
pub fn clear(&self) {
self.sessions.clear();
}
#[must_use]
pub fn len(&self) -> usize {
self.sessions.len()
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.sessions.is_empty()
}
}
#[cfg(test)]
impl Default for InMemorySessionStore {
fn default() -> Self {
Self::new()
}
}
#[cfg(test)]
#[async_trait]
impl SessionStore for InMemorySessionStore {
async fn create_session(&self, user_id: &str, expires_at: u64) -> Result<TokenPair> {
let refresh_token = generate_refresh_token();
let refresh_token_hash = hash_token(&refresh_token);
let now = unix_now()?;
let session = SessionData {
user_id: user_id.to_string(),
issued_at: now,
expires_at,
refresh_token_hash: refresh_token_hash.clone(),
};
self.sessions.insert(refresh_token_hash, session);
let expires_in = expires_at.saturating_sub(now);
let access_token = format!("access_token_{}", refresh_token);
Ok(TokenPair {
access_token,
refresh_token,
expires_in,
})
}
async fn get_session(&self, refresh_token_hash: &str) -> Result<SessionData> {
self.sessions
.get(refresh_token_hash)
.map(|entry| entry.clone())
.ok_or(AuthError::TokenNotFound)
}
async fn revoke_session(&self, refresh_token_hash: &str) -> Result<()> {
self.sessions.remove(refresh_token_hash).ok_or(AuthError::SessionError {
message: "Session not found".to_string(),
})?;
Ok(())
}
async fn revoke_all_sessions(&self, user_id: &str) -> Result<()> {
let mut to_remove = Vec::new();
for entry in self.sessions.iter() {
if entry.user_id == user_id {
to_remove.push(entry.key().clone());
}
}
for key in to_remove {
self.sessions.remove(&key);
}
Ok(())
}
}