use super::super::DbType;
use super::common_helpers;
use super::{SqlExecutor, SqlStatement};
use crate::impl_insert_conflict_methods;
use crate::model::{FromRowValues, Model};
use crate::query::insert::InsertConflict;
use crate::raw_sql::{IntoRawSql, RawSql};
#[cfg(any(feature = "sqlite", feature = "mssql"))]
use std::collections::VecDeque;
use std::marker::PhantomData;
#[cfg(any(feature = "sqlite", feature = "mssql"))]
use std::sync::Arc;
#[cfg(any(feature = "sqlite", feature = "mssql"))]
use std::sync::atomic::{AtomicU32, Ordering};
#[cfg(any(feature = "sqlite", feature = "mssql"))]
use tokio::sync::Mutex;
#[cfg(feature = "postgresql")]
use bb8_postgres::PostgresConnectionManager;
#[cfg(feature = "postgresql")]
use tokio_postgres::NoTls;
#[cfg(any(
feature = "sqlite",
feature = "postgresql",
feature = "mysql",
feature = "mssql"
))]
use super::unified::{CreateTableExecutor, DropTableExecutor};
pub struct PooledInsertExecutor<'a, I: crate::model::Insertable> {
pooled_conn: &'a PooledConnection<'a>,
models: I,
conflict: Option<InsertConflict>,
_marker: PhantomData<I>,
}
impl_insert_conflict_methods!(PooledInsertExecutor);
impl<'a, I: crate::model::Insertable> PooledInsertExecutor<'a, I> {
pub fn to_sql(&self) -> crate::Result<SqlStatement> {
let refs = self.models.as_refs();
if refs.is_empty() {
return Ok(SqlStatement::batch(
db_type_for_connection(self.pooled_conn.get_connection()),
Vec::new(),
));
}
match self.pooled_conn.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(_) => {
let (sql, all_values) =
common_helpers::build_insert_statement_with_conflict_and_auto_increment_returning::<I::Model>(
DbType::Sqlite,
&refs,
self.conflict.as_ref(),
)?;
Ok(SqlStatement::single(DbType::Sqlite, sql, all_values))
}
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(_) => {
let (sql, all_values) =
common_helpers::build_insert_statement_with_conflict_and_auto_increment_returning::<I::Model>(
DbType::PostgreSQL,
&refs,
self.conflict.as_ref(),
)?;
let rust_types = postgresql_backend::pg_insert_param_rust_types::<I::Model>(
refs.len(),
self.conflict.as_ref(),
);
Ok(SqlStatement::batch(
DbType::PostgreSQL,
vec![
super::SingleSqlStatement::new(sql, all_values)
.with_param_rust_types(rust_types),
],
))
}
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(_) => {
let (sql, all_values) =
common_helpers::build_insert_statement_with_conflict_and_auto_increment_returning::<I::Model>(
DbType::MySQL,
&refs,
self.conflict.as_ref(),
)?;
Ok(SqlStatement::single(DbType::MySQL, sql, all_values))
}
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(_) => {
let (sql, all_values) =
common_helpers::build_insert_statement_with_conflict_and_auto_increment_returning::<I::Model>(
DbType::MSSQL,
&refs,
self.conflict.as_ref(),
)?;
Ok(SqlStatement::single(DbType::MSSQL, sql, all_values))
}
}
}
pub async fn execute(
self,
) -> crate::Result<<I::Model as crate::model::Model>::AutoIncrementKeyType>
where
I: Send + Sync,
{
<Self as SqlExecutor>::execute(self).await
}
}
impl<'a, I: crate::model::Insertable + Send + Sync> SqlExecutor for PooledInsertExecutor<'a, I> {
type Output = <I::Model as crate::model::Model>::AutoIncrementKeyType;
fn to_sql(&self) -> crate::Result<SqlStatement> {
PooledInsertExecutor::to_sql(self)
}
async fn execute_with_sql(self, _sql: SqlStatement) -> crate::Result<Self::Output> {
if self
.conflict
.as_ref()
.is_some_and(|conflict| conflict.is_configured())
{
return match self.pooled_conn.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(db) => {
db.insert(self.models)
.with_conflict(self.conflict)
.execute()
.await
}
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(db) => {
db.insert(self.models)
.with_conflict(self.conflict)
.execute()
.await
}
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(db) => {
db.insert(self.models)
.with_conflict(self.conflict)
.execute()
.await
}
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(db) => {
db.insert(self.models)
.with_conflict(self.conflict)
.execute()
.await
}
};
}
let refs = self.models.as_refs();
match self.pooled_conn.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(db) => db.insert_impl::<I::Model>(&refs).await,
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(db) => db.insert_impl::<I::Model>(&refs).await,
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(db) => db.insert_impl::<I::Model>(&refs).await,
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(db) => db.insert_impl::<I::Model>(&refs).await,
}
}
}
pub struct PooledInsertOrUpdateExecutor<'a, I: crate::model::Insertable> {
pooled_conn: &'a PooledConnection<'a>,
models: I,
_marker: PhantomData<I>,
}
impl<'a, I: crate::model::Insertable> PooledInsertOrUpdateExecutor<'a, I> {
pub fn to_sql(&self) -> crate::Result<SqlStatement> {
let refs = self.models.as_refs();
if refs.is_empty() {
return Ok(SqlStatement::batch(
db_type_for_connection(self.pooled_conn.get_connection()),
Vec::new(),
));
}
match self.pooled_conn.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(_) => {
let columns = I::Model::insert_columns();
let primary_key_columns = I::Model::primary_key_columns();
let primary_key = primary_key_columns.join(", ");
let (mut sql, all_values) = common_helpers::build_batch_insert_statement::<I::Model>(
DbType::Sqlite,
"INSERT INTO",
<I::Model as Model>::table_name_for_db(DbType::Sqlite),
&columns,
&refs,
common_helpers::BatchInsertValuesMode::WithoutAutoIncrement,
);
sql.push_str(&format!(" ON CONFLICT ({}) DO UPDATE SET ", primary_key));
let mut first = true;
for col_name in columns.iter() {
if primary_key_columns.contains(col_name) {
continue;
}
if !first {
sql.push_str(", ");
}
sql.push_str(&format!("{col_name} = excluded.{col_name}"));
first = false;
}
Ok(SqlStatement::single(DbType::Sqlite, sql, all_values))
}
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(_) => {
let columns = I::Model::insert_columns();
let primary_key_columns = I::Model::primary_key_columns();
let primary_key = primary_key_columns.join(", ");
let (mut sql, all_values) = common_helpers::build_batch_insert_statement::<I::Model>(
DbType::PostgreSQL,
"INSERT INTO",
<I::Model as Model>::table_name_for_db(DbType::PostgreSQL),
&columns,
&refs,
common_helpers::BatchInsertValuesMode::WithoutAutoIncrement,
);
sql.push_str(&format!(" ON CONFLICT ({}) DO UPDATE SET ", primary_key));
let mut first = true;
for col_name in columns.iter() {
if primary_key_columns.contains(col_name) {
continue;
}
if !first {
sql.push_str(", ");
}
sql.push_str(&format!("{col_name} = EXCLUDED.{col_name}"));
first = false;
}
let rust_types: Vec<&str> = I::Model::COLUMN_SCHEMA
.iter()
.filter(|col| !col.is_auto_increment)
.map(|col| col.data_type.unwrap_or(col.rust_type))
.collect();
Ok(SqlStatement::batch(
DbType::PostgreSQL,
vec![
super::SingleSqlStatement::new(sql, all_values)
.with_param_rust_types(rust_types),
],
))
}
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(_) => {
let (mut sql, all_values) = common_helpers::build_batch_insert_statement::<I::Model>(
DbType::MySQL,
"INSERT INTO",
<I::Model as Model>::table_name_for_db(DbType::MySQL),
I::Model::COLUMNS,
&refs,
common_helpers::BatchInsertValuesMode::All,
);
sql.push_str(" ON DUPLICATE KEY UPDATE ");
let mut first = true;
for col_name in I::Model::COLUMNS.iter() {
if !first {
sql.push_str(", ");
}
sql.push_str(&format!("{col_name} = VALUES({col_name})"));
first = false;
}
Ok(SqlStatement::single(DbType::MySQL, sql, all_values))
}
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(_) => {
let (mut sql, all_values) =
common_helpers::build_mssql_merge_source::<I::Model>(&refs);
common_helpers::append_mssql_merge_update_clause::<I::Model>(&mut sql);
common_helpers::append_mssql_merge_insert_clause::<I::Model>(&mut sql);
Ok(SqlStatement::single(DbType::MSSQL, sql, all_values))
}
}
}
pub async fn execute(self) -> crate::Result<()> {
<Self as SqlExecutor>::execute(self).await
}
}
impl<'a, I: crate::model::Insertable> SqlExecutor for PooledInsertOrUpdateExecutor<'a, I> {
type Output = ();
fn to_sql(&self) -> crate::Result<SqlStatement> {
PooledInsertOrUpdateExecutor::to_sql(self)
}
async fn execute_with_sql(self, _sql: SqlStatement) -> crate::Result<Self::Output> {
let refs = self.models.as_refs();
match self.pooled_conn.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(db) => db.insert_or_update_batch::<I::Model>(&refs).await,
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(db) => db.insert_or_update_batch::<I::Model>(&refs).await,
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(db) => db.insert_or_update_batch::<I::Model>(&refs).await,
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(db) => db.insert_or_update_impl::<I::Model>(&refs).await,
}
}
}
pub struct PooledInsertOrIgnoreExecutor<'a, I: crate::model::Insertable> {
pooled_conn: &'a PooledConnection<'a>,
models: I,
_marker: PhantomData<I>,
}
impl<'a, I: crate::model::Insertable> PooledInsertOrIgnoreExecutor<'a, I> {
pub fn to_sql(&self) -> crate::Result<SqlStatement> {
let refs = self.models.as_refs();
if refs.is_empty() {
return Ok(SqlStatement::batch(
db_type_for_connection(self.pooled_conn.get_connection()),
Vec::new(),
));
}
match self.pooled_conn.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(_) => {
let columns = I::Model::insert_columns();
let primary_key_columns = I::Model::primary_key_columns();
let primary_key = primary_key_columns.join(", ");
let (mut sql, all_values) = common_helpers::build_batch_insert_statement::<I::Model>(
DbType::Sqlite,
"INSERT INTO",
<I::Model as Model>::table_name_for_db(DbType::Sqlite),
&columns,
&refs,
common_helpers::BatchInsertValuesMode::WithoutAutoIncrement,
);
sql.push_str(&format!(" ON CONFLICT ({}) DO NOTHING", primary_key));
Ok(SqlStatement::single(DbType::Sqlite, sql, all_values))
}
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(_) => {
let columns = I::Model::insert_columns();
let primary_key_columns = I::Model::primary_key_columns();
let primary_key = primary_key_columns.join(", ");
let (mut sql, all_values) = common_helpers::build_batch_insert_statement::<I::Model>(
DbType::PostgreSQL,
"INSERT INTO",
<I::Model as Model>::table_name_for_db(DbType::PostgreSQL),
&columns,
&refs,
common_helpers::BatchInsertValuesMode::WithoutAutoIncrement,
);
sql.push_str(&format!(" ON CONFLICT ({}) DO NOTHING", primary_key));
let rust_types: Vec<&str> = I::Model::COLUMN_SCHEMA
.iter()
.filter(|col| !col.is_auto_increment)
.map(|col| col.data_type.unwrap_or(col.rust_type))
.collect();
Ok(SqlStatement::batch(
DbType::PostgreSQL,
vec![
super::SingleSqlStatement::new(sql, all_values)
.with_param_rust_types(rust_types),
],
))
}
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(_) => {
let (sql, all_values) = common_helpers::build_batch_insert_statement::<I::Model>(
DbType::MySQL,
"INSERT IGNORE INTO",
<I::Model as Model>::table_name_for_db(DbType::MySQL),
I::Model::COLUMNS,
&refs,
common_helpers::BatchInsertValuesMode::All,
);
Ok(SqlStatement::single(DbType::MySQL, sql, all_values))
}
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(_) => {
let (mut sql, all_values) =
common_helpers::build_mssql_merge_source::<I::Model>(&refs);
common_helpers::append_mssql_merge_insert_clause::<I::Model>(&mut sql);
Ok(SqlStatement::single(DbType::MSSQL, sql, all_values))
}
}
}
pub async fn execute(self) -> crate::Result<()> {
<Self as SqlExecutor>::execute(self).await
}
}
fn db_type_for_connection(connection: &ConnectionWrapper) -> DbType {
match connection {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(_) => DbType::Sqlite,
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(_) => DbType::PostgreSQL,
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(_) => DbType::MySQL,
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(_) => DbType::MSSQL,
}
}
impl<'a, I: crate::model::Insertable> SqlExecutor for PooledInsertOrIgnoreExecutor<'a, I> {
type Output = ();
fn to_sql(&self) -> crate::Result<SqlStatement> {
PooledInsertOrIgnoreExecutor::to_sql(self)
}
async fn execute_with_sql(self, _sql: SqlStatement) -> crate::Result<Self::Output> {
let refs = self.models.as_refs();
match self.pooled_conn.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(db) => db.insert_or_ignore_batch::<I::Model>(&refs).await,
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(db) => db.insert_or_ignore_batch::<I::Model>(&refs).await,
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(db) => db.insert_or_ignore_batch::<I::Model>(&refs).await,
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(db) => db
.insert_or_ignore_impl::<I::Model>(&refs)
.await
.map(|_| ()),
}
}
}
#[cfg(feature = "sqlite")]
use super::super::sqlite_backend;
#[cfg(feature = "postgresql")]
use super::super::postgresql_backend;
#[cfg(feature = "mysql")]
use super::super::mysql_backend;
#[cfg(feature = "mssql")]
use super::super::mssql_backend;
#[allow(clippy::upper_case_acronyms)]
enum ConnectionWrapper {
#[cfg(feature = "sqlite")]
Sqlite(sqlite_backend::Database),
#[cfg(feature = "postgresql")]
PostgreSQL(postgresql_backend::Database),
#[cfg(feature = "mysql")]
MySQL(mysql_backend::Database),
#[cfg(feature = "mssql")]
MSSQL(mssql_backend::Database),
}
#[cfg(any(feature = "sqlite", feature = "mssql"))]
impl ConnectionWrapper {
#[cfg(any(feature = "sqlite", feature = "mssql"))]
async fn is_valid(&self) -> bool {
match self {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(db) => db.is_valid().await,
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(db) => db.is_valid().await,
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(db) => db.is_valid().await,
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(db) => db.is_valid(),
}
}
}
#[cfg(any(feature = "sqlite", feature = "mssql"))]
pub struct ManualPool {
idle_connections: Mutex<VecDeque<ConnectionWrapper>>,
total_connections: AtomicU32,
config: PoolConfig,
db_type: DbType,
connection_string: String,
}
#[cfg(any(feature = "sqlite", feature = "mssql"))]
impl ManualPool {
fn new(db_type: DbType, connection_string: String, config: PoolConfig) -> Arc<Self> {
Arc::new(Self {
idle_connections: Mutex::new(VecDeque::new()),
total_connections: AtomicU32::new(0),
config,
db_type,
connection_string,
})
}
async fn create_connection(&self) -> crate::Result<ConnectionWrapper> {
match self.db_type {
#[cfg(feature = "sqlite")]
DbType::Sqlite => {
let db = crate::utils::FutureTraceExt::trace(sqlite_backend::Database::connect(
self.db_type,
&self.connection_string,
))
.await?;
Ok(ConnectionWrapper::Sqlite(db))
}
#[cfg(feature = "postgresql")]
DbType::PostgreSQL => Err(crate::ormer_error!(
"Build the pool through PoolBuilder/ConnectionPool"
)),
#[cfg(feature = "mysql")]
DbType::MySQL => Err(crate::ormer_error!(
"Build the pool through PoolBuilder/ConnectionPool"
)),
#[cfg(feature = "mssql")]
DbType::MSSQL => {
let db = crate::utils::FutureTraceExt::trace(mssql_backend::Database::connect(
self.db_type,
&self.connection_string,
))
.await?;
Ok(ConnectionWrapper::MSSQL(db))
}
}
}
async fn get(&self) -> crate::Result<ConnectionWrapper> {
{
let mut idle = self.idle_connections.lock().await;
if let Some(conn) = idle.pop_front() {
if conn.is_valid().await {
return Ok(conn);
}
self.total_connections.fetch_sub(1, Ordering::SeqCst);
}
}
let current_total = self.total_connections.load(Ordering::SeqCst);
if current_total < self.config.max_size {
let conn = crate::utils::FutureTraceExt::trace(self.create_connection()).await?;
self.total_connections.fetch_add(1, Ordering::SeqCst);
return Ok(conn);
}
loop {
tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
let mut idle = self.idle_connections.lock().await;
if let Some(conn) = idle.pop_front() {
if conn.is_valid().await {
return Ok(conn);
}
self.total_connections.fetch_sub(1, Ordering::SeqCst);
}
}
}
async fn return_connection(&self, conn: ConnectionWrapper) {
if conn.is_valid().await {
let mut idle = self.idle_connections.lock().await;
idle.push_back(conn);
} else {
self.total_connections.fetch_sub(1, Ordering::SeqCst);
}
}
async fn maintain_min_connections(&self) {
let current_total = self.total_connections.load(Ordering::SeqCst);
let target = self.config.min_size;
if current_total < target {
let to_create = target - current_total;
for _ in 0..to_create {
if let Ok(conn) = self.create_connection().await {
self.total_connections.fetch_add(1, Ordering::SeqCst);
let mut idle = self.idle_connections.lock().await;
idle.push_back(conn);
}
}
}
}
}
#[derive(Clone)]
pub struct PoolConfig {
min_size: u32,
max_size: u32,
}
impl Default for PoolConfig {
fn default() -> Self {
Self {
min_size: 0,
max_size: 10,
}
}
}
pub struct PoolBuilder {
db_type: DbType,
connection_string: String,
config: PoolConfig,
}
impl PoolBuilder {
pub fn new(db_type: DbType, connection_string: &str) -> Self {
Self {
db_type,
connection_string: connection_string.to_string(),
config: PoolConfig::default(),
}
}
pub fn range(mut self, range: std::ops::Range<u32>) -> Self {
self.config.min_size = range.start;
self.config.max_size = range.end;
self
}
pub async fn build(self) -> crate::Result<ConnectionPool> {
match self.db_type {
#[cfg(feature = "sqlite")]
DbType::Sqlite => {
let pool =
ManualPool::new(self.db_type, self.connection_string, self.config.clone());
if self.config.min_size > 0 {
pool.maintain_min_connections().await;
}
Ok(ConnectionPool::Sqlite(pool))
}
#[cfg(feature = "postgresql")]
DbType::PostgreSQL => {
let manager = crate::utils::ResultTraceExt::trace_for(
PostgresConnectionManager::new_from_stringlike(&self.connection_string, NoTls),
"bb8_postgres::PostgresConnectionManager::new_from_stringlike",
)?;
let mut builder = bb8::Pool::builder();
builder = builder.max_size(self.config.max_size);
if self.config.min_size > 0 {
builder = builder.min_idle(Some(self.config.min_size));
}
let pool = crate::utils::FutureTraceExt::trace(builder.build(manager)).await?;
Ok(ConnectionPool::PostgreSQL(pool))
}
#[cfg(feature = "mysql")]
DbType::MySQL => {
let opts = crate::utils::ResultTraceExt::trace_for(
mysql_async::Opts::from_url(&self.connection_string),
"mysql_async::Opts::from_url",
)?;
let pool = mysql_async::Pool::new(opts);
Ok(ConnectionPool::MySQL(pool))
}
#[cfg(feature = "mssql")]
DbType::MSSQL => {
let pool =
ManualPool::new(self.db_type, self.connection_string, self.config.clone());
if self.config.min_size > 0 {
pool.maintain_min_connections().await;
}
Ok(ConnectionPool::MSSQL(pool))
}
}
}
}
pub enum ConnectionPool {
#[cfg(feature = "sqlite")]
Sqlite(Arc<ManualPool>),
#[cfg(feature = "postgresql")]
PostgreSQL(bb8::Pool<PostgresConnectionManager<NoTls>>),
#[cfg(feature = "mysql")]
MySQL(mysql_async::Pool),
#[cfg(feature = "mssql")]
MSSQL(Arc<ManualPool>),
}
impl ConnectionPool {
pub async fn get(&self) -> crate::Result<PooledConnection<'_>> {
match self {
#[cfg(feature = "sqlite")]
ConnectionPool::Sqlite(pool) => {
let conn = crate::utils::FutureTraceExt::trace(pool.get()).await?;
Ok(PooledConnection {
inner: PooledConnectionInner::Sqlite(pool.clone()),
connection: Some(conn),
_marker: PhantomData,
})
}
#[cfg(feature = "postgresql")]
ConnectionPool::PostgreSQL(pool) => {
let pooled = crate::utils::FutureTraceExt::trace(pool.get()).await?;
let db = postgresql_backend::Database::from_pooled_connection(pooled);
Ok(PooledConnection {
inner: PooledConnectionInner::PostgreSQL,
connection: Some(ConnectionWrapper::PostgreSQL(db)),
_marker: PhantomData,
})
}
#[cfg(feature = "mysql")]
ConnectionPool::MySQL(pool) => {
let db = mysql_backend::Database::from_pool(pool.clone());
Ok(PooledConnection {
inner: PooledConnectionInner::MySQL,
connection: Some(ConnectionWrapper::MySQL(db)),
_marker: PhantomData,
})
}
#[cfg(feature = "mssql")]
ConnectionPool::MSSQL(pool) => {
let conn = crate::utils::FutureTraceExt::trace(pool.get()).await?;
Ok(PooledConnection {
inner: PooledConnectionInner::MSSQL(pool.clone()),
connection: Some(conn),
_marker: PhantomData,
})
}
}
}
}
#[derive(Clone)]
#[allow(clippy::upper_case_acronyms)]
enum PooledConnectionInner {
#[cfg(feature = "sqlite")]
Sqlite(Arc<ManualPool>),
#[cfg(feature = "postgresql")]
PostgreSQL,
#[cfg(feature = "mysql")]
MySQL,
#[cfg(feature = "mssql")]
MSSQL(Arc<ManualPool>),
}
impl PooledConnectionInner {
async fn return_connection(&self, conn: ConnectionWrapper) {
match self {
#[cfg(feature = "sqlite")]
PooledConnectionInner::Sqlite(pool) => pool.return_connection(conn).await,
#[cfg(feature = "postgresql")]
PooledConnectionInner::PostgreSQL => {
let _ = conn;
}
#[cfg(feature = "mysql")]
PooledConnectionInner::MySQL => {
let _ = conn;
}
#[cfg(feature = "mssql")]
PooledConnectionInner::MSSQL(pool) => pool.return_connection(conn).await,
}
}
}
pub struct PooledRawSelectExecutor<'conn, 'pool, T> {
pooled_conn: &'conn PooledConnection<'pool>,
sql: RawSql,
_marker: PhantomData<T>,
}
impl<'conn, 'pool, T> PooledRawSelectExecutor<'conn, 'pool, T> {
pub async fn collect<C>(self) -> crate::Result<C>
where
T: FromRowValues,
C: FromIterator<T>,
{
let raw_sql = self.sql;
match self.pooled_conn.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(db) => {
let (sql, params) = raw_sql.render(DbType::Sqlite)?;
db.select_raw::<T, C>(&sql, params).await
}
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(db) => {
let (sql, params) = raw_sql.render(DbType::PostgreSQL)?;
db.select_raw::<T, C>(&sql, params).await
}
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(db) => {
let (sql, params) = raw_sql.render(DbType::MySQL)?;
db.select_raw::<T, C>(&sql, params).await
}
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(db) => {
let (sql, params) = raw_sql.render(DbType::MSSQL)?;
db.select_raw::<T, C>(&sql, params).await
}
}
}
}
pub struct PooledConnection<'a> {
inner: PooledConnectionInner,
connection: Option<ConnectionWrapper>,
_marker: PhantomData<&'a ()>,
}
impl<'a> Drop for PooledConnection<'a> {
fn drop(&mut self) {
if let Some(conn) = self.connection.take() {
let inner = self.inner.clone();
match tokio::runtime::Handle::try_current() {
Ok(handle) => {
handle.spawn(async move {
inner.return_connection(conn).await;
});
}
Err(_) => {
eprintln!(
"Warning: PooledConnection dropped outside tokio runtime, connection may be leaked"
);
}
}
}
}
}
impl<'a> PooledConnection<'a> {
fn get_connection(&self) -> &ConnectionWrapper {
self.connection.as_ref().expect("Connection already taken")
}
pub fn create_table<T: Model>(&self) -> CreateTableExecutor<'_, T> {
match self.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(db) => CreateTableExecutor::Sqlite(db.create_table::<T>()),
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(db) => {
CreateTableExecutor::PostgreSQL(db.create_table::<T>())
}
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(db) => CreateTableExecutor::MySQL(db.create_table::<T>()),
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(db) => CreateTableExecutor::MSSQL(db.create_table::<T>()),
}
}
pub async fn validate_table<T: Model>(&self) -> crate::Result<()> {
match self.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(db) => db.validate_table::<T>().await,
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(db) => db.validate_table::<T>().await,
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(db) => db.validate_table::<T>().await,
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(db) => db.validate_table::<T>().await,
}
}
pub fn insert<I: crate::model::Insertable>(&self, models: I) -> PooledInsertExecutor<'_, I> {
PooledInsertExecutor {
pooled_conn: self,
models,
conflict: None,
_marker: PhantomData,
}
}
pub fn insert_or_update<I: crate::model::Insertable>(
&self,
models: I,
) -> PooledInsertOrUpdateExecutor<'_, I> {
PooledInsertOrUpdateExecutor {
pooled_conn: self,
models,
_marker: PhantomData,
}
}
pub fn upsert<I: crate::model::Insertable>(
&self,
models: I,
) -> PooledInsertOrUpdateExecutor<'_, I> {
self.insert_or_update(models)
}
pub fn insert_or_ignore<I: crate::model::Insertable>(
&self,
models: I,
) -> PooledInsertOrIgnoreExecutor<'_, I> {
PooledInsertOrIgnoreExecutor {
pooled_conn: self,
models,
_marker: PhantomData,
}
}
pub fn select<T: Model>(&self) -> super::unified::SelectExecutor<'_, T> {
match self.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(db) => {
super::unified::SelectExecutor::Sqlite(db.select::<T>())
}
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(db) => {
super::unified::SelectExecutor::PostgreSQL(db.select::<T>())
}
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(db) => super::unified::SelectExecutor::MySQL(db.select::<T>()),
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(db) => super::unified::SelectExecutor::MSSQL(db.select::<T>()),
}
}
pub fn stream<T: Model>(&self) -> super::unified::SelectStream<'_, T> {
self.select::<T>().stream()
}
pub fn delete<T: Model>(&self) -> super::unified::DeleteExecutor<'_, T> {
match self.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(db) => {
super::unified::DeleteExecutor::Sqlite(db.delete::<T>(), std::marker::PhantomData)
}
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(db) => {
super::unified::DeleteExecutor::PostgreSQL(db.delete::<T>())
}
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(db) => super::unified::DeleteExecutor::MySQL(db.delete::<T>()),
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(db) => super::unified::DeleteExecutor::MSSQL(db.delete::<T>()),
}
}
pub fn update<T: Model>(&self) -> super::unified::UpdateExecutor<'_, T> {
match self.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(db) => {
super::unified::UpdateExecutor::Sqlite(db.update::<T>(), std::marker::PhantomData)
}
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(db) => {
super::unified::UpdateExecutor::PostgreSQL(db.update::<T>())
}
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(db) => super::unified::UpdateExecutor::MySQL(db.update::<T>()),
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(db) => super::unified::UpdateExecutor::MSSQL(db.update::<T>()),
}
}
pub fn related<T: Model + 'static, R: Model>(
&self,
) -> super::unified::RelatedSelectExecutor<'_, T, R> {
match self.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(db) => super::unified::RelatedSelectExecutor::Sqlite(
db.related::<T, R>(),
std::marker::PhantomData,
),
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(db) => {
super::unified::RelatedSelectExecutor::PostgreSQL(db.related::<T, R>())
}
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(db) => {
super::unified::RelatedSelectExecutor::MySQL(db.related::<T, R>())
}
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(db) => {
super::unified::RelatedSelectExecutor::MSSQL(db.related::<T, R>())
}
}
}
pub async fn begin(&self) -> crate::Result<super::unified::Transaction<'_>> {
match self.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(db) => {
let txn = crate::utils::FutureTraceExt::trace(db.begin()).await?;
Ok(super::unified::Transaction::Sqlite(txn))
}
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(db) => {
let txn = crate::utils::FutureTraceExt::trace(db.begin()).await?;
Ok(super::unified::Transaction::PostgreSQL(txn))
}
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(db) => {
let txn = crate::utils::FutureTraceExt::trace(db.begin()).await?;
Ok(super::unified::Transaction::MySQL(txn))
}
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(db) => {
let txn = crate::utils::FutureTraceExt::trace(db.begin()).await?;
Ok(super::unified::Transaction::MSSQL(txn))
}
}
}
pub fn drop_table<T: Model>(&self) -> DropTableExecutor<'_, T> {
match self.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(db) => DropTableExecutor::Sqlite(db.drop_table::<T>()),
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(db) => {
DropTableExecutor::PostgreSQL(db.drop_table::<T>())
}
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(db) => DropTableExecutor::MySQL(db.drop_table::<T>()),
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(db) => DropTableExecutor::MSSQL(db.drop_table::<T>()),
}
}
pub fn select_sql<T>(&self, sql: impl IntoRawSql) -> PooledRawSelectExecutor<'_, 'a, T> {
PooledRawSelectExecutor {
pooled_conn: self,
sql: sql.into_raw_sql(),
_marker: PhantomData,
}
}
pub async fn execute_sql(&self, sql: impl IntoRawSql) -> crate::Result<u64> {
let raw_sql = sql.into_raw_sql();
match self.get_connection() {
#[cfg(feature = "sqlite")]
ConnectionWrapper::Sqlite(db) => {
let (sql, params) = raw_sql.render(DbType::Sqlite)?;
db.exec_raw(&sql, params).await
}
#[cfg(feature = "postgresql")]
ConnectionWrapper::PostgreSQL(db) => {
let (sql, params) = raw_sql.render(DbType::PostgreSQL)?;
db.exec_raw(&sql, params).await
}
#[cfg(feature = "mysql")]
ConnectionWrapper::MySQL(db) => {
let (sql, params) = raw_sql.render(DbType::MySQL)?;
db.exec_raw(&sql, params).await
}
#[cfg(feature = "mssql")]
ConnectionWrapper::MSSQL(db) => {
let (sql, params) = raw_sql.render(DbType::MSSQL)?;
db.exec_raw(&sql, params).await
}
}
}
}