use anyhow::{Context, Result};
use rusqlite::types::ValueRef;
use rusqlite::{params, Connection, OptionalExtension};
use std::path::{Path, PathBuf};
pub struct SqliteStore {
conn: Connection,
pub path: PathBuf,
pub tables: Vec<String>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct Sort {
pub column: String,
pub desc: bool,
}
impl Sort {
fn order_by(&self) -> String {
let dir = self.dir();
format!("\"{}\" {dir}, rowid {dir}", esc(&self.column))
}
fn order_by_reversed(&self) -> String {
let dir = if self.desc { "ASC" } else { "DESC" };
format!("\"{}\" {dir}, rowid {dir}", esc(&self.column))
}
fn dir(&self) -> &'static str {
if self.desc {
"DESC"
} else {
"ASC"
}
}
fn compare_to_marker(&self, table: &str, later: bool) -> String {
let op = if later != self.desc { ">" } else { "<" };
let col = esc(&self.column);
format!(
"(\"{col}\", rowid) {op} (SELECT \"{col}\", rowid FROM \"{t}\" WHERE rowid = ?2)",
t = esc(table)
)
}
}
pub struct RowsView {
pub columns: Vec<String>,
pub rows: Vec<Vec<String>>,
pub rowids: Vec<Option<i64>>,
pub total: i64,
}
impl SqliteStore {
pub fn open(path: &Path) -> Result<Self> {
let conn =
Connection::open(path).with_context(|| format!("open sqlite {}", path.display()))?;
let tables = list_tables(&conn)?;
Ok(Self {
conn,
path: path.to_path_buf(),
tables,
})
}
fn filter_clause(columns: &[String], term: &str) -> Option<(String, String)> {
if term.is_empty() || columns.is_empty() {
return None;
}
let likes = columns
.iter()
.map(|c| format!("CAST(\"{}\" AS TEXT) LIKE ?1 ESCAPE '\\'", esc(c)))
.collect::<Vec<_>>()
.join(" OR ");
Some((
format!(" WHERE ({likes})"),
format!("%{}%", like_escape(term)),
))
}
pub fn count_filtered(&self, table: &str, filter: &str) -> Result<i64> {
let columns = self.columns(table)?;
match Self::filter_clause(&columns, filter) {
None => self.count(table),
Some((where_sql, pattern)) => {
let sql = format!("SELECT COUNT(*) FROM \"{}\"{}", esc(table), where_sql);
Ok(self.conn.query_row(&sql, params![pattern], |r| r.get(0))?)
}
}
}
pub fn count(&self, table: &str) -> Result<i64> {
let n: i64 = self.conn.query_row(
&format!("SELECT COUNT(*) FROM \"{}\"", esc(table)),
[],
|r| r.get(0),
)?;
Ok(n)
}
pub fn columns(&self, table: &str) -> Result<Vec<String>> {
let mut stmt = self
.conn
.prepare(&format!("PRAGMA table_info(\"{}\")", esc(table)))?;
let cols = stmt
.query_map([], |r| r.get::<_, String>(1))?
.collect::<Result<Vec<_>, _>>()?;
Ok(cols)
}
pub fn rows(
&self,
table: &str,
limit: i64,
offset: i64,
sort: Option<&Sort>,
filter: &str,
) -> Result<RowsView> {
let columns = self.columns(table)?;
let total = self.count_filtered(table, filter)?;
let ncols = columns.len();
let clause = Self::filter_clause(&columns, filter);
let where_sql = clause.as_ref().map(|(w, _)| w.as_str()).unwrap_or("");
let with_rowid = self
.conn
.prepare(&format!(
"SELECT rowid, * FROM \"{}\" LIMIT {} OFFSET {}",
esc(table),
limit,
offset
))
.is_ok();
let sql = if with_rowid {
let order = match sort {
Some(s) if columns.contains(&s.column) => s.order_by(),
_ => "rowid".to_string(),
};
format!(
"SELECT rowid, * FROM \"{}\"{} ORDER BY {} LIMIT {} OFFSET {}",
esc(table),
where_sql,
order,
limit,
offset
)
} else {
let order = match sort {
Some(s) if columns.contains(&s.column) => {
format!("\"{}\" {}", esc(&s.column), s.dir())
}
_ => "1".to_string(),
};
format!(
"SELECT * FROM \"{}\"{} ORDER BY {} LIMIT {} OFFSET {}",
esc(table),
where_sql,
order,
limit,
offset
)
};
let mut stmt = self.conn.prepare(&sql)?;
let mut rows_out = Vec::new();
let mut rowids = Vec::new();
let mut q = match &clause {
Some((_, pattern)) => stmt.query(params![pattern])?,
None => stmt.query([])?,
};
while let Some(row) = q.next()? {
let (base, rid) = if with_rowid {
(1usize, row.get::<_, i64>(0).ok())
} else {
(0usize, None)
};
rowids.push(rid);
let mut cells = Vec::with_capacity(ncols);
for i in 0..ncols {
cells.push(value_to_string(row, base + i));
}
rows_out.push(cells);
}
Ok(RowsView {
columns,
rows: rows_out,
rowids,
total,
})
}
pub fn find_row(
&self,
table: &str,
columns: &[String],
term: &str,
from_rowid: i64,
forward: bool,
sort: Option<&Sort>,
) -> Result<Option<i64>> {
if columns.is_empty() {
return Ok(None);
}
let likes = columns
.iter()
.map(|c| format!("CAST(\"{}\" AS TEXT) LIKE ?1 ESCAPE '\\'", esc(c)))
.collect::<Vec<_>>()
.join(" OR ");
let sql = match sort {
Some(s) if columns.contains(&s.column) => format!(
"SELECT rowid FROM \"{}\" WHERE {} AND ({}) ORDER BY {} LIMIT 1",
esc(table),
s.compare_to_marker(table, forward),
likes,
if forward {
s.order_by()
} else {
s.order_by_reversed()
}
),
_ => {
let (cmp, ord) = if forward { (">", "ASC") } else { ("<", "DESC") };
format!(
"SELECT rowid FROM \"{}\" WHERE rowid {} ?2 AND ({}) ORDER BY rowid {} LIMIT 1",
esc(table),
cmp,
likes,
ord
)
}
};
let pattern = format!("%{}%", like_escape(term));
let mut stmt = self.conn.prepare(&sql)?;
let rid = stmt
.query_row(params![pattern, from_rowid], |r| r.get::<_, i64>(0))
.optional()?;
Ok(rid)
}
pub fn find_row_edge(
&self,
table: &str,
columns: &[String],
term: &str,
forward: bool,
sort: Option<&Sort>,
) -> Result<Option<i64>> {
if columns.is_empty() {
return Ok(None);
}
let likes = columns
.iter()
.map(|c| format!("CAST(\"{}\" AS TEXT) LIKE ?1 ESCAPE '\\'", esc(c)))
.collect::<Vec<_>>()
.join(" OR ");
let order = match sort {
Some(s) if columns.contains(&s.column) => {
if forward {
s.order_by()
} else {
s.order_by_reversed()
}
}
_ => format!("rowid {}", if forward { "ASC" } else { "DESC" }),
};
let sql = format!(
"SELECT rowid FROM \"{}\" WHERE ({}) ORDER BY {} LIMIT 1",
esc(table),
likes,
order
);
let pattern = format!("%{}%", like_escape(term));
let mut stmt = self.conn.prepare(&sql)?;
Ok(stmt
.query_row(params![pattern], |r| r.get::<_, i64>(0))
.optional()?)
}
pub fn rowid_ordinal(&self, table: &str, rowid: i64, sort: Option<&Sort>) -> Result<i64> {
let sql = match sort {
Some(s) => format!(
"SELECT COUNT(*) FROM \"{}\" WHERE NOT ({})",
esc(table),
s.compare_to_marker(table, true)
),
None => format!("SELECT COUNT(*) FROM \"{}\" WHERE rowid <= ?2", esc(table)),
};
let n: i64 = self
.conn
.query_row(&sql, params![rowid, rowid], |r| r.get(0))?;
Ok(n)
}
pub fn cell_bytes(&self, table: &str, rowid: i64, col: &str) -> Result<Vec<u8>> {
let v: Vec<u8> = self.conn.query_row(
&format!(
"SELECT \"{}\" FROM \"{}\" WHERE rowid = ?1",
esc(col),
esc(table)
),
params![rowid],
|r| {
use rusqlite::types::ValueRef;
Ok(match r.get_ref(0)? {
ValueRef::Blob(b) => b.to_vec(),
ValueRef::Text(t) => t.to_vec(),
ValueRef::Integer(i) => i.to_string().into_bytes(),
ValueRef::Real(f) => f.to_string().into_bytes(),
ValueRef::Null => Vec::new(),
})
},
)?;
Ok(v)
}
pub fn schema(&self) -> Result<Vec<(String, String, String)>> {
let mut stmt = self.conn.prepare(
"SELECT type, name, COALESCE(sql, '') FROM sqlite_master \
WHERE name NOT LIKE 'sqlite_%' ORDER BY type, name",
)?;
let out = stmt
.query_map([], |r| {
Ok((
r.get::<_, String>(0)?,
r.get::<_, String>(1)?,
r.get::<_, String>(2)?,
))
})?
.collect::<Result<Vec<_>, _>>()?;
Ok(out)
}
pub fn update_cell(&self, table: &str, rowid: i64, col: &str, val: &str) -> Result<()> {
self.conn.execute(
&format!(
"UPDATE \"{}\" SET \"{}\" = ?1 WHERE rowid = ?2",
esc(table),
esc(col)
),
params![val, rowid],
)?;
Ok(())
}
pub fn delete_row(&self, table: &str, rowid: i64) -> Result<()> {
self.conn.execute(
&format!("DELETE FROM \"{}\" WHERE rowid = ?1", esc(table)),
params![rowid],
)?;
Ok(())
}
pub fn insert_blank(&self, table: &str) -> Result<()> {
self.conn.execute(
&format!("INSERT INTO \"{}\" DEFAULT VALUES", esc(table)),
[],
)?;
Ok(())
}
pub fn exec(&self, sql: &str) -> Result<usize> {
Ok(self.conn.execute(sql, [])?)
}
}
fn value_to_string(row: &rusqlite::Row, idx: usize) -> String {
match row.get_ref(idx) {
Ok(ValueRef::Null) => "NULL".into(),
Ok(ValueRef::Integer(i)) => i.to_string(),
Ok(ValueRef::Real(f)) => f.to_string(),
Ok(ValueRef::Text(t)) => String::from_utf8_lossy(t).into_owned(),
Ok(ValueRef::Blob(b)) => format!("<blob {} bytes>", b.len()),
Err(_) => "?".into(),
}
}
fn list_tables(conn: &Connection) -> Result<Vec<String>> {
let mut stmt = conn.prepare(
"SELECT name FROM sqlite_master \
WHERE type IN ('table','view') AND name NOT LIKE 'sqlite_%' \
ORDER BY name",
)?;
let names = stmt
.query_map([], |r| r.get::<_, String>(0))?
.collect::<Result<Vec<_>, _>>()?;
Ok(names)
}
fn esc(ident: &str) -> String {
ident.replace('"', "\"\"")
}
fn like_escape(term: &str) -> String {
let mut out = String::with_capacity(term.len());
for c in term.chars() {
if matches!(c, '\\' | '%' | '_') {
out.push('\\');
}
out.push(c);
}
out
}