use std::collections::HashMap;
use std::path::Path;
use std::sync::Arc;
use std::time::Duration;
use serde::{Deserialize, Serialize};
use tokio::sync::OnceCell;
use turso::Builder;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct Notification {
pub id: i64,
pub kind: String,
pub message: String,
pub sent: bool,
pub created_at: String,
}
impl Notification {
pub fn new(
id: i64,
kind: impl Into<String>,
message: impl Into<String>,
sent: bool,
created_at: impl Into<String>,
) -> Self {
Self {
id,
kind: kind.into(),
message: message.into(),
sent,
created_at: created_at.into(),
}
}
}
#[derive(Clone)]
pub struct Database {
db: Arc<OnceCell<turso::Database>>,
db_path: String,
}
impl Database {
pub fn new(db_path: &str) -> Self {
Self {
db: Arc::new(OnceCell::new()),
db_path: db_path.to_owned(),
}
}
#[cfg(test)]
pub fn db_path(&self) -> &str {
&self.db_path
}
pub async fn conn(&self) -> Result<turso::Connection, String> {
let db = self
.db
.get_or_try_init(|| async {
let db_path = Path::new(&self.db_path);
let parent = db_path.parent().unwrap_or(Path::new("."));
if !parent.exists() {
log::info!("Creating parent directory: {}", parent.display());
std::fs::create_dir_all(parent).map_err(|e| e.to_string())?;
}
let db_path_str = db_path.to_str().ok_or_else(|| {
"Invalid database path: contains non-UTF8 characters".to_string()
})?;
let db = Builder::new_local(db_path_str).build().await.map_err(|e| {
log::error!("Failed to create database: {e}");
e.to_string()
})?;
log::info!("Created db");
Ok::<turso::Database, String>(db)
})
.await?;
let conn = db
.connect()
.map_err(|e| format!("Failed to connect to database: {e}"))?;
conn.busy_timeout(Duration::from_secs(30))
.map_err(|e| format!("Failed to set busy timeout: {e}"))?;
Ok(conn)
}
pub async fn init(&self) -> Result<(), String> {
let conn = self.conn().await?;
conn.busy_timeout(Duration::from_secs(30))
.map_err(|e| format!("Failed to set busy timeout: {e}"))?;
conn.query("PRAGMA journal_mode = WAL;", ())
.await
.map_err(|e| format!("Failed to set WAL mode: {e}"))?;
conn.execute(
"CREATE TABLE IF NOT EXISTS notifications (
id INTEGER PRIMARY KEY AUTOINCREMENT,
kind TEXT NOT NULL,
message TEXT NOT NULL,
sent INTEGER NOT NULL DEFAULT 0,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);",
(),
)
.await
.map_err(|e| format!("Failed to create notifications table: {e}"))?;
conn.execute(
"CREATE TABLE IF NOT EXISTS sessions (
token TEXT PRIMARY KEY,
expires_at INTEGER NOT NULL
);",
(),
)
.await
.map_err(|e| format!("Failed to create sessions table: {e}"))?;
log::info!("Database initialized successfully");
Ok(())
}
/// Inserts or replaces a session token with its expiry (Unix seconds).
pub async fn save_session(&self, token: &str, expires_at: u64) -> Result<(), String> {
let conn = self.conn().await?;
conn.execute(
"INSERT OR REPLACE INTO sessions (token, expires_at) VALUES (?, ?)",
(token, expires_at as i64),
)
.await
.map_err(|e| format!("Failed to save session: {e}"))?;
Ok(())
}
/// Loads every stored session as a token -> expiry map.
pub async fn load_sessions(&self) -> Result<HashMap<String, u64>, String> {
let conn = self.conn().await?;
let mut rows = conn
.query("SELECT token, expires_at FROM sessions", ())
.await
.map_err(|e| format!("Failed to query sessions: {e}"))?;
let mut sessions = HashMap::new();
while let Some(row) = rows
.next()
.await
.map_err(|e| format!("Failed to read session row: {e}"))?
{
let token: String = row
.get(0)
.map_err(|e| format!("Failed to get session token: {e}"))?;
let expires_at: i64 = row
.get(1)
.map_err(|e| format!("Failed to get session expiry: {e}"))?;
sessions.insert(token, expires_at.max(0) as u64);
}
Ok(sessions)
}
/// Deletes all sessions that expired at or before `now` (Unix seconds).
pub async fn delete_expired_sessions(&self, now: u64) -> Result<(), String> {
let conn = self.conn().await?;
conn.execute("DELETE FROM sessions WHERE expires_at <= ?", (now as i64,))
.await
.map_err(|e| format!("Failed to delete expired sessions: {e}"))?;
Ok(())
}
pub async fn save_notification(
&self,
kind: &str,
message: &str,
sent: bool,
) -> Result<Notification, String> {
let conn = self.conn().await?;
let sent_val: i64 = if sent { 1 } else { 0 };
conn.execute(
"INSERT INTO notifications (kind, message, sent) VALUES (?, ?, ?)",
(kind, message, sent_val),
)
.await
.map_err(|e| format!("Failed to insert notification: {e}"))?;
let id = conn.last_insert_rowid();
let mut rows = conn
.query(
"SELECT id, kind, message, sent, created_at FROM notifications WHERE id = ?",
(id,),
)
.await
.map_err(|e| format!("Failed to query notification by id: {e}"))?;
if let Some(row) = rows
.next()
.await
.map_err(|e| format!("Failed to read notification row: {e}"))?
{
let id: i64 = row
.get(0)
.map_err(|e| format!("Failed to get notification id: {e}"))?;
let kind: String = row
.get(1)
.map_err(|e| format!("Failed to get notification kind: {e}"))?;
let message: String = row
.get(2)
.map_err(|e| format!("Failed to get notification message: {e}"))?;
let sent_val: i64 = row
.get(3)
.map_err(|e| format!("Failed to get notification sent: {e}"))?;
let created_at: String = row
.get(4)
.map_err(|e| format!("Failed to get notification created_at: {e}"))?;
Ok(Notification::new(
id,
kind,
message,
sent_val != 0,
created_at,
))
} else {
Err(format!("Notification with id {id} not found after insert"))
}
}
#[cfg(test)]
pub async fn set_notification_sent(&self, id: i64, sent: bool) -> Result<(), String> {
let conn = self.conn().await?;
let sent_val: i64 = if sent { 1 } else { 0 };
conn.execute(
"UPDATE notifications SET sent = ? WHERE id = ?",
(sent_val, id),
)
.await
.map_err(|e| format!("Failed to update notification sent status: {e}"))?;
Ok(())
}
#[cfg(test)]
pub async fn get_notification(&self, id: i64) -> Result<Option<Notification>, String> {
let conn = self.conn().await?;
let mut rows = conn
.query(
"SELECT id, kind, message, sent, created_at FROM notifications WHERE id = ?",
(id,),
)
.await
.map_err(|e| format!("Failed to query notification: {e}"))?;
if let Some(row) = rows
.next()
.await
.map_err(|e| format!("Failed to read notification row: {e}"))?
{
let id: i64 = row
.get(0)
.map_err(|e| format!("Failed to get notification id: {e}"))?;
let kind: String = row
.get(1)
.map_err(|e| format!("Failed to get notification kind: {e}"))?;
let message: String = row
.get(2)
.map_err(|e| format!("Failed to get notification message: {e}"))?;
let sent_val: i64 = row
.get(3)
.map_err(|e| format!("Failed to get notification sent: {e}"))?;
let created_at: String = row
.get(4)
.map_err(|e| format!("Failed to get notification created_at: {e}"))?;
Ok(Some(Notification::new(
id,
kind,
message,
sent_val != 0,
created_at,
)))
} else {
Ok(None)
}
}
pub async fn list_notifications(&self) -> Result<Vec<Notification>, String> {
let conn = self.conn().await?;
let mut rows = conn
.query(
"SELECT id, kind, message, sent, created_at FROM notifications ORDER BY id DESC",
(),
)
.await
.map_err(|e| format!("Failed to query notifications: {e}"))?;
let mut list = Vec::new();
while let Some(row) = rows
.next()
.await
.map_err(|e| format!("Failed to read notification row: {e}"))?
{
let id: i64 = row
.get(0)
.map_err(|e| format!("Failed to get notification id: {e}"))?;
let kind: String = row
.get(1)
.map_err(|e| format!("Failed to get notification kind: {e}"))?;
let message: String = row
.get(2)
.map_err(|e| format!("Failed to get notification message: {e}"))?;
let sent_val: i64 = row
.get(3)
.map_err(|e| format!("Failed to get notification sent: {e}"))?;
let created_at: String = row
.get(4)
.map_err(|e| format!("Failed to get notification created_at: {e}"))?;
list.push(Notification::new(
id,
kind,
message,
sent_val != 0,
created_at,
));
}
Ok(list)
}
pub async fn get_pending_notifications(
&self,
limit: Option<usize>,
) -> Result<Vec<Notification>, String> {
let conn = self.conn().await?;
let query_sql = match limit {
Some(l) => format!(
"SELECT id, kind, message, sent, created_at FROM notifications WHERE sent = 0 ORDER BY id ASC LIMIT {l}"
),
None => {
"SELECT id, kind, message, sent, created_at FROM notifications WHERE sent = 0 ORDER BY id ASC".to_string()
}
};
let mut rows = conn
.query(&query_sql, ())
.await
.map_err(|e| format!("Failed to query unsent notifications: {e}"))?;
let mut list = Vec::new();
while let Some(row) = rows
.next()
.await
.map_err(|e| format!("Failed to read notification row: {e}"))?
{
let id: i64 = row
.get(0)
.map_err(|e| format!("Failed to get notification id: {e}"))?;
let kind: String = row
.get(1)
.map_err(|e| format!("Failed to get notification kind: {e}"))?;
let message: String = row
.get(2)
.map_err(|e| format!("Failed to get notification message: {e}"))?;
let sent_val: i64 = row
.get(3)
.map_err(|e| format!("Failed to get notification sent: {e}"))?;
let created_at: String = row
.get(4)
.map_err(|e| format!("Failed to get notification created_at: {e}"))?;
list.push(Notification::new(
id,
kind,
message,
sent_val != 0,
created_at,
));
}
Ok(list)
}
pub async fn mark_as_sent(&self, ids: &[i64]) -> Result<Vec<i64>, String> {
let conn = self.conn().await?;
if ids.is_empty() {
let mut rows = conn
.query("SELECT id FROM notifications WHERE sent = 0", ())
.await
.map_err(|e| format!("Failed to query unsent notification IDs: {e}"))?;
let mut all_ids = Vec::new();
while let Some(row) = rows
.next()
.await
.map_err(|e| format!("Failed to read notification id: {e}"))?
{
let id: i64 = row
.get(0)
.map_err(|e| format!("Failed to get notification id: {e}"))?;
all_ids.push(id);
}
if all_ids.is_empty() {
return Ok(Vec::new());
}
conn.execute("UPDATE notifications SET sent = 1 WHERE sent = 0", ())
.await
.map_err(|e| format!("Failed to mark all notifications sent: {e}"))?;
Ok(all_ids)
} else {
let mut updated_ids = Vec::new();
for &id in ids {
let affected = conn
.execute("UPDATE notifications SET sent = 1 WHERE id = ?", (id,))
.await
.map_err(|e| format!("Failed to mark notification {id} as sent: {e}"))?;
if affected > 0 {
updated_ids.push(id);
}
}
Ok(updated_ids)
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_database_new_stores_path() {
let db = Database::new("/tmp/test_notify_new.db");
assert_eq!(db.db_path(), "/tmp/test_notify_new.db");
}
#[tokio::test]
async fn test_database_conn_creates_file() {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let db_path = format!("/tmp/test_notify_conn_{now}.db");
let db = Database::new(&db_path);
let conn = db.conn().await;
assert!(conn.is_ok(), "conn should succeed: {conn:?}");
assert!(
std::path::Path::new(&db_path).exists(),
"db file should be created"
);
let _ = std::fs::remove_file(&db_path);
}
#[tokio::test]
async fn test_database_init_succeeds() {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let db_path = format!("/tmp/test_notify_init_{now}.db");
let db = Database::new(&db_path);
let result = db.init().await;
assert!(result.is_ok(), "init should succeed: {result:?}");
let _ = std::fs::remove_file(&db_path);
}
#[tokio::test]
async fn test_save_and_get_notification() {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let db_path = format!("/tmp/test_notify_crud_{now}.db");
let db = Database::new(&db_path);
db.init().await.unwrap();
// Save notification with sent = false
let saved = db
.save_notification("heart", "Look at your phone", false)
.await
.expect("save_notification should succeed");
assert_eq!(saved.kind, "heart");
assert_eq!(saved.message, "Look at your phone");
assert!(!saved.sent);
assert!(!saved.created_at.is_empty());
// Get by ID
let fetched = db
.get_notification(saved.id)
.await
.expect("get_notification should succeed");
assert_eq!(fetched, Some(saved.clone()));
// Update sent to true
db.set_notification_sent(saved.id, true)
.await
.expect("set_notification_sent should succeed");
let updated = db
.get_notification(saved.id)
.await
.expect("get_notification should succeed")
.expect("notification should exist");
assert!(updated.sent);
// List notifications
let list = db
.list_notifications()
.await
.expect("list_notifications should succeed");
assert_eq!(list.len(), 1);
assert_eq!(list[0].id, saved.id);
assert!(list[0].sent);
let _ = std::fs::remove_file(&db_path);
}
#[tokio::test]
async fn test_scan_and_mark_sent() {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos();
let db_path = format!("/tmp/test_notify_scan_{now}.db");
let db = Database::new(&db_path);
db.init().await.unwrap();
// Initially no pending notifications
let pending = db.get_pending_notifications(None).await.unwrap();
assert!(pending.is_empty());
// Insert 2 notifications with sent = false
let n1 = db
.save_notification("heart", "Look at your phone", false)
.await
.unwrap();
let n2 = db
.save_notification("cat", "Missing you", false)
.await
.unwrap();
assert!(!n1.sent);
assert!(!n2.sent);
// get_pending_notifications scans them but does NOT mark them as sent
let pending = db.get_pending_notifications(None).await.unwrap();
assert_eq!(pending.len(), 2);
assert_eq!(pending[0].id, n1.id);
assert_eq!(pending[0].kind, "heart");
assert!(!pending[0].sent);
assert_eq!(pending[1].id, n2.id);
assert_eq!(pending[1].kind, "cat");
assert!(!pending[1].sent);
// Repeated scan still returns both because they were NOT marked sent
let pending2 = db.get_pending_notifications(None).await.unwrap();
assert_eq!(pending2.len(), 2);
// Now explicitly mark them as sent
let updated = db.mark_as_sent(&[n1.id]).await.unwrap();
assert_eq!(updated, vec![n1.id]);
let remaining = db.get_pending_notifications(None).await.unwrap();
assert_eq!(remaining.len(), 1);
assert_eq!(remaining[0].id, n2.id);
// Mark remaining sent
let updated_all = db.mark_as_sent(&[]).await.unwrap();
assert_eq!(updated_all, vec![n2.id]);
let empty = db.get_pending_notifications(None).await.unwrap();
assert!(empty.is_empty());
let _ = std::fs::remove_file(&db_path);
}
}