Skip to content
use std::sync::Mutex;

use rusqlite::{Connection, params};
use thiserror::Error;

use crate::auth::{Session, Ticket, User};
use crate::code::{Code, CodeVersion, Lang, TestVersion, MAX_VERSIONS};

#[derive(Debug, Error)]
pub enum StoreError {
    #[error("Database error: {0}")]
    Database(#[from] rusqlite::Error),
    #[error("Not found: {0}")]
    NotFound(String),
}

pub struct DB {
    conn: Mutex<Connection>,
}

impl DB {
    pub fn new(path: &str) -> Result<Self, StoreError> {
        let conn = Connection::open(path)?;
        conn.execute_batch(
            "CREATE TABLE IF NOT EXISTS endpoints (
                id TEXT PRIMARY KEY,
                path TEXT UNIQUE NOT NULL,
                lang TEXT NOT NULL,
                raw_code TEXT NOT NULL,
                wasm BLOB NOT NULL,
                hurl TEXT NOT NULL DEFAULT '',
                user_id TEXT NOT NULL DEFAULT '',
                created_at INTEGER NOT NULL,
                updated_at INTEGER NOT NULL
            );",
        )?;

        let _ = conn.execute_batch("ALTER TABLE endpoints ADD COLUMN hurl TEXT NOT NULL DEFAULT ''");
        let _ = conn.execute_batch("ALTER TABLE endpoints ADD COLUMN user_id TEXT NOT NULL DEFAULT ''");

        conn.execute_batch(
            "CREATE TABLE IF NOT EXISTS users (
                id TEXT PRIMARY KEY,
                email TEXT UNIQUE NOT NULL,
                password_hash TEXT NOT NULL,
                created_at INTEGER NOT NULL,
                updated_at INTEGER NOT NULL
            );",
        )?;

        conn.execute_batch(
            "CREATE TABLE IF NOT EXISTS tickets (
                id TEXT PRIMARY KEY,
                code TEXT UNIQUE NOT NULL,
                created_by TEXT NOT NULL,
                used_by TEXT,
                used_at INTEGER,
                created_at INTEGER NOT NULL
            );",
        )?;

        conn.execute_batch(
            "CREATE TABLE IF NOT EXISTS sessions (
                token TEXT PRIMARY KEY,
                user_id TEXT NOT NULL,
                created_at INTEGER NOT NULL,
                expires_at INTEGER NOT NULL
            );",
        )?;

        conn.execute_batch(
            "CREATE TABLE IF NOT EXISTS code_versions (
                id TEXT PRIMARY KEY,
                endpoint_id TEXT NOT NULL,
                version INTEGER NOT NULL,
                raw_code TEXT NOT NULL,
                lang TEXT NOT NULL,
                wasm BLOB NOT NULL,
                created_at INTEGER NOT NULL
            );",
        )?;

        conn.execute_batch(
            "CREATE TABLE IF NOT EXISTS test_versions (
                id TEXT PRIMARY KEY,
                endpoint_id TEXT NOT NULL,
                version INTEGER NOT NULL,
                hurl TEXT NOT NULL,
                created_at INTEGER NOT NULL
            );",
        )?;

        Ok(Self {
            conn: Mutex::new(conn),
        })
    }

    // --- Endpoint operations ---

    pub async fn save_code(&self, code: &Code) -> Result<(), StoreError> {
        let conn = self.conn.lock().unwrap();
        conn.execute(
            "INSERT OR REPLACE INTO endpoints (id, path, lang, raw_code, wasm, hurl, user_id, created_at, updated_at)
             VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7, ?8, ?9)",
            params![
                code.id(),
                code.path(),
                code.lang().as_str(),
                code.raw(),
                code.wasm(),
                code.hurl(),
                code.user_id(),
                code.created_at() as i64,
                code.updated_at() as i64,
            ],
        )?;
        Ok(())
    }

    pub async fn load_code(&self, path: &str) -> Option<Code> {
        let conn = self.conn.lock().unwrap();
        let mut stmt = conn
            .prepare("SELECT id, path, lang, raw_code, wasm, hurl, user_id, created_at, updated_at FROM endpoints WHERE path = ?1")
            .ok()?;
        let result = stmt.query_row(params![path], |row| Ok(row_to_code(row)));
        result.ok()
    }

    pub async fn load_code_by_id(&self, id: &str) -> Option<Code> {
        let conn = self.conn.lock().unwrap();
        let mut stmt = conn
            .prepare("SELECT id, path, lang, raw_code, wasm, hurl, user_id, created_at, updated_at FROM endpoints WHERE id = ?1")
            .ok()?;
        let result = stmt.query_row(params![id], |row| Ok(row_to_code(row)));
        result.ok()
    }

    pub async fn list_codes(&self) -> Vec<Code> {
        let conn = match self.conn.lock() {
            Ok(c) => c,
            Err(_) => return Vec::new(),
        };
        let mut stmt = match conn.prepare(
            "SELECT id, path, lang, raw_code, wasm, hurl, user_id, created_at, updated_at FROM endpoints ORDER BY path",
        ) {
            Ok(s) => s,
            Err(_) => return Vec::new(),
        };

        let result = stmt.query_map(params![], |row| Ok(row_to_code(row)));

        match result {
            Ok(iter) => iter.filter_map(|r| r.ok()).collect(),
            Err(_) => Vec::new(),
        }
    }

    pub async fn list_codes_by_user(&self, user_id: &str) -> Vec<Code> {
        let conn = match self.conn.lock() {
            Ok(c) => c,
            Err(_) => return Vec::new(),
        };
        let mut stmt = match conn.prepare(
            "SELECT id, path, lang, raw_code, wasm, hurl, user_id, created_at, updated_at FROM endpoints WHERE user_id = ?1 OR user_id = '' ORDER BY path",
        ) {
            Ok(s) => s,
            Err(_) => return Vec::new(),
        };

        let result = stmt.query_map(params![user_id], |row| Ok(row_to_code(row)));

        match result {
            Ok(iter) => iter.filter_map(|r| r.ok()).collect(),
            Err(_) => Vec::new(),
        }
    }

    pub async fn delete_code(&self, path: &str) -> bool {
        let conn = match self.conn.lock() {
            Ok(c) => c,
            Err(_) => return false,
        };
        match conn.execute("DELETE FROM endpoints WHERE path = ?1", params![path]) {
            Ok(rows_affected) => rows_affected > 0,
            Err(_) => false,
        }
    }

    // --- User operations ---

    pub async fn save_user(&self, user: &User) -> Result<(), StoreError> {
        let conn = self.conn.lock().unwrap();
        conn.execute(
            "INSERT OR REPLACE INTO users (id, email, password_hash, created_at, updated_at)
             VALUES (?1, ?2, ?3, ?4, ?5)",
            params![
                user.id,
                user.email,
                user.password_hash,
                user.created_at as i64,
                user.updated_at as i64,
            ],
        )?;
        Ok(())
    }

    pub async fn get_user_by_id(&self, id: &str) -> Option<User> {
        let conn = self.conn.lock().unwrap();
        let mut stmt = conn
            .prepare("SELECT id, email, password_hash, created_at, updated_at FROM users WHERE id = ?1")
            .ok()?;
        let result = stmt.query_row(params![id], |row| Ok(row_to_user(row)));
        result.ok()
    }

    pub async fn get_user_by_email(&self, email: &str) -> Option<User> {
        let conn = self.conn.lock().unwrap();
        let mut stmt = conn
            .prepare("SELECT id, email, password_hash, created_at, updated_at FROM users WHERE email = ?1")
            .ok()?;
        let result = stmt.query_row(params![email], |row| Ok(row_to_user(row)));
        result.ok()
    }

    pub async fn update_user_password(&self, id: &str, new_hash: &str) -> Result<(), StoreError> {
        let conn = self.conn.lock().unwrap();
        let now = std::time::SystemTime::now()
            .duration_since(std::time::UNIX_EPOCH)
            .unwrap_or_default()
            .as_secs();
        conn.execute(
            "UPDATE users SET password_hash = ?1, updated_at = ?2 WHERE id = ?3",
            params![new_hash, now as i64, id],
        )?;
        Ok(())
    }

    // --- Ticket operations ---

    pub async fn save_ticket(&self, ticket: &Ticket) -> Result<(), StoreError> {
        let conn = self.conn.lock().unwrap();
        conn.execute(
            "INSERT OR REPLACE INTO tickets (id, code, created_by, used_by, used_at, created_at)
             VALUES (?1, ?2, ?3, ?4, ?5, ?6)",
            params![
                ticket.id,
                ticket.code,
                ticket.created_by,
                ticket.used_by,
                ticket.used_at.map(|t| t as i64),
                ticket.created_at as i64,
            ],
        )?;
        Ok(())
    }

    pub async fn get_ticket_by_code(&self, code: &str) -> Option<Ticket> {
        let conn = self.conn.lock().unwrap();
        let mut stmt = conn
            .prepare("SELECT id, code, created_by, used_by, used_at, created_at FROM tickets WHERE code = ?1")
            .ok()?;
        let result = stmt.query_row(params![code], |row| Ok(row_to_ticket(row)));
        result.ok()
    }

    pub async fn use_ticket(&self, code: &str, used_by: &str) -> Result<(), StoreError> {
        let conn = self.conn.lock().unwrap();
        let now = std::time::SystemTime::now()
            .duration_since(std::time::UNIX_EPOCH)
            .unwrap_or_default()
            .as_secs();
        conn.execute(
            "UPDATE tickets SET used_by = ?1, used_at = ?2 WHERE code = ?3",
            params![used_by, now as i64, code],
        )?;
        Ok(())
    }

    pub async fn list_tickets(&self) -> Vec<Ticket> {
        let conn = match self.conn.lock() {
            Ok(c) => c,
            Err(_) => return Vec::new(),
        };
        let mut stmt = match conn.prepare(
            "SELECT id, code, created_by, used_by, used_at, created_at FROM tickets ORDER BY created_at DESC",
        ) {
            Ok(s) => s,
            Err(_) => return Vec::new(),
        };

        let result = stmt.query_map(params![], |row| Ok(row_to_ticket(row)));

        match result {
            Ok(iter) => iter.filter_map(|r| r.ok()).collect(),
            Err(_) => Vec::new(),
        }
    }

    pub async fn delete_ticket(&self, id: &str) -> bool {
        let conn = match self.conn.lock() {
            Ok(c) => c,
            Err(_) => return false,
        };
        match conn.execute("DELETE FROM tickets WHERE id = ?1", params![id]) {
            Ok(rows_affected) => rows_affected > 0,
            Err(_) => false,
        }
    }

    // --- Session operations ---

    pub async fn save_session(&self, session: &Session) -> Result<(), StoreError> {
        let conn = self.conn.lock().unwrap();
        conn.execute(
            "INSERT OR REPLACE INTO sessions (token, user_id, created_at, expires_at)
             VALUES (?1, ?2, ?3, ?4)",
            params![
                session.token,
                session.user_id,
                session.created_at as i64,
                session.expires_at as i64,
            ],
        )?;
        Ok(())
    }

    pub async fn get_session(&self, token: &str) -> Option<Session> {
        let conn = self.conn.lock().unwrap();
        let mut stmt = conn
            .prepare("SELECT token, user_id, created_at, expires_at FROM sessions WHERE token = ?1")
            .ok()?;
        let result = stmt.query_row(params![token], |row| Ok(row_to_session(row)));
        result.ok()
    }

    pub async fn delete_session(&self, token: &str) -> bool {
        let conn = match self.conn.lock() {
            Ok(c) => c,
            Err(_) => return false,
        };
        match conn.execute("DELETE FROM sessions WHERE token = ?1", params![token]) {
            Ok(rows_affected) => rows_affected > 0,
            Err(_) => false,
        }
    }

    pub async fn cleanup_expired_sessions(&self) {
        let conn = match self.conn.lock() {
            Ok(c) => c,
            Err(_) => return,
        };
        let now = std::time::SystemTime::now()
            .duration_since(std::time::UNIX_EPOCH)
            .unwrap_or_default()
            .as_secs();
        let _ = conn.execute("DELETE FROM sessions WHERE expires_at < ?1", params![now as i64]);
    }

    pub async fn save_code_version(&self, version: &CodeVersion) -> Result<(), StoreError> {
        let conn = self.conn.lock().unwrap();
        let existing: Vec<CodeVersion> = match conn.prepare(
            "SELECT id, endpoint_id, version, raw_code, lang, wasm, created_at FROM code_versions WHERE endpoint_id = ?1 ORDER BY version ASC",
        ) {
            Ok(mut stmt) => stmt.query_map(params![version.endpoint_id], |row| Ok(row_to_code_version(row)))
                .ok()
                .map(|iter| iter.filter_map(|r| r.ok()).collect())
                .unwrap_or_default(),
            Err(_) => Vec::new(),
        };

        if existing.len() >= MAX_VERSIONS as usize {
            if let Some(oldest) = existing.first() {
                let _ = conn.execute("DELETE FROM code_versions WHERE id = ?1", params![oldest.id]);
            }
        }

        conn.execute(
            "INSERT OR REPLACE INTO code_versions (id, endpoint_id, version, raw_code, lang, wasm, created_at)
             VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)",
            params![
                version.id,
                version.endpoint_id,
                version.version,
                version.raw,
                version.lang.as_str(),
                version.wasm,
                version.created_at as i64,
            ],
        )?;
        Ok(())
    }

    pub async fn list_code_versions(&self, endpoint_id: &str) -> Vec<CodeVersion> {
        let conn = match self.conn.lock() {
            Ok(c) => c,
            Err(_) => return Vec::new(),
        };
        let mut stmt = match conn.prepare(
            "SELECT id, endpoint_id, version, raw_code, lang, wasm, created_at FROM code_versions WHERE endpoint_id = ?1 ORDER BY version DESC",
        ) {
            Ok(s) => s,
            Err(_) => return Vec::new(),
        };
        let result = stmt.query_map(params![endpoint_id], |row| Ok(row_to_code_version(row)));
        match result {
            Ok(iter) => iter.filter_map(|r| r.ok()).collect(),
            Err(_) => Vec::new(),
        }
    }

    pub async fn next_code_version_number(&self, endpoint_id: &str) -> i32 {
        let conn = match self.conn.lock() {
            Ok(c) => c,
            Err(_) => return 1,
        };
        let result: Option<i32> = conn.query_row(
            "SELECT MAX(version) FROM code_versions WHERE endpoint_id = ?1",
            params![endpoint_id],
            |row| row.get(0),
        ).ok().flatten();
        result.unwrap_or(0) + 1
    }

    pub async fn save_test_version(&self, version: &TestVersion) -> Result<(), StoreError> {
        let conn = self.conn.lock().unwrap();
        let existing: Vec<TestVersion> = match conn.prepare(
            "SELECT id, endpoint_id, version, hurl, created_at FROM test_versions WHERE endpoint_id = ?1 ORDER BY version ASC",
        ) {
            Ok(mut stmt) => stmt.query_map(params![version.endpoint_id], |row| Ok(row_to_test_version(row)))
                .ok()
                .map(|iter| iter.filter_map(|r| r.ok()).collect())
                .unwrap_or_default(),
            Err(_) => Vec::new(),
        };

        if existing.len() >= MAX_VERSIONS as usize {
            if let Some(oldest) = existing.first() {
                let _ = conn.execute("DELETE FROM test_versions WHERE id = ?1", params![oldest.id]);
            }
        }

        conn.execute(
            "INSERT OR REPLACE INTO test_versions (id, endpoint_id, version, hurl, created_at)
             VALUES (?1, ?2, ?3, ?4, ?5)",
            params![
                version.id,
                version.endpoint_id,
                version.version,
                version.hurl,
                version.created_at as i64,
            ],
        )?;
        Ok(())
    }

    pub async fn list_test_versions(&self, endpoint_id: &str) -> Vec<TestVersion> {
        let conn = match self.conn.lock() {
            Ok(c) => c,
            Err(_) => return Vec::new(),
        };
        let mut stmt = match conn.prepare(
            "SELECT id, endpoint_id, version, hurl, created_at FROM test_versions WHERE endpoint_id = ?1 ORDER BY version DESC",
        ) {
            Ok(s) => s,
            Err(_) => return Vec::new(),
        };
        let result = stmt.query_map(params![endpoint_id], |row| Ok(row_to_test_version(row)));
        match result {
            Ok(iter) => iter.filter_map(|r| r.ok()).collect(),
            Err(_) => Vec::new(),
        }
    }

    pub async fn next_test_version_number(&self, endpoint_id: &str) -> i32 {
        let conn = match self.conn.lock() {
            Ok(c) => c,
            Err(_) => return 1,
        };
        let result: Option<i32> = conn.query_row(
            "SELECT MAX(version) FROM test_versions WHERE endpoint_id = ?1",
            params![endpoint_id],
            |row| row.get(0),
        ).ok().flatten();
        result.unwrap_or(0) + 1
    }

    pub async fn get_code_version(&self, id: &str) -> Option<CodeVersion> {
        let conn = self.conn.lock().unwrap();
        let mut stmt = conn.prepare(
            "SELECT id, endpoint_id, version, raw_code, lang, wasm, created_at FROM code_versions WHERE id = ?1",
        ).ok()?;
        let result = stmt.query_row(params![id], |row| Ok(row_to_code_version(row)));
        result.ok()
    }

    pub async fn get_test_version(&self, id: &str) -> Option<TestVersion> {
        let conn = self.conn.lock().unwrap();
        let mut stmt = conn.prepare(
            "SELECT id, endpoint_id, version, hurl, created_at FROM test_versions WHERE id = ?1",
        ).ok()?;
        let result = stmt.query_row(params![id], |row| Ok(row_to_test_version(row)));
        result.ok()
    }
}

fn row_to_code(row: &rusqlite::Row<'_>) -> Code {
    let id: String = row.get(0).unwrap_or_default();
    let path: String = row.get(1).unwrap_or_default();
    let lang_str: String = row.get(2).unwrap_or_default();
    let raw: String = row.get(3).unwrap_or_default();
    let wasm: Vec<u8> = row.get(4).unwrap_or_default();
    let hurl: String = row.get(5).unwrap_or_default();
    let user_id: String = row.get(6).unwrap_or_default();
    let created_at: i64 = row.get(7).unwrap_or_default();
    let updated_at: i64 = row.get(8).unwrap_or_default();

    let mut code = Code::new(path, raw, Lang::from_str(&lang_str), wasm);
    code.id = id;
    code.hurl = hurl;
    code.user_id = user_id;
    code.created_at = created_at as u64;
    code.updated_at = updated_at as u64;
    code
}

fn row_to_user(row: &rusqlite::Row<'_>) -> User {
    let id: String = row.get(0).unwrap_or_default();
    let email: String = row.get(1).unwrap_or_default();
    let password_hash: String = row.get(2).unwrap_or_default();
    let created_at: i64 = row.get(3).unwrap_or_default();
    let updated_at: i64 = row.get(4).unwrap_or_default();

    User {
        id,
        email,
        password_hash,
        created_at: created_at as u64,
        updated_at: updated_at as u64,
    }
}

fn row_to_ticket(row: &rusqlite::Row<'_>) -> Ticket {
    let id: String = row.get(0).unwrap_or_default();
    let code: String = row.get(1).unwrap_or_default();
    let created_by: String = row.get(2).unwrap_or_default();
    let used_by: Option<String> = row.get(3).unwrap_or(None);
    let used_at: Option<i64> = row.get(4).unwrap_or(None);
    let created_at: i64 = row.get(5).unwrap_or_default();

    Ticket {
        id,
        code,
        created_by,
        used_by,
        used_at: used_at.map(|t| t as u64),
        created_at: created_at as u64,
    }
}

fn row_to_session(row: &rusqlite::Row<'_>) -> Session {
    let token: String = row.get(0).unwrap_or_default();
    let user_id: String = row.get(1).unwrap_or_default();
    let created_at: i64 = row.get(2).unwrap_or_default();
    let expires_at: i64 = row.get(3).unwrap_or_default();

    Session {
        token,
        user_id,
        created_at: created_at as u64,
        expires_at: expires_at as u64,
    }
}

fn row_to_code_version(row: &rusqlite::Row<'_>) -> CodeVersion {
    let id: String = row.get(0).unwrap_or_default();
    let endpoint_id: String = row.get(1).unwrap_or_default();
    let version: i32 = row.get(2).unwrap_or_default();
    let raw: String = row.get(3).unwrap_or_default();
    let lang_str: String = row.get(4).unwrap_or_default();
    let wasm: Vec<u8> = row.get(5).unwrap_or_default();
    let created_at: i64 = row.get(6).unwrap_or_default();

    CodeVersion {
        id,
        endpoint_id,
        version,
        raw,
        lang: Lang::from_str(&lang_str),
        wasm,
        created_at: created_at as u64,
    }
}

fn row_to_test_version(row: &rusqlite::Row<'_>) -> TestVersion {
    let id: String = row.get(0).unwrap_or_default();
    let endpoint_id: String = row.get(1).unwrap_or_default();
    let version: i32 = row.get(2).unwrap_or_default();
    let hurl: String = row.get(3).unwrap_or_default();
    let created_at: i64 = row.get(4).unwrap_or_default();

    TestVersion {
        id,
        endpoint_id,
        version,
        hurl,
        created_at: created_at as u64,
    }
}