use crate as rltbl;
use rltbl::{
core::{self, RelatableError, NEW_ORDER_MULTIPLIER},
table::{Column, Table},
};
use anyhow::Result;
use async_std::task::block_on;
use indexmap::IndexMap;
use lazy_static::lazy_static;
use regex::Regex;
use serde::{Deserialize, Serialize};
use serde_json::{json, Map as JsonMap, Value as JsonValue};
use std::{fmt::Display, str::FromStr};
#[cfg(feature = "rusqlite")]
use rusqlite;
#[cfg(feature = "sqlx")]
use bigdecimal::{BigDecimal, ToPrimitive};
#[cfg(feature = "sqlx")]
use sqlx::{
any::{install_default_drivers, Any, AnyArguments, AnyRow},
postgres::{PgArguments, PgConnectOptions, PgPool, PgPoolOptions, PgRow, Postgres},
query::Query,
Acquire as _, AnyPool, Column as _, Row as _, Transaction, TypeInfo as _,
};
pub static DB_OBJECT_MATCH_STR: &str = r"^[\w_]+$";
lazy_static! {
pub static ref DB_OBJECT_REGEX: Regex = Regex::new(DB_OBJECT_MATCH_STR).unwrap();
}
pub static MAX_DB_CONNECTIONS: u32 = 5;
pub static MAX_PARAMS_SQLITE: usize = 32766;
pub static MAX_PARAMS_POSTGRES: usize = 65535;
pub static DEFAULT_MEMORY_CACHE_SIZE: usize = 1000;
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum CachingStrategy {
None,
TruncateAll,
Truncate,
Trigger,
Memory(usize),
}
#[derive(Clone, Debug, Eq, PartialEq, Hash)]
pub struct MemoryCacheKey {
pub tables: String,
pub statement: String,
pub parameters: String,
}
impl FromStr for CachingStrategy {
type Err = anyhow::Error;
fn from_str(strategy: &str) -> Result<Self> {
tracing::trace!("CachingStrategy::from_str({strategy:?})");
match strategy.to_lowercase().as_str() {
"none" => Ok(CachingStrategy::None),
"truncate_all" => Ok(CachingStrategy::TruncateAll),
"truncate" => Ok(CachingStrategy::Truncate),
"trigger" => Ok(CachingStrategy::Trigger),
strategy if strategy.starts_with("memory") => {
let elems = strategy.split(":").collect::<Vec<_>>();
let cache_size = {
if elems.len() < 2 {
DEFAULT_MEMORY_CACHE_SIZE
} else {
let cache_size = elems[1];
let cache_size = cache_size.parse::<usize>()?;
match cache_size {
0 => DEFAULT_MEMORY_CACHE_SIZE,
size => size,
}
}
};
tracing::debug!("Using memory cache with size: {cache_size}");
Ok(CachingStrategy::Memory(cache_size))
}
_ => {
return Err(RelatableError::InputError(format!(
"Unrecognized strategy: {strategy}"
))
.into());
}
}
}
}
impl Display for CachingStrategy {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
CachingStrategy::None => write!(f, "none"),
CachingStrategy::TruncateAll => write!(f, "truncate_all"),
CachingStrategy::Truncate => write!(f, "truncate"),
CachingStrategy::Trigger => write!(f, "trigger"),
CachingStrategy::Memory(size) => write!(f, "memory:{size}"),
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum DbKind {
Postgres,
Sqlite,
}
#[derive(Clone, Copy, Debug)]
pub struct SqlParam {
pub kind: DbKind,
pub index: usize,
}
impl SqlParam {
pub fn new(kind: &DbKind) -> Self {
Self {
kind: *kind,
index: 0,
}
}
pub fn next(&mut self) -> String {
match self.kind {
DbKind::Postgres => {
self.index += 1;
format!("${}", self.index)
}
DbKind::Sqlite => "?".to_string(),
}
}
pub fn get(&mut self, amount: usize) -> Vec<String> {
let mut params = vec![];
let mut made = 0;
while made < amount {
params.push(self.next());
made += 1;
}
params
}
pub fn get_as_list(&mut self, amount: usize) -> String {
self.get(amount).join(", ")
}
pub fn reset(&mut self) {
self.index = 0;
}
}
#[cfg(feature = "sqlx")]
#[derive(Debug)]
pub enum DbPool {
Sqlite(AnyPool),
Postgres(PgPool),
}
#[derive(Debug)]
pub enum DbActiveConnection {
#[cfg(feature = "rusqlite")]
Rusqlite(rusqlite::Connection),
}
#[derive(Debug)]
pub enum DbConnection {
#[cfg(feature = "sqlx")]
Sqlx(DbPool, DbKind),
#[cfg(feature = "rusqlite")]
Rusqlite(String),
}
impl DbConnection {
pub fn kind(&self) -> DbKind {
tracing::trace!("DbConnection::kind()");
match self {
#[cfg(feature = "sqlx")]
DbConnection::Sqlx(_, kind) => *kind,
#[cfg(feature = "rusqlite")]
DbConnection::Rusqlite(_) => DbKind::Sqlite,
}
}
pub async fn connect(database: &str) -> Result<(Self, Option<DbActiveConnection>)> {
tracing::trace!("DbConnection::connect({database})");
let is_postgresql = database.starts_with("postgresql://");
match is_postgresql {
true => {
#[cfg(not(feature = "sqlx"))]
return Err(RelatableError::InputError(
"rltbl was built without the sqlx feature, which is required for PostgreSQL \
support. To build rltbl with sqlx enabled, run \
`cargo build --features sqlx`"
.to_string(),
)
.into());
#[cfg(feature = "sqlx")]
{
let connection_options = PgConnectOptions::from_str(database)?;
let db_kind = DbKind::Postgres;
let pool = PgPoolOptions::new()
.max_connections(MAX_DB_CONNECTIONS)
.connect_with(connection_options)
.await?;
let connection = DbConnection::Sqlx(DbPool::Postgres(pool), db_kind);
Ok((connection, None))
}
}
false => {
#[allow(unused_variables)]
#[cfg(feature = "rusqlite")]
let tuple = (
DbConnection::Rusqlite(database.to_string()),
Some(DbActiveConnection::Rusqlite(rusqlite::Connection::open(
database,
)?)),
);
#[cfg(feature = "sqlx")]
let tuple = {
let url = {
if database.starts_with("sqlite://") {
database.to_string()
} else {
format!("sqlite://{database}?mode=rwc")
}
};
install_default_drivers();
let pool = AnyPool::connect(&url).await?;
let connection = DbConnection::Sqlx(DbPool::Sqlite(pool), DbKind::Sqlite);
(connection, None)
};
Ok(tuple)
}
}
}
pub fn reconnect(&self) -> Result<Option<DbActiveConnection>> {
tracing::trace!("DbConnection::reconnect()");
match self {
#[cfg(feature = "sqlx")]
DbConnection::Sqlx(_, _) => Ok(None),
#[cfg(feature = "rusqlite")]
DbConnection::Rusqlite(path) => Ok(Some(DbActiveConnection::Rusqlite(
rusqlite::Connection::open(path)?,
))),
}
}
pub async fn begin<'a>(
&self,
conn: &'a mut Option<DbActiveConnection>,
) -> Result<DbTransaction<'a>> {
tracing::trace!("DbConnection::begin({self:?}, {conn:?})");
match self {
#[cfg(feature = "sqlx")]
DbConnection::Sqlx(db_pool, kind) => match db_pool {
DbPool::Sqlite(pool) => {
let tx = pool.begin().await?;
Ok(DbTransaction::Sqlx(SqlxDbTransaction::Sqlite(tx), *kind))
}
DbPool::Postgres(pool) => {
let tx = pool.begin().await?;
Ok(DbTransaction::Sqlx(SqlxDbTransaction::Postgres(tx), *kind))
}
},
#[cfg(feature = "rusqlite")]
DbConnection::Rusqlite(_) => match conn {
None => {
return Err(RelatableError::InputError(
"Can't begin Rusqlite transaction: No connection provided".to_string(),
)
.into())
}
Some(DbActiveConnection::Rusqlite(ref mut conn)) => {
let tx = conn.transaction()?;
Ok(DbTransaction::Rusqlite(tx))
}
},
}
}
pub async fn query(&self, statement: &str, params: Option<&JsonValue>) -> Result<Vec<JsonRow>> {
tracing::trace!("DbConnection::query({self:?}, {statement}, {params:?})");
if !valid_params(params) {
tracing::warn!("Invalid parameter argument");
return Ok(vec![]);
}
match self {
#[cfg(feature = "sqlx")]
DbConnection::Sqlx(db_pool, _) => match db_pool {
DbPool::Sqlite(pool) => {
let query = prepare_sqlx_sqlite_query(&statement, params)?;
let mut rows = vec![];
for row in query.fetch_all(pool).await? {
rows.push(JsonRow::try_from(row)?);
}
Ok(rows)
}
DbPool::Postgres(pool) => {
let query = prepare_sqlx_pg_query(&statement, params)?;
let mut rows = vec![];
for row in query.fetch_all(pool).await? {
rows.push(JsonRow::try_from(row)?);
}
Ok(rows)
}
},
#[cfg(feature = "rusqlite")]
DbConnection::Rusqlite(path) => {
let conn = self.reconnect()?;
match conn {
Some(DbActiveConnection::Rusqlite(conn)) => {
let mut stmt = conn.prepare(&statement)?;
submit_rusqlite_statement(&mut stmt, params)
}
None => Err(RelatableError::DataError(format!(
"Unable to connect to the db at '{path}'"
))
.into()),
}
}
}
}
pub async fn query_one(
&self,
statement: &str,
params: Option<&JsonValue>,
) -> Result<Option<JsonRow>> {
tracing::trace!("DbConnection::query_one({statement}, {params:?})");
let rows = self.query(&statement, params).await?;
match rows.iter().next() {
Some(row) => Ok(Some(row.clone())),
None => Ok(None),
}
}
pub async fn query_value(
&self,
statement: &str,
params: Option<&JsonValue>,
) -> Result<Option<JsonValue>> {
tracing::trace!("DbConnection::query_value({statement}, {params:?})");
let rows = self.query(statement, params).await?;
Ok(extract_value(&rows))
}
pub async fn cache(
&self,
sql: &str,
params: Option<&JsonValue>,
tables: &Vec<String>,
strategy: &CachingStrategy,
) -> Result<Vec<JsonRow>> {
tracing::trace!("cache({sql}, {params:?}, {strategy:?})");
for t in vec!["message", "history", "change", "user"] {
if tables.contains(&t.to_string()) {
return self.query(&sql, params).await;
}
}
async fn _cache(
conn: &DbConnection,
tables: &Vec<String>,
sql: &str,
params: Option<&JsonValue>,
) -> Result<Vec<JsonRow>> {
let tables = tables
.iter()
.map(|t| json!(t).to_string())
.collect::<Vec<_>>()
.join(", ");
let (cache_sql, tables) = {
let mut sql_param = SqlParam::new(&conn.kind());
match conn.kind() {
DbKind::Postgres => {
let sql = format!(
r#"SELECT {}||rtrim(ltrim("value", '['), ']')||{} AS "value"
FROM "cache"
WHERE "tables"::TEXT = {}
AND "statement" = {}
AND "parameters" = {}
LIMIT 1"#,
sql_param.next(),
sql_param.next(),
sql_param.next(),
sql_param.next(),
sql_param.next()
);
(sql, format!("[{tables}]"))
}
DbKind::Sqlite => {
let sql = format!(
r#"SELECT {}||rtrim(ltrim("value", '['), ']')||{} AS "value"
FROM "cache"
WHERE CAST("tables" AS TEXT) = {}
AND "statement" = {}
AND "parameters" = {}
LIMIT 1"#,
sql_param.next(),
sql_param.next(),
sql_param.next(),
sql_param.next(),
sql_param.next()
);
(sql, format!("[{tables}]"))
}
}
};
let empty = json!("[]");
let json_params = params.unwrap_or(&empty);
let cache_params = json!([r#"[{"content": "#, "}]", tables, sql, json_params]);
match conn.query_one(&cache_sql, Some(&cache_params)).await? {
Some(json_row) => {
tracing::debug!("Cache hit for tables {tables}");
let value = json_row.get_string("value")?;
let json_rows: Vec<JsonRow> = serde_json::from_str(&value)?;
Ok(json_rows)
}
None => {
tracing::debug!("Cache miss for tables {tables}");
let json_rows = conn.query(sql, params).await?;
let json_rows_content = json_rows
.iter()
.map(|r| r.content.clone())
.collect::<Vec<_>>();
let mut sql_param = SqlParam::new(&conn.kind());
let update_cache_sql = match conn.kind() {
DbKind::Postgres => {
format!(
r#"INSERT INTO "cache"
("tables", "statement", "parameters", "value")
VALUES ({}::JSONB, {}, {}, {})"#,
sql_param.next(),
sql_param.next(),
sql_param.next(),
sql_param.next(),
)
}
DbKind::Sqlite => {
format!(
r#"INSERT INTO "cache"
("tables", "statement", "parameters", "value")
VALUES ({}, {}, {}, {})"#,
sql_param.next(),
sql_param.next(),
sql_param.next(),
sql_param.next(),
)
}
};
let update_cache_params = json!([tables, sql, json_params, json_rows_content]);
conn.query(&update_cache_sql, Some(&update_cache_params))
.await?;
Ok(json_rows)
}
}
}
match strategy {
CachingStrategy::None => self.query(sql, params).await,
CachingStrategy::TruncateAll | CachingStrategy::Truncate | CachingStrategy::Trigger => {
_cache(self, tables, sql, params).await
}
CachingStrategy::Memory(cache_size) => {
let mut cache = core::CACHE.lock().expect("Could not lock cache");
let keys = cache.keys().map(|key| key.clone()).collect::<Vec<_>>();
for (i, key) in keys.iter().enumerate().rev() {
if i >= *cache_size {
tracing::debug!("Removing {key:?} ({i}th entry) from cache");
cache.remove(&key);
} else {
break;
}
}
let tables = tables
.iter()
.map(|t| json!(t).to_string())
.collect::<Vec<_>>()
.join(", ");
let mem_key = MemoryCacheKey {
tables: tables.to_string(),
statement: sql.to_string(),
parameters: format!("{params:?}"),
};
match cache.get(&mem_key) {
Some(json_rows) => {
tracing::debug!("Cache hit for tables {tables}");
Ok(json_rows.to_vec())
}
None => {
tracing::debug!("Cache miss for tables {tables}");
let json_rows = block_on(self.query(sql, params))?;
cache.insert(
MemoryCacheKey {
tables: tables.to_string(),
statement: sql.to_string(),
parameters: format!("{params:?}"),
},
json_rows.to_vec(),
);
Ok(json_rows)
}
}
}
}
}
}
#[cfg(feature = "sqlx")]
#[derive(Debug)]
pub enum SqlxDbTransaction<'a> {
Sqlite(Transaction<'a, Any>),
Postgres(Transaction<'a, Postgres>),
}
#[derive(Debug)]
pub enum DbTransaction<'a> {
#[cfg(feature = "sqlx")]
Sqlx(SqlxDbTransaction<'a>, DbKind),
#[cfg(feature = "rusqlite")]
Rusqlite(rusqlite::Transaction<'a>),
}
impl DbTransaction<'_> {
pub fn kind(&self) -> DbKind {
tracing::trace!("DbTransaction::kind({self:?})");
match self {
#[cfg(feature = "sqlx")]
DbTransaction::Sqlx(_, kind) => *kind,
#[cfg(feature = "rusqlite")]
DbTransaction::Rusqlite(_) => DbKind::Sqlite,
}
}
pub fn commit(self) -> Result<()> {
tracing::trace!("DbTransaction::commit({self:?})");
match self {
#[cfg(feature = "sqlx")]
DbTransaction::Sqlx(tx, _) => match tx {
SqlxDbTransaction::Sqlite(tx) => block_on(tx.commit())?,
SqlxDbTransaction::Postgres(tx) => block_on(tx.commit())?,
},
#[cfg(feature = "rusqlite")]
DbTransaction::Rusqlite(tx) => tx.commit()?,
};
Ok(())
}
pub fn rollback(self) -> Result<()> {
tracing::trace!("DbTransaction::rollback({self:?})");
match self {
#[cfg(feature = "sqlx")]
DbTransaction::Sqlx(tx, _) => match tx {
SqlxDbTransaction::Sqlite(tx) => block_on(tx.rollback())?,
SqlxDbTransaction::Postgres(tx) => block_on(tx.rollback())?,
},
#[cfg(feature = "rusqlite")]
DbTransaction::Rusqlite(tx) => tx.rollback()?,
};
Ok(())
}
pub fn query(&mut self, statement: &str, params: Option<&JsonValue>) -> Result<Vec<JsonRow>> {
tracing::trace!("DbTransaction::query({self:?}, {statement}, {params:?})");
if !valid_params(params) {
tracing::warn!("invalid parameter argument");
return Ok(vec![]);
}
match self {
#[cfg(feature = "sqlx")]
DbTransaction::Sqlx(tx, _) => match tx {
SqlxDbTransaction::Sqlite(tx) => {
let query = prepare_sqlx_sqlite_query(&statement, params)?;
let mut rows = vec![];
for row in block_on(query.fetch_all(block_on(tx.acquire())?))? {
rows.push(JsonRow::try_from(row)?);
}
Ok(rows)
}
SqlxDbTransaction::Postgres(tx) => {
let query = prepare_sqlx_pg_query(&statement, params)?;
let mut rows = vec![];
for row in block_on(query.fetch_all(block_on(tx.acquire())?))? {
rows.push(JsonRow::try_from(row)?);
}
Ok(rows)
}
},
#[cfg(feature = "rusqlite")]
DbTransaction::Rusqlite(tx) => {
let mut stmt = tx.prepare(&statement)?;
submit_rusqlite_statement(&mut stmt, params)
}
}
}
pub fn query_one(
&mut self,
statement: &str,
params: Option<&JsonValue>,
) -> Result<Option<JsonRow>> {
tracing::trace!("DbTransaction::query_one({self:?}, {statement}, {params:?})");
let rows = self.query(&statement, params)?;
match rows.iter().next() {
Some(row) => Ok(Some(row.clone())),
None => Ok(None),
}
}
pub fn query_value(
&mut self,
statement: &str,
params: Option<&JsonValue>,
) -> Result<Option<JsonValue>> {
tracing::trace!("DbTransaction::query_value({self:?}, {statement}, {params:?})");
let rows = self.query(statement, params)?;
Ok(extract_value(&rows))
}
}
pub fn interpolate_sql(sql: &str, params: Option<&JsonValue>, kind: &DbKind) -> Result<String> {
tracing::trace!("interpolate_sql({sql}, {params:?}, {kind:?})");
let params = match params {
Some(JsonValue::Array(params)) => params.iter().collect::<Vec<_>>(),
None => vec![],
Some(params) => {
tracing::warn!("Invalid parameter list: {params:?}");
vec![]
}
};
let mut final_sql = String::from("");
let mut saved_start = 0;
let quotes = r#"('[^'\\]*(?:\\.[^'\\]*)*'|"[^"\\]*(?:\\.[^"\\]*)*")"#;
let rx = match kind {
DbKind::Sqlite => Regex::new(&format!(r#"{}|\B[?]\B"#, quotes))?,
DbKind::Postgres => Regex::new(&format!(r#"{}|\B[$]\d+\b"#, quotes))?,
};
let mut param_index = 0;
for m in rx.find_iter(&sql) {
let this_match = &sql[m.start()..m.end()];
final_sql.push_str(&sql[saved_start..m.start()]);
if !((this_match.starts_with("\"") && this_match.ends_with("\""))
|| (this_match.starts_with("'") && this_match.ends_with("'")))
{
let param = params.get(param_index);
match param {
None => {
return Err(RelatableError::InputError(format!(
"No parameter at index {param_index}"
))
.into())
}
Some(param) => {
match param {
JsonValue::String(param) => final_sql.push_str(&format!("'{param}'")),
JsonValue::Number(param) => final_sql.push_str(&format!("{param}")),
JsonValue::Bool(param) => final_sql.push_str(¶m.to_string()),
JsonValue::Array(param) => final_sql.push_str(&format!("{param:?}")),
JsonValue::Object(param) => final_sql.push_str(&format!("{param:?}")),
JsonValue::Null => final_sql.push_str(&"NULL".to_string()),
};
}
};
param_index += 1;
} else {
final_sql.push_str(&format!("{}", this_match));
}
saved_start = m.start() + this_match.len();
}
final_sql.push_str(&sql[saved_start..]);
Ok(final_sql)
}
pub fn is_simple(db_object_name: &str) -> Result<(), String> {
tracing::trace!("is_simple({db_object_name})");
let db_object_root = db_object_name.splitn(2, ".").collect::<Vec<_>>()[0];
if !DB_OBJECT_REGEX.is_match(&db_object_root) {
Err(format!(
"Illegal database object name: '{}' in '{}'. Does not match: /{}/",
db_object_root, db_object_name, DB_OBJECT_MATCH_STR,
))
} else {
Ok(())
}
}
pub fn is_clause(db_kind: &DbKind) -> String {
tracing::trace!("is_clause({db_kind:?})");
match db_kind {
DbKind::Sqlite => "IS".into(),
DbKind::Postgres => "IS NOT DISTINCT FROM".into(),
}
}
pub fn is_not_clause(db_kind: &DbKind) -> String {
tracing::trace!("is_not_clause({db_kind:?})");
match db_kind {
DbKind::Sqlite => "IS NOT".into(),
DbKind::Postgres => "IS DISTINCT FROM".into(),
}
}
#[cfg(feature = "sqlx")]
pub fn prepare_sqlx_sqlite_query<'a>(
statement: &'a str,
params: Option<&'a JsonValue>,
) -> Result<Query<'a, Any, AnyArguments<'a>>> {
tracing::trace!("prepare_sqlx_query({statement}, {params:?})");
let mut query = sqlx::query::<Any>(&statement);
if let Some(params) = params {
for param in params.as_array().unwrap() {
match param {
JsonValue::Number(n) => match n.as_i64() {
Some(p) => query = query.bind(p),
None => match n.as_f64() {
Some(p) => query = query.bind(p),
None => panic!(),
},
},
JsonValue::String(s) => query = query.bind(s),
_ => query = query.bind(param.to_string()),
};
}
}
Ok(query)
}
#[cfg(feature = "sqlx")]
pub fn prepare_sqlx_pg_query<'a>(
statement: &'a str,
params: Option<&'a JsonValue>,
) -> Result<Query<'a, Postgres, PgArguments>> {
tracing::trace!("prepare_sqlx_query({statement}, {params:?})");
let mut query = sqlx::query::<Postgres>(&statement);
if let Some(params) = params {
for param in params.as_array().unwrap() {
match param {
JsonValue::Number(n) => match n.as_i64() {
Some(p) => query = query.bind(p),
None => match n.as_f64() {
Some(p) => query = query.bind(p),
None => panic!(),
},
},
JsonValue::String(s) => query = query.bind(s),
_ => query = query.bind(param.to_string()),
};
}
}
Ok(query)
}
#[cfg(feature = "rusqlite")]
fn submit_rusqlite_statement(
stmt: &mut rusqlite::Statement<'_>,
params: Option<&JsonValue>,
) -> Result<Vec<JsonRow>> {
tracing::trace!("submit_rusqlite_statement({stmt:?}, {params:?})");
let column_names = stmt
.column_names()
.iter()
.map(|c| c.to_string())
.collect::<Vec<_>>();
let column_names = column_names.iter().map(|c| c.as_str()).collect::<Vec<_>>();
if let Some(params) = params {
for (i, param) in params.as_array().unwrap().iter().enumerate() {
let param = match param {
JsonValue::String(s) => s,
_ => ¶m.to_string(),
};
stmt.raw_bind_parameter(i + 1, param)?;
}
}
let mut rows = stmt.raw_query();
let mut result = Vec::new();
while let Some(row) = rows.next()? {
result.push(JsonRow::from_rusqlite(&column_names, row));
}
Ok(result)
}
fn valid_params(params: Option<&JsonValue>) -> bool {
tracing::trace!("valid_params({params:?})");
if let Some(params) = params {
match params {
JsonValue::Array(_) => true,
_ => false,
}
} else {
true
}
}
fn extract_value(rows: &Vec<JsonRow>) -> Option<JsonValue> {
tracing::trace!("extract_value({rows:?})");
match rows.iter().next() {
Some(row) => match row.content.values().next() {
Some(value) => Some(value.clone()),
None => None,
},
None => None,
}
}
pub fn generate_table_ddl(
table: &Table,
force: bool,
db_kind: &DbKind,
caching_strategy: &CachingStrategy,
) -> Result<Vec<String>> {
tracing::trace!("generate_table_ddl({table:?}, {force}, {db_kind:?}, {caching_strategy:?})");
if table.has_meta {
for (cname, col) in table.columns.iter() {
if cname == "_id" || cname == "_order" {
return Err(RelatableError::InputError(format!(
"column {cname} conflicts with has_meta == {has_meta}",
has_meta = table.has_meta,
))
.into());
}
if col.primary_key {
return Err(RelatableError::InputError(format!(
"Primary key on column {cname} conflicts with has_meta == {has_meta}",
has_meta = table.has_meta,
))
.into());
}
}
}
let mut ddl = vec![];
let mut column_clauses = vec![];
for (cname, col) in table.columns.iter() {
if col.table != table.name {
return Err(RelatableError::InputError(format!(
"Table name mismatch: '{}' != '{}'",
col.table, table.name,
))
.into());
}
let sql_type = col.datatype.infer_sql_type(&col.datatype_hierarchy);
let clause = format!(
r#""{cname}" {sql_type}{unique}"#,
unique = match col.unique {
true => " UNIQUE",
false => "",
},
);
column_clauses.push(clause);
}
if force {
match db_kind {
DbKind::Postgres => {
ddl.push(format!(r#"DROP TABLE IF EXISTS "{}" CASCADE"#, table.name))
}
DbKind::Sqlite => ddl.push(format!(r#"DROP TABLE IF EXISTS "{}""#, table.name)),
}
}
let mut sql = format!(r#"CREATE TABLE "{}" ( "#, table.name);
if table.has_meta {
sql.push_str(match db_kind {
DbKind::Sqlite => {
"_id INTEGER PRIMARY KEY AUTOINCREMENT, \
_order INTEGER UNIQUE, "
}
DbKind::Postgres => {
"_id SERIAL PRIMARY KEY, \
_order BIGINT UNIQUE, "
}
});
}
sql.push_str(&format!(" {})", column_clauses.join(", ")));
ddl.push(sql);
if table.has_meta {
add_metacolumn_trigger_ddl(&mut ddl, &table.name, db_kind);
}
if let CachingStrategy::Trigger = caching_strategy {
add_caching_trigger_ddl(&mut ddl, &table.name, db_kind);
}
Ok(ddl)
}
pub fn add_metacolumn_trigger_ddl(ddl: &mut Vec<String>, table: &str, db_kind: &DbKind) {
let update_stmt = format!(
r#"UPDATE "{table}" SET _order = ({NEW_ORDER_MULTIPLIER} * NEW._id)
WHERE _id = NEW._id;"#
);
match db_kind {
DbKind::Sqlite => {
ddl.push(format!(
r#"CREATE TRIGGER "{table}_order"
AFTER INSERT ON "{table}"
WHEN NEW._order IS NULL
BEGIN
{update_stmt}
END"#
));
}
DbKind::Postgres => {
ddl.push(format!(
r#"CREATE OR REPLACE FUNCTION "update_order_and_nextval_{table}"()
RETURNS TRIGGER
LANGUAGE PLPGSQL
AS
$$
BEGIN
IF NEW._order IS NOT DISTINCT FROM NULL THEN
{update_stmt}
END IF;
IF NEW._id > (SELECT MAX(last_value) FROM "{table}__id_seq") THEN
PERFORM setval('{table}__id_seq', NEW._id);
END IF;
RETURN NEW;
END;
$$"#
));
ddl.push(format!(
r#"CREATE TRIGGER "{table}_order"
AFTER INSERT ON "{table}"
FOR EACH ROW
EXECUTE FUNCTION "update_order_and_nextval_{table}"()"#
));
}
};
}
pub fn add_caching_trigger_ddl(ddl: &mut Vec<String>, table: &str, db_kind: &DbKind) {
match db_kind {
DbKind::Sqlite => {
ddl.push(format!(
r#"CREATE TRIGGER "{table}_cache_after_insert"
AFTER INSERT ON "{table}"
BEGIN
DELETE FROM "cache" WHERE "tables" LIKE '%"{table}"%';
END"#
));
ddl.push(format!(
r#"CREATE TRIGGER "{table}_cache_after_update"
AFTER UPDATE ON "{table}"
BEGIN
DELETE FROM "cache" WHERE "tables" LIKE '%"{table}"%';
END"#
));
ddl.push(format!(
r#"CREATE TRIGGER "{table}_cache_after_delete"
AFTER DELETE ON "{table}"
BEGIN
DELETE FROM "cache" WHERE "tables" LIKE '%"{table}"%';
END"#
));
}
DbKind::Postgres => {
ddl.push(format!(
r#"CREATE OR REPLACE FUNCTION "clean_cache_for_{table}"()
RETURNS TRIGGER
LANGUAGE PLPGSQL
AS
$$
BEGIN
DELETE FROM "cache" WHERE "tables" ? '{table}';
RETURN NEW;
END;
$$"#
));
ddl.push(format!(
r#"CREATE TRIGGER "{table}_cache_after_insert"
AFTER INSERT ON "{table}"
EXECUTE FUNCTION "clean_cache_for_{table}"()"#
));
ddl.push(format!(
r#"CREATE TRIGGER "{table}_cache_after_update"
AFTER UPDATE ON "{table}"
EXECUTE FUNCTION "clean_cache_for_{table}"()"#
));
ddl.push(format!(
r#"CREATE TRIGGER "{table}_cache_after_delete"
AFTER DELETE ON "{table}"
EXECUTE FUNCTION "clean_cache_for_{table}"()"#
));
}
};
}
pub(crate) fn generate_default_view_ddl(
table_name: &str,
id_col: &str,
order_col: &str,
columns: &Vec<Column>,
kind: &DbKind,
) -> Vec<String> {
tracing::trace!(
"generate_default_view_ddl({table_name}, {id_col}, {order_col}, {columns:?}, {kind:?})"
);
let view_name = format!("{table_name}_default_view");
match kind {
DbKind::Sqlite => vec![
format!(r#"DROP VIEW IF EXISTS "{}""#, view_name),
format!(
r#"CREATE VIEW "{view}" AS
SELECT
{id_col} AS _id,
{order_col} AS _order,
(SELECT "change_id"
FROM "history"
WHERE "table" = '{table}'
AND "row" = {id_col}
ORDER BY "change_id" DESC
LIMIT 1
) AS _change_id,
(SELECT '[' || GROUP_CONCAT("after") || ']'
FROM (
SELECT "after"
FROM "history"
WHERE "table" = '{table}'
AND "after" IS NOT NULL
AND "row" = {id_col}
ORDER BY "history_id"
)
) AS "_history",
(SELECT NULLIF(
JSON_GROUP_ARRAY(
JSON_OBJECT(
'column', "column",
'value', "value",
'level', "level",
'rule', "rule",
'message', "message"
)
),
'[]'
) AS "_message"
FROM "message"
WHERE "table" = '{table}'
AND "row" = {id_col}
ORDER BY "column", "message_id"
) AS "_message",
{columns}
FROM "{table}""#,
table = table_name,
view = view_name,
columns = columns
.iter()
.map(|c| format!(r#""{}""#, c.name))
.collect::<Vec<_>>()
.join(", "),
),
],
DbKind::Postgres => vec![format!(
r#"CREATE OR REPLACE VIEW "{view}" AS
SELECT
"{id_col}" AS _id,
"{order_col}" AS _order,
(
SELECT "change_id"
FROM "history"
WHERE "table" = '{table}'
AND "row" = {id_col}
ORDER BY "change_id" DESC
LIMIT 1
) AS _change_id,
(
SELECT ('['::TEXT || string_agg(h.after, ','::TEXT)) || ']'::TEXT
FROM ( SELECT "history"."after"
FROM "history"
WHERE "history"."table" = '{table}'
AND "after" IS DISTINCT FROM NULL
AND "row" = "{id_col}"
ORDER BY "history_id" ) h
) AS "_history",
(
SELECT json_agg(m.*)::TEXT AS json_agg
FROM ( SELECT "message"."column",
"message"."value",
"message"."level",
"message"."rule",
"message"."message"
FROM "message"
WHERE "message"."table" = '{table}' AND "message"."row" = "{id_col}"
ORDER BY "message"."column", "message"."message_id") m
) AS "_message",
{columns}
FROM "{table}""#,
table = table_name,
view = view_name,
columns = columns
.iter()
.map(|c| format!(r#""{}""#, c.name))
.collect::<Vec<_>>()
.join(", "),
)],
}
}
pub fn sprintf_to_pg_char(
flag_opt: &str,
width_opt: &str,
precision_opt: &str,
format_type: &str,
) -> String {
tracing::trace!("sprintf_to_pg_char({flag_opt}, {width_opt}, {precision_opt}, {format_type})");
if format_type == "s" {
if flag_opt != "" || width_opt != "" || precision_opt != "" {
tracing::warn!(
"Ignoring options: flag: '{flag_opt}', width: '{width_opt}', precision: \
'{precision_opt}' for format type '{format_type}'"
);
}
return "".to_string();
}
let default_width = 99;
let default_precision = 6;
let mut zero_pad = false;
let mut pm_sign = false;
let mut comma_sep = false;
match flag_opt {
"" => (),
"0" => zero_pad = true,
"+" => pm_sign = true,
"," => comma_sep = true,
"-" => tracing::warn!("Flag '-' is unsupported"),
" " => tracing::warn!("Flag ' ' is unsupported"),
"#" => tracing::warn!("Flag '#' is unsupported"),
"!" => tracing::warn!("Flag '!' is unsupported"),
invalid => tracing::warn!("Invalid flag: '{invalid}'"),
};
let width = match width_opt {
"" => default_width,
width => match width.parse::<usize>() {
Ok(width) => width,
Err(err) => {
tracing::warn!("Could not parse width: {err}");
default_width
}
},
};
let precision = match precision_opt {
"" => default_precision,
precision => match precision.parse::<usize>() {
Ok(precision) => precision,
Err(err) => {
tracing::warn!("Could not parse precision: {err}");
default_precision
}
},
};
let mut to_char_format = "".to_string();
if zero_pad {
to_char_format.push_str("0");
} else if pm_sign {
to_char_format.push_str("SG");
}
let mut digits = "".to_string();
for (i, _) in (0..width).enumerate() {
digits.push_str("9");
let i = i + 1;
if comma_sep && i != width {
if i % 3 == 0 {
digits.push_str(",");
}
}
}
let digits = digits.chars().rev().collect::<String>();
to_char_format.push_str(&digits);
if precision > 0 {
to_char_format.push_str(".");
for _ in 0..precision {
to_char_format.push_str("9");
}
}
tracing::debug!("Using to_char() format for PostgreSQL: '{to_char_format}'");
to_char_format
}
pub fn split_sprintf_format(sprintf_format: &str) -> (String, String, String, String) {
tracing::trace!("split_sprintf_format({sprintf_format:?})");
let sprintf_regex = Regex::new(r#"^%([\-+ 0#,!])?([1-9]+)?((.)([0-9]+))?(\w)$"#).unwrap();
let valid_format_types = ["d", "i", "c", "o", "u", "x", "e", "f", "g", "a", "s"];
match sprintf_format {
"" => (
"".to_string(),
"".to_string(),
"".to_string(),
"s".to_string(),
),
sprintf_format => match sprintf_regex.captures(sprintf_format) {
None => {
tracing::warn!("Illegal format: '{}'", sprintf_format);
(
"".to_string(),
"".to_string(),
"".to_string(),
"s".to_string(),
)
}
Some(captures) => {
let flag_opt = captures
.get(1)
.and_then(|c| Some(c.as_str().to_string()))
.unwrap_or("".to_string());
let width_opt = captures
.get(2)
.and_then(|c| Some(c.as_str().to_string()))
.unwrap_or("".to_string());
let precision_opt = captures
.get(5)
.and_then(|c| Some(c.as_str().to_string()))
.unwrap_or("".to_string());
let mut format_type = &captures[6];
if !valid_format_types.contains(&format_type) {
tracing::warn!("Invalid format type: '{format_type}'");
format_type = "s";
}
(flag_opt, width_opt, precision_opt, format_type.to_string())
}
},
}
}
pub(crate) fn generate_text_view_ddl(
table_name: &str,
id_col: &str,
order_col: &str,
columns: &Vec<Column>,
kind: &DbKind,
) -> Vec<String> {
tracing::trace!(
"generate_text_view_ddl({table_name}, {id_col}, {order_col}, {columns:?}, {kind:?})"
);
let view_name = format!("{table_name}_text_view");
let mut inner_columns = columns
.iter()
.map(|column| {
let column_cast = {
let (flag_opt, width_opt, precision_opt, format_type) =
split_sprintf_format(column.datatype.format.as_ref());
if *kind == DbKind::Sqlite {
let dt_format = format!(
"%{flag_opt}{width_opt}{precision_opt}{format_type}",
precision_opt = match precision_opt.as_str() {
"" => "".to_string(),
_ => format!(".{precision_opt}"),
}
);
tracing::debug!("Formatting column '{}' using '{dt_format}'", column.name);
format!(r#"FORMAT('{}', "{}")"#, dt_format, column.name)
} else {
match sprintf_to_pg_char(&flag_opt, &width_opt, &precision_opt, &format_type)
.as_str()
{
"" => format!(r#""{}"::TEXT"#, column.name),
dt_format => {
format!(r#"LTRIM(TO_CHAR("{}", '{dt_format}'), ' ')"#, column.name)
}
}
}
};
format!(
r#"CASE
WHEN "{column}" {is_clause} NULL THEN (
SELECT "value"
FROM "message"
WHERE "row" = "_id"
AND "column" = '{column}'
AND "table" = '{table_name}'
ORDER BY "message_id" DESC
LIMIT 1
)
ELSE {column_cast}
END AS "{column}""#,
column = column.name,
is_clause = is_clause(kind)
)
})
.collect::<Vec<_>>();
let inner_columns = {
let mut v = vec![
"_id".to_string(),
"_order".to_string(),
"_message".to_string(),
"_history".to_string(),
];
v.append(&mut inner_columns);
v
};
let mut outer_columns = columns
.iter()
.map(|column| format!(r#"t."{}""#, column.name))
.collect::<Vec<_>>();
let outer_columns = {
let mut v = vec![
"t._id".to_string(),
"t._order".to_string(),
"t._message".to_string(),
"t._history".to_string(),
];
v.append(&mut outer_columns);
v
};
let create_view_sql = format!(
r#"CREATE VIEW "{view_name}" AS
SELECT {outer_columns}
FROM (
SELECT {inner_columns}
FROM "{table_name}_default_view"
) t"#,
outer_columns = outer_columns.join(", "),
inner_columns = inner_columns.join(", "),
);
vec![
format!(r#"DROP VIEW IF EXISTS "{}""#, view_name),
create_view_sql,
]
}
pub fn generate_table_table_ddl(force: bool, db_kind: &DbKind) -> Vec<String> {
tracing::trace!("generate_table_table_ddl({force}, {db_kind:?})");
let mut ddl = vec![];
if force {
if let DbKind::Postgres = db_kind {
ddl.push(format!(r#"DROP TABLE IF EXISTS "table" CASCADE"#));
}
}
let pkey_clause = match db_kind {
DbKind::Sqlite => "INTEGER PRIMARY KEY AUTOINCREMENT",
DbKind::Postgres => "SERIAL PRIMARY KEY",
};
ddl.push(format!(
r#"CREATE TABLE "table" (
"_id" {pkey_clause},
"_order" BIGINT UNIQUE,
"table" TEXT UNIQUE,
"path" TEXT UNIQUE
)"#
));
add_metacolumn_trigger_ddl(&mut ddl, "table", db_kind);
ddl
}
pub fn generate_cache_table_ddl(force: bool, db_kind: &DbKind) -> Vec<String> {
tracing::trace!("generate_cache_table_ddl({force}, {db_kind:?})");
let mut ddl = vec![];
if force {
if let DbKind::Postgres = db_kind {
ddl.push(format!(r#"DROP TABLE IF EXISTS "cache" CASCADE"#));
}
}
let json_type = match db_kind {
DbKind::Postgres => "JSONB",
DbKind::Sqlite => "JSON",
};
ddl.push(format!(
r#"CREATE TABLE "cache" (
"tables" {json_type},
"statement" TEXT,
"parameters" TEXT,
"value" TEXT,
PRIMARY KEY ("tables", "statement", "parameters")
)"#
));
ddl
}
pub fn generate_user_table_ddl(force: bool, db_kind: &DbKind) -> Vec<String> {
tracing::trace!("generate_user_table_ddl({force}, {db_kind:?})");
let mut ddl = vec![];
if force {
if let DbKind::Postgres = db_kind {
ddl.push(format!(r#"DROP TABLE IF EXISTS "user" CASCADE"#));
}
}
ddl.push(format!(
r#"CREATE TABLE "user" (
"name" TEXT PRIMARY KEY,
"color" TEXT,
"cursor" TEXT,
"datetime" TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)"#
));
ddl
}
pub fn generate_change_table_ddl(force: bool, db_kind: &DbKind) -> Vec<String> {
tracing::trace!("generate_change_table_ddl({force}, {db_kind:?})");
match db_kind {
DbKind::Sqlite => {
vec![r#"CREATE TABLE "change" (
change_id INTEGER PRIMARY KEY AUTOINCREMENT,
"datetime" TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
"user" TEXT NOT NULL,
"action" TEXT NOT NULL,
"table" TEXT NOT NULL,
"description" TEXT,
"content" TEXT,
FOREIGN KEY ("user") REFERENCES "user"("name")
)"#
.to_string()]
}
DbKind::Postgres => {
let mut ddl = vec![];
if force {
if let DbKind::Postgres = db_kind {
ddl.push(format!(r#"DROP TABLE IF EXISTS "change" CASCADE"#));
}
}
ddl.push(format!(
r#"CREATE TABLE "change" (
change_id SERIAL PRIMARY KEY,
"datetime" TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
"user" TEXT NOT NULL,
"action" TEXT NOT NULL,
"table" TEXT NOT NULL,
"description" TEXT,
"content" TEXT,
FOREIGN KEY ("user") REFERENCES "user"("name")
)"#
));
ddl
}
}
}
pub fn generate_history_table_ddl(force: bool, db_kind: &DbKind) -> Vec<String> {
tracing::trace!("generate_history_table_ddl({force}, {db_kind:?})");
match db_kind {
DbKind::Sqlite => {
vec![r#"CREATE TABLE "history" (
history_id INTEGER PRIMARY KEY AUTOINCREMENT,
change_id INTEGER NOT NULL,
"table" TEXT NOT NULL,
"row" BIGINT NOT NULL,
"before" TEXT,
"after" TEXT,
FOREIGN KEY ("change_id") REFERENCES "change"("change_id"),
FOREIGN KEY ("table") REFERENCES "table"("table")
)"#
.to_string()]
}
DbKind::Postgres => {
let mut ddl = vec![];
if force {
if let DbKind::Postgres = db_kind {
ddl.push(format!(r#"DROP TABLE IF EXISTS "history" CASCADE"#));
}
}
ddl.push(format!(
r#"CREATE TABLE "history" (
history_id SERIAL PRIMARY KEY,
change_id INTEGER NOT NULL,
"table" TEXT NOT NULL,
"row" BIGINT NOT NULL,
"before" TEXT,
"after" TEXT,
FOREIGN KEY ("change_id") REFERENCES "change"("change_id"),
FOREIGN KEY ("table") REFERENCES "table"("table")
)"#
));
ddl
}
}
}
pub fn generate_message_table_ddl(force: bool, db_kind: &DbKind) -> Vec<String> {
tracing::trace!("generate_message_table_ddl({force}, {db_kind:?})");
match db_kind {
DbKind::Sqlite => {
vec![r#"CREATE TABLE "message" (
"message_id" INTEGER PRIMARY KEY AUTOINCREMENT,
"added_by" TEXT,
"table" TEXT NOT NULL,
"row" BIGINT NOT NULL,
"column" TEXT NOT NULL,
"value" TEXT,
"level" TEXT,
"rule" TEXT,
"message" TEXT,
FOREIGN KEY ("table") REFERENCES "table"("table")
)"#
.to_string()]
}
DbKind::Postgres => {
let mut ddl = vec![];
if force {
if let DbKind::Postgres = db_kind {
ddl.push(format!(r#"DROP TABLE IF EXISTS "message" CASCADE"#));
}
}
ddl.push(format!(
r#"CREATE TABLE "message" (
"message_id" SERIAL PRIMARY KEY,
"added_by" TEXT,
"table" TEXT NOT NULL,
"row" BIGINT NOT NULL,
"column" TEXT NOT NULL,
"value" TEXT,
"level" TEXT,
"rule" TEXT,
"message" TEXT,
FOREIGN KEY ("table") REFERENCES "table"("table")
)"#
));
ddl
}
}
}
pub fn generate_meta_tables_ddl(force: bool, db_kind: &DbKind) -> Vec<String> {
tracing::trace!("generate_meta_tables_ddl({force}, {db_kind:?})");
let mut ddl = generate_table_table_ddl(force, db_kind);
ddl.append(&mut generate_cache_table_ddl(force, db_kind));
ddl.append(&mut generate_user_table_ddl(force, db_kind));
ddl.append(&mut generate_change_table_ddl(force, db_kind));
ddl.append(&mut generate_history_table_ddl(force, db_kind));
ddl.append(&mut generate_message_table_ddl(force, db_kind));
ddl
}
pub fn json_to_string(value: &JsonValue) -> String {
tracing::trace!("json_to_string({value:?})");
match value {
JsonValue::Null => "".to_string(),
JsonValue::Bool(value) => value.to_string(),
JsonValue::Number(value) => value.to_string(),
JsonValue::String(value) => value.to_string(),
JsonValue::Array(value) => format!("{value:?}"),
JsonValue::Object(value) => format!("{value:?}"),
}
}
pub fn json_to_unsigned(value: &JsonValue) -> Result<u64> {
tracing::trace!("json_to_unsigned({value:?})");
match value {
JsonValue::Bool(flag) => match flag {
true => Ok(1),
false => Ok(0),
},
JsonValue::Number(value) => match value.as_u64() {
Some(unsigned) => Ok(unsigned as u64),
None => Err(
RelatableError::InputError(format!("{value} is not an unsigned integer")).into(),
),
},
JsonValue::String(value_str) => match value_str.parse::<u64>() {
Ok(unsigned) => Ok(unsigned),
Err(err) => Err(RelatableError::InputError(format!(
"{value} could not be parsed as an unsigned integer: {err}"
))
.into()),
},
_ => Err(RelatableError::InputError(format!(
"{value} could not be parsed as an unsigned integer"
))
.into()),
}
}
pub trait VecInto<D> {
fn vec_into(self) -> Vec<D>;
}
impl<E, D> VecInto<D> for Vec<E>
where
D: From<E>,
{
fn vec_into(self) -> Vec<D> {
self.into_iter().map(std::convert::Into::into).collect()
}
}
#[derive(Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct JsonRow {
pub content: JsonMap<String, JsonValue>,
}
impl JsonRow {
pub fn new() -> Self {
Self {
content: JsonMap::new(),
}
}
pub fn nullify(row: &Self, table: &Table) -> Self {
tracing::trace!("JsonRow::nullify({row:?}, {table:?})");
let mut nullified_row = JsonRow::new();
let default_col = Column::default();
for (column, value) in row.content.iter() {
match &table.columns.get(column).unwrap_or(&default_col).nulltype {
Some(supported) if supported.name == "empty" => match value {
JsonValue::String(s) if s == "" => {
nullified_row
.content
.insert(column.to_string(), JsonValue::Null);
}
value => {
nullified_row
.content
.insert(column.to_string(), value.clone());
}
},
Some(unsupported) => {
tracing::warn!("Unsupported nulltype: '{}'", unsupported.name);
nullified_row
.content
.insert(column.to_string(), value.clone());
}
None => {
nullified_row
.content
.insert(column.to_string(), value.clone());
}
};
}
tracing::debug!("Nullified row: {row:?} to: {nullified_row:?}");
nullified_row
}
pub fn nullify_value(table: &Table, column: &str, value: &JsonValue) -> JsonValue {
tracing::trace!("JsonRow::nullify_value({table:?}, {column}, {value:?})");
let default_col = Column::default();
match &table.columns.get(column).unwrap_or(&default_col).nulltype {
Some(supported) if supported.name == "empty" => match value {
JsonValue::String(s) if s == "" => JsonValue::Null,
_ => value.clone(),
},
Some(unsupported) => {
tracing::warn!("Unsupported nulltype: '{}'", unsupported.name);
value.clone()
}
None => value.clone(),
}
}
pub fn get_value(&self, column_name: &str) -> Result<JsonValue> {
tracing::trace!("JsonRow::get_value({self:?}, {column_name})");
let value = self.content.get(column_name);
match value {
Some(value) => Ok(value.clone()),
None => Err(RelatableError::DataError("missing value".to_string()).into()),
}
}
pub fn get_string(&self, column_name: &str) -> Result<String> {
tracing::trace!("JsonRow::get_string({self:?}, {column_name})");
let value = self.content.get(column_name);
match value {
Some(value) => Ok(json_to_string(&value)),
None => Err(RelatableError::DataError("missing value".to_string()).into()),
}
}
pub fn get_unsigned(&self, column_name: &str) -> Result<u64> {
tracing::trace!("JsonRow::get_unsigned({self:?}, {column_name})");
let value = self.content.get(column_name);
match value {
Some(value) => json_to_unsigned(&value),
None => Err(RelatableError::DataError("missing value".to_string()).into()),
}
}
pub fn from_strings(strings: &Vec<&str>) -> Self {
tracing::trace!("JsonRow::from_strings({strings:?})");
let mut json_row = JsonRow::new();
for string in strings {
json_row.content.insert(string.to_string(), JsonValue::Null);
}
json_row
}
pub fn to_strings(&self) -> Vec<String> {
tracing::trace!("JsonRow::to_strings({self:?})");
let mut result = vec![];
for column_name in self.content.keys() {
result.push(self.get_string(column_name).expect("Column not found"));
}
result
}
pub fn to_string_map(&self) -> IndexMap<String, String> {
tracing::trace!("JsonRow::to_string_map({self:?})");
let mut result = IndexMap::new();
for column_name in self.content.keys() {
result.insert(
column_name.clone(),
self.get_string(column_name).expect("Column not found"),
);
}
result
}
#[cfg(feature = "rusqlite")]
pub fn from_rusqlite(column_names: &Vec<&str>, row: &rusqlite::Row) -> Self {
tracing::trace!("JsonRow::from_rusqlite({column_names:?}, {row:?})");
let mut content = JsonMap::new();
for column_name in column_names {
let value = match row.get_ref(*column_name) {
Ok(value) => match value {
rusqlite::types::ValueRef::Null => JsonValue::Null,
rusqlite::types::ValueRef::Integer(value) => JsonValue::from(value),
rusqlite::types::ValueRef::Real(value) => JsonValue::from(value),
rusqlite::types::ValueRef::Text(value)
| rusqlite::types::ValueRef::Blob(value) => {
let value = std::str::from_utf8(value).unwrap_or_default();
JsonValue::from(value)
}
},
Err(_) => JsonValue::Null,
};
content.insert(column_name.to_string(), value);
}
Self { content }
}
}
#[cfg(feature = "sqlx")]
impl TryFrom<AnyRow> for JsonRow {
type Error = anyhow::Error;
fn try_from(row: AnyRow) -> Result<Self> {
tracing::trace!("JsonRow::try_from::<AnyRow>(row)");
let mut content = JsonMap::new();
for column in row.columns() {
let mut value: JsonValue = JsonValue::Null;
if value.is_null() {
let x: Result<i64, sqlx::Error> = row.try_get(column.ordinal());
if let Ok(x) = x {
value = JsonValue::from(x);
}
}
if value.is_null() {
let x: Result<f64, sqlx::Error> = row.try_get(column.ordinal());
if let Ok(x) = x {
value = JsonValue::from(x);
}
}
if value.is_null() {
let x: Result<String, sqlx::Error> = row.try_get(column.ordinal());
if let Ok(x) = x {
value = JsonValue::from(x);
}
}
if value.is_null() {
let x: Result<bool, sqlx::Error> = row.try_get(column.ordinal());
if let Ok(x) = x {
value = JsonValue::from(x);
}
}
content.insert(column.name().into(), value);
}
Ok(Self { content })
}
}
#[cfg(feature = "sqlx")]
impl TryFrom<PgRow> for JsonRow {
type Error = anyhow::Error;
fn try_from(row: PgRow) -> Result<Self> {
tracing::trace!("JsonRow::try_from::<PgRow>(row)");
let mut content = JsonMap::new();
for column in row.columns() {
let column_type = column.type_info().name();
let value = match column_type {
"INT4" => {
let value: Result<i32, sqlx::Error> = row.try_get(column.ordinal());
match value {
Ok(value) => JsonValue::from(value),
Err(_) => JsonValue::Null,
}
}
"INT8" => {
let value: Result<i64, sqlx::Error> = row.try_get(column.ordinal());
match value {
Ok(value) => JsonValue::from(value),
Err(_) => JsonValue::Null,
}
}
"FLOAT4" => {
let value: Result<f32, sqlx::Error> = row.try_get(column.ordinal());
match value {
Ok(value) => JsonValue::from(value),
Err(_) => JsonValue::Null,
}
}
"FLOAT8" => {
let value: Result<f64, sqlx::Error> = row.try_get(column.ordinal());
match value {
Ok(value) => JsonValue::from(value),
Err(_) => JsonValue::Null,
}
}
"NUMERIC" => {
let value: Result<BigDecimal, sqlx::Error> = row.try_get(column.ordinal());
match value {
Ok(value) => {
let value = value.to_f64();
JsonValue::from(value)
}
Err(_) => JsonValue::Null,
}
}
"TEXT" => {
let value: Result<String, sqlx::Error> = row.try_get(column.ordinal());
match value {
Ok(value) => JsonValue::from(value),
Err(_) => JsonValue::Null,
}
}
"BOOL" => {
let value: Result<bool, sqlx::Error> = row.try_get(column.ordinal());
match value {
Ok(value) => JsonValue::from(value),
Err(_) => JsonValue::Null,
}
}
unsupported => {
tracing::warn!(
"Got unsupported column '{}' with type '{}'",
column.name(),
unsupported
);
let value: Result<String, sqlx::Error> = row.try_get(column.ordinal());
match value {
Ok(value) => JsonValue::from(value),
Err(_) => JsonValue::Null,
}
}
};
content.insert(column.name().into(), value);
}
Ok(Self { content })
}
}
impl std::fmt::Display for JsonRow {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{:?}", self.to_strings().join("\t"))
}
}
impl std::fmt::Debug for JsonRow {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{:?}", self.to_string_map())
}
}
impl From<JsonRow> for Vec<String> {
fn from(row: JsonRow) -> Self {
row.to_strings()
}
}
#[cfg(test)]
mod tests {
use crate::{core::Relatable, select::Select, sql::CachingStrategy};
use async_std::task::block_on;
use pretty_assertions::assert_eq;
#[test]
fn test_cache() {
let rltbl = block_on(Relatable::build_demo(
Some("build/test_cache.db"),
&true,
10,
&CachingStrategy::Trigger,
))
.unwrap();
let select = Select::from("penguin")
.filters(&vec![format!("island = Dream")])
.unwrap();
let count = block_on(rltbl.count(&select)).unwrap();
assert_eq!(count, 2);
let select = Select::from("penguin")
.filters(&vec![format!("island = Torgersen")])
.unwrap();
let count = block_on(rltbl.count(&select)).unwrap();
assert_eq!(count, 5);
}
}