use super::{Client, Param};
use crate::errors::Error::PoolError;
use crate::errors::Result;
use async_mutex::Mutex as AsyncMutex;
use std::ops::DerefMut;
use std::sync::Arc;
use std::sync::Mutex;
use std::sync::mpsc::{Receiver, Sender, channel};
use std::time::Duration;
use tokio::task::yield_now;
use tokio::time::sleep;
#[cfg(feature = "unstable-api")]
pub(crate) mod pooled_stream;
pub(crate) type TiberiusConn = tiberius::Client<tokio_util::compat::Compat<tokio::net::TcpStream>>;
mod pooledconnection;
pub(crate) use pooledconnection::PooledConnection;
pub(crate) struct Pool {
ado_connection_string: String,
slots: Vec<Mutex<Slot>>,
round_robin_next: Mutex<usize>,
tx: Sender<(TiberiusConn, ConnectionStatus)>,
}
impl Pool {
pub fn new(ado: impl Into<String>) -> Arc<Self> {
let size = 10;
let mut slots = Vec::with_capacity(size);
for _ in 0..size {
slots.push(Mutex::new(Slot::Empty));
}
let (tx, rx) = channel();
let me = Arc::new(Self {
ado_connection_string: ado.into(),
slots,
round_robin_next: Mutex::new(0),
tx,
});
let return_ref = me.clone();
tokio::spawn(async move { pool_return(return_ref, rx).await });
me
}
pub async fn get(&self) -> Result<PooledConnection> {
let mut slot_index: usize = {
let guard = self.round_robin_next.lock().map_err(|_| PoolError)?;
*guard
};
loop {
let cs = self.ado_connection_string.as_str();
let checked_out = try_checkout(slot_index, &self.slots, cs).await?;
if let Some(tiberius_conn) = checked_out {
slot_index += 1;
slot_index %= self.slots.len();
{
let mut guard = self.round_robin_next.lock().unwrap();
*guard = slot_index;
}
let tiberius_conn = AsyncMutex::new(Some(tiberius_conn));
return Ok(PooledConnection {
status: ConnectionStatus::Clean,
tiberius_conn,
conn_return: self.tx.clone(),
});
}
slot_index += 1;
slot_index %= self.slots.len();
yield_now().await;
}
}
pub async fn status(&self) -> String {
let mut display = Vec::with_capacity(self.slots.len() + 2);
display.push('[');
for slot in &self.slots {
let guard = slot.lock().unwrap();
let slot_guard: &Slot = &guard;
match slot_guard {
Slot::Avalable(_) => display.push('A'),
Slot::Empty => display.push('.'),
Slot::Checkedout => display.push('C'),
}
}
display.push(']');
display.iter().collect()
}
}
async fn try_checkout(
index: usize,
slots: &[Mutex<Slot>],
ado: &str,
) -> Result<Option<TiberiusConn>> {
let slot: Slot = {
let mut slot_guard = slots[index].lock().unwrap();
let slot_guard: &mut Slot = slot_guard.deref_mut();
if let Slot::Checkedout = slot_guard {
return Ok(None);
}
let mut slot = Slot::Checkedout;
std::mem::swap(&mut slot, slot_guard);
slot
};
let conn: Option<TiberiusConn> = match slot {
Slot::Avalable(c) => Some(c),
Slot::Empty => None,
Slot::Checkedout => panic!("double checkout"),
};
match conn {
None => {
log::debug!("MSSQL POOL adding Connection");
let new_conn = build_connection(ado).await?;
Ok(Some(new_conn))
}
Some(mut conn) => {
if conn.execute("SELECT 1", &[]).await.is_err() {
log::debug!("MSSQL POOL rebuild Connection");
let new_conn = build_connection(ado).await?;
Ok(Some(new_conn))
} else {
Ok(Some(conn))
}
}
}
}
async fn build_connection(ado: &str) -> Result<TiberiusConn> {
let config = tiberius::Config::from_ado_string(ado)?;
use crate::errors::Error;
use tiberius::Client;
use tokio::net::TcpStream;
use tokio_util::compat::TokioAsyncWriteCompatExt;
let tcp = TcpStream::connect(config.get_addr())
.await
.map_err(Error::TiberiusConnPool)?;
tcp.set_nodelay(true).map_err(Error::TiberiusConnPool)?;
let mut client = Client::connect(config, tcp.compat_write())
.await
.map_err(Error::Tiberius)?;
let _ = client.simple_query("SELECT 1").await?;
Ok(client)
}
pub(crate) enum Slot {
Avalable(TiberiusConn),
Checkedout,
Empty,
}
#[derive(Clone, PartialEq, Eq)]
pub(crate) enum ConnectionStatus {
Clean,
NeedsRollback(String),
}
async fn pool_return(pool: Arc<Pool>, mut rx: Receiver<(TiberiusConn, ConnectionStatus)>) {
loop {
if let Ok(returned_count) = pool_return_inner(pool.clone(), &mut rx).await
&& returned_count == 0
{
sleep(Duration::from_millis(10)).await;
}
if Arc::strong_count(&pool) == 1 {
return;
}
yield_now().await;
}
}
async fn pool_return_inner(
pool: Arc<Pool>,
rx: &mut Receiver<(TiberiusConn, ConnectionStatus)>,
) -> Result<usize> {
let tuple = match rx.try_recv().ok() {
None => return Ok(0),
Some(conn) => conn,
};
tokio::spawn(async move {
let (mut conn, status) = tuple;
if let ConnectionStatus::NeedsRollback(_trans_name) = status {
let sql = "WHILE @@TRANCOUNT > 0 BEGIN ROLLBACK TRANSACTION; END";
let _ = conn.simple_query(sql).await;
}
let slot_index: usize = {
let guard = pool.round_robin_next.lock().unwrap();
*guard
};
let mut index = slot_index;
let itor = pool
.slots
.iter()
.cycle()
.skip(slot_index)
.take(pool.slots.len());
for slot in itor {
let mut slot_guard = slot.lock().unwrap();
let slot_guard: &mut Slot = slot_guard.deref_mut();
let mut is_checkout = false;
if let Slot::Checkedout = slot_guard {
is_checkout = true;
}
if is_checkout {
let mut returning = Slot::Avalable(conn);
std::mem::swap(&mut returning, slot_guard);
{
let mut guard = pool.round_robin_next.lock().unwrap();
*guard = index;
}
return;
}
index += 1;
index %= pool.slots.len();
}
panic!("unable to return a connection to the connection pool, pool is full");
});
Ok(1)
}