use std::fmt::Debug;
use std::ops::{Deref, DerefMut, Index, IndexMut};
use crate::core::TdsResult;
use crate::error::Error;
pub(crate) struct TdsReadBuffer {
pub(crate) buffer_position: usize,
pub(crate) buffer_length: usize,
pub(crate) max_packet_size: usize,
pub(crate) working_buffer: Vec<u8>,
pub(crate) pending_bytes: usize,
pub(crate) pending_bytes_offset: usize,
pub(crate) end_of_message: bool,
}
impl TdsReadBuffer {
pub(crate) fn new(packet_size: usize) -> Self {
let packet_storage = packet_size * 2;
Self {
buffer_position: 0,
buffer_length: 0,
max_packet_size: packet_size,
working_buffer: vec![0; packet_storage],
pending_bytes: 0,
pending_bytes_offset: 0,
end_of_message: false,
}
}
pub(crate) fn change_packet_size(&mut self, packet_size: u32) {
if packet_size != self.max_packet_size as u32 {
self.max_packet_size = packet_size as usize;
self.working_buffer.resize(packet_size as usize * 2, 0);
self.buffer_position = 0;
self.buffer_length = 0;
self.pending_bytes = 0;
self.pending_bytes_offset = 0;
self.end_of_message = false;
}
}
pub(crate) fn do_we_have_enough_data(&self, byte_count: usize) -> bool {
self.get_remaining_byte_count() >= byte_count
}
pub(crate) fn get_remaining_byte_count(&self) -> usize {
self.buffer_length.saturating_sub(self.buffer_position)
}
#[inline(always)]
fn try_read_array<const N: usize>(&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");
let consumed = self.consume_bytes(N);
debug_assert!(consumed.is_ok(), "capacity is checked at the top");
consumed.ok()?;
Some(bytes)
}
#[inline(always)]
pub(crate) fn try_read_byte(&mut self) -> Option<u8> {
self.try_read_array().map(|[value]| value)
}
#[inline(always)]
pub(crate) fn try_read_slice(&mut self, length: usize) -> Option<&[u8]> {
if !self.do_we_have_enough_data(length) {
return None;
}
let start = self.buffer_position;
let consumed = self.consume_bytes(length);
debug_assert!(consumed.is_ok(), "capacity is checked at the top");
consumed.ok()?;
Some(&self.working_buffer[start..start + length])
}
#[inline(always)]
pub(crate) fn try_read_int16(&mut self) -> Option<i16> {
self.try_read_array().map(i16::from_le_bytes)
}
#[inline(always)]
pub(crate) fn try_read_uint16(&mut self) -> Option<u16> {
self.try_read_array().map(u16::from_le_bytes)
}
#[inline(always)]
pub(crate) fn try_read_uint24(&mut self) -> Option<u32> {
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<i32> {
self.try_read_array().map(i32::from_le_bytes)
}
#[inline(always)]
pub(crate) fn try_read_uint32(&mut self) -> Option<u32> {
self.try_read_array().map(u32::from_le_bytes)
}
#[inline(always)]
pub(crate) fn try_read_uint40(&mut self) -> Option<u64> {
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<i64> {
self.try_read_array().map(i64::from_le_bytes)
}
#[inline(always)]
pub(crate) fn try_read_float32(&mut self) -> Option<f32> {
self.try_read_array().map(f32::from_le_bytes)
}
#[inline(always)]
pub(crate) fn try_read_float64(&mut self) -> Option<f64> {
self.try_read_array().map(f64::from_le_bytes)
}
pub(crate) fn consume_bytes(&mut self, byte_count: usize) -> TdsResult<()> {
let remaining = self.get_remaining_byte_count();
if byte_count > remaining {
return Err(Error::ProtocolError(format!(
"Cannot consume {byte_count} byte(s) from the packet buffer: only {remaining} remain"
)));
}
self.buffer_position += byte_count;
if self.buffer_length == self.buffer_position {
self.buffer_length = 0;
self.buffer_position = 0;
}
Ok(())
}
pub(crate) fn reset_to_length(&mut self, length: usize) {
self.buffer_position = 0;
self.buffer_length = length;
self.end_of_message = false;
}
pub(crate) fn shift_data_to_front(&mut self) {
let remaining = self.get_remaining_byte_count();
self.working_buffer
.copy_within(self.buffer_position..self.buffer_length, 0);
self.buffer_position = 0;
self.buffer_length = remaining;
if self.pending_bytes > 0 {
let pending_src_start = self.pending_bytes_offset;
let pending_src_end = self.pending_bytes_offset + self.pending_bytes;
let pending_dest = remaining;
self.working_buffer
.copy_within(pending_src_start..pending_src_end, pending_dest);
self.pending_bytes_offset = remaining;
}
}
pub(crate) fn remove_header_from_packet(&mut self, new_packet_size: usize) {
self.working_buffer.copy_within(
self.buffer_length + 8..self.buffer_length + new_packet_size,
self.buffer_length,
);
self.buffer_length += new_packet_size - 8;
}
pub(crate) fn get_slice(&self) -> &[u8] {
&self.working_buffer[self.buffer_position..]
}
pub(crate) fn get_buffered_slice(&self) -> &[u8] {
&self.working_buffer[self.buffer_position..self.buffer_length]
}
}
impl Debug for TdsReadBuffer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TdsReadBuffer")
.field("buffer_position", &self.buffer_position)
.field("buffer_length", &self.buffer_length)
.field("max_packet_size", &self.max_packet_size)
.finish()
}
}
impl Index<usize> for TdsReadBuffer {
type Output = u8;
fn index(&self, index: usize) -> &Self::Output {
&self.working_buffer[index]
}
}
impl IndexMut<usize> for TdsReadBuffer {
fn index_mut(&mut self, index: usize) -> &mut Self::Output {
&mut self.working_buffer[index]
}
}
impl Deref for TdsReadBuffer {
type Target = Vec<u8>;
fn deref(&self) -> &Vec<u8> {
&self.working_buffer
}
}
impl DerefMut for TdsReadBuffer {
fn deref_mut(&mut self) -> &mut Vec<u8> {
&mut self.working_buffer
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_buffer_resize_after_packet_size_change() {
let initial_packet_size: usize = 4096;
let mut buffer = TdsReadBuffer::new(initial_packet_size);
assert_eq!(buffer.working_buffer.len(), 8192);
assert_eq!(buffer.max_packet_size, 4096);
let negotiated_packet_size: u32 = 8000;
buffer.reset_to_length(0);
assert_eq!(buffer.working_buffer.len(), 8192);
assert_eq!(buffer.max_packet_size, 4096);
buffer.change_packet_size(negotiated_packet_size);
assert_eq!(buffer.working_buffer.len(), 16000);
assert_eq!(buffer.max_packet_size, 8000);
assert!(buffer.working_buffer.len() >= negotiated_packet_size as usize * 2);
}
#[test]
fn test_change_packet_size_same_size_is_noop() {
let packet_size: usize = 4096;
let mut buffer = TdsReadBuffer::new(packet_size);
buffer.buffer_position = 100;
buffer.buffer_length = 500;
buffer.change_packet_size(packet_size as u32);
assert_eq!(buffer.buffer_position, 100);
assert_eq!(buffer.buffer_length, 500);
assert_eq!(buffer.working_buffer.len(), 8192);
}
#[test]
fn test_change_packet_size_resets_state_on_size_change() {
let initial_size: usize = 4096;
let mut buffer = TdsReadBuffer::new(initial_size);
buffer.buffer_position = 100;
buffer.buffer_length = 500;
buffer.change_packet_size(8000);
assert_eq!(buffer.buffer_position, 0);
assert_eq!(buffer.buffer_length, 0);
assert_eq!(buffer.working_buffer.len(), 16000);
assert_eq!(buffer.max_packet_size, 8000);
}
#[test]
fn test_shift_data_to_front_with_pending_bytes_no_corruption() {
let mut buf = TdsReadBuffer::new(4096);
for i in 82..4088 {
buf.working_buffer[i] = (i % 256) as u8;
}
buf.buffer_position = 82;
buf.buffer_length = 4088;
let pending_start = 4088;
let pending_len = 4096;
for i in 0..pending_len {
buf.working_buffer[pending_start + i] = 0xAA;
}
buf.pending_bytes = pending_len;
buf.pending_bytes_offset = pending_start;
let expected_remaining: Vec<u8> = buf.working_buffer[82..4088].to_vec();
buf.shift_data_to_front();
assert_eq!(
&buf.working_buffer[..4006],
&expected_remaining[..],
"remaining data corrupted after shift_data_to_front"
);
assert!(
buf.working_buffer[4006..4006 + pending_len]
.iter()
.all(|&b| b == 0xAA),
"pending data not correctly placed after remaining"
);
assert_eq!(buf.buffer_position, 0);
assert_eq!(buf.buffer_length, 4006);
assert_eq!(buf.pending_bytes_offset, 4006);
assert_eq!(buf.pending_bytes, pending_len);
}
#[test]
fn test_shift_data_to_front_consumed_remaining_and_pending() {
let mut buf = TdsReadBuffer::new(4096);
for i in 500..2000 {
buf.working_buffer[i] = (i % 256) as u8;
}
buf.buffer_position = 500;
buf.buffer_length = 2000;
let pending_start = 4088;
let pending_len = 200;
for i in 0..pending_len {
buf.working_buffer[pending_start + i] = 0xDD;
}
buf.pending_bytes = pending_len;
buf.pending_bytes_offset = pending_start;
let expected_remaining: Vec<u8> = buf.working_buffer[500..2000].to_vec();
buf.shift_data_to_front();
let remaining = 1500;
assert_eq!(buf.buffer_position, 0);
assert_eq!(buf.buffer_length, remaining);
assert_eq!(
&buf.working_buffer[..remaining],
&expected_remaining[..],
"remaining data corrupted"
);
assert_eq!(buf.pending_bytes_offset, remaining);
assert_eq!(buf.pending_bytes, pending_len);
assert!(
buf.working_buffer[remaining..remaining + pending_len]
.iter()
.all(|&b| b == 0xDD),
"pending data corrupted or misplaced"
);
}
#[test]
fn test_shift_data_to_front_no_pending_bytes() {
let mut buf = TdsReadBuffer::new(4096);
for i in 100..500 {
buf.working_buffer[i] = (i % 256) as u8;
}
buf.buffer_position = 100;
buf.buffer_length = 500;
let expected: Vec<u8> = buf.working_buffer[100..500].to_vec();
buf.shift_data_to_front();
assert_eq!(&buf.working_buffer[..400], &expected[..]);
assert_eq!(buf.buffer_position, 0);
assert_eq!(buf.buffer_length, 400);
}
#[test]
fn test_shift_data_to_front_already_at_zero() {
let mut buf = TdsReadBuffer::new(4096);
for i in 0..200 {
buf.working_buffer[i] = (i % 256) as u8;
}
buf.buffer_position = 0;
buf.buffer_length = 200;
let expected: Vec<u8> = buf.working_buffer[..200].to_vec();
buf.shift_data_to_front();
assert_eq!(&buf.working_buffer[..200], &expected[..]);
assert_eq!(buf.buffer_position, 0);
assert_eq!(buf.buffer_length, 200);
}
#[test]
fn test_shift_data_to_front_no_remaining_with_pending() {
let mut buf = TdsReadBuffer::new(4096);
buf.buffer_position = 0;
buf.buffer_length = 0;
let pending_start = 4088;
for i in 0..100 {
buf.working_buffer[pending_start + i] = 0xBB;
}
buf.pending_bytes = 100;
buf.pending_bytes_offset = pending_start;
buf.shift_data_to_front();
assert_eq!(buf.buffer_position, 0);
assert_eq!(buf.buffer_length, 0);
assert_eq!(buf.pending_bytes_offset, 0);
assert!(buf.working_buffer[..100].iter().all(|&b| b == 0xBB));
}
#[test]
fn test_consume_bytes_partial() {
let mut buf = TdsReadBuffer::new(4096);
buf.buffer_length = 500;
buf.buffer_position = 0;
buf.consume_bytes(200)
.expect("200 of 500 bytes must consume");
assert_eq!(buf.buffer_position, 200);
assert_eq!(buf.buffer_length, 500);
}
#[test]
fn test_consume_bytes_exact_resets() {
let mut buf = TdsReadBuffer::new(4096);
buf.buffer_length = 500;
buf.buffer_position = 100;
buf.consume_bytes(400)
.expect("the exact remaining count must consume");
assert_eq!(buf.buffer_position, 0);
assert_eq!(buf.buffer_length, 0);
}
#[test]
fn test_consume_bytes_over_errors() {
let mut buf = TdsReadBuffer::new(4096);
buf.buffer_length = 500;
buf.buffer_position = 100;
let error = buf
.consume_bytes(401)
.expect_err("consuming past the end must be an error");
assert!(
matches!(error, Error::ProtocolError(ref message) if message.contains("only 400 remain")),
"expected a protocol error naming the shortfall, got {error:?}"
);
assert_eq!(buf.buffer_position, 100, "a rejected consume must not move");
assert_eq!(buf.buffer_length, 500);
}
#[test]
fn remaining_byte_count_saturates_when_position_leads_length() {
let mut buf = TdsReadBuffer::new(4096);
buf.buffer_length = 100;
buf.buffer_position = 250;
assert_eq!(buf.get_remaining_byte_count(), 0);
assert!(!buf.do_we_have_enough_data(1));
assert!(buf.consume_bytes(1).is_err());
}
#[test]
fn test_do_we_have_enough_data() {
let mut buf = TdsReadBuffer::new(4096);
buf.buffer_length = 500;
buf.buffer_position = 100;
assert!(buf.do_we_have_enough_data(400));
assert!(buf.do_we_have_enough_data(1));
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_slice_probe_reads_and_consumes() {
let payload: Vec<u8> = (0..64u8).collect();
let mut buf = TdsReadBuffer::new(4096);
buf.working_buffer[..payload.len()].copy_from_slice(&payload);
buf.reset_to_length(payload.len());
assert_eq!(buf.try_read_slice(16), Some(&payload[..16]));
assert_eq!(buf.buffer_position, 16);
assert_eq!(buf.get_remaining_byte_count(), 48);
assert_eq!(buf.try_read_slice(48), Some(&payload[16..]));
assert_eq!(buf.get_remaining_byte_count(), 0);
}
#[test]
fn test_slice_probe_zero_length_is_a_hit() {
let mut buf = TdsReadBuffer::new(4096);
buf.reset_to_length(0);
assert_eq!(buf.try_read_slice(0), Some(&[][..]));
assert_eq!(buf.buffer_position, 0);
}
#[test]
fn test_slice_probe_miss_does_not_consume() {
let payload: Vec<u8> = (0..8u8).collect();
let mut buf = TdsReadBuffer::new(4096);
buf.working_buffer[..payload.len()].copy_from_slice(&payload);
buf.reset_to_length(payload.len());
assert_eq!(buf.try_read_slice(9), None);
assert_eq!(buf.buffer_position, 0);
assert_eq!(buf.get_remaining_byte_count(), payload.len());
assert_eq!(buf.try_read_slice(8), Some(&payload[..]));
}
#[test]
fn test_slice_probe_miss_preserves_partial_bytes_mid_buffer() {
let payload: Vec<u8> = (0..32u8).collect();
let mut buf = TdsReadBuffer::new(4096);
buf.working_buffer[..payload.len()].copy_from_slice(&payload);
buf.reset_to_length(payload.len());
assert_eq!(buf.try_read_slice(20), Some(&payload[..20]));
assert_eq!(buf.try_read_slice(13), None);
assert_eq!(buf.buffer_position, 20);
assert_eq!(buf.get_remaining_byte_count(), 12);
assert_eq!(buf.try_read_slice(12), Some(&payload[20..]));
}
#[test]
fn test_get_remaining_byte_count() {
let mut buf = TdsReadBuffer::new(4096);
buf.buffer_length = 500;
buf.buffer_position = 100;
assert_eq!(buf.get_remaining_byte_count(), 400);
buf.consume_bytes(200)
.expect("200 of 400 bytes must consume");
assert_eq!(buf.get_remaining_byte_count(), 200);
}
#[test]
fn test_reset_to_length() {
let mut buf = TdsReadBuffer::new(4096);
buf.buffer_position = 250;
buf.buffer_length = 500;
buf.reset_to_length(1000);
assert_eq!(buf.buffer_position, 0);
assert_eq!(buf.buffer_length, 1000);
}
#[test]
fn test_remove_header_from_packet() {
let mut buf = TdsReadBuffer::new(4096);
buf.buffer_length = 100;
let header = [0x04, 0x00, 0x00, 0x64, 0x00, 0x00, 0x01, 0x00];
buf.working_buffer[100..108].copy_from_slice(&header);
for i in 108..200 {
buf.working_buffer[i] = 0xCC;
}
buf.remove_header_from_packet(100);
assert_eq!(buf.buffer_length, 192);
assert!(buf.working_buffer[100..192].iter().all(|&b| b == 0xCC));
}
#[test]
fn test_get_slice_returns_from_position() {
let mut buf = TdsReadBuffer::new(4096);
buf.working_buffer[50] = 0xDE;
buf.working_buffer[51] = 0xAD;
buf.buffer_position = 50;
let slice = buf.get_slice();
assert_eq!(slice[0], 0xDE);
assert_eq!(slice[1], 0xAD);
}
}