use sqlx::postgres::{PgConnection, PgRow};
use sqlx::sqlite::{SqliteConnection, SqliteRow};
use sqlx::{Postgres, Row as _, SqlSafeStr as _, Sqlite};
use uuid::Uuid;
use crate::db::Database;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Dialect {
Sqlite,
Postgres,
}
impl Dialect {
#[must_use]
pub fn json_array_source(self, column: &str) -> String {
match self {
Dialect::Sqlite => format!("json_each({column}) AS ident"),
Dialect::Postgres => format!("jsonb_array_elements({column}::jsonb) AS ident"),
}
}
#[must_use]
pub fn json_member(self, member: &str) -> String {
match self {
Dialect::Sqlite => format!("json_extract(ident.value, '$.{member}')"),
Dialect::Postgres => format!("ident.value ->> '{member}'"),
}
}
#[must_use]
pub fn substring_position(self) -> &'static str {
match self {
Dialect::Sqlite => "instr",
Dialect::Postgres => "strpos",
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum Value {
Null(NullKind),
Bool(bool),
I64(i64),
Text(String),
Blob(Vec<u8>),
Uuid(Uuid),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum NullKind {
Bool,
I64,
Text,
Blob,
Uuid,
}
pub trait Bind {
fn to_value(self) -> Value;
}
macro_rules! bind {
($($ty:ty => $kind:ident, |$v:ident| $body:expr),* $(,)?) => {$(
impl Bind for $ty {
fn to_value(self) -> Value {
let $v = self;
$body
}
}
impl Bind for Option<$ty> {
fn to_value(self) -> Value {
match self {
Some($v) => { $body }
None => Value::Null(NullKind::$kind),
}
}
}
)*};
}
bind! {
bool => Bool, |v| Value::Bool(v),
i64 => I64, |v| Value::I64(v),
i32 => I64, |v| Value::I64(i64::from(v)),
u32 => I64, |v| Value::I64(i64::from(v)),
String => Text, |v| Value::Text(v),
&str => Text, |v| Value::Text(v.to_string()),
&String => Text, |v| Value::Text(v.clone()),
Vec<u8> => Blob, |v| Value::Blob(v),
&[u8] => Blob, |v| Value::Blob(v.to_vec()),
&Vec<u8> => Blob, |v| Value::Blob(v.clone()),
Uuid => Uuid, |v| Value::Uuid(v),
&Uuid => Uuid, |v| Value::Uuid(*v),
}
impl Bind for Value {
fn to_value(self) -> Value {
self
}
}
impl<T: Clone> Bind for &Option<T>
where
Option<T>: Bind,
{
fn to_value(self) -> Value {
self.clone().to_value()
}
}
#[must_use]
pub fn to_dollar_placeholders(sql: &str) -> String {
let mut out = String::with_capacity(sql.len() + 16);
let mut next = 1u32;
let mut in_literal = false;
let mut chars = sql.chars().peekable();
while let Some(c) = chars.next() {
match c {
'\'' => {
out.push(c);
if in_literal && chars.peek() == Some(&'\'') {
out.push('\'');
chars.next();
} else {
in_literal = !in_literal;
}
}
'?' if !in_literal => {
out.push('$');
out.push_str(&next.to_string());
next += 1;
}
_ => out.push(c),
}
}
out
}
pub struct Query {
sql: sqlx::SqlStr,
args: Vec<Value>,
}
#[must_use]
pub fn query(sql: impl sqlx::SqlSafeStr) -> Query {
Query {
sql: sql.into_sql_str(),
args: Vec::new(),
}
}
#[derive(Debug, Clone, Copy)]
pub struct QueryResult {
rows_affected: u64,
}
impl QueryResult {
#[must_use]
pub fn rows_affected(&self) -> u64 {
self.rows_affected
}
}
impl Query {
#[must_use]
pub fn bind(mut self, value: impl Bind) -> Self {
self.args.push(value.to_value());
self
}
#[must_use]
pub fn sql_for(&self, dialect: Dialect) -> String {
match dialect {
Dialect::Sqlite => self.sql.as_str().to_string(),
Dialect::Postgres => to_dollar_placeholders(self.sql.as_str()),
}
}
pub async fn execute<'a>(self, exec: impl Into<Exec<'a>>) -> Result<QueryResult, sqlx::Error> {
let rows_affected = match exec.into() {
Exec::SqlitePool(pool) => self.sqlite().execute(pool).await?.rows_affected(),
Exec::SqliteConn(conn) => self.sqlite().execute(conn).await?.rows_affected(),
Exec::PgPool(pool) => self.postgres().execute(pool).await?.rows_affected(),
Exec::PgConn(conn) => self.postgres().execute(conn).await?.rows_affected(),
};
Ok(QueryResult { rows_affected })
}
pub async fn fetch_one<'a>(self, exec: impl Into<Exec<'a>>) -> Result<Row, sqlx::Error> {
Ok(match exec.into() {
Exec::SqlitePool(pool) => Row::Sqlite(self.sqlite().fetch_one(pool).await?),
Exec::SqliteConn(conn) => Row::Sqlite(self.sqlite().fetch_one(conn).await?),
Exec::PgPool(pool) => Row::Postgres(self.postgres().fetch_one(pool).await?),
Exec::PgConn(conn) => Row::Postgres(self.postgres().fetch_one(conn).await?),
})
}
pub async fn fetch_optional<'a>(
self,
exec: impl Into<Exec<'a>>,
) -> Result<Option<Row>, sqlx::Error> {
Ok(match exec.into() {
Exec::SqlitePool(pool) => self.sqlite().fetch_optional(pool).await?.map(Row::Sqlite),
Exec::SqliteConn(conn) => self.sqlite().fetch_optional(conn).await?.map(Row::Sqlite),
Exec::PgPool(pool) => self
.postgres()
.fetch_optional(pool)
.await?
.map(Row::Postgres),
Exec::PgConn(conn) => self
.postgres()
.fetch_optional(conn)
.await?
.map(Row::Postgres),
})
}
pub async fn fetch_all<'a>(self, exec: impl Into<Exec<'a>>) -> Result<Vec<Row>, sqlx::Error> {
Ok(match exec.into() {
Exec::SqlitePool(pool) => self
.sqlite()
.fetch_all(pool)
.await?
.into_iter()
.map(Row::Sqlite)
.collect(),
Exec::SqliteConn(conn) => self
.sqlite()
.fetch_all(conn)
.await?
.into_iter()
.map(Row::Sqlite)
.collect(),
Exec::PgPool(pool) => self
.postgres()
.fetch_all(pool)
.await?
.into_iter()
.map(Row::Postgres)
.collect(),
Exec::PgConn(conn) => self
.postgres()
.fetch_all(conn)
.await?
.into_iter()
.map(Row::Postgres)
.collect(),
})
}
fn sqlite(self) -> sqlx::query::Query<'static, Sqlite, sqlx::sqlite::SqliteArguments> {
let mut q = sqlx::query(self.sql);
for value in self.args {
q = match value {
Value::Null(NullKind::Bool) => q.bind(None::<bool>),
Value::Null(NullKind::I64) => q.bind(None::<i64>),
Value::Null(NullKind::Text) => q.bind(None::<String>),
Value::Null(NullKind::Blob) => q.bind(None::<Vec<u8>>),
Value::Null(NullKind::Uuid) => q.bind(None::<Uuid>),
Value::Bool(v) => q.bind(v),
Value::I64(v) => q.bind(v),
Value::Text(v) => q.bind(v),
Value::Blob(v) => q.bind(v),
Value::Uuid(v) => q.bind(v),
};
}
q
}
fn postgres(self) -> sqlx::query::Query<'static, Postgres, sqlx::postgres::PgArguments> {
let mut q = sqlx::query(sqlx::AssertSqlSafe(to_dollar_placeholders(
self.sql.as_str(),
)));
for value in self.args {
q = match value {
Value::Null(NullKind::Bool) => q.bind(None::<bool>),
Value::Null(NullKind::I64) => q.bind(None::<i64>),
Value::Null(NullKind::Text) => q.bind(None::<String>),
Value::Null(NullKind::Blob) => q.bind(None::<Vec<u8>>),
Value::Null(NullKind::Uuid) => q.bind(None::<Uuid>),
Value::Bool(v) => q.bind(v),
Value::I64(v) => q.bind(v),
Value::Text(v) => q.bind(v),
Value::Blob(v) => q.bind(v),
Value::Uuid(v) => q.bind(v),
};
}
q
}
}
#[derive(Debug)]
pub enum Row {
Sqlite(SqliteRow),
Postgres(PgRow),
}
pub enum Idx<'a> {
Name(&'a str),
Position(usize),
}
impl<'a> From<&'a str> for Idx<'a> {
fn from(name: &'a str) -> Self {
Idx::Name(name)
}
}
impl From<usize> for Idx<'_> {
fn from(position: usize) -> Self {
Idx::Position(position)
}
}
pub trait Decode: Sized {
fn from_sqlite(row: &SqliteRow, idx: &Idx<'_>) -> Result<Self, sqlx::Error>;
fn from_pg(row: &PgRow, idx: &Idx<'_>) -> Result<Self, sqlx::Error>;
}
macro_rules! decode {
($($ty:ty),* $(,)?) => {$(
impl Decode for $ty {
fn from_sqlite(row: &SqliteRow, idx: &Idx<'_>) -> Result<Self, sqlx::Error> {
match idx {
Idx::Name(name) => row.try_get(*name),
Idx::Position(i) => row.try_get(*i),
}
}
fn from_pg(row: &PgRow, idx: &Idx<'_>) -> Result<Self, sqlx::Error> {
match idx {
Idx::Name(name) => row.try_get(*name),
Idx::Position(i) => row.try_get(*i),
}
}
}
impl Decode for Option<$ty> {
fn from_sqlite(row: &SqliteRow, idx: &Idx<'_>) -> Result<Self, sqlx::Error> {
match idx {
Idx::Name(name) => row.try_get(*name),
Idx::Position(i) => row.try_get(*i),
}
}
fn from_pg(row: &PgRow, idx: &Idx<'_>) -> Result<Self, sqlx::Error> {
match idx {
Idx::Name(name) => row.try_get(*name),
Idx::Position(i) => row.try_get(*i),
}
}
}
)*};
}
decode!(bool, i64, String, Vec<u8>, Uuid);
impl Row {
pub fn try_get<'i, T: Decode>(&self, idx: impl Into<Idx<'i>>) -> Result<T, sqlx::Error> {
let idx = idx.into();
match self {
Row::Sqlite(row) => T::from_sqlite(row, &idx),
Row::Postgres(row) => T::from_pg(row, &idx),
}
}
}
pub enum Exec<'a> {
SqlitePool(&'a sqlx::Pool<Sqlite>),
SqliteConn(&'a mut SqliteConnection),
PgPool(&'a sqlx::Pool<Postgres>),
PgConn(&'a mut PgConnection),
}
impl<'a> From<&'a Database> for Exec<'a> {
fn from(database: &'a Database) -> Self {
database.exec()
}
}
impl<'a> From<&'a std::sync::Arc<Database>> for Exec<'a> {
fn from(database: &'a std::sync::Arc<Database>) -> Self {
database.exec()
}
}
impl<'a> From<&'a sqlx::Pool<Sqlite>> for Exec<'a> {
fn from(pool: &'a sqlx::Pool<Sqlite>) -> Self {
Exec::SqlitePool(pool)
}
}
impl<'a> From<&'a sqlx::Pool<Postgres>> for Exec<'a> {
fn from(pool: &'a sqlx::Pool<Postgres>) -> Self {
Exec::PgPool(pool)
}
}
impl<'a> From<&'a mut SqliteConnection> for Exec<'a> {
fn from(conn: &'a mut SqliteConnection) -> Self {
Exec::SqliteConn(conn)
}
}
impl<'a> From<&'a mut PgConnection> for Exec<'a> {
fn from(conn: &'a mut PgConnection) -> Self {
Exec::PgConn(conn)
}
}
impl Exec<'_> {
#[must_use]
pub fn dialect(&self) -> Dialect {
match self {
Exec::SqlitePool(_) | Exec::SqliteConn(_) => Dialect::Sqlite,
Exec::PgPool(_) | Exec::PgConn(_) => Dialect::Postgres,
}
}
pub fn reborrow(&mut self) -> Exec<'_> {
match self {
Exec::SqlitePool(pool) => Exec::SqlitePool(pool),
Exec::SqliteConn(conn) => Exec::SqliteConn(conn),
Exec::PgPool(pool) => Exec::PgPool(pool),
Exec::PgConn(conn) => Exec::PgConn(conn),
}
}
}
#[must_use]
pub fn is_unique_violation_on(
error: &sqlx::Error,
sqlite_columns: &str,
pg_constraint: &str,
) -> bool {
let sqlx::Error::Database(db) = error else {
return false;
};
if !db.is_unique_violation() {
return false;
}
match db.constraint() {
Some(name) => name == pg_constraint,
None => db.message().contains(sqlite_columns),
}
}
#[must_use]
pub fn is_check_violation(error: &sqlx::Error) -> bool {
matches!(error, sqlx::Error::Database(db) if db.is_check_violation())
}
#[must_use]
pub fn is_foreign_key_violation(error: &sqlx::Error) -> bool {
matches!(error, sqlx::Error::Database(db) if db.is_foreign_key_violation())
}
#[must_use]
pub fn is_unique_violation(error: &sqlx::Error) -> bool {
matches!(error, sqlx::Error::Database(db) if db.is_unique_violation())
}
pub struct Builder {
dialect: Dialect,
sql: String,
args: Vec<Value>,
}
impl Builder {
#[must_use]
pub fn new(dialect: Dialect, sql: impl Into<String>) -> Self {
Builder {
dialect,
sql: sql.into(),
args: Vec::new(),
}
}
#[must_use]
pub fn dialect(&self) -> Dialect {
self.dialect
}
pub fn push(&mut self, sql: impl AsRef<str>) -> &mut Self {
self.sql.push_str(sql.as_ref());
self
}
pub fn push_bind(&mut self, value: impl Bind) -> &mut Self {
self.sql.push('?');
self.args.push(value.to_value());
self
}
pub fn separated<'b>(&'b mut self, sep: &'static str) -> Separated<'b> {
Separated {
builder: self,
sep,
first: true,
}
}
#[must_use]
pub fn sql(&self) -> &str {
&self.sql
}
#[must_use]
pub fn build(self) -> Query {
Query {
sql: sqlx::AssertSqlSafe(self.sql).into_sql_str(),
args: self.args,
}
}
}
pub struct Separated<'b> {
builder: &'b mut Builder,
sep: &'static str,
first: bool,
}
impl Separated<'_> {
pub fn push_bind(&mut self, value: impl Bind) -> &mut Self {
if !self.first {
self.builder.push(self.sep);
}
self.first = false;
self.builder.push_bind(value);
self
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_separated_run_joins_only_between_values() {
let mut builder = Builder::new(Dialect::Sqlite, "SELECT 1 WHERE id IN (");
let mut list = builder.separated(", ");
for id in [1i64, 2, 3] {
list.push_bind(id);
}
builder.push(")");
assert_eq!(builder.sql(), "SELECT 1 WHERE id IN (?, ?, ?)");
}
#[test]
fn markers_are_numbered_in_order() {
assert_eq!(
to_dollar_placeholders("SELECT a FROM t WHERE b = ? AND c = ?;"),
"SELECT a FROM t WHERE b = $1 AND c = $2;"
);
}
#[test]
fn nothing_to_rewrite_is_returned_unchanged() {
let sql = "SELECT COUNT(*) FROM nonces;";
assert_eq!(to_dollar_placeholders(sql), sql);
}
#[test]
fn numbering_is_contiguous_from_one() {
let rewritten = to_dollar_placeholders("? ? ? ? ? ? ? ? ? ? ?");
let numbers: Vec<u32> = rewritten
.split_whitespace()
.map(|marker| marker.trim_start_matches('$').parse().expect("a number"))
.collect();
assert_eq!(numbers, (1..=11).collect::<Vec<u32>>());
}
#[test]
fn a_marker_inside_a_string_literal_is_left_alone() {
assert_eq!(
to_dollar_placeholders("UPDATE t SET a = 'what?' WHERE b = ?;"),
"UPDATE t SET a = 'what?' WHERE b = $1;"
);
}
#[test]
fn an_escaped_quote_does_not_end_a_literal() {
assert_eq!(
to_dollar_placeholders("SELECT 'it''s ? fine' WHERE a = ?;"),
"SELECT 'it''s ? fine' WHERE a = $1;"
);
}
#[test]
fn a_builder_pushes_markers_and_values_together() {
let mut builder = Builder::new(Dialect::Sqlite, "SELECT 1 FROM t");
builder.push(" WHERE a = ").push_bind("x");
builder.push(" AND b = ").push_bind(7i64);
assert_eq!(builder.sql(), "SELECT 1 FROM t WHERE a = ? AND b = ?");
let query = builder.build();
assert_eq!(
query.sql_for(Dialect::Postgres),
"SELECT 1 FROM t WHERE a = $1 AND b = $2"
);
assert_eq!(
query.args,
vec![Value::Text("x".to_string()), Value::I64(7)]
);
}
#[test]
fn an_absent_optional_binds_null_at_its_own_type() {
assert_eq!(
super::query("SELECT ?").bind(None::<String>).args,
vec![Value::Null(NullKind::Text)]
);
assert_eq!(
super::query("SELECT ?").bind(None::<Uuid>).args,
vec![Value::Null(NullKind::Uuid)]
);
assert_eq!(
super::query("SELECT ?").bind(None::<bool>).args,
vec![Value::Null(NullKind::Bool)]
);
}
#[test]
fn a_present_optional_binds_its_value() {
let query = super::query("SELECT ?").bind(Some(3i64));
assert_eq!(query.args, vec![Value::I64(3)]);
}
}