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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions mssql-tds/fuzz/fuzz_targets/fuzz_connection_provider.rs
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@

use libfuzzer_sys::fuzz_target;
use mssql_tds::connection::client_context::ClientContext;
use mssql_tds::fuzz_support::{FuzzReader, MockTransport, TdsConnectionProvider};
use mssql_tds::fuzz_support::{FuzzPacketReader, MockTransport, TdsConnectionProvider};

fuzz_target!(|data: &[u8]| {
// Need at least some data to work with
Expand All @@ -37,7 +37,7 @@ fuzz_target!(|data: &[u8]| {

async fn fuzz_connection_provider(data: &[u8]) {
// Create a fuzz reader with the input data
let reader = Box::new(FuzzReader::new(data));
let reader = FuzzPacketReader::from_data(data);
let packet_size = 4096;

// Create a mock transport with fuzzed data
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@ use arbitrary::Arbitrary;
use libfuzzer_sys::fuzz_target;
use mssql_tds::connection::client_context::{ClientContext, TdsAuthenticationMethod};
use mssql_tds::message::login_options::ApplicationIntent;
use mssql_tds::fuzz_support::{EmptyReader, MockTransport, TdsConnectionProvider};
use mssql_tds::fuzz_support::{FuzzPacketReader, MockTransport, TdsConnectionProvider};

#[derive(Debug, Arbitrary)]
struct FuzzClientContext {
Expand Down Expand Up @@ -101,7 +101,7 @@ fuzz_target!(|fuzz_context: FuzzClientContext| {
});

async fn fuzz_client_context(fuzz_context: FuzzClientContext) {
let reader = Box::new(EmptyReader);
let reader = FuzzPacketReader::empty();
let packet_size = 4096;
let transport = MockTransport::new(reader, packet_size);

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@

use libfuzzer_sys::fuzz_target;
use mssql_tds::connection::client_context::ClientContext;
use mssql_tds::fuzz_support::{FuzzReader, MockTransport, TdsConnectionProvider};
use mssql_tds::fuzz_support::{FuzzPacketReader, MockTransport, TdsConnectionProvider};

fuzz_target!(|data: &[u8]| {
if data.is_empty() {
Expand All @@ -32,7 +32,7 @@ fuzz_target!(|data: &[u8]| {
});

async fn fuzz_network_response(data: &[u8]) {
let reader = Box::new(FuzzReader::new(data));
let reader = FuzzPacketReader::from_data(data);
let packet_size = 4096;
let transport = MockTransport::new(reader, packet_size);

Expand Down
4 changes: 2 additions & 2 deletions mssql-tds/fuzz/fuzz_targets/fuzz_tds_client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@
#![no_main]

use libfuzzer_sys::fuzz_target;
use mssql_tds::fuzz_support::{FuzzReader, create_fuzz_tds_client};
use mssql_tds::fuzz_support::{FuzzPacketReader, create_fuzz_tds_client};

fuzz_target!(|data: &[u8]| {
// Need at least 2 bytes: 1 for scenario, 1+ for token data
Expand All @@ -39,7 +39,7 @@ fuzz_target!(|data: &[u8]| {

rt.block_on(async {
// Create TdsClient with mock transport using fuzzer data
let packet_reader = Box::new(FuzzReader::new(token_data));
let packet_reader = FuzzPacketReader::from_data(token_data);
let mut client = create_fuzz_tds_client(packet_reader, 4096);

// Execute scenario based on fuzzer input
Expand Down
1 change: 0 additions & 1 deletion mssql-tds/src/connection/tds_client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4880,7 +4880,6 @@ mod tests {
}
}

#[async_trait::async_trait]
impl crate::io::packet_reader::TdsPacketReader for TestTransport {
async fn read_byte(&mut self) -> TdsResult<u8> {
Ok(self.take_packet_bytes(1)?[0])
Expand Down
1 change: 0 additions & 1 deletion mssql-tds/src/connection/transport/network_transport.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1099,7 +1099,6 @@ impl TransportSslHandler for NetworkTransport {
}
}

#[async_trait]
impl TdsPacketReader for NetworkTransport {
fn reset_reader(&mut self) {
// Make sure that we have read all the data from the buffer.
Expand Down
95 changes: 58 additions & 37 deletions mssql-tds/src/datatypes/decoder.rs
Original file line number Diff line number Diff line change
@@ -1,10 +1,11 @@
// Copyright (c) Microsoft Corporation.
// Licensed under the MIT License.

use async_trait::async_trait;
use bigdecimal::num_bigint::BigUint;
use bigdecimal::num_traits::ToPrimitive;
use core::fmt;
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::{fmt::Debug, io::Error, vec};

Expand Down Expand Up @@ -144,9 +145,12 @@ macro_rules! safe_vec {
}};
}

#[async_trait]
pub(crate) trait SqlTypeDecode {
async fn decode<T>(&self, reader: &mut T, metadata: &ColumnMetadata) -> TdsResult<ColumnValues>
fn decode<T>(
&self,
reader: &mut T,
metadata: &ColumnMetadata,
) -> impl Future<Output = TdsResult<ColumnValues>> + Send
where
T: TdsPacketReader + Send + Sync;
}
Expand Down Expand Up @@ -201,9 +205,10 @@ impl PlpChunkStreamReader {
}
}

pub(crate) async fn begin(
reader: &mut (dyn TdsPacketReader + Send + Sync),
) -> TdsResult<Option<Self>> {
pub(crate) async fn begin<T>(reader: &mut T) -> TdsResult<Option<Self>>
where
T: TdsPacketReader + Send + Sync,
{
let raw_len_i64 = reader.read_int64().await?;
let raw_len = raw_len_i64 as u64;
let raw_len_usize = raw_len as usize;
Expand Down Expand Up @@ -247,10 +252,10 @@ impl PlpChunkStreamReader {
self.reached_end
}

async fn ensure_active_chunk(
&mut self,
reader: &mut (dyn TdsPacketReader + Send + Sync),
) -> TdsResult<bool> {
async fn ensure_active_chunk<T>(&mut self, reader: &mut T) -> TdsResult<bool>
where
T: TdsPacketReader + Send + Sync,
{
if self.reached_end {
return Ok(false);
}
Expand Down Expand Up @@ -314,11 +319,10 @@ impl PlpChunkStreamReader {
Ok(true)
}

pub(crate) async fn read_into(
&mut self,
reader: &mut (dyn TdsPacketReader + Send + Sync),
out: &mut [u8],
) -> TdsResult<usize> {
pub(crate) async fn read_into<T>(&mut self, reader: &mut T, out: &mut [u8]) -> TdsResult<usize>
where
T: TdsPacketReader + Send + Sync,
{
// Supports the msodbcsql-style cbRequest==0 pattern to consume a
// pending terminator after all data bytes were already read.
if out.is_empty() {
Expand Down Expand Up @@ -356,10 +360,10 @@ impl PlpChunkStreamReader {
Ok(written)
}

pub(crate) async fn skip_to_end(
&mut self,
reader: &mut (dyn TdsPacketReader + Send + Sync),
) -> TdsResult<()> {
pub(crate) async fn skip_to_end<T>(&mut self, reader: &mut T) -> TdsResult<()>
where
T: TdsPacketReader + Send + Sync,
{
while self.ensure_active_chunk(reader).await? {
if self.chunk_remaining > 0 {
reader.skip_bytes(self.chunk_remaining).await?;
Expand Down Expand Up @@ -417,10 +421,13 @@ impl PlpColumnStream {
/// - `Ok(None)` for SQL NULL
/// - `Ok(Some(stream))` ready for incremental reads
/// - `Err` if the column is not PLP-typed or the header is malformed
pub(crate) async fn begin(
pub(crate) async fn begin<T>(
metadata: &ColumnMetadata,
reader: &mut (dyn TdsPacketReader + Send + Sync),
) -> TdsResult<Option<Self>> {
reader: &mut T,
) -> TdsResult<Option<Self>>
where
T: TdsPacketReader + Send + Sync,
{
let (plp_type, collation) = Self::type_from_metadata(metadata)?;
let inner = match PlpChunkStreamReader::begin(reader).await? {
None => return Ok(None),
Expand Down Expand Up @@ -460,19 +467,18 @@ impl PlpColumnStream {
}

/// Incrementally reads PLP payload bytes into `out`.
pub(crate) async fn read_into(
&mut self,
reader: &mut (dyn TdsPacketReader + Send + Sync),
out: &mut [u8],
) -> TdsResult<usize> {
pub(crate) async fn read_into<T>(&mut self, reader: &mut T, out: &mut [u8]) -> TdsResult<usize>
where
T: TdsPacketReader + Send + Sync,
{
self.inner.read_into(reader, out).await
}

/// Discards all remaining PLP payload and terminator bytes.
pub(crate) async fn skip_to_end(
&mut self,
reader: &mut (dyn TdsPacketReader + Send + Sync),
) -> TdsResult<()> {
pub(crate) async fn skip_to_end<T>(&mut self, reader: &mut T) -> TdsResult<()>
where
T: TdsPacketReader + Send + Sync,
{
self.inner.skip_to_end(reader).await
}

Expand Down Expand Up @@ -508,6 +514,24 @@ impl GenericDecoder {
#[cfg(not(fuzzing))]
const MAX_PLP_CHUNK_SIZE: usize = 16 * 1024 * 1024;

/// Boxed re-entry into [`SqlTypeDecode::decode`], used only by SQL_VARIANT.
Comment thread
saurabh500 marked this conversation as resolved.
///
/// SQL_VARIANT embeds a base type, so decoding one re-enters the type switch. Native
/// `async fn` cannot express that cycle: the opaque return type would contain itself.
/// Naming a concrete `Pin<Box<dyn Future>>` in the signature severs the dependency, at
/// the cost of one allocation per nested variant column — a rare type, never on the hot
/// path of ordinary scalar columns.
fn decode_boxed<'a, T>(
&'a self,
reader: &'a mut T,
metadata: &'a ColumnMetadata,
) -> Pin<Box<dyn Future<Output = TdsResult<ColumnValues>> + Send + 'a>>
where
T: TdsPacketReader + Send + Sync,
{
Box::pin(self.decode(reader, metadata))
}

// Reads a SQL_VARIANT type from the TDS stream.
async fn read_sql_variant<T>(&self, reader: &mut T) -> TdsResult<ColumnValues>
where
Expand Down Expand Up @@ -584,7 +608,7 @@ impl GenericDecoder {
multi_part_name: None,
crypto_metadata: None,
};
self.decode(reader, &variant_actual_type_md).await
self.decode_boxed(reader, &variant_actual_type_md).await
}
_ => {
// If the type is not a fixed length type, we should not reach here.
Expand Down Expand Up @@ -1383,7 +1407,6 @@ impl GenericDecoder {
}
}

#[async_trait]
impl SqlTypeDecode for GenericDecoder {
async fn decode<T>(&self, reader: &mut T, metadata: &ColumnMetadata) -> TdsResult<ColumnValues>
where
Expand Down Expand Up @@ -1774,7 +1797,6 @@ impl StringDecoder {
}
}

#[async_trait]
impl SqlTypeDecode for StringDecoder {
async fn decode<T>(&self, reader: &mut T, metadata: &ColumnMetadata) -> TdsResult<ColumnValues>
where
Expand Down Expand Up @@ -1846,7 +1868,7 @@ impl SqlTypeDecode for StringDecoder {
} else {
let length = reader.read_uint16().await? as usize;
if length == 0xFFFF {
return Ok(ColumnValues::Null);
Ok(ColumnValues::Null)
} else {
let mut buffer = vec![0u8; length];
reader.read_bytes(&mut buffer).await?;
Expand Down Expand Up @@ -3229,7 +3251,7 @@ mod test {
}

mod decode_into_tests {
use async_trait::async_trait;

use byteorder::{ByteOrder, LittleEndian};

use crate::core::TdsResult;
Expand Down Expand Up @@ -3270,7 +3292,6 @@ mod test {
}
}

#[async_trait]
impl TdsPacketReader for ByteReader {
async fn read_byte(&mut self) -> TdsResult<u8> {
Ok(self.take(1)?[0])
Expand Down
Loading
Loading