From d1ca43acc35635fe136359a27af1858126dcaa78 Mon Sep 17 00:00:00 2001 From: Evgeny Khodzitsky <45710942+ekhodzitsky@users.noreply.github.com> Date: Thu, 10 Sep 2026 16:12:02 +0300 Subject: [PATCH 1/2] fix: release message memory reservations on cancellation --- Makefile | 2 +- src/messages/socket.rs | 20 +++---- tests/bdd/extended/session_management.rs | 23 ++++++++ .../features/stale-server-detection.feature | 41 +++++++++++++ tests/bdd/pg_connection.rs | 7 +++ tests/message_memory.rs | 58 +++++++++++++++++++ 6 files changed, 139 insertions(+), 12 deletions(-) create mode 100644 tests/message_memory.rs diff --git a/Makefile b/Makefile index 53580a21f..f5dd05b02 100644 --- a/Makefile +++ b/Makefile @@ -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" diff --git a/src/messages/socket.rs b/src/messages/socket.rs index 787e2d477..d37153742 100644 --- a/src/messages/socket.rs +++ b/src/messages/socket.rs @@ -136,13 +136,13 @@ where { 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 { + scopeguard::defer! { 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 @@ -177,8 +177,10 @@ where } let prev = CURRENT_MEMORY.fetch_add(len as i64, Ordering::Relaxed); - if (prev + len as i64) as u64 > max_memory_usage { + scopeguard::defer! { CURRENT_MEMORY.fetch_sub(len as i64, Ordering::Relaxed); + } + if (prev + len as i64) as u64 > max_memory_usage { return Err(Error::CurrentMemoryUsage); } @@ -190,16 +192,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. diff --git a/tests/bdd/extended/session_management.rs b/tests/bdd/extended/session_management.rs index 1fb91ac33..40cd8fd08 100644 --- a/tests/bdd/extended/session_management.rs +++ b/tests/bdd/extended/session_management.rs @@ -540,6 +540,29 @@ 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, +) { + assert!(len >= query.len() + 5); + let padded = format!("{query}{}", " ".repeat(len - query.len() - 5)); + let conn = super::helpers::get_session(&mut world.named_sessions, &session_name); + conn.send_simple_query(&padded) + .await + .expect("Failed to send padded query"); +} + #[when( regex = r#"^we send SimpleQuery "([^"]+)" to (\d+) sessions with prefix "([^"]+)" without waiting$"# )] diff --git a/tests/bdd/features/stale-server-detection.feature b/tests/bdd/features/stale-server-detection.feature index cf8ea9052..581764f0f 100644 --- a/tests/bdd/features/stale-server-detection.feature +++ b/tests/bdd/features/stale-server-detection.feature @@ -58,6 +58,47 @@ 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 + 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 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, diff --git a/tests/bdd/pg_connection.rs b/tests/bdd/pg_connection.rs index 20be85fed..1f7f0f11f 100644 --- a/tests/bdd/pg_connection.rs +++ b/tests/bdd/pg_connection.rs @@ -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)>> { diff --git a/tests/message_memory.rs b/tests/message_memory.rs new file mode 100644 index 000000000..acbb75e66 --- /dev/null +++ b/tests/message_memory.rs @@ -0,0 +1,58 @@ +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 tokio::io::{AsyncRead, AsyncWriteExt}; + +async fn read( + stream: &mut S, + buffer: Option<&mut BytesMut>, + limit: u64, +) -> Result { + 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] +async fn cancelled_reads_release_memory() { + let mut message = vec![b'Q']; + message.extend_from_slice(&100_i32.to_be_bytes()); + message.extend_from_slice(&[b'x'; 96]); + + for mut buffer in [None, Some(BytesMut::new())] { + for prefix_len in [0, 1, 3, 5, 6, message.len() - 1] { + let (mut writer, mut reader) = tokio::io::duplex(message.len()); + writer.write_all(&message[..prefix_len]).await.unwrap(); + + { + let mut pending = std::pin::pin!(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 } + ); + } + + 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); + } + + 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); + } +} From 7ebc78ee2b8a7497395ab9308fae7887bdd93d20 Mon Sep 17 00:00:00 2001 From: Evgeny Khodzitsky <45710942+ekhodzitsky@users.noreply.github.com> Date: Thu, 10 Sep 2026 16:38:51 +0300 Subject: [PATCH 2/2] refactor: own message reservations and cover cancellation paths --- src/messages/socket.rs | 13 +- tests/bdd/extended/session_management.rs | 39 +++++- .../features/stale-server-detection.feature | 3 + tests/message_memory.rs | 130 +++++++++++++++++- 4 files changed, 173 insertions(+), 12 deletions(-) diff --git a/src/messages/socket.rs b/src/messages/socket.rs index d37153742..0a5c1ead9 100644 --- a/src/messages/socket.rs +++ b/src/messages/socket.rs @@ -129,6 +129,8 @@ 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(stream: &mut S, max_memory_usage: u64) -> Result where @@ -136,9 +138,9 @@ where { let (code, len) = read_message_header(stream).await?; let prev = CURRENT_MEMORY.fetch_add(len as i64, Ordering::Relaxed); - scopeguard::defer! { + 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); } @@ -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( stream: &mut S, @@ -177,9 +182,9 @@ where } let prev = CURRENT_MEMORY.fetch_add(len as i64, Ordering::Relaxed); - scopeguard::defer! { + 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); } diff --git a/tests/bdd/extended/session_management.rs b/tests/bdd/extended/session_management.rs index 40cd8fd08..f4e336089 100644 --- a/tests/bdd/extended/session_management.rs +++ b/tests/bdd/extended/session_management.rs @@ -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; @@ -555,14 +556,48 @@ pub async fn send_padded_simple_query( len: usize, session_name: String, ) { - assert!(len >= query.len() + 5); - let padded = format!("{query}{}", " ".repeat(len - query.len() - 5)); + 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$"# )] diff --git a/tests/bdd/features/stale-server-detection.feature b/tests/bdd/features/stale-server-detection.feature index 581764f0f..22e6307d6 100644 --- a/tests/bdd/features/stale-server-detection.feature +++ b/tests/bdd/features/stale-server-detection.feature @@ -69,6 +69,7 @@ Feature: Detect stale server connections during client idle in transaction 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] @@ -91,6 +92,8 @@ Feature: Detect stale server connections during client idle in transaction 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" diff --git a/tests/message_memory.rs b/tests/message_memory.rs index acbb75e66..12bc5440f 100644 --- a/tests/message_memory.rs +++ b/tests/message_memory.rs @@ -4,7 +4,17 @@ 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 { + 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( stream: &mut S, @@ -19,22 +29,27 @@ async fn read( // A separate test binary keeps CURRENT_MEMORY isolated from other socket tests. #[tokio::test] +#[serial] async fn cancelled_reads_release_memory() { - let mut message = vec![b'Q']; - message.extend_from_slice(&100_i32.to_be_bytes()); - message.extend_from_slice(&[b'x'; 96]); + let message = message(100); for mut buffer in [None, Some(BytesMut::new())] { - for prefix_len in [0, 1, 3, 5, 6, message.len() - 1] { + 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(); { - let mut pending = std::pin::pin!(read(&mut reader, buffer.as_mut(), 100)); + // 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 } + if prefix_len >= 5 { 100 } else { 0 }, + "prefix length: {prefix_len}" ); } @@ -46,6 +61,20 @@ async fn cancelled_reads_release_memory() { 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))); @@ -56,3 +85,92 @@ async fn cancelled_reads_release_memory() { 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); + } +}