use async_trait::async_trait;
use futures::StreamExt;
use sqlx::{Column, Executor, Row, TypeInfo};
use std::collections::HashMap;
use std::future::Future;
use std::pin::Pin;
use std::str::FromStr;
use std::sync::Arc;
use sz_orm_core::{ColType, Connection, ConnectionFactory, DbError, QueryRows, QueryValues, Value};
use crate::error::map_sqlx_error;
fn needs_raw_sql(sql: &str) -> bool {
let trimmed = sql.trim_start();
let upper = trimmed.to_uppercase();
upper.starts_with("BEGIN")
|| upper.starts_with("COMMIT")
|| upper.starts_with("ROLLBACK")
|| upper.starts_with("SAVEPOINT")
|| upper.starts_with("RELEASE")
|| upper.starts_with("SET ")
|| upper.starts_with("USE ")
|| upper.starts_with("START TRANSACTION")
}
fn row_to_value_with_coltype_sqlite(
row: &sqlx::sqlite::SqliteRow,
ordinal: usize,
col_type: ColType,
) -> Value {
match col_type {
ColType::Bool => match row.try_get::<Option<bool>, usize>(ordinal) {
Ok(v) => v.map(Value::Bool).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::I8 => match row.try_get::<Option<i8>, usize>(ordinal) {
Ok(v) => v.map(Value::I8).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::I16 => match row.try_get::<Option<i16>, usize>(ordinal) {
Ok(v) => v.map(Value::I16).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::I32 => match row.try_get::<Option<i32>, usize>(ordinal) {
Ok(v) => v.map(Value::I32).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::I64 => match row.try_get::<Option<i64>, usize>(ordinal) {
Ok(v) => v.map(Value::I64).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::U8 => match row.try_get::<Option<u8>, usize>(ordinal) {
Ok(v) => v.map(Value::U8).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::U16 => match row.try_get::<Option<u16>, usize>(ordinal) {
Ok(v) => v.map(Value::U16).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::U32 => match row.try_get::<Option<u32>, usize>(ordinal) {
Ok(v) => v.map(Value::U32).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::U64 => match row.try_get::<Option<i64>, usize>(ordinal) {
Ok(v) => v.map(Value::I64).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::F32 => match row.try_get::<Option<f32>, usize>(ordinal) {
Ok(v) => v.map(Value::F32).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::F64 => match row.try_get::<Option<f64>, usize>(ordinal) {
Ok(v) => v.map(Value::F64).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::Decimal => match row.try_get::<Option<String>, usize>(ordinal) {
Ok(v) => v.map(Value::Decimal).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::String => match row.try_get::<Option<String>, usize>(ordinal) {
Ok(v) => v.map(Value::String).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::Bytes => match row.try_get::<Option<Vec<u8>>, usize>(ordinal) {
Ok(v) => v.map(Value::Bytes).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::Date | ColType::DateTime | ColType::Time | ColType::Json | ColType::Uuid => {
match row.try_get::<Option<String>, usize>(ordinal) {
Ok(v) => v.map(Value::String).unwrap_or(Value::Null),
Err(_) => Value::Null,
}
}
ColType::Unknown => {
if let Ok(v) = row.try_get::<Option<bool>, usize>(ordinal) {
return v.map(Value::Bool).unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<Option<i64>, usize>(ordinal) {
return v.map(Value::I64).unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<Option<f64>, usize>(ordinal) {
return v.map(Value::F64).unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<Option<String>, usize>(ordinal) {
return v.map(Value::String).unwrap_or(Value::Null);
}
Value::Null
}
_ => Value::Null,
}
}
pub struct SqlitePoolHandle {
pool: sqlx::SqlitePool,
}
impl SqlitePoolHandle {
pub async fn connect(url: &str) -> Result<Self, DbError> {
let opts = sqlx::sqlite::SqliteConnectOptions::from_str(url)
.map_err(map_sqlx_error)?
.journal_mode(sqlx::sqlite::SqliteJournalMode::Wal)
.synchronous(sqlx::sqlite::SqliteSynchronous::Normal)
.busy_timeout(std::time::Duration::from_secs(5))
.pragma("mmap_size", "268435456");
let pool = sqlx::sqlite::SqlitePoolOptions::new()
.max_connections(10)
.acquire_timeout(std::time::Duration::from_secs(30))
.idle_timeout(Some(std::time::Duration::from_secs(600)))
.max_lifetime(Some(std::time::Duration::from_secs(1800)))
.connect_with(opts)
.await
.map_err(map_sqlx_error)?;
Ok(Self { pool })
}
pub fn from_pool(pool: sqlx::SqlitePool) -> Self {
Self { pool }
}
pub fn pool(&self) -> &sqlx::SqlitePool {
&self.pool
}
}
pub struct SqlxSqliteConnectionFactory {
pool: Arc<SqlitePoolHandle>,
}
impl SqlxSqliteConnectionFactory {
pub fn new(pool: Arc<SqlitePoolHandle>) -> Self {
Self { pool }
}
}
#[async_trait]
impl ConnectionFactory for SqlxSqliteConnectionFactory {
async fn create(&self) -> Result<Box<dyn Connection>, DbError> {
let conn = self.pool.pool.acquire().await.map_err(map_sqlx_error)?;
Ok(Box::new(SqlxSqliteConnection {
conn: Some(conn),
connected: true,
in_transaction: false,
}))
}
}
pub struct SqlxSqliteConnection {
conn: Option<sqlx::pool::PoolConnection<sqlx::Sqlite>>,
connected: bool,
in_transaction: bool,
}
impl Connection for SqlxSqliteConnection {
fn execute<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<u64, DbError>> + Send + 'a>> {
Box::pin(async move {
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let result = if needs_raw_sql(sql) {
(&mut *pool_conn)
.execute(sqlx::raw_sql(sqlx::AssertSqlSafe(sql)))
.await
} else {
(&mut *pool_conn).execute(sqlx::AssertSqlSafe(sql)).await
};
self.conn = Some(pool_conn);
match result {
Ok(r) => Ok(r.rows_affected()),
Err(e) => {
let db_err = map_sqlx_error(e);
if matches!(db_err, DbError::ConnectionError(_) | DbError::IoError(_)) {
self.connected = false;
}
Err(db_err)
}
}
})
}
fn query<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<Vec<HashMap<String, Value>>, DbError>> + Send + 'a>>
{
Box::pin(async move {
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let rows_result = (&mut *pool_conn).fetch_all(sqlx::AssertSqlSafe(sql)).await;
self.conn = Some(pool_conn);
let rows = rows_result.map_err(map_sqlx_error)?;
if rows.is_empty() {
return Ok(Vec::new());
}
let col_types: Vec<ColType> = rows[0]
.columns()
.iter()
.map(|col| ColType::parse_sqlite(col.type_info().name()))
.collect();
let mut result = Vec::with_capacity(rows.len());
for row in &rows {
let mut record = HashMap::with_capacity(col_types.len());
for (i, col) in row.columns().iter().enumerate() {
let name = col.name().to_string();
let value = row_to_value_with_coltype_sqlite(row, i, col_types[i]);
record.insert(name, value);
}
result.push(record);
}
Ok(result)
})
}
fn begin_transaction<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
if self.in_transaction {
return Err(DbError::Internal("transaction already started".to_string()));
}
self.execute("BEGIN").await?;
self.in_transaction = true;
Ok(())
})
}
fn commit<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
if self.in_transaction {
self.execute("COMMIT").await?;
self.in_transaction = false;
}
Ok(())
})
}
fn rollback<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
if self.in_transaction {
let result = self.execute("ROLLBACK").await;
self.in_transaction = false;
result.map(|_| ())
} else {
Ok(())
}
})
}
fn is_connected(&self) -> bool {
self.connected
}
fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
Box::pin(async move {
match self.execute("SELECT 1").await {
Ok(_) => true,
Err(_) => {
self.connected = false;
false
}
}
})
}
fn close<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
if let Some(conn) = self.conn.take() {
drop(conn);
}
self.connected = false;
self.in_transaction = false;
Ok(())
})
}
fn execute_with_params<'a>(
&'a mut self,
sql: &'a str,
params: &'a [Value],
) -> Pin<Box<dyn Future<Output = Result<u64, DbError>> + Send + 'a>> {
Box::pin(async move {
if needs_raw_sql(sql) || params.is_empty() {
return self.execute(sql).await;
}
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let mut q = sqlx::query(sqlx::AssertSqlSafe(sql));
for v in params {
q = match v {
Value::Null => q.bind(None::<i64>),
Value::Bool(b) => q.bind(*b),
Value::I8(n) => q.bind(*n),
Value::I16(n) => q.bind(*n),
Value::I32(n) => q.bind(*n),
Value::I64(n) => q.bind(*n),
Value::U8(n) => q.bind(*n),
Value::U16(n) => q.bind(*n),
Value::U32(n) => q.bind(*n),
Value::U64(n) => q.bind(*n as i64),
Value::F32(f) => q.bind(*f),
Value::F64(f) => q.bind(*f),
Value::String(s) => q.bind(s.as_str()),
Value::Decimal(s) => q.bind(s.as_str()),
Value::Bytes(b) => q.bind(b.as_slice()),
other => q.bind(other.to_string()),
};
}
let result = q.execute(&mut *pool_conn).await;
self.conn = Some(pool_conn);
match result {
Ok(r) => Ok(r.rows_affected()),
Err(e) => {
let db_err = map_sqlx_error(e);
if matches!(db_err, DbError::ConnectionError(_) | DbError::IoError(_)) {
self.connected = false;
}
Err(db_err)
}
}
})
}
fn query_with_params<'a>(
&'a mut self,
sql: &'a str,
params: &'a [Value],
) -> Pin<Box<dyn Future<Output = Result<QueryRows, DbError>> + Send + 'a>> {
Box::pin(async move {
if params.is_empty() {
return self.query(sql).await;
}
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let mut q = sqlx::query(sqlx::AssertSqlSafe(sql));
for v in params {
q = match v {
Value::Null => q.bind(None::<i64>),
Value::Bool(b) => q.bind(*b),
Value::I8(n) => q.bind(*n),
Value::I16(n) => q.bind(*n),
Value::I32(n) => q.bind(*n),
Value::I64(n) => q.bind(*n),
Value::U8(n) => q.bind(*n),
Value::U16(n) => q.bind(*n),
Value::U32(n) => q.bind(*n),
Value::U64(n) => q.bind(*n as i64),
Value::F32(f) => q.bind(*f),
Value::F64(f) => q.bind(*f),
Value::String(s) => q.bind(s.as_str()),
Value::Decimal(s) => q.bind(s.as_str()),
Value::Bytes(b) => q.bind(b.as_slice()),
other => q.bind(other.to_string()),
};
}
let rows_result = q.fetch_all(&mut *pool_conn).await;
self.conn = Some(pool_conn);
let rows = rows_result.map_err(map_sqlx_error)?;
if rows.is_empty() {
return Ok(Vec::new());
}
let col_types: Vec<ColType> = rows[0]
.columns()
.iter()
.map(|col| ColType::parse_sqlite(col.type_info().name()))
.collect();
let mut result = Vec::with_capacity(rows.len());
for row in &rows {
let mut record = HashMap::with_capacity(col_types.len());
for (i, col) in row.columns().iter().enumerate() {
let name = col.name().to_string();
let value = row_to_value_with_coltype_sqlite(row, i, col_types[i]);
record.insert(name, value);
}
result.push(record);
}
Ok(result)
})
}
fn query_values<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<QueryValues, DbError>> + Send + 'a>> {
Box::pin(async move {
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let rows_result = (&mut *pool_conn).fetch_all(sqlx::AssertSqlSafe(sql)).await;
self.conn = Some(pool_conn);
let rows = rows_result.map_err(map_sqlx_error)?;
if rows.is_empty() {
return Ok((Vec::new(), Vec::new()));
}
let cols = rows[0].columns();
let mut col_names: Vec<String> = Vec::with_capacity(cols.len());
let mut col_types: Vec<ColType> = Vec::with_capacity(cols.len());
for col in cols {
col_names.push(col.name().to_string());
col_types.push(ColType::parse_sqlite(col.type_info().name()));
}
let mut result_rows: Vec<Vec<Value>> = Vec::with_capacity(rows.len());
for row in &rows {
let mut row_values: Vec<Value> = Vec::with_capacity(col_names.len());
for (idx, _) in col_names.iter().enumerate() {
row_values.push(row_to_value_with_coltype_sqlite(row, idx, col_types[idx]));
}
result_rows.push(row_values);
}
Ok((col_names, result_rows))
})
}
fn query_values_with_params<'a>(
&'a mut self,
sql: &'a str,
params: &'a [Value],
) -> Pin<Box<dyn Future<Output = Result<QueryValues, DbError>> + Send + 'a>> {
Box::pin(async move {
if params.is_empty() {
return self.query_values(sql).await;
}
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let mut q = sqlx::query(sqlx::AssertSqlSafe(sql));
for v in params {
q = match v {
Value::Null => q.bind(None::<i64>),
Value::Bool(b) => q.bind(*b),
Value::I8(n) => q.bind(*n),
Value::I16(n) => q.bind(*n),
Value::I32(n) => q.bind(*n),
Value::I64(n) => q.bind(*n),
Value::U8(n) => q.bind(*n),
Value::U16(n) => q.bind(*n),
Value::U32(n) => q.bind(*n),
Value::U64(n) => q.bind(*n as i64),
Value::F32(f) => q.bind(*f),
Value::F64(f) => q.bind(*f),
Value::String(s) => q.bind(s.as_str()),
Value::Decimal(s) => q.bind(s.as_str()),
Value::Bytes(b) => q.bind(b.as_slice()),
other => q.bind(other.to_string()),
};
}
let rows_result = q.fetch_all(&mut *pool_conn).await;
self.conn = Some(pool_conn);
let rows = rows_result.map_err(map_sqlx_error)?;
if rows.is_empty() {
return Ok((Vec::new(), Vec::new()));
}
let cols = rows[0].columns();
let mut col_names: Vec<String> = Vec::with_capacity(cols.len());
let mut col_types: Vec<ColType> = Vec::with_capacity(cols.len());
for col in cols {
col_names.push(col.name().to_string());
col_types.push(ColType::parse_sqlite(col.type_info().name()));
}
let mut result_rows: Vec<Vec<Value>> = Vec::with_capacity(rows.len());
for row in &rows {
let mut row_values: Vec<Value> = Vec::with_capacity(col_names.len());
for (idx, _) in col_names.iter().enumerate() {
row_values.push(row_to_value_with_coltype_sqlite(row, idx, col_types[idx]));
}
result_rows.push(row_values);
}
Ok((col_names, result_rows))
})
}
fn query_stream<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn futures::Stream<Item = Result<HashMap<String, Value>, DbError>> + Send + 'a>>
{
Box::pin(async_stream::try_stream! {
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let mut row_stream = sqlx::query(sqlx::AssertSqlSafe(sql)).fetch(&mut *pool_conn);
let mut col_types: Vec<ColType> = Vec::new();
let mut col_names: Vec<String> = Vec::new();
let mut first_row = true;
while let Some(row_result) = row_stream.next().await {
let row = row_result.map_err(map_sqlx_error)?;
if first_row {
for col in row.columns() {
col_names.push(col.name().to_string());
col_types.push(ColType::parse_sqlite(col.type_info().name()));
}
first_row = false;
}
let mut record = HashMap::with_capacity(col_names.len());
for (i, name) in col_names.iter().enumerate() {
let value = row_to_value_with_coltype_sqlite(&row, i, col_types[i]);
record.insert(name.clone(), value);
}
yield record;
}
drop(row_stream);
self.conn = Some(pool_conn);
})
}
}
impl Drop for SqlxSqliteConnection {
fn drop(&mut self) {
if let Some(conn) = self.conn.take() {
drop(conn);
}
}
}
pub async fn sqlite_backup(
conn: &mut SqlxSqliteConnection,
dest_path: &str,
) -> Result<(), DbError> {
let escaped_path = dest_path.replace('\'', "''");
let sql = format!("VACUUM INTO '{}'", escaped_path);
conn.execute(&sql).await?;
Ok(())
}
fn row_to_value_mysql(row: &sqlx::mysql::MySqlRow, ordinal: usize) -> Value {
use sqlx::TypeInfo;
let type_name = row.columns()[ordinal].type_info().name();
match type_name {
"BOOLEAN" => match row.try_get::<Option<bool>, usize>(ordinal) {
Ok(v) => v.map(Value::Bool).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
"TINYINT" | "TINYINT UNSIGNED" => match row.try_get::<Option<i8>, usize>(ordinal) {
Ok(v) => v.map(Value::I8).unwrap_or(Value::Null),
Err(_) => match row.try_get::<Option<u8>, usize>(ordinal) {
Ok(v) => v.map(Value::U8).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
},
"SMALLINT" | "SMALLINT UNSIGNED" => match row.try_get::<Option<i16>, usize>(ordinal) {
Ok(v) => v.map(Value::I16).unwrap_or(Value::Null),
Err(_) => match row.try_get::<Option<u16>, usize>(ordinal) {
Ok(v) => v.map(Value::U16).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
},
"INT" | "INT UNSIGNED" | "MEDIUMINT" | "MEDIUMINT UNSIGNED" => {
match row.try_get::<Option<i32>, usize>(ordinal) {
Ok(v) => v.map(Value::I32).unwrap_or(Value::Null),
Err(_) => match row.try_get::<Option<u32>, usize>(ordinal) {
Ok(v) => v.map(Value::U32).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
}
}
"BIGINT" | "BIGINT UNSIGNED" => match row.try_get::<Option<i64>, usize>(ordinal) {
Ok(v) => v.map(Value::I64).unwrap_or(Value::Null),
Err(_) => match row.try_get::<Option<u64>, usize>(ordinal) {
Ok(v) => v.map(Value::U64).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
},
"FLOAT" => match row.try_get::<Option<f32>, usize>(ordinal) {
Ok(v) => v.map(Value::F32).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
"DOUBLE" => match row.try_get::<Option<f64>, usize>(ordinal) {
Ok(v) => v.map(Value::F64).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
"VARCHAR" | "TEXT" | "CHAR" | "TINYTEXT" | "MEDIUMTEXT" | "LONGTEXT" | "ENUM" => {
match row.try_get::<Option<String>, usize>(ordinal) {
Ok(v) => v.map(Value::String).unwrap_or(Value::Null),
Err(_) => Value::Null,
}
}
"BLOB" | "TINYBLOB" | "MEDIUMBLOB" | "LONGBLOB" | "BINARY" | "VARBINARY" => {
match row.try_get::<Option<Vec<u8>>, usize>(ordinal) {
Ok(v) => v.map(Value::Bytes).unwrap_or(Value::Null),
Err(_) => Value::Null,
}
}
"DECIMAL" | "NUMERIC" | "NEWDECIMAL" => {
match row.try_get::<Option<rust_decimal::Decimal>, usize>(ordinal) {
Ok(Some(v)) => Value::Decimal(v.to_string()),
Ok(None) => Value::Null,
Err(_) => match row.try_get::<Option<String>, usize>(ordinal) {
Ok(v) => v.map(Value::String).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
}
}
_ => {
if let Ok(v) = row.try_get::<Option<i64>, usize>(ordinal) {
return v.map(Value::I64).unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<Option<f64>, usize>(ordinal) {
return v.map(Value::F64).unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<Option<bool>, usize>(ordinal) {
return v.map(Value::Bool).unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<Option<String>, usize>(ordinal) {
return v.map(Value::String).unwrap_or(Value::Null);
}
Value::Null
}
}
}
fn row_to_value_with_coltype_mysql(
row: &sqlx::mysql::MySqlRow,
ordinal: usize,
col_type: ColType,
) -> Value {
match col_type {
ColType::Bool => match row.try_get::<Option<bool>, usize>(ordinal) {
Ok(v) => v.map(Value::Bool).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::I8 => match row.try_get::<Option<i8>, usize>(ordinal) {
Ok(v) => v.map(Value::I8).unwrap_or(Value::Null),
Err(_) => match row.try_get::<Option<u8>, usize>(ordinal) {
Ok(v) => v.map(Value::U8).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
},
ColType::I16 => match row.try_get::<Option<i16>, usize>(ordinal) {
Ok(v) => v.map(Value::I16).unwrap_or(Value::Null),
Err(_) => match row.try_get::<Option<u16>, usize>(ordinal) {
Ok(v) => v.map(Value::U16).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
},
ColType::I32 => match row.try_get::<Option<i32>, usize>(ordinal) {
Ok(v) => v.map(Value::I32).unwrap_or(Value::Null),
Err(_) => match row.try_get::<Option<u32>, usize>(ordinal) {
Ok(v) => v.map(Value::U32).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
},
ColType::I64 => match row.try_get::<Option<i64>, usize>(ordinal) {
Ok(v) => v.map(Value::I64).unwrap_or(Value::Null),
Err(_) => match row.try_get::<Option<u64>, usize>(ordinal) {
Ok(v) => v.map(Value::U64).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
},
ColType::U8 => match row.try_get::<Option<u8>, usize>(ordinal) {
Ok(v) => v.map(Value::U8).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::U16 => match row.try_get::<Option<u16>, usize>(ordinal) {
Ok(v) => v.map(Value::U16).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::U32 => match row.try_get::<Option<u32>, usize>(ordinal) {
Ok(v) => v.map(Value::U32).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::U64 => match row.try_get::<Option<u64>, usize>(ordinal) {
Ok(v) => v.map(Value::U64).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::F32 => match row.try_get::<Option<f32>, usize>(ordinal) {
Ok(v) => v.map(Value::F32).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::F64 => match row.try_get::<Option<f64>, usize>(ordinal) {
Ok(v) => v.map(Value::F64).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::Decimal => match row.try_get::<Option<rust_decimal::Decimal>, usize>(ordinal) {
Ok(Some(v)) => Value::Decimal(v.to_string()),
Ok(None) => Value::Null,
Err(_) => match row.try_get::<Option<String>, usize>(ordinal) {
Ok(v) => v.map(Value::String).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
},
ColType::String => match row.try_get::<Option<String>, usize>(ordinal) {
Ok(v) => v.map(Value::String).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::Bytes => match row.try_get::<Option<Vec<u8>>, usize>(ordinal) {
Ok(v) => v.map(Value::Bytes).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::Date | ColType::DateTime | ColType::Time | ColType::Json | ColType::Uuid => {
match row.try_get::<Option<String>, usize>(ordinal) {
Ok(v) => v.map(Value::String).unwrap_or(Value::Null),
Err(_) => Value::Null,
}
}
ColType::Unknown => {
if let Ok(v) = row.try_get::<Option<i64>, usize>(ordinal) {
return v.map(Value::I64).unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<Option<f64>, usize>(ordinal) {
return v.map(Value::F64).unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<Option<bool>, usize>(ordinal) {
return v.map(Value::Bool).unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<Option<String>, usize>(ordinal) {
return v.map(Value::String).unwrap_or(Value::Null);
}
Value::Null
}
_ => Value::Null,
}
}
pub struct MySqlPoolHandle {
pool: sqlx::MySqlPool,
}
impl MySqlPoolHandle {
pub async fn connect(url: &str) -> Result<Self, DbError> {
let pool = sqlx::pool::PoolOptions::<sqlx::MySql>::new()
.max_connections(10)
.acquire_timeout(std::time::Duration::from_secs(30))
.idle_timeout(Some(std::time::Duration::from_secs(600)))
.max_lifetime(Some(std::time::Duration::from_secs(1800)))
.connect(url)
.await
.map_err(map_sqlx_error)?;
Ok(Self { pool })
}
pub fn from_pool(pool: sqlx::MySqlPool) -> Self {
Self { pool }
}
pub fn pool(&self) -> &sqlx::MySqlPool {
&self.pool
}
}
pub struct SqlxMySqlConnectionFactory {
pool: Arc<MySqlPoolHandle>,
}
impl SqlxMySqlConnectionFactory {
pub fn new(pool: Arc<MySqlPoolHandle>) -> Self {
Self { pool }
}
}
#[async_trait]
impl ConnectionFactory for SqlxMySqlConnectionFactory {
async fn create(&self) -> Result<Box<dyn Connection>, DbError> {
let conn = self.pool.pool.acquire().await.map_err(map_sqlx_error)?;
Ok(Box::new(SqlxMySqlConnection {
conn: Some(conn),
connected: true,
in_transaction: false,
}))
}
}
pub struct SqlxMySqlConnection {
conn: Option<sqlx::pool::PoolConnection<sqlx::MySql>>,
connected: bool,
in_transaction: bool,
}
impl Connection for SqlxMySqlConnection {
fn execute<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<u64, DbError>> + Send + 'a>> {
Box::pin(async move {
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let result = if needs_raw_sql(sql) {
(&mut *pool_conn)
.execute(sqlx::raw_sql(sqlx::AssertSqlSafe(sql)))
.await
} else {
(&mut *pool_conn).execute(sqlx::AssertSqlSafe(sql)).await
};
self.conn = Some(pool_conn);
match result {
Ok(r) => Ok(r.rows_affected()),
Err(e) => {
let db_err = map_sqlx_error(e);
if matches!(db_err, DbError::ConnectionError(_) | DbError::IoError(_)) {
self.connected = false;
}
Err(db_err)
}
}
})
}
fn query<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<Vec<HashMap<String, Value>>, DbError>> + Send + 'a>>
{
Box::pin(async move {
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let rows_result = (&mut *pool_conn).fetch_all(sqlx::AssertSqlSafe(sql)).await;
self.conn = Some(pool_conn);
let rows = rows_result.map_err(map_sqlx_error)?;
if rows.is_empty() {
return Ok(Vec::new());
}
let col_types: Vec<ColType> = rows[0]
.columns()
.iter()
.map(|col| ColType::parse_mysql(col.type_info().name()))
.collect();
let mut result = Vec::with_capacity(rows.len());
for row in &rows {
let mut record = HashMap::with_capacity(col_types.len());
for (i, col) in row.columns().iter().enumerate() {
let name = col.name().to_string();
let value = row_to_value_with_coltype_mysql(row, i, col_types[i]);
record.insert(name, value);
}
result.push(record);
}
Ok(result)
})
}
fn begin_transaction<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
if self.in_transaction {
return Err(DbError::Internal("transaction already started".to_string()));
}
self.execute("BEGIN").await?;
self.in_transaction = true;
Ok(())
})
}
fn commit<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
if self.in_transaction {
self.execute("COMMIT").await?;
self.in_transaction = false;
}
Ok(())
})
}
fn rollback<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
if self.in_transaction {
let result = self.execute("ROLLBACK").await;
self.in_transaction = false;
result.map(|_| ())
} else {
Ok(())
}
})
}
fn is_connected(&self) -> bool {
self.connected
}
fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
Box::pin(async move {
match self.execute("SELECT 1").await {
Ok(_) => true,
Err(_) => {
self.connected = false;
false
}
}
})
}
fn close<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
if let Some(conn) = self.conn.take() {
drop(conn);
}
self.connected = false;
self.in_transaction = false;
Ok(())
})
}
fn execute_with_params<'a>(
&'a mut self,
sql: &'a str,
params: &'a [Value],
) -> Pin<Box<dyn Future<Output = Result<u64, DbError>> + Send + 'a>> {
Box::pin(async move {
if needs_raw_sql(sql) || params.is_empty() {
return self.execute(sql).await;
}
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let mut q = sqlx::query(sqlx::AssertSqlSafe(sql));
for v in params {
q = match v {
Value::Null => q.bind(None::<i64>),
Value::Bool(b) => q.bind(*b),
Value::I8(n) => q.bind(*n),
Value::I16(n) => q.bind(*n),
Value::I32(n) => q.bind(*n),
Value::I64(n) => q.bind(*n),
Value::U8(n) => q.bind(*n),
Value::U16(n) => q.bind(*n),
Value::U32(n) => q.bind(*n),
Value::U64(n) => q.bind(*n as i64),
Value::F32(f) => q.bind(*f),
Value::F64(f) => q.bind(*f),
Value::String(s) => q.bind(s.as_str()),
Value::Decimal(s) => q.bind(s.as_str()),
Value::Bytes(b) => q.bind(b.as_slice()),
other => q.bind(other.to_string()),
};
}
let result = q.execute(&mut *pool_conn).await;
self.conn = Some(pool_conn);
match result {
Ok(r) => Ok(r.rows_affected()),
Err(e) => {
let db_err = map_sqlx_error(e);
if matches!(db_err, DbError::ConnectionError(_) | DbError::IoError(_)) {
self.connected = false;
}
Err(db_err)
}
}
})
}
fn query_with_params<'a>(
&'a mut self,
sql: &'a str,
params: &'a [Value],
) -> Pin<Box<dyn Future<Output = Result<QueryRows, DbError>> + Send + 'a>> {
Box::pin(async move {
if params.is_empty() {
return self.query(sql).await;
}
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let mut q = sqlx::query(sqlx::AssertSqlSafe(sql));
for v in params {
q = match v {
Value::Null => q.bind(None::<i64>),
Value::Bool(b) => q.bind(*b),
Value::I8(n) => q.bind(*n),
Value::I16(n) => q.bind(*n),
Value::I32(n) => q.bind(*n),
Value::I64(n) => q.bind(*n),
Value::U8(n) => q.bind(*n),
Value::U16(n) => q.bind(*n),
Value::U32(n) => q.bind(*n),
Value::U64(n) => q.bind(*n as i64),
Value::F32(f) => q.bind(*f),
Value::F64(f) => q.bind(*f),
Value::String(s) => q.bind(s.as_str()),
Value::Decimal(s) => q.bind(s.as_str()),
Value::Bytes(b) => q.bind(b.as_slice()),
other => q.bind(other.to_string()),
};
}
let rows_result = q.fetch_all(&mut *pool_conn).await;
self.conn = Some(pool_conn);
let rows = rows_result.map_err(map_sqlx_error)?;
let mut result = Vec::with_capacity(rows.len());
for row in rows {
let columns = row.columns();
let mut record = HashMap::with_capacity(columns.len());
for col in columns {
let name = col.name().to_string();
let ordinal = col.ordinal();
let value = row_to_value_mysql(&row, ordinal);
record.insert(name, value);
}
result.push(record);
}
Ok(result)
})
}
fn query_values<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<QueryValues, DbError>> + Send + 'a>> {
Box::pin(async move {
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let rows_result = (&mut *pool_conn).fetch_all(sqlx::AssertSqlSafe(sql)).await;
self.conn = Some(pool_conn);
let rows = rows_result.map_err(map_sqlx_error)?;
if rows.is_empty() {
return Ok((Vec::new(), Vec::new()));
}
let cols = rows[0].columns();
let mut col_names: Vec<String> = Vec::with_capacity(cols.len());
let mut col_types: Vec<ColType> = Vec::with_capacity(cols.len());
for col in cols {
col_names.push(col.name().to_string());
col_types.push(ColType::parse_mysql(col.type_info().name()));
}
let mut result_rows: Vec<Vec<Value>> = Vec::with_capacity(rows.len());
for row in &rows {
let mut row_values: Vec<Value> = Vec::with_capacity(col_names.len());
for (idx, _) in col_names.iter().enumerate() {
row_values.push(row_to_value_with_coltype_mysql(row, idx, col_types[idx]));
}
result_rows.push(row_values);
}
Ok((col_names, result_rows))
})
}
fn query_values_with_params<'a>(
&'a mut self,
sql: &'a str,
params: &'a [Value],
) -> Pin<Box<dyn Future<Output = Result<QueryValues, DbError>> + Send + 'a>> {
Box::pin(async move {
if params.is_empty() {
return self.query_values(sql).await;
}
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let mut q = sqlx::query(sqlx::AssertSqlSafe(sql));
for v in params {
q = match v {
Value::Null => q.bind(None::<i64>),
Value::Bool(b) => q.bind(*b),
Value::I8(n) => q.bind(*n),
Value::I16(n) => q.bind(*n),
Value::I32(n) => q.bind(*n),
Value::I64(n) => q.bind(*n),
Value::U8(n) => q.bind(*n),
Value::U16(n) => q.bind(*n),
Value::U32(n) => q.bind(*n),
Value::U64(n) => q.bind(*n as i64),
Value::F32(f) => q.bind(*f),
Value::F64(f) => q.bind(*f),
Value::String(s) => q.bind(s.as_str()),
Value::Decimal(s) => q.bind(s.as_str()),
Value::Bytes(b) => q.bind(b.as_slice()),
other => q.bind(other.to_string()),
};
}
let rows_result = q.fetch_all(&mut *pool_conn).await;
self.conn = Some(pool_conn);
let rows = rows_result.map_err(map_sqlx_error)?;
if rows.is_empty() {
return Ok((Vec::new(), Vec::new()));
}
let cols = rows[0].columns();
let mut col_names: Vec<String> = Vec::with_capacity(cols.len());
let mut col_types: Vec<ColType> = Vec::with_capacity(cols.len());
for col in cols {
col_names.push(col.name().to_string());
col_types.push(ColType::parse_mysql(col.type_info().name()));
}
let mut result_rows: Vec<Vec<Value>> = Vec::with_capacity(rows.len());
for row in &rows {
let mut row_values: Vec<Value> = Vec::with_capacity(col_names.len());
for (idx, _) in col_names.iter().enumerate() {
row_values.push(row_to_value_with_coltype_mysql(row, idx, col_types[idx]));
}
result_rows.push(row_values);
}
Ok((col_names, result_rows))
})
}
fn query_stream<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn futures::Stream<Item = Result<HashMap<String, Value>, DbError>> + Send + 'a>>
{
Box::pin(async_stream::try_stream! {
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let mut row_stream = sqlx::query(sqlx::AssertSqlSafe(sql)).fetch(&mut *pool_conn);
let mut col_types: Vec<ColType> = Vec::new();
let mut col_names: Vec<String> = Vec::new();
let mut first_row = true;
while let Some(row_result) = row_stream.next().await {
let row = row_result.map_err(map_sqlx_error)?;
if first_row {
for col in row.columns() {
col_names.push(col.name().to_string());
col_types.push(ColType::parse_mysql(col.type_info().name()));
}
first_row = false;
}
let mut record = HashMap::with_capacity(col_names.len());
for (i, name) in col_names.iter().enumerate() {
let value = row_to_value_with_coltype_mysql(&row, i, col_types[i]);
record.insert(name.clone(), value);
}
yield record;
}
drop(row_stream);
self.conn = Some(pool_conn);
})
}
}
impl Drop for SqlxMySqlConnection {
fn drop(&mut self) {
if let Some(conn) = self.conn.take() {
drop(conn);
}
}
}
pub async fn mysql_bulk_insert(
conn: &mut SqlxMySqlConnection,
table: &str,
columns: &[&str],
rows: &[Vec<Value>],
) -> Result<u64, DbError> {
if rows.is_empty() {
return Ok(0);
}
let col_list = columns.join(", ");
let cols_per_row = columns.len();
let row_placeholder = format!("({})", vec!["?"; cols_per_row].join(", "));
let placeholders = vec![row_placeholder; rows.len()].join(", ");
let sql = format!(
"INSERT INTO {} ({}) VALUES {}",
table, col_list, placeholders
);
let mut params: Vec<Value> = Vec::with_capacity(rows.len() * cols_per_row);
for row in rows {
for v in row {
params.push(v.clone());
}
}
conn.execute_with_params(&sql, ¶ms).await
}
fn row_to_value_pg(row: &sqlx::postgres::PgRow, ordinal: usize) -> Value {
use sqlx::TypeInfo;
let type_name = row.columns()[ordinal].type_info().name();
match type_name {
"BOOL" => match row.try_get::<Option<bool>, usize>(ordinal) {
Ok(v) => v.map(Value::Bool).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
"INT2" => match row.try_get::<Option<i16>, usize>(ordinal) {
Ok(v) => v.map(Value::I16).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
"INT4" | "OID" => match row.try_get::<Option<i32>, usize>(ordinal) {
Ok(v) => v.map(Value::I32).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
"INT8" => match row.try_get::<Option<i64>, usize>(ordinal) {
Ok(v) => v.map(Value::I64).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
"FLOAT4" => match row.try_get::<Option<f32>, usize>(ordinal) {
Ok(v) => v.map(Value::F32).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
"FLOAT8" => match row.try_get::<Option<f64>, usize>(ordinal) {
Ok(v) => v.map(Value::F64).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
"TEXT" | "VARCHAR" | "CHAR" | "NAME" => match row.try_get::<Option<String>, usize>(ordinal)
{
Ok(v) => v.map(Value::String).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
"BYTEA" => match row.try_get::<Option<Vec<u8>>, usize>(ordinal) {
Ok(v) => v.map(Value::Bytes).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
"NUMERIC" => match row.try_get::<Option<rust_decimal::Decimal>, usize>(ordinal) {
Ok(Some(v)) => Value::Decimal(v.to_string()),
Ok(None) => Value::Null,
Err(_) => match row.try_get::<Option<String>, usize>(ordinal) {
Ok(v) => v.map(Value::String).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
},
"UUID" => match row.try_get::<Option<sqlx::types::Uuid>, usize>(ordinal) {
Ok(v) => v
.map(|uuid| Value::String(uuid.to_string()))
.unwrap_or(Value::Null),
Err(_) => match row.try_get::<Option<String>, usize>(ordinal) {
Ok(v) => v.map(Value::String).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
},
_ => {
if let Ok(v) = row.try_get::<Option<i64>, usize>(ordinal) {
return v.map(Value::I64).unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<Option<f64>, usize>(ordinal) {
return v.map(Value::F64).unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<Option<bool>, usize>(ordinal) {
return v.map(Value::Bool).unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<Option<String>, usize>(ordinal) {
return v.map(Value::String).unwrap_or(Value::Null);
}
Value::Null
}
}
}
fn row_to_value_with_coltype_pg(
row: &sqlx::postgres::PgRow,
ordinal: usize,
col_type: ColType,
) -> Value {
match col_type {
ColType::Bool => match row.try_get::<Option<bool>, usize>(ordinal) {
Ok(v) => v.map(Value::Bool).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::I8 => match row.try_get::<Option<i8>, usize>(ordinal) {
Ok(v) => v.map(Value::I8).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::I16 => match row.try_get::<Option<i16>, usize>(ordinal) {
Ok(v) => v.map(Value::I16).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::I32 => match row.try_get::<Option<i32>, usize>(ordinal) {
Ok(v) => v.map(Value::I32).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::I64 => match row.try_get::<Option<i64>, usize>(ordinal) {
Ok(v) => v.map(Value::I64).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::U8 => match row.try_get::<Option<i16>, usize>(ordinal) {
Ok(v) => v.map(Value::I16).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::U16 => match row.try_get::<Option<i32>, usize>(ordinal) {
Ok(v) => v.map(Value::I32).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::U32 => match row.try_get::<Option<i64>, usize>(ordinal) {
Ok(v) => v.map(Value::I64).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::U64 => match row.try_get::<Option<i64>, usize>(ordinal) {
Ok(v) => v.map(Value::I64).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::F32 => match row.try_get::<Option<f32>, usize>(ordinal) {
Ok(v) => v.map(Value::F32).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::F64 => match row.try_get::<Option<f64>, usize>(ordinal) {
Ok(v) => v.map(Value::F64).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::Decimal => match row.try_get::<Option<rust_decimal::Decimal>, usize>(ordinal) {
Ok(Some(v)) => Value::Decimal(v.to_string()),
Ok(None) => Value::Null,
Err(_) => match row.try_get::<Option<String>, usize>(ordinal) {
Ok(v) => v.map(Value::String).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
},
ColType::String => match row.try_get::<Option<String>, usize>(ordinal) {
Ok(v) => v.map(Value::String).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::Bytes => match row.try_get::<Option<Vec<u8>>, usize>(ordinal) {
Ok(v) => v.map(Value::Bytes).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
ColType::Date | ColType::DateTime | ColType::Time | ColType::Json => {
match row.try_get::<Option<String>, usize>(ordinal) {
Ok(v) => v.map(Value::String).unwrap_or(Value::Null),
Err(_) => Value::Null,
}
}
ColType::Uuid => match row.try_get::<Option<sqlx::types::Uuid>, usize>(ordinal) {
Ok(v) => v
.map(|uuid| Value::String(uuid.to_string()))
.unwrap_or(Value::Null),
Err(_) => match row.try_get::<Option<String>, usize>(ordinal) {
Ok(v) => v.map(Value::String).unwrap_or(Value::Null),
Err(_) => Value::Null,
},
},
ColType::Unknown => {
if let Ok(v) = row.try_get::<Option<i64>, usize>(ordinal) {
return v.map(Value::I64).unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<Option<f64>, usize>(ordinal) {
return v.map(Value::F64).unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<Option<bool>, usize>(ordinal) {
return v.map(Value::Bool).unwrap_or(Value::Null);
}
if let Ok(v) = row.try_get::<Option<String>, usize>(ordinal) {
return v.map(Value::String).unwrap_or(Value::Null);
}
Value::Null
}
_ => Value::Null,
}
}
pub struct PgPoolHandle {
pool: sqlx::PgPool,
}
impl PgPoolHandle {
pub async fn connect(url: &str) -> Result<Self, DbError> {
let pool = sqlx::pool::PoolOptions::<sqlx::Postgres>::new()
.max_connections(10)
.acquire_timeout(std::time::Duration::from_secs(30))
.idle_timeout(Some(std::time::Duration::from_secs(600)))
.max_lifetime(Some(std::time::Duration::from_secs(1800)))
.connect(url)
.await
.map_err(map_sqlx_error)?;
Ok(Self { pool })
}
pub fn from_pool(pool: sqlx::PgPool) -> Self {
Self { pool }
}
pub fn pool(&self) -> &sqlx::PgPool {
&self.pool
}
}
pub struct SqlxPgConnectionFactory {
pool: Arc<PgPoolHandle>,
}
impl SqlxPgConnectionFactory {
pub fn new(pool: Arc<PgPoolHandle>) -> Self {
Self { pool }
}
}
#[async_trait]
impl ConnectionFactory for SqlxPgConnectionFactory {
async fn create(&self) -> Result<Box<dyn Connection>, DbError> {
let conn = self.pool.pool.acquire().await.map_err(map_sqlx_error)?;
Ok(Box::new(SqlxPgConnection {
conn: Some(conn),
connected: true,
in_transaction: false,
}))
}
}
pub struct SqlxPgConnection {
conn: Option<sqlx::pool::PoolConnection<sqlx::Postgres>>,
connected: bool,
in_transaction: bool,
}
impl Connection for SqlxPgConnection {
fn execute<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<u64, DbError>> + Send + 'a>> {
Box::pin(async move {
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let result = if needs_raw_sql(sql) {
(&mut *pool_conn)
.execute(sqlx::raw_sql(sqlx::AssertSqlSafe(sql)))
.await
} else {
(&mut *pool_conn).execute(sqlx::AssertSqlSafe(sql)).await
};
self.conn = Some(pool_conn);
match result {
Ok(r) => Ok(r.rows_affected()),
Err(e) => {
let db_err = map_sqlx_error(e);
if matches!(db_err, DbError::ConnectionError(_) | DbError::IoError(_)) {
self.connected = false;
}
Err(db_err)
}
}
})
}
fn query<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<Vec<HashMap<String, Value>>, DbError>> + Send + 'a>>
{
Box::pin(async move {
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let rows_result = (&mut *pool_conn).fetch_all(sqlx::AssertSqlSafe(sql)).await;
self.conn = Some(pool_conn);
let rows = rows_result.map_err(map_sqlx_error)?;
if rows.is_empty() {
return Ok(Vec::new());
}
let col_types: Vec<ColType> = rows[0]
.columns()
.iter()
.map(|col| ColType::parse_postgres(col.type_info().name()))
.collect();
let mut result = Vec::with_capacity(rows.len());
for row in &rows {
let mut record = HashMap::with_capacity(col_types.len());
for (i, col) in row.columns().iter().enumerate() {
let name = col.name().to_string();
let value = row_to_value_with_coltype_pg(row, i, col_types[i]);
record.insert(name, value);
}
result.push(record);
}
Ok(result)
})
}
fn begin_transaction<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
if self.in_transaction {
return Err(DbError::Internal("transaction already started".to_string()));
}
self.execute("BEGIN").await?;
self.in_transaction = true;
Ok(())
})
}
fn commit<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
if self.in_transaction {
self.execute("COMMIT").await?;
self.in_transaction = false;
}
Ok(())
})
}
fn rollback<'a>(
&'a mut self,
) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
if self.in_transaction {
let result = self.execute("ROLLBACK").await;
self.in_transaction = false;
result.map(|_| ())
} else {
Ok(())
}
})
}
fn is_connected(&self) -> bool {
self.connected
}
fn ping<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = bool> + Send + 'a>> {
Box::pin(async move {
match self.execute("SELECT 1").await {
Ok(_) => true,
Err(_) => {
self.connected = false;
false
}
}
})
}
fn close<'a>(&'a mut self) -> Pin<Box<dyn Future<Output = Result<(), DbError>> + Send + 'a>> {
Box::pin(async move {
if let Some(conn) = self.conn.take() {
drop(conn);
}
self.connected = false;
self.in_transaction = false;
Ok(())
})
}
fn execute_with_params<'a>(
&'a mut self,
sql: &'a str,
params: &'a [Value],
) -> Pin<Box<dyn Future<Output = Result<u64, DbError>> + Send + 'a>> {
Box::pin(async move {
if needs_raw_sql(sql) || params.is_empty() {
return self.execute(sql).await;
}
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let mut q = sqlx::query(sqlx::AssertSqlSafe(sql));
for v in params {
q = match v {
Value::Null => q.bind(None::<i64>),
Value::Bool(b) => q.bind(*b),
Value::I8(n) => q.bind(*n),
Value::I16(n) => q.bind(*n),
Value::I32(n) => q.bind(*n),
Value::I64(n) => q.bind(*n),
Value::U8(n) => q.bind(*n as i16),
Value::U16(n) => q.bind(*n as i32),
Value::U32(n) => q.bind(*n as i64),
Value::U64(n) => q.bind(*n as i64),
Value::F32(f) => q.bind(*f),
Value::F64(f) => q.bind(*f),
Value::String(s) => q.bind(s.as_str()),
Value::Decimal(s) => q.bind(s.as_str()),
Value::Bytes(b) => q.bind(b.as_slice()),
other => q.bind(other.to_string()),
};
}
let result = q.execute(&mut *pool_conn).await;
self.conn = Some(pool_conn);
match result {
Ok(r) => Ok(r.rows_affected()),
Err(e) => {
let db_err = map_sqlx_error(e);
if matches!(db_err, DbError::ConnectionError(_) | DbError::IoError(_)) {
self.connected = false;
}
Err(db_err)
}
}
})
}
fn query_with_params<'a>(
&'a mut self,
sql: &'a str,
params: &'a [Value],
) -> Pin<Box<dyn Future<Output = Result<QueryRows, DbError>> + Send + 'a>> {
Box::pin(async move {
if params.is_empty() {
return self.query(sql).await;
}
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let mut q = sqlx::query(sqlx::AssertSqlSafe(sql));
for v in params {
q = match v {
Value::Null => q.bind(None::<i64>),
Value::Bool(b) => q.bind(*b),
Value::I8(n) => q.bind(*n),
Value::I16(n) => q.bind(*n),
Value::I32(n) => q.bind(*n),
Value::I64(n) => q.bind(*n),
Value::U8(n) => q.bind(*n as i16),
Value::U16(n) => q.bind(*n as i32),
Value::U32(n) => q.bind(*n as i64),
Value::U64(n) => q.bind(*n as i64),
Value::F32(f) => q.bind(*f),
Value::F64(f) => q.bind(*f),
Value::String(s) => q.bind(s.as_str()),
Value::Decimal(s) => q.bind(s.as_str()),
Value::Bytes(b) => q.bind(b.as_slice()),
other => q.bind(other.to_string()),
};
}
let rows_result = q.fetch_all(&mut *pool_conn).await;
self.conn = Some(pool_conn);
let rows = rows_result.map_err(map_sqlx_error)?;
if rows.is_empty() {
return Ok(Vec::new());
}
let col_types: Vec<ColType> = rows[0]
.columns()
.iter()
.map(|col| ColType::parse_postgres(col.type_info().name()))
.collect();
let mut result = Vec::with_capacity(rows.len());
for row in &rows {
let mut record = HashMap::with_capacity(col_types.len());
for (i, col) in row.columns().iter().enumerate() {
let name = col.name().to_string();
let value = row_to_value_with_coltype_pg(row, i, col_types[i]);
record.insert(name, value);
}
result.push(record);
}
Ok(result)
})
}
fn query_values<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn Future<Output = Result<QueryValues, DbError>> + Send + 'a>> {
Box::pin(async move {
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let rows_result = (&mut *pool_conn).fetch_all(sqlx::AssertSqlSafe(sql)).await;
self.conn = Some(pool_conn);
let rows = rows_result.map_err(map_sqlx_error)?;
if rows.is_empty() {
return Ok((Vec::new(), Vec::new()));
}
let cols = rows[0].columns();
let mut col_names: Vec<String> = Vec::with_capacity(cols.len());
let mut col_types: Vec<ColType> = Vec::with_capacity(cols.len());
for col in cols {
col_names.push(col.name().to_string());
col_types.push(ColType::parse_postgres(col.type_info().name()));
}
let mut result_rows: Vec<Vec<Value>> = Vec::with_capacity(rows.len());
for row in &rows {
let mut row_values: Vec<Value> = Vec::with_capacity(col_names.len());
for (idx, _) in col_names.iter().enumerate() {
row_values.push(row_to_value_with_coltype_pg(row, idx, col_types[idx]));
}
result_rows.push(row_values);
}
Ok((col_names, result_rows))
})
}
fn query_values_with_params<'a>(
&'a mut self,
sql: &'a str,
params: &'a [Value],
) -> Pin<Box<dyn Future<Output = Result<QueryValues, DbError>> + Send + 'a>> {
Box::pin(async move {
if params.is_empty() {
return self.query_values(sql).await;
}
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let mut q = sqlx::query(sqlx::AssertSqlSafe(sql));
for v in params {
q = match v {
Value::Null => q.bind(None::<i64>),
Value::Bool(b) => q.bind(*b),
Value::I8(n) => q.bind(*n),
Value::I16(n) => q.bind(*n),
Value::I32(n) => q.bind(*n),
Value::I64(n) => q.bind(*n),
Value::U8(n) => q.bind(*n as i16),
Value::U16(n) => q.bind(*n as i32),
Value::U32(n) => q.bind(*n as i64),
Value::U64(n) => q.bind(*n as i64),
Value::F32(f) => q.bind(*f),
Value::F64(f) => q.bind(*f),
Value::String(s) => q.bind(s.as_str()),
Value::Decimal(s) => q.bind(s.as_str()),
Value::Bytes(b) => q.bind(b.as_slice()),
other => q.bind(other.to_string()),
};
}
let rows_result = q.fetch_all(&mut *pool_conn).await;
self.conn = Some(pool_conn);
let rows = rows_result.map_err(map_sqlx_error)?;
if rows.is_empty() {
return Ok((Vec::new(), Vec::new()));
}
let cols = rows[0].columns();
let mut col_names: Vec<String> = Vec::with_capacity(cols.len());
for col in cols {
col_names.push(col.name().to_string());
}
let mut result_rows: Vec<Vec<Value>> = Vec::with_capacity(rows.len());
for row in rows {
let mut row_values: Vec<Value> = Vec::with_capacity(col_names.len());
for (idx, _) in col_names.iter().enumerate() {
let ordinal = row.columns()[idx].ordinal();
row_values.push(row_to_value_pg(&row, ordinal));
}
result_rows.push(row_values);
}
Ok((col_names, result_rows))
})
}
fn query_stream<'a>(
&'a mut self,
sql: &'a str,
) -> Pin<Box<dyn futures::Stream<Item = Result<HashMap<String, Value>, DbError>> + Send + 'a>>
{
Box::pin(async_stream::try_stream! {
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let mut row_stream = sqlx::query(sqlx::AssertSqlSafe(sql)).fetch(&mut *pool_conn);
while let Some(row_result) = row_stream.next().await {
let row = row_result.map_err(map_sqlx_error)?;
let cols = row.columns();
let mut record = HashMap::with_capacity(cols.len());
for (i, col) in cols.iter().enumerate() {
let name = col.name().to_string();
let value = row_to_value_pg(&row, i);
record.insert(name, value);
}
yield record;
}
drop(row_stream);
self.conn = Some(pool_conn);
})
}
}
impl Drop for SqlxPgConnection {
fn drop(&mut self) {
if let Some(conn) = self.conn.take() {
drop(conn);
}
}
}
#[async_trait]
pub trait PgExtensions: Send + Sync {
async fn listen(&mut self, channel: &str) -> Result<(), DbError>;
async fn notify(&mut self, channel: &str, payload: &str) -> Result<(), DbError>;
async fn copy_from_stdin(&mut self, sql: &str, data: &[u8]) -> Result<u64, DbError>;
}
fn validate_pg_channel_name(channel: &str) -> Result<(), DbError> {
if channel.is_empty() {
return Err(DbError::Internal(
"PG channel name must not be empty".to_string(),
));
}
if !channel
.chars()
.all(|c| c.is_ascii_alphanumeric() || c == '_')
{
return Err(DbError::Internal(format!(
"invalid PG channel name: {} (only alphanumeric and underscore allowed)",
channel
)));
}
Ok(())
}
#[async_trait]
impl PgExtensions for SqlxPgConnection {
async fn listen(&mut self, channel: &str) -> Result<(), DbError> {
validate_pg_channel_name(channel)?;
self.execute(&format!("LISTEN {}", channel)).await?;
Ok(())
}
async fn notify(&mut self, channel: &str, payload: &str) -> Result<(), DbError> {
validate_pg_channel_name(channel)?;
let escaped_payload = payload.replace('\'', "''");
self.execute(&format!("NOTIFY {}, '{}'", channel, escaped_payload))
.await?;
Ok(())
}
async fn copy_from_stdin(&mut self, sql: &str, data: &[u8]) -> Result<u64, DbError> {
let mut pool_conn = self
.conn
.take()
.ok_or_else(|| DbError::Internal("connection already closed".to_string()))?;
let mut copy = (*pool_conn)
.copy_in_raw(sql)
.await
.map_err(map_sqlx_error)?;
copy.send(data).await.map_err(map_sqlx_error)?;
let result = copy.finish().await.map_err(map_sqlx_error)?;
self.conn = Some(pool_conn);
Ok(result)
}
}
pub async fn pg_bulk_insert(
conn: &mut SqlxPgConnection,
table: &str,
columns: &[&str],
rows: &[Vec<Value>],
) -> Result<u64, DbError> {
if rows.is_empty() {
return Ok(0);
}
let col_list = columns.join(", ");
let cols_per_row = columns.len();
let placeholders: Vec<String> = rows
.iter()
.enumerate()
.map(|(row_idx, _)| {
let base = row_idx * cols_per_row;
let ph: Vec<String> = (0..cols_per_row)
.map(|i| format!("${}", base + i + 1))
.collect();
format!("({})", ph.join(", "))
})
.collect();
let sql = format!(
"INSERT INTO {} ({}) VALUES {}",
table,
col_list,
placeholders.join(", ")
);
let mut params: Vec<Value> = Vec::with_capacity(rows.len() * cols_per_row);
for row in rows {
for v in row {
params.push(v.clone());
}
}
conn.execute_with_params(&sql, ¶ms).await
}