use std::path::Path;
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::OnceCell;
use turso::Builder;
use super::{ColumnInfo, TablePreview};
#[derive(Clone)]
pub struct SqliteDb {
db: Arc<OnceCell<turso::Database>>,
db_path: String,
}
impl SqliteDb {
#[must_use]
pub fn new(db_path: &str) -> Self {
Self {
db: Arc::new(OnceCell::new()),
db_path: db_path.to_owned(),
}
}
#[must_use]
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);
if let Some(parent) = db_path
.parent()
.filter(|p| !p.as_os_str().is_empty() && !p.exists())
{
let _ = std::fs::create_dir_all(parent);
}
let db = Builder::new_local(&self.db_path)
.build()
.await
.map_err(|e| {
format!("Failed to create local SQLite db at {}: {e}", self.db_path)
})?;
Ok::<turso::Database, String>(db)
})
.await?;
let conn = db
.connect()
.map_err(|e| format!("Failed to connect to SQLite: {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.query("PRAGMA journal_mode = WAL;", ())
.await
.map_err(|e| format!("Failed to set WAL mode: {e}"))?;
conn.execute(
"CREATE TABLE IF NOT EXISTS tournaments (
id INTEGER PRIMARY KEY AUTOINCREMENT,
name TEXT NOT NULL,
status TEXT NOT NULL DEFAULT 'draft',
format TEXT NOT NULL DEFAULT 'elimination',
turn_time_seconds INTEGER NOT NULL DEFAULT 20,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);",
(),
)
.await
.map_err(|e| format!("Failed to create tournaments table: {e}"))?;
conn.execute(
"CREATE TABLE IF NOT EXISTS participants (
id INTEGER PRIMARY KEY AUTOINCREMENT,
tournament_id INTEGER NOT NULL,
player_name TEXT NOT NULL,
pokemon_id INTEGER NOT NULL,
pin TEXT NOT NULL,
is_bot INTEGER NOT NULL DEFAULT 0,
status TEXT NOT NULL DEFAULT 'registered',
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);",
(),
)
.await
.map_err(|e| format!("Failed to create participants table: {e}"))?;
conn.execute(
"CREATE TABLE IF NOT EXISTS battles (
id INTEGER PRIMARY KEY AUTOINCREMENT,
tournament_id INTEGER NOT NULL,
round INTEGER NOT NULL DEFAULT 1,
status TEXT NOT NULL DEFAULT 'pending',
player1_name TEXT NOT NULL,
player2_name TEXT NOT NULL,
winner_name TEXT,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);",
(),
)
.await
.map_err(|e| format!("Failed to create battles table: {e}"))?;
conn.execute(
"CREATE TABLE IF NOT EXISTS battle_moves (
id INTEGER PRIMARY KEY AUTOINCREMENT,
battle_id INTEGER NOT NULL,
turn_number INTEGER NOT NULL,
actor TEXT NOT NULL,
move_name TEXT NOT NULL,
damage INTEGER NOT NULL DEFAULT 0,
message TEXT NOT NULL,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);",
(),
)
.await
.map_err(|e| format!("Failed to create battle_moves table: {e}"))?;
log::info!("SQLite schema initialized successfully");
Ok(())
}
pub async fn list_tables(&self) -> Result<Vec<String>, String> {
let conn = self.conn().await?;
let mut rows = conn
.query(
"SELECT name FROM sqlite_master
WHERE type='table'
AND name NOT LIKE 'sqlite_%'
AND name NOT LIKE '_litestream_%'
ORDER BY name;",
(),
)
.await
.map_err(|e| format!("Failed to list sqlite tables: {e}"))?;
let mut tables = Vec::new();
while let Some(row) = rows
.next()
.await
.map_err(|e| format!("Failed to read table row: {e}"))?
{
let name: String = row
.get(0)
.map_err(|e| format!("Failed to get table name: {e}"))?;
tables.push(name);
}
Ok(tables)
}
pub async fn describe_table(&self, table: &str) -> Result<Vec<ColumnInfo>, String> {
if !table.chars().all(|c| c.is_alphanumeric() || c == '_') {
return Err("Invalid table name".to_string());
}
let conn = self.conn().await?;
let pragma_sql = format!("PRAGMA table_info({table});");
let mut rows = conn
.query(&pragma_sql, ())
.await
.map_err(|e| format!("Failed to get schema for {table}: {e}"))?;
let mut columns = Vec::new();
while let Some(row) = rows
.next()
.await
.map_err(|e| format!("Failed to read column row: {e}"))?
{
// pragma table_info columns: cid(0), name(1), type(2), notnull(3), dflt_value(4), pk(5)
let name: String = row.get(1).unwrap_or_default();
let data_type: String = row.get(2).unwrap_or_default();
let notnull: i64 = row.get(3).unwrap_or(0);
let pk: i64 = row.get(5).unwrap_or(0);
columns.push(ColumnInfo {
name,
data_type,
is_nullable: notnull == 0,
is_pk: pk > 0,
});
}
Ok(columns)
}
pub async fn table_row_count(&self, table: &str) -> Result<i64, String> {
if !table.chars().all(|c| c.is_alphanumeric() || c == '_') {
return Err("Invalid table name".to_string());
}
let conn = self.conn().await?;
let count_sql = format!("SELECT count(*) FROM {table};");
let mut rows = conn
.query(&count_sql, ())
.await
.map_err(|e| format!("Failed to get count for {table}: {e}"))?;
if let Some(row) = rows
.next()
.await
.map_err(|e| format!("Failed to read count row: {e}"))?
{
let count: i64 = row.get(0).unwrap_or(0);
Ok(count)
} else {
Ok(0)
}
}
pub async fn preview_table(&self, table: &str, limit: usize) -> Result<TablePreview, String> {
if !table.chars().all(|c| c.is_alphanumeric() || c == '_') {
return Err("Invalid table name".to_string());
}
let total_rows = self.table_row_count(table).await?;
let columns_info = self.describe_table(table).await?;
let col_names: Vec<String> = columns_info.into_iter().map(|c| c.name).collect();
let conn = self.conn().await?;
let preview_sql = format!("SELECT * FROM {table} LIMIT {limit};");
let mut rows = conn
.query(&preview_sql, ())
.await
.map_err(|e| format!("Failed to query preview for {table}: {e}"))?;
let mut data_rows = Vec::new();
while let Some(row) = rows
.next()
.await
.map_err(|e| format!("Failed to read preview row: {e}"))?
{
let mut row_vals = Vec::new();
for (idx, _col) in col_names.iter().enumerate() {
// Try reading as string, int, or float
let val_str = if let Ok(s) = row.get::<String>(idx) {
s
} else if let Ok(i) = row.get::<i64>(idx) {
i.to_string()
} else if let Ok(f) = row.get::<f64>(idx) {
format!("{f:.2}")
} else {
"NULL".to_string()
};
row_vals.push(val_str);
}
data_rows.push(row_vals);
}
Ok(TablePreview {
table_name: table.to_string(),
total_rows,
columns: col_names,
rows: data_rows,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_sqlite_db_init_and_inspection() {
let temp_path = format!(
"/tmp/test_sqlite_{}.db",
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
);
let db = SqliteDb::new(&temp_path);
db.init().await.expect("init sqlite");
let tables = db.list_tables().await.expect("list tables");
assert!(tables.contains(&"tournaments".to_string()));
assert!(tables.contains(&"battles".to_string()));
assert!(tables.contains(&"participants".to_string()));
assert!(tables.contains(&"battle_moves".to_string()));
let cols = db.describe_table("tournaments").await.expect("describe");
assert!(cols.iter().any(|c| c.name == "turn_time_seconds"));
let count = db.table_row_count("tournaments").await.expect("count");
assert_eq!(count, 0);
let preview = db.preview_table("tournaments", 10).await.expect("preview");
assert_eq!(preview.total_rows, 0);
assert!(!preview.columns.is_empty());
let _ = std::fs::remove_file(&temp_path);
}
}