use turso_orm::sql::{
AlterTable, Build, CreateIndex, CreateTable, DropIndex, DropTable, Statement,
};
use turso_orm::{ConnectionTrait, DbErr, Transaction};
#[derive(Debug)]
pub struct SchemaManager<'c> {
conn: &'c Transaction,
}
impl<'c> SchemaManager<'c> {
pub fn new(conn: &'c Transaction) -> Self {
Self { conn }
}
pub fn get_connection(&self) -> &'c Transaction {
self.conn
}
pub async fn exec_stmt(&self, stmt: impl Build) -> Result<(), DbErr> {
self.conn.execute(stmt.to_statement()).await?;
Ok(())
}
pub async fn create_table(&self, stmt: CreateTable) -> Result<(), DbErr> {
self.exec_stmt(stmt).await
}
pub async fn alter_table(&self, stmt: AlterTable) -> Result<(), DbErr> {
self.exec_stmt(stmt).await
}
pub async fn drop_table(&self, stmt: DropTable) -> Result<(), DbErr> {
self.exec_stmt(stmt).await
}
pub async fn create_index(&self, stmt: CreateIndex) -> Result<(), DbErr> {
self.exec_stmt(stmt).await
}
pub async fn drop_index(&self, stmt: DropIndex) -> Result<(), DbErr> {
self.exec_stmt(stmt).await
}
pub async fn has_table(&self, table: &str) -> Result<bool, DbErr> {
has_table(self.conn, table).await
}
pub async fn has_column(&self, table: &str, column: &str) -> Result<bool, DbErr> {
count_positive(
self.conn,
Statement::from_sql_and_values(
"SELECT COUNT(*) AS n FROM pragma_table_info(?) WHERE name = ?",
[table, column],
),
)
.await
}
pub async fn has_index(&self, index: &str) -> Result<bool, DbErr> {
count_positive(
self.conn,
Statement::from_sql_and_values(
"SELECT COUNT(*) AS n FROM sqlite_schema WHERE type = 'index' AND name = ?",
[index],
),
)
.await
}
}
pub(crate) async fn has_table<C: ConnectionTrait>(conn: &C, table: &str) -> Result<bool, DbErr> {
count_positive(
conn,
Statement::from_sql_and_values(
"SELECT COUNT(*) AS n FROM sqlite_schema WHERE type = 'table' AND name = ?",
[table],
),
)
.await
}
async fn count_positive<C: ConnectionTrait>(conn: &C, stmt: Statement) -> Result<bool, DbErr> {
let row = conn
.query_one(stmt)
.await?
.ok_or_else(|| DbErr::Migration("catalog query returned no row".into()))?;
Ok(row.get::<i64>("n")? > 0)
}