diff --git a/mssql-tds/src/connection/transport/buffers.rs b/mssql-tds/src/connection/transport/buffers.rs index 6ed265f9..99ff5eff 100644 --- a/mssql-tds/src/connection/transport/buffers.rs +++ b/mssql-tds/src/connection/transport/buffers.rs @@ -49,6 +49,72 @@ impl TdsReadBuffer { self.buffer_length - self.buffer_position } + #[inline(always)] + fn try_read_array(&mut self) -> Option<[u8; N]> { + if !self.do_we_have_enough_data(N) { + return None; + } + + let position = self.buffer_position; + let bytes = self.working_buffer[position..position + N] + .try_into() + .expect("slice length is fixed by N"); + self.consume_bytes(N); + Some(bytes) + } + + #[inline(always)] + pub(crate) fn try_read_byte(&mut self) -> Option { + self.try_read_array().map(|[value]| value) + } + + #[inline(always)] + pub(crate) fn try_read_int16(&mut self) -> Option { + self.try_read_array().map(i16::from_le_bytes) + } + + #[inline(always)] + pub(crate) fn try_read_uint16(&mut self) -> Option { + self.try_read_array().map(u16::from_le_bytes) + } + + #[inline(always)] + pub(crate) fn try_read_uint24(&mut self) -> Option { + let [b0, b1, b2] = self.try_read_array()?; + Some(u32::from_le_bytes([b0, b1, b2, 0])) + } + + #[inline(always)] + pub(crate) fn try_read_int32(&mut self) -> Option { + self.try_read_array().map(i32::from_le_bytes) + } + + #[inline(always)] + pub(crate) fn try_read_uint32(&mut self) -> Option { + self.try_read_array().map(u32::from_le_bytes) + } + + #[inline(always)] + pub(crate) fn try_read_uint40(&mut self) -> Option { + let [b0, b1, b2, b3, b4] = self.try_read_array()?; + Some(u64::from_le_bytes([b0, b1, b2, b3, b4, 0, 0, 0])) + } + + #[inline(always)] + pub(crate) fn try_read_int64(&mut self) -> Option { + self.try_read_array().map(i64::from_le_bytes) + } + + #[inline(always)] + pub(crate) fn try_read_float32(&mut self) -> Option { + self.try_read_array().map(f32::from_le_bytes) + } + + #[inline(always)] + pub(crate) fn try_read_float64(&mut self) -> Option { + self.try_read_array().map(f64::from_le_bytes) + } + pub(crate) fn consume_bytes(&mut self, byte_count: usize) { if byte_count > (self.buffer_length - self.buffer_position) { panic!("Not enough data to consume"); @@ -426,6 +492,74 @@ mod tests { assert!(!buf.do_we_have_enough_data(401)); } + #[test] + fn test_fixed_scalar_probes_read_complete_values() { + let expected_byte = 0xAB; + let expected_int16 = -0x1234i16; + let expected_uint16 = 0x1234u16; + let expected_uint24 = 0x00A1_B2C3u32; + let expected_int32 = -0x0123_4567i32; + let expected_uint32 = 0x89AB_CDEFu32; + let expected_uint40 = 0xAB_CDEF_0123u64; + let expected_int64 = -0x0102_0304_0506_0708i64; + let expected_float32 = 1.5f32; + let expected_float64 = -2.25f64; + + let mut bytes = Vec::new(); + bytes.push(expected_byte); + bytes.extend_from_slice(&expected_int16.to_le_bytes()); + bytes.extend_from_slice(&expected_uint16.to_le_bytes()); + bytes.extend_from_slice(&expected_uint24.to_le_bytes()[..3]); + bytes.extend_from_slice(&expected_int32.to_le_bytes()); + bytes.extend_from_slice(&expected_uint32.to_le_bytes()); + bytes.extend_from_slice(&expected_uint40.to_le_bytes()[..5]); + bytes.extend_from_slice(&expected_int64.to_le_bytes()); + bytes.extend_from_slice(&expected_float32.to_le_bytes()); + bytes.extend_from_slice(&expected_float64.to_le_bytes()); + + let mut buf = TdsReadBuffer::new(4096); + buf.working_buffer[..bytes.len()].copy_from_slice(&bytes); + buf.reset_to_length(bytes.len()); + + assert_eq!(buf.try_read_byte(), Some(expected_byte)); + assert_eq!(buf.try_read_int16(), Some(expected_int16)); + assert_eq!(buf.try_read_uint16(), Some(expected_uint16)); + assert_eq!(buf.try_read_uint24(), Some(expected_uint24)); + assert_eq!(buf.try_read_int32(), Some(expected_int32)); + assert_eq!(buf.try_read_uint32(), Some(expected_uint32)); + assert_eq!(buf.try_read_uint40(), Some(expected_uint40)); + assert_eq!(buf.try_read_int64(), Some(expected_int64)); + assert_eq!(buf.try_read_float32(), Some(expected_float32)); + assert_eq!(buf.try_read_float64(), Some(expected_float64)); + assert_eq!(buf.get_remaining_byte_count(), 0); + } + + #[test] + fn test_fixed_scalar_probe_misses_do_not_consume() { + let mut buf = TdsReadBuffer::new(4096); + + macro_rules! assert_miss_does_not_consume { + ($partial_len:expr, $method:ident) => {{ + buf.working_buffer[..$partial_len].fill(0xA5); + buf.reset_to_length($partial_len); + assert_eq!(buf.$method(), None); + assert_eq!(buf.buffer_position, 0); + assert_eq!(buf.get_remaining_byte_count(), $partial_len); + }}; + } + + assert_miss_does_not_consume!(0, try_read_byte); + assert_miss_does_not_consume!(1, try_read_int16); + assert_miss_does_not_consume!(1, try_read_uint16); + assert_miss_does_not_consume!(2, try_read_uint24); + assert_miss_does_not_consume!(3, try_read_int32); + assert_miss_does_not_consume!(3, try_read_uint32); + assert_miss_does_not_consume!(4, try_read_uint40); + assert_miss_does_not_consume!(7, try_read_int64); + assert_miss_does_not_consume!(3, try_read_float32); + assert_miss_does_not_consume!(7, try_read_float64); + } + #[test] fn test_get_remaining_byte_count() { let mut buf = TdsReadBuffer::new(4096); diff --git a/mssql-tds/src/connection/transport/network_transport.rs b/mssql-tds/src/connection/transport/network_transport.rs index 536f50bb..d336152b 100644 --- a/mssql-tds/src/connection/transport/network_transport.rs +++ b/mssql-tds/src/connection/transport/network_transport.rs @@ -1115,13 +1115,63 @@ impl TdsPacketReader for NetworkTransport { self.tds_read_buffer.reset_to_length(0); } + #[inline(always)] + fn try_read_byte(&mut self) -> Option { + self.tds_read_buffer.try_read_byte() + } + + #[inline(always)] + fn try_read_int16(&mut self) -> Option { + self.tds_read_buffer.try_read_int16() + } + + #[inline(always)] + fn try_read_uint16(&mut self) -> Option { + self.tds_read_buffer.try_read_uint16() + } + + #[inline(always)] + fn try_read_uint24(&mut self) -> Option { + self.tds_read_buffer.try_read_uint24() + } + + #[inline(always)] + fn try_read_int32(&mut self) -> Option { + self.tds_read_buffer.try_read_int32() + } + + #[inline(always)] + fn try_read_uint32(&mut self) -> Option { + self.tds_read_buffer.try_read_uint32() + } + + #[inline(always)] + fn try_read_uint40(&mut self) -> Option { + self.tds_read_buffer.try_read_uint40() + } + + #[inline(always)] + fn try_read_int64(&mut self) -> Option { + self.tds_read_buffer.try_read_int64() + } + + #[inline(always)] + fn try_read_float32(&mut self) -> Option { + self.tds_read_buffer.try_read_float32() + } + + #[inline(always)] + fn try_read_float64(&mut self) -> Option { + self.tds_read_buffer.try_read_float64() + } + async fn read_byte(&mut self) -> TdsResult { - while !self.tds_read_buffer.do_we_have_enough_data(1) { + loop { + if let Some(value) = self.try_read_byte() { + return Ok(value); + } self.read_tds_packet().await?; } - let result: u8 = self.tds_read_buffer.working_buffer[self.tds_read_buffer.buffer_position]; - self.tds_read_buffer.consume_bytes(1); - Ok(result) } async fn read_int16_big_endian(&mut self) -> TdsResult { @@ -1142,80 +1192,79 @@ impl TdsPacketReader for NetworkTransport { } async fn read_uint40(&mut self) -> TdsResult { - while !self.tds_read_buffer.do_we_have_enough_data(5) { + loop { + if let Some(value) = self.try_read_uint40() { + return Ok(value); + } self.read_tds_packet().await?; } - - let result = LittleEndian::read_uint(self.tds_read_buffer.get_slice(), 5); - self.tds_read_buffer.consume_bytes(5); - Ok(result) } async fn read_float32(&mut self) -> TdsResult { - while !self.tds_read_buffer.do_we_have_enough_data(4) { + loop { + if let Some(value) = self.try_read_float32() { + return Ok(value); + } self.read_tds_packet().await?; } - let result = LittleEndian::read_f32(self.tds_read_buffer.get_slice()); - self.tds_read_buffer.consume_bytes(4); - Ok(result) } async fn read_float64(&mut self) -> TdsResult { - while !self.tds_read_buffer.do_we_have_enough_data(8) { + loop { + if let Some(value) = self.try_read_float64() { + return Ok(value); + } self.read_tds_packet().await?; } - let result = LittleEndian::read_f64(self.tds_read_buffer.get_slice()); - self.tds_read_buffer.consume_bytes(8); - Ok(result) } async fn read_int16(&mut self) -> TdsResult { - while !self.tds_read_buffer.do_we_have_enough_data(2) { + loop { + if let Some(value) = self.try_read_int16() { + return Ok(value); + } self.read_tds_packet().await?; } - let result = LittleEndian::read_i16(self.tds_read_buffer.get_slice()); - self.tds_read_buffer.consume_bytes(2); - Ok(result) } async fn read_uint16(&mut self) -> TdsResult { - while !self.tds_read_buffer.do_we_have_enough_data(2) { + loop { + if let Some(value) = self.try_read_uint16() { + return Ok(value); + } self.read_tds_packet().await?; } - let result = LittleEndian::read_u16(self.tds_read_buffer.get_slice()); - self.tds_read_buffer.consume_bytes(2); - Ok(result) } async fn read_uint24(&mut self) -> TdsResult { - while !self.tds_read_buffer.do_we_have_enough_data(3) { + loop { + if let Some(value) = self.try_read_uint24() { + return Ok(value); + } self.read_tds_packet().await?; } - let result = LittleEndian::read_u24(self.tds_read_buffer.get_slice()); - self.tds_read_buffer.consume_bytes(3); - Ok(result) } async fn read_int32(&mut self) -> TdsResult { - while !self.tds_read_buffer.do_we_have_enough_data(4) { + loop { + if let Some(value) = self.try_read_int32() { + return Ok(value); + } self.read_tds_packet().await?; } - let result = LittleEndian::read_i32(self.tds_read_buffer.get_slice()); - self.tds_read_buffer.consume_bytes(4); - Ok(result) } async fn read_uint32(&mut self) -> TdsResult { - while !self.tds_read_buffer.do_we_have_enough_data(4) { + loop { + if let Some(value) = self.try_read_uint32() { + return Ok(value); + } self.read_tds_packet().await?; } - let result = LittleEndian::read_u32(self.tds_read_buffer.get_slice()); - self.tds_read_buffer.consume_bytes(4); - Ok(result) } async fn read_int64(&mut self) -> TdsResult { - while !self.tds_read_buffer.do_we_have_enough_data(8) { + loop { + if let Some(value) = self.try_read_int64() { + return Ok(value); + } self.read_tds_packet().await?; } - let result = LittleEndian::read_i64(self.tds_read_buffer.get_slice()); - self.tds_read_buffer.consume_bytes(8); - Ok(result) } async fn read_uint64(&mut self) -> TdsResult { while !self.tds_read_buffer.do_we_have_enough_data(8) { @@ -2806,6 +2855,94 @@ pub(crate) mod tests { assert_eq!(reader.read_uint32().await.unwrap(), 0x4433_2211); } + #[tokio::test] + async fn test_sync_scalar_probe_fallback_across_packet_boundaries() { + let expected_uint16 = 0x1234u16; + let expected_int16 = -0x1234i16; + let expected_uint24 = 0x00A1_B2C3u32; + let expected_int32 = -0x0123_4567i32; + let expected_uint32 = 0x89AB_CDEFu32; + let expected_uint40 = 0xAB_CDEF_0123u64; + let expected_int64 = -0x0102_0304_0506_0708i64; + let expected_float32 = 1.5f32; + let expected_float64 = -2.25f64; + + let uint16 = expected_uint16.to_le_bytes(); + let int16 = expected_int16.to_le_bytes(); + let uint24 = expected_uint24.to_le_bytes(); + let int32 = expected_int32.to_le_bytes(); + let uint32 = expected_uint32.to_le_bytes(); + let uint40 = expected_uint40.to_le_bytes(); + let int64 = expected_int64.to_le_bytes(); + let float32 = expected_float32.to_le_bytes(); + let float64 = expected_float64.to_le_bytes(); + + let payloads = [ + vec![0xAB, uint16[0]], + vec![uint16[1], int16[0]], + vec![int16[1], uint24[0], uint24[1]], + vec![uint24[2], int32[0], int32[1], int32[2]], + vec![int32[3], uint32[0], uint32[1], uint32[2]], + vec![uint32[3], uint40[0], uint40[1], uint40[2], uint40[3]], + vec![ + uint40[4], int64[0], int64[1], int64[2], int64[3], int64[4], int64[5], int64[6], + ], + vec![int64[7], float32[0], float32[1], float32[2]], + vec![ + float32[3], float64[0], float64[1], float64[2], float64[3], float64[4], float64[5], + float64[6], + ], + vec![float64[7]], + ]; + + let mut stream = Vec::new(); + for payload in payloads { + let mut packet = TestPacketBuilder::new(PacketType::TabularResult); + stream.extend_from_slice(&packet.append_bytes(&payload).build()); + } + + let mut reader = create_network_transport_with_data(&stream); + + assert_eq!(reader.try_read_byte(), None); + assert_eq!(reader.read_byte().await.unwrap(), 0xAB); + + assert_eq!(reader.try_read_uint16(), None); + assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 1); + assert_eq!(reader.read_uint16().await.unwrap(), expected_uint16); + + assert_eq!(reader.try_read_int16(), None); + assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 1); + assert_eq!(reader.read_int16().await.unwrap(), expected_int16); + + assert_eq!(reader.try_read_uint24(), None); + assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 2); + assert_eq!(reader.read_uint24().await.unwrap(), expected_uint24); + + assert_eq!(reader.try_read_int32(), None); + assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 3); + assert_eq!(reader.read_int32().await.unwrap(), expected_int32); + + assert_eq!(reader.try_read_uint32(), None); + assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 3); + assert_eq!(reader.read_uint32().await.unwrap(), expected_uint32); + + assert_eq!(reader.try_read_uint40(), None); + assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 4); + assert_eq!(reader.read_uint40().await.unwrap(), expected_uint40); + + assert_eq!(reader.try_read_int64(), None); + assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 7); + assert_eq!(reader.read_int64().await.unwrap(), expected_int64); + + assert_eq!(reader.try_read_float32(), None); + assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 3); + assert_eq!(reader.read_float32().await.unwrap(), expected_float32); + + assert_eq!(reader.try_read_float64(), None); + assert_eq!(reader.tds_read_buffer.get_remaining_byte_count(), 7); + assert_eq!(reader.read_float64().await.unwrap(), expected_float64); + } + /// A payload-free non-EOM packet is malformed: it neither carries payload /// nor terminates a message. #[tokio::test] diff --git a/mssql-tds/src/datatypes/decoder.rs b/mssql-tds/src/datatypes/decoder.rs index 94fd6d2e..2072e493 100644 --- a/mssql-tds/src/datatypes/decoder.rs +++ b/mssql-tds/src/datatypes/decoder.rs @@ -30,6 +30,17 @@ use crate::{query::metadata::ColumnMetadata, token::tokens::SqlCollation}; use super::row_writer::{RowWriter, write_column_value}; +// Avoid constructing a read future when the complete scalar is already buffered. +// Probe misses consume nothing; the async method remains the authoritative refill path. +macro_rules! read_sync_first { + ($reader:expr, $try_method:ident, $read_method:ident) => { + match ($reader).$try_method() { + Some(value) => value, + None => ($reader).$read_method().await?, + } + }; +} + /// Reads an encrypted column's cipher bytes from the wire and turns them back /// into a plaintext [`ColumnValues`]. /// @@ -209,7 +220,7 @@ impl PlpChunkStreamReader { where T: TdsPacketReader + Send + Sync, { - let raw_len_i64 = reader.read_int64().await?; + let raw_len_i64 = read_sync_first!(reader, try_read_int64, read_int64); let raw_len = raw_len_i64 as u64; let raw_len_usize = raw_len as usize; @@ -264,7 +275,7 @@ impl PlpChunkStreamReader { return Ok(true); } - let chunk_len = reader.read_uint32().await? as usize; + let chunk_len = read_sync_first!(reader, try_read_uint32, read_uint32) as usize; if chunk_len == 0 { self.reached_end = true; if let PlpChunkReadLength::Known(known_len) = self.length @@ -537,10 +548,10 @@ impl GenericDecoder { where T: TdsPacketReader + Send + Sync, { - let length = reader.read_uint32().await?; - let variant_base_type = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_uint32, read_uint32); + let variant_base_type = read_sync_first!(reader, try_read_byte, read_byte); let tds_type = TdsDataType::try_from(variant_base_type)?; - let variant_prop_bytes = reader.read_byte().await?; + let variant_prop_bytes = read_sync_first!(reader, try_read_byte, read_byte); let bytes_for_type_and_properties_byte = 2; // Use checked arithmetic to prevent integer underflow @@ -632,7 +643,7 @@ impl GenericDecoder { where T: TdsPacketReader + Send + Sync, { - let scale = reader.read_byte().await?; + let scale = read_sync_first!(reader, try_read_byte, read_byte); Ok(match tds_type { TdsDataType::TimeN => { let time_nanos = self.read_time(reader, data_length as u8, scale).await?; @@ -663,7 +674,7 @@ impl GenericDecoder { T: TdsPacketReader + Send + Sync, { // Decimal/numeric data type has 1 byte length. - let length = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_byte, read_byte); let TypeInfoVariant::VarLenPrecisionScale(_, _, precision, scale) = metadata.type_info.type_info_variant else { @@ -688,7 +699,7 @@ impl GenericDecoder { if length == 0 { return Ok(None); } - let sign = reader.read_byte().await?; + let sign = read_sync_first!(reader, try_read_byte, read_byte); let is_positive = sign == 1; // Round up: a declared length that does not cover whole 32-bit words @@ -729,8 +740,8 @@ impl GenericDecoder { where T: TdsPacketReader + Send + Sync, { - let days = reader.read_int32().await?; - let ticks = reader.read_uint32().await?; + let days = read_sync_first!(reader, try_read_int32, read_int32); + let ticks = read_sync_first!(reader, try_read_uint32, read_uint32); Ok(SqlDateTime { days, time: ticks }) } @@ -739,8 +750,8 @@ impl GenericDecoder { where T: TdsPacketReader + Send + Sync, { - let days = reader.read_uint16().await?; - let minutes = reader.read_uint16().await?; + let days = read_sync_first!(reader, try_read_uint16, read_uint16); + let minutes = read_sync_first!(reader, try_read_uint16, read_uint16); Ok(SqlSmallDateTime { days, time: minutes, @@ -751,7 +762,7 @@ impl GenericDecoder { where T: TdsPacketReader + Send + Sync, { - let days = reader.read_uint24().await?; + let days = read_sync_first!(reader, try_read_uint24, read_uint24); Ok(SqlDate::unchecked_create(days)) } @@ -760,9 +771,9 @@ impl GenericDecoder { T: TdsPacketReader + Send + Sync, { let scaled_value = match byte_len { - 3 => reader.read_uint24().await? as u64, - 4 => reader.read_uint32().await? as u64, - _ => reader.read_uint40().await?, + 3 => read_sync_first!(reader, try_read_uint24, read_uint24) as u64, + 4 => read_sync_first!(reader, try_read_uint32, read_uint32) as u64, + _ => read_sync_first!(reader, try_read_uint40, read_uint40), }; // The value from SQL Server is in scaled units based on the scale: @@ -841,7 +852,7 @@ impl GenericDecoder { ))); } }; - let offset = reader.read_int16().await?; + let offset = read_sync_first!(reader, try_read_int16, read_int16); let datetime_offset = SqlDateTimeOffset { datetime2, offset }; Ok(ColumnValues::DateTimeOffset(datetime_offset)) } @@ -851,10 +862,10 @@ impl GenericDecoder { T: TdsPacketReader + Send + Sync, { let value: ColumnValues = match byte_len { - 1 => ColumnValues::TinyInt(reader.read_byte().await?), // Some(reader.read_byte().await? as i64), - 2 => ColumnValues::SmallInt(reader.read_int16().await?), // Some(reader.read_int16().await? as i64), - 4 => ColumnValues::Int(reader.read_int32().await?), - 8 => ColumnValues::BigInt(reader.read_int64().await?), + 1 => ColumnValues::TinyInt(read_sync_first!(reader, try_read_byte, read_byte)), + 2 => ColumnValues::SmallInt(read_sync_first!(reader, try_read_int16, read_int16)), + 4 => ColumnValues::Int(read_sync_first!(reader, try_read_int32, read_int32)), + 8 => ColumnValues::BigInt(read_sync_first!(reader, try_read_int64, read_int64)), 0 => ColumnValues::Null, _ => { return Err(crate::error::Error::from(Error::new( @@ -870,7 +881,7 @@ impl GenericDecoder { where T: TdsPacketReader + Send + Sync, { - let small_money_val = reader.read_int32().await?; + let small_money_val = read_sync_first!(reader, try_read_int32, read_int32); Ok(small_money_val.into()) } @@ -880,8 +891,8 @@ impl GenericDecoder { where T: TdsPacketReader + Send + Sync, { - let msb = reader.read_int32().await?; - let lsb = reader.read_int32().await?; + let msb = read_sync_first!(reader, try_read_int32, read_int32); + let lsb = read_sync_first!(reader, try_read_int32, read_int32); Ok(SqlMoney { lsb_part: lsb, msb_part: msb, @@ -937,7 +948,7 @@ impl GenericDecoder { }; // Read length prefix (USHORTLEN format) - let length_prefix_value = reader.read_uint16().await? as usize; + let length_prefix_value = read_sync_first!(reader, try_read_uint16, read_uint16) as usize; // Handle NULL (length = 0xFFFF) if length_prefix_value == 0xFFFF { @@ -961,13 +972,13 @@ impl GenericDecoder { } // Read 8-byte header - let layout_format_byte = reader.read_byte().await?; - let layout_version_byte = reader.read_byte().await?; - let dimension_count = reader.read_uint16().await?; - let base_type_byte = reader.read_byte().await?; - let _reserved1 = reader.read_byte().await?; // Reserved - let _reserved2 = reader.read_byte().await?; // Reserved - let _reserved3 = reader.read_byte().await?; // Reserved + let layout_format_byte = read_sync_first!(reader, try_read_byte, read_byte); + let layout_version_byte = read_sync_first!(reader, try_read_byte, read_byte); + let dimension_count = read_sync_first!(reader, try_read_uint16, read_uint16); + let base_type_byte = read_sync_first!(reader, try_read_byte, read_byte); + let _reserved1 = read_sync_first!(reader, try_read_byte, read_byte); // Reserved + let _reserved2 = read_sync_first!(reader, try_read_byte, read_byte); // Reserved + let _reserved3 = read_sync_first!(reader, try_read_byte, read_byte); // Reserved // Validate header using enum conversions let _layout_format = VectorLayoutFormat::try_from(layout_format_byte)?; @@ -1031,7 +1042,7 @@ impl GenericDecoder { where T: TdsPacketReader + Send + Sync, { - let long_len_i64 = reader.read_int64().await?; + let long_len_i64 = read_sync_first!(reader, try_read_int64, read_int64); let long_len = long_len_i64 as u64; // If the length is SQL_PLP_NULL, it means the value is NULL. @@ -1055,7 +1066,7 @@ impl GenericDecoder { 0 }; let mut plp_buffer = vec![0u8; vector_capacity]; - let mut chunk_len = reader.read_uint32().await? as usize; + let mut chunk_len = read_sync_first!(reader, try_read_uint32, read_uint32) as usize; let mut offset: usize = 0; let mut chunk_count = 0u32; @@ -1116,7 +1127,7 @@ impl GenericDecoder { .read_bytes(&mut plp_buffer[offset..offset + chunk_len]) .await?; offset += chunk_size_read; - chunk_len = reader.read_uint32().await? as usize; + chunk_len = read_sync_first!(reader, try_read_uint32, read_uint32) as usize; } Ok(Some(plp_buffer)) } @@ -1140,24 +1151,30 @@ impl GenericDecoder { match metadata.data_type { // === Fixed-length integer types === TdsDataType::Int1 => { - writer.write_u8(col, reader.read_byte().await?); + writer.write_u8(col, read_sync_first!(reader, try_read_byte, read_byte)); } TdsDataType::Int2 => { - writer.write_i16(col, reader.read_int16().await?); + writer.write_i16(col, read_sync_first!(reader, try_read_int16, read_int16)); } TdsDataType::Int4 => { - writer.write_i32(col, reader.read_int32().await?); + writer.write_i32(col, read_sync_first!(reader, try_read_int32, read_int32)); } TdsDataType::Int8 => { - writer.write_i64(col, reader.read_int64().await?); + writer.write_i64(col, read_sync_first!(reader, try_read_int64, read_int64)); } TdsDataType::IntN => { - let byte_len = reader.read_byte().await?; + let byte_len = read_sync_first!(reader, try_read_byte, read_byte); match byte_len { - 1 => writer.write_u8(col, reader.read_byte().await?), - 2 => writer.write_i16(col, reader.read_int16().await?), - 4 => writer.write_i32(col, reader.read_int32().await?), - 8 => writer.write_i64(col, reader.read_int64().await?), + 1 => writer.write_u8(col, read_sync_first!(reader, try_read_byte, read_byte)), + 2 => { + writer.write_i16(col, read_sync_first!(reader, try_read_int16, read_int16)) + } + 4 => { + writer.write_i32(col, read_sync_first!(reader, try_read_int32, read_int32)) + } + 8 => { + writer.write_i64(col, read_sync_first!(reader, try_read_int64, read_int64)) + } 0 => writer.write_null(col), _ => { return Err(crate::error::Error::from(Error::new( @@ -1170,28 +1187,40 @@ impl GenericDecoder { // === Fixed-length float types === TdsDataType::Flt4 => { - writer.write_f32(col, reader.read_float32().await?); + writer.write_f32( + col, + read_sync_first!(reader, try_read_float32, read_float32), + ); } TdsDataType::Flt8 => { - writer.write_f64(col, reader.read_float64().await?); + writer.write_f64( + col, + read_sync_first!(reader, try_read_float64, read_float64), + ); } TdsDataType::FltN => { - let length = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_byte, read_byte); match length { 0 => writer.write_null(col), - 4 => writer.write_f32(col, reader.read_float32().await?), - _ => writer.write_f64(col, reader.read_float64().await?), + 4 => writer.write_f32( + col, + read_sync_first!(reader, try_read_float32, read_float32), + ), + _ => writer.write_f64( + col, + read_sync_first!(reader, try_read_float64, read_float64), + ), } } // === Bit types === TdsDataType::Bit => { - writer.write_bool(col, reader.read_byte().await? == 1); + writer.write_bool(col, read_sync_first!(reader, try_read_byte, read_byte) == 1); } TdsDataType::BitN => { - let byte_len = reader.read_byte().await?; + let byte_len = read_sync_first!(reader, try_read_byte, read_byte); if byte_len > 0 { - writer.write_bool(col, reader.read_byte().await? == 1); + writer.write_bool(col, read_sync_first!(reader, try_read_byte, read_byte) == 1); } else { writer.write_null(col); } @@ -1205,7 +1234,7 @@ impl GenericDecoder { writer.write_money(col, self.read_money8(reader).await?); } TdsDataType::MoneyN => { - let byte_len = reader.read_byte().await?; + let byte_len = read_sync_first!(reader, try_read_byte, read_byte); match byte_len { 4 => writer.write_smallmoney(col, self.read_money4(reader).await?), 8 => writer.write_money(col, self.read_money8(reader).await?), @@ -1244,7 +1273,7 @@ impl GenericDecoder { // === Binary types === TdsDataType::BigBinary => { - let length = reader.read_uint16().await?; + let length = read_sync_first!(reader, try_read_uint16, read_uint16); // 0xFFFF is the USHORTLEN NULL marker (CHARBIN_NULL). if length == 0xFFFF { writer.write_null(col); @@ -1266,7 +1295,7 @@ impl GenericDecoder { None => writer.write_null(col), } } else { - let length = reader.read_uint16().await?; + let length = read_sync_first!(reader, try_read_uint16, read_uint16); // 0xFFFF is the USHORTLEN NULL marker (CHARBIN_NULL). if length == 0xFFFF { writer.write_null(col); @@ -1288,8 +1317,8 @@ impl GenericDecoder { writer.write_datetime(col, self.read_datetime(reader).await?); } TdsDataType::DateTim4 => { - let daypart = reader.read_uint16().await?; - let timepart = reader.read_uint16().await?; + let daypart = read_sync_first!(reader, try_read_uint16, read_uint16); + let timepart = read_sync_first!(reader, try_read_uint16, read_uint16); writer.write_smalldatetime( col, SqlSmallDateTime { @@ -1299,7 +1328,7 @@ impl GenericDecoder { ); } TdsDataType::DateTimeN => { - let length = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_byte, read_byte); match length { 0 => writer.write_null(col), 4 => writer.write_smalldatetime(col, self.read_small_datetime(reader).await?), @@ -1307,7 +1336,7 @@ impl GenericDecoder { } } TdsDataType::DateN => { - let length = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_byte, read_byte); if length == 0 { writer.write_null(col); } else { @@ -1315,7 +1344,7 @@ impl GenericDecoder { } } TdsDataType::TimeN => { - let length = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_byte, read_byte); if length == 0 { writer.write_null(col); } else { @@ -1335,7 +1364,7 @@ impl GenericDecoder { } } TdsDataType::DateTime2N => { - let length = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_byte, read_byte); if length == 0 { writer.write_null(col); } else { @@ -1356,7 +1385,7 @@ impl GenericDecoder { } } TdsDataType::DateTimeOffsetN => { - let length = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_byte, read_byte); if length == 0 { writer.write_null(col); } else { @@ -1379,7 +1408,7 @@ impl GenericDecoder { // === GUID === TdsDataType::Guid => { - let length = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_byte, read_byte); if length == 0 { writer.write_null(col); } else { @@ -1414,33 +1443,33 @@ impl SqlTypeDecode for GenericDecoder { { let result = match metadata.data_type { TdsDataType::Int1 => { - let value = reader.read_byte().await?; + let value = read_sync_first!(reader, try_read_byte, read_byte); ColumnValues::from(value) } TdsDataType::Int2 => { - let value = reader.read_int16().await?; + let value = read_sync_first!(reader, try_read_int16, read_int16); ColumnValues::SmallInt(value) } TdsDataType::Int4 => { - let value = reader.read_int32().await?; + let value = read_sync_first!(reader, try_read_int32, read_int32); ColumnValues::from(value) } TdsDataType::Int8 => { - let value = reader.read_int64().await?; + let value = read_sync_first!(reader, try_read_int64, read_int64); ColumnValues::BigInt(value) } TdsDataType::Flt4 => { - let value = reader.read_float32().await?; + let value = read_sync_first!(reader, try_read_float32, read_float32); ColumnValues::Real(value) } TdsDataType::Flt8 => { - let value = reader.read_float64().await?; + let value = read_sync_first!(reader, try_read_float64, read_float64); ColumnValues::Float(value) } TdsDataType::Money4 => ColumnValues::SmallMoney(self.read_money4(reader).await?), TdsDataType::Money => ColumnValues::Money(self.read_money8(reader).await?), TdsDataType::MoneyN => { - let byte_len = reader.read_byte().await?; + let byte_len = read_sync_first!(reader, try_read_byte, read_byte); match byte_len { 4 => ColumnValues::SmallMoney(self.read_money4(reader).await?), 8 => ColumnValues::Money(self.read_money8(reader).await?), @@ -1467,7 +1496,7 @@ impl SqlTypeDecode for GenericDecoder { } } TdsDataType::Bit => { - let value = reader.read_byte().await?; + let value = read_sync_first!(reader, try_read_byte, read_byte); ColumnValues::Bit(value == 1) } TdsDataType::NChar @@ -1483,11 +1512,11 @@ impl SqlTypeDecode for GenericDecoder { ColumnValues::DateTime(value) } TdsDataType::IntN => { - let byte_len = reader.read_byte().await?; + let byte_len = read_sync_first!(reader, try_read_byte, read_byte); self.read_intn(reader, byte_len).await? } TdsDataType::BigBinary => { - let length = reader.read_uint16().await?; + let length = read_sync_first!(reader, try_read_uint16, read_uint16); // 0xFFFF is the USHORTLEN NULL marker (CHARBIN_NULL). if length == 0xFFFF { ColumnValues::Null @@ -1510,7 +1539,7 @@ impl SqlTypeDecode for GenericDecoder { None => ColumnValues::Null, } } else { - let length = reader.read_uint16().await?; + let length = read_sync_first!(reader, try_read_uint16, read_uint16); // 0xFFFF is the USHORTLEN NULL marker (CHARBIN_NULL). if length == 0xFFFF { ColumnValues::Null @@ -1552,34 +1581,34 @@ impl SqlTypeDecode for GenericDecoder { } TdsDataType::Vector => self.decode_vector(reader, metadata).await?, TdsDataType::BitN => { - let byte_len = reader.read_byte().await?; + let byte_len = read_sync_first!(reader, try_read_byte, read_byte); if byte_len > 0 { - let value = reader.read_byte().await?; + let value = read_sync_first!(reader, try_read_byte, read_byte); ColumnValues::Bit(value == 1) } else { ColumnValues::Null } } TdsDataType::Guid => { - let length = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_byte, read_byte); Self::read_guid(reader, length).await? } TdsDataType::FltN => { // This is variable length float, hence the length needs to be read first - let length = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_byte, read_byte); if length == 0 { return Ok(ColumnValues::Null); } if length == 4 { - let value = reader.read_float32().await?; + let value = read_sync_first!(reader, try_read_float32, read_float32); ColumnValues::Real(value) } else { - let value = reader.read_float64().await?; + let value = read_sync_first!(reader, try_read_float64, read_float64); ColumnValues::Float(value) } } TdsDataType::DateTimeN => { - let length = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_byte, read_byte); // If length is 0, then it is NULL if length == 0 { return Ok(ColumnValues::Null); @@ -1593,11 +1622,11 @@ impl SqlTypeDecode for GenericDecoder { } } TdsDataType::DateN => { - let length = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_byte, read_byte); return Self::read_daten(reader, length).await; } TdsDataType::TimeN => { - let length = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_byte, read_byte); match length { 0 => return Ok(ColumnValues::Null), _ => { @@ -1617,7 +1646,7 @@ impl SqlTypeDecode for GenericDecoder { } } TdsDataType::DateTime2N => { - let length = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_byte, read_byte); match length { 0 => Ok(ColumnValues::Null), _ => { @@ -1635,7 +1664,7 @@ impl SqlTypeDecode for GenericDecoder { } }?, TdsDataType::DateTimeOffsetN => { - let length = reader.read_byte().await?; + let length = read_sync_first!(reader, try_read_byte, read_byte); match length { 0 => Ok(ColumnValues::Null), _ => { @@ -1653,13 +1682,13 @@ impl SqlTypeDecode for GenericDecoder { } }?, TdsDataType::Image => { - let text_ptr_len = reader.read_byte().await? as usize; + let text_ptr_len = read_sync_first!(reader, try_read_byte, read_byte) as usize; let length = if text_ptr_len > 0 { const TIMESTAMP_BYTE_COUNT: usize = 8; reader.skip_bytes(text_ptr_len).await?; reader.skip_bytes(TIMESTAMP_BYTE_COUNT).await?; - reader.read_uint32().await? as usize + read_sync_first!(reader, try_read_uint32, read_uint32) as usize } else { 0 }; @@ -1691,8 +1720,8 @@ impl SqlTypeDecode for GenericDecoder { } TdsDataType::SsVariant => self.read_sql_variant(reader).await?, TdsDataType::DateTim4 => { - let daypart = reader.read_uint16().await?; - let timepart = reader.read_uint16().await?; + let daypart = read_sync_first!(reader, try_read_uint16, read_uint16); + let timepart = read_sync_first!(reader, try_read_uint16, read_uint16); ColumnValues::SmallDateTime(SqlSmallDateTime { days: daypart, time: timepart, @@ -1757,7 +1786,7 @@ impl StringDecoder { None => writer.write_null(col), } } else if Self::is_long_len_type(metadata.data_type) { - let text_ptr_len = reader.read_byte().await? as usize; + let text_ptr_len = read_sync_first!(reader, try_read_byte, read_byte) as usize; if text_ptr_len == 0 { writer.write_null(col); @@ -1767,7 +1796,7 @@ impl StringDecoder { const TIMESTAMP_BYTE_COUNT: usize = 8; reader.skip_bytes(text_ptr_len).await?; reader.skip_bytes(TIMESTAMP_BYTE_COUNT).await?; - let length = reader.read_uint32().await? as usize; + let length = read_sync_first!(reader, try_read_uint32, read_uint32) as usize; if length > MAX_ALLOC_SIZE { return Err(crate::error::Error::ProtocolError(format!( @@ -1784,7 +1813,7 @@ impl StringDecoder { }; writer.write_string(col, sql_string); } else { - let length = reader.read_uint16().await? as usize; + let length = read_sync_first!(reader, try_read_uint16, read_uint16) as usize; if length == 0xFFFF { writer.write_null(col); } else { @@ -1837,13 +1866,13 @@ impl SqlTypeDecode for StringDecoder { // Creates SqlString with appropriate encoding type // NULL handling works (textptr_len = 0) // LCID-based decoding implemented (see sql_string.rs) - let text_ptr_len = reader.read_byte().await? as usize; + let text_ptr_len = read_sync_first!(reader, try_read_byte, read_byte) as usize; let length = if text_ptr_len > 0 { const TIMESTAMP_BYTE_COUNT: usize = 8; reader.skip_bytes(text_ptr_len).await?; reader.skip_bytes(TIMESTAMP_BYTE_COUNT).await?; - reader.read_uint32().await? as usize + read_sync_first!(reader, try_read_uint32, read_uint32) as usize } else { // text_ptr_len == 0 means NULL value return Ok(ColumnValues::Null); @@ -1866,7 +1895,7 @@ impl SqlTypeDecode for StringDecoder { }; Ok(ColumnValues::String(sql_string)) } else { - let length = reader.read_uint16().await? as usize; + let length = read_sync_first!(reader, try_read_uint16, read_uint16) as usize; if length == 0xFFFF { Ok(ColumnValues::Null) } else { @@ -2167,7 +2196,7 @@ where Ok(match tds_type { // BIGVARBINARYTYPE, BIGBINARYTYPE TdsDataType::BigVarBinary | TdsDataType::BigBinary => { - let _max_length: u16 = reader.read_uint16().await?; + let _max_length: u16 = read_sync_first!(reader, try_read_uint16, read_uint16); if data_length as usize > MAX_ALLOC_SIZE { return Err(crate::error::Error::ProtocolError(format!( "SQL Variant binary data length {data_length} exceeds maximum allowed size of {MAX_ALLOC_SIZE} bytes" @@ -2178,8 +2207,8 @@ where ColumnValues::Bytes(buffer) } TdsDataType::NumericN | TdsDataType::DecimalN => { - let precision = reader.read_byte().await?; - let scale = reader.read_byte().await?; + let precision = read_sync_first!(reader, try_read_byte, read_byte); + let scale = read_sync_first!(reader, try_read_byte, read_byte); let decimal_parts = GenericDecoder::read_decimal_data(reader, data_length as u8, precision, scale) .await?; @@ -2222,7 +2251,7 @@ where } let mut collation_bytes = vec![0u8; 5]; reader.read_bytes(&mut collation_bytes).await?; - let _max_length = reader.read_uint16().await? as usize; + let _max_length = read_sync_first!(reader, try_read_uint16, read_uint16) as usize; let collation: SqlCollation = collation_bytes.as_slice().try_into()?; if data_length as usize > MAX_ALLOC_SIZE { return Err(crate::error::Error::ProtocolError(format!( diff --git a/mssql-tds/src/io/packet_reader.rs b/mssql-tds/src/io/packet_reader.rs index cf9c9329..617f60f3 100644 --- a/mssql-tds/src/io/packet_reader.rs +++ b/mssql-tds/src/io/packet_reader.rs @@ -10,6 +10,66 @@ pub(crate) const LENGTH_NULL: u16 = 0xffff; #[cfg(not(fuzzing))] pub(crate) trait TdsPacketReader { + /// Returns a buffered byte, or `None` without consuming data if one is unavailable. + #[inline] + fn try_read_byte(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian `i16`, or `None` without consuming partial data. + #[inline] + fn try_read_int16(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian `u16`, or `None` without consuming partial data. + #[inline] + fn try_read_uint16(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian 24-bit integer, or `None` without consuming partial data. + #[inline] + fn try_read_uint24(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian `i32`, or `None` without consuming partial data. + #[inline] + fn try_read_int32(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian `u32`, or `None` without consuming partial data. + #[inline] + fn try_read_uint32(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian 40-bit integer, or `None` without consuming partial data. + #[inline] + fn try_read_uint40(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian `i64`, or `None` without consuming partial data. + #[inline] + fn try_read_int64(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian `f32`, or `None` without consuming partial data. + #[inline] + fn try_read_float32(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian `f64`, or `None` without consuming partial data. + #[inline] + fn try_read_float64(&mut self) -> Option { + None + } + fn read_byte(&mut self) -> impl Future> + Send; fn read_int16_big_endian(&mut self) -> impl Future> + Send; fn read_int32_big_endian(&mut self) -> impl Future> + Send; @@ -49,6 +109,66 @@ pub(crate) trait TdsPacketReader { /// Low-level TDS packet reading operations (public under `fuzzing` cfg). #[cfg(fuzzing)] pub trait TdsPacketReader { + /// Returns a buffered byte, or `None` without consuming data if one is unavailable. + #[inline] + fn try_read_byte(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian `i16`, or `None` without consuming partial data. + #[inline] + fn try_read_int16(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian `u16`, or `None` without consuming partial data. + #[inline] + fn try_read_uint16(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian 24-bit integer, or `None` without consuming partial data. + #[inline] + fn try_read_uint24(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian `i32`, or `None` without consuming partial data. + #[inline] + fn try_read_int32(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian `u32`, or `None` without consuming partial data. + #[inline] + fn try_read_uint32(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian 40-bit integer, or `None` without consuming partial data. + #[inline] + fn try_read_uint40(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian `i64`, or `None` without consuming partial data. + #[inline] + fn try_read_int64(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian `f32`, or `None` without consuming partial data. + #[inline] + fn try_read_float32(&mut self) -> Option { + None + } + + /// Returns a buffered little-endian `f64`, or `None` without consuming partial data. + #[inline] + fn try_read_float64(&mut self) -> Option { + None + } + fn read_byte(&mut self) -> impl Future> + Send; fn read_int16_big_endian(&mut self) -> impl Future> + Send; fn read_int32_big_endian(&mut self) -> impl Future> + Send;