Skip to content
use std::collections::HashMap;

use actix_session::storage::{LoadError, SaveError, SessionKey, SessionStore, UpdateError};
use actix_web::cookie::time::Duration;
use anyhow::anyhow;

use crate::db::Database;

#[derive(Clone)]
pub struct SqliteSessionStore {
    db: Database,
}

impl SqliteSessionStore {
    pub fn new(db: Database) -> Self {
        Self { db }
    }

    fn expires_at(ttl: &Duration) -> i64 {
        chrono::Utc::now()
            .timestamp()
            .saturating_add(ttl.whole_seconds())
    }
}

impl SessionStore for SqliteSessionStore {
    async fn load(
        &self,
        session_key: &SessionKey,
    ) -> Result<Option<HashMap<String, String>>, LoadError> {
        let conn = self
            .db
            .conn()
            .await
            .map_err(|e| LoadError::Other(anyhow!(e)))?;
        let mut rows = conn
            .query(
                "SELECT state FROM sessions WHERE session_key = ?1 AND expires_at > ?2",
                turso::params![session_key.as_ref(), chrono::Utc::now().timestamp()],
            )
            .await
            .map_err(|e| LoadError::Other(anyhow!(e.to_string())))?;

        let Some(row) = rows
            .next()
            .await
            .map_err(|e| LoadError::Other(anyhow!(e.to_string())))?
        else {
            return Ok(None);
        };

        let state: String = row
            .get(0)
            .map_err(|e| LoadError::Other(anyhow!(e.to_string())))?;
        serde_json::from_str(&state)
            .map(Some)
            .map_err(|e| LoadError::Deserialization(anyhow!(e)))
    }

    async fn save(
        &self,
        session_state: HashMap<String, String>,
        ttl: &Duration,
    ) -> Result<SessionKey, SaveError> {
        let session_key = uuid::Uuid::new_v4().to_string();
        let serialized = serde_json::to_string(&session_state)
            .map_err(|e| SaveError::Serialization(anyhow!(e)))?;
        let conn = self
            .db
            .conn()
            .await
            .map_err(|e| SaveError::Other(anyhow!(e)))?;
        let now = chrono::Utc::now().timestamp();
        conn.execute(
            "DELETE FROM sessions WHERE expires_at <= ?1",
            turso::params![now],
        )
        .await
        .map_err(|e| SaveError::Other(anyhow!(e.to_string())))?;
        conn.execute(
            "INSERT INTO sessions (session_key, state, expires_at) VALUES (?1, ?2, ?3)",
            turso::params![session_key.clone(), serialized, Self::expires_at(ttl)],
        )
        .await
        .map_err(|e| SaveError::Other(anyhow!(e.to_string())))?;

        Ok(SessionKey::try_from(session_key).expect("UUID is a valid session key"))
    }

    async fn update(
        &self,
        session_key: SessionKey,
        session_state: HashMap<String, String>,
        ttl: &Duration,
    ) -> Result<SessionKey, UpdateError> {
        let serialized = serde_json::to_string(&session_state)
            .map_err(|e| UpdateError::Serialization(anyhow!(e)))?;
        let updated = self
            .db
            .conn()
            .await
            .map_err(|e| UpdateError::Other(anyhow!(e)))?
            .execute(
                "UPDATE sessions SET state = ?1, expires_at = ?2 WHERE session_key = ?3",
                turso::params![serialized, Self::expires_at(ttl), session_key.as_ref()],
            )
            .await
            .map_err(|e| UpdateError::Other(anyhow!(e.to_string())))?;
        if updated == 0 {
            return self
                .save(session_state, ttl)
                .await
                .map_err(|err| match err {
                    SaveError::Serialization(err) => UpdateError::Serialization(err),
                    SaveError::Other(err) => UpdateError::Other(err),
                });
        }
        Ok(session_key)
    }

    async fn update_ttl(
        &self,
        session_key: &SessionKey,
        ttl: &Duration,
    ) -> Result<(), anyhow::Error> {
        self.db
            .conn()
            .await
            .map_err(|e| anyhow!(e))?
            .execute(
                "UPDATE sessions SET expires_at = ?1 WHERE session_key = ?2",
                turso::params![Self::expires_at(ttl), session_key.as_ref()],
            )
            .await
            .map_err(|e| anyhow!(e.to_string()))?;
        Ok(())
    }

    async fn delete(&self, session_key: &SessionKey) -> Result<(), anyhow::Error> {
        self.db
            .conn()
            .await
            .map_err(|e| anyhow!(e))?
            .execute(
                "DELETE FROM sessions WHERE session_key = ?1",
                turso::params![session_key.as_ref()],
            )
            .await
            .map_err(|e| anyhow!(e.to_string()))?;
        Ok(())
    }
}