use super::config::PoolConfig;
use super::connection::PooledConn;
use super::lifecycle::PgPool;
use super::tests::socket_pair_connection;
use std::future::Future;
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
fn answer_queries(mut peer: tokio::net::UnixStream) {
tokio::spawn(async move {
let mut head = [0u8; 5];
while peer.read_exact(&mut head).await.is_ok() {
let len = u32::from_be_bytes([head[1], head[2], head[3], head[4]]) as usize;
let mut payload = vec![0u8; len.saturating_sub(4)];
if peer.read_exact(&mut payload).await.is_err() {
return;
}
if head[0] != b'Q' {
continue;
}
let sql =
std::str::from_utf8(&payload[..payload.len().saturating_sub(1)]).unwrap_or("");
let marker = sql
.strip_prefix("SELECT '")
.and_then(|rest| rest.strip_suffix('\''));
let tag: &[u8] = if marker.is_some() {
b"SELECT 1\0"
} else {
b"ROLLBACK\0"
};
let mut reply = Vec::new();
if let Some(token) = marker {
reply.push(b'D');
reply.extend_from_slice(&((4 + 2 + 4 + token.len()) as u32).to_be_bytes());
reply.extend_from_slice(&1i16.to_be_bytes());
reply.extend_from_slice(&(token.len() as i32).to_be_bytes());
reply.extend_from_slice(token.as_bytes());
}
reply.push(b'C');
reply.extend_from_slice(&(4 + tag.len() as u32).to_be_bytes());
reply.extend_from_slice(tag);
reply.extend_from_slice(&[b'Z', 0, 0, 0, 5, b'I']);
if peer.write_all(&reply).await.is_err() {
return;
}
}
});
}
async fn pool(max: usize, reserve: &[usize], acquire_timeout: Duration) -> PgPool {
let pool = PgPool::connect(
PoolConfig::new_dev("localhost", 5432, "user", "db")
.min_connections(0)
.max_connections(max)
.nested_reserve(reserve)
.acquire_timeout(acquire_timeout),
)
.await
.expect("pool init");
let mut idle = pool.inner.connections.lock().await;
for _ in 0..max {
let (conn, peer) = socket_pair_connection();
answer_queries(peer);
idle.push(PooledConn {
conn,
created_at: Instant::now(),
last_used: Instant::now(),
});
}
drop(idle);
pool
}
fn free(pool: &PgPool) -> (usize, Vec<usize>) {
(
pool.inner.semaphore.available_permits(),
pool.inner
.nested_semaphores
.iter()
.map(|reserve| reserve.available_permits())
.collect(),
)
}
fn registered_tasks(pool: &PgPool) -> usize {
pool.inner.level_holders.lock().expect("holders").len()
}
async fn in_task<F>(scenario: F) -> F::Output
where
F: Future + Send + 'static,
F::Output: Send + 'static,
{
tokio::spawn(scenario).await.expect("scenario task")
}
async fn hold_one_then_ask_for_a_second(pool: &PgPool) -> Vec<Result<(), String>> {
let barrier = Arc::new(tokio::sync::Barrier::new(2));
let mut tasks = Vec::new();
for _ in 0..2 {
let pool = pool.clone();
let barrier = Arc::clone(&barrier);
tasks.push(tokio::spawn(async move {
let first = pool.acquire_raw().await.map_err(|e| e.to_string())?;
barrier.wait().await;
let outcome = match pool.acquire_raw().await {
Ok(second) => {
second.release().await;
Ok(())
}
Err(e) => Err(e.to_string()),
};
first.release().await;
outcome
}));
}
let mut outcomes = Vec::new();
for task in tasks {
outcomes.push(task.await.expect("task"));
}
outcomes
}
#[tokio::test]
async fn test_holders_waiting_for_a_second_connection_deadlock_without_a_reserve() {
let pool = pool(2, &[], Duration::from_millis(300)).await;
let outcomes = hold_one_then_ask_for_a_second(&pool).await;
for outcome in &outcomes {
let err = outcome
.as_ref()
.expect_err("each second acquire waits on a slot the other task holds");
assert!(err.contains("pool acquire after"), "{err}");
}
}
#[tokio::test]
async fn test_a_reserve_lets_both_holders_get_their_second_connection() {
let pool = pool(3, &[1], Duration::from_secs(2)).await;
let outcomes = hold_one_then_ask_for_a_second(&pool).await;
assert_eq!(outcomes, vec![Ok(()), Ok(())]);
assert_eq!(
free(&pool),
(2, vec![1]),
"every slot came back to its level"
);
assert_eq!(registered_tasks(&pool), 0, "no task is left registered");
}
#[tokio::test]
async fn test_a_nested_acquire_takes_the_reserve_and_leaves_the_shared_slots() {
let pool = pool(3, &[1], Duration::from_secs(2)).await;
let scenario = pool.clone();
in_task(async move {
let pool = scenario;
let first = pool.acquire_raw().await.expect("first");
assert_eq!(free(&pool), (1, vec![1]));
let second = pool.acquire_raw().await.expect("second, nested");
assert_eq!(
free(&pool),
(1, vec![0]),
"the nested acquire took the reserve"
);
let other = pool.clone();
let shared = in_task(async move {
let conn = other.acquire_raw().await.map_err(|e| e.to_string())?;
conn.release().await;
Ok::<(), String>(())
})
.await;
assert_eq!(shared, Ok(()));
second.release().await;
first.release().await;
})
.await;
assert_eq!(free(&pool), (2, vec![1]));
assert_eq!(registered_tasks(&pool), 0);
}
#[tokio::test]
async fn test_join_in_one_task_claims_distinct_levels() {
let pool = pool(2, &[1], Duration::from_millis(500)).await;
let scenario = pool.clone();
in_task(async move {
let pool = scenario;
let (a, b) = tokio::join!(pool.acquire_raw(), pool.acquire_raw());
let (a, b) = (a.expect("first"), b.expect("second"));
assert_eq!(free(&pool), (0, vec![0]));
a.release().await;
b.release().await;
})
.await;
assert_eq!(free(&pool), (1, vec![1]));
assert_eq!(registered_tasks(&pool), 0);
}
#[tokio::test]
async fn test_a_cancelled_nested_acquire_leaves_no_claim() {
let pool = pool(3, &[1], Duration::from_secs(5)).await;
let scenario = pool.clone();
in_task(async move {
let pool = scenario;
let first = pool.acquire_raw().await.expect("first");
let second = pool.acquire_raw().await.expect("second, nested");
let third = tokio::time::timeout(Duration::from_millis(100), pool.acquire_raw()).await;
assert!(third.is_err(), "the last reserve is held");
{
let registered = pool.inner.level_holders.lock().expect("holders");
let counts = registered.values().next().expect("this task");
assert_eq!(counts, &vec![1, 1], "the cancelled claim is gone");
}
second.release().await;
first.release().await;
})
.await;
assert_eq!(registered_tasks(&pool), 0);
assert_eq!(free(&pool), (2, vec![1]));
}
#[tokio::test]
async fn test_acquires_past_the_reserves_share_the_last_one() {
let pool = pool(3, &[1], Duration::from_millis(200)).await;
let err = in_task(async move {
let first = pool.acquire_raw().await.expect("first");
let second = pool.acquire_raw().await.expect("second, nested");
let err = pool
.acquire_raw()
.await
.err()
.expect("the last reserve is held by this task")
.to_string();
second.release().await;
first.release().await;
err
})
.await;
assert!(
err.contains("pool acquire after") && err.contains("nested level 1"),
"{err}"
);
}
#[tokio::test]
async fn test_nested_reserve_is_validated() {
let zero = PgPool::connect(
PoolConfig::new_dev("localhost", 5432, "user", "db")
.min_connections(0)
.max_connections(4)
.nested_reserve(&[1, 0]),
)
.await;
assert!(
zero.is_err(),
"a level with no slot never serves an acquire"
);
let all = PgPool::connect(
PoolConfig::new_dev("localhost", 5432, "user", "db")
.min_connections(0)
.max_connections(3)
.nested_reserve(&[2, 1]),
)
.await;
assert!(all.is_err(), "a reserve must leave shared slots");
}
#[tokio::test]
async fn test_without_a_reserve_nothing_is_tracked() {
let pool = pool(2, &[], Duration::from_millis(500)).await;
let scenario = pool.clone();
in_task(async move {
let pool = scenario;
let first = pool.acquire_raw().await.expect("first");
let second = pool.acquire_raw().await.expect("second");
assert_eq!(free(&pool), (0, vec![]));
assert_eq!(registered_tasks(&pool), 0);
second.release().await;
first.release().await;
})
.await;
}