use std::sync::{Arc, Mutex as StdMutex, MutexGuard, PoisonError};
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
use crate::client::{Client, ClientConfig, ClientError};
use crate::wire::Config;
fn lock<T>(mutex: &StdMutex<T>) -> MutexGuard<'_, T> {
mutex.lock().unwrap_or_else(PoisonError::into_inner)
}
pub struct Pool {
endpoint: String,
config: Config,
client_config: ClientConfig,
permits: Arc<Semaphore>,
idle: Arc<StdMutex<Vec<Client>>>,
}
impl Pool {
pub fn new(
endpoint: impl Into<String>,
config: Config,
client_config: ClientConfig,
max_connections: usize,
) -> Self {
let max = max_connections.max(1);
Self {
endpoint: endpoint.into(),
config,
client_config,
permits: Arc::new(Semaphore::new(max)),
idle: Arc::new(StdMutex::new(Vec::with_capacity(max))),
}
}
pub async fn acquire(&self) -> Result<PooledConn, ClientError> {
let permit = Arc::clone(&self.permits)
.acquire_owned()
.await
.map_err(|_| ClientError::Connection {
message: "connection pool is closed".to_owned(),
})?;
let reused = {
let mut idle = lock(&self.idle);
loop {
match idle.pop() {
Some(client) if client.is_alive() => break Some(client),
Some(_dead) => continue,
None => break None,
}
}
};
let client = match reused {
Some(client) => client,
None => {
Client::connect_with(
&self.endpoint,
self.config.clone(),
self.client_config.clone(),
)
.await?
}
};
Ok(PooledConn {
inner: Some(client),
idle: Arc::clone(&self.idle),
_permit: permit,
})
}
pub fn idle_count(&self) -> usize {
lock(&self.idle).len()
}
}
pub struct PooledConn {
inner: Option<Client>,
idle: Arc<StdMutex<Vec<Client>>>,
_permit: OwnedSemaphorePermit,
}
impl PooledConn {
pub fn client(&self) -> &Client {
match &self.inner {
Some(client) => client,
None => unreachable!("PooledConn::client after drop"),
}
}
}
impl std::ops::Deref for PooledConn {
type Target = Client;
fn deref(&self) -> &Client {
self.client()
}
}
impl Drop for PooledConn {
fn drop(&mut self) {
if let Some(client) = self.inner.take() {
if client.is_alive() {
lock(&self.idle).push(client);
}
}
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
#[test]
fn new_does_not_dial_and_clamps_capacity() {
let pool = Pool::new(
"test://127.0.0.1:0",
Config::standard(),
ClientConfig::new(),
0,
);
assert_eq!(pool.idle_count(), 0);
assert_eq!(pool.permits.available_permits(), 1);
}
}