use super::connection_trait::ConnectionProvider;
use super::deadpool_connection::{Pool, TcpManager, TcpManagerOptions};
use super::health_check::{HealthCheckMetrics, check_date_response};
use crate::pool::PoolStatus;
use crate::tls::TlsConfig;
use anyhow::Result;
use deadpool::managed;
use futures::StreamExt;
use futures::stream::FuturesUnordered;
use std::sync::Arc;
use std::sync::atomic::Ordering;
use tokio::sync::broadcast;
use tracing::{debug, info, warn};
#[derive(Debug, Clone)]
pub struct DeadpoolConnectionProvider {
pool: Pool,
name: Arc<str>,
shutdown_tx: Option<broadcast::Sender<()>>,
pub health_check_metrics: Arc<HealthCheckMetrics>,
original_max_size: usize,
active_cooldowns: Arc<std::sync::atomic::AtomicUsize>,
replacement_cooldown: Option<std::time::Duration>,
is_shutting_down: Arc<std::sync::atomic::AtomicBool>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct DeadpoolStatusCounts {
pub available: usize,
pub size: usize,
pub max_size: usize,
pub waiting: usize,
}
pub struct Builder {
host: String,
port: u16,
name: Option<String>,
max_size: usize,
username: Option<String>,
password: Option<String>,
tls_config: Option<TlsConfig>,
}
impl Builder {
#[must_use]
pub fn new(host: impl Into<String>, port: u16) -> Self {
Self {
host: host.into(),
port,
name: None,
max_size: 10, username: None,
password: None,
tls_config: None,
}
}
#[must_use]
pub fn name(mut self, name: impl Into<String>) -> Self {
self.name = Some(name.into());
self
}
#[must_use]
pub const fn max_connections(mut self, max_size: usize) -> Self {
self.max_size = max_size;
self
}
#[must_use]
pub fn username(mut self, username: impl Into<String>) -> Self {
self.username = Some(username.into());
self
}
#[must_use]
pub fn password(mut self, password: impl Into<String>) -> Self {
self.password = Some(password.into());
self
}
#[must_use]
pub fn tls_config(mut self, config: TlsConfig) -> Self {
self.tls_config = Some(config);
self
}
pub fn build(self) -> Result<DeadpoolConnectionProvider> {
let name = self
.name
.unwrap_or_else(|| format!("{}:{}", self.host, self.port));
let manager = TcpManager::new(
self.host,
self.port,
name.clone(),
TcpManagerOptions {
username: self.username,
password: self.password,
tls_config: self.tls_config,
..TcpManagerOptions::default()
},
)?;
Ok(DeadpoolConnectionProvider::from_manager(
manager,
name,
self.max_size,
))
}
}
impl DeadpoolConnectionProvider {
#[must_use]
pub fn builder(host: impl Into<String>, port: u16) -> Builder {
Builder::new(host, port)
}
pub fn simple(host: impl Into<String>, port: u16) -> Result<Self> {
Self::builder(host, port).build()
}
pub fn with_auth(
host: impl Into<String>,
port: u16,
username: impl Into<String>,
password: impl Into<String>,
) -> Result<Self> {
Self::builder(host, port)
.username(username)
.password(password)
.build()
}
pub fn with_tls(host: impl Into<String>, port: u16) -> Result<Self> {
Self::builder(host, port)
.tls_config(TlsConfig::default())
.build()
}
pub fn with_tls_auth(
host: impl Into<String>,
port: u16,
username: impl Into<String>,
password: impl Into<String>,
) -> Result<Self> {
Self::builder(host, port)
.username(username)
.password(password)
.tls_config(TlsConfig::default())
.build()
}
#[must_use]
pub fn new(
host: String,
port: u16,
name: String,
max_size: usize,
username: Option<String>,
password: Option<String>,
) -> Self {
Self::from_manager(
TcpManager::new(
host,
port,
name.clone(),
TcpManagerOptions {
username,
password,
..TcpManagerOptions::default()
},
)
.expect("Plain TCP TcpManager creation cannot fail"),
name,
max_size,
)
}
pub fn new_with_tls(
host: String,
port: u16,
name: String,
max_size: usize,
username: Option<String>,
password: Option<String>,
tls_config: TlsConfig,
) -> Result<Self> {
let manager = TcpManager::new(
host,
port,
name.clone(),
TcpManagerOptions {
username,
password,
tls_config: Some(tls_config),
..TcpManagerOptions::default()
},
)?;
Ok(Self::from_manager(manager, name, max_size))
}
fn from_manager(manager: TcpManager, name: String, max_size: usize) -> Self {
let pool = Pool::builder(manager)
.max_size(max_size)
.build()
.expect("Failed to create connection pool");
Self {
pool,
name: Arc::from(name),
shutdown_tx: None,
health_check_metrics: Arc::new(HealthCheckMetrics::new()),
original_max_size: max_size,
active_cooldowns: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
replacement_cooldown: None, is_shutting_down: Arc::new(std::sync::atomic::AtomicBool::new(false)),
}
}
pub fn from_server_config(
server: &crate::config::Server,
recv_buffer_size: usize,
send_buffer_size: usize,
) -> Result<Self> {
let tls_builder = TlsConfig::builder()
.enabled(server.use_tls)
.verify_cert(server.tls_verify_cert);
let tls_builder = server
.tls_cert_path
.as_ref()
.map(|cert_path| tls_builder.clone().cert_path(cert_path.as_str()))
.unwrap_or(tls_builder);
let tls_config = tls_builder.build();
let manager = TcpManager::new(
server.host.to_string(),
server.port.get(),
server.name.to_string(),
TcpManagerOptions {
username: server.username.clone(),
password: server.password.clone(),
tls_config: Some(tls_config),
recv_buffer_size,
send_buffer_size,
compress: server.compress,
compress_level: server.compress_level,
..TcpManagerOptions::default()
},
)?;
let max_size = server.max_connections.get();
let pool = Pool::builder(manager)
.max_size(max_size)
.build()
.expect("Failed to create connection pool");
let keepalive_interval = server.connection_keepalive;
let metrics = Arc::new(HealthCheckMetrics::new());
let shutdown_tx = keepalive_interval.map_or_else(
|| None,
|interval| {
let (tx, rx) = broadcast::channel(1);
let pool_clone = pool.clone();
let name_clone = server.name.to_string();
let metrics_clone = metrics.clone();
tokio::spawn(async move {
Self::run_periodic_health_checks(
pool_clone,
name_clone,
interval,
rx,
metrics_clone,
)
.await;
});
Some(tx)
},
);
Ok(Self {
pool,
name: Arc::from(server.name.to_string()),
shutdown_tx,
health_check_metrics: metrics,
original_max_size: max_size,
active_cooldowns: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
replacement_cooldown: server.replacement_cooldown,
is_shutting_down: Arc::new(std::sync::atomic::AtomicBool::new(false)),
})
}
pub async fn get_pooled_connection(
&self,
) -> Result<managed::Object<TcpManager>, crate::connection_error::ConnectionError> {
use crate::connection_error::ConnectionError;
self.pool.get().await.map_err(|e| {
let status = self.pool.status();
let err = match e {
deadpool::managed::PoolError::Backend(conn_err) => conn_err,
_other => ConnectionError::PoolExhausted {
backend: self.name.to_string(),
max_size: status.max_size,
},
};
warn!(
pool = %self.name,
max_size = status.max_size,
available = status.available,
current_size = status.size,
waiting = status.waiting,
error = %err,
"Connection acquisition failed"
);
err
})
}
pub fn clear_idle_connections(&self) {
let available = self.pool.status().available;
if available > 0 {
debug!(
pool = %self.name,
available = available,
"Clearing idle connections from pool"
);
let cooldowns = self.active_cooldowns.load(Ordering::Acquire);
let target_max = self.original_max_size.saturating_sub(cooldowns).max(1);
self.pool.resize(0);
self.pool.resize(target_max);
}
}
pub fn remove_without_cooldown(&self, conn: managed::Object<TcpManager>) {
shutdown_and_drop(conn);
}
pub fn remove_with_cooldown(&self, conn: managed::Object<TcpManager>) {
if self.is_shutting_down.load(Ordering::Acquire) {
shutdown_and_drop(conn);
return;
}
let Some(cooldown) = self.replacement_cooldown.filter(|d| !d.is_zero()) else {
shutdown_and_drop(conn);
return;
};
let max_reduction = self.original_max_size / 2;
let current = self.active_cooldowns.load(Ordering::Acquire);
if current >= max_reduction {
shutdown_and_drop(conn);
return;
}
let cooldowns = self.active_cooldowns.fetch_add(1, Ordering::AcqRel) + 1;
let new_max = self.original_max_size.saturating_sub(cooldowns).max(1);
resize_then_drop(&self.pool, conn, new_max);
warn!(
pool = %self.name,
new_max_size = new_max,
original_max_size = self.original_max_size,
active_cooldowns = cooldowns,
cooldown_secs = cooldown.as_secs(),
"Connection removed, pool size temporarily reduced"
);
let active_cooldowns = self.active_cooldowns.clone();
let pool = self.pool.clone();
let original_max = self.original_max_size;
let name = self.name.clone();
tokio::spawn(async move {
struct CooldownGuard {
active_cooldowns: Arc<std::sync::atomic::AtomicUsize>,
pool: Pool,
original_max: usize,
name: Arc<str>,
}
impl Drop for CooldownGuard {
fn drop(&mut self) {
let remaining = self.active_cooldowns.fetch_sub(1, Ordering::AcqRel) - 1;
let restored_max = self.original_max.saturating_sub(remaining);
self.pool.resize(restored_max);
debug!(
pool = %self.name,
restored_max_size = restored_max,
remaining_cooldowns = remaining,
"Connection cooldown expired, pool size restored"
);
}
}
let _guard = CooldownGuard {
active_cooldowns,
pool,
original_max,
name,
};
tokio::time::sleep(cooldown).await;
});
}
#[must_use]
#[inline]
pub fn max_size(&self) -> usize {
self.pool.status().max_size
}
#[must_use]
#[inline]
pub fn name(&self) -> &str {
self.name.as_ref()
}
#[must_use]
pub fn status_counts(&self) -> DeadpoolStatusCounts {
let status = self.pool.status();
DeadpoolStatusCounts {
available: status.available,
size: status.size,
max_size: status.max_size,
waiting: status.waiting,
}
}
#[must_use]
#[inline]
pub fn host(&self) -> &str {
&self.pool.manager().host
}
#[must_use]
#[inline]
pub fn port(&self) -> u16 {
self.pool.manager().port
}
#[must_use]
pub fn health_check_metrics(&self) -> &HealthCheckMetrics {
&self.health_check_metrics
}
pub fn shutdown(&self) {
self.is_shutting_down.store(true, Ordering::Release);
if let Some(tx) = &self.shutdown_tx {
let _ = tx.send(());
}
}
async fn run_periodic_health_checks(
pool: Pool,
name: String,
interval: std::time::Duration,
mut shutdown_rx: broadcast::Receiver<()>,
metrics: Arc<HealthCheckMetrics>,
) {
use crate::constants::pool::{
HEALTH_CHECK_POOL_TIMEOUT_MS, MAX_CONNECTIONS_PER_HEALTH_CHECK_CYCLE,
};
use tokio::time::{Duration, sleep};
info!(
pool = %name,
interval_secs = interval.as_secs(),
"Starting periodic health checks"
);
loop {
tokio::select! {
() = sleep(interval) => {
}
_ = shutdown_rx.recv() => {
info!(pool = %name, "Shutting down periodic health check task");
break;
}
}
let status = pool.status();
if status.available == 0 {
continue;
}
debug!(
pool = %name,
available = status.available,
max_check = MAX_CONNECTIONS_PER_HEALTH_CHECK_CYCLE,
"Running health check cycle"
);
let check_count =
std::cmp::min(status.available, MAX_CONNECTIONS_PER_HEALTH_CHECK_CYCLE);
let mut checked = 0;
let mut failed = 0;
let mut timeouts = managed::Timeouts::new();
timeouts.wait = Some(Duration::from_millis(HEALTH_CHECK_POOL_TIMEOUT_MS));
for _ in 0..check_count {
if let Ok(mut conn_obj) = pool.timeout_get(&timeouts).await {
checked += 1;
if let Err(e) = check_date_response(&mut *conn_obj).await {
failed += 1;
warn!(
pool = %name,
error = %e,
"Health check failed, discarding connection"
);
let _ = socket2::SockRef::from(conn_obj.underlying_tcp_stream())
.shutdown(std::net::Shutdown::Both);
}
drop(conn_obj);
} else {
break;
}
}
if checked > 0 {
metrics.record_cycle(checked, failed);
debug!(
pool = %name,
checked = checked,
failed = failed,
"Health check cycle complete"
);
}
}
info!(pool = %name, "Periodic health check task terminated");
}
pub async fn graceful_shutdown(&self) {
use deadpool::managed::Object;
use tokio::io::AsyncWriteExt;
self.shutdown();
let status = self.pool.status();
info!(
"Shutting down pool '{}' ({} idle connections)",
self.name, status.available
);
let mut timeouts = managed::Timeouts::new();
timeouts.wait = Some(crate::constants::timeout::SHUTDOWN_POOL_GET);
let mut idle_connections = Vec::with_capacity(status.available);
for _ in 0..status.available {
if let Ok(conn_obj) = self.pool.timeout_get(&timeouts).await {
idle_connections.push(Object::take(conn_obj));
} else {
break;
}
}
shutdown_connections_concurrently(idle_connections, |mut conn| async move {
let _ = tokio::time::timeout(
crate::constants::timeout::SHUTDOWN_QUIT_WRITE,
conn.write_all(b"QUIT\r\n"),
)
.await;
})
.await;
self.pool.close();
}
}
async fn shutdown_connections_concurrently<T, F, Fut>(connections: Vec<T>, shutdown_one: F)
where
F: FnMut(T) -> Fut,
Fut: std::future::Future<Output = ()>,
{
let mut pending = connections
.into_iter()
.map(shutdown_one)
.collect::<FuturesUnordered<_>>();
while pending.next().await.is_some() {}
}
fn resize_then_drop(pool: &Pool, conn: managed::Object<TcpManager>, new_max: usize) {
pool.resize(new_max);
shutdown_and_drop(conn);
}
fn shutdown_and_drop(conn: managed::Object<TcpManager>) {
let _ = socket2::SockRef::from(conn.underlying_tcp_stream()).shutdown(std::net::Shutdown::Both);
drop(conn);
}
impl ConnectionProvider for DeadpoolConnectionProvider {
fn status(&self) -> PoolStatus {
use crate::types::{AvailableConnections, CreatedConnections, MaxPoolSize};
let status = self.pool.status();
PoolStatus {
available: AvailableConnections::new(status.available),
max_size: MaxPoolSize::new(status.max_size),
created: CreatedConnections::new(status.size),
}
}
}
#[cfg(test)]
#[allow(clippy::float_cmp)] mod tests {
use super::*;
#[test]
fn test_builder_new() {
let builder = Builder::new("news.example.com", 119);
assert_eq!(builder.host, "news.example.com");
assert_eq!(builder.port, 119);
assert_eq!(builder.max_size, 10); assert!(builder.name.is_none());
assert!(builder.username.is_none());
assert!(builder.password.is_none());
assert!(builder.tls_config.is_none());
}
#[test]
fn test_builder_with_name() {
let builder = Builder::new("example.com", 119).name("Test Server");
assert_eq!(builder.name, Some("Test Server".to_string()));
}
#[test]
fn test_builder_with_max_connections() {
let builder = Builder::new("example.com", 119).max_connections(25);
assert_eq!(builder.max_size, 25);
}
#[test]
fn test_builder_with_username() {
let builder = Builder::new("example.com", 119).username("testuser");
assert_eq!(builder.username, Some("testuser".to_string()));
}
#[test]
fn test_builder_with_password() {
let builder = Builder::new("example.com", 119).password("testpass");
assert_eq!(builder.password, Some("testpass".to_string()));
}
#[test]
fn test_builder_with_tls_config() {
let tls_config = TlsConfig::builder().enabled(true).build();
let builder = Builder::new("example.com", 563).tls_config(tls_config);
assert!(builder.tls_config.is_some());
}
#[test]
fn test_builder_chaining() {
let builder = Builder::new("news.example.com", 119)
.name("Chained Server")
.max_connections(30)
.username("user")
.password("pass");
assert_eq!(builder.name, Some("Chained Server".to_string()));
assert_eq!(builder.max_size, 30);
assert_eq!(builder.username, Some("user".to_string()));
assert_eq!(builder.password, Some("pass".to_string()));
}
#[test]
fn test_builder_default_name_from_host_port() {
let provider = Builder::new("test.example.com", 8119)
.max_connections(5)
.build()
.unwrap();
assert_eq!(provider.name(), "test.example.com:8119");
}
#[test]
fn test_builder_custom_name_used() {
let provider = Builder::new("test.example.com", 8119)
.name("Custom Name")
.build()
.unwrap();
assert_eq!(provider.name(), "Custom Name");
}
#[test]
fn test_provider_builder_method() {
let builder = DeadpoolConnectionProvider::builder("example.com", 119);
assert_eq!(builder.host, "example.com");
assert_eq!(builder.port, 119);
}
#[test]
fn test_provider_status_conversion() {
let provider = DeadpoolConnectionProvider::builder("localhost", 119)
.max_connections(15)
.build()
.unwrap();
let status = ConnectionProvider::status(&provider);
assert_eq!(status.max_size.get(), 15);
assert_eq!(status.created.get(), 0);
}
#[test]
fn test_provider_inherent_methods() {
let provider = DeadpoolConnectionProvider::builder("localhost", 119)
.build()
.unwrap();
assert_eq!(provider.name(), "localhost:119");
assert_eq!(provider.host(), "localhost");
assert_eq!(provider.port(), 119);
}
#[test]
fn test_provider_with_all_builder_options() {
let tls_config = TlsConfig::builder().enabled(false).build();
let provider = DeadpoolConnectionProvider::builder("news.test.com", 563)
.name("Full Test")
.max_connections(42)
.username("testuser")
.password("testpass")
.tls_config(tls_config)
.build()
.unwrap();
assert_eq!(provider.name(), "Full Test");
assert_eq!(provider.host(), "news.test.com");
assert_eq!(provider.port(), 563);
let status = ConnectionProvider::status(&provider);
assert_eq!(status.max_size.get(), 42);
}
#[test]
fn test_health_check_metrics_initialization() {
let provider = DeadpoolConnectionProvider::builder("localhost", 119)
.build()
.unwrap();
let metrics = &provider.health_check_metrics;
assert_eq!(metrics.cycles_run(), 0);
assert_eq!(metrics.connections_checked(), 0);
assert_eq!(metrics.connections_failed(), 0);
assert_eq!(metrics.failure_rate(), 0.0);
}
#[test]
fn test_builder_accepts_string_types() {
let _ = Builder::new("example.com", 119);
let _ = Builder::new(String::from("example.com"), 119);
let _ = Builder::new("example.com", 119).name("test");
let _ = Builder::new("example.com", 119).name(String::from("test"));
}
#[test]
fn test_builder_zero_max_connections() {
let provider = Builder::new("localhost", 119)
.max_connections(0)
.build()
.unwrap();
let status = ConnectionProvider::status(&provider);
assert_eq!(status.max_size.get(), 0);
}
#[test]
fn test_builder_large_max_connections() {
let provider = Builder::new("localhost", 119)
.max_connections(1000)
.build()
.unwrap();
let status = ConnectionProvider::status(&provider);
assert_eq!(status.max_size.get(), 1000);
}
#[test]
fn test_provider_name_special_characters() {
let provider = Builder::new("example.com", 119)
.name("Server-123_Test.Name")
.build()
.unwrap();
assert_eq!(provider.name(), "Server-123_Test.Name");
}
#[test]
fn test_provider_name_unicode() {
let provider = Builder::new("example.com", 119)
.name("测试服务器")
.build()
.unwrap();
assert_eq!(provider.name(), "测试服务器");
}
#[test]
fn test_provider_empty_name() {
let provider = Builder::new("example.com", 119).name("").build().unwrap();
assert_eq!(provider.name(), "");
}
#[test]
fn test_builder_idempotent_chaining() {
let builder = Builder::new("example.com", 119)
.name("First")
.name("Second")
.max_connections(10)
.max_connections(20);
assert_eq!(builder.name, Some("Second".to_string()));
assert_eq!(builder.max_size, 20);
}
async fn spawn_mock_nntp_server() -> std::net::SocketAddr {
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
while let Ok((stream, _)) = listener.accept().await {
tokio::spawn(async move {
let (read_half, mut write_half) = stream.into_split();
let mut reader = BufReader::new(read_half);
let _ = write_half.write_all(b"200 mock\r\n").await;
let mut line = String::new();
loop {
line.clear();
match reader.read_line(&mut line).await {
Ok(0) | Err(_) => break,
Ok(_) => {
let cmd = line.trim().to_ascii_uppercase();
if cmd == "COMPRESS DEFLATE" {
let _ = write_half.write_all(b"500 Not supported\r\n").await;
} else if cmd.starts_with("MODE") {
let _ = write_half.write_all(b"200 Posting allowed\r\n").await;
} else if cmd.starts_with("QUIT") {
let _ = write_half.write_all(b"205 Goodbye\r\n").await;
break;
} else if cmd.starts_with("DATE") {
let _ = write_half.write_all(b"111 20240101000000\r\n").await;
} else {
let _ = write_half.write_all(b"200 OK\r\n").await;
}
}
}
}
});
}
});
addr
}
fn provider_with_cooldown(
addr: std::net::SocketAddr,
max_size: usize,
cooldown: Option<std::time::Duration>,
) -> DeadpoolConnectionProvider {
let manager = TcpManager::new(
addr.ip().to_string(),
addr.port(),
format!("test-{}", addr.port()),
TcpManagerOptions {
compress: Some(false), ..TcpManagerOptions::default()
},
)
.unwrap();
let pool = Pool::builder(manager).max_size(max_size).build().unwrap();
DeadpoolConnectionProvider {
pool,
name: Arc::from(format!("test-{}", addr.port())),
shutdown_tx: None,
health_check_metrics: Arc::new(crate::pool::health_check::HealthCheckMetrics::new()),
original_max_size: max_size,
active_cooldowns: Arc::new(std::sync::atomic::AtomicUsize::new(0)),
replacement_cooldown: cooldown,
is_shutting_down: Arc::new(std::sync::atomic::AtomicBool::new(false)),
}
}
#[tokio::test]
async fn test_remove_with_cooldown_resize_before_drop() {
let addr = spawn_mock_nntp_server().await;
let cooldown = std::time::Duration::from_secs(10);
let max_size = 4;
let provider = provider_with_cooldown(addr, max_size, Some(cooldown));
let conn = provider.get_pooled_connection().await.unwrap();
assert_eq!(provider.pool.status().max_size, max_size);
provider.remove_with_cooldown(conn);
let status = provider.pool.status();
assert_eq!(
status.max_size,
max_size - 1,
"Pool max_size should be reduced to {} after remove_with_cooldown, got {}",
max_size - 1,
status.max_size
);
assert_eq!(provider.active_cooldowns.load(Ordering::Acquire), 1);
}
#[tokio::test]
async fn test_remove_with_cooldown_cap() {
let addr = spawn_mock_nntp_server().await;
let cooldown = std::time::Duration::from_secs(10);
let max_size = 4;
let provider = provider_with_cooldown(addr, max_size, Some(cooldown));
let conn1 = provider.get_pooled_connection().await.unwrap();
let conn2 = provider.get_pooled_connection().await.unwrap();
let conn3 = provider.get_pooled_connection().await.unwrap();
provider.remove_with_cooldown(conn1);
assert_eq!(provider.pool.status().max_size, 3);
provider.remove_with_cooldown(conn2);
assert_eq!(provider.pool.status().max_size, 2);
provider.remove_with_cooldown(conn3);
assert_eq!(
provider.pool.status().max_size,
2,
"Pool max_size should not drop below half (2) of original (4)"
);
}
#[tokio::test]
async fn test_remove_with_cooldown_disabled() {
let addr = spawn_mock_nntp_server().await;
let max_size = 4;
let provider = provider_with_cooldown(addr, max_size, None);
let conn = provider.get_pooled_connection().await.unwrap();
provider.remove_with_cooldown(conn);
assert_eq!(provider.pool.status().max_size, max_size);
assert_eq!(provider.active_cooldowns.load(Ordering::Acquire), 0);
}
#[tokio::test]
async fn test_remove_without_cooldown_preserves_pool_size() {
let addr = spawn_mock_nntp_server().await;
let cooldown = std::time::Duration::from_secs(10);
let max_size = 4;
let provider = provider_with_cooldown(addr, max_size, Some(cooldown));
let conn = provider.get_pooled_connection().await.unwrap();
provider.remove_without_cooldown(conn);
assert_eq!(provider.pool.status().max_size, max_size);
assert_eq!(provider.active_cooldowns.load(Ordering::Acquire), 0);
}
#[tokio::test]
async fn test_remove_with_cooldown_restores_after_timer() {
let addr = spawn_mock_nntp_server().await;
let cooldown = std::time::Duration::from_millis(100);
let max_size = 4;
let provider = provider_with_cooldown(addr, max_size, Some(cooldown));
let conn = provider.get_pooled_connection().await.unwrap();
provider.remove_with_cooldown(conn);
assert_eq!(provider.pool.status().max_size, 3);
assert_eq!(provider.active_cooldowns.load(Ordering::Acquire), 1);
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
assert_eq!(
provider.pool.status().max_size,
max_size,
"Pool max_size should be restored after cooldown expires"
);
assert_eq!(provider.active_cooldowns.load(Ordering::Acquire), 0);
}
#[tokio::test]
async fn test_remove_with_cooldown_skips_reduction_during_shutdown() {
let addr = spawn_mock_nntp_server().await;
let cooldown = std::time::Duration::from_secs(10);
let max_size = 4;
let provider = provider_with_cooldown(addr, max_size, Some(cooldown));
provider.shutdown();
let conn = provider.get_pooled_connection().await.unwrap();
provider.remove_with_cooldown(conn);
assert_eq!(provider.pool.status().max_size, max_size);
assert_eq!(provider.active_cooldowns.load(Ordering::Acquire), 0);
assert!(provider.is_shutting_down.load(Ordering::Acquire));
}
#[tokio::test(start_paused = true)]
async fn test_shutdown_connections_concurrently_runs_all_writes_together() {
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering};
let completed = Arc::new(AtomicUsize::new(0));
let completed_clone = completed.clone();
let task = tokio::spawn(async move {
shutdown_connections_concurrently(vec![1_u8, 2, 3], move |_| {
let completed = completed_clone.clone();
async move {
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
completed.fetch_add(1, AtomicOrdering::Relaxed);
}
})
.await;
});
tokio::task::yield_now().await;
tokio::time::advance(std::time::Duration::from_millis(499)).await;
assert_eq!(completed.load(AtomicOrdering::Relaxed), 0);
assert!(!task.is_finished());
tokio::time::advance(std::time::Duration::from_millis(1)).await;
task.await.unwrap();
assert_eq!(completed.load(AtomicOrdering::Relaxed), 3);
}
}