use deadpool::managed::Object;
use crate::constants::pool::HEALTH_CHECK_TIMEOUT;
use crate::pool::deadpool_connection::TcpManager;
use crate::pool::provider::DeadpoolConnectionProvider;
pub const MAX_CONNECTION_SALVAGE_MS: u64 = 1_000;
#[allow(dead_code)] const fn assert_no_timeout_loop(max_iterations: usize, timeout_per_iteration_ms: u64) {
assert!(
max_iterations == 1,
"Connection salvage MUST NOT use timeout loops (max_iterations must be 1)"
);
assert!(
timeout_per_iteration_ms <= MAX_CONNECTION_SALVAGE_MS,
"Single timeout must be <= MAX_CONNECTION_SALVAGE_MS"
);
}
const _SALVAGE_NO_LOOP: () = {
assert!(
HEALTH_CHECK_TIMEOUT.as_millis() <= MAX_CONNECTION_SALVAGE_MS as u128,
"HEALTH_CHECK_TIMEOUT must be <= MAX_CONNECTION_SALVAGE_MS"
);
assert_no_timeout_loop(1, MAX_CONNECTION_SALVAGE_MS);
};
pub struct ConnectionGuard {
conn: Option<Object<TcpManager>>,
provider: DeadpoolConnectionProvider,
released: bool,
}
impl ConnectionGuard {
pub const fn new(conn: Object<TcpManager>, provider: DeadpoolConnectionProvider) -> Self {
Self {
conn: Some(conn),
provider,
released: false,
}
}
pub fn release(mut self) -> Object<TcpManager> {
self.released = true;
self.conn
.take()
.expect("ConnectionGuard::release() called on consumed guard")
}
pub fn retire_without_cooldown(mut self) {
self.released = true;
let conn = self
.conn
.take()
.expect("ConnectionGuard::retire_without_cooldown() called on consumed guard");
self.provider.remove_without_cooldown(conn);
}
pub fn retire_with_cooldown(mut self) {
self.released = true;
let conn = self
.conn
.take()
.expect("ConnectionGuard::retire_with_cooldown() called on consumed guard");
self.provider.remove_with_cooldown(conn);
}
pub const fn get_mut(&mut self) -> &mut Object<TcpManager> {
self.conn
.as_mut()
.expect("ConnectionGuard already consumed")
}
pub const fn get(&self) -> &Object<TcpManager> {
self.conn
.as_ref()
.expect("ConnectionGuard already consumed")
}
#[must_use]
pub fn connection_type(&self) -> &'static str {
self.get().connection_type()
}
#[must_use]
pub fn pending_bytes_len(&self) -> usize {
self.get().pending_bytes_len()
}
#[must_use]
pub fn provider_status_counts(&self) -> crate::pool::provider::DeadpoolStatusCounts {
self.provider.status_counts()
}
#[must_use]
pub fn provider_name(&self) -> &str {
self.provider.name()
}
}
impl Drop for ConnectionGuard {
fn drop(&mut self) {
if !self.released
&& let Some(conn) = self.conn.take()
{
tracing::debug!(
connection_type = conn.connection_type(),
pending_bytes = conn.pending_bytes_len(),
"ConnectionGuard dropping unreleased pooled connection; removing backend connection with cooldown"
);
self.provider.remove_with_cooldown(conn);
}
}
}
impl std::ops::Deref for ConnectionGuard {
type Target = Object<TcpManager>;
fn deref(&self) -> &Self::Target {
self.get()
}
}
impl std::ops::DerefMut for ConnectionGuard {
fn deref_mut(&mut self) -> &mut Self::Target {
self.get_mut()
}
}
pub async fn salvage_with_health_check(
mut conn: Object<TcpManager>,
provider: DeadpoolConnectionProvider,
) {
use tracing::{debug, warn};
match crate::pool::health_check::check_date_response(&mut *conn).await {
Ok(()) => {
debug!("Connection salvaged after Invalid response - DATE check passed");
drop(conn); }
Err(e) => {
warn!("DATE health check failed after Invalid response: {}", e);
provider.remove_with_cooldown(conn);
}
}
}
#[cfg(test)]
mod tests {
use super::ConnectionGuard;
use std::sync::Arc;
use std::sync::atomic::{AtomicUsize, Ordering};
use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader};
use tokio::net::TcpListener;
async fn spawn_greeting_server() -> (u16, Arc<AtomicUsize>) {
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let port = listener.local_addr().unwrap().port();
let accept_count = Arc::new(AtomicUsize::new(0));
let count = Arc::clone(&accept_count);
tokio::spawn(async move {
while let Ok((stream, _)) = listener.accept().await {
count.fetch_add(1, Ordering::SeqCst);
tokio::spawn(async move {
let (read_half, mut write_half) = stream.into_split();
let mut reader = BufReader::new(read_half);
if write_half.write_all(b"200 Ready\r\n").await.is_err() {
return;
}
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;
}
}
}
}
});
}
});
(port, accept_count)
}
fn make_provider(port: u16) -> crate::pool::DeadpoolConnectionProvider {
crate::pool::DeadpoolConnectionProvider::builder("127.0.0.1", port)
.max_connections(5)
.build()
.unwrap()
}
#[tokio::test]
async fn release_reuses_pool_connection() {
let (port, accept_count) = spawn_greeting_server().await;
let provider = make_provider(port);
let conn = provider.get_pooled_connection().await.unwrap();
assert_eq!(accept_count.load(Ordering::SeqCst), 1);
let guard = ConnectionGuard::new(conn, provider.clone());
drop(guard.release());
let _conn2 = provider.get_pooled_connection().await.unwrap();
assert_eq!(
accept_count.load(Ordering::SeqCst),
1,
"release() must return connection to pool; next get() must reuse it without \
creating a new TCP connection"
);
}
#[tokio::test]
async fn drop_without_release_forces_new_connection() {
let (port, accept_count) = spawn_greeting_server().await;
let provider = make_provider(port);
let conn = provider.get_pooled_connection().await.unwrap();
assert_eq!(accept_count.load(Ordering::SeqCst), 1);
let guard = ConnectionGuard::new(conn, provider.clone());
drop(guard);
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
let _conn2 = provider.get_pooled_connection().await.unwrap();
assert_eq!(
accept_count.load(Ordering::SeqCst),
2,
"drop without release() must remove connection; next get() must create \
a new TCP connection"
);
}
}