use crate::cook::execution::errors::{MapReduceError, MapReduceResult};
use async_trait::async_trait;
use std::collections::VecDeque;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::{Mutex, Semaphore};
use tracing::{debug, info};
#[derive(Debug, Clone, Default)]
pub struct PoolMetrics {
pub total_created: usize,
pub in_use: usize,
pub available: usize,
pub total_acquisitions: usize,
pub reuse_count: usize,
pub avg_wait_time_ms: u64,
}
#[async_trait]
pub trait ResourcePool<T>: Send + Sync {
async fn acquire(&self) -> MapReduceResult<super::ResourceGuard<T>>;
fn release(&self, resource: T);
fn metrics(&self) -> PoolMetrics;
async fn clear(&self);
}
pub struct GenericResourcePool<T, F>
where
T: Send + 'static,
F: Fn() -> futures::future::BoxFuture<'static, MapReduceResult<T>> + Send + Sync,
{
available: Arc<Mutex<VecDeque<T>>>,
factory: Arc<F>,
#[allow(dead_code)]
max_size: usize,
semaphore: Arc<Semaphore>,
metrics: Arc<Mutex<PoolMetrics>>,
cleanup: Arc<dyn Fn(T) + Send + Sync>,
}
impl<T, F> GenericResourcePool<T, F>
where
T: Send + 'static,
F: Fn() -> futures::future::BoxFuture<'static, MapReduceResult<T>> + Send + Sync,
{
pub fn new(max_size: usize, factory: F) -> Self {
Self::with_cleanup(max_size, factory, |_| {})
}
pub fn with_cleanup<C>(max_size: usize, factory: F, cleanup: C) -> Self
where
C: Fn(T) + Send + Sync + 'static,
{
Self {
available: Arc::new(Mutex::new(VecDeque::new())),
factory: Arc::new(factory),
max_size,
semaphore: Arc::new(Semaphore::new(max_size)),
metrics: Arc::new(Mutex::new(PoolMetrics::default())),
cleanup: Arc::new(cleanup),
}
}
async fn try_get_available(&self) -> Option<T> {
let mut available = self.available.lock().await;
available.pop_front()
}
#[allow(dead_code)]
async fn return_to_pool(&self, resource: T) {
let mut available = self.available.lock().await;
let mut metrics = self.metrics.lock().await;
if available.len() < self.max_size {
available.push_back(resource);
metrics.available = available.len();
metrics.in_use = metrics.in_use.saturating_sub(1);
} else {
(self.cleanup)(resource);
metrics.in_use = metrics.in_use.saturating_sub(1);
}
}
fn update_acquisition_metrics(metrics: &mut PoolMetrics, start: Instant, is_reuse: bool) {
metrics.in_use += 1;
metrics.total_acquisitions += 1;
if is_reuse {
metrics.reuse_count += 1;
metrics.available = metrics.available.saturating_sub(1);
} else {
metrics.total_created += 1;
}
let wait_time = start.elapsed();
metrics.avg_wait_time_ms = ((metrics.avg_wait_time_ms
* (metrics.total_acquisitions - 1) as u64)
+ wait_time.as_millis() as u64)
/ metrics.total_acquisitions as u64;
}
fn create_resource_guard(
resource: T,
pool: Arc<Mutex<VecDeque<T>>>,
cleanup: Arc<dyn Fn(T) + Send + Sync>,
) -> super::ResourceGuard<T> {
let pool_weak = Arc::downgrade(&pool);
super::ResourceGuard::new(resource, move |r| {
if let Some(pool) = pool_weak.upgrade() {
tokio::spawn(async move {
let mut available = pool.lock().await;
available.push_back(r);
});
} else {
cleanup(r);
}
})
}
}
#[async_trait]
impl<T, F> ResourcePool<T> for GenericResourcePool<T, F>
where
T: Send + 'static,
F: Fn() -> futures::future::BoxFuture<'static, MapReduceResult<T>> + Send + Sync,
{
async fn acquire(&self) -> MapReduceResult<super::ResourceGuard<T>> {
let start = Instant::now();
if let Some(resource) = self.try_get_available().await {
let mut metrics = self.metrics.lock().await;
Self::update_acquisition_metrics(&mut metrics, start, true);
debug!("Reused resource from pool");
return Ok(Self::create_resource_guard(
resource,
self.available.clone(),
self.cleanup.clone(),
));
}
let _permit = self
.semaphore
.acquire()
.await
.map_err(|e| MapReduceError::General {
message: format!("Failed to acquire pool semaphore: {}", e),
source: None,
})?;
let resource = (self.factory)().await?;
let mut metrics = self.metrics.lock().await;
Self::update_acquisition_metrics(&mut metrics, start, false);
info!("Created new resource (total: {})", metrics.total_created);
Ok(Self::create_resource_guard(
resource,
self.available.clone(),
self.cleanup.clone(),
))
}
fn release(&self, resource: T) {
let available = self.available.clone();
let metrics = self.metrics.clone();
tokio::spawn(async move {
let mut avail = available.lock().await;
let mut m = metrics.lock().await;
avail.push_back(resource);
m.available = avail.len();
m.in_use = m.in_use.saturating_sub(1);
});
}
fn metrics(&self) -> PoolMetrics {
PoolMetrics::default()
}
async fn clear(&self) {
let mut available = self.available.lock().await;
let cleanup = self.cleanup.clone();
while let Some(resource) = available.pop_front() {
cleanup(resource);
}
let mut metrics = self.metrics.lock().await;
metrics.available = 0;
}
}
pub struct BoundedResourcePool<T>
where
T: Send + 'static,
{
inner: Arc<dyn ResourcePool<T>>,
acquire_timeout: Duration,
}
impl<T> BoundedResourcePool<T>
where
T: Send + 'static,
{
pub fn new(inner: Arc<dyn ResourcePool<T>>, acquire_timeout: Duration) -> Self {
Self {
inner,
acquire_timeout,
}
}
}
#[async_trait]
impl<T> ResourcePool<T> for BoundedResourcePool<T>
where
T: Send + 'static,
{
async fn acquire(&self) -> MapReduceResult<super::ResourceGuard<T>> {
tokio::time::timeout(self.acquire_timeout, self.inner.acquire())
.await
.map_err(|_| MapReduceError::General {
message: format!(
"Resource acquisition timed out after {:?}",
self.acquire_timeout
),
source: None,
})?
}
fn release(&self, resource: T) {
self.inner.release(resource)
}
fn metrics(&self) -> PoolMetrics {
self.inner.metrics()
}
async fn clear(&self) {
self.inner.clear().await
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Debug, Clone)]
struct CounterResource {
id: usize,
}
fn create_counter_factory(
counter: Arc<AtomicUsize>,
) -> impl Fn() -> futures::future::BoxFuture<'static, MapReduceResult<CounterResource>> {
move || {
let counter = counter.clone();
Box::pin(async move {
let id = counter.fetch_add(1, Ordering::Relaxed);
Ok(CounterResource { id })
})
}
}
#[tokio::test]
async fn test_pool_creates_new_resource() {
let counter = Arc::new(AtomicUsize::new(0));
let factory = create_counter_factory(counter.clone());
let pool = GenericResourcePool::new(5, factory);
let guard = pool.acquire().await.expect("Failed to acquire resource");
let resource = guard.get().expect("Resource should be present");
assert_eq!(resource.id, 0);
assert_eq!(counter.load(Ordering::Relaxed), 1);
}
#[tokio::test]
async fn test_pool_reuses_resource() {
let counter = Arc::new(AtomicUsize::new(0));
let factory = create_counter_factory(counter.clone());
let pool = Arc::new(GenericResourcePool::new(5, factory));
{
let _guard = pool.acquire().await.expect("Failed to acquire");
}
tokio::time::sleep(Duration::from_millis(50)).await;
let guard = pool.acquire().await.expect("Failed to reacquire");
let resource = guard.get().expect("Resource should be present");
assert_eq!(resource.id, 0);
assert_eq!(counter.load(Ordering::Relaxed), 1);
}
#[tokio::test]
async fn test_pool_respects_max_size() {
let counter = Arc::new(AtomicUsize::new(0));
let factory = create_counter_factory(counter.clone());
let pool = Arc::new(GenericResourcePool::new(2, factory));
let guard1 = pool.acquire().await.expect("Failed to acquire 1");
let guard2 = pool.acquire().await.expect("Failed to acquire 2");
assert_eq!(counter.load(Ordering::Relaxed), 2);
drop(guard1);
tokio::time::sleep(Duration::from_millis(50)).await;
let guard3 = pool.acquire().await.expect("Failed to acquire 3");
assert_eq!(counter.load(Ordering::Relaxed), 2);
drop(guard2);
drop(guard3);
}
#[tokio::test]
async fn test_update_acquisition_metrics_reuse() {
let mut metrics = PoolMetrics::default();
let start = Instant::now();
type TestPool = GenericResourcePool<
CounterResource,
fn() -> futures::future::BoxFuture<'static, MapReduceResult<CounterResource>>,
>;
TestPool::update_acquisition_metrics(&mut metrics, start, true);
assert_eq!(metrics.in_use, 1);
assert_eq!(metrics.total_acquisitions, 1);
assert_eq!(metrics.reuse_count, 1);
assert_eq!(metrics.total_created, 0);
}
#[tokio::test]
async fn test_update_acquisition_metrics_new() {
let mut metrics = PoolMetrics::default();
let start = Instant::now();
type TestPool = GenericResourcePool<
CounterResource,
fn() -> futures::future::BoxFuture<'static, MapReduceResult<CounterResource>>,
>;
TestPool::update_acquisition_metrics(&mut metrics, start, false);
assert_eq!(metrics.in_use, 1);
assert_eq!(metrics.total_acquisitions, 1);
assert_eq!(metrics.reuse_count, 0);
assert_eq!(metrics.total_created, 1);
}
#[tokio::test]
async fn test_update_acquisition_metrics_multiple() {
let mut metrics = PoolMetrics::default();
let start = Instant::now();
type TestPool = GenericResourcePool<
CounterResource,
fn() -> futures::future::BoxFuture<'static, MapReduceResult<CounterResource>>,
>;
TestPool::update_acquisition_metrics(&mut metrics, start, false);
TestPool::update_acquisition_metrics(&mut metrics, start, true);
assert_eq!(metrics.in_use, 2);
assert_eq!(metrics.total_acquisitions, 2);
assert_eq!(metrics.reuse_count, 1);
assert_eq!(metrics.total_created, 1);
}
#[tokio::test]
async fn test_resource_guard_returns_to_pool() {
let counter = Arc::new(AtomicUsize::new(0));
let factory = create_counter_factory(counter.clone());
let pool = Arc::new(GenericResourcePool::new(5, factory));
let guard = pool.acquire().await.expect("Failed to acquire");
let initial_id = guard.get().expect("Resource should be present").id;
drop(guard);
tokio::time::sleep(Duration::from_millis(50)).await;
let guard2 = pool.acquire().await.expect("Failed to reacquire");
let reused_id = guard2.get().expect("Resource should be present").id;
assert_eq!(initial_id, reused_id);
}
#[tokio::test]
async fn test_resource_cleanup_called() {
let counter = Arc::new(AtomicUsize::new(0));
let cleanup_counter = Arc::new(AtomicUsize::new(0));
let cleanup_counter_clone = cleanup_counter.clone();
let factory = create_counter_factory(counter.clone());
let pool = GenericResourcePool::with_cleanup(5, factory, move |_resource| {
cleanup_counter_clone.fetch_add(1, Ordering::Relaxed);
});
for _ in 0..3 {
let guard = pool.acquire().await.expect("Failed to acquire");
drop(guard);
}
tokio::time::sleep(Duration::from_millis(50)).await;
pool.clear().await;
assert_eq!(cleanup_counter.load(Ordering::Relaxed), 3);
}
#[tokio::test]
async fn test_concurrent_acquisitions() {
let counter = Arc::new(AtomicUsize::new(0));
let factory = create_counter_factory(counter.clone());
let pool = Arc::new(GenericResourcePool::new(10, factory));
let mut handles = vec![];
for _ in 0..5 {
let pool_clone = pool.clone();
let handle = tokio::spawn(async move {
let guard = pool_clone.acquire().await.expect("Failed to acquire");
tokio::time::sleep(Duration::from_millis(10)).await;
drop(guard);
});
handles.push(handle);
}
for handle in handles.drain(..) {
handle.await.expect("Task panicked");
}
tokio::time::sleep(Duration::from_millis(50)).await;
let created_after_first_round = counter.load(Ordering::Relaxed);
for _ in 0..5 {
let pool_clone = pool.clone();
let handle = tokio::spawn(async move {
let guard = pool_clone.acquire().await.expect("Failed to acquire");
tokio::time::sleep(Duration::from_millis(10)).await;
drop(guard);
});
handles.push(handle);
}
for handle in handles {
handle.await.expect("Task panicked");
}
let created_after_second_round = counter.load(Ordering::Relaxed);
assert_eq!(
created_after_first_round, created_after_second_round,
"Resources should have been reused in second round"
);
}
}