use rusqlite::Connection;
use serde::Serialize;
use crate::core::errors::{Result, TgaError};
use super::text_columns::{classify, TextClass};
#[derive(Debug, Clone, Serialize)]
#[non_exhaustive]
pub struct ColumnInfo {
pub name: String,
pub declared_type: String,
pub not_null: bool,
pub pk_position: i64,
pub text_class: Option<TextClass>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize)]
#[serde(rename_all = "lowercase")]
pub enum ObjectKind {
Table,
View,
}
#[derive(Debug, Clone, Serialize)]
#[non_exhaustive]
pub struct TableInfo {
pub name: String,
pub kind: ObjectKind,
pub columns: Vec<ColumnInfo>,
pub row_count: Option<i64>,
}
#[derive(Debug, Clone, Serialize)]
#[non_exhaustive]
pub struct SchemaSnapshot {
pub schema_version: Option<i64>,
pub objects: Vec<TableInfo>,
}
impl SchemaSnapshot {
pub fn text_columns(&self) -> Vec<(&TableInfo, &ColumnInfo)> {
self.objects
.iter()
.filter(|o| o.kind == ObjectKind::Table)
.flat_map(|t| {
t.columns
.iter()
.filter(|c| c.declared_type.eq_ignore_ascii_case("TEXT"))
.map(move |c| (t, c))
})
.collect()
}
}
pub fn snapshot(conn: &Connection) -> Result<SchemaSnapshot> {
let mut objects = Vec::new();
for (kind, name) in object_names(conn)? {
let columns = columns_of(conn, &name)?;
let row_count = match kind {
ObjectKind::Table => Some(count_rows(conn, &name)?),
ObjectKind::View => None,
};
objects.push(TableInfo {
name,
kind,
columns,
row_count,
});
}
Ok(SchemaSnapshot {
schema_version: schema_version(conn)?,
objects,
})
}
fn object_names(conn: &Connection) -> Result<Vec<(ObjectKind, String)>> {
let mut stmt = conn
.prepare(
"SELECT type, name FROM sqlite_master \
WHERE type IN ('table', 'view') AND name NOT LIKE 'sqlite_%' \
ORDER BY type = 'view', name",
)
.map_err(TgaError::from)?;
let rows = stmt
.query_map([], |row| {
let kind: String = row.get(0)?;
let name: String = row.get(1)?;
Ok((kind, name))
})
.map_err(TgaError::from)?;
let mut out = Vec::new();
for row in rows {
let (kind, name) = row.map_err(TgaError::from)?;
let kind = if kind == "view" {
ObjectKind::View
} else {
ObjectKind::Table
};
out.push((kind, name));
}
Ok(out)
}
fn columns_of(conn: &Connection, table: &str) -> Result<Vec<ColumnInfo>> {
let quoted = table.replace('"', "\"\"");
let mut stmt = conn
.prepare(&format!("PRAGMA table_info(\"{quoted}\")"))
.map_err(TgaError::from)?;
let rows = stmt
.query_map([], |row| {
Ok(ColumnInfo {
name: row.get::<_, String>(1)?,
declared_type: row.get::<_, String>(2)?,
not_null: row.get::<_, i64>(3)? != 0,
pk_position: row.get::<_, i64>(5)?,
text_class: None,
})
})
.map_err(TgaError::from)?;
let mut out = Vec::new();
for row in rows {
let mut col = row.map_err(TgaError::from)?;
if col.declared_type.eq_ignore_ascii_case("TEXT") {
col.text_class = Some(classify(table, &col.name));
}
out.push(col);
}
Ok(out)
}
fn count_rows(conn: &Connection, table: &str) -> Result<i64> {
let quoted = table.replace('"', "\"\"");
conn.query_row(&format!("SELECT COUNT(*) FROM \"{quoted}\""), [], |row| {
row.get(0)
})
.map_err(TgaError::from)
}
fn schema_version(conn: &Connection) -> Result<Option<i64>> {
let present: i64 = conn
.query_row(
"SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = 'schema_migrations'",
[],
|row| row.get(0),
)
.map_err(TgaError::from)?;
if present == 0 {
return Ok(None);
}
let version: Option<i64> = conn
.query_row("SELECT MAX(version) FROM schema_migrations", [], |row| {
row.get(0)
})
.map_err(TgaError::from)?;
Ok(version)
}