From f840bb96a44681e52f30a2320a525890b797a028 Mon Sep 17 00:00:00 2001 From: Vincent Gao Date: Wed, 29 Jul 2026 12:32:52 +0200 Subject: [PATCH 1/3] Reject malformed element counts in message decoders Several decoders read a repeated-element count as a signed i16/i32 and pass it to Vec::with_capacity before reading any element. A negative count (0xffff read as i16 = -1) sign-extends to usize::MAX and panics with "capacity overflow", so a short crafted message crashes the peer decoding it (the shipped client/proxy on backend messages, the server on Bind). #192 fixed this for Parse, Bind and ParameterDescription by widening the count to u16 but left RowDescription, the three Copy*Response messages, NegotiateProtocolVersion, and Bind's result_column_format_codes. Read those counts as unsigned and route them through a shared codec::read_count that rejects a count exceeding the bytes remaining, turning a malformed message into a decode error instead of a panic. --- src/error.rs | 5 ++++ src/messages/codec.rs | 13 ++++++++++ src/messages/copy.rs | 18 +++++++------- src/messages/data.rs | 4 ++-- src/messages/extendedquery.rs | 5 ++-- src/messages/mod.rs | 45 +++++++++++++++++++++++++++++++++++ src/messages/startup.rs | 4 ++-- 7 files changed, 78 insertions(+), 16 deletions(-) diff --git a/src/error.rs b/src/error.rs index 1e2a649e..7026250f 100644 --- a/src/error.rs +++ b/src/error.rs @@ -17,6 +17,8 @@ pub enum PgWireError { InvalidMessageType(u8), #[error("Invalid message length, expected max {0}, actual: {1}")] MessageTooLarge(usize, usize), + #[error("Invalid element count {0}, exceeds {1} remaining bytes")] + InvalidElementCount(usize, usize), #[error("Invalid target type, received {0}")] InvalidTargetType(u8), #[error("Invalid transaction status, received {0}")] @@ -327,6 +329,9 @@ impl From for ErrorInfo { PgWireError::MessageTooLarge(..) => { ErrorInfo::new("FATAL".to_owned(), "08P01".to_owned(), error.to_string()) } + PgWireError::InvalidElementCount(..) => { + ErrorInfo::new("FATAL".to_owned(), "08P01".to_owned(), error.to_string()) + } PgWireError::InvalidTransactionStatus(_) => { ErrorInfo::new("FATAL".to_owned(), "08P01".to_owned(), error.to_string()) } diff --git a/src/messages/codec.rs b/src/messages/codec.rs index b58f42be..f5f05203 100644 --- a/src/messages/codec.rs +++ b/src/messages/codec.rs @@ -58,6 +58,19 @@ pub(crate) fn get_length(buf: &BytesMut, offset: usize) -> Option { } } +/// Validate an on-wire element count before it is used to pre-allocate a +/// collection. Counts are unsigned; a value larger than the bytes left in the +/// buffer cannot describe a real message (each element takes at least one byte), +/// so it is rejected instead of driving an oversized `Vec::with_capacity`. +pub(crate) fn read_count(count: usize, buf: &BytesMut) -> PgWireResult { + let remaining = buf.remaining(); + if count > remaining { + Err(PgWireError::InvalidElementCount(count, remaining)) + } else { + Ok(count) + } +} + /// Check if message_length matches and move the cursor to right position then /// call the `decode_fn` for the body pub(crate) fn decode_packet( diff --git a/src/messages/copy.rs b/src/messages/copy.rs index 8473af90..4cd8678a 100644 --- a/src/messages/copy.rs +++ b/src/messages/copy.rs @@ -135,13 +135,13 @@ impl Message for CopyInResponse { fn decode_body(buf: &mut BytesMut, _len: usize, _ctx: &DecodeContext) -> PgWireResult { let format = buf.get_i8(); - let columns = buf.get_i16(); - let mut column_formats = Vec::with_capacity(columns as usize); + let columns = codec::read_count(buf.get_u16() as usize, buf)?; + let mut column_formats = Vec::with_capacity(columns); for _ in 0..columns { column_formats.push(buf.get_i16()); } - Ok(Self::new(format, columns, column_formats)) + Ok(Self::new(format, columns as i16, column_formats)) } } @@ -183,13 +183,13 @@ impl Message for CopyOutResponse { fn decode_body(buf: &mut BytesMut, _len: usize, _ctx: &DecodeContext) -> PgWireResult { let format = buf.get_i8(); - let columns = buf.get_i16(); - let mut column_formats = Vec::with_capacity(columns as usize); + let columns = codec::read_count(buf.get_u16() as usize, buf)?; + let mut column_formats = Vec::with_capacity(columns); for _ in 0..columns { column_formats.push(buf.get_i16()); } - Ok(Self::new(format, columns, column_formats)) + Ok(Self::new(format, columns as i16, column_formats)) } } @@ -231,12 +231,12 @@ impl Message for CopyBothResponse { fn decode_body(buf: &mut BytesMut, _len: usize, _ctx: &DecodeContext) -> PgWireResult { let format = buf.get_i8(); - let columns = buf.get_i16(); - let mut column_formats = Vec::with_capacity(columns as usize); + let columns = codec::read_count(buf.get_u16() as usize, buf)?; + let mut column_formats = Vec::with_capacity(columns); for _ in 0..columns { column_formats.push(buf.get_i16()); } - Ok(Self::new(format, columns, column_formats)) + Ok(Self::new(format, columns as i16, column_formats)) } } diff --git a/src/messages/data.rs b/src/messages/data.rs index 241e0f97..6c219ce6 100644 --- a/src/messages/data.rs +++ b/src/messages/data.rs @@ -75,8 +75,8 @@ impl Message for RowDescription { } fn decode_body(buf: &mut BytesMut, _: usize, _ctx: &DecodeContext) -> PgWireResult { - let fields_len = buf.get_i16(); - let mut fields = Vec::with_capacity(fields_len as usize); + let fields_len = codec::read_count(buf.get_u16() as usize, buf)?; + let mut fields = Vec::with_capacity(fields_len); for _ in 0..fields_len { let field = FieldDescription { diff --git a/src/messages/extendedquery.rs b/src/messages/extendedquery.rs index ff8129d3..9181bfe5 100644 --- a/src/messages/extendedquery.rs +++ b/src/messages/extendedquery.rs @@ -280,9 +280,8 @@ impl Message for Bind { } } - let result_column_format_code_len = buf.get_i16(); - let mut result_column_format_codes = - Vec::with_capacity(result_column_format_code_len as usize); + let result_column_format_code_len = codec::read_count(buf.get_u16() as usize, buf)?; + let mut result_column_format_codes = Vec::with_capacity(result_column_format_code_len); for _ in 0..result_column_format_code_len { result_column_format_codes.push(buf.get_i16()); } diff --git a/src/messages/mod.rs b/src/messages/mod.rs index a334aba2..fb17d24a 100644 --- a/src/messages/mod.rs +++ b/src/messages/mod.rs @@ -1029,4 +1029,49 @@ mod test { assert_eq!(196608i32, i32::from(ProtocolVersion::PROTOCOL3_0)); assert_eq!(196610i32, i32::from(ProtocolVersion::PROTOCOL3_2)); } + + // A count read as signed (0xffff -> -1 -> usize::MAX) must become a decode + // error, not a `capacity overflow` panic. Completes #192 for the decoders it + // did not reach. + #[test] + fn test_element_count_capacity_overflow_rejected() { + let ctx = DecodeContext::new(ProtocolVersion::PROTOCOL3_0); + + // Backend messages decoded by the client/proxy from an untrusted server. + for bytes in [ + &b"T\0\0\0\x06\xff\xff"[..], // RowDescription + &b"G\0\0\0\x07\0\xff\xff"[..], // CopyInResponse + &b"H\0\0\0\x07\0\xff\xff"[..], // CopyOutResponse + &b"W\0\0\0\x07\0\xff\xff"[..], // CopyBothResponse + &b"v\0\0\0\x0c\0\0\0\0\xff\xff\xff\xff"[..], // NegotiateProtocolVersion + ] { + let mut buf = BytesMut::from(bytes); + assert!(super::PgWireBackendMessage::decode(&mut buf, &ctx).is_err()); + } + + // Bind decoded by the server from an untrusted client. #192 widened the + // first two of Bind's three counts; this covers the third. + let mut fctx = DecodeContext::new(ProtocolVersion::PROTOCOL3_0); + fctx.awaiting_frontend_ssl = false; + fctx.awaiting_frontend_startup = false; + let mut buf = BytesMut::from(&b"B\0\0\0\x0c\0\0\0\0\0\0\xff\xff"[..]); + assert!(super::PgWireFrontendMessage::decode(&mut buf, &fctx).is_err()); + } + + #[test] + fn test_valid_element_count_still_decodes() { + let ctx = DecodeContext::new(ProtocolVersion::PROTOCOL3_0); + + // Zero count decodes to an empty collection. + let mut buf = BytesMut::from(&b"T\0\0\0\x06\0\0"[..]); + let row = RowDescription::decode(&mut buf, &ctx).unwrap().unwrap(); + assert!(row.fields.is_empty()); + + // A well-formed positive count round-trips. + let row = RowDescription::new(vec![ + FieldDescription::new("id".to_owned(), 0, 0, 23, 4, -1, 0), + FieldDescription::new("name".to_owned(), 0, 0, 25, -1, -1, 0), + ]); + roundtrip!(row, RowDescription, &ctx); + } } diff --git a/src/messages/startup.rs b/src/messages/startup.rs index 97844120..34d41fd1 100644 --- a/src/messages/startup.rs +++ b/src/messages/startup.rs @@ -820,8 +820,8 @@ impl Message for NegotiateProtocolVersion { _ctx: &DecodeContext, ) -> PgWireResult { let version = buf.get_i32(); - let option_count = buf.get_i32(); - let mut options = Vec::with_capacity(option_count as usize); + let option_count = codec::read_count(buf.get_u32() as usize, buf)?; + let mut options = Vec::with_capacity(option_count); for _ in 0..option_count { options.push(codec::get_cstring(buf).unwrap_or_else(|| "".to_owned())) From 074ca5744b62d9cfa1f018e5bfe64c80acb958c2 Mon Sep 17 00:00:00 2001 From: Vincent Gao Date: Fri, 31 Jul 2026 20:19:34 +0200 Subject: [PATCH 2/3] Ensure element counts without converting them twice --- src/messages/codec.rs | 4 ++-- src/messages/copy.rs | 15 +++++++++------ src/messages/data.rs | 5 +++-- src/messages/extendedquery.rs | 6 ++++-- src/messages/startup.rs | 5 +++-- 5 files changed, 21 insertions(+), 14 deletions(-) diff --git a/src/messages/codec.rs b/src/messages/codec.rs index f5f05203..64c56e08 100644 --- a/src/messages/codec.rs +++ b/src/messages/codec.rs @@ -62,12 +62,12 @@ pub(crate) fn get_length(buf: &BytesMut, offset: usize) -> Option { /// collection. Counts are unsigned; a value larger than the bytes left in the /// buffer cannot describe a real message (each element takes at least one byte), /// so it is rejected instead of driving an oversized `Vec::with_capacity`. -pub(crate) fn read_count(count: usize, buf: &BytesMut) -> PgWireResult { +pub(crate) fn ensure_count(count: usize, buf: &BytesMut) -> PgWireResult<()> { let remaining = buf.remaining(); if count > remaining { Err(PgWireError::InvalidElementCount(count, remaining)) } else { - Ok(count) + Ok(()) } } diff --git a/src/messages/copy.rs b/src/messages/copy.rs index 4cd8678a..ce1ca9be 100644 --- a/src/messages/copy.rs +++ b/src/messages/copy.rs @@ -135,8 +135,9 @@ impl Message for CopyInResponse { fn decode_body(buf: &mut BytesMut, _len: usize, _ctx: &DecodeContext) -> PgWireResult { let format = buf.get_i8(); - let columns = codec::read_count(buf.get_u16() as usize, buf)?; - let mut column_formats = Vec::with_capacity(columns); + let columns = buf.get_u16(); + codec::ensure_count(columns as usize, buf)?; + let mut column_formats = Vec::with_capacity(columns as usize); for _ in 0..columns { column_formats.push(buf.get_i16()); } @@ -183,8 +184,9 @@ impl Message for CopyOutResponse { fn decode_body(buf: &mut BytesMut, _len: usize, _ctx: &DecodeContext) -> PgWireResult { let format = buf.get_i8(); - let columns = codec::read_count(buf.get_u16() as usize, buf)?; - let mut column_formats = Vec::with_capacity(columns); + let columns = buf.get_u16(); + codec::ensure_count(columns as usize, buf)?; + let mut column_formats = Vec::with_capacity(columns as usize); for _ in 0..columns { column_formats.push(buf.get_i16()); } @@ -231,8 +233,9 @@ impl Message for CopyBothResponse { fn decode_body(buf: &mut BytesMut, _len: usize, _ctx: &DecodeContext) -> PgWireResult { let format = buf.get_i8(); - let columns = codec::read_count(buf.get_u16() as usize, buf)?; - let mut column_formats = Vec::with_capacity(columns); + let columns = buf.get_u16(); + codec::ensure_count(columns as usize, buf)?; + let mut column_formats = Vec::with_capacity(columns as usize); for _ in 0..columns { column_formats.push(buf.get_i16()); } diff --git a/src/messages/data.rs b/src/messages/data.rs index 6c219ce6..fa0e33b3 100644 --- a/src/messages/data.rs +++ b/src/messages/data.rs @@ -75,8 +75,9 @@ impl Message for RowDescription { } fn decode_body(buf: &mut BytesMut, _: usize, _ctx: &DecodeContext) -> PgWireResult { - let fields_len = codec::read_count(buf.get_u16() as usize, buf)?; - let mut fields = Vec::with_capacity(fields_len); + let fields_len = buf.get_u16(); + codec::ensure_count(fields_len as usize, buf)?; + let mut fields = Vec::with_capacity(fields_len as usize); for _ in 0..fields_len { let field = FieldDescription { diff --git a/src/messages/extendedquery.rs b/src/messages/extendedquery.rs index 9181bfe5..f53b9298 100644 --- a/src/messages/extendedquery.rs +++ b/src/messages/extendedquery.rs @@ -280,8 +280,10 @@ impl Message for Bind { } } - let result_column_format_code_len = codec::read_count(buf.get_u16() as usize, buf)?; - let mut result_column_format_codes = Vec::with_capacity(result_column_format_code_len); + let result_column_format_code_len = buf.get_u16(); + codec::ensure_count(result_column_format_code_len as usize, buf)?; + let mut result_column_format_codes = + Vec::with_capacity(result_column_format_code_len as usize); for _ in 0..result_column_format_code_len { result_column_format_codes.push(buf.get_i16()); } diff --git a/src/messages/startup.rs b/src/messages/startup.rs index 34d41fd1..153c25e7 100644 --- a/src/messages/startup.rs +++ b/src/messages/startup.rs @@ -820,8 +820,9 @@ impl Message for NegotiateProtocolVersion { _ctx: &DecodeContext, ) -> PgWireResult { let version = buf.get_i32(); - let option_count = codec::read_count(buf.get_u32() as usize, buf)?; - let mut options = Vec::with_capacity(option_count); + let option_count = buf.get_u32(); + codec::ensure_count(option_count as usize, buf)?; + let mut options = Vec::with_capacity(option_count as usize); for _ in 0..option_count { options.push(codec::get_cstring(buf).unwrap_or_else(|| "".to_owned())) From e68abf0b38178cd505b2f0ce5fcfc54cb795d9d0 Mon Sep 17 00:00:00 2001 From: Vincent Gao Date: Sat, 1 Aug 2026 01:29:42 +0200 Subject: [PATCH 3/3] Make element count validation element-size aware Counts on the wire are element counts, not byte counts: Bind's result format codes are Int16s, so a count that fits the remaining bytes can still overrun the buffer when read as 2-byte elements. ensure_count now takes the per-element size and rejects count * size > remaining. --- src/error.rs | 4 ++-- src/messages/codec.rs | 12 ++++++------ src/messages/copy.rs | 10 +++++----- src/messages/data.rs | 3 ++- src/messages/extendedquery.rs | 2 +- src/messages/mod.rs | 6 ++++++ src/messages/startup.rs | 2 +- 7 files changed, 23 insertions(+), 16 deletions(-) diff --git a/src/error.rs b/src/error.rs index 7026250f..77b7c8c7 100644 --- a/src/error.rs +++ b/src/error.rs @@ -17,8 +17,8 @@ pub enum PgWireError { InvalidMessageType(u8), #[error("Invalid message length, expected max {0}, actual: {1}")] MessageTooLarge(usize, usize), - #[error("Invalid element count {0}, exceeds {1} remaining bytes")] - InvalidElementCount(usize, usize), + #[error("Invalid element count {0}: {1} bytes required, {2} remaining")] + InvalidElementCount(usize, usize, usize), #[error("Invalid target type, received {0}")] InvalidTargetType(u8), #[error("Invalid transaction status, received {0}")] diff --git a/src/messages/codec.rs b/src/messages/codec.rs index 64c56e08..72558ea6 100644 --- a/src/messages/codec.rs +++ b/src/messages/codec.rs @@ -59,13 +59,13 @@ pub(crate) fn get_length(buf: &BytesMut, offset: usize) -> Option { } /// Validate an on-wire element count before it is used to pre-allocate a -/// collection. Counts are unsigned; a value larger than the bytes left in the -/// buffer cannot describe a real message (each element takes at least one byte), -/// so it is rejected instead of driving an oversized `Vec::with_capacity`. -pub(crate) fn ensure_count(count: usize, buf: &BytesMut) -> PgWireResult<()> { +/// collection. Counts are unsigned; `count` elements of `elem_size` bytes each +/// must fit in the remaining buffer, or the message cannot be real. +pub(crate) fn ensure_count(count: usize, elem_size: usize, buf: &BytesMut) -> PgWireResult<()> { let remaining = buf.remaining(); - if count > remaining { - Err(PgWireError::InvalidElementCount(count, remaining)) + let needed = count.saturating_mul(elem_size); + if needed > remaining { + Err(PgWireError::InvalidElementCount(count, needed, remaining)) } else { Ok(()) } diff --git a/src/messages/copy.rs b/src/messages/copy.rs index ce1ca9be..d31921d2 100644 --- a/src/messages/copy.rs +++ b/src/messages/copy.rs @@ -135,14 +135,14 @@ impl Message for CopyInResponse { fn decode_body(buf: &mut BytesMut, _len: usize, _ctx: &DecodeContext) -> PgWireResult { let format = buf.get_i8(); - let columns = buf.get_u16(); - codec::ensure_count(columns as usize, buf)?; + let columns = buf.get_i16(); + codec::ensure_count(columns as usize, 2, buf)?; let mut column_formats = Vec::with_capacity(columns as usize); for _ in 0..columns { column_formats.push(buf.get_i16()); } - Ok(Self::new(format, columns as i16, column_formats)) + Ok(Self::new(format, columns, column_formats)) } } @@ -185,7 +185,7 @@ impl Message for CopyOutResponse { fn decode_body(buf: &mut BytesMut, _len: usize, _ctx: &DecodeContext) -> PgWireResult { let format = buf.get_i8(); let columns = buf.get_u16(); - codec::ensure_count(columns as usize, buf)?; + codec::ensure_count(columns as usize, 2, buf)?; let mut column_formats = Vec::with_capacity(columns as usize); for _ in 0..columns { column_formats.push(buf.get_i16()); @@ -234,7 +234,7 @@ impl Message for CopyBothResponse { fn decode_body(buf: &mut BytesMut, _len: usize, _ctx: &DecodeContext) -> PgWireResult { let format = buf.get_i8(); let columns = buf.get_u16(); - codec::ensure_count(columns as usize, buf)?; + codec::ensure_count(columns as usize, 2, buf)?; let mut column_formats = Vec::with_capacity(columns as usize); for _ in 0..columns { column_formats.push(buf.get_i16()); diff --git a/src/messages/data.rs b/src/messages/data.rs index fa0e33b3..58328972 100644 --- a/src/messages/data.rs +++ b/src/messages/data.rs @@ -76,7 +76,8 @@ impl Message for RowDescription { fn decode_body(buf: &mut BytesMut, _: usize, _ctx: &DecodeContext) -> PgWireResult { let fields_len = buf.get_u16(); - codec::ensure_count(fields_len as usize, buf)?; + // Each field: C-string name (>= 1 byte) + 18 fixed bytes. + codec::ensure_count(fields_len as usize, 19, buf)?; let mut fields = Vec::with_capacity(fields_len as usize); for _ in 0..fields_len { diff --git a/src/messages/extendedquery.rs b/src/messages/extendedquery.rs index f53b9298..be0fe8ed 100644 --- a/src/messages/extendedquery.rs +++ b/src/messages/extendedquery.rs @@ -281,7 +281,7 @@ impl Message for Bind { } let result_column_format_code_len = buf.get_u16(); - codec::ensure_count(result_column_format_code_len as usize, buf)?; + codec::ensure_count(result_column_format_code_len as usize, 2, buf)?; let mut result_column_format_codes = Vec::with_capacity(result_column_format_code_len as usize); for _ in 0..result_column_format_code_len { diff --git a/src/messages/mod.rs b/src/messages/mod.rs index fb17d24a..c23a0f4b 100644 --- a/src/messages/mod.rs +++ b/src/messages/mod.rs @@ -1056,6 +1056,12 @@ mod test { fctx.awaiting_frontend_startup = false; let mut buf = BytesMut::from(&b"B\0\0\0\x0c\0\0\0\0\0\0\xff\xff"[..]); assert!(super::PgWireFrontendMessage::decode(&mut buf, &fctx).is_err()); + + // Bind with 5 result format codes but only 6 bytes left: the codes are + // Int16s, so the element-aware bound must reject (5 * 2 > 6) instead of + // passing a byte-count check and panicking on the truncated read. + let mut buf = BytesMut::from(&b"B\0\0\0\x0e\0\0\0\0\0\0\0\x05\0\0\0\0\0\0"[..]); + assert!(super::PgWireFrontendMessage::decode(&mut buf, &fctx).is_err()); } #[test] diff --git a/src/messages/startup.rs b/src/messages/startup.rs index 153c25e7..1330134b 100644 --- a/src/messages/startup.rs +++ b/src/messages/startup.rs @@ -821,7 +821,7 @@ impl Message for NegotiateProtocolVersion { ) -> PgWireResult { let version = buf.get_i32(); let option_count = buf.get_u32(); - codec::ensure_count(option_count as usize, buf)?; + codec::ensure_count(option_count as usize, 1, buf)?; let mut options = Vec::with_capacity(option_count as usize); for _ in 0..option_count {