use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use async_trait::async_trait;
use tokio::sync::OwnedSemaphorePermit;
use tokio::sync::Semaphore;
static ACQUIRE_COUNT: AtomicUsize = AtomicUsize::new(0);
#[derive(Debug, thiserror::Error)]
pub enum PoolError {
#[error("pool exhausted")]
Exhausted,
#[error(transparent)]
Other(#[from] Box<dyn std::error::Error + Send + Sync>),
}
pub struct Pooled<T> {
pub resource: T,
_permit: OwnedSemaphorePermit,
}
#[async_trait]
pub trait Pool<T>: Send + Sync {
async fn acquire(&self) -> Result<Pooled<T>, PoolError>;
}
#[derive(Clone)]
pub struct SemaphorePool<T: Clone> {
semaphore: Arc<Semaphore>,
items: Arc<Vec<T>>,
}
impl<T: Clone> SemaphorePool<T> {
pub fn items(&self) -> &[T] {
&self.items
}
pub fn new(items: Vec<T>) -> Self {
let permits = items.len().max(1);
Self {
semaphore: Arc::new(Semaphore::new(permits)),
items: Arc::new(items),
}
}
}
#[async_trait]
impl<T: Clone + Send + Sync + 'static> Pool<T> for SemaphorePool<T> {
async fn acquire(&self) -> Result<Pooled<T>, PoolError> {
let permit = self.semaphore.clone().acquire_owned().await
.map_err(|_| PoolError::Exhausted)?;
let idx = ACQUIRE_COUNT.fetch_add(1, Ordering::Relaxed) % self.items.len();
Ok(Pooled {
resource: self.items[idx].clone(),
_permit: permit,
})
}
}
#[async_trait]
pub trait Reconnectable {
type Item;
fn is_healthy(&self, item: &Self::Item) -> bool;
async fn reconnect(&self) -> Result<Self::Item, Box<dyn std::error::Error + Send + Sync>>;
}
pub struct AutoReconnectPool<P, R> {
inner: P,
reconnect: R,
}
impl<P, R> AutoReconnectPool<P, R> {
pub fn new(inner: P, reconnect: R) -> Self {
Self { inner, reconnect }
}
}
#[async_trait]
impl<P, R> Pool<R::Item> for AutoReconnectPool<P, R>
where
P: Pool<R::Item> + Send + Sync,
R: Reconnectable + Send + Sync,
R::Item: Send,
{
async fn acquire(&self) -> Result<Pooled<R::Item>, PoolError> {
let pooled = { self.inner.acquire().await? };
let item = pooled.resource;
if self.reconnect.is_healthy(&item) {
Ok(Pooled {
resource: item,
_permit: pooled._permit,
})
} else {
let fresh = self
.reconnect
.reconnect()
.await
.map_err(PoolError::Other)?;
Ok(Pooled {
resource: fresh,
_permit: pooled._permit,
})
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn pool_returns_item() {
let pool = SemaphorePool::new(vec![42u32, 84u32]);
let item = pool.acquire().await.unwrap();
assert!(item.resource == 42 || item.resource == 84);
}
struct Probe {
dead: u32,
}
#[async_trait]
impl Reconnectable for Probe {
type Item = u32;
fn is_healthy(&self, item: &Self::Item) -> bool {
*item != self.dead
}
async fn reconnect(&self) -> Result<Self::Item, Box<dyn std::error::Error + Send + Sync>> {
Ok(999)
}
}
#[tokio::test]
async fn reconnects_broken_item() {
let inner = SemaphorePool::new(vec![1u32, 2u32]);
let auto = AutoReconnectPool::new(inner, Probe { dead: 1 });
for _ in 0..10 {
let p = auto.acquire().await.unwrap();
assert_ne!(p.resource, 1); }
}
}