Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@ install: build
install -c -m 755 ./target/release/pg_doorman $(DESTDIR)/usr/bin/

test:
cargo test --lib
cargo test --lib --test message_memory

clippy:
cargo clippy -- --deny "warnings"
Expand Down
25 changes: 14 additions & 11 deletions src/messages/socket.rs
Original file line number Diff line number Diff line change
Expand Up @@ -129,20 +129,22 @@ where
}
}

/// Read a message, reserving memory until the read completes or is dropped.
/// If the future is dropped, discard the stream: part of a message may have been read.
#[inline]
pub async fn read_message<S>(stream: &mut S, max_memory_usage: u64) -> Result<BytesMut, Error>
where
S: tokio::io::AsyncRead + std::marker::Unpin,
{
let (code, len) = read_message_header(stream).await?;
let prev = CURRENT_MEMORY.fetch_add(len as i64, Ordering::Relaxed);
if (prev + len as i64) as u64 > max_memory_usage {
let _reservation = scopeguard::guard(len, |len| {
CURRENT_MEMORY.fetch_sub(len as i64, Ordering::Relaxed);
});
if (prev + len as i64) as u64 > max_memory_usage {
return Err(Error::CurrentMemoryUsage);
}
let result = read_message_data(stream, code, len).await;
CURRENT_MEMORY.fetch_sub(len as i64, Ordering::Relaxed);
result
read_message_data(stream, code, len).await
}

/// Read a message into a reusable buffer. Returns the owned `BytesMut` via
Expand All @@ -154,6 +156,9 @@ where
/// allocation until exhausted. A buffer that grew past
/// `REUSE_BUF_SHRINK_THRESHOLD` is dropped before the next read, so a single
/// oversized message does not pin its allocation across the connection.
///
/// Dropping the future releases its memory reservation, but may leave a partial
/// message consumed. Discard the stream afterwards.
#[inline]
pub async fn read_message_reuse<S>(
stream: &mut S,
Expand All @@ -177,8 +182,10 @@ where
}

let prev = CURRENT_MEMORY.fetch_add(len as i64, Ordering::Relaxed);
if (prev + len as i64) as u64 > max_memory_usage {
let _reservation = scopeguard::guard(len, |len| {
CURRENT_MEMORY.fetch_sub(len as i64, Ordering::Relaxed);
});
if (prev + len as i64) as u64 > max_memory_usage {
return Err(Error::CurrentMemoryUsage);
}

Expand All @@ -190,16 +197,12 @@ where
buf.put_i32(len);
buf.resize(total_len, 0);

let result = match stream.read_exact(&mut buf[5..]).await {
match stream.read_exact(&mut buf[5..]).await {
Ok(_) => Ok(buf.split()),
Err(err) => Err(Error::SocketError(format!(
"Error reading message data from socket - Code: {code:?}, Error: {err:?}"
))),
};

CURRENT_MEMORY.fetch_sub(len as i64, Ordering::Relaxed);

result
}
}

/// Read message body into a reusable buffer when header is already consumed.
Expand Down
58 changes: 58 additions & 0 deletions tests/bdd/extended/session_management.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ use crate::pg_connection::PgConnection;
use crate::world::DoormanWorld;
use cucumber::{then, when};
use log::info;
use pg_doorman::messages::PgErrorMsg;
use std::time::Duration;
use tokio::task::JoinSet;
use tokio::time::timeout;
Expand Down Expand Up @@ -540,6 +541,63 @@ pub async fn send_simple_query_to_session_without_waiting(
.expect("Failed to send query");
}

#[when(regex = r#"^we send a SimpleQuery header with length (\d+) to session "([^"]+)"$"#)]
pub async fn send_simple_query_header(world: &mut DoormanWorld, len: i32, session_name: String) {
let conn = super::helpers::get_session(&mut world.named_sessions, &session_name);
conn.send_message_header(b'Q', len)
.await
.expect("Failed to send query header");
}

#[when(regex = r#"^we send SimpleQuery "([^"]+)" padded to (\d+) bytes to session "([^"]+)"$"#)]
pub async fn send_padded_simple_query(
world: &mut DoormanWorld,
query: String,
len: usize,
session_name: String,
) {
let padded = pad_simple_query(&query, len);
let conn = super::helpers::get_session(&mut world.named_sessions, &session_name);
conn.send_simple_query(&padded)
.await
.expect("Failed to send padded query");
}

fn pad_simple_query(query: &str, len: usize) -> String {
// The protocol length includes its four bytes and the query's trailing NUL.
assert!(len >= query.len() + 5);
format!("{query}{}", " ".repeat(len - query.len() - 5))
}

#[when(
regex = r#"^we wait for a (\d+)-byte query on session "([^"]+)" to exceed the memory limit$"#
)]
pub async fn wait_for_memory_limit(world: &mut DoormanWorld, len: usize, session_name: String) {
let conn = super::helpers::get_session(&mut world.named_sessions, &session_name);
let query = pad_simple_query("SELECT 1", len);
timeout(Duration::from_secs(5), async {
loop {
conn.send_simple_query(&query)
.await
.expect("Failed to send memory limit probe");
let messages = conn
.read_all_messages_until_ready()
.await
.expect("Failed to read memory limit probe response");
for (code, data) in messages {
if code == 'E' {
let error = PgErrorMsg::parse(&data).expect("Invalid ErrorResponse");
assert_eq!(error.code, "53200", "Unexpected error: {error}");
return;
}
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
})
.await
.expect("Memory limit was not reached within 5 seconds");
}

#[when(
regex = r#"^we send SimpleQuery "([^"]+)" to (\d+) sessions with prefix "([^"]+)" without waiting$"#
)]
Expand Down
44 changes: 44 additions & 0 deletions tests/bdd/features/stale-server-detection.feature
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,50 @@ Feature: Detect stale server connections during client idle in transaction
Then we read SimpleQuery response from session "client2" within 5000ms
Then session "client2" should receive DataRow with "2"

@cancelled-read-memory
Scenario: A backend exit releases memory reserved for an incomplete client message
Given pg_doorman started with config:
"""
[general]
host = "127.0.0.1"
port = ${DOORMAN_PORT}
admin_username = "admin"
admin_password = "admin"
pg_hba.content = "host all all 127.0.0.1/32 trust"
max_memory_usage = 1536
worker_threads = 1
server_idle_check_timeout = 0

[pools.example_db]
server_host = "127.0.0.1"
server_port = ${PG_PORT}

[[pools.example_db.users]]
username = "example_user_1"
password = ""
pool_size = 1

[[pools.example_db.users]]
username = "postgres"
password = ""
pool_size = 1
"""
When we create session "victim" to pg_doorman as "example_user_1" with password "" and database "example_db"
And we send SimpleQuery "BEGIN" to session "victim" without waiting
Then we read SimpleQuery response from session "victim" within 2000ms
When we send SimpleQuery "SELECT pg_backend_pid()" to session "victim" and store backend_pid as "victim_pid"
And we create session "killer" to pg_doorman as "postgres" with password "" and database "example_db"
And we send a SimpleQuery header with length 1024 to session "victim"
And we create session "pressure" to pg_doorman as "postgres" with password "" and database "example_db"
And we wait for a 1024-byte query on session "pressure" to exceed the memory limit
And we terminate backend "victim_pid" from session "victim" via session "killer"
Then we read SimpleQuery response from session "victim" within 5000ms
And session "victim" should receive error containing "server closed the connection unexpectedly" with code "08006"
When we create session "next" to pg_doorman as "example_user_1" with password "" and database "example_db"
And we send SimpleQuery "SELECT 2" padded to 1024 bytes to session "next"
Then we read SimpleQuery response from session "next" within 5000ms
And session "next" should receive DataRow with "2"

@stale-server-victim-connection-closed
Scenario: Victim client connection is closed after server is killed
# pg_doorman detects the dead server, sends ErrorResponse to the victim,
Expand Down
7 changes: 7 additions & 0 deletions tests/bdd/pg_connection.rs
Original file line number Diff line number Diff line change
Expand Up @@ -620,6 +620,13 @@ impl PgConnection {
Ok(())
}

pub async fn send_message_header(&mut self, code: u8, len: i32) -> tokio::io::Result<()> {
let mut header = [0; 5];
header[0] = code;
header[1..].copy_from_slice(&len.to_be_bytes());
self.stream.write_all(&header).await
}

pub async fn read_all_messages_until_ready(
&mut self,
) -> tokio::io::Result<Vec<(char, Vec<u8>)>> {
Expand Down
176 changes: 176 additions & 0 deletions tests/message_memory.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,176 @@
use std::io::Cursor;
use std::sync::atomic::Ordering;

use bytes::BytesMut;
use pg_doorman::errors::Error;
use pg_doorman::messages::{read_message, read_message_reuse, CURRENT_MEMORY};
use serial_test::serial;
use tokio::io::{AsyncRead, AsyncWriteExt};
use tokio::time::{timeout, Duration};

fn message(len: i32) -> Vec<u8> {
assert!(len >= 4);
let mut message = vec![b'Q'];
message.extend_from_slice(&len.to_be_bytes());
message.resize(len as usize + 1, b'x');
message
}

async fn read<S: AsyncRead + Unpin>(
stream: &mut S,
buffer: Option<&mut BytesMut>,
limit: u64,
) -> Result<BytesMut, Error> {
match buffer {
Some(buffer) => read_message_reuse(stream, buffer, limit).await,
None => read_message(stream, limit).await,
}
}

// A separate test binary keeps CURRENT_MEMORY isolated from other socket tests.
#[tokio::test]
#[serial]
async fn cancelled_reads_release_memory() {
let message = message(100);

for mut buffer in [None, Some(BytesMut::new())] {
for prefix_len in 0..message.len() {
let (mut writer, mut reader) = tokio::io::duplex(message.len());
writer.write_all(&message[..prefix_len]).await.unwrap();

{
// Pending must mean missing input, not an exhausted Tokio task budget.
let mut pending = std::pin::pin!(tokio::task::unconstrained(read(
&mut reader,
buffer.as_mut(),
100
)));
assert!(futures::poll!(pending.as_mut()).is_pending());
assert_eq!(
CURRENT_MEMORY.load(Ordering::Relaxed),
if prefix_len >= 5 { 100 } else { 0 },
"prefix length: {prefix_len}"
);
}

assert_eq!(CURRENT_MEMORY.load(Ordering::Relaxed), 0);

let result = read(&mut Cursor::new(&message), buffer.as_mut(), 100)
.await
.unwrap();
assert_eq!(result.as_ref(), message.as_slice());
assert_eq!(CURRENT_MEMORY.load(Ordering::Relaxed), 0);
}
}
}

#[tokio::test]
#[serial]
async fn completed_and_failed_reads_release_memory() {
let message = message(100);

for mut buffer in [None, Some(BytesMut::new())] {
let result = read(&mut Cursor::new(&message), buffer.as_mut(), 100)
.await
.unwrap();
assert_eq!(result.as_ref(), message.as_slice());
assert_eq!(CURRENT_MEMORY.load(Ordering::Relaxed), 0);

let result = read(&mut Cursor::new(&message), buffer.as_mut(), 99).await;
assert!(matches!(result, Err(Error::CurrentMemoryUsage)));
assert_eq!(CURRENT_MEMORY.load(Ordering::Relaxed), 0);

let result = read(&mut Cursor::new(&message[..6]), buffer.as_mut(), 100).await;
assert!(matches!(result, Err(Error::SocketError(_))));
assert_eq!(CURRENT_MEMORY.load(Ordering::Relaxed), 0);
}
}

#[tokio::test]
#[serial]
async fn cancelling_one_read_preserves_other_reservations() {
for reuse in [false, true] {
let (mut first_writer, mut first_reader) = tokio::io::duplex(101);
let (mut second_writer, mut second_reader) = tokio::io::duplex(101);
first_writer.write_all(&message(60)[..5]).await.unwrap();
second_writer.write_all(&message(100)[..5]).await.unwrap();
let mut first_buffer = reuse.then(BytesMut::new);
let mut second_buffer = reuse.then(BytesMut::new);
let mut first = Box::pin(read(&mut first_reader, first_buffer.as_mut(), 160));
let mut second = Box::pin(read(&mut second_reader, second_buffer.as_mut(), 160));

assert!(futures::poll!(first.as_mut()).is_pending());
assert!(futures::poll!(second.as_mut()).is_pending());
assert_eq!(CURRENT_MEMORY.load(Ordering::Relaxed), 160);

drop(first);
assert_eq!(CURRENT_MEMORY.load(Ordering::Relaxed), 100);

let mut next_buffer = reuse.then(BytesMut::new);
let rejected = read(&mut Cursor::new(message(61)), next_buffer.as_mut(), 160).await;
assert!(matches!(rejected, Err(Error::CurrentMemoryUsage)));
assert_eq!(CURRENT_MEMORY.load(Ordering::Relaxed), 100);

let next = message(60);
assert_eq!(
read(&mut Cursor::new(&next), next_buffer.as_mut(), 160)
.await
.unwrap()
.as_ref(),
next.as_slice()
);
assert_eq!(CURRENT_MEMORY.load(Ordering::Relaxed), 100);

second_writer.write_all(&message(100)[5..]).await.unwrap();
assert_eq!(second.await.unwrap().as_ref(), message(100));
assert_eq!(CURRENT_MEMORY.load(Ordering::Relaxed), 0);
}
}

#[tokio::test]
#[serial]
async fn resumed_reads_keep_their_reservation() {
let message = message(100);
for mut buffer in [None, Some(BytesMut::new())] {
let (mut writer, mut reader) = tokio::io::duplex(message.len());
writer.write_all(&message[..6]).await.unwrap();
let mut pending = std::pin::pin!(read(&mut reader, buffer.as_mut(), 100));

tokio::select! {
biased;
result = &mut pending => panic!("incomplete read finished: {result:?}"),
_ = std::future::ready(()) => {}
}
assert_eq!(CURRENT_MEMORY.load(Ordering::Relaxed), 100);

writer.write_all(&message[6..]).await.unwrap();
assert_eq!(pending.await.unwrap().as_ref(), message.as_slice());
assert_eq!(CURRENT_MEMORY.load(Ordering::Relaxed), 0);
}
}

#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
#[serial]
async fn aborted_tasks_release_memory() {
for mut buffer in [None, Some(BytesMut::new())] {
let (mut writer, mut reader) = tokio::io::duplex(101);
writer.write_all(&message(100)[..6]).await.unwrap();
let (ready_tx, ready_rx) = tokio::sync::oneshot::channel();
let task = tokio::spawn(async move {
let mut pending = std::pin::pin!(read(&mut reader, buffer.as_mut(), 100));
assert!(futures::poll!(pending.as_mut()).is_pending());
ready_tx.send(()).unwrap();
pending.await
});

timeout(Duration::from_secs(5), ready_rx)
.await
.unwrap()
.unwrap();
assert_eq!(CURRENT_MEMORY.load(Ordering::Relaxed), 100);
task.abort();
let result = timeout(Duration::from_secs(5), task).await.unwrap();
assert!(result.unwrap_err().is_cancelled());
assert_eq!(CURRENT_MEMORY.load(Ordering::Relaxed), 0);
}
}