Skip to content
use std::collections::HashMap;
use std::sync::{Arc, RwLock};
use std::time::{SystemTime, UNIX_EPOCH};

use actix_web::HttpRequest;

use crate::db::Database;

pub const SESSION_COOKIE_NAME: &str = "notify_session";

/// How long a session stays valid after it is created: one year.
pub const SESSION_TTL_SECS: u64 = 365 * 24 * 60 * 60;

/// Session store backed by the SQLite database so sessions survive restarts.
///
/// Rows live in the `sessions` table; an in-memory cache mirrors them so the
/// synchronous auth extractor can check a cookie without awaiting the database.
/// The cache is populated by [`SessionStore::load`] at startup and updated
/// whenever a session is created.
#[derive(Clone)]
pub struct SessionStore {
    /// Backing database. `None` for the in-memory test store.
    db: Option<Database>,
    /// Session token -> expiry as a Unix timestamp in seconds.
    cache: Arc<RwLock<HashMap<String, u64>>>,
}

impl SessionStore {
    /// Creates a store backed by `db`. Call [`SessionStore::load`] once the
    /// database is initialized to restore previously persisted sessions.
    pub fn new(db: Database) -> Self {
        Self {
            db: Some(db),
            cache: Arc::new(RwLock::new(HashMap::new())),
        }
    }

    /// Creates a purely in-memory store (used by tests).
    #[cfg(test)]
    pub fn in_memory() -> Self {
        Self {
            db: None,
            cache: Arc::new(RwLock::new(HashMap::new())),
        }
    }

    /// Loads persisted sessions into the cache, discarding expired ones.
    pub async fn load(&self) -> Result<(), String> {
        let Some(db) = &self.db else {
            return Ok(());
        };

        let now = now_secs();
        let sessions = db.load_sessions().await?;
        if let Ok(mut cache) = self.cache.write() {
            *cache = sessions
                .into_iter()
                .filter(|(_, expiry)| *expiry > now)
                .collect();
        }

        if let Err(e) = db.delete_expired_sessions(now).await {
            log::warn!("Failed to prune expired sessions: {e}");
        }
        Ok(())
    }

    pub async fn create_session(&self) -> String {
        self.create_session_with_ttl(SESSION_TTL_SECS).await
    }

    /// Creates a session that expires after `ttl_secs`, returning its token.
    pub async fn create_session_with_ttl(&self, ttl_secs: u64) -> String {
        let token = generate_token();
        let expires_at = now_secs().saturating_add(ttl_secs);

        if let Ok(mut cache) = self.cache.write() {
            let now = now_secs();
            cache.retain(|_, expiry| *expiry > now);
            cache.insert(token.clone(), expires_at);
        }

        if let Some(db) = &self.db
            && let Err(e) = db.save_session(&token, expires_at).await
        {
            log::warn!("Failed to persist session: {e}");
        }

        token
    }

    pub fn is_authenticated(&self, req: &HttpRequest) -> bool {
        let Some(cookie) = req.cookie(SESSION_COOKIE_NAME) else {
            return false;
        };
        let Ok(cache) = self.cache.read() else {
            return false;
        };
        cache
            .get(cookie.value())
            .is_some_and(|&expires_at| expires_at > now_secs())
    }
}

fn now_secs() -> u64 {
    SystemTime::now()
        .duration_since(UNIX_EPOCH)
        .map(|d| d.as_secs())
        .unwrap_or(0)
}

fn generate_token() -> String {
    let mut bytes = [0u8; 16];
    if let Ok(mut file) = std::fs::File::open("/dev/urandom") {
        use std::io::Read;
        let _ = file.read_exact(&mut bytes);
    } else {
        let now = std::time::SystemTime::now()
            .duration_since(std::time::UNIX_EPOCH)
            .map(|d| d.as_nanos())
            .unwrap_or(0);
        return format!("{now:x}");
    }
    bytes.iter().map(|b| format!("{b:02x}")).collect()
}

#[cfg(test)]
mod tests {
    use super::*;
    use actix_web::cookie::Cookie;
    use actix_web::test::TestRequest;

    fn temp_db_path(tag: &str) -> String {
        let now = SystemTime::now()
            .duration_since(UNIX_EPOCH)
            .unwrap()
            .as_nanos();
        format!("/tmp/notify_sessions_{tag}_{now}.db")
    }

    fn auth_request(token: &str) -> HttpRequest {
        let cookie = Cookie::build(SESSION_COOKIE_NAME, token).finish();
        TestRequest::default().cookie(cookie).to_http_request()
    }

    #[actix_web::test]
    async fn test_session_store_lifecycle() {
        let store = SessionStore::in_memory();
        let req_unauth = TestRequest::default().to_http_request();
        assert!(!store.is_authenticated(&req_unauth));

        let token = store.create_session().await;
        assert!(store.is_authenticated(&auth_request(&token)));
    }

    #[actix_web::test]
    async fn test_session_expires() {
        let store = SessionStore::in_memory();
        let token = store.create_session_with_ttl(0).await;
        assert!(!store.is_authenticated(&auth_request(&token)));
    }

    #[actix_web::test]
    async fn test_persistent_sessions_survive_reopen() {
        let path = temp_db_path("persist");

        let db = Database::new(&path);
        db.init().await.unwrap();

        let token = {
            let store = SessionStore::new(db.clone());
            store.create_session().await
        };

        // A fresh store loading the same database still authenticates the token.
        let reopened = SessionStore::new(db.clone());
        reopened.load().await.unwrap();
        assert!(reopened.is_authenticated(&auth_request(&token)));

        let _ = std::fs::remove_file(&path);
    }

    #[actix_web::test]
    async fn test_expired_sessions_pruned_on_load() {
        let path = temp_db_path("prune");

        let db = Database::new(&path);
        db.init().await.unwrap();

        let token = {
            let store = SessionStore::new(db.clone());
            store.create_session_with_ttl(0).await
        };

        let reopened = SessionStore::new(db.clone());
        reopened.load().await.unwrap();
        assert!(!reopened.is_authenticated(&auth_request(&token)));

        // Expired rows are removed from the database too.
        let remaining = db.load_sessions().await.unwrap();
        assert!(remaining.is_empty());

        let _ = std::fs::remove_file(&path);
    }

    #[test]
    fn test_in_memory_store_is_not_persistent() {
        let store = SessionStore::in_memory();
        assert!(store.db.is_none());
    }
}