use crate::dialect::Dialect;
use crate::driver::{Driver, DriverConnection};
use rustlavel_core::{Error, Result};
use std::collections::VecDeque;
use std::sync::Arc;
use tokio::sync::{Mutex, Semaphore};
struct Inner {
driver: Arc<dyn Driver>,
idle: Mutex<VecDeque<(u64, Box<dyn DriverConnection>)>>,
permits: Arc<Semaphore>,
}
#[derive(Clone)]
pub struct Pool {
inner: Arc<Inner>,
}
impl Pool {
pub fn new(driver: Arc<dyn Driver>) -> Self {
let permits = Arc::new(Semaphore::new(driver.max_connections().max(1)));
Pool { inner: Arc::new(Inner { driver, idle: Mutex::new(VecDeque::new()), permits }) }
}
pub async fn verify(&self) -> Result<()> {
let mut connection = self.acquire().await?;
connection.simple_query("select 1").await?;
Ok(())
}
pub fn driver(&self) -> &Arc<dyn Driver> {
&self.inner.driver
}
pub fn dialect(&self) -> Arc<dyn Dialect> {
self.inner.driver.dialect()
}
pub async fn acquire(&self) -> Result<PooledConnection> {
let permit = Arc::clone(&self.inner.permits)
.acquire_owned()
.await
.map_err(|_| Error::msg("the database pool has been closed"))?;
let generation = self.inner.driver.generation();
loop {
let Some((opened_under, connection)) = self.inner.idle.lock().await.pop_front() else {
break;
};
if opened_under == generation {
return Ok(PooledConnection {
connection: Some(connection),
generation,
pool: Arc::clone(&self.inner),
_permit: permit,
});
}
connection.close().await;
}
let connection = self.inner.driver.connect().await?;
Ok(PooledConnection {
connection: Some(connection),
generation,
pool: Arc::clone(&self.inner),
_permit: permit,
})
}
pub async fn idle_count(&self) -> usize {
self.inner.idle.lock().await.len()
}
pub async fn open_count(&self) -> usize {
let borrowed = self
.inner
.driver
.max_connections()
.max(1)
.saturating_sub(self.inner.permits.available_permits());
self.inner.idle.lock().await.len() + borrowed
}
pub async fn close_idle(&self, limit: usize) -> usize {
let mut closed = 0;
while closed < limit {
let Some((_, connection)) = self.inner.idle.lock().await.pop_front() else { break };
connection.close().await;
closed += 1;
}
closed
}
pub async fn close(&self) {
let mut idle = self.inner.idle.lock().await;
while let Some((_, connection)) = idle.pop_front() {
connection.close().await;
}
}
pub async fn retire_superseded(&self) -> usize {
let generation = self.inner.driver.generation();
let mut idle = self.inner.idle.lock().await;
let mut keeping = VecDeque::with_capacity(idle.len());
let mut closed = 0;
while let Some((opened_under, connection)) = idle.pop_front() {
if opened_under == generation {
keeping.push_back((opened_under, connection));
} else {
connection.close().await;
closed += 1;
}
}
*idle = keeping;
closed
}
}
pub struct PooledConnection {
connection: Option<Box<dyn DriverConnection>>,
generation: u64,
pool: Arc<Inner>,
_permit: tokio::sync::OwnedSemaphorePermit,
}
impl std::ops::Deref for PooledConnection {
type Target = dyn DriverConnection;
fn deref(&self) -> &(dyn DriverConnection + 'static) {
self.connection.as_deref().expect("connection is present until drop")
}
}
impl std::ops::DerefMut for PooledConnection {
fn deref_mut(&mut self) -> &mut (dyn DriverConnection + 'static) {
self.connection.as_deref_mut().expect("connection is present until drop")
}
}
impl Drop for PooledConnection {
fn drop(&mut self) {
let Some(connection) = self.connection.take() else { return };
if connection.is_broken() || connection.in_transaction() {
tokio::spawn(async move { connection.close().await });
return;
}
let pool = Arc::clone(&self.pool);
let generation = self.generation;
tokio::spawn(async move {
if generation != pool.driver.generation() {
connection.close().await;
return;
}
pool.idle.lock().await.push_back((generation, connection));
});
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::dialect::Postgres;
use crate::driver::BoxFuture;
struct Counting {
opened: Arc<std::sync::atomic::AtomicUsize>,
closed: Arc<std::sync::atomic::AtomicUsize>,
generation: Arc<std::sync::atomic::AtomicU64>,
}
struct Nothing(Arc<std::sync::atomic::AtomicUsize>);
impl DriverConnection for Nothing {
fn query<'a>(
&'a mut self,
_sql: &'a str,
_params: &'a [crate::value::Value],
) -> BoxFuture<'a, Result<crate::driver::QueryResult>> {
Box::pin(async { Err(Error::msg("not a real connection")) })
}
fn simple_query<'a>(
&'a mut self,
_sql: &'a str,
) -> BoxFuture<'a, Result<crate::driver::QueryResult>> {
Box::pin(async { Err(Error::msg("not a real connection")) })
}
fn is_broken(&self) -> bool {
false
}
fn in_transaction(&self) -> bool {
false
}
fn close(self: Box<Self>) -> BoxFuture<'static, ()> {
self.0.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Box::pin(async {})
}
}
impl Driver for Counting {
fn dialect(&self) -> Arc<dyn Dialect> {
Arc::new(Postgres)
}
fn connect(&self) -> BoxFuture<'_, Result<Box<dyn DriverConnection>>> {
self.opened.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
let closed = Arc::clone(&self.closed);
Box::pin(async move { Ok(Box::new(Nothing(closed)) as Box<dyn DriverConnection>) })
}
fn describe(&self) -> String {
"test://counting".into()
}
fn generation(&self) -> u64 {
self.generation.load(std::sync::atomic::Ordering::Acquire)
}
}
fn counting() -> (Pool, Arc<std::sync::atomic::AtomicUsize>, Arc<std::sync::atomic::AtomicUsize>, Arc<std::sync::atomic::AtomicU64>) {
let opened = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let closed = Arc::new(std::sync::atomic::AtomicUsize::new(0));
let generation = Arc::new(std::sync::atomic::AtomicU64::new(1));
let driver = Counting {
opened: Arc::clone(&opened),
closed: Arc::clone(&closed),
generation: Arc::clone(&generation),
};
(Pool::new(Arc::new(driver)), opened, closed, generation)
}
async fn settle() {
for _ in 0..8 {
tokio::task::yield_now().await;
}
}
#[tokio::test]
async fn a_connection_comes_back_and_is_reused() {
let (pool, opened, _, _) = counting();
drop(pool.acquire().await.unwrap());
settle().await;
drop(pool.acquire().await.unwrap());
settle().await;
assert_eq!(opened.load(std::sync::atomic::Ordering::SeqCst), 1, "the second borrow reused it");
}
#[tokio::test]
async fn an_idle_connection_from_a_rotated_credential_is_never_handed_out() {
let (pool, opened, closed, generation) = counting();
drop(pool.acquire().await.unwrap());
settle().await;
assert_eq!(pool.idle_count().await, 1);
generation.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
drop(pool.acquire().await.unwrap());
settle().await;
assert_eq!(opened.load(std::sync::atomic::Ordering::SeqCst), 2, "a fresh connection");
assert_eq!(closed.load(std::sync::atomic::Ordering::SeqCst), 1, "the stale one was closed");
}
#[tokio::test]
async fn a_borrowed_connection_is_retired_when_it_comes_back() {
let (pool, _, closed, generation) = counting();
let borrowed = pool.acquire().await.unwrap();
generation.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
drop(borrowed);
settle().await;
assert_eq!(pool.idle_count().await, 0);
assert_eq!(closed.load(std::sync::atomic::Ordering::SeqCst), 1);
}
#[tokio::test]
async fn retiring_early_closes_the_stale_and_keeps_the_current() {
let (pool, _, closed, generation) = counting();
let first = pool.acquire().await.unwrap();
let second = pool.acquire().await.unwrap();
drop(first);
drop(second);
settle().await;
assert_eq!(pool.idle_count().await, 2);
generation.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
assert_eq!(pool.retire_superseded().await, 2);
assert_eq!(pool.idle_count().await, 0);
assert_eq!(closed.load(std::sync::atomic::Ordering::SeqCst), 2);
drop(pool.acquire().await.unwrap());
settle().await;
assert_eq!(pool.retire_superseded().await, 0);
assert_eq!(pool.idle_count().await, 1);
}
#[tokio::test]
async fn a_pool_with_static_credentials_never_retires_anything() {
let (pool, opened, closed, _) = counting();
for _ in 0..5 {
drop(pool.acquire().await.unwrap());
settle().await;
}
assert_eq!(opened.load(std::sync::atomic::Ordering::SeqCst), 1);
assert_eq!(closed.load(std::sync::atomic::Ordering::SeqCst), 0);
}
struct Unreachable;
impl Driver for Unreachable {
fn dialect(&self) -> Arc<dyn Dialect> {
Arc::new(Postgres)
}
fn connect(&self) -> BoxFuture<'_, Result<Box<dyn DriverConnection>>> {
Box::pin(async { Err(Error::msg("nothing is listening")) })
}
fn describe(&self) -> String {
"test://unreachable".into()
}
fn max_connections(&self) -> usize {
3
}
}
#[tokio::test]
async fn a_pool_opens_nothing_until_it_is_used() {
let pool = Pool::new(Arc::new(Unreachable));
assert_eq!(pool.idle_count().await, 0);
}
#[tokio::test]
async fn acquiring_reports_the_drivers_failure() {
let pool = Pool::new(Arc::new(Unreachable));
let error = match pool.acquire().await {
Err(error) => error.to_string(),
Ok(_) => panic!("this driver cannot connect"),
};
assert!(error.contains("nothing is listening"), "{error}");
}
#[tokio::test]
async fn the_pool_carries_its_drivers_dialect() {
let pool = Pool::new(Arc::new(Unreachable));
assert_eq!(pool.dialect().name(), "postgres");
}
}