use sqlx::any::{AnyArguments, AnyRow};
use sqlx::{Any, Encode, Row, Type};
use crate::migration::DbBackend;
use crate::Db;
pub trait AnyRowExt {
fn get_bool(&self, column: &str) -> Result<bool, sqlx::Error>;
fn get_text(&self, column: &str) -> Result<String, sqlx::Error>;
fn get_text_opt(&self, column: &str) -> Result<Option<String>, sqlx::Error>;
}
impl AnyRowExt for AnyRow {
fn get_bool(&self, column: &str) -> Result<bool, sqlx::Error> {
Ok(self.try_get::<i32, _>(column)? != 0)
}
fn get_text(&self, column: &str) -> Result<String, sqlx::Error> {
match self.try_get::<String, _>(column) {
Ok(s) => Ok(s),
Err(_) => decode_utf8(column, self.try_get::<Vec<u8>, _>(column)?),
}
}
fn get_text_opt(&self, column: &str) -> Result<Option<String>, sqlx::Error> {
match self.try_get::<Option<String>, _>(column) {
Ok(s) => Ok(s),
Err(_) => self
.try_get::<Option<Vec<u8>>, _>(column)?
.map(|b| decode_utf8(column, b))
.transpose(),
}
}
}
fn decode_utf8(column: &str, bytes: Vec<u8>) -> Result<String, sqlx::Error> {
String::from_utf8(bytes).map_err(|e| sqlx::Error::ColumnDecode {
index: column.to_string(),
source: Box::new(e),
})
}
pub fn on_conflict_ignore<C>(keys: impl IntoIterator<Item = C>) -> sea_query::OnConflict
where
C: sea_query::IntoIden,
{
use sea_query::IntoIden;
let keys: Vec<sea_query::DynIden> = keys.into_iter().map(IntoIden::into_iden).collect();
let first = keys[0].clone();
sea_query::OnConflict::columns(keys)
.update_column(first)
.to_owned()
}
pub fn text_cast(backend: DbBackend) -> &'static str {
match backend {
DbBackend::Mysql => "char",
DbBackend::Postgres | DbBackend::Sqlite => "text",
}
}
pub fn build<S>(backend: DbBackend, stmt: S) -> (String, sea_query::Values)
where
S: sea_query::QueryStatementWriter,
{
match backend {
DbBackend::Postgres => stmt.build(sea_query::PostgresQueryBuilder),
DbBackend::Mysql => stmt.build(sea_query::MysqlQueryBuilder),
DbBackend::Sqlite => stmt.build(sea_query::SqliteQueryBuilder),
}
}
pub fn insert_returning_id<I>(
db: &Db,
stmt: sea_query::InsertStatement,
id: I,
) -> impl std::future::Future<Output = Result<i64, sqlx::Error>> + Send + '_
where
I: sea_query::IntoIden + 'static,
{
let (sql, values, returning) = render_insert(db.backend, stmt, id);
async move {
if returning {
bind_values(sqlx::query(&sql), values)
.fetch_one(&db.pool)
.await?
.try_get::<i64, _>(0)
} else {
bind_values(sqlx::query(&sql), values)
.execute(&db.pool)
.await?
.last_insert_id()
.ok_or(sqlx::Error::RowNotFound)
}
}
}
fn render_insert<I>(
backend: DbBackend,
mut stmt: sea_query::InsertStatement,
id: I,
) -> (String, sea_query::Values, bool)
where
I: sea_query::IntoIden + 'static,
{
let returning = matches!(backend, DbBackend::Postgres | DbBackend::Sqlite);
if returning {
stmt.returning_col(id);
}
let (sql, values) = build(backend, stmt);
(sql, values, returning)
}
type AnyQuery<'q> = sqlx::query::Query<'q, Any, AnyArguments<'q>>;
type AnyQueryAs<'q, O> = sqlx::query::QueryAs<'q, Any, O, AnyArguments<'q>>;
fn bind_one<'q, T>(query: AnyQuery<'q>, value: T) -> AnyQuery<'q>
where
T: 'q + Send + Encode<'q, Any> + Type<Any>,
{
query.bind(value)
}
pub fn bind_values(mut query: AnyQuery<'_>, values: sea_query::Values) -> AnyQuery<'_> {
use sea_query::Value;
for value in values.0 {
query = match value {
Value::Bool(v) => bind_one(query, v.map(i32::from)),
Value::TinyInt(v) => bind_one(query, v.map(i32::from)),
Value::SmallInt(v) => bind_one(query, v),
Value::Int(v) => bind_one(query, v),
Value::BigInt(v) => bind_one(query, v),
Value::TinyUnsigned(v) => bind_one(query, v.map(i32::from)),
Value::SmallUnsigned(v) => bind_one(query, v.map(i32::from)),
Value::Unsigned(v) => bind_one(query, v.map(i64::from)),
Value::BigUnsigned(v) => bind_one(query, v.map(|n| n as i64)),
Value::Float(v) => bind_one(query, v),
Value::Double(v) => bind_one(query, v),
Value::String(v) => bind_one(query, v.map(|b| *b)),
Value::Char(v) => bind_one(query, v.map(|c| c.to_string())),
Value::Bytes(v) => bind_one(query, v.map(|b| *b)),
#[allow(unreachable_patterns)]
other => panic!("unsupported portable bind value: {other:?}"),
};
}
query
}
pub fn bind_values_as<O>(
mut query: AnyQueryAs<'_, O>,
values: sea_query::Values,
) -> AnyQueryAs<'_, O> {
use sea_query::Value;
for value in values.0 {
query = match value {
Value::Bool(v) => query.bind(v.map(i32::from)),
Value::TinyInt(v) => query.bind(v.map(i32::from)),
Value::SmallInt(v) => query.bind(v),
Value::Int(v) => query.bind(v),
Value::BigInt(v) => query.bind(v),
Value::TinyUnsigned(v) => query.bind(v.map(i32::from)),
Value::SmallUnsigned(v) => query.bind(v.map(i32::from)),
Value::Unsigned(v) => query.bind(v.map(i64::from)),
Value::BigUnsigned(v) => query.bind(v.map(|n| n as i64)),
Value::Float(v) => query.bind(v),
Value::Double(v) => query.bind(v),
Value::String(v) => query.bind(v.map(|b| *b)),
Value::Char(v) => query.bind(v.map(|c| c.to_string())),
Value::Bytes(v) => query.bind(v.map(|b| *b)),
#[allow(unreachable_patterns)]
other => panic!("unsupported portable bind value: {other:?}"),
};
}
query
}
#[cfg(test)]
mod tests {
use super::*;
use sea_query::{Alias, Expr, Iden, Query};
use sqlx::AnyPool;
#[derive(Iden)]
enum Widget {
Table,
Id,
Label,
Qty,
}
async fn pool() -> AnyPool {
sqlx::any::install_default_drivers();
let pool = sqlx::any::AnyPoolOptions::new()
.max_connections(1)
.connect("sqlite::memory:")
.await
.unwrap();
sqlx::raw_sql(
"create table widget (id text primary key, label text not null, qty integer not null)",
)
.execute(&pool)
.await
.unwrap();
pool
}
#[tokio::test]
async fn binds_parameters_on_insert_and_select() {
let pool = pool().await;
let backend = DbBackend::Sqlite;
let insert = Query::insert()
.into_table(Widget::Table)
.columns([Widget::Id, Widget::Label, Widget::Qty])
.values_panic(["w-1".into(), "Sprocket".into(), 7.into()])
.to_owned();
let (sql, values) = build(backend, insert);
bind_values(sqlx::query(&sql), values)
.execute(&pool)
.await
.unwrap();
let select = Query::select()
.column(Widget::Label)
.from(Widget::Table)
.and_where(Expr::col(Widget::Id).eq("w-1"))
.to_owned();
let (sql, values) = build(backend, select);
let label: String = bind_values_as(sqlx::query_as::<_, (String,)>(&sql), values)
.fetch_one(&pool)
.await
.unwrap()
.0;
assert_eq!(label, "Sprocket");
let count_stmt = Query::select()
.expr(Expr::col(Widget::Id).count())
.from(Widget::Table)
.and_where(Expr::col(Alias::new("qty")).eq(7))
.to_owned();
let (sql, values) = build(backend, count_stmt);
let count: i64 = bind_values_as(sqlx::query_as::<_, (i64,)>(&sql), values)
.fetch_one(&pool)
.await
.unwrap()
.0;
assert_eq!(count, 1);
}
}