use core::fmt;
use std::borrow::Cow;
use std::future::Future;
use std::mem::MaybeUninit;
use std::pin::Pin;
use std::sync::Arc;
use std::{fmt::Debug, io::Error, vec};
use super::{
sql_string::{SqlString, get_encoding_type},
sqldatatypes::{PartialLengthType, TdsDataType, TypeInfoVariant},
};
use crate::datatypes::sqldatatypes::TypeInfo;
use crate::security::cell_decryptor::CellDecryptor;
use crate::{
core::TdsResult,
datatypes::{sql_json::SqlJson, sql_string::EncodingType, sqldatatypes::FixedLengthTypes},
};
use crate::{
datatypes::column_values::{
ColumnValues, SqlDate, SqlDateTime, SqlDateTime2, SqlDateTimeOffset, SqlMoney,
SqlSmallDateTime, SqlSmallMoney, SqlTime, SqlXml,
},
io::packet_reader::TdsPacketReader,
};
use crate::{query::metadata::ColumnMetadata, token::tokens::SqlCollation};
use super::row_writer::{RowWriter, ValueKind, write_column_value};
macro_rules! read_sync_first {
($reader:expr, try_read_byte, read_byte) => {
read_sync_first!(@pair $reader, try_read_byte, read_byte)
};
($reader:expr, try_read_int16, read_int16) => {
read_sync_first!(@pair $reader, try_read_int16, read_int16)
};
($reader:expr, try_read_uint16, read_uint16) => {
read_sync_first!(@pair $reader, try_read_uint16, read_uint16)
};
($reader:expr, try_read_uint24, read_uint24) => {
read_sync_first!(@pair $reader, try_read_uint24, read_uint24)
};
($reader:expr, try_read_int32, read_int32) => {
read_sync_first!(@pair $reader, try_read_int32, read_int32)
};
($reader:expr, try_read_uint32, read_uint32) => {
read_sync_first!(@pair $reader, try_read_uint32, read_uint32)
};
($reader:expr, try_read_uint40, read_uint40) => {
read_sync_first!(@pair $reader, try_read_uint40, read_uint40)
};
($reader:expr, try_read_int64, read_int64) => {
read_sync_first!(@pair $reader, try_read_int64, read_int64)
};
($reader:expr, try_read_float32, read_float32) => {
read_sync_first!(@pair $reader, try_read_float32, read_float32)
};
($reader:expr, try_read_float64, read_float64) => {
read_sync_first!(@pair $reader, try_read_float64, read_float64)
};
(@pair $reader:expr, $try_method:ident, $read_method:ident) => {{
let reader = &mut *($reader);
match reader.$try_method() {
Some(value) => value,
None => reader.$read_method().await?,
}
}};
}
macro_rules! read_value_into {
(
$writer:expr, $col:expr, $kind:expr, $length:expr,
|$dest:ident| $fill:expr,
|$bytes:ident| $owned:expr $(,)?
) => {{
let col = $col;
let length = $length;
match $writer.value_destination(col, $kind, length) {
Some($dest) => {
debug_assert_eq!(
$dest.len(),
length,
"value_destination returned {} bytes for a {length}-byte value",
$dest.len()
);
let outcome = $fill;
$writer.commit_value(col, outcome.is_ok());
outcome?;
}
None => {
let mut $bytes: Vec<u8> = Vec::with_capacity(length);
{
let $dest: &mut [MaybeUninit<u8>] = &mut $bytes.spare_capacity_mut()[..length];
$fill?;
}
unsafe { $bytes.set_len(length) };
$owned;
}
}
}};
}
pub(crate) async fn decrypt_encrypted_column<D, T>(
decoder: &D,
reader: &mut T,
metadata: &ColumnMetadata,
decryptor: &Arc<dyn CellDecryptor>,
) -> TdsResult<ColumnValues>
where
D: SqlTypeDecode,
T: TdsPacketReader + Send + Sync,
{
let cipher = match decoder.decode(reader, metadata).await? {
ColumnValues::Null => return Ok(ColumnValues::Null),
ColumnValues::Bytes(bytes) => bytes,
other => {
return Err(crate::error::Error::ColumnEncryptionError(format!(
"Encrypted column '{}' was expected to arrive as varbinary cipher bytes, but \
decoded as {other:?}",
metadata.column_name
)));
}
};
let crypto_metadata = metadata.crypto_metadata.as_ref().ok_or_else(|| {
crate::error::Error::ColumnEncryptionError(format!(
"decrypt_encrypted_column called for non-encrypted column '{}'",
metadata.column_name
))
})?;
decryptor.decrypt(crypto_metadata, &cipher)
}
#[cfg(fuzzing)]
const MAX_ALLOC_SIZE: usize = 64 * 1024; #[cfg(not(fuzzing))]
const MAX_ALLOC_SIZE: usize = 100 * 1024 * 1024;
#[cfg(fuzzing)]
const MAX_PLP_SIZE: usize = 64 * 1024; #[cfg(not(fuzzing))]
const MAX_PLP_SIZE: usize = i32::MAX as usize;
const MAX_DECIMAL_INT_PARTS: usize = 4;
const MAX_DECIMAL_PRECISION: u8 = 38;
pub(crate) const fn decimal_metadata_is_valid(precision: u8, scale: u8) -> bool {
precision > 0 && precision <= MAX_DECIMAL_PRECISION && scale <= precision
}
#[inline]
const fn scale_time_value(value: u64, scale: u8) -> u64 {
match scale {
0 => value * 10_000_000,
1 => value * 1_000_000,
2 => value * 100_000,
3 => value * 10_000,
4 => value * 1_000,
5 => value * 100,
6 => value * 10,
_ => value,
}
}
pub(crate) const DECIMAL_MAGNITUDE_BYTES: usize = MAX_DECIMAL_INT_PARTS * 4;
#[inline]
fn validate_alloc_size(size: usize, context: &str) -> TdsResult<()> {
if size > MAX_ALLOC_SIZE {
#[cfg(fuzzing)]
{
use std::io::Write;
let _ = writeln!(
std::io::stderr(),
"[ALLOC-REJECT] {} requesting {} bytes (max {})",
context,
size,
MAX_ALLOC_SIZE
);
}
return Err(crate::error::Error::ProtocolError(format!(
"{context}: allocation size {size} exceeds maximum allowed {MAX_ALLOC_SIZE} bytes"
)));
}
#[cfg(fuzzing)]
{
use std::io::Write;
let _ = writeln!(
std::io::stderr(),
"[ALLOC-OK] {} requesting {} bytes",
context,
size
);
}
Ok(())
}
#[cfg(fuzzing)]
macro_rules! safe_vec {
($elem:expr; $size:expr, $context:expr) => {{
let size = $size;
validate_alloc_size(size, $context)?;
vec![$elem; size]
}};
}
#[cfg(not(fuzzing))]
macro_rules! safe_vec {
($elem:expr; $size:expr, $context:expr) => {{
let size = $size;
validate_alloc_size(size, $context)?;
vec![$elem; size]
}};
}
pub(crate) trait SqlTypeDecode {
fn decode<T>(
&self,
reader: &mut T,
metadata: &ColumnMetadata,
) -> impl Future<Output = TdsResult<ColumnValues>> + Send
where
T: TdsPacketReader + Send + Sync;
}
impl From<u8> for ColumnValues {
fn from(value: u8) -> Self {
ColumnValues::TinyInt(value)
}
}
impl From<i32> for ColumnValues {
fn from(value: i32) -> Self {
ColumnValues::Int(value)
}
}
#[derive(Debug, Default)]
pub(crate) struct GenericDecoder {
string_decoder: StringDecoder,
}
struct BufferedSlice<'a> {
bytes: &'a [u8],
position: usize,
}
impl<'a> BufferedSlice<'a> {
fn new(bytes: &'a [u8]) -> Self {
Self { bytes, position: 0 }
}
fn take<const N: usize>(&mut self) -> Option<[u8; N]> {
let end = self.position.checked_add(N)?;
let value = self.bytes.get(self.position..end)?.try_into().ok()?;
self.position = end;
Some(value)
}
fn take_bytes(&mut self, len: usize) -> Option<Vec<u8>> {
self.take_slice(len).map(<[u8]>::to_vec)
}
fn take_slice(&mut self, len: usize) -> Option<&'a [u8]> {
let end = self.position.checked_add(len)?;
let value = self.bytes.get(self.position..end)?;
self.position = end;
Some(value)
}
fn byte(&mut self) -> Option<u8> {
self.take().map(|[value]| value)
}
fn i16(&mut self) -> Option<i16> {
self.take().map(i16::from_le_bytes)
}
fn u16(&mut self) -> Option<u16> {
self.take().map(u16::from_le_bytes)
}
fn u24(&mut self) -> Option<u32> {
let [b0, b1, b2] = self.take()?;
Some(u32::from_le_bytes([b0, b1, b2, 0]))
}
fn i32(&mut self) -> Option<i32> {
self.take().map(i32::from_le_bytes)
}
fn u32(&mut self) -> Option<u32> {
self.take().map(u32::from_le_bytes)
}
fn u40(&mut self) -> Option<u64> {
let [b0, b1, b2, b3, b4] = self.take()?;
Some(u64::from_le_bytes([b0, b1, b2, b3, b4, 0, 0, 0]))
}
fn i64(&mut self) -> Option<i64> {
self.take().map(i64::from_le_bytes)
}
fn f32(&mut self) -> Option<f32> {
self.take().map(f32::from_le_bytes)
}
fn f64(&mut self) -> Option<f64> {
self.take().map(f64::from_le_bytes)
}
}
#[cfg_attr(not(test), allow(dead_code))]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum PlpChunkReadLength {
Unknown,
Known(u64),
}
#[cfg_attr(not(test), allow(dead_code))]
#[derive(Debug, Clone)]
pub(crate) struct PlpChunkStreamReader {
length: PlpChunkReadLength,
chunk_remaining: usize,
reached_end: bool,
total_read: usize,
}
#[cfg_attr(not(test), allow(dead_code))]
impl PlpChunkStreamReader {
fn new(length: PlpChunkReadLength) -> Self {
Self {
length,
chunk_remaining: 0,
reached_end: false,
total_read: 0,
}
}
pub(crate) async fn begin<T>(reader: &mut T) -> TdsResult<Option<Self>>
where
T: TdsPacketReader + Send + Sync,
{
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;
if raw_len_usize == GenericDecoder::SQL_PLP_NULL {
return Ok(None);
}
let length = if raw_len_usize == GenericDecoder::SQL_PLP_UNKNOWNLEN
|| raw_len_usize == GenericDecoder::SQL_PLP_MAXLEN
{
PlpChunkReadLength::Unknown
} else {
let declared_len = raw_len as usize;
if raw_len_i64 < 0 || declared_len > MAX_PLP_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"PLP length {declared_len} (raw i64: {raw_len_i64}) exceeds maximum allowed size of {MAX_PLP_SIZE} bytes"
)));
}
PlpChunkReadLength::Known(raw_len)
};
Ok(Some(Self::new(length)))
}
fn try_begin_buffered(bytes: &[u8]) -> TdsResult<Option<(Option<Self>, usize)>> {
let Some(header) = bytes.get(..8) else {
return Ok(None);
};
let raw_len_i64 = i64::from_le_bytes(header.try_into().map_err(|_| {
crate::error::Error::ProtocolError("Invalid buffered PLP header".to_string())
})?);
let raw_len = raw_len_i64 as u64;
let raw_len_usize = raw_len as usize;
if raw_len_usize == GenericDecoder::SQL_PLP_NULL {
return Ok(Some((None, 8)));
}
let length = if raw_len_usize == GenericDecoder::SQL_PLP_UNKNOWNLEN
|| raw_len_usize == GenericDecoder::SQL_PLP_MAXLEN
{
PlpChunkReadLength::Unknown
} else {
let declared_len = raw_len as usize;
if raw_len_i64 < 0 || declared_len > MAX_PLP_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"PLP length {declared_len} (raw i64: {raw_len_i64}) exceeds maximum allowed size of {MAX_PLP_SIZE} bytes"
)));
}
PlpChunkReadLength::Known(raw_len)
};
Ok(Some((Some(Self::new(length)), 8)))
}
fn try_read_complete_buffered(
&mut self,
bytes: &[u8],
out: &mut [u8],
) -> TdsResult<Option<(usize, usize)>> {
let PlpChunkReadLength::Known(known_len) = self.length else {
return Ok(None);
};
if self.reached_end || self.chunk_remaining != 0 || self.total_read != 0 {
return Ok(None);
}
let length = usize::try_from(known_len).map_err(|_| {
crate::error::Error::ProtocolError("PLP length does not fit usize".to_string())
})?;
if length > out.len() {
return Ok(None);
}
let Some(chunk_header) = bytes.get(..4) else {
return Ok(None);
};
let chunk_len = u32::from_le_bytes(chunk_header.try_into().map_err(|_| {
crate::error::Error::ProtocolError("Invalid buffered PLP chunk header".to_string())
})?) as usize;
if length == 0 && chunk_len == 0 {
self.reached_end = true;
return Ok(Some((4, 0)));
}
if chunk_len != length {
return Ok(None);
}
let payload_end = 4usize.checked_add(length).ok_or_else(|| {
crate::error::Error::ProtocolError("Buffered PLP length overflowed".to_string())
})?;
let terminator_end = payload_end.checked_add(4).ok_or_else(|| {
crate::error::Error::ProtocolError("Buffered PLP terminator overflowed".to_string())
})?;
let Some(payload) = bytes.get(4..payload_end) else {
return Ok(None);
};
let Some(terminator) = bytes.get(payload_end..terminator_end) else {
return Ok(None);
};
if u32::from_le_bytes(terminator.try_into().map_err(|_| {
crate::error::Error::ProtocolError("Invalid buffered PLP terminator".to_string())
})?) != 0
{
return Ok(None);
}
out.get_mut(..length)
.ok_or_else(|| {
crate::error::Error::ProtocolError(
"Buffered PLP output range was unavailable".to_string(),
)
})?
.copy_from_slice(payload);
self.total_read = length;
self.reached_end = true;
Ok(Some((terminator_end, length)))
}
fn accept_chunk_length(&mut self, chunk_len: usize) -> TdsResult<bool> {
if chunk_len == 0 {
self.reached_end = true;
if let PlpChunkReadLength::Known(known_len) = self.length
&& self.total_read != known_len as usize
{
return Err(crate::error::Error::ProtocolError(format!(
"PLP stream ended before declared length was reached: total_read={}, declared_len={known_len}",
self.total_read
)));
}
return Ok(false);
}
if chunk_len > GenericDecoder::MAX_PLP_CHUNK_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"PLP chunk size {chunk_len} exceeds maximum allowed chunk size of {} bytes",
GenericDecoder::MAX_PLP_CHUNK_SIZE
)));
}
let next_total = self.total_read.checked_add(chunk_len).ok_or_else(|| {
crate::error::Error::ProtocolError(format!(
"PLP chunk accumulation would overflow capacity: {} + {chunk_len}",
self.total_read
))
})?;
if next_total > MAX_PLP_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"PLP accumulated size {next_total} exceeds maximum allowed size of {MAX_PLP_SIZE} bytes (SQL Server limit: 2GB)"
)));
}
if let PlpChunkReadLength::Known(known_len) = self.length
&& next_total > known_len as usize
{
return Err(crate::error::Error::ProtocolError(format!(
"PLP chunk exceeds declared length: accumulated={next_total}, declared_len={known_len}"
)));
}
self.chunk_remaining = chunk_len;
Ok(true)
}
fn try_ensure_active_buffered_chunk(
&mut self,
bytes: &[u8],
position: &mut usize,
) -> TdsResult<Option<bool>> {
if self.reached_end {
return Ok(Some(false));
}
if self.chunk_remaining > 0 {
return Ok(Some(true));
}
let Some(end) = position.checked_add(4) else {
return Err(crate::error::Error::ProtocolError(
"Buffered PLP position overflowed".to_string(),
));
};
let Some(header) = bytes.get(*position..end) else {
return Ok(None);
};
let chunk_len = u32::from_le_bytes(header.try_into().map_err(|_| {
crate::error::Error::ProtocolError("Invalid buffered PLP chunk header".to_string())
})?) as usize;
*position = end;
self.accept_chunk_length(chunk_len).map(Some)
}
fn try_read_buffered_inner(
&mut self,
bytes: &[u8],
mut out: Option<&mut [u8]>,
out_len: usize,
) -> TdsResult<Option<(usize, usize)>> {
let mut position = 0;
if out_len == 0 {
return match self.try_ensure_active_buffered_chunk(bytes, &mut position)? {
Some(_) => Ok(Some((position, 0))),
None => Ok(None),
};
}
let mut written = 0;
while written < out_len {
let Some(active) = self.try_ensure_active_buffered_chunk(bytes, &mut position)? else {
return Ok(None);
};
if !active {
break;
}
let to_read = std::cmp::min(out_len - written, self.chunk_remaining);
let Some(end) = position.checked_add(to_read) else {
return Err(crate::error::Error::ProtocolError(
"Buffered PLP position overflowed".to_string(),
));
};
let Some(payload) = bytes.get(position..end) else {
return Ok(None);
};
if let Some(output) = out.as_deref_mut() {
let Some(target) = output.get_mut(written..written + to_read) else {
return Err(crate::error::Error::ProtocolError(
"Buffered PLP output range was unavailable".to_string(),
));
};
target.copy_from_slice(payload);
}
position = end;
self.chunk_remaining -= to_read;
self.total_read += to_read;
written += to_read;
}
if written == out_len && self.chunk_remaining == 0 && !self.reached_end {
let Some(_) = self.try_ensure_active_buffered_chunk(bytes, &mut position)? else {
return Ok(None);
};
}
Ok(Some((position, written)))
}
fn try_read_buffered(
&mut self,
bytes: &[u8],
out: &mut [u8],
) -> TdsResult<Option<(usize, usize)>> {
let out_len = out.len();
let mut probe = self.clone();
if probe
.try_read_buffered_inner(bytes, None, out_len)?
.is_none()
{
return Ok(None);
}
self.try_read_buffered_inner(bytes, Some(out), out_len)
}
pub(crate) fn total_read(&self) -> usize {
self.total_read
}
pub(crate) fn known_len(&self) -> Option<u64> {
match self.length {
PlpChunkReadLength::Known(n) => Some(n),
PlpChunkReadLength::Unknown => None,
}
}
pub(crate) fn reached_end(&self) -> bool {
self.reached_end
}
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);
}
if self.chunk_remaining > 0 {
return Ok(true);
}
let chunk_len = read_sync_first!(reader, try_read_uint32, read_uint32) as usize;
self.accept_chunk_length(chunk_len)
}
pub(crate) async fn read_into<T>(&mut self, reader: &mut T, out: &mut [u8]) -> TdsResult<usize>
where
T: TdsPacketReader + Send + Sync,
{
if out.is_empty() {
let _ = self.ensure_active_chunk(reader).await?;
return Ok(0);
}
let mut written = 0;
while written < out.len() {
if !self.ensure_active_chunk(reader).await? {
break;
}
let to_read = std::cmp::min(out.len() - written, self.chunk_remaining);
let slice = &mut out[written..written + to_read];
let bytes_read = reader.read_bytes(slice).await?;
if to_read > 0 && bytes_read == 0 {
return Err(crate::error::Error::ProtocolError(
"PLP stream read made no progress while bytes were requested".to_string(),
));
}
self.chunk_remaining -= bytes_read;
self.total_read += bytes_read;
written += bytes_read;
}
if written == out.len() && self.chunk_remaining == 0 && !self.reached_end {
let _ = self.ensure_active_chunk(reader).await?;
}
Ok(written)
}
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?;
self.total_read += self.chunk_remaining;
self.chunk_remaining = 0;
}
}
Ok(())
}
}
#[cfg_attr(not(test), allow(dead_code))]
#[derive(Debug, Clone)]
pub(crate) struct PlpColumnStream {
plp_type: PartialLengthType,
collation: Option<SqlCollation>,
inner: PlpChunkStreamReader,
}
#[cfg_attr(not(test), allow(dead_code))]
impl PlpColumnStream {
pub(crate) async fn begin<T>(
metadata: &ColumnMetadata,
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),
Some(r) => r,
};
Ok(Some(Self {
plp_type,
collation,
inner,
}))
}
pub(crate) fn try_begin_buffered(
metadata: &ColumnMetadata,
bytes: &[u8],
) -> TdsResult<Option<(Option<Self>, usize)>> {
let (plp_type, collation) = Self::type_from_metadata(metadata)?;
let Some((inner, used)) = PlpChunkStreamReader::try_begin_buffered(bytes)? else {
return Ok(None);
};
Ok(Some((
inner.map(|inner| Self {
plp_type,
collation,
inner,
}),
used,
)))
}
pub(crate) fn try_read_complete_buffered(
&mut self,
bytes: &[u8],
out: &mut [u8],
) -> TdsResult<Option<(usize, usize)>> {
self.inner.try_read_complete_buffered(bytes, out)
}
pub(crate) fn try_read_buffered(
&mut self,
bytes: &[u8],
out: &mut [u8],
) -> TdsResult<Option<(usize, usize)>> {
self.inner.try_read_buffered(bytes, out)
}
pub(crate) fn plp_type(&self) -> PartialLengthType {
self.plp_type
}
pub(crate) fn collation(&self) -> Option<SqlCollation> {
self.collation
}
pub(crate) fn total_read(&self) -> usize {
self.inner.total_read()
}
pub(crate) fn known_len(&self) -> Option<u64> {
self.inner.known_len()
}
pub(crate) fn reached_end(&self) -> bool {
self.inner.reached_end()
}
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
}
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
}
fn type_from_metadata(
metadata: &ColumnMetadata,
) -> TdsResult<(PartialLengthType, Option<SqlCollation>)> {
match metadata.type_info.type_info_variant {
TypeInfoVariant::PartialLen(pt, _, collation, _, _) => Ok((pt, collation)),
_ => Err(crate::error::Error::ProtocolError(format!(
"Column '{}' (type {:?}) is not a PLP type",
metadata.column_name, metadata.data_type
))),
}
}
}
impl GenericDecoder {
pub(crate) fn try_decode_buffered_variant(
&self,
bytes: &[u8],
) -> TdsResult<Option<(Option<TdsDataType>, ColumnValues, usize)>> {
let mut outer = BufferedSlice::new(bytes);
let Some(length) = outer.u32() else {
return Ok(None);
};
if length == 0 {
return Ok(Some((None, ColumnValues::Null, outer.position)));
}
let Some(payload) = outer.take_slice(length as usize) else {
return Ok(None);
};
let mut reader = BufferedSlice::new(payload);
let Some(base_byte) = reader.byte() else {
return Ok(None);
};
let base = TdsDataType::try_from(base_byte)?;
let Some(property_bytes) = reader.byte() else {
return Ok(None);
};
let data_length = usize::try_from(length)
.ok()
.and_then(|length| length.checked_sub(2 + usize::from(property_bytes)))
.ok_or_else(|| {
crate::error::Error::ProtocolError(format!(
"SQL_VARIANT data length calculation underflow: length={length}, prop_bytes={property_bytes}"
))
})?;
let value = match property_bytes {
0 => {
let Ok(fixed_type) = FixedLengthTypes::try_from(base) else {
return Ok(None);
};
let metadata = ColumnMetadata {
user_type: 0,
flags: 0,
type_info: TypeInfo {
tds_type: base,
length: data_length,
type_info_variant: TypeInfoVariant::FixedLen(fixed_type),
},
data_type: base,
column_name: String::new(),
multi_part_name: None,
crypto_metadata: None,
};
let Some((value, used)) =
self.try_decode_buffered(&payload[reader.position..], &metadata)?
else {
return Ok(None);
};
if used != data_length {
return Ok(None);
}
reader.position += used;
value
}
7 if matches!(
base,
TdsDataType::BigVarChar
| TdsDataType::BigChar
| TdsDataType::NVarChar
| TdsDataType::NChar
) =>
{
let Some(collation_bytes) = reader.take::<5>() else {
return Ok(None);
};
let Some(_max_length) = reader.u16() else {
return Ok(None);
};
let Some(value_bytes) = reader.take_slice(data_length) else {
return Ok(None);
};
let encoding = if matches!(base, TdsDataType::NVarChar | TdsDataType::NChar) {
EncodingType::Utf16
} else {
let collation: SqlCollation = collation_bytes.as_slice().try_into()?;
if collation.utf8() {
EncodingType::Utf8
} else {
EncodingType::LcidBased(collation)
}
};
ColumnValues::String(SqlString::new(value_bytes.to_vec(), encoding))
}
_ => return Ok(None),
};
if reader.position != payload.len() {
return Ok(None);
}
Ok(Some((Some(base), value, outer.position)))
}
pub(crate) fn try_decode_buffered(
&self,
bytes: &[u8],
metadata: &ColumnMetadata,
) -> TdsResult<Option<(ColumnValues, usize)>> {
if metadata.is_plp() || metadata.crypto_metadata.is_some() {
return Ok(None);
}
let mut reader = BufferedSlice::new(bytes);
let value = match metadata.data_type {
TdsDataType::Int1 => ColumnValues::TinyInt(match reader.byte() {
Some(value) => value,
None => return Ok(None),
}),
TdsDataType::Int2 => ColumnValues::SmallInt(match reader.i16() {
Some(value) => value,
None => return Ok(None),
}),
TdsDataType::Int4 => ColumnValues::Int(match reader.i32() {
Some(value) => value,
None => return Ok(None),
}),
TdsDataType::Int8 => ColumnValues::BigInt(match reader.i64() {
Some(value) => value,
None => return Ok(None),
}),
TdsDataType::Flt4 => ColumnValues::Real(match reader.f32() {
Some(value) => value,
None => return Ok(None),
}),
TdsDataType::Flt8 => ColumnValues::Float(match reader.f64() {
Some(value) => value,
None => return Ok(None),
}),
TdsDataType::Bit => ColumnValues::Bit(match reader.byte() {
Some(value) => value == 1,
None => return Ok(None),
}),
TdsDataType::IntN => {
let Some(length) = reader.byte() else {
return Ok(None);
};
match length {
0 => ColumnValues::Null,
1 => ColumnValues::TinyInt(match reader.byte() {
Some(value) => value,
None => return Ok(None),
}),
2 => ColumnValues::SmallInt(match reader.i16() {
Some(value) => value,
None => return Ok(None),
}),
4 => ColumnValues::Int(match reader.i32() {
Some(value) => value,
None => return Ok(None),
}),
8 => ColumnValues::BigInt(match reader.i64() {
Some(value) => value,
None => return Ok(None),
}),
_ => {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid IntN length: {length}"
)));
}
}
}
TdsDataType::FltN => {
let Some(length) = reader.byte() else {
return Ok(None);
};
match length {
0 => ColumnValues::Null,
4 => ColumnValues::Real(match reader.f32() {
Some(value) => value,
None => return Ok(None),
}),
_ => ColumnValues::Float(match reader.f64() {
Some(value) => value,
None => return Ok(None),
}),
}
}
TdsDataType::BitN => {
let Some(length) = reader.byte() else {
return Ok(None);
};
match length {
0 => ColumnValues::Null,
_ => ColumnValues::Bit(match reader.byte() {
Some(value) => value == 1,
None => return Ok(None),
}),
}
}
TdsDataType::Money4 => ColumnValues::SmallMoney(SqlSmallMoney {
int_val: match reader.i32() {
Some(value) => value,
None => return Ok(None),
},
}),
TdsDataType::Money => {
let Some(msb_part) = reader.i32() else {
return Ok(None);
};
let Some(lsb_part) = reader.i32() else {
return Ok(None);
};
ColumnValues::Money(SqlMoney { lsb_part, msb_part })
}
TdsDataType::MoneyN => {
let Some(length) = reader.byte() else {
return Ok(None);
};
match length {
0 => ColumnValues::Null,
4 => ColumnValues::SmallMoney(SqlSmallMoney {
int_val: match reader.i32() {
Some(value) => value,
None => return Ok(None),
},
}),
8 => {
let Some(msb_part) = reader.i32() else {
return Ok(None);
};
let Some(lsb_part) = reader.i32() else {
return Ok(None);
};
ColumnValues::Money(SqlMoney { lsb_part, msb_part })
}
_ => {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid MoneyN length: {length}"
)));
}
}
}
TdsDataType::DateTim4 => {
let Some(days) = reader.u16() else {
return Ok(None);
};
let Some(time) = reader.u16() else {
return Ok(None);
};
ColumnValues::SmallDateTime(SqlSmallDateTime { days, time })
}
TdsDataType::DateTime => {
let Some(days) = reader.i32() else {
return Ok(None);
};
let Some(time) = reader.u32() else {
return Ok(None);
};
ColumnValues::DateTime(SqlDateTime { days, time })
}
TdsDataType::DateTimeN => {
let Some(length) = reader.byte() else {
return Ok(None);
};
match length {
0 => ColumnValues::Null,
4 => {
let Some(days) = reader.u16() else {
return Ok(None);
};
let Some(time) = reader.u16() else {
return Ok(None);
};
ColumnValues::SmallDateTime(SqlSmallDateTime { days, time })
}
_ => {
let Some(days) = reader.i32() else {
return Ok(None);
};
let Some(time) = reader.u32() else {
return Ok(None);
};
ColumnValues::DateTime(SqlDateTime { days, time })
}
}
}
TdsDataType::DateN => {
let Some(length) = reader.byte() else {
return Ok(None);
};
match length {
0 => ColumnValues::Null,
_ => ColumnValues::Date(SqlDate::unchecked_create(match reader.u24() {
Some(value) => value,
None => return Ok(None),
})),
}
}
TdsDataType::TimeN | TdsDataType::DateTime2N | TdsDataType::DateTimeOffsetN => {
let Some(length) = reader.byte() else {
return Ok(None);
};
if length == 0 {
ColumnValues::Null
} else {
let scale = metadata.get_scale().ok_or_else(|| {
crate::error::Error::ImplementationError(format!(
"{:?} type should have scale",
metadata.data_type
))
})?;
let trailing = match metadata.data_type {
TdsDataType::TimeN => 0,
TdsDataType::DateTime2N => 3,
TdsDataType::DateTimeOffsetN => 5,
_ => {
return Err(crate::error::Error::ImplementationError(
"unexpected buffered time type".to_string(),
));
}
};
let time_length = length.checked_sub(trailing).ok_or_else(|| {
crate::error::Error::ProtocolError(format!(
"Invalid {:?} length: {length}",
metadata.data_type
))
})?;
let scaled = match time_length {
3 => reader.u24().map(u64::from),
4 => reader.u32().map(u64::from),
_ => reader.u40(),
};
let Some(scaled) = scaled else {
return Ok(None);
};
let time = SqlTime {
time_nanoseconds: scale_time_value(scaled, scale),
scale,
};
match metadata.data_type {
TdsDataType::TimeN => ColumnValues::Time(time),
TdsDataType::DateTime2N => {
let Some(days) = reader.u24() else {
return Ok(None);
};
ColumnValues::DateTime2(SqlDateTime2 { days, time })
}
TdsDataType::DateTimeOffsetN => {
let Some(days) = reader.u24() else {
return Ok(None);
};
let Some(offset) = reader.i16() else {
return Ok(None);
};
ColumnValues::DateTimeOffset(SqlDateTimeOffset {
datetime2: SqlDateTime2 { days, time },
offset,
})
}
_ => {
return Err(crate::error::Error::ImplementationError(
"unexpected buffered time type".to_string(),
));
}
}
}
}
TdsDataType::Guid => {
let Some(length) = reader.byte() else {
return Ok(None);
};
if length == 0 {
ColumnValues::Null
} else {
if length != 16 {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid GUID length: expected 16 bytes, got {length}"
)));
}
let Some(bytes) = reader.take::<16>() else {
return Ok(None);
};
ColumnValues::Uuid(uuid::Uuid::from_bytes_le(bytes))
}
}
TdsDataType::DecimalN | TdsDataType::NumericN => {
let Some(length) = reader.byte() else {
return Ok(None);
};
let TypeInfoVariant::VarLenPrecisionScale(_, _, precision, scale) =
metadata.type_info.type_info_variant
else {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid type info variant for Decimal/Numeric: expected VarLenPrecisionScale, got: {:?}",
metadata.type_info.type_info_variant
)));
};
if !decimal_metadata_is_valid(precision, scale) {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid decimal precision {precision} / scale {scale}"
)));
}
if length == 0 {
ColumnValues::Null
} else {
let Some(sign) = reader.byte() else {
return Ok(None);
};
let magnitude_len = usize::from(length - 1);
let number_of_int_parts = magnitude_len.div_ceil(4);
if number_of_int_parts > MAX_DECIMAL_INT_PARTS {
return Err(crate::error::Error::ProtocolError(format!(
"Decimal int parts {number_of_int_parts} exceeds maximum allowed {MAX_DECIMAL_INT_PARTS} (length was {length})"
)));
}
let Some(source) = reader.take_slice(magnitude_len) else {
return Ok(None);
};
let mut magnitude = [0u8; DECIMAL_MAGNITUDE_BYTES];
magnitude[..magnitude_len].copy_from_slice(source);
let value = DecimalParts::new(
sign == 1,
precision,
scale,
u128::from_le_bytes(magnitude),
);
if metadata.data_type == TdsDataType::DecimalN {
ColumnValues::Decimal(value)
} else {
ColumnValues::Numeric(value)
}
}
}
TdsDataType::NChar
| TdsDataType::NVarChar
| TdsDataType::BigChar
| TdsDataType::BigVarChar
| TdsDataType::Char
| TdsDataType::VarChar => {
let Some(length) = reader.u16() else {
return Ok(None);
};
if length == u16::MAX {
ColumnValues::Null
} else {
let Some(bytes) = reader.take_bytes(usize::from(length)) else {
return Ok(None);
};
ColumnValues::String(SqlString::new(bytes, get_encoding_type(metadata)))
}
}
TdsDataType::BigBinary | TdsDataType::BigVarBinary => {
let Some(length) = reader.u16() else {
return Ok(None);
};
if length == u16::MAX {
ColumnValues::Null
} else {
let Some(bytes) = reader.take_bytes(usize::from(length)) else {
return Ok(None);
};
ColumnValues::Bytes(bytes)
}
}
_ => return Ok(None),
};
Ok(Some((value, reader.position)))
}
pub(crate) fn try_decode_buffered_into<W: RowWriter + ?Sized>(
&self,
bytes: &[u8],
metadata: &ColumnMetadata,
col: usize,
writer: &mut W,
) -> TdsResult<Option<usize>> {
if metadata.is_plp() || metadata.crypto_metadata.is_some() {
return Ok(None);
}
macro_rules! read {
($reader:ident.$method:ident()) => {
match $reader.$method() {
Some(value) => value,
None => return Ok(None),
}
};
}
let mut reader = BufferedSlice::new(bytes);
match metadata.data_type {
TdsDataType::Int1 => writer.write_u8(col, read!(reader.byte())),
TdsDataType::Int2 => writer.write_i16(col, read!(reader.i16())),
TdsDataType::Int4 => writer.write_i32(col, read!(reader.i32())),
TdsDataType::Int8 => writer.write_i64(col, read!(reader.i64())),
TdsDataType::Flt4 => writer.write_f32(col, read!(reader.f32())),
TdsDataType::Flt8 => writer.write_f64(col, read!(reader.f64())),
TdsDataType::Bit => writer.write_bool(col, read!(reader.byte()) == 1),
TdsDataType::IntN => match read!(reader.byte()) {
0 => writer.write_null(col),
1 => writer.write_u8(col, read!(reader.byte())),
2 => writer.write_i16(col, read!(reader.i16())),
4 => writer.write_i32(col, read!(reader.i32())),
8 => writer.write_i64(col, read!(reader.i64())),
length => {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid IntN length: {length}"
)));
}
},
TdsDataType::FltN => match read!(reader.byte()) {
0 => writer.write_null(col),
4 => writer.write_f32(col, read!(reader.f32())),
_ => writer.write_f64(col, read!(reader.f64())),
},
TdsDataType::BitN => match read!(reader.byte()) {
0 => writer.write_null(col),
_ => writer.write_bool(col, read!(reader.byte()) == 1),
},
TdsDataType::Money4 => writer.write_smallmoney(
col,
SqlSmallMoney {
int_val: read!(reader.i32()),
},
),
TdsDataType::Money => {
let msb_part = read!(reader.i32());
let lsb_part = read!(reader.i32());
writer.write_money(col, SqlMoney { lsb_part, msb_part });
}
TdsDataType::MoneyN => match read!(reader.byte()) {
0 => writer.write_null(col),
4 => writer.write_smallmoney(
col,
SqlSmallMoney {
int_val: read!(reader.i32()),
},
),
8 => {
let msb_part = read!(reader.i32());
let lsb_part = read!(reader.i32());
writer.write_money(col, SqlMoney { lsb_part, msb_part });
}
length => {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid MoneyN length: {length}"
)));
}
},
TdsDataType::DateTim4 => {
let days = read!(reader.u16());
let time = read!(reader.u16());
writer.write_smalldatetime(col, SqlSmallDateTime { days, time });
}
TdsDataType::DateTime => {
let days = read!(reader.i32());
let time = read!(reader.u32());
writer.write_datetime(col, SqlDateTime { days, time });
}
TdsDataType::DateTimeN => match read!(reader.byte()) {
0 => writer.write_null(col),
4 => {
let days = read!(reader.u16());
let time = read!(reader.u16());
writer.write_smalldatetime(col, SqlSmallDateTime { days, time });
}
_ => {
let days = read!(reader.i32());
let time = read!(reader.u32());
writer.write_datetime(col, SqlDateTime { days, time });
}
},
TdsDataType::DateN => match read!(reader.byte()) {
0 => writer.write_null(col),
_ => writer.write_date(col, SqlDate::unchecked_create(read!(reader.u24()))),
},
TdsDataType::TimeN | TdsDataType::DateTime2N | TdsDataType::DateTimeOffsetN => {
let length = read!(reader.byte());
if length == 0 {
writer.write_null(col);
} else {
let scale = metadata.get_scale().ok_or_else(|| {
crate::error::Error::ImplementationError(format!(
"{:?} type should have scale",
metadata.data_type
))
})?;
let trailing = match metadata.data_type {
TdsDataType::TimeN => 0,
TdsDataType::DateTime2N => 3,
TdsDataType::DateTimeOffsetN => 5,
_ => {
return Err(crate::error::Error::ImplementationError(
"unexpected buffered time type".to_string(),
));
}
};
let time_length = length.checked_sub(trailing).ok_or_else(|| {
crate::error::Error::ProtocolError(format!(
"Invalid {:?} length: {length}",
metadata.data_type
))
})?;
let scaled = match time_length {
3 => read!(reader.u24()).into(),
4 => read!(reader.u32()).into(),
_ => read!(reader.u40()),
};
let time = SqlTime {
time_nanoseconds: scale_time_value(scaled, scale),
scale,
};
match metadata.data_type {
TdsDataType::TimeN => writer.write_time(col, time),
TdsDataType::DateTime2N => writer.write_datetime2(
col,
SqlDateTime2 {
days: read!(reader.u24()),
time,
},
),
TdsDataType::DateTimeOffsetN => writer.write_datetimeoffset(
col,
SqlDateTimeOffset {
datetime2: SqlDateTime2 {
days: read!(reader.u24()),
time,
},
offset: read!(reader.i16()),
},
),
_ => {
return Err(crate::error::Error::ImplementationError(
"unexpected buffered time type".to_string(),
));
}
}
}
}
TdsDataType::Guid => {
let length = read!(reader.byte());
if length == 0 {
writer.write_null(col);
} else {
if length != 16 {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid GUID length: expected 16 bytes, got {length}"
)));
}
let Some(bytes) = reader.take::<16>() else {
return Ok(None);
};
writer.write_uuid(col, uuid::Uuid::from_bytes_le(bytes));
}
}
TdsDataType::DecimalN | TdsDataType::NumericN => {
let length = read!(reader.byte());
let TypeInfoVariant::VarLenPrecisionScale(_, _, precision, scale) =
metadata.type_info.type_info_variant
else {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid type info variant for Decimal/Numeric: expected VarLenPrecisionScale, got: {:?}",
metadata.type_info.type_info_variant
)));
};
if !decimal_metadata_is_valid(precision, scale) {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid decimal precision {precision} / scale {scale}"
)));
}
if length == 0 {
writer.write_null(col);
} else {
let sign = read!(reader.byte());
let magnitude_len = usize::from(length - 1);
let number_of_int_parts = magnitude_len.div_ceil(4);
if number_of_int_parts > MAX_DECIMAL_INT_PARTS {
return Err(crate::error::Error::ProtocolError(format!(
"Decimal int parts {number_of_int_parts} exceeds maximum allowed {MAX_DECIMAL_INT_PARTS} (length was {length})"
)));
}
let Some(source) = reader.take_slice(magnitude_len) else {
return Ok(None);
};
let mut magnitude = [0u8; DECIMAL_MAGNITUDE_BYTES];
magnitude[..magnitude_len].copy_from_slice(source);
let value = DecimalParts::new(
sign == 1,
precision,
scale,
u128::from_le_bytes(magnitude),
);
if metadata.data_type == TdsDataType::DecimalN {
writer.write_decimal(col, value);
} else {
writer.write_numeric(col, value);
}
}
}
TdsDataType::NChar
| TdsDataType::NVarChar
| TdsDataType::BigChar
| TdsDataType::BigVarChar
| TdsDataType::Char
| TdsDataType::VarChar => {
let length = read!(reader.u16());
if length == u16::MAX {
writer.write_null(col);
} else {
let Some(bytes) = reader.take_slice(usize::from(length)) else {
return Ok(None);
};
writer.write_string(col, Cow::Borrowed(bytes), get_encoding_type(metadata));
}
}
TdsDataType::BigBinary | TdsDataType::BigVarBinary => {
let length = read!(reader.u16());
if length == u16::MAX {
writer.write_null(col);
} else {
let Some(bytes) = reader.take_slice(usize::from(length)) else {
return Ok(None);
};
writer.write_bytes(col, Cow::Borrowed(bytes));
}
}
TdsDataType::SsVariant => {
let Some((base, value, used)) = self.try_decode_buffered_variant(bytes)? else {
return Ok(None);
};
if let Some(base) = base {
writer.write_variant_base_type(col, base);
}
write_column_value(writer, col, value);
return Ok(Some(used));
}
_ => {
let Some((value, used)) = self.try_decode_buffered(bytes, metadata)? else {
return Ok(None);
};
write_column_value(writer, col, value);
return Ok(Some(used));
}
}
Ok(Some(reader.position))
}
#[cfg(test)]
const SHORTLEN_MAXVALUE: usize = 65535;
const SQL_PLP_NULL: usize = 0xffffffffffffffff;
const SQL_PLP_UNKNOWNLEN: usize = 0xfffffffffffffffe;
#[cfg_attr(not(test), allow(dead_code))]
const SQL_PLP_MAXLEN: usize = 0xfffffffffffffffd;
#[cfg(fuzzing)]
const MAX_PLP_CHUNK_SIZE: usize = 8 * 1024;
#[cfg(not(fuzzing))]
const MAX_PLP_CHUNK_SIZE: usize = 16 * 1024 * 1024;
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))
}
async fn read_sql_variant<T>(&self, reader: &mut T) -> TdsResult<ColumnValues>
where
T: TdsPacketReader + Send + Sync,
{
self.read_sql_variant_with_base(reader)
.await
.map(|(_, value)| value)
}
async fn read_sql_variant_with_base<T>(
&self,
reader: &mut T,
) -> TdsResult<(Option<TdsDataType>, ColumnValues)>
where
T: TdsPacketReader + Send + Sync,
{
let length = read_sync_first!(reader, try_read_uint32, read_uint32);
if length == 0 {
return Ok((None, ColumnValues::Null));
}
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 = read_sync_first!(reader, try_read_byte, read_byte);
let bytes_for_type_and_properties_byte = 2;
let data_length = length
.checked_sub(bytes_for_type_and_properties_byte)
.and_then(|v| v.checked_sub(variant_prop_bytes as u32))
.ok_or_else(|| {
crate::error::Error::ProtocolError(format!(
"SQL_VARIANT data length calculation underflow: length={length}, prop_bytes={variant_prop_bytes}"
))
})?;
let col_value = match variant_prop_bytes {
0 => {
self.decode_zero_propbyte_variant(reader, tds_type, data_length)
.await?
}
1 => {
self.decode_one_byte_variant(reader, tds_type, data_length)
.await?
}
2 => {
decode_two_propbyte_variant(reader, variant_base_type, tds_type, data_length)
.await?
}
7 => {
decode_seven_propbyte_variant(reader, tds_type, data_length).await?
}
_ => {
return Err(crate::error::Error::ProtocolError(format!(
"Unexpected SQL variant properties length: {variant_prop_bytes}. Expected 0, 1, 2, or 7. This indicates malformed or invalid data."
)));
}
};
Ok((Some(tds_type), col_value))
}
async fn decode_zero_propbyte_variant<T>(
&self,
reader: &mut T,
tds_type: TdsDataType,
data_length: u32,
) -> Result<ColumnValues, crate::error::Error>
where
T: TdsPacketReader + Send + Sync,
{
let fixed_length_type_result = FixedLengthTypes::try_from(tds_type);
match fixed_length_type_result {
Ok(fixed_length_type) => {
let type_info = TypeInfo {
tds_type,
length: data_length as usize,
type_info_variant: TypeInfoVariant::FixedLen(fixed_length_type),
};
let variant_actual_type_md = ColumnMetadata {
user_type: 0,
flags: 0,
type_info,
data_type: tds_type,
column_name: "".to_string(),
multi_part_name: None,
crypto_metadata: None,
};
self.decode_boxed(reader, &variant_actual_type_md).await
}
_ => {
match tds_type {
TdsDataType::Guid => Self::read_guid(reader, data_length as u8).await,
TdsDataType::DateN => Self::read_daten(reader, data_length as u8).await,
_ => Err(crate::error::Error::ProtocolError(format!(
"For 0 byte property, only Guid and DateN are expected, but got: {tds_type:?}"
))),
}
}
}
}
async fn decode_one_byte_variant<T>(
&self,
reader: &mut T,
tds_type: TdsDataType,
data_length: u32,
) -> TdsResult<ColumnValues>
where
T: TdsPacketReader + Send + Sync,
{
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?;
ColumnValues::Time(time_nanos)
}
TdsDataType::DateTime2N => {
self.read_datetime2(reader, data_length as u8, scale)
.await?
}
TdsDataType::DateTimeOffsetN => {
self.read_datetime_offset(reader, data_length as u8, scale)
.await?
}
_ => {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid SQL_VARIANT: 1-byte property is only valid for TimeN, DateTime2N, and DateTimeOffsetN types, but got: {tds_type:?}"
)));
}
})
}
async fn read_decimal<T>(
&self,
reader: &mut T,
metadata: &ColumnMetadata,
) -> TdsResult<Option<DecimalParts>>
where
T: TdsPacketReader + Send + Sync,
{
let length = read_sync_first!(reader, try_read_byte, read_byte);
let TypeInfoVariant::VarLenPrecisionScale(_, _, precision, scale) =
metadata.type_info.type_info_variant
else {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid type info variant for Decimal/Numeric: expected VarLenPrecisionScale, got: {:?}",
metadata.type_info.type_info_variant
)));
};
GenericDecoder::read_decimal_data(reader, length, precision, scale).await
}
async fn read_decimal_data<T>(
reader: &mut T,
length: u8,
precision: u8,
scale: u8,
) -> TdsResult<Option<DecimalParts>>
where
T: TdsPacketReader + Send + Sync,
{
if !decimal_metadata_is_valid(precision, scale) {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid decimal precision {precision} / scale {scale}"
)));
}
if length == 0 {
return Ok(None);
}
let sign = read_sync_first!(reader, try_read_byte, read_byte);
let is_positive = sign == 1;
let magnitude_len = (length - 1) as usize;
let number_of_int_parts = magnitude_len.div_ceil(4);
if number_of_int_parts > MAX_DECIMAL_INT_PARTS {
return Err(crate::error::Error::ProtocolError(format!(
"Decimal int parts {number_of_int_parts} exceeds maximum allowed {MAX_DECIMAL_INT_PARTS} (length was {length})"
)));
}
let mut magnitude = [0u8; DECIMAL_MAGNITUDE_BYTES];
reader.read_bytes(&mut magnitude[..magnitude_len]).await?;
Ok(Some(DecimalParts::new(
is_positive,
precision,
scale,
u128::from_le_bytes(magnitude),
)))
}
async fn read_datetime<T>(&self, reader: &mut T) -> TdsResult<SqlDateTime>
where
T: TdsPacketReader + Send + Sync,
{
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 })
}
async fn read_small_datetime<T>(&self, reader: &mut T) -> TdsResult<SqlSmallDateTime>
where
T: TdsPacketReader + Send + Sync,
{
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,
})
}
async fn read_date<T>(reader: &mut T) -> TdsResult<SqlDate>
where
T: TdsPacketReader + Send + Sync,
{
let days = read_sync_first!(reader, try_read_uint24, read_uint24);
Ok(SqlDate::unchecked_create(days))
}
async fn read_time<T>(&self, reader: &mut T, byte_len: u8, scale: u8) -> TdsResult<SqlTime>
where
T: TdsPacketReader + Send + Sync,
{
let scaled_value = match byte_len {
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),
};
Ok(SqlTime {
time_nanoseconds: scale_time_value(scaled_value, scale),
scale,
})
}
async fn read_datetime2<T>(
&self,
reader: &mut T,
byte_len: u8,
scale: u8,
) -> TdsResult<ColumnValues>
where
T: TdsPacketReader + Send + Sync,
{
let time_byte_len = byte_len.checked_sub(3).ok_or_else(|| {
crate::error::Error::ProtocolError(format!(
"Invalid DateTime2 byte length: {byte_len}. Expected at least 3 bytes for date component."
))
})?;
let time_nanos = self.read_time(reader, time_byte_len, scale).await?;
let sql_date = Self::read_date(reader).await?;
let datetime2 = SqlDateTime2 {
days: sql_date.get_days(),
time: time_nanos,
};
Ok(ColumnValues::DateTime2(datetime2))
}
async fn read_datetime_offset<T>(
&self,
reader: &mut T,
byte_len: u8,
scale: u8,
) -> TdsResult<ColumnValues>
where
T: TdsPacketReader + Send + Sync,
{
let datetime2_byte_len = byte_len.checked_sub(2).ok_or_else(|| {
crate::error::Error::ProtocolError(format!(
"Invalid DateTimeOffset byte length: {byte_len}. Expected at least 2 bytes for offset component."
))
})?;
let datetime2 = self
.read_datetime2(reader, datetime2_byte_len, scale)
.await?;
let datetime2 = match datetime2 {
ColumnValues::DateTime2(dt2) => dt2,
_ => {
return Err(crate::error::Error::ProtocolError(format!(
"Internal error: read_datetime2 returned unexpected type: {datetime2:?}"
)));
}
};
let offset = read_sync_first!(reader, try_read_int16, read_int16);
let datetime_offset = SqlDateTimeOffset { datetime2, offset };
Ok(ColumnValues::DateTimeOffset(datetime_offset))
}
async fn read_intn<T>(&self, reader: &mut T, byte_len: u8) -> TdsResult<ColumnValues>
where
T: TdsPacketReader + Send + Sync,
{
let value: ColumnValues = match byte_len {
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(
std::io::ErrorKind::InvalidData,
"Invalid IntN length",
)));
}
};
Ok(value)
}
async fn read_money4<T>(&self, reader: &mut T) -> TdsResult<SqlSmallMoney>
where
T: TdsPacketReader + Send + Sync,
{
let small_money_val = read_sync_first!(reader, try_read_int32, read_int32);
Ok(small_money_val.into())
}
async fn read_money8<T>(&self, reader: &mut T) -> TdsResult<SqlMoney>
where
T: TdsPacketReader + Send + Sync,
{
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,
})
}
async fn read_daten<T>(reader: &mut T, length: u8) -> TdsResult<ColumnValues>
where
T: TdsPacketReader + Send + Sync,
{
if length == 0 {
Ok(ColumnValues::Null)
} else {
Ok(ColumnValues::Date(Self::read_date(reader).await?))
}
}
async fn read_guid<T>(reader: &mut T, length: u8) -> TdsResult<ColumnValues>
where
T: TdsPacketReader + Send + Sync,
{
if length > 0 {
if length != 16 {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid GUID length: expected 16 bytes, got {length}"
)));
}
let mut bytes = safe_vec![0u8; length as usize, "read_guid"];
reader.read_bytes(&mut bytes).await?;
let unique_id = uuid::Uuid::from_slice_le(&bytes).map_err(|e| {
crate::error::Error::ProtocolError(format!("Failed to parse UUID: {e}"))
})?;
Ok(ColumnValues::Uuid(unique_id))
} else {
Ok(ColumnValues::Null)
}
}
async fn decode_vector<T>(
&self,
reader: &mut T,
metadata: &ColumnMetadata,
) -> TdsResult<ColumnValues>
where
T: TdsPacketReader + Send + Sync,
{
use crate::datatypes::sql_vector::SqlVector;
use crate::datatypes::sqldatatypes::{
VECTOR_HEADER_SIZE, VECTOR_MAX_SIZE, VectorBaseType, VectorLayoutFormat,
VectorLayoutVersion,
};
let length_prefix_value = read_sync_first!(reader, try_read_uint16, read_uint16) as usize;
if length_prefix_value == 0xFFFF {
return Ok(ColumnValues::Null);
}
if length_prefix_value > VECTOR_MAX_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"Vector length {} exceeds maximum of {} bytes",
length_prefix_value, VECTOR_MAX_SIZE
)));
}
if length_prefix_value < VECTOR_HEADER_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"Vector length {} is less than minimum header size of {} bytes",
length_prefix_value, VECTOR_HEADER_SIZE
)));
}
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); let _reserved2 = read_sync_first!(reader, try_read_byte, read_byte); let _reserved3 = read_sync_first!(reader, try_read_byte, read_byte);
let _layout_format = VectorLayoutFormat::try_from(layout_format_byte)?;
let _layout_version = VectorLayoutVersion::try_from(layout_version_byte)?;
let base_type_in_metadata = match &metadata.type_info.type_info_variant {
TypeInfoVariant::VarLenScale(_, scale) => *scale,
_ => {
return Err(crate::error::Error::ProtocolError(
"Vector metadata missing scale (base type)".to_string(),
));
}
};
if base_type_byte != base_type_in_metadata {
return Err(crate::error::Error::ProtocolError(format!(
"Vector base type mismatch: metadata has 0x{:02X}, vector header has 0x{:02X}",
base_type_in_metadata, base_type_byte
)));
}
let base_type = VectorBaseType::try_from(base_type_byte)?;
let length_in_metadata = metadata.type_info.length;
let element_size = base_type.element_size_bytes();
let length_from_vector_header =
VECTOR_HEADER_SIZE + (dimension_count as usize * element_size);
if length_prefix_value != length_from_vector_header
|| length_prefix_value != length_in_metadata
{
return Err(crate::error::Error::ProtocolError(format!(
"Vector length mismatch: length in prefix {} bytes, length from vector header {} bytes, length in metadata {} bytes, for {} dimensions (element size: {} bytes)",
length_prefix_value,
length_from_vector_header,
length_in_metadata,
dimension_count,
element_size
)));
}
let element_bytes = length_prefix_value - VECTOR_HEADER_SIZE;
let mut raw_bytes = vec![0u8; element_bytes];
reader.read_bytes(&mut raw_bytes).await?;
let vector = SqlVector::try_from_raw(
layout_format_byte,
layout_version_byte,
base_type_byte,
raw_bytes,
)?;
Ok(ColumnValues::Vector(vector))
}
async fn read_plp_bytes<T>(reader: &mut T) -> TdsResult<Option<Vec<u8>>>
where
T: TdsPacketReader + Send + Sync,
{
match Self::read_plp_framing(reader).await? {
PlpFraming::Null => Ok(None),
PlpFraming::Known(length) => {
let mut plp_buffer = vec![0u8; length];
let dest = unsafe {
std::slice::from_raw_parts_mut(
plp_buffer.as_mut_ptr().cast::<MaybeUninit<u8>>(),
plp_buffer.len(),
)
};
Self::read_plp_chunks_into_slice(reader, dest, length).await?;
Ok(Some(plp_buffer))
}
PlpFraming::Unknown => Ok(Some(Self::read_plp_chunks_unknown_len(reader).await?)),
}
}
async fn read_plp_chunks_unknown_len<T>(reader: &mut T) -> TdsResult<Vec<u8>>
where
T: TdsPacketReader + Send + Sync,
{
let mut plp_buffer = Vec::new();
let mut total_len: usize = 0;
let mut chunk_len = read_sync_first!(reader, try_read_uint32, read_uint32) as usize;
#[cfg(fuzzing)]
let mut chunk_count = 0u32;
while chunk_len > 0 {
#[cfg(fuzzing)]
{
chunk_count += 1;
eprintln!(
"[ALLOC] read_plp_chunks_unknown_len: chunk #{chunk_count}, chunk_len={chunk_len}, total_len={total_len}"
);
}
Self::check_plp_chunk(chunk_len)?;
total_len = total_len.checked_add(chunk_len).ok_or_else(|| {
crate::error::Error::ProtocolError(format!(
"PLP chunk accumulation would overflow capacity: {total_len} + {chunk_len}"
))
})?;
if total_len > MAX_PLP_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"PLP accumulated size {total_len} exceeds maximum allowed size of {MAX_PLP_SIZE} bytes (SQL Server limit: 2GB)"
)));
}
plp_buffer.reserve(chunk_len);
let read = reader
.read_bytes_uninit(&mut plp_buffer.spare_capacity_mut()[..chunk_len])
.await?;
if read != chunk_len {
return Err(crate::error::Error::ProtocolError(format!(
"PLP chunk read returned {read} byte(s) but the chunk header declared {chunk_len}"
)));
}
let new_len = plp_buffer.len() + chunk_len;
unsafe { plp_buffer.set_len(new_len) };
chunk_len = read_sync_first!(reader, try_read_uint32, read_uint32) as usize;
}
Ok(plp_buffer)
}
fn check_plp_chunk(chunk_len: usize) -> TdsResult<()> {
if chunk_len > Self::MAX_PLP_CHUNK_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"PLP chunk size {chunk_len} exceeds maximum allowed chunk size of {} bytes",
Self::MAX_PLP_CHUNK_SIZE
)));
}
Ok(())
}
async fn read_plp_framing<T>(reader: &mut T) -> TdsResult<PlpFraming>
where
T: TdsPacketReader + Send + Sync,
{
let long_len_i64 = read_sync_first!(reader, try_read_int64, read_int64);
let long_len = long_len_i64 as u64;
if long_len as usize == Self::SQL_PLP_NULL {
return Ok(PlpFraming::Null);
}
if long_len as usize == Self::SQL_PLP_UNKNOWNLEN {
return Ok(PlpFraming::Unknown);
}
let capacity = long_len as usize;
if long_len_i64 < 0 || capacity > MAX_PLP_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"PLP length {capacity} (raw i64: {long_len_i64}) exceeds maximum allowed size of {MAX_PLP_SIZE} bytes"
)));
}
Ok(PlpFraming::Known(capacity))
}
async fn read_plp_chunks_into_slice<T>(
reader: &mut T,
dest: &mut [MaybeUninit<u8>],
declared_len: usize,
) -> TdsResult<()>
where
T: TdsPacketReader + Send + Sync,
{
let mut chunk_len = read_sync_first!(reader, try_read_uint32, read_uint32) as usize;
let mut offset: usize = 0;
#[cfg(fuzzing)]
let mut chunk_count = 0u32;
while chunk_len > 0 {
#[cfg(fuzzing)]
{
chunk_count += 1;
eprintln!(
"[ALLOC] read_plp_chunks_into_slice: chunk #{chunk_count}, chunk_len={chunk_len}, wire_declared_len={declared_len}, destination_len={}",
dest.len()
);
}
Self::check_plp_chunk(chunk_len)?;
let end_offset = offset.checked_add(chunk_len).ok_or_else(|| {
crate::error::Error::ProtocolError(format!(
"PLP chunk offset would overflow: {offset} + {chunk_len}"
))
})?;
if end_offset > declared_len || end_offset > dest.len() {
return Err(crate::error::Error::ProtocolError(format!(
"PLP chunk exceeds wire-declared length: offset={offset}, chunk_len={chunk_len}, wire_declared_len={declared_len}, destination_len={}",
dest.len(),
)));
}
offset += reader
.read_bytes_uninit(&mut dest[offset..end_offset])
.await?;
chunk_len = read_sync_first!(reader, try_read_uint32, read_uint32) as usize;
}
dest[offset..].fill(MaybeUninit::new(0));
Ok(())
}
pub(crate) async fn decode_into<T, W>(
&self,
reader: &mut T,
metadata: &ColumnMetadata,
col: usize,
writer: &mut W,
) -> TdsResult<()>
where
T: TdsPacketReader + Send + Sync,
W: RowWriter + ?Sized,
{
match metadata.data_type {
TdsDataType::Int1 => {
writer.write_u8(col, read_sync_first!(reader, try_read_byte, read_byte));
}
TdsDataType::Int2 => {
writer.write_i16(col, read_sync_first!(reader, try_read_int16, read_int16));
}
TdsDataType::Int4 => {
writer.write_i32(col, read_sync_first!(reader, try_read_int32, read_int32));
}
TdsDataType::Int8 => {
writer.write_i64(col, read_sync_first!(reader, try_read_int64, read_int64));
}
TdsDataType::IntN => {
let byte_len = read_sync_first!(reader, try_read_byte, read_byte);
match byte_len {
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(
std::io::ErrorKind::InvalidData,
"Invalid IntN length",
)));
}
}
}
TdsDataType::Flt4 => {
writer.write_f32(
col,
read_sync_first!(reader, try_read_float32, read_float32),
);
}
TdsDataType::Flt8 => {
writer.write_f64(
col,
read_sync_first!(reader, try_read_float64, read_float64),
);
}
TdsDataType::FltN => {
let length = read_sync_first!(reader, try_read_byte, read_byte);
match length {
0 => writer.write_null(col),
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),
),
}
}
TdsDataType::Bit => {
writer.write_bool(col, read_sync_first!(reader, try_read_byte, read_byte) == 1);
}
TdsDataType::BitN => {
let byte_len = read_sync_first!(reader, try_read_byte, read_byte);
if byte_len > 0 {
writer.write_bool(col, read_sync_first!(reader, try_read_byte, read_byte) == 1);
} else {
writer.write_null(col);
}
}
TdsDataType::Money4 => {
writer.write_smallmoney(col, self.read_money4(reader).await?);
}
TdsDataType::Money => {
writer.write_money(col, self.read_money8(reader).await?);
}
TdsDataType::MoneyN => {
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?),
0 => writer.write_null(col),
_ => {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid MoneyN length - {byte_len}"
)));
}
}
}
TdsDataType::DecimalN => match self.read_decimal(reader, metadata).await? {
Some(val) => writer.write_decimal(col, val),
None => writer.write_null(col),
},
TdsDataType::NumericN => match self.read_decimal(reader, metadata).await? {
Some(val) => writer.write_numeric(col, val),
None => writer.write_null(col),
},
TdsDataType::NChar
| TdsDataType::NVarChar
| TdsDataType::BigChar
| TdsDataType::BigVarChar
| TdsDataType::Char
| TdsDataType::VarChar
| TdsDataType::NText
| TdsDataType::Text => {
self.string_decoder
.decode_string_into(reader, metadata, col, writer)
.await?;
}
TdsDataType::BigBinary => {
let length = read_sync_first!(reader, try_read_uint16, read_uint16);
if length == 0xFFFF {
writer.write_null(col);
} else {
if length as usize > MAX_ALLOC_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"BigBinary length {length} exceeds maximum allowed size of {MAX_ALLOC_SIZE} bytes"
)));
}
if let Some(bytes) = reader.try_read_slice(length as usize) {
writer.write_bytes(col, Cow::Borrowed(bytes));
} else {
let mut bytes = vec![0u8; length as usize];
reader.read_bytes(&mut bytes).await?;
writer.write_bytes(col, Cow::Owned(bytes));
}
}
}
TdsDataType::BigVarBinary => {
if metadata.is_plp() {
match GenericDecoder::read_plp_framing(reader).await? {
PlpFraming::Null => writer.write_null(col),
PlpFraming::Known(length) => {
read_value_into!(
writer,
col,
ValueKind::Bytes,
length,
|dest| GenericDecoder::read_plp_chunks_into_slice(
reader, dest, length,
)
.await,
|bytes| writer.write_bytes(col, Cow::Owned(bytes)),
);
}
PlpFraming::Unknown => {
writer.write_bytes(
col,
Cow::Owned(
GenericDecoder::read_plp_chunks_unknown_len(reader).await?,
),
);
}
}
} else {
let length = read_sync_first!(reader, try_read_uint16, read_uint16);
if length == 0xFFFF {
writer.write_null(col);
} else {
if length as usize > MAX_ALLOC_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"BigVarBinary length {length} exceeds maximum allowed size of {MAX_ALLOC_SIZE} bytes"
)));
}
if let Some(bytes) = reader.try_read_slice(length as usize) {
writer.write_bytes(col, Cow::Borrowed(bytes));
} else {
let mut bytes = vec![0u8; length as usize];
reader.read_bytes(&mut bytes).await?;
writer.write_bytes(col, Cow::Owned(bytes));
}
}
}
}
TdsDataType::DateTime => {
writer.write_datetime(col, self.read_datetime(reader).await?);
}
TdsDataType::DateTim4 => {
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 {
days: daypart,
time: timepart,
},
);
}
TdsDataType::DateTimeN => {
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?),
_ => writer.write_datetime(col, self.read_datetime(reader).await?),
}
}
TdsDataType::DateN => {
let length = read_sync_first!(reader, try_read_byte, read_byte);
if length == 0 {
writer.write_null(col);
} else {
writer.write_date(col, Self::read_date(reader).await?);
}
}
TdsDataType::TimeN => {
let length = read_sync_first!(reader, try_read_byte, read_byte);
if length == 0 {
writer.write_null(col);
} else {
writer.write_time(
col,
self.read_time(
reader,
length,
metadata.get_scale().ok_or_else(|| {
crate::error::Error::ImplementationError(
"TimeN type should have scale".to_string(),
)
})?,
)
.await?,
);
}
}
TdsDataType::DateTime2N => {
let length = read_sync_first!(reader, try_read_byte, read_byte);
if length == 0 {
writer.write_null(col);
} else {
let cv = self
.read_datetime2(
reader,
length,
metadata.get_scale().ok_or_else(|| {
crate::error::Error::ImplementationError(
"DateTime2N type should have scale".to_string(),
)
})?,
)
.await?;
if let ColumnValues::DateTime2(dt2) = cv {
writer.write_datetime2(col, dt2);
}
}
}
TdsDataType::DateTimeOffsetN => {
let length = read_sync_first!(reader, try_read_byte, read_byte);
if length == 0 {
writer.write_null(col);
} else {
let cv = self
.read_datetime_offset(
reader,
length,
metadata.get_scale().ok_or_else(|| {
crate::error::Error::ImplementationError(
"DateTimeOffsetN type should have scale".to_string(),
)
})?,
)
.await?;
if let ColumnValues::DateTimeOffset(dto) = cv {
writer.write_datetimeoffset(col, dto);
}
}
}
TdsDataType::Guid => {
let length = read_sync_first!(reader, try_read_byte, read_byte);
if length == 0 {
writer.write_null(col);
} else {
if length != 16 {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid GUID length: expected 16 bytes, got {length}"
)));
}
let mut bytes = [0u8; 16];
reader.read_bytes(&mut bytes).await?;
let uuid = uuid::Uuid::from_slice_le(&bytes).map_err(|e| {
crate::error::Error::ProtocolError(format!("Failed to parse UUID: {e}"))
})?;
writer.write_uuid(col, uuid);
}
}
TdsDataType::SsVariant => {
let (base, value) = self.read_sql_variant_with_base(reader).await?;
if let Some(base) = base {
writer.write_variant_base_type(col, base);
}
write_column_value(writer, col, value);
}
_ => {
let value = self.decode(reader, metadata).await?;
write_column_value(writer, col, value);
}
}
Ok(())
}
}
enum PlpFraming {
Null,
Known(usize),
Unknown,
}
impl SqlTypeDecode for GenericDecoder {
async fn decode<T>(&self, reader: &mut T, metadata: &ColumnMetadata) -> TdsResult<ColumnValues>
where
T: TdsPacketReader + Send + Sync,
{
let result = match metadata.data_type {
TdsDataType::Int1 => {
let value = read_sync_first!(reader, try_read_byte, read_byte);
ColumnValues::from(value)
}
TdsDataType::Int2 => {
let value = read_sync_first!(reader, try_read_int16, read_int16);
ColumnValues::SmallInt(value)
}
TdsDataType::Int4 => {
let value = read_sync_first!(reader, try_read_int32, read_int32);
ColumnValues::from(value)
}
TdsDataType::Int8 => {
let value = read_sync_first!(reader, try_read_int64, read_int64);
ColumnValues::BigInt(value)
}
TdsDataType::Flt4 => {
let value = read_sync_first!(reader, try_read_float32, read_float32);
ColumnValues::Real(value)
}
TdsDataType::Flt8 => {
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 = 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?),
0 => ColumnValues::Null,
_ => {
return Err(crate::error::Error::ProtocolError(format!(
"Invalid MoneyN length - {byte_len}"
)));
}
}
}
TdsDataType::DecimalN => {
let value = self.read_decimal(reader, metadata).await?;
match value {
Some(value) => ColumnValues::Decimal(value),
None => ColumnValues::Null,
}
}
TdsDataType::NumericN => {
let value = self.read_decimal(reader, metadata).await?;
match value {
Some(value) => ColumnValues::Numeric(value),
None => ColumnValues::Null,
}
}
TdsDataType::Bit => {
let value = read_sync_first!(reader, try_read_byte, read_byte);
ColumnValues::Bit(value == 1)
}
TdsDataType::NChar
| TdsDataType::NVarChar
| TdsDataType::BigChar
| TdsDataType::BigVarChar
| TdsDataType::Char
| TdsDataType::VarChar
| TdsDataType::NText
| TdsDataType::Text => self.string_decoder.decode(reader, metadata).await?,
TdsDataType::DateTime => {
let value = self.read_datetime(reader).await?;
ColumnValues::DateTime(value)
}
TdsDataType::IntN => {
let byte_len = read_sync_first!(reader, try_read_byte, read_byte);
self.read_intn(reader, byte_len).await?
}
TdsDataType::BigBinary => {
let length = read_sync_first!(reader, try_read_uint16, read_uint16);
if length == 0xFFFF {
ColumnValues::Null
} else {
if length as usize > MAX_ALLOC_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"BigBinary length {length} exceeds maximum allowed size of {MAX_ALLOC_SIZE} bytes"
)));
}
let mut bytes = vec![0u8; length as usize];
reader.read_bytes(&mut bytes).await?;
ColumnValues::Bytes(bytes)
}
}
TdsDataType::BigVarBinary => {
if metadata.is_plp() {
let some_bytes = GenericDecoder::read_plp_bytes(reader).await?;
match some_bytes {
Some(bytes) => ColumnValues::Bytes(bytes),
None => ColumnValues::Null,
}
} else {
let length = read_sync_first!(reader, try_read_uint16, read_uint16);
if length == 0xFFFF {
ColumnValues::Null
} else {
if length as usize > MAX_ALLOC_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"BigVarBinary length {length} exceeds maximum allowed size of {MAX_ALLOC_SIZE} bytes"
)));
}
let mut bytes = vec![0u8; length as usize];
reader.read_bytes(&mut bytes).await?;
ColumnValues::Bytes(bytes)
}
}
}
TdsDataType::Xml => {
if !metadata.is_plp() {
return Err(crate::error::Error::ProtocolError(
"XML column metadata is not partially-length-prefixed".to_string(),
));
}
let some_bytes = GenericDecoder::read_plp_bytes(reader).await?;
match some_bytes {
Some(bytes) => ColumnValues::Xml(SqlXml { bytes }),
None => ColumnValues::Null,
}
}
TdsDataType::Json => {
if !metadata.is_plp() {
return Err(crate::error::Error::ProtocolError(
"JSON column metadata is not partially-length-prefixed".to_string(),
));
}
let some_bytes = GenericDecoder::read_plp_bytes(reader).await?;
match some_bytes {
Some(bytes) => ColumnValues::Json(SqlJson::new(bytes)),
None => ColumnValues::Null,
}
}
TdsDataType::Vector => self.decode_vector(reader, metadata).await?,
TdsDataType::BitN => {
let byte_len = read_sync_first!(reader, try_read_byte, read_byte);
if byte_len > 0 {
let value = read_sync_first!(reader, try_read_byte, read_byte);
ColumnValues::Bit(value == 1)
} else {
ColumnValues::Null
}
}
TdsDataType::Guid => {
let length = read_sync_first!(reader, try_read_byte, read_byte);
Self::read_guid(reader, length).await?
}
TdsDataType::FltN => {
let length = read_sync_first!(reader, try_read_byte, read_byte);
if length == 0 {
return Ok(ColumnValues::Null);
}
if length == 4 {
let value = read_sync_first!(reader, try_read_float32, read_float32);
ColumnValues::Real(value)
} else {
let value = read_sync_first!(reader, try_read_float64, read_float64);
ColumnValues::Float(value)
}
}
TdsDataType::DateTimeN => {
let length = read_sync_first!(reader, try_read_byte, read_byte);
if length == 0 {
return Ok(ColumnValues::Null);
} else if length == 4 {
let smalldatetime = self.read_small_datetime(reader).await?;
return Ok(ColumnValues::SmallDateTime(smalldatetime));
} else {
return Ok(ColumnValues::DateTime(self.read_datetime(reader).await?));
}
}
TdsDataType::DateN => {
let length = read_sync_first!(reader, try_read_byte, read_byte);
return Self::read_daten(reader, length).await;
}
TdsDataType::TimeN => {
let length = read_sync_first!(reader, try_read_byte, read_byte);
match length {
0 => return Ok(ColumnValues::Null),
_ => {
return Ok(ColumnValues::Time(
self.read_time(
reader,
length,
metadata.get_scale().ok_or_else(|| {
crate::error::Error::ImplementationError(
"TimeN type should have scale".to_string(),
)
})?,
)
.await?,
));
}
}
}
TdsDataType::DateTime2N => {
let length = read_sync_first!(reader, try_read_byte, read_byte);
match length {
0 => Ok(ColumnValues::Null),
_ => {
self.read_datetime2(
reader,
length,
metadata.get_scale().ok_or_else(|| {
crate::error::Error::ImplementationError(
"DateTime2N type should have scale".to_string(),
)
})?,
)
.await
}
}
}?,
TdsDataType::DateTimeOffsetN => {
let length = read_sync_first!(reader, try_read_byte, read_byte);
match length {
0 => Ok(ColumnValues::Null),
_ => {
self.read_datetime_offset(
reader,
length,
metadata.get_scale().ok_or_else(|| {
crate::error::Error::ImplementationError(
"DateTimeOffsetN type should have scale".to_string(),
)
})?,
)
.await
}
}
}?,
TdsDataType::Image => {
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?;
read_sync_first!(reader, try_read_uint32, read_uint32) as usize
} else {
0
};
if length == 0 {
ColumnValues::Null
} else {
if length > MAX_ALLOC_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"Image length {length} exceeds maximum allowed size of {MAX_ALLOC_SIZE} bytes"
)));
}
let mut buffer = vec![0u8; length];
reader.read_bytes(&mut buffer).await?;
ColumnValues::Bytes(buffer)
}
}
TdsDataType::Udt => {
if !metadata.is_plp() {
return Err(crate::error::Error::ProtocolError(
"UDT column metadata is not partially-length-prefixed".to_string(),
));
}
let some_bytes = GenericDecoder::read_plp_bytes(reader).await?;
match some_bytes {
Some(bytes) => ColumnValues::Bytes(bytes),
None => ColumnValues::Null,
}
}
TdsDataType::SsVariant => self.read_sql_variant(reader).await?,
TdsDataType::DateTim4 => {
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,
})
}
TdsDataType::Decimal => {
return Err(crate::error::Error::UnimplementedFeature {
feature: "Fixed-length Decimal type".to_string(),
context: format!(
"Data type {:?} (0x{:02X}) is not implemented. Use DecimalN instead.",
metadata.data_type, metadata.data_type as u8
),
});
}
TdsDataType::Numeric => {
return Err(crate::error::Error::UnimplementedFeature {
feature: "Fixed-length Numeric type".to_string(),
context: format!(
"Data type {:?} (0x{:02X}) is not implemented. Use NumericN instead.",
metadata.data_type, metadata.data_type as u8
),
});
}
_ => {
return Err(crate::error::Error::UnimplementedFeature {
feature: format!("Data type {:?}", metadata.data_type),
context: format!(
"Data type {:?} (0x{:02X}) is not yet supported in the decoder",
metadata.data_type, metadata.data_type as u8
),
});
}
};
Ok(result)
}
}
#[derive(Debug, Default)]
struct StringDecoder;
impl StringDecoder {
fn is_long_len_type(data_type: TdsDataType) -> bool {
matches!(data_type, TdsDataType::NText | TdsDataType::Text)
}
async fn decode_string_into<T, W>(
&self,
reader: &mut T,
metadata: &ColumnMetadata,
col: usize,
writer: &mut W,
) -> TdsResult<()>
where
T: TdsPacketReader + Send + Sync,
W: RowWriter + ?Sized,
{
let encoding_type = get_encoding_type(metadata);
if metadata.is_plp() {
match GenericDecoder::read_plp_framing(reader).await? {
PlpFraming::Null => writer.write_null(col),
PlpFraming::Known(length) => {
read_value_into!(
writer,
col,
ValueKind::String(&encoding_type),
length,
|dest| GenericDecoder::read_plp_chunks_into_slice(reader, dest, length)
.await,
|bytes| writer.write_string(col, Cow::Owned(bytes), encoding_type),
);
}
PlpFraming::Unknown => {
let bytes = GenericDecoder::read_plp_chunks_unknown_len(reader).await?;
writer.write_string(col, Cow::Owned(bytes), encoding_type);
}
}
} else if Self::is_long_len_type(metadata.data_type) {
let text_ptr_len = read_sync_first!(reader, try_read_byte, read_byte) as usize;
if text_ptr_len == 0 {
writer.write_null(col);
return Ok(());
}
const TIMESTAMP_BYTE_COUNT: usize = 8;
reader.skip_bytes(text_ptr_len).await?;
reader.skip_bytes(TIMESTAMP_BYTE_COUNT).await?;
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!(
"Text data length {length} exceeds maximum allowed size of {MAX_ALLOC_SIZE} bytes"
)));
}
let bytes = if length == 0 {
Vec::new()
} else {
let mut buffer = vec![0u8; length];
reader.read_bytes(&mut buffer).await?;
buffer
};
writer.write_string(col, Cow::Owned(bytes), encoding_type);
} else {
let length = read_sync_first!(reader, try_read_uint16, read_uint16) as usize;
if length == 0xFFFF {
writer.write_null(col);
} else if let Some(bytes) = reader.try_read_slice(length) {
writer.write_string(col, Cow::Borrowed(bytes), encoding_type);
} else {
let mut buffer = vec![0u8; length];
reader.read_bytes(&mut buffer).await?;
writer.write_string(col, Cow::Owned(buffer), encoding_type);
}
}
Ok(())
}
}
impl SqlTypeDecode for StringDecoder {
async fn decode<T>(&self, reader: &mut T, metadata: &ColumnMetadata) -> TdsResult<ColumnValues>
where
T: TdsPacketReader + Send + Sync,
{
let encoding_type = get_encoding_type(metadata);
if metadata.is_plp() {
let some_bytes = GenericDecoder::read_plp_bytes(reader).await?;
match some_bytes {
Some(bytes) => Ok(ColumnValues::String(SqlString::new(bytes, encoding_type))),
None => Ok(ColumnValues::Null),
}
} else if Self::is_long_len_type(metadata.data_type) {
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?;
read_sync_first!(reader, try_read_uint32, read_uint32) as usize
} else {
return Ok(ColumnValues::Null);
};
if length > MAX_ALLOC_SIZE {
return Err(crate::error::Error::ProtocolError(format!(
"Text data length {length} exceeds maximum allowed size of {MAX_ALLOC_SIZE} bytes"
)));
}
let sql_string = if length == 0 {
SqlString::new(Vec::new(), encoding_type)
} else {
let mut buffer = vec![0u8; length];
reader.read_bytes(&mut buffer).await?;
SqlString::new(buffer, encoding_type)
};
Ok(ColumnValues::String(sql_string))
} else {
let length = read_sync_first!(reader, try_read_uint16, read_uint16) as usize;
if length == 0xFFFF {
Ok(ColumnValues::Null)
} else {
let mut buffer = vec![0u8; length];
reader.read_bytes(&mut buffer).await?;
let sql_string = SqlString::new(buffer, encoding_type);
Ok(ColumnValues::String(sql_string))
}
}
}
}
pub const DECIMAL_STR_LEN: usize = 48;
const _: () = assert!(
DECIMAL_STR_LEN >= 42,
"buffer must hold 39 digits of u128::MAX plus a decimal point and sign"
);
#[derive(Clone, Copy, PartialEq, Eq)]
pub struct DecimalParts {
pub is_positive: bool,
pub scale: u8,
pub precision: u8,
magnitude: [u64; 2],
}
impl fmt::Display for DecimalParts {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.format_into(&mut [0u8; DECIMAL_STR_LEN]))
}
}
impl DecimalParts {
fn clamped_scale(&self) -> usize {
(self.scale as usize).min(DECIMAL_STR_LEN - 3)
}
pub const fn new(is_positive: bool, precision: u8, scale: u8, magnitude: u128) -> Self {
DecimalParts {
is_positive,
scale,
precision,
magnitude: [magnitude as u64, (magnitude >> 64) as u64],
}
}
pub const fn magnitude(&self) -> u128 {
(self.magnitude[0] as u128) | ((self.magnitude[1] as u128) << 64)
}
pub fn from_string(s: &str, precision: u8, scale: u8) -> TdsResult<Self> {
use bigdecimal::num_bigint::{BigInt, Sign};
use bigdecimal::{BigDecimal, Zero};
use std::str::FromStr;
let trimmed = s.trim();
let has_negative_sign = trimmed.starts_with('-');
let decimal = BigDecimal::from_str(trimmed).map_err(|e| {
crate::error::Error::TypeConversionError(format!(
"Invalid decimal string '{}': {}",
s, e
))
})?;
let input_scale = decimal.fractional_digit_count();
if input_scale > scale as i64 {
return Err(crate::error::Error::TypeConversionError(format!(
"Input decimal scale {} exceeds target scale {}",
input_scale, scale
)));
}
if decimal.is_zero() {
return Ok(Self::new(!has_negative_sign, precision, scale, 0));
}
let is_positive = decimal.sign() != Sign::Minus;
let abs_decimal = decimal.abs();
let scale_factor = BigDecimal::from(10u64).powi(scale as i64);
let scaled = abs_decimal * scale_factor;
let rounded = scaled.round(0);
let (bigint, exponent) = rounded.into_bigint_and_exponent();
let final_bigint = if exponent > 0 {
bigint * BigInt::from(10u64).pow(exponent as u32)
} else if exponent < 0 {
bigint / BigInt::from(10u64).pow((-exponent) as u32)
} else {
bigint
};
let digits_str = final_bigint.to_string();
if digits_str.len() > precision as usize {
return Err(crate::error::Error::TypeConversionError(format!(
"Decimal value has {} digits, exceeds target precision {}",
digits_str.len(),
precision
)));
}
let bytes = final_bigint.to_signed_bytes_le();
if bytes.len() > DECIMAL_MAGNITUDE_BYTES {
return Err(crate::error::Error::TypeConversionError(format!(
"Decimal value '{s}' does not fit in a 128-bit magnitude"
)));
}
let mut magnitude_bytes = [0u8; DECIMAL_MAGNITUDE_BYTES];
magnitude_bytes[..bytes.len()].copy_from_slice(&bytes);
Ok(Self::new(
is_positive,
precision,
scale,
u128::from_le_bytes(magnitude_bytes),
))
}
pub fn from_i64(value: i64, precision: u8, scale: u8) -> TdsResult<Self> {
let is_positive = value >= 0;
let magnitude = 10u128
.checked_pow(u32::from(scale))
.and_then(|factor| u128::from(value.unsigned_abs()).checked_mul(factor))
.ok_or_else(|| {
crate::error::Error::TypeConversionError(format!(
"Decimal value {value} scaled by 10^{scale} does not fit in a 128-bit magnitude"
))
})?;
Ok(Self::new(is_positive, precision, scale, magnitude))
}
pub fn from_f64(value: f64, precision: u8, scale: u8) -> TdsResult<Self> {
let s = format!("{:.prec$}", value, prec = scale as usize);
Self::from_string(&s, precision, scale)
}
pub fn format_into<'a>(&self, buf: &'a mut [u8; DECIMAL_STR_LEN]) -> &'a str {
const LIMB: u128 = 10_000_000_000_000_000_000;
const LIMB_DIGITS: usize = 19;
let magnitude = self.magnitude();
let mut limbs = [0u64; 3];
let limb_count;
if magnitude <= u64::MAX as u128 {
limbs[0] = magnitude as u64;
limb_count = 1;
} else {
limbs[0] = (magnitude % LIMB) as u64;
let rest = magnitude / LIMB;
if rest <= u64::MAX as u128 {
limbs[1] = rest as u64;
limb_count = 2;
} else {
limbs[1] = (rest % LIMB) as u64;
limbs[2] = (rest / LIMB) as u64;
limb_count = 3;
}
}
let scale = self.clamped_scale();
let mut pos = buf.len();
let mut emitted = 0usize;
let mut limb = 0usize;
let mut taken = 0usize;
let mut current = limbs[0];
loop {
if scale > 0 && emitted == scale {
pos -= 1;
buf[pos] = b'.';
}
let digit = (current % 10) as u8;
current /= 10;
taken += 1;
pos -= 1;
buf[pos] = b'0' + digit;
emitted += 1;
if taken == LIMB_DIGITS && limb + 1 < limb_count {
limb += 1;
current = limbs[limb];
taken = 0;
}
if limb + 1 == limb_count && current == 0 && emitted > scale {
break;
}
}
if !self.is_positive {
pos -= 1;
buf[pos] = b'-';
}
core::str::from_utf8(&buf[pos..]).unwrap_or_default()
}
pub fn to_decimal_string(&self) -> String {
self.format_into(&mut [0u8; DECIMAL_STR_LEN]).to_owned()
}
fn to_f64(self) -> f64 {
let magnitude = self.magnitude() as f64;
let d_ret = magnitude / 10.0_f64.powi(self.clamped_scale() as i32);
if self.is_positive { d_ret } else { -d_ret }
}
pub const fn word(&self, index: usize) -> i32 {
if index >= MAX_DECIMAL_INT_PARTS {
return 0;
}
(self.magnitude() >> (index * 32)) as u32 as i32
}
pub fn word_count(&self) -> usize {
(128 - self.magnitude().leading_zeros() as usize)
.div_ceil(32)
.max(1)
}
pub(crate) fn from_words(is_positive: bool, precision: u8, scale: u8, words: &[i32]) -> Self {
debug_assert!(
words
.iter()
.skip(MAX_DECIMAL_INT_PARTS)
.all(|&word| word == 0),
"magnitude wider than 128 bits truncated: {words:?}"
);
let magnitude = words
.iter()
.take(MAX_DECIMAL_INT_PARTS)
.enumerate()
.fold(0u128, |acc, (i, &word)| {
acc | ((word as u32 as u128) << (i * 32))
});
Self::new(is_positive, precision, scale, magnitude)
}
pub fn try_from_words(
is_positive: bool,
precision: u8,
scale: u8,
words: &[i32],
) -> Option<Self> {
words
.iter()
.skip(MAX_DECIMAL_INT_PARTS)
.all(|&word| word == 0)
.then(|| Self::from_words(is_positive, precision, scale, words))
}
}
impl Debug for DecimalParts {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
write!(
f,
"Decimal: {} Precision {} Scale {} F64 value: {}",
self,
self.precision,
self.scale,
self.to_f64()
)
}
}
async fn decode_two_propbyte_variant<T>(
reader: &mut T,
variant_base_type: u8,
tds_type: TdsDataType,
data_length: u32,
) -> TdsResult<ColumnValues>
where
T: TdsPacketReader + Send + Sync,
{
Ok(match tds_type {
TdsDataType::BigVarBinary | TdsDataType::BigBinary => {
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"
)));
}
let mut buffer = vec![0u8; data_length as usize];
reader.read_bytes(&mut buffer).await?;
ColumnValues::Bytes(buffer)
}
TdsDataType::NumericN | TdsDataType::DecimalN => {
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?;
if matches!(tds_type, TdsDataType::NumericN) {
match decimal_parts {
Some(value) => ColumnValues::Numeric(value),
None => ColumnValues::Null,
}
} else {
match decimal_parts {
Some(value) => ColumnValues::Decimal(value),
None => ColumnValues::Null,
}
}
}
_ => {
return Err(crate::error::Error::ProtocolError(format!(
"Unexpected SQL variant base type for len(2) prop bytes: {variant_base_type:#04X}. Expected binary or numeric types."
)));
}
})
}
async fn decode_seven_propbyte_variant<T>(
reader: &mut T,
tds_type: TdsDataType,
data_length: u32,
) -> TdsResult<ColumnValues>
where
T: TdsPacketReader + Send + Sync,
{
if !matches!(
tds_type,
TdsDataType::BigVarChar | TdsDataType::BigChar | TdsDataType::NVarChar | TdsDataType::NChar
) {
return Err(crate::error::Error::ProtocolError(format!(
"Unexpected SQL variant base type for len(7) prop bytes: {tds_type:?}. Expected a character type."
)));
}
let mut collation_bytes = vec![0u8; 5];
reader.read_bytes(&mut collation_bytes).await?;
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!(
"SQL Variant string data length {data_length} exceeds maximum allowed size of {MAX_ALLOC_SIZE} bytes"
)));
}
let mut buffer = vec![0u8; data_length as usize];
reader.read_bytes(&mut buffer).await?;
let encoding = if matches!(tds_type, TdsDataType::NVarChar | TdsDataType::NChar) {
EncodingType::Utf16
} else if collation.utf8() {
EncodingType::Utf8
} else {
EncodingType::LcidBased(collation)
};
let sql_string = SqlString::new(buffer, encoding);
Ok(ColumnValues::String(sql_string))
}
#[cfg(test)]
mod test {
use crate::datatypes::{
column_values::ColumnValues,
decoder::{
DecimalParts, GenericDecoder, MAX_ALLOC_SIZE, PlpChunkReadLength, PlpChunkStreamReader,
StringDecoder, validate_alloc_size,
},
sqldatatypes::TdsDataType,
};
#[test]
fn test_f64_conversion() {
let expected: f64 = 123456.322;
let magnitude = 12345632200;
let parts = DecimalParts::new(true, 18, 5, magnitude);
assert_eq!(expected, parts.to_f64());
}
#[test]
fn test_f64_conversion_negative() {
let expected: f64 = -123456.322;
let magnitude = 12345632200;
let parts = DecimalParts::new(false, 18, 5, magnitude);
assert_eq!(expected, parts.to_f64());
}
#[test]
fn test_f64_conversion_zero() {
let expected: f64 = 0.0;
let magnitude = 0;
let parts = DecimalParts::new(true, 1, 0, magnitude);
assert_eq!(expected, parts.to_f64());
}
#[test]
fn empty_buffered_plp_consumes_only_its_single_terminator() {
let mut stream = PlpChunkStreamReader::new(PlpChunkReadLength::Known(0));
let mut out = [];
assert_eq!(
stream
.try_read_complete_buffered(&[0; 8], &mut out)
.unwrap(),
Some((4, 0))
);
assert!(stream.reached_end());
}
#[test]
fn buffered_plp_chunks_continue_an_active_stream() {
let mut stream = PlpChunkStreamReader::new(PlpChunkReadLength::Known(10));
let mut first = Vec::from(10_u32.to_le_bytes());
first.extend(0_u8..10);
first.extend(0_u32.to_le_bytes());
let mut out = [0_u8; 4];
assert_eq!(
stream.try_read_buffered(&first, &mut out).unwrap(),
Some((8, 4))
);
assert_eq!(out, [0, 1, 2, 3]);
assert_eq!(stream.total_read(), 4);
assert!(!stream.reached_end());
let remaining = &first[8..];
assert_eq!(
stream.try_read_buffered(remaining, &mut out).unwrap(),
Some((4, 4))
);
assert_eq!(out, [4, 5, 6, 7]);
assert_eq!(
stream.try_read_buffered(&remaining[4..], &mut out).unwrap(),
Some((6, 2))
);
assert_eq!(&out[..2], &[8, 9]);
assert_eq!(stream.total_read(), 10);
assert!(stream.reached_end());
}
#[test]
fn incomplete_buffered_plp_chunk_does_not_advance() {
let mut stream = PlpChunkStreamReader::new(PlpChunkReadLength::Known(4));
let mut bytes = Vec::from(4_u32.to_le_bytes());
bytes.extend([1, 2]);
let mut out = [0_u8; 4];
assert_eq!(stream.try_read_buffered(&bytes, &mut out).unwrap(), None);
assert_eq!(stream.total_read(), 0);
assert!(!stream.reached_end());
bytes.extend([3, 4]);
bytes.extend(0_u32.to_le_bytes());
assert_eq!(
stream.try_read_buffered(&bytes, &mut out).unwrap(),
Some((12, 4))
);
assert_eq!(out, [1, 2, 3, 4]);
assert!(stream.reached_end());
}
#[test]
fn test_f64_conversion_large_number() {
let magnitude = 100000;
let parts = DecimalParts::new(true, 7, 2, magnitude);
let result = parts.to_f64();
assert!((result - 1000.0).abs() < 0.01);
}
#[test]
fn test_decimal_parts_with_multi_word_magnitude() {
let magnitude = 5294967296;
let parts = DecimalParts::new(true, 19, 0, magnitude);
let result = parts.to_f64();
assert!(result > 0.0);
}
#[test]
fn test_u8_to_column_values() {
let value: u8 = 123;
let col_val: ColumnValues = value.into();
match col_val {
ColumnValues::TinyInt(v) => assert_eq!(v, 123),
_ => panic!("Expected TinyInt variant"),
}
}
#[test]
fn test_i32_to_column_values() {
let value: i32 = 12345;
let col_val: ColumnValues = value.into();
match col_val {
ColumnValues::Int(v) => assert_eq!(v, 12345),
_ => panic!("Expected Int variant"),
}
}
#[test]
fn test_i32_negative_to_column_values() {
let value: i32 = -12345;
let col_val: ColumnValues = value.into();
match col_val {
ColumnValues::Int(v) => assert_eq!(v, -12345),
_ => panic!("Expected Int variant"),
}
}
#[test]
fn test_generic_decoder_default() {
let _decoder = GenericDecoder::default();
}
#[test]
fn test_string_decoder_default() {
let _decoder = StringDecoder;
}
#[test]
fn test_decimal_parts_debug() {
let parts = DecimalParts::new(true, 10, 2, 1958505087099);
let debug_str = format!("{parts:?}");
assert!(!debug_str.is_empty());
}
#[test]
fn test_generic_decoder_constants() {
assert_eq!(GenericDecoder::SHORTLEN_MAXVALUE, 65535);
assert_eq!(GenericDecoder::SQL_PLP_NULL, 0xffffffffffffffff);
assert_eq!(GenericDecoder::SQL_PLP_UNKNOWNLEN, 0xfffffffffffffffe);
assert_eq!(GenericDecoder::SQL_PLP_MAXLEN, 0xfffffffffffffffd);
}
#[test]
fn test_decimal_parts_scale_precision() {
let parts = DecimalParts::new(true, 18, 5, 100000);
let result = parts.to_f64();
assert!((result - 1.0).abs() < 0.00001);
}
#[test]
fn test_decimal_parts_zero_magnitude() {
let parts = DecimalParts::new(true, 1, 0, 0);
let result = parts.to_f64();
assert_eq!(result, 0.0);
}
#[test]
fn test_decimal_parts_single_int_part() {
let parts = DecimalParts::new(true, 5, 0, 12345);
let result = parts.to_f64();
assert_eq!(result, 12345.0);
}
#[test]
fn test_column_values_from_u8_zero() {
let value: u8 = 0;
let col_val: ColumnValues = value.into();
match col_val {
ColumnValues::TinyInt(v) => assert_eq!(v, 0),
_ => panic!("Expected TinyInt variant"),
}
}
#[test]
fn test_column_values_from_u8_max() {
let value: u8 = 255;
let col_val: ColumnValues = value.into();
match col_val {
ColumnValues::TinyInt(v) => assert_eq!(v, 255),
_ => panic!("Expected TinyInt variant"),
}
}
#[test]
fn test_column_values_from_i32_zero() {
let value: i32 = 0;
let col_val: ColumnValues = value.into();
match col_val {
ColumnValues::Int(v) => assert_eq!(v, 0),
_ => panic!("Expected Int variant"),
}
}
#[test]
fn test_column_values_from_i32_max() {
let value: i32 = i32::MAX;
let col_val: ColumnValues = value.into();
match col_val {
ColumnValues::Int(v) => assert_eq!(v, i32::MAX),
_ => panic!("Expected Int variant"),
}
}
#[test]
fn test_column_values_from_i32_min() {
let value: i32 = i32::MIN;
let col_val: ColumnValues = value.into();
match col_val {
ColumnValues::Int(v) => assert_eq!(v, i32::MIN),
_ => panic!("Expected Int variant"),
}
}
#[test]
fn test_validate_alloc_size_within_limit() {
let result = validate_alloc_size(1024, "test_allocation");
assert!(result.is_ok());
}
#[test]
fn test_validate_alloc_size_at_limit() {
let result = validate_alloc_size(MAX_ALLOC_SIZE, "test_at_limit");
assert!(result.is_ok());
}
#[test]
fn test_validate_alloc_size_exceeds_limit() {
let result = validate_alloc_size(MAX_ALLOC_SIZE + 1, "test_exceeds");
assert!(result.is_err());
if let Err(e) = result {
let error_msg = format!("{e:?}");
assert!(error_msg.contains("exceeds maximum allowed"));
}
}
#[test]
fn test_validate_alloc_size_zero() {
let result = validate_alloc_size(0, "test_zero");
assert!(result.is_ok());
}
#[test]
fn test_decimal_parts_equality_same() {
let parts1 = DecimalParts::new(true, 10, 2, 858993459300);
let parts2 = DecimalParts::new(true, 10, 2, 858993459300);
assert_eq!(parts1, parts2);
}
#[test]
fn test_decimal_parts_equality_different_sign() {
let parts1 = DecimalParts::new(true, 10, 2, 100);
let parts2 = DecimalParts::new(false, 10, 2, 100);
assert_ne!(parts1, parts2);
}
#[test]
fn test_decimal_parts_equality_different_scale() {
let parts1 = DecimalParts::new(true, 10, 2, 100);
let parts2 = DecimalParts::new(true, 10, 3, 100);
assert_ne!(parts1, parts2);
}
#[test]
fn test_decimal_parts_equality_different_precision() {
let parts1 = DecimalParts::new(true, 10, 2, 100);
let parts2 = DecimalParts::new(true, 12, 2, 100);
assert_ne!(parts1, parts2);
}
#[test]
fn test_decimal_parts_equality_different_length_with_zeros() {
let parts1 = DecimalParts::new(true, 10, 2, 100);
let parts2 = DecimalParts::new(true, 10, 2, 100);
assert_eq!(parts1, parts2);
}
#[test]
fn test_decimal_parts_equality_different_length_with_nonzeros() {
let parts1 = DecimalParts::new(true, 10, 2, 858993459300);
let parts2 = DecimalParts::new(true, 10, 2, 100);
assert_ne!(parts1, parts2);
}
#[test]
fn test_decimal_parts_debug_format_positive() {
let parts = DecimalParts::new(true, 10, 2, 12345);
let debug_str = format!("{parts:?}");
assert!(debug_str.starts_with("Decimal: 123.45 "));
assert!(debug_str.contains("Precision 10"));
assert!(debug_str.contains("Scale 2"));
}
#[test]
fn test_decimal_parts_debug_format_negative() {
let parts = DecimalParts::new(false, 15, 3, 54321);
let debug_str = format!("{parts:?}");
assert!(debug_str.starts_with("Decimal: -54.321 "));
assert!(debug_str.contains("Precision 15"));
assert!(debug_str.contains("Scale 3"));
}
#[test]
fn test_decimal_parts_debug_format_multi_word_magnitude() {
let parts = DecimalParts::new(true, 20, 0, 5534023222971858944100);
let debug_str = format!("{parts:?}");
assert!(debug_str.starts_with("Decimal: 5534023222971858944100 "));
}
#[test]
fn test_f64_conversion_high_scale() {
let magnitude = 12345;
let parts = DecimalParts::new(true, 15, 10, magnitude);
let result = parts.to_f64();
assert!((result - 0.0000012345).abs() < 0.0000000001);
}
#[test]
fn test_f64_conversion_single_zero() {
let parts = DecimalParts::new(true, 10, 5, 0);
let result = parts.to_f64();
assert_eq!(result, 0.0);
}
#[test]
fn test_f64_conversion_negative_zero() {
let parts = DecimalParts::new(false, 1, 0, 0);
let result = parts.to_f64();
assert_eq!(result, -0.0);
}
#[test]
fn test_string_decoder_is_long_len_type_ntext() {
assert!(StringDecoder::is_long_len_type(TdsDataType::NText));
}
#[test]
fn test_string_decoder_is_long_len_type_text() {
assert!(StringDecoder::is_long_len_type(TdsDataType::Text));
}
#[test]
fn test_string_decoder_is_long_len_type_not_long() {
assert!(!StringDecoder::is_long_len_type(TdsDataType::NVarChar));
assert!(!StringDecoder::is_long_len_type(TdsDataType::BigVarChar));
assert!(!StringDecoder::is_long_len_type(TdsDataType::Int4));
}
#[test]
fn test_decimal_parts_f64_conversion_with_wide_magnitude() {
let parts = DecimalParts::new(true, 30, 0, 55340232229718589441);
let result = parts.to_f64();
assert!(result > 0.0);
}
#[test]
fn test_decimal_parts_equality_reversed_order() {
let parts1 = DecimalParts::new(true, 10, 2, 100);
let parts2 = DecimalParts::new(true, 10, 2, 429496729600);
assert_ne!(parts1, parts2);
}
#[test]
fn test_decimal_parts_equality_both_empty() {
let parts1 = DecimalParts::new(true, 1, 0, 0);
let parts2 = DecimalParts::new(true, 1, 0, 0);
assert_eq!(parts1, parts2);
}
#[test]
fn test_validate_alloc_size_mid_range() {
let result = validate_alloc_size(MAX_ALLOC_SIZE / 2, "test_mid_range");
assert!(result.is_ok());
}
#[test]
fn test_decimal_parts_f64_negative_with_scale() {
let parts = DecimalParts::new(false, 10, 3, 123456);
let result = parts.to_f64();
assert!((result + 123.456).abs() < 0.001);
}
#[test]
fn test_decimal_parts_equality_one_empty_one_zero() {
let parts1 = DecimalParts::new(true, 1, 0, 0);
let parts2 = DecimalParts::new(true, 1, 0, 0);
assert_eq!(parts1, parts2);
}
#[test]
fn test_decimal_parts_debug_with_zero() {
let parts = DecimalParts::new(true, 1, 0, 0);
let debug_str = format!("{parts:?}");
assert!(debug_str.contains("0"));
assert!(debug_str.contains("F64 value: 0"));
}
#[test]
fn test_from_string_positive_decimal() {
let result = DecimalParts::from_string("123.45", 10, 2);
assert!(result.is_ok());
let parts = result.unwrap();
assert!(parts.is_positive);
assert_eq!(parts.scale, 2);
assert_eq!(parts.precision, 10);
assert_eq!(parts.to_decimal_string(), "123.45");
}
#[test]
fn test_from_string_negative_decimal() {
let result = DecimalParts::from_string("-123.45", 10, 2);
assert!(result.is_ok());
let parts = result.unwrap();
assert!(!parts.is_positive);
assert_eq!(parts.scale, 2);
assert_eq!(parts.precision, 10);
assert_eq!(parts.to_decimal_string(), "-123.45");
}
#[test]
fn test_from_string_integer_no_decimal_point() {
let result = DecimalParts::from_string("12345", 10, 0);
assert!(result.is_ok());
let parts = result.unwrap();
assert!(parts.is_positive);
assert_eq!(parts.scale, 0);
assert_eq!(parts.to_decimal_string(), "12345");
}
#[test]
fn test_from_string_with_leading_zeros() {
let result = DecimalParts::from_string("00123.45", 10, 2);
assert!(result.is_ok());
let parts = result.unwrap();
assert_eq!(parts.to_decimal_string(), "123.45");
}
#[test]
fn test_from_string_small_fractional_value() {
let result = DecimalParts::from_string("0.01", 10, 2);
assert!(result.is_ok());
let parts = result.unwrap();
assert_eq!(parts.to_decimal_string(), "0.01");
}
#[test]
fn test_from_string_zero() {
let result = DecimalParts::from_string("0", 10, 0);
assert!(result.is_ok());
let parts = result.unwrap();
assert_eq!(parts.to_decimal_string(), "0");
}
#[test]
fn test_from_string_zero_with_scale() {
let result = DecimalParts::from_string("0.00", 10, 2);
assert!(result.is_ok());
let parts = result.unwrap();
assert_eq!(parts.to_decimal_string(), "0.00");
}
#[test]
fn test_from_string_fractional_padding() {
let result = DecimalParts::from_string("1.5", 10, 3);
assert!(result.is_ok());
let parts = result.unwrap();
assert_eq!(parts.to_decimal_string(), "1.500");
}
#[test]
fn test_from_string_max_precision_38_digits() {
let value = "12345678901234567890123456789012345678";
let result = DecimalParts::from_string(value, 38, 0);
assert!(result.is_ok());
let parts = result.unwrap();
assert_eq!(parts.to_decimal_string(), value);
}
#[test]
fn test_from_string_high_scale() {
let result = DecimalParts::from_string("123.456789", 10, 6);
assert!(result.is_ok());
let parts = result.unwrap();
assert_eq!(parts.to_decimal_string(), "123.456789");
}
#[test]
fn test_from_string_leading_zeros_precision_check() {
let result = DecimalParts::from_string("00001.00", 5, 2);
assert!(result.is_ok());
let parts = result.unwrap();
assert_eq!(parts.to_decimal_string(), "1.00");
}
#[test]
fn test_from_string_with_positive_sign() {
let result = DecimalParts::from_string("+123.45", 10, 2);
assert!(result.is_ok());
let parts = result.unwrap();
assert!(parts.is_positive);
assert_eq!(parts.to_decimal_string(), "123.45");
}
#[test]
fn test_from_string_invalid_characters() {
let result = DecimalParts::from_string("not_a_number", 10, 2);
assert!(result.is_err());
let error_msg = format!("{:?}", result.unwrap_err());
assert!(error_msg.contains("invalid digit"));
}
#[test]
fn test_from_string_multiple_decimal_points() {
let result = DecimalParts::from_string("123.45.67", 10, 2);
assert!(result.is_err());
let error_msg = format!("{:?}", result.unwrap_err());
assert!(error_msg.contains("Invalid decimal string"));
}
#[test]
fn test_from_string_scale_exceeded() {
let result = DecimalParts::from_string("123.456", 10, 2);
assert!(result.is_err());
let error_msg = format!("{:?}", result.unwrap_err());
assert!(error_msg.contains("scale") && error_msg.contains("exceeds"));
}
#[test]
fn test_from_string_precision_exceeded() {
let result = DecimalParts::from_string("12345", 4, 0);
assert!(result.is_err());
let error_msg = format!("{:?}", result.unwrap_err());
assert!(error_msg.contains("precision") && error_msg.contains("exceeds"));
}
#[test]
fn test_from_string_precision_exceeded_with_decimal() {
let result = DecimalParts::from_string("123.45", 4, 2);
assert!(result.is_err());
let error_msg = format!("{:?}", result.unwrap_err());
assert!(error_msg.contains("precision") && error_msg.contains("exceeds"));
}
#[test]
fn test_from_string_invalid_digit_in_integer_part() {
let result = DecimalParts::from_string("12a34", 10, 0);
assert!(result.is_err());
let error_msg = format!("{:?}", result.unwrap_err());
assert!(error_msg.contains("invalid digit"));
}
#[test]
fn test_from_string_invalid_digit_in_fractional_part() {
let result = DecimalParts::from_string("123.4x5", 10, 2);
assert!(result.is_err());
let error_msg = format!("{:?}", result.unwrap_err());
assert!(error_msg.contains("invalid digit"));
}
#[test]
fn test_from_string_leading_zeros_not_counted_in_precision() {
let result = DecimalParts::from_string("0000123", 3, 0);
assert!(result.is_ok());
let parts = result.unwrap();
assert_eq!(parts.to_decimal_string(), "123");
}
#[test]
fn test_from_string_whitespace_trimmed() {
let result = DecimalParts::from_string(" 123.45 ", 10, 2);
assert!(result.is_ok());
let parts = result.unwrap();
assert_eq!(parts.to_decimal_string(), "123.45");
}
#[test]
fn test_from_string_negative_zero() {
let result = DecimalParts::from_string("-0", 10, 0);
assert!(result.is_ok());
let parts = result.unwrap();
assert!(!parts.is_positive);
assert_eq!(parts.to_decimal_string(), "-0");
}
#[test]
fn test_from_string_only_zeros_with_decimal() {
let result = DecimalParts::from_string("0.0", 10, 1);
assert!(result.is_ok());
let parts = result.unwrap();
assert_eq!(parts.to_decimal_string(), "0.0");
}
#[test]
fn test_from_string_exact_precision_match() {
let result = DecimalParts::from_string("12345", 5, 0);
assert!(result.is_ok());
let parts = result.unwrap();
assert_eq!(parts.to_decimal_string(), "12345");
}
#[test]
fn test_from_string_exact_scale_match() {
let result = DecimalParts::from_string("123.45", 5, 2);
assert!(result.is_ok());
let parts = result.unwrap();
assert_eq!(parts.to_decimal_string(), "123.45");
}
#[test]
fn test_decimal_parts_max_width_magnitude() {
let parts = DecimalParts::new(true, 38, 0, u128::MAX);
assert_eq!(parts.to_decimal_string(), u128::MAX.to_string());
assert_eq!(parts.word_count(), 4);
assert_eq!(parts.word(3), -1);
}
#[test]
fn test_decimal_parts_try_from_words_rejects_oversized_magnitude() {
assert!(DecimalParts::try_from_words(true, 38, 0, &[1, 0, 0, 0, 1]).is_none());
}
#[test]
fn test_decimal_parts_from_words_with_trailing_zero_words() {
let parts = DecimalParts::try_from_words(false, 38, 2, &[12345, 0, 0, 0, 0, 0]).unwrap();
assert_eq!(parts.magnitude(), 12345);
assert_eq!(parts.to_decimal_string(), "-123.45");
}
#[test]
fn test_decimal_parts_empty_magnitude_is_zero() {
let parts = DecimalParts::from_words(true, 38, 0, &[]);
assert_eq!(parts.magnitude(), 0);
assert_eq!(parts.word_count(), 1);
assert_eq!(parts.to_decimal_string(), "0");
}
mod decimal_format_tests {
use super::*;
use crate::datatypes::decoder::DECIMAL_STR_LEN;
use rand::{Rng, SeedableRng, rngs::StdRng};
fn reference(is_positive: bool, scale: u8, magnitude: u128) -> String {
let value_str = magnitude.to_string();
let result = if scale == 0 {
value_str
} else {
let scale_pos = scale as usize;
if value_str.len() <= scale_pos {
format!("0.{}{}", "0".repeat(scale_pos - value_str.len()), value_str)
} else {
let split_pos = value_str.len() - scale_pos;
format!("{}.{}", &value_str[..split_pos], &value_str[split_pos..])
}
};
if is_positive {
result
} else {
format!("-{result}")
}
}
#[track_caller]
fn assert_matches_reference(is_positive: bool, precision: u8, scale: u8, magnitude: u128) {
let parts = DecimalParts::new(is_positive, precision, scale, magnitude);
let expected = reference(is_positive, scale, magnitude);
let mut buf = [0u8; DECIMAL_STR_LEN];
assert_eq!(
parts.format_into(&mut buf),
expected,
"format_into({is_positive}, {scale}, {magnitude})"
);
assert_eq!(parts.to_decimal_string(), expected);
assert_eq!(parts.to_string(), expected);
}
#[test]
fn exhaustive_small_magnitudes_match_reference() {
for magnitude in 0u128..=4096 {
for scale in 0u8..=6 {
assert_matches_reference(true, 38, scale, magnitude);
assert_matches_reference(false, 38, scale, magnitude);
}
}
}
#[test]
fn powers_of_ten_match_reference_at_every_scale() {
let mut boundaries = vec![0u128, u128::MAX];
for k in 1..=38u32 {
let power = 10u128.pow(k);
boundaries.push(power);
boundaries.push(power - 1);
}
boundaries.push(u64::MAX as u128);
boundaries.push(u64::MAX as u128 + 1);
boundaries.push(10_000_000_000_000_000_000u128);
for &magnitude in &boundaries {
for scale in 0..=38u8 {
assert_matches_reference(true, 38, scale, magnitude);
assert_matches_reference(false, 38, scale, magnitude);
}
}
}
#[test]
fn randomized_magnitudes_match_reference() {
let mut rng = StdRng::seed_from_u64(0x5EED_D3C1);
for _ in 0..20_000 {
let bits = rng.random_range(1..=128u32);
let magnitude = if bits == 128 {
rng.random::<u128>()
} else {
rng.random::<u128>() & ((1u128 << bits) - 1)
};
let scale = rng.random_range(0..=38u8);
let precision = rng.random_range(1..=38u8);
assert_matches_reference(rng.random(), precision, scale, magnitude);
}
}
#[test]
fn precision_one_to_thirty_eight_match_reference() {
for precision in 1..=38u8 {
let max = 10u128.pow(precision as u32) - 1;
for scale in 0..=precision {
assert_matches_reference(true, precision, scale, max);
assert_matches_reference(false, precision, scale, max);
assert_matches_reference(true, precision, scale, max / 10);
}
}
}
#[test]
fn negative_zero_renders_with_sign() {
let parts = DecimalParts::new(false, 18, 0, 0);
assert_eq!(parts.to_string(), "-0");
assert_eq!(DecimalParts::new(false, 18, 4, 0).to_string(), "-0.0000");
}
#[test]
fn zero_with_scale_renders_padded() {
let parts = DecimalParts::new(true, 18, 6, 0);
assert_eq!(parts.to_string(), "0.000000");
}
#[test]
fn magnitude_narrower_than_scale_is_zero_padded() {
let parts = DecimalParts::new(true, 18, 3, 5);
assert_eq!(parts.to_string(), "0.005");
}
#[test]
fn out_of_domain_scale_does_not_overrun_the_buffer() {
let parts = DecimalParts::new(false, 38, u8::MAX, u128::MAX);
let mut buf = [0u8; DECIMAL_STR_LEN];
assert_eq!(parts.format_into(&mut buf).len(), DECIMAL_STR_LEN);
}
#[test]
fn format_into_leaves_the_buffer_reusable() {
let mut buf = [0u8; DECIMAL_STR_LEN];
for (magnitude, scale, expected) in [
(1u128, 0u8, "1"),
(123456u128, 4u8, "12.3456"),
(0u128, 0u8, "0"),
] {
let parts = DecimalParts::new(true, 38, scale, magnitude);
assert_eq!(parts.format_into(&mut buf), expected);
}
}
}
mod decimal_constructor_tests {
use super::*;
#[test]
fn from_i64_rejects_a_scale_that_overflows_the_magnitude() {
assert!(DecimalParts::from_i64(1, 38, 38).is_ok());
let err = DecimalParts::from_i64(i64::MAX, 38, 38).unwrap_err();
assert!(matches!(err, crate::error::Error::TypeConversionError(_)));
assert!(DecimalParts::from_i64(1, 38, 39).is_err());
}
#[test]
fn from_i64_round_trips_through_the_formatter() {
let parts = DecimalParts::from_i64(-12345, 18, 3).unwrap();
assert_eq!(parts.magnitude(), 12_345_000);
assert_eq!(parts.to_string(), "-12345.000");
}
#[test]
fn from_string_rejects_a_magnitude_wider_than_128_bits() {
let wide = "1".repeat(40);
assert!(DecimalParts::from_string(&wide, 40, 0).is_err());
}
#[test]
fn from_string_accepts_the_widest_valid_decimal() {
let max = "9".repeat(38);
let parts = DecimalParts::from_string(&max, 38, 0).unwrap();
assert_eq!(parts.magnitude(), 10u128.pow(38) - 1);
assert_eq!(parts.to_string(), max);
}
}
mod decimal_word_tests {
use super::*;
#[test]
fn words_round_trip_through_from_words() {
for magnitude in [
0u128,
1,
u32::MAX as u128,
u32::MAX as u128 + 1,
u64::MAX as u128,
u64::MAX as u128 + 1,
u128::MAX,
12_345_678_901_234_567_890,
] {
let parts = DecimalParts::new(false, 38, 7, magnitude);
let words: Vec<i32> = (0..parts.word_count()).map(|i| parts.word(i)).collect();
assert_eq!(DecimalParts::from_words(false, 38, 7, &words), parts);
}
}
#[test]
fn word_count_covers_the_significant_words() {
let cases = [
(0u128, 1usize),
(1, 1),
(u32::MAX as u128, 1),
(u32::MAX as u128 + 1, 2),
(u64::MAX as u128, 2),
(u64::MAX as u128 + 1, 3),
(u128::MAX, 4),
];
for (magnitude, expected) in cases {
let parts = DecimalParts::new(true, 38, 0, magnitude);
assert_eq!(parts.word_count(), expected, "magnitude {magnitude}");
}
}
#[test]
fn words_past_the_fourth_are_zero() {
let parts = DecimalParts::new(true, 38, 0, u128::MAX);
assert_eq!(parts.word(4), 0);
assert_eq!(parts.word(usize::MAX), 0);
}
}
mod vector_tests {
use super::*;
use crate::datatypes::{
sql_vector::SqlVector,
sqldatatypes::{VectorBaseType, VectorLayoutFormat, VectorLayoutVersion},
};
#[test]
fn test_vector_creation_and_validation() {
let dimensions = vec![1.0, 2.0, 3.0];
let vector = SqlVector::try_from_f32(dimensions.clone()).unwrap();
assert_eq!(vector.as_f32(), Some(dimensions.as_slice()));
assert_eq!(vector.dimension_count(), 3);
}
#[test]
fn test_vector_single_dimension() {
let vector = SqlVector::try_from_f32(vec![42.5]).unwrap();
assert_eq!(vector.as_f32(), Some(&[42.5][..]));
assert_eq!(vector.dimension_count(), 1);
}
#[test]
fn test_vector_max_dimensions() {
let max_dim = VectorBaseType::Float32.max_dimensions();
let dimensions: Vec<f32> = (0..max_dim).map(|i| i as f32).collect();
let vector = SqlVector::try_from_f32(dimensions).unwrap();
assert_eq!(vector.dimension_count(), max_dim);
}
#[test]
fn test_vector_from_raw_valid() {
let values = vec![1.0_f32, 2.0, 3.0];
let mut raw_bytes = Vec::new();
for val in &values {
raw_bytes.extend_from_slice(&val.to_le_bytes());
}
let vector = SqlVector::try_from_raw(
VectorLayoutFormat::V1 as u8,
VectorLayoutVersion::V1 as u8,
VectorBaseType::Float32 as u8,
raw_bytes,
);
assert!(vector.is_ok());
let vector = vector.unwrap();
assert_eq!(vector.as_f32(), Some(values.as_slice()));
}
#[test]
fn test_vector_from_raw_invalid_layout_format() {
let values = vec![1.0_f32, 2.0];
let mut raw_bytes = Vec::new();
for val in &values {
raw_bytes.extend_from_slice(&val.to_le_bytes());
}
let result = SqlVector::try_from_raw(
0x00, VectorLayoutVersion::V1 as u8,
VectorBaseType::Float32 as u8,
raw_bytes,
);
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("layout format"));
}
#[test]
fn test_vector_from_raw_invalid_layout_version() {
let values = vec![1.0_f32, 2.0];
let mut raw_bytes = Vec::new();
for val in &values {
raw_bytes.extend_from_slice(&val.to_le_bytes());
}
let result = SqlVector::try_from_raw(
VectorLayoutFormat::V1 as u8,
0x99, VectorBaseType::Float32 as u8,
raw_bytes,
);
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("layout version"));
}
#[test]
fn test_vector_from_raw_invalid_base_type() {
let values = vec![1.0_f32, 2.0];
let mut raw_bytes = Vec::new();
for val in &values {
raw_bytes.extend_from_slice(&val.to_le_bytes());
}
let result = SqlVector::try_from_raw(
VectorLayoutFormat::V1 as u8,
VectorLayoutVersion::V1 as u8,
0x99, raw_bytes,
);
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("base type"));
}
#[test]
fn test_vector_empty_dimensions() {
let result = SqlVector::try_from_f32(vec![]);
assert!(result.is_err());
assert!(
result
.unwrap_err()
.to_string()
.contains("at least one dimension")
);
}
#[test]
fn test_vector_too_many_dimensions() {
let max_dim = VectorBaseType::Float32.max_dimensions();
let dimensions: Vec<f32> = (0..(max_dim + 1)).map(|i| i as f32).collect();
let result = SqlVector::try_from_f32(dimensions);
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("exceeds maximum"));
}
#[test]
fn test_vector_total_size() {
let vector = SqlVector::try_from_f32(vec![1.0, 2.0, 3.0]).unwrap();
assert_eq!(vector.total_size(), 8 + 3 * 4); }
#[test]
fn test_column_values_vector_variant() {
let vector = SqlVector::try_from_f32(vec![1.0, 2.0, 3.0]).unwrap();
let col_val = ColumnValues::Vector(vector);
match col_val {
ColumnValues::Vector(v) => {
assert_eq!(v.dimension_count(), 3);
assert_eq!(v.as_f32(), Some(&[1.0, 2.0, 3.0][..]));
}
_ => panic!("Expected Vector variant"),
}
}
}
mod decode_into_tests {
use byteorder::{ByteOrder, LittleEndian};
use std::borrow::Cow;
use crate::core::TdsResult;
use crate::datatypes::column_values::{
ColumnValues, SqlDateTime, SqlMoney, SqlSmallDateTime, SqlTime,
};
use crate::datatypes::decoder::{
DecimalParts, GenericDecoder, MAX_PLP_SIZE, PlpChunkReadLength, PlpChunkStreamReader,
PlpColumnStream, SqlTypeDecode,
};
use crate::datatypes::row_writer::{DefaultRowWriter, RowWriter};
use crate::datatypes::sql_string::EncodingType;
use crate::datatypes::sqldatatypes::{
PartialLengthType, TdsDataType, TypeInfo, TypeInfoVariant, VariableLengthTypes,
};
use crate::io::packet_reader::TdsPacketReader;
use crate::query::metadata::ColumnMetadata;
use crate::token::tokens::SqlCollation;
pub(super) struct ByteReader {
data: Vec<u8>,
pos: usize,
deny_slices: bool,
}
impl ByteReader {
pub(super) fn new(data: Vec<u8>) -> Self {
Self {
data,
pos: 0,
deny_slices: false,
}
}
pub(super) fn new_unbuffered(data: Vec<u8>) -> Self {
Self {
data,
pos: 0,
deny_slices: true,
}
}
fn take(&mut self, n: usize) -> TdsResult<&[u8]> {
if self.pos + n > self.data.len() {
return Err(crate::error::Error::Io(std::io::Error::new(
std::io::ErrorKind::UnexpectedEof,
"End of data",
)));
}
let slice = &self.data[self.pos..self.pos + n];
self.pos += n;
Ok(slice)
}
fn try_take<const N: usize>(&mut self) -> Option<[u8; N]> {
let end = self.pos.checked_add(N)?;
let bytes = self.data.get(self.pos..end)?.try_into().ok()?;
self.pos = end;
Some(bytes)
}
}
impl TdsPacketReader for ByteReader {
fn try_read_slice(&mut self, length: usize) -> Option<&[u8]> {
if self.deny_slices {
return None;
}
let end = self.pos.checked_add(length)?;
if end > self.data.len() {
return None;
}
let start = self.pos;
self.pos = end;
Some(&self.data[start..end])
}
fn try_read_byte(&mut self) -> Option<u8> {
self.try_take().map(|[value]| value)
}
fn try_read_int16(&mut self) -> Option<i16> {
self.try_take().map(i16::from_le_bytes)
}
fn try_read_uint16(&mut self) -> Option<u16> {
self.try_take().map(u16::from_le_bytes)
}
fn try_read_uint24(&mut self) -> Option<u32> {
let [b0, b1, b2] = self.try_take()?;
Some(u32::from_le_bytes([b0, b1, b2, 0]))
}
fn try_read_int32(&mut self) -> Option<i32> {
self.try_take().map(i32::from_le_bytes)
}
fn try_read_uint32(&mut self) -> Option<u32> {
self.try_take().map(u32::from_le_bytes)
}
fn try_read_uint40(&mut self) -> Option<u64> {
let [b0, b1, b2, b3, b4] = self.try_take()?;
Some(u64::from_le_bytes([b0, b1, b2, b3, b4, 0, 0, 0]))
}
fn try_read_int64(&mut self) -> Option<i64> {
self.try_take().map(i64::from_le_bytes)
}
fn try_read_float32(&mut self) -> Option<f32> {
self.try_take().map(f32::from_le_bytes)
}
fn try_read_float64(&mut self) -> Option<f64> {
self.try_take().map(f64::from_le_bytes)
}
async fn read_byte(&mut self) -> TdsResult<u8> {
Ok(self.take(1)?[0])
}
async fn read_int16(&mut self) -> TdsResult<i16> {
Ok(LittleEndian::read_i16(self.take(2)?))
}
async fn read_uint16(&mut self) -> TdsResult<u16> {
Ok(LittleEndian::read_u16(self.take(2)?))
}
async fn read_int32(&mut self) -> TdsResult<i32> {
Ok(LittleEndian::read_i32(self.take(4)?))
}
async fn read_uint32(&mut self) -> TdsResult<u32> {
Ok(LittleEndian::read_u32(self.take(4)?))
}
async fn read_int64(&mut self) -> TdsResult<i64> {
Ok(LittleEndian::read_i64(self.take(8)?))
}
async fn read_uint64(&mut self) -> TdsResult<u64> {
Ok(LittleEndian::read_u64(self.take(8)?))
}
async fn read_float32(&mut self) -> TdsResult<f32> {
Ok(LittleEndian::read_f32(self.take(4)?))
}
async fn read_float64(&mut self) -> TdsResult<f64> {
Ok(LittleEndian::read_f64(self.take(8)?))
}
async fn read_uint24(&mut self) -> TdsResult<u32> {
let b = self.take(3)?;
Ok(b[0] as u32 | (b[1] as u32) << 8 | (b[2] as u32) << 16)
}
async fn read_uint40(&mut self) -> TdsResult<u64> {
let b = self.take(5)?;
Ok(b[0] as u64
| (b[1] as u64) << 8
| (b[2] as u64) << 16
| (b[3] as u64) << 24
| (b[4] as u64) << 32)
}
async fn read_bytes(&mut self, buffer: &mut [u8]) -> TdsResult<usize> {
let slice = self.take(buffer.len())?;
buffer.copy_from_slice(slice);
Ok(buffer.len())
}
async fn skip_bytes(&mut self, count: usize) -> TdsResult<()> {
self.take(count)?;
Ok(())
}
async fn read_int16_big_endian(&mut self) -> TdsResult<i16> {
unimplemented!()
}
async fn read_int32_big_endian(&mut self) -> TdsResult<i32> {
unimplemented!()
}
async fn read_varchar_u16_length(&mut self) -> TdsResult<Option<String>> {
unimplemented!()
}
async fn read_varchar_u8_length(&mut self) -> TdsResult<String> {
unimplemented!()
}
async fn read_u8_varbyte(&mut self) -> TdsResult<Vec<u8>> {
unimplemented!()
}
async fn read_u16_varbyte(&mut self) -> TdsResult<Vec<u8>> {
unimplemented!()
}
async fn read_unicode(&mut self, _len: usize) -> TdsResult<String> {
unimplemented!()
}
async fn read_unicode_with_byte_length(&mut self, _len: usize) -> TdsResult<String> {
unimplemented!()
}
async fn cancel_read_stream(&mut self) -> TdsResult<()> {
unimplemented!()
}
fn reset_reader(&mut self) {
self.pos = 0;
}
}
fn fixed_metadata(data_type: TdsDataType, length: usize) -> ColumnMetadata {
ColumnMetadata {
user_type: 0,
flags: 0,
data_type,
type_info: TypeInfo {
tds_type: data_type,
length,
type_info_variant: TypeInfoVariant::FixedLen(
crate::datatypes::sqldatatypes::FixedLengthTypes::try_from(data_type)
.unwrap_or(crate::datatypes::sqldatatypes::FixedLengthTypes::Int4),
),
},
column_name: String::new(),
multi_part_name: None,
crypto_metadata: None,
}
}
pub(super) fn varlen_metadata(data_type: TdsDataType, length: usize) -> ColumnMetadata {
ColumnMetadata {
user_type: 0,
flags: 0,
data_type,
type_info: TypeInfo {
tds_type: data_type,
length,
type_info_variant: TypeInfoVariant::VarLen(
VariableLengthTypes::try_from(data_type)
.unwrap_or(VariableLengthTypes::IntN),
length,
),
},
column_name: String::new(),
multi_part_name: None,
crypto_metadata: None,
}
}
fn buffered_value(bytes: &[u8], metadata: &ColumnMetadata) -> (ColumnValues, usize) {
GenericDecoder::default()
.try_decode_buffered(bytes, metadata)
.unwrap()
.expect("complete buffered value")
}
#[test]
fn buffered_decode_intn_handles_null_value_and_partial_payload() {
let metadata = varlen_metadata(TdsDataType::IntN, 4);
assert_eq!(buffered_value(&[0], &metadata), (ColumnValues::Null, 1));
assert_eq!(
GenericDecoder::default()
.try_decode_buffered(&[4, 0x78, 0x56], &metadata)
.unwrap(),
None
);
assert_eq!(
buffered_value(&[4, 0x78, 0x56, 0x34, 0x12], &metadata),
(ColumnValues::Int(0x1234_5678), 5)
);
}
#[test]
fn buffered_decode_money_preserves_wire_word_order() {
let metadata = fixed_metadata(TdsDataType::Money, 8);
let mut bytes = 7_i32.to_le_bytes().to_vec();
bytes.extend_from_slice(&11_i32.to_le_bytes());
assert_eq!(
buffered_value(&bytes, &metadata),
(
ColumnValues::Money(SqlMoney {
lsb_part: 11,
msb_part: 7,
}),
8,
)
);
}
#[test]
fn buffered_decode_time_applies_fractional_scale() {
let metadata = ColumnMetadata {
user_type: 0,
flags: 0,
data_type: TdsDataType::TimeN,
type_info: TypeInfo::var_len_scale(TdsDataType::TimeN, 4, 3).unwrap(),
column_name: String::new(),
multi_part_name: None,
crypto_metadata: None,
};
let mut bytes = vec![4];
bytes.extend_from_slice(&1234_u32.to_le_bytes());
assert_eq!(
buffered_value(&bytes, &metadata),
(
ColumnValues::Time(SqlTime {
time_nanoseconds: 12_340_000,
scale: 3,
}),
5,
)
);
}
#[test]
fn buffered_decode_guid_uses_tds_little_endian_layout() {
let metadata = varlen_metadata(TdsDataType::Guid, 16);
let expected = uuid::Uuid::from_u128(0x0011_2233_4455_6677_8899_aabb_ccdd_eeff);
let mut bytes = vec![16];
bytes.extend_from_slice(&expected.to_bytes_le());
assert_eq!(
buffered_value(&bytes, &metadata),
(ColumnValues::Uuid(expected), 17)
);
}
#[test]
fn buffered_decode_nvarchar_waits_for_the_complete_payload() {
let metadata = ColumnMetadata {
user_type: 0,
flags: 0,
data_type: TdsDataType::NVarChar,
type_info: TypeInfo::var_len_string(TdsDataType::NVarChar, 100, None).unwrap(),
column_name: String::new(),
multi_part_name: None,
crypto_metadata: None,
};
let payload = "row_1"
.encode_utf16()
.flat_map(u16::to_le_bytes)
.collect::<Vec<_>>();
let mut bytes = u16::try_from(payload.len()).unwrap().to_le_bytes().to_vec();
bytes.extend_from_slice(&payload);
assert_eq!(
GenericDecoder::default()
.try_decode_buffered(&bytes[..bytes.len() - 1], &metadata)
.unwrap(),
None
);
let (value, used) = buffered_value(&bytes, &metadata);
let ColumnValues::String(value) = value else {
panic!("expected string");
};
assert_eq!(value.to_utf8_string(), "row_1");
assert_eq!(used, bytes.len());
}
#[test]
fn buffered_decoders_reject_invalid_or_incomplete_values() {
let decoder = GenericDecoder::default();
let plp = ColumnMetadata {
user_type: 0,
flags: 0,
data_type: TdsDataType::BigVarBinary,
type_info: TypeInfo::partial_len(TdsDataType::BigVarBinary, usize::MAX, None)
.unwrap(),
column_name: String::new(),
multi_part_name: None,
crypto_metadata: None,
};
assert_eq!(decoder.try_decode_buffered(&[], &plp).unwrap(), None);
let mut writer = DefaultRowWriter::new(1);
assert_eq!(
decoder
.try_decode_buffered_into(&[], &plp, 0, &mut writer)
.unwrap(),
None
);
let invalid = [
(vec![3, 0, 0, 0], varlen_metadata(TdsDataType::IntN, 8)),
(
vec![5, 0, 0, 0, 0, 0],
varlen_metadata(TdsDataType::MoneyN, 8),
),
(vec![15], varlen_metadata(TdsDataType::Guid, 16)),
(
vec![4, 0, 0, 0, 0],
scale_metadata(TdsDataType::DateTimeOffsetN, 10, 7),
),
];
for (bytes, metadata) in invalid {
assert!(decoder.try_decode_buffered(&bytes, &metadata).is_err());
let mut writer = DefaultRowWriter::new(1);
assert!(
decoder
.try_decode_buffered_into(&bytes, &metadata, 0, &mut writer)
.is_err()
);
}
let incomplete = [
(vec![0; 7], fixed_metadata(TdsDataType::Int8, 8)),
(vec![0; 3], fixed_metadata(TdsDataType::Money, 8)),
(vec![0; 3], fixed_metadata(TdsDataType::DateTim4, 4)),
(vec![0; 7], fixed_metadata(TdsDataType::DateTime, 8)),
(vec![4, 0, 0], varlen_metadata(TdsDataType::IntN, 4)),
(vec![8, 0, 0, 0, 0], varlen_metadata(TdsDataType::MoneyN, 8)),
(vec![2, 0, b'a'], varlen_metadata(TdsDataType::NVarChar, 10)),
(
vec![2, 0, 1],
varlen_metadata(TdsDataType::BigVarBinary, 10),
),
];
for (bytes, metadata) in incomplete {
assert_eq!(
decoder.try_decode_buffered(&bytes, &metadata).unwrap(),
None
);
let mut writer = DefaultRowWriter::new(1);
assert_eq!(
decoder
.try_decode_buffered_into(&bytes, &metadata, 0, &mut writer)
.unwrap(),
None
);
}
let missing_scale = varlen_metadata(TdsDataType::TimeN, 5);
assert!(
decoder
.try_decode_buffered(&[3, 1, 0, 0], &missing_scale)
.is_err()
);
let mut writer = DefaultRowWriter::new(1);
assert!(
decoder
.try_decode_buffered_into(&[3, 1, 0, 0], &missing_scale, 0, &mut writer,)
.is_err()
);
let invalid_decimal = varlen_metadata(TdsDataType::DecimalN, 9);
let mut writer = DefaultRowWriter::new(1);
assert!(
decoder
.try_decode_buffered_into(
&[5, 1, 1, 0, 0, 0],
&invalid_decimal,
0,
&mut writer,
)
.is_err()
);
let invalid_precision = precision_scale_metadata(TdsDataType::DecimalN, 9, 39, 0);
let mut writer = DefaultRowWriter::new(1);
assert!(
decoder
.try_decode_buffered_into(
&[5, 1, 1, 0, 0, 0],
&invalid_precision,
0,
&mut writer,
)
.is_err()
);
let oversized_decimal = precision_scale_metadata(TdsDataType::DecimalN, 21, 38, 0);
let mut writer = DefaultRowWriter::new(1);
assert!(
decoder
.try_decode_buffered_into(&[21, 1], &oversized_decimal, 0, &mut writer,)
.is_err()
);
let variant = fixed_metadata(TdsDataType::SsVariant, 0);
let mut writer = DefaultRowWriter::new(1);
assert_eq!(
decoder
.try_decode_buffered_into(&[0; 3], &variant, 0, &mut writer)
.unwrap(),
None
);
assert!(writer.take_row().is_empty());
}
#[test]
fn buffered_variant_into_preserves_base_type() {
let decoder = GenericDecoder::default();
let variant = fixed_metadata(TdsDataType::SsVariant, 0);
let mut writer = DefaultRowWriter::new(1);
let wire = [6, 0, 0, 0, TdsDataType::Int4 as u8, 0, 42, 0, 0, 0];
assert_eq!(
decoder
.try_decode_buffered_into(&wire, &variant, 0, &mut writer)
.unwrap(),
Some(wire.len())
);
assert_eq!(writer.variant_base(0), Some(TdsDataType::Int4));
assert_eq!(writer.take_row(), vec![ColumnValues::Int(42)]);
}
#[tokio::test]
async fn buffered_decoders_cover_wide_nullable_values() {
let int = varlen_metadata(TdsDataType::IntN, 8);
let mut int_wire = vec![8];
int_wire.extend_from_slice(&i64::MIN.to_le_bytes());
assert_eq!(
assert_decode_equivalence(int_wire, &int).await,
ColumnValues::BigInt(i64::MIN)
);
let time = scale_metadata(TdsDataType::TimeN, 5, 7);
assert_eq!(
assert_decode_equivalence(vec![0], &time).await,
ColumnValues::Null
);
let numeric = precision_scale_metadata(TdsDataType::NumericN, 9, 18, 4);
let mut numeric_wire = vec![9, 1];
numeric_wire.extend_from_slice(&1_234_500_u64.to_le_bytes());
let buffered_numeric = GenericDecoder::default()
.try_decode_buffered(&numeric_wire, &numeric)
.unwrap();
assert!(matches!(
buffered_numeric,
Some((ColumnValues::Numeric(_), 10))
));
assert!(matches!(
assert_decode_equivalence(numeric_wire, &numeric).await,
ColumnValues::Numeric(_)
));
}
async fn assert_decode_equivalence(
bytes: Vec<u8>,
metadata: &ColumnMetadata,
) -> ColumnValues {
let decoder = GenericDecoder::default();
let mut reader1 = ByteReader::new(bytes.clone());
let expected = decoder.decode(&mut reader1, metadata).await.unwrap();
let consumed = reader1.pos;
let mut reader2 = ByteReader::new(bytes.clone());
let mut writer = DefaultRowWriter::new(1);
decoder
.decode_into(&mut reader2, metadata, 0, &mut writer)
.await
.unwrap();
let row = writer.take_row();
assert_eq!(row.len(), 1);
assert_eq!(
row[0], expected,
"decode_into mismatch for {:?}",
metadata.data_type
);
assert_eq!(reader2.pos, consumed);
if let Some((value, used)) = decoder.try_decode_buffered(&bytes, metadata).unwrap() {
assert_eq!(value, expected, "buffered value mismatch");
assert_eq!(used, consumed, "buffered value consumed wrong width");
assert_eq!(
decoder
.try_decode_buffered(&bytes[..bytes.len() - 1], metadata)
.unwrap(),
None,
"buffered value accepted one-byte-short input"
);
}
let mut buffered_writer = DefaultRowWriter::new(1);
if let Some(used) = decoder
.try_decode_buffered_into(&bytes, metadata, 0, &mut buffered_writer)
.unwrap()
{
assert_eq!(
buffered_writer.take_row()[0],
expected,
"buffered writer mismatch"
);
assert_eq!(used, consumed, "buffered writer consumed wrong width");
let mut short_writer = DefaultRowWriter::new(1);
assert_eq!(
decoder
.try_decode_buffered_into(
&bytes[..bytes.len() - 1],
metadata,
0,
&mut short_writer,
)
.unwrap(),
None,
"buffered writer accepted one-byte-short input"
);
}
expected
}
#[tokio::test]
async fn decode_into_int1() {
let md = fixed_metadata(TdsDataType::Int1, 1);
let val = assert_decode_equivalence(vec![42], &md).await;
assert_eq!(val, ColumnValues::TinyInt(42));
}
#[tokio::test]
async fn decode_into_null_sql_variant_leaves_following_bytes() {
let md = fixed_metadata(TdsDataType::SsVariant, 0);
let mut reader = ByteReader::new(vec![0, 0, 0, 0, 0xAB, 0xCD]);
let decoder = GenericDecoder::default();
let mut writer = DefaultRowWriter::new(1);
decoder
.decode_into(&mut reader, &md, 0, &mut writer)
.await
.unwrap();
assert_eq!(writer.take_row()[0], ColumnValues::Null);
assert_eq!(
reader.read_byte().await.unwrap(),
0xAB,
"the NULL variant consumed the following column's bytes"
);
assert_eq!(
decoder.try_decode_buffered_variant(&[0, 0, 0, 0]).unwrap(),
Some((None, ColumnValues::Null, 4))
);
}
#[tokio::test]
async fn decode_into_sql_variant_reports_the_base_type() {
let md = fixed_metadata(TdsDataType::SsVariant, 0);
let wire = vec![6, 0, 0, 0, 0x38, 0x00, 42, 0, 0, 0];
let mut with_sentinel = wire.clone();
with_sentinel.push(0xAB);
let mut reader = ByteReader::new(with_sentinel);
let decoder = GenericDecoder::default();
let mut writer = DefaultRowWriter::new(1);
decoder
.decode_into(&mut reader, &md, 0, &mut writer)
.await
.unwrap();
assert_eq!(writer.variant_base(0), Some(TdsDataType::Int4));
assert_eq!(writer.take_row()[0], ColumnValues::Int(42));
assert_eq!(reader.read_byte().await.unwrap(), 0xAB);
assert_eq!(
decoder.try_decode_buffered_variant(&wire).unwrap(),
Some((Some(TdsDataType::Int4), ColumnValues::Int(42), wire.len()))
);
assert_eq!(
decoder
.try_decode_buffered_variant(&wire[..wire.len() - 1])
.unwrap(),
None
);
}
#[test]
fn buffered_sql_variant_decodes_nvarchar_and_bigint() {
let text = "ODBCVARIANT"
.encode_utf16()
.flat_map(u16::to_le_bytes)
.collect::<Vec<_>>();
let length = 2 + 7 + text.len();
let mut wire = (length as u32).to_le_bytes().to_vec();
wire.extend_from_slice(&[TdsDataType::NVarChar as u8, 7, 0x09, 0x04, 0, 0, 0, 64, 0]);
wire.extend_from_slice(&text);
let decoder = GenericDecoder::default();
let (base, value, used) = decoder.try_decode_buffered_variant(&wire).unwrap().unwrap();
assert_eq!(base, Some(TdsDataType::NVarChar));
assert_eq!(used, wire.len());
let ColumnValues::String(value) = value else {
panic!("expected string variant");
};
assert_eq!(value.to_utf8_string(), "ODBCVARIANT");
let mut bigint = 10_u32.to_le_bytes().to_vec();
bigint.extend_from_slice(&[TdsDataType::Int8 as u8, 0]);
bigint.extend_from_slice(&i64::MIN.to_le_bytes());
assert_eq!(
decoder.try_decode_buffered_variant(&bigint).unwrap(),
Some((
Some(TdsDataType::Int8),
ColumnValues::BigInt(i64::MIN),
bigint.len(),
))
);
}
#[tokio::test]
async fn buffered_sql_variant_preserves_narrow_character_collations() {
let decoder = GenericDecoder::default();
let cases = [
(
[0x09, 0x04, 0x00, 0x04, 0x00],
"caf\u{e9}".as_bytes().to_vec(),
"caf\u{e9}",
),
(
[0x09, 0x04, 0x00, 0x00, 0x00],
vec![b'c', b'a', b'f', 0xE9],
"caf\u{e9}",
),
];
for (collation_bytes, text, expected_text) in cases {
let collation = SqlCollation::try_from(collation_bytes.as_slice()).unwrap();
let expected_encoding = if collation.utf8() {
EncodingType::Utf8
} else {
EncodingType::LcidBased(collation)
};
let length = 2 + 7 + text.len();
let mut wire = (length as u32).to_le_bytes().to_vec();
wire.extend_from_slice(&[TdsDataType::BigVarChar as u8, 7]);
wire.extend_from_slice(&collation_bytes);
wire.extend_from_slice(&64_u16.to_le_bytes());
wire.extend_from_slice(&text);
let (buffered_base, buffered_value, used) =
decoder.try_decode_buffered_variant(&wire).unwrap().unwrap();
let mut reader = ByteReader::new(wire.clone());
let (async_base, async_value) = decoder
.read_sql_variant_with_base(&mut reader)
.await
.unwrap();
assert_eq!(buffered_base, Some(TdsDataType::BigVarChar));
assert_eq!(buffered_base, async_base);
assert_eq!(used, wire.len());
assert_eq!(reader.pos, wire.len());
let ColumnValues::String(buffered_text) = buffered_value else {
panic!("expected buffered string variant");
};
let ColumnValues::String(async_text) = async_value else {
panic!("expected async string variant");
};
assert_eq!(buffered_text.encoding_type(), &expected_encoding);
assert_eq!(async_text.encoding_type(), &expected_encoding);
assert_eq!(buffered_text, async_text);
assert_eq!(buffered_text.to_utf8_string(), expected_text);
}
}
#[tokio::test]
async fn decode_into_int2() {
let md = fixed_metadata(TdsDataType::Int2, 2);
let mut buf = [0u8; 2];
LittleEndian::write_i16(&mut buf, -1234);
let val = assert_decode_equivalence(buf.to_vec(), &md).await;
assert_eq!(val, ColumnValues::SmallInt(-1234));
}
#[tokio::test]
async fn decode_into_bigbinary_null() {
let md = varlen_metadata(TdsDataType::BigBinary, 8);
let decoder = GenericDecoder::default();
let mut reader = ByteReader::new(vec![0xFF, 0xFF]);
let mut writer = DefaultRowWriter::new(1);
decoder
.decode_into(&mut reader, &md, 0, &mut writer)
.await
.unwrap();
assert_eq!(writer.take_row()[0], ColumnValues::Null);
}
#[tokio::test]
async fn decode_into_bigbinary_value() {
let md = varlen_metadata(TdsDataType::BigBinary, 8);
let decoder = GenericDecoder::default();
let mut bytes = vec![4, 0]; bytes.extend_from_slice(&[0xDE, 0xAD, 0xBE, 0xEF]);
let mut reader = ByteReader::new(bytes);
let mut writer = DefaultRowWriter::new(1);
decoder
.decode_into(&mut reader, &md, 0, &mut writer)
.await
.unwrap();
assert_eq!(
writer.take_row()[0],
ColumnValues::Bytes(vec![0xDE, 0xAD, 0xBE, 0xEF])
);
}
#[tokio::test]
async fn decode_into_bigvarbinary_null_non_plp() {
let md = varlen_metadata(TdsDataType::BigVarBinary, 8);
let decoder = GenericDecoder::default();
let mut reader = ByteReader::new(vec![0xFF, 0xFF]);
let mut writer = DefaultRowWriter::new(1);
decoder
.decode_into(&mut reader, &md, 0, &mut writer)
.await
.unwrap();
assert_eq!(writer.take_row()[0], ColumnValues::Null);
}
#[tokio::test]
async fn decode_into_bigvarbinary_value_non_plp() {
let md = varlen_metadata(TdsDataType::BigVarBinary, 8);
let decoder = GenericDecoder::default();
let mut bytes = vec![3, 0]; bytes.extend_from_slice(&[0x01, 0x02, 0x03]);
let mut reader = ByteReader::new(bytes);
let mut writer = DefaultRowWriter::new(1);
decoder
.decode_into(&mut reader, &md, 0, &mut writer)
.await
.unwrap();
assert_eq!(
writer.take_row()[0],
ColumnValues::Bytes(vec![0x01, 0x02, 0x03])
);
}
struct BorrowSpy(DefaultRowWriter, Vec<bool>);
impl RowWriter for BorrowSpy {
fn write_null(&mut self, col: usize) {
self.0.write_null(col);
}
crate::connection::transport::any_transport::forward_row_values!(
write_bool: bool,
write_u8: u8,
write_i16: i16,
write_i32: i32,
write_i64: i64,
write_f32: f32,
write_f64: f64,
write_decimal: DecimalParts,
write_numeric: DecimalParts,
write_date: crate::datatypes::column_values::SqlDate,
write_time: crate::datatypes::column_values::SqlTime,
write_datetime: crate::datatypes::column_values::SqlDateTime,
write_smalldatetime: crate::datatypes::column_values::SqlSmallDateTime,
write_datetime2: crate::datatypes::column_values::SqlDateTime2,
write_datetimeoffset: crate::datatypes::column_values::SqlDateTimeOffset,
write_money: crate::datatypes::column_values::SqlMoney,
write_smallmoney: crate::datatypes::column_values::SqlSmallMoney,
write_uuid: uuid::Uuid,
write_xml: crate::datatypes::column_values::SqlXml,
write_json: crate::datatypes::sql_json::SqlJson,
write_vector: crate::datatypes::sql_vector::SqlVector,
);
fn write_string(&mut self, col: usize, bytes: Cow<'_, [u8]>, encoding: EncodingType) {
self.1.push(matches!(&bytes, Cow::Borrowed(_)));
self.0.write_string(col, bytes, encoding);
}
fn write_bytes(&mut self, col: usize, bytes: Cow<'_, [u8]>) {
self.1.push(matches!(&bytes, Cow::Borrowed(_)));
self.0.write_bytes(col, bytes);
}
fn end_row(&mut self) {
self.0.end_row();
}
}
async fn borrowed_and_owned_agree(wire: Vec<u8>, md: &ColumnMetadata) -> ColumnValues {
let decoder = GenericDecoder::default();
let mut borrowed_writer = BorrowSpy(DefaultRowWriter::new(1), Vec::new());
let mut borrowed_reader = ByteReader::new(wire.clone());
decoder
.decode_into(&mut borrowed_reader, md, 0, &mut borrowed_writer)
.await
.unwrap();
let borrowed_flags = std::mem::take(&mut borrowed_writer.1);
let borrowed = borrowed_writer.0.take_row()[0].clone();
let mut owned_writer = BorrowSpy(DefaultRowWriter::new(1), Vec::new());
let mut owned_reader = ByteReader::new_unbuffered(wire);
decoder
.decode_into(&mut owned_reader, md, 0, &mut owned_writer)
.await
.unwrap();
let owned_flags = std::mem::take(&mut owned_writer.1);
let owned = owned_writer.0.take_row()[0].clone();
assert_eq!(borrowed, owned, "borrowed arm disagreed with the owned arm");
assert_eq!(
borrowed_flags.len(),
owned_flags.len(),
"the two arms took different numbers of string/binary writes"
);
assert!(
borrowed_flags.iter().all(|borrowed| *borrowed),
"a slice-offering reader still produced an owned value, so the \
zero-copy path this PR exists for did not run: {borrowed_flags:?}"
);
assert!(
owned_flags.iter().all(|borrowed| !*borrowed),
"a reader that never hands out slices somehow produced a borrow: {owned_flags:?}"
);
borrowed
}
fn nvarchar_wire(text: &str) -> Vec<u8> {
let payload: Vec<u8> = text
.encode_utf16()
.flat_map(|unit| unit.to_le_bytes())
.collect();
let mut wire = (payload.len() as u16).to_le_bytes().to_vec();
wire.extend_from_slice(&payload);
wire
}
#[tokio::test]
async fn short_string_borrowed_matches_owned() {
let md = varlen_metadata(TdsDataType::NVarChar, 100);
let value = borrowed_and_owned_agree(nvarchar_wire("abcdef"), &md).await;
match value {
ColumnValues::String(s) => assert_eq!(s.to_utf8_string(), "abcdef"),
other => panic!("expected a string, got {other:?}"),
}
}
#[tokio::test]
async fn short_string_empty_borrowed_matches_owned() {
let md = varlen_metadata(TdsDataType::NVarChar, 100);
let value = borrowed_and_owned_agree(nvarchar_wire(""), &md).await;
match value {
ColumnValues::String(s) => assert_eq!(s.to_utf8_string(), ""),
other => panic!("expected a string, got {other:?}"),
}
}
#[tokio::test]
async fn short_string_null_borrowed_matches_owned() {
let md = varlen_metadata(TdsDataType::NVarChar, 100);
let value = borrowed_and_owned_agree(vec![0xFF, 0xFF], &md).await;
assert_eq!(value, ColumnValues::Null);
}
#[tokio::test]
async fn short_binary_borrowed_matches_owned() {
let md = varlen_metadata(TdsDataType::BigVarBinary, 8);
let mut wire = vec![3, 0];
wire.extend_from_slice(&[0x01, 0x02, 0x03]);
let value = borrowed_and_owned_agree(wire, &md).await;
assert_eq!(value, ColumnValues::Bytes(vec![0x01, 0x02, 0x03]));
}
#[tokio::test]
async fn bigbinary_borrowed_matches_owned() {
let md = varlen_metadata(TdsDataType::BigBinary, 8);
let mut wire = vec![4, 0];
wire.extend_from_slice(&[0xDE, 0xAD, 0xBE, 0xEF]);
let value = borrowed_and_owned_agree(wire, &md).await;
assert_eq!(value, ColumnValues::Bytes(vec![0xDE, 0xAD, 0xBE, 0xEF]));
}
#[tokio::test]
async fn short_string_truncated_errors_on_both_arms() {
let md = varlen_metadata(TdsDataType::NVarChar, 100);
let wire = vec![12, 0, b'a', 0];
let decoder = GenericDecoder::default();
for mut reader in [
ByteReader::new(wire.clone()),
ByteReader::new_unbuffered(wire),
] {
let mut writer = DefaultRowWriter::new(1);
let result = decoder.decode_into(&mut reader, &md, 0, &mut writer).await;
assert!(result.is_err(), "truncated value should not decode");
}
}
#[tokio::test]
async fn decode_into_int4() {
let md = fixed_metadata(TdsDataType::Int4, 4);
let mut buf = [0u8; 4];
LittleEndian::write_i32(&mut buf, 99999);
let val = assert_decode_equivalence(buf.to_vec(), &md).await;
assert_eq!(val, ColumnValues::Int(99999));
}
#[tokio::test]
async fn decode_into_int8() {
let md = fixed_metadata(TdsDataType::Int8, 8);
let mut buf = [0u8; 8];
LittleEndian::write_i64(&mut buf, i64::MAX);
let val = assert_decode_equivalence(buf.to_vec(), &md).await;
assert_eq!(val, ColumnValues::BigInt(i64::MAX));
}
#[tokio::test]
async fn decode_into_intn_null() {
let md = varlen_metadata(TdsDataType::IntN, 4);
let val = assert_decode_equivalence(vec![0], &md).await;
assert_eq!(val, ColumnValues::Null);
}
#[tokio::test]
async fn decode_into_intn_i32() {
let md = varlen_metadata(TdsDataType::IntN, 4);
let mut buf = vec![4u8]; let mut i32_buf = [0u8; 4];
LittleEndian::write_i32(&mut i32_buf, 777);
buf.extend_from_slice(&i32_buf);
let val = assert_decode_equivalence(buf, &md).await;
assert_eq!(val, ColumnValues::Int(777));
}
#[tokio::test]
async fn decode_into_flt4() {
let md = fixed_metadata(TdsDataType::Flt4, 4);
let mut buf = [0u8; 4];
LittleEndian::write_f32(&mut buf, 1.5);
let val = assert_decode_equivalence(buf.to_vec(), &md).await;
assert_eq!(val, ColumnValues::Real(1.5));
}
#[tokio::test]
async fn decode_into_flt8() {
let md = fixed_metadata(TdsDataType::Flt8, 8);
let mut buf = [0u8; 8];
LittleEndian::write_f64(&mut buf, 99.25);
let val = assert_decode_equivalence(buf.to_vec(), &md).await;
assert_eq!(val, ColumnValues::Float(99.25));
}
#[tokio::test]
async fn decode_into_fltn_null() {
let md = varlen_metadata(TdsDataType::FltN, 8);
let val = assert_decode_equivalence(vec![0], &md).await;
assert_eq!(val, ColumnValues::Null);
}
#[tokio::test]
async fn decode_into_fltn_f32() {
let md = varlen_metadata(TdsDataType::FltN, 4);
let mut buf = vec![4u8];
let mut f32_buf = [0u8; 4];
LittleEndian::write_f32(&mut f32_buf, 2.5);
buf.extend_from_slice(&f32_buf);
let val = assert_decode_equivalence(buf, &md).await;
assert_eq!(val, ColumnValues::Real(2.5));
}
#[tokio::test]
async fn decode_into_bit() {
let md = fixed_metadata(TdsDataType::Bit, 1);
let val = assert_decode_equivalence(vec![1], &md).await;
assert_eq!(val, ColumnValues::Bit(true));
}
#[tokio::test]
async fn decode_into_bitn_null() {
let md = varlen_metadata(TdsDataType::BitN, 1);
let val = assert_decode_equivalence(vec![0], &md).await;
assert_eq!(val, ColumnValues::Null);
}
#[tokio::test]
async fn decode_into_bitn_true() {
let md = varlen_metadata(TdsDataType::BitN, 1);
let val = assert_decode_equivalence(vec![1, 1], &md).await;
assert_eq!(val, ColumnValues::Bit(true));
}
#[tokio::test]
async fn decode_into_money4() {
let md = fixed_metadata(TdsDataType::Money4, 4);
let mut buf = [0u8; 4];
LittleEndian::write_i32(&mut buf, 10000); let val = assert_decode_equivalence(buf.to_vec(), &md).await;
assert!(matches!(val, ColumnValues::SmallMoney(_)));
}
#[tokio::test]
async fn decode_into_money8() {
let md = fixed_metadata(TdsDataType::Money, 8);
let mut buf = [0u8; 8];
LittleEndian::write_i32(&mut buf[0..4], 0); LittleEndian::write_i32(&mut buf[4..8], 10000); let val = assert_decode_equivalence(buf.to_vec(), &md).await;
assert!(matches!(val, ColumnValues::Money(_)));
}
#[tokio::test]
async fn decode_into_moneyn_null() {
let md = varlen_metadata(TdsDataType::MoneyN, 8);
let val = assert_decode_equivalence(vec![0], &md).await;
assert_eq!(val, ColumnValues::Null);
}
#[tokio::test]
async fn decode_into_datetime() {
let md = fixed_metadata(TdsDataType::DateTime, 8);
let mut buf = [0u8; 8];
LittleEndian::write_i32(&mut buf[0..4], 43000); LittleEndian::write_u32(&mut buf[4..8], 100); let val = assert_decode_equivalence(buf.to_vec(), &md).await;
assert!(matches!(
val,
ColumnValues::DateTime(SqlDateTime {
days: 43000,
time: 100
})
));
}
#[tokio::test]
async fn decode_into_smalldatetime() {
let md = fixed_metadata(TdsDataType::DateTim4, 4);
let mut buf = [0u8; 4];
LittleEndian::write_u16(&mut buf[0..2], 1000); LittleEndian::write_u16(&mut buf[2..4], 60); let val = assert_decode_equivalence(buf.to_vec(), &md).await;
assert!(matches!(
val,
ColumnValues::SmallDateTime(SqlSmallDateTime {
days: 1000,
time: 60
})
));
}
#[tokio::test]
async fn decode_into_datetimen_null() {
let md = varlen_metadata(TdsDataType::DateTimeN, 8);
let val = assert_decode_equivalence(vec![0], &md).await;
assert_eq!(val, ColumnValues::Null);
}
#[tokio::test]
async fn decode_into_daten() {
let md = varlen_metadata(TdsDataType::DateN, 3);
let val = assert_decode_equivalence(vec![3, 0x01, 0x00, 0x00], &md).await;
assert!(matches!(val, ColumnValues::Date(_)));
}
#[tokio::test]
async fn decode_into_daten_null() {
let md = varlen_metadata(TdsDataType::DateN, 3);
let val = assert_decode_equivalence(vec![0], &md).await;
assert_eq!(val, ColumnValues::Null);
}
#[tokio::test]
async fn decode_into_guid() {
let md = varlen_metadata(TdsDataType::Guid, 16);
let mut buf = vec![16u8]; buf.extend_from_slice(&[1u8; 16]); let val = assert_decode_equivalence(buf, &md).await;
assert!(matches!(val, ColumnValues::Uuid(_)));
}
#[tokio::test]
async fn decode_into_guid_null() {
let md = varlen_metadata(TdsDataType::Guid, 16);
let val = assert_decode_equivalence(vec![0], &md).await;
assert_eq!(val, ColumnValues::Null);
}
#[tokio::test]
async fn plp_chunk_stream_reader_supports_incremental_reads() {
let mut buf = Vec::new();
buf.extend_from_slice(&0xFFFFFFFFFFFFFFFEu64.to_le_bytes());
buf.extend_from_slice(&3u32.to_le_bytes());
buf.extend_from_slice(b"abc");
buf.extend_from_slice(&2u32.to_le_bytes());
buf.extend_from_slice(b"de");
buf.extend_from_slice(&0u32.to_le_bytes());
let mut reader = ByteReader::new(buf);
let mut stream = PlpChunkStreamReader::begin(&mut reader)
.await
.unwrap()
.expect("not null");
let mut out = [0u8; 2];
let n1 = stream.read_into(&mut reader, &mut out).await.unwrap();
assert_eq!(n1, 2);
assert_eq!(&out, b"ab");
let n2 = stream.read_into(&mut reader, &mut out).await.unwrap();
assert_eq!(n2, 2);
assert_eq!(&out, b"cd");
let mut out_last = [0u8; 2];
let n3 = stream.read_into(&mut reader, &mut out_last).await.unwrap();
assert_eq!(n3, 1);
assert_eq!(out_last[0], b'e');
assert_eq!(stream.total_read(), 5);
assert!(stream.reached_end());
let mut empty: [u8; 0] = [];
let n4 = stream.read_into(&mut reader, &mut empty).await.unwrap();
assert_eq!(n4, 0);
assert!(stream.reached_end());
}
#[tokio::test]
async fn plp_chunk_stream_reader_skip_to_end_flushes_remaining_chunks() {
let mut buf = Vec::new();
buf.extend_from_slice(&5u64.to_le_bytes());
buf.extend_from_slice(&2u32.to_le_bytes());
buf.extend_from_slice(b"ab");
buf.extend_from_slice(&3u32.to_le_bytes());
buf.extend_from_slice(b"cde");
buf.extend_from_slice(&0u32.to_le_bytes());
let mut reader = ByteReader::new(buf);
let mut stream = PlpChunkStreamReader::begin(&mut reader)
.await
.unwrap()
.expect("not null");
let mut one = [0u8; 1];
let n = stream.read_into(&mut reader, &mut one).await.unwrap();
assert_eq!(n, 1);
assert_eq!(one[0], b'a');
stream.skip_to_end(&mut reader).await.unwrap();
assert!(stream.reached_end());
assert_eq!(stream.total_read(), 5);
}
#[tokio::test]
async fn plp_chunk_stream_reader_partial_read_then_skip_drains_all_bytes() {
let mut buf = Vec::new();
buf.extend_from_slice(&0xFFFFFFFFFFFFFFFEu64.to_le_bytes()); buf.extend_from_slice(&10u32.to_le_bytes()); buf.extend_from_slice(b"0123456789");
buf.extend_from_slice(&0u32.to_le_bytes());
let mut reader = ByteReader::new(buf);
let mut stream = PlpChunkStreamReader::begin(&mut reader)
.await
.unwrap()
.expect("not null");
let mut partial = [0u8; 3];
let n = stream.read_into(&mut reader, &mut partial).await.unwrap();
assert_eq!(n, 3);
assert_eq!(&partial, b"012");
assert!(!stream.reached_end());
stream.skip_to_end(&mut reader).await.unwrap();
assert!(stream.reached_end());
assert_eq!(stream.total_read(), 10);
}
#[tokio::test]
async fn plp_column_stream_repeated_small_reads_exhaust_payload() {
let payload = b"hello world";
let md = plp_metadata(
TdsDataType::BigVarBinary,
PartialLengthType::BigVarBinary,
None,
);
let mut reader = ByteReader::new(plp_wire(payload));
let mut stream = PlpColumnStream::begin(&md, &mut reader)
.await
.unwrap()
.unwrap();
let mut collected = Vec::new();
let mut chunk = [0u8; 4];
loop {
let n = stream.read_into(&mut reader, &mut chunk).await.unwrap();
if n == 0 {
break;
}
collected.extend_from_slice(&chunk[..n]);
}
assert_eq!(collected, payload);
assert!(stream.reached_end());
}
#[test]
fn try_begin_buffered_returns_none_for_a_short_header() {
assert!(
PlpChunkStreamReader::try_begin_buffered(&[0; 7])
.unwrap()
.is_none()
);
}
#[test]
fn try_begin_buffered_reports_a_null_plp_value() {
let header = 0xFFFF_FFFF_FFFF_FFFFu64.to_le_bytes();
let result = PlpChunkStreamReader::try_begin_buffered(&header).unwrap();
assert!(matches!(result, Some((None, 8))));
}
#[test]
fn try_begin_buffered_rejects_a_declared_length_over_the_maximum() {
let header = (MAX_PLP_SIZE as u64 + 1).to_le_bytes();
assert!(PlpChunkStreamReader::try_begin_buffered(&header).is_err());
}
#[test]
fn try_read_complete_buffered_ignores_unknown_length_streams() {
let mut stream = PlpChunkStreamReader::new(PlpChunkReadLength::Unknown);
let mut out = [0u8; 4];
assert_eq!(
stream
.try_read_complete_buffered(&[0; 8], &mut out)
.unwrap(),
None
);
}
#[test]
fn try_read_complete_buffered_skips_once_already_finished() {
let mut stream = PlpChunkStreamReader::new(PlpChunkReadLength::Known(0));
let mut out = [];
assert!(
stream
.try_read_complete_buffered(&[0; 8], &mut out)
.unwrap()
.is_some()
);
assert_eq!(
stream
.try_read_complete_buffered(&[0; 8], &mut out)
.unwrap(),
None
);
}
#[test]
fn try_read_complete_buffered_rejects_an_output_buffer_too_small() {
let mut stream = PlpChunkStreamReader::new(PlpChunkReadLength::Known(4));
let mut out = [0u8; 2];
let mut bytes = Vec::from(4u32.to_le_bytes());
bytes.extend([1, 2, 3, 4]);
bytes.extend(0u32.to_le_bytes());
assert_eq!(
stream.try_read_complete_buffered(&bytes, &mut out).unwrap(),
None
);
}
#[test]
fn try_read_complete_buffered_needs_a_full_chunk_header() {
let mut stream = PlpChunkStreamReader::new(PlpChunkReadLength::Known(4));
let mut out = [0u8; 4];
assert_eq!(
stream
.try_read_complete_buffered(&[1, 2, 3], &mut out)
.unwrap(),
None
);
}
#[test]
fn try_read_complete_buffered_rejects_a_chunk_length_mismatch() {
let mut stream = PlpChunkStreamReader::new(PlpChunkReadLength::Known(4));
let mut out = [0u8; 4];
let mut bytes = Vec::from(5u32.to_le_bytes()); bytes.extend([1, 2, 3, 4, 5]);
bytes.extend(0u32.to_le_bytes());
assert_eq!(
stream.try_read_complete_buffered(&bytes, &mut out).unwrap(),
None
);
}
#[test]
fn try_read_complete_buffered_needs_the_full_payload() {
let mut stream = PlpChunkStreamReader::new(PlpChunkReadLength::Known(4));
let mut out = [0u8; 4];
let mut bytes = Vec::from(4u32.to_le_bytes());
bytes.extend([1, 2]); assert_eq!(
stream.try_read_complete_buffered(&bytes, &mut out).unwrap(),
None
);
}
#[test]
fn try_read_complete_buffered_needs_the_full_terminator() {
let mut stream = PlpChunkStreamReader::new(PlpChunkReadLength::Known(4));
let mut out = [0u8; 4];
let mut bytes = Vec::from(4u32.to_le_bytes());
bytes.extend([1, 2, 3, 4]);
bytes.extend([0, 0]); assert_eq!(
stream.try_read_complete_buffered(&bytes, &mut out).unwrap(),
None
);
}
#[test]
fn try_read_complete_buffered_rejects_a_non_zero_terminator() {
let mut stream = PlpChunkStreamReader::new(PlpChunkReadLength::Known(4));
let mut out = [0u8; 4];
let mut bytes = Vec::from(4u32.to_le_bytes());
bytes.extend([1, 2, 3, 4]);
bytes.extend(1u32.to_le_bytes()); assert_eq!(
stream.try_read_complete_buffered(&bytes, &mut out).unwrap(),
None
);
}
#[test]
fn try_ensure_active_buffered_chunk_detects_position_overflow() {
let mut stream = PlpChunkStreamReader::new(PlpChunkReadLength::Known(4));
let mut position = usize::MAX - 2;
let err = stream
.try_ensure_active_buffered_chunk(&[], &mut position)
.unwrap_err();
assert!(matches!(err, crate::error::Error::ProtocolError(_)));
}
#[test]
fn try_read_buffered_inner_zero_output_reports_missing_header() {
let mut stream = PlpChunkStreamReader::new(PlpChunkReadLength::Known(4));
assert_eq!(stream.try_read_buffered_inner(&[], None, 0).unwrap(), None);
}
#[test]
fn try_read_buffered_inner_rejects_an_inconsistent_output_length() {
let mut stream = PlpChunkStreamReader::new(PlpChunkReadLength::Known(4));
let mut bytes = Vec::from(4u32.to_le_bytes());
bytes.extend([1, 2, 3, 4]);
bytes.extend(0u32.to_le_bytes());
let mut out = [0u8; 2];
let err = stream
.try_read_buffered_inner(&bytes, Some(&mut out), 4)
.unwrap_err();
assert!(matches!(err, crate::error::Error::ProtocolError(_)));
}
#[test]
fn try_read_buffered_inner_reports_missing_trailing_header_after_exact_fill() {
let mut stream = PlpChunkStreamReader::new(PlpChunkReadLength::Known(4));
let mut bytes = Vec::from(4u32.to_le_bytes());
bytes.extend([1, 2, 3, 4]); let mut out = [0u8; 4];
assert_eq!(
stream
.try_read_buffered_inner(&bytes, Some(&mut out), 4)
.unwrap(),
None
);
}
pub(super) fn plp_metadata(
data_type: TdsDataType,
partial_type: PartialLengthType,
collation: Option<crate::token::tokens::SqlCollation>,
) -> ColumnMetadata {
ColumnMetadata {
user_type: 0,
flags: 0,
data_type,
type_info: TypeInfo {
tds_type: data_type,
length: 0xFFFF,
type_info_variant: TypeInfoVariant::PartialLen(
partial_type,
Some(0xFFFF),
collation,
None,
None,
),
},
column_name: "col".to_string(),
multi_part_name: None,
crypto_metadata: None,
}
}
fn plp_wire(payload: &[u8]) -> Vec<u8> {
let mut buf = Vec::new();
buf.extend_from_slice(&0xFFFFFFFFFFFFFFFEu64.to_le_bytes()); buf.extend_from_slice(&(payload.len() as u32).to_le_bytes());
buf.extend_from_slice(payload);
buf.extend_from_slice(&0u32.to_le_bytes()); buf
}
#[tokio::test]
async fn plp_column_stream_null_returns_none() {
let mut buf = Vec::new();
buf.extend_from_slice(&0xFFFFFFFFFFFFFFFFu64.to_le_bytes()); let md = plp_metadata(
TdsDataType::BigVarBinary,
PartialLengthType::BigVarBinary,
None,
);
let mut reader = ByteReader::new(buf);
let result = PlpColumnStream::begin(&md, &mut reader).await.unwrap();
assert!(result.is_none());
}
#[test]
fn plp_column_stream_try_begin_buffered_returns_none_for_a_short_header() {
let md = plp_metadata(
TdsDataType::BigVarBinary,
PartialLengthType::BigVarBinary,
None,
);
assert!(
PlpColumnStream::try_begin_buffered(&md, &[0; 4])
.unwrap()
.is_none()
);
}
#[tokio::test]
async fn plp_column_stream_binary_kind() {
let payload = b"binarydata";
let md = plp_metadata(
TdsDataType::BigVarBinary,
PartialLengthType::BigVarBinary,
None,
);
let mut reader = ByteReader::new(plp_wire(payload));
let mut stream = PlpColumnStream::begin(&md, &mut reader)
.await
.unwrap()
.unwrap();
assert_eq!(stream.plp_type(), PartialLengthType::BigVarBinary);
let mut out = vec![0u8; payload.len()];
let n = stream.read_into(&mut reader, &mut out).await.unwrap();
assert_eq!(n, payload.len());
assert_eq!(&out, payload);
assert!(stream.reached_end());
}
#[tokio::test]
async fn plp_column_stream_unicode_text_kind() {
let payload = b"hi"; let md = plp_metadata(TdsDataType::NVarChar, PartialLengthType::NVarChar, None);
let mut reader = ByteReader::new(plp_wire(payload));
let mut stream = PlpColumnStream::begin(&md, &mut reader)
.await
.unwrap()
.unwrap();
assert_eq!(stream.plp_type(), PartialLengthType::NVarChar);
let mut out = vec![0u8; 2];
let n = stream.read_into(&mut reader, &mut out).await.unwrap();
assert_eq!(n, 2);
assert!(stream.reached_end());
}
#[tokio::test]
async fn plp_column_stream_bigvarchar_kind_carries_collation() {
let payload = b"hello";
let col = crate::token::tokens::SqlCollation {
info: 0x0409_0034,
lcid_language_id: 0x0409,
col_flags: 0,
sort_id: 52,
};
let md = plp_metadata(
TdsDataType::BigVarChar,
PartialLengthType::BigVarChar,
Some(col),
);
let mut reader = ByteReader::new(plp_wire(payload));
let mut stream = PlpColumnStream::begin(&md, &mut reader)
.await
.unwrap()
.unwrap();
assert_eq!(stream.plp_type(), PartialLengthType::BigVarChar);
assert!(stream.collation().is_some());
let mut out = vec![0u8; 5];
let n = stream.read_into(&mut reader, &mut out).await.unwrap();
assert_eq!(n, 5);
assert_eq!(&out, b"hello");
assert!(stream.reached_end());
}
#[tokio::test]
async fn plp_column_stream_xml_kind() {
let payload = b"<r/>";
let md = plp_metadata(TdsDataType::Xml, PartialLengthType::Xml, None);
let mut reader = ByteReader::new(plp_wire(payload));
let mut stream = PlpColumnStream::begin(&md, &mut reader)
.await
.unwrap()
.unwrap();
assert_eq!(stream.plp_type(), PartialLengthType::Xml);
stream.skip_to_end(&mut reader).await.unwrap();
assert!(stream.reached_end());
assert_eq!(stream.total_read(), 4);
}
#[tokio::test]
async fn plp_column_stream_json_kind() {
let payload = b"{}";
let md = plp_metadata(TdsDataType::Json, PartialLengthType::Json, None);
let mut reader = ByteReader::new(plp_wire(payload));
let mut stream = PlpColumnStream::begin(&md, &mut reader)
.await
.unwrap()
.unwrap();
assert_eq!(stream.plp_type(), PartialLengthType::Json);
let mut out = vec![0u8; 2];
let n = stream.read_into(&mut reader, &mut out).await.unwrap();
assert_eq!(n, 2);
assert!(stream.reached_end());
}
#[tokio::test]
async fn plp_column_stream_udt_kind() {
let payload = b"\x01\x02\x03";
let md = plp_metadata(TdsDataType::Udt, PartialLengthType::Udt, None);
let mut reader = ByteReader::new(plp_wire(payload));
let mut stream = PlpColumnStream::begin(&md, &mut reader)
.await
.unwrap()
.unwrap();
assert_eq!(stream.plp_type(), PartialLengthType::Udt);
let mut out = vec![0u8; 3];
let n = stream.read_into(&mut reader, &mut out).await.unwrap();
assert_eq!(n, 3);
assert!(stream.reached_end());
}
#[tokio::test]
async fn plp_column_stream_rejects_non_plp_metadata() {
let md = varlen_metadata(TdsDataType::NVarChar, 100); let buf = plp_wire(b"x"); let mut reader = ByteReader::new(buf);
let err = PlpColumnStream::begin(&md, &mut reader).await.unwrap_err();
assert!(
err.to_string().contains("is not a PLP type"),
"unexpected: {err}"
);
}
#[tokio::test]
async fn decode_xml_rejects_non_plp_metadata() {
let md = varlen_metadata(TdsDataType::Xml, 100);
let decoder = GenericDecoder::default();
let mut reader = ByteReader::new(plp_wire(b"x"));
let err = decoder.decode(&mut reader, &md).await.unwrap_err();
assert!(
err.to_string()
.contains("XML column metadata is not partially-length-prefixed"),
"unexpected: {err}"
);
}
#[tokio::test]
async fn decode_json_rejects_non_plp_metadata() {
let md = varlen_metadata(TdsDataType::Json, 100);
let decoder = GenericDecoder::default();
let mut reader = ByteReader::new(plp_wire(b"x"));
let err = decoder.decode(&mut reader, &md).await.unwrap_err();
assert!(
err.to_string()
.contains("JSON column metadata is not partially-length-prefixed"),
"unexpected: {err}"
);
}
#[tokio::test]
async fn decode_udt_rejects_non_plp_metadata() {
let md = varlen_metadata(TdsDataType::Udt, 100);
let decoder = GenericDecoder::default();
let mut reader = ByteReader::new(plp_wire(b"x"));
let err = decoder.decode(&mut reader, &md).await.unwrap_err();
assert!(
err.to_string()
.contains("UDT column metadata is not partially-length-prefixed"),
"unexpected: {err}"
);
}
#[tokio::test]
async fn plp_chunk_stream_reader_known_length_overflow_errors() {
let mut buf = Vec::new();
buf.extend_from_slice(&4u64.to_le_bytes());
buf.extend_from_slice(&5u32.to_le_bytes());
let mut reader = ByteReader::new(buf);
let mut stream = PlpChunkStreamReader::begin(&mut reader)
.await
.unwrap()
.expect("not null");
let mut out = [0u8; 1];
let err = stream.read_into(&mut reader, &mut out).await.unwrap_err();
assert!(
err.to_string()
.contains("PLP chunk exceeds declared length"),
"unexpected error: {err}"
);
}
#[tokio::test]
async fn plp_chunk_stream_reader_known_length_early_terminator_errors() {
let mut buf = Vec::new();
buf.extend_from_slice(&4u64.to_le_bytes());
buf.extend_from_slice(&2u32.to_le_bytes());
buf.extend_from_slice(b"ab");
buf.extend_from_slice(&0u32.to_le_bytes());
let mut reader = ByteReader::new(buf);
let mut stream = PlpChunkStreamReader::begin(&mut reader)
.await
.unwrap()
.expect("not null");
let mut out = [0u8; 4];
let err = stream.read_into(&mut reader, &mut out).await.unwrap_err();
assert!(
err.to_string()
.contains("PLP stream ended before declared length was reached"),
"unexpected error: {err}"
);
}
#[cfg(not(fuzzing))]
#[tokio::test]
async fn plp_chunk_stream_reader_accepts_more_chunks_than_the_removed_count_cap() {
let mut reader = ByteReader::new(unknown_len_plp_wire(REMOVED_MAX_PLP_CHUNKS + 1));
let mut stream = PlpChunkStreamReader::begin(&mut reader)
.await
.expect("begin")
.expect("not null");
let mut total = 0usize;
let mut out = [0u8; 1];
loop {
let read = stream.read_into(&mut reader, &mut out).await.expect(
"streaming past the removed chunk-count cap must not fail: a chunk count \
limit rejects large LOBs that the protocol allows",
);
if read == 0 {
break;
}
total += read;
}
assert_eq!(total, REMOVED_MAX_PLP_CHUNKS + 1);
}
#[tokio::test]
async fn plp_chunk_stream_reader_accumulated_size_limit_errors() {
let mut stream = PlpChunkStreamReader {
length: PlpChunkReadLength::Known((MAX_PLP_SIZE as u64) + 1),
chunk_remaining: 0,
reached_end: false,
total_read: MAX_PLP_SIZE,
};
let mut reader = ByteReader::new(vec![1, 0, 0, 0]);
let mut out = [0u8; 1];
let err = stream.read_into(&mut reader, &mut out).await.unwrap_err();
assert!(
err.to_string().contains("PLP accumulated size"),
"unexpected error: {err}"
);
}
#[tokio::test]
async fn read_plp_bytes_rejects_oversized_chunk_in_existing_path() {
let mut buf = Vec::new();
buf.extend_from_slice(&0xFFFFFFFFFFFFFFFEu64.to_le_bytes());
let oversized = (GenericDecoder::MAX_PLP_CHUNK_SIZE + 1) as u32;
buf.extend_from_slice(&oversized.to_le_bytes());
let mut reader = ByteReader::new(buf);
let err = GenericDecoder::read_plp_bytes(&mut reader)
.await
.unwrap_err();
assert!(
err.to_string()
.contains("exceeds maximum allowed chunk size"),
"unexpected error: {err}"
);
}
#[cfg(not(fuzzing))]
const REMOVED_MAX_PLP_CHUNKS: usize = 100_000;
#[cfg(not(fuzzing))]
fn unknown_len_plp_wire(chunks: usize) -> Vec<u8> {
let mut buf = Vec::with_capacity(8 + chunks * 5 + 4);
buf.extend_from_slice(&0xFFFFFFFFFFFFFFFEu64.to_le_bytes());
for i in 0..chunks {
buf.extend_from_slice(&1u32.to_le_bytes());
buf.push(i as u8);
}
buf.extend_from_slice(&0u32.to_le_bytes());
buf
}
#[cfg(not(fuzzing))]
#[tokio::test]
async fn read_plp_bytes_accepts_more_chunks_than_the_removed_count_cap() {
let chunks = REMOVED_MAX_PLP_CHUNKS + 1;
let mut reader = ByteReader::new(unknown_len_plp_wire(chunks));
let value = GenericDecoder::read_plp_bytes(&mut reader)
.await
.expect(
"a PLP value split into more chunks than the removed cap must decode: \
rejecting it caps readable LOB size far below the 2 GB the protocol allows",
)
.expect("not null");
assert_eq!(value.len(), chunks);
assert_eq!(value[0], 0);
assert_eq!(value[chunks - 1], ((chunks - 1) % 256) as u8);
}
#[cfg(not(fuzzing))]
#[tokio::test]
async fn read_plp_bytes_accepts_many_chunks_for_a_known_length_value() {
let chunks = REMOVED_MAX_PLP_CHUNKS + 1;
let mut buf = Vec::with_capacity(8 + chunks * 5 + 4);
buf.extend_from_slice(&(chunks as u64).to_le_bytes());
for i in 0..chunks {
buf.extend_from_slice(&1u32.to_le_bytes());
buf.push(i as u8);
}
buf.extend_from_slice(&0u32.to_le_bytes());
let mut reader = ByteReader::new(buf);
let value = GenericDecoder::read_plp_bytes(&mut reader)
.await
.expect("known-length PLP value with many chunks must decode")
.expect("not null");
assert_eq!(value.len(), chunks);
assert_eq!(value[chunks - 1], ((chunks - 1) % 256) as u8);
}
#[tokio::test]
async fn decode_into_bigbinary() {
let md = varlen_metadata(TdsDataType::BigBinary, 4);
let mut buf = Vec::new();
buf.extend_from_slice(&[4, 0]);
buf.extend_from_slice(&[0xDE, 0xAD, 0xBE, 0xEF]);
let val = assert_decode_equivalence(buf, &md).await;
assert_eq!(val, ColumnValues::Bytes(vec![0xDE, 0xAD, 0xBE, 0xEF]));
}
#[tokio::test]
async fn decode_into_nvarchar() {
let md = varlen_metadata(TdsDataType::NVarChar, 100);
let text_bytes = b"hi";
let mut buf = Vec::new();
buf.extend_from_slice(&[2, 0]);
buf.extend_from_slice(text_bytes);
let val = assert_decode_equivalence(buf, &md).await;
assert!(matches!(val, ColumnValues::String(_)));
}
#[tokio::test]
async fn decode_into_nvarchar_null() {
let md = varlen_metadata(TdsDataType::NVarChar, 100);
let val = assert_decode_equivalence(vec![0xFF, 0xFF], &md).await;
assert_eq!(val, ColumnValues::Null);
}
#[tokio::test]
async fn decode_into_decimaln_null() {
let md = ColumnMetadata {
user_type: 0,
flags: 0,
data_type: TdsDataType::DecimalN,
type_info: TypeInfo {
tds_type: TdsDataType::DecimalN,
length: 9,
type_info_variant: TypeInfoVariant::VarLenPrecisionScale(
VariableLengthTypes::DecimalN,
9,
18,
5,
),
},
column_name: String::new(),
multi_part_name: None,
crypto_metadata: None,
};
let val = assert_decode_equivalence(vec![0], &md).await;
assert_eq!(val, ColumnValues::Null);
}
#[tokio::test]
async fn decode_into_decimaln_value() {
let md = ColumnMetadata {
user_type: 0,
flags: 0,
data_type: TdsDataType::DecimalN,
type_info: TypeInfo {
tds_type: TdsDataType::DecimalN,
length: 9,
type_info_variant: TypeInfoVariant::VarLenPrecisionScale(
VariableLengthTypes::DecimalN,
9,
18,
2,
),
},
column_name: String::new(),
multi_part_name: None,
crypto_metadata: None,
};
let mut buf = vec![5u8, 1u8];
let mut part = [0u8; 4];
LittleEndian::write_i32(&mut part, 12345);
buf.extend_from_slice(&part);
let buffered_decimal = GenericDecoder::default()
.try_decode_buffered(&buf, &md)
.unwrap();
assert!(matches!(
buffered_decimal,
Some((ColumnValues::Decimal(_), 6))
));
let val = assert_decode_equivalence(buf, &md).await;
assert!(matches!(val, ColumnValues::Decimal(_)));
}
fn precision_scale_metadata(
data_type: TdsDataType,
length: usize,
precision: u8,
scale: u8,
) -> ColumnMetadata {
ColumnMetadata {
user_type: 0,
flags: 0,
data_type,
type_info: TypeInfo {
tds_type: data_type,
length,
type_info_variant: TypeInfoVariant::VarLenPrecisionScale(
VariableLengthTypes::try_from(data_type)
.unwrap_or(VariableLengthTypes::DecimalN),
length,
precision,
scale,
),
},
column_name: String::new(),
multi_part_name: None,
crypto_metadata: None,
}
}
fn scale_metadata(data_type: TdsDataType, length: usize, scale: u8) -> ColumnMetadata {
ColumnMetadata {
user_type: 0,
flags: 0,
data_type,
type_info: TypeInfo {
tds_type: data_type,
length,
type_info_variant: TypeInfoVariant::VarLenScale(
VariableLengthTypes::try_from(data_type)
.unwrap_or(VariableLengthTypes::TimeN),
scale,
),
},
column_name: String::new(),
multi_part_name: None,
crypto_metadata: None,
}
}
fn ssvariant_metadata() -> ColumnMetadata {
ColumnMetadata {
user_type: 0,
flags: 0,
data_type: TdsDataType::SsVariant,
type_info: TypeInfo {
tds_type: TdsDataType::SsVariant,
length: 8009,
type_info_variant: TypeInfoVariant::VarLen(
VariableLengthTypes::SsVariant,
8009,
),
},
column_name: String::new(),
multi_part_name: None,
crypto_metadata: None,
}
}
async fn assert_decode_err(bytes: Vec<u8>, metadata: &ColumnMetadata) {
let decoder = GenericDecoder::default();
let mut reader = ByteReader::new(bytes);
assert!(decoder.decode(&mut reader, metadata).await.is_err());
}
#[tokio::test]
async fn ssvariant_length_underflow() {
let md = ssvariant_metadata();
let mut buf = Vec::new();
LittleEndian::write_u32(&mut [0u8; 4], 1);
buf.extend_from_slice(&1u32.to_le_bytes());
buf.push(TdsDataType::Int4 as u8); buf.push(0); assert_decode_err(buf, &md).await;
}
#[tokio::test]
async fn ssvariant_unexpected_prop_bytes() {
let md = ssvariant_metadata();
let mut buf = Vec::new();
buf.extend_from_slice(&10u32.to_le_bytes()); buf.push(TdsDataType::Int4 as u8); buf.push(5); assert_decode_err(buf, &md).await;
}
#[tokio::test]
async fn ssvariant_zero_prop_unexpected_type() {
let md = ssvariant_metadata();
let mut buf = Vec::new();
buf.extend_from_slice(&6u32.to_le_bytes()); buf.push(TdsDataType::IntN as u8); buf.push(0); buf.extend_from_slice(&[0; 4]);
assert_decode_err(buf, &md).await;
}
#[tokio::test]
async fn ssvariant_one_prop_unexpected_type() {
let md = ssvariant_metadata();
let mut buf = Vec::new();
buf.extend_from_slice(&20u32.to_le_bytes()); buf.push(TdsDataType::Guid as u8); buf.push(1); buf.push(0); buf.extend_from_slice(&[0; 16]); assert_decode_err(buf, &md).await;
}
#[tokio::test]
async fn ssvariant_two_prop_unexpected_type() {
let md = ssvariant_metadata();
let mut buf = Vec::new();
buf.extend_from_slice(&6u32.to_le_bytes()); buf.push(TdsDataType::IntN as u8); buf.push(2); buf.extend_from_slice(&[0; 4]); assert_decode_err(buf, &md).await;
}
#[tokio::test]
async fn ssvariant_seven_prop_unexpected_type() {
let md = ssvariant_metadata();
let mut buf = Vec::new();
buf.extend_from_slice(&20u32.to_le_bytes()); buf.push(TdsDataType::IntN as u8); buf.push(7); buf.extend_from_slice(&[0; 11]); assert_decode_err(buf, &md).await;
}
#[test]
fn buffered_decimal_into_waits_for_the_complete_value() {
let metadata = precision_scale_metadata(TdsDataType::DecimalN, 9, 18, 4);
let mut wire = vec![9, 1];
wire.extend_from_slice(&123_4500_u64.to_le_bytes());
let decoder = GenericDecoder::default();
let mut incomplete = DefaultRowWriter::new(1);
assert_eq!(
decoder
.try_decode_buffered_into(&wire[..5], &metadata, 0, &mut incomplete)
.unwrap(),
None
);
assert!(incomplete.take_row().is_empty());
let mut complete = DefaultRowWriter::new(1);
assert_eq!(
decoder
.try_decode_buffered_into(&wire, &metadata, 0, &mut complete)
.unwrap(),
Some(wire.len())
);
match &complete.take_row()[0] {
ColumnValues::Decimal(parts) => assert_eq!(parts.to_string(), "123.4500"),
other => panic!("expected Decimal, got {other:?}"),
}
let mut null = DefaultRowWriter::new(1);
assert_eq!(
decoder
.try_decode_buffered_into(&[0], &metadata, 0, &mut null)
.unwrap(),
Some(1)
);
assert_eq!(null.take_row(), vec![ColumnValues::Null]);
}
#[tokio::test]
async fn decimal_wrong_type_info_variant() {
let md = varlen_metadata(TdsDataType::DecimalN, 9);
let mut buf = vec![5u8, 1u8];
buf.extend_from_slice(&42i32.to_le_bytes());
assert_decode_err(buf, &md).await;
}
#[tokio::test]
async fn decimal_rejects_invalid_precision_and_scale() {
for (precision, scale) in [(0, 0), (39, 0), (18, 19)] {
let md = precision_scale_metadata(TdsDataType::DecimalN, 5, precision, scale);
assert_decode_err(vec![0], &md).await;
}
}
#[tokio::test]
async fn decimal_large_valid() {
let md = precision_scale_metadata(TdsDataType::DecimalN, 17, 38, 0);
let mut buf = vec![17u8, 1u8]; for _ in 0..4 {
buf.extend_from_slice(&1i32.to_le_bytes());
}
let val = assert_decode_equivalence(buf, &md).await;
assert!(matches!(val, ColumnValues::Decimal(_)));
}
#[tokio::test]
async fn decimal_too_many_int_parts_rejected() {
let md = precision_scale_metadata(TdsDataType::DecimalN, 21, 38, 0);
let mut buf = vec![21u8, 1u8];
for _ in 0..5 {
buf.extend_from_slice(&1i32.to_le_bytes());
}
assert_decode_err(buf, &md).await;
}
#[tokio::test]
async fn decimal_oversized_partial_word_length_rejected() {
let md = precision_scale_metadata(TdsDataType::DecimalN, 18, 38, 0);
let mut buf = vec![18u8, 1u8];
buf.extend_from_slice(&[1u8; 17]);
assert_decode_err(buf, &md).await;
}
#[tokio::test]
async fn decimal_partial_trailing_word_is_fully_consumed() {
let md = precision_scale_metadata(TdsDataType::DecimalN, 9, 18, 2);
let mut buf = vec![7u8, 1u8];
buf.extend_from_slice(&12345i32.to_le_bytes());
buf.extend_from_slice(&[0u8, 0u8]);
let trailing = 0x5A5A_5A5Ai32;
buf.extend_from_slice(&trailing.to_le_bytes());
let decoder = GenericDecoder::default();
let mut reader = ByteReader::new(buf);
let val = decoder.decode(&mut reader, &md).await.unwrap();
match val {
ColumnValues::Decimal(parts) => assert_eq!(parts.to_string(), "123.45"),
other => panic!("expected Decimal, got {other:?}"),
}
assert_eq!(reader.read_int32().await.unwrap(), trailing);
}
#[tokio::test]
async fn decimal_wire_widths_reassemble_the_magnitude() {
let cases: [(u8, u128); 4] = [
(5, 0xDEAD_BEEF),
(9, 0x0123_4567_89AB_CDEF),
(13, 0x0000_0009_8765_4321_0FED_CBA9),
(17, u128::MAX >> 1),
];
for (length, magnitude) in cases {
let md = precision_scale_metadata(TdsDataType::DecimalN, length.into(), 38, 0);
let mut buf = vec![length, 1u8];
let bytes = magnitude.to_le_bytes();
buf.extend_from_slice(&bytes[..(length - 1) as usize]);
match assert_decode_equivalence(buf, &md).await {
ColumnValues::Decimal(parts) => {
assert_eq!(parts.magnitude(), magnitude, "wire length {length}");
assert!(parts.is_positive);
}
other => panic!("expected Decimal, got {other:?}"),
}
}
}
#[tokio::test]
async fn decimal_negative_sign_byte_is_honored() {
let md = precision_scale_metadata(TdsDataType::DecimalN, 5, 18, 2);
let mut buf = vec![5u8, 0u8];
buf.extend_from_slice(&12345i32.to_le_bytes());
match assert_decode_equivalence(buf, &md).await {
ColumnValues::Decimal(parts) => assert_eq!(parts.to_string(), "-123.45"),
other => panic!("expected Decimal, got {other:?}"),
}
}
#[tokio::test]
async fn time_scale_0() {
let md = scale_metadata(TdsDataType::TimeN, 3, 0);
let mut buf = vec![3u8]; buf.extend_from_slice(&[1, 0, 0]); let val = assert_decode_equivalence(buf, &md).await;
match val {
ColumnValues::Time(t) => {
assert_eq!(t.time_nanoseconds, 10_000_000);
assert_eq!(t.scale, 0);
}
_ => panic!("expected Time"),
}
}
#[tokio::test]
async fn time_scale_1() {
let md = scale_metadata(TdsDataType::TimeN, 3, 1);
let mut buf = vec![3u8];
buf.extend_from_slice(&[1, 0, 0]);
let val = assert_decode_equivalence(buf, &md).await;
match val {
ColumnValues::Time(t) => {
assert_eq!(t.time_nanoseconds, 1_000_000);
assert_eq!(t.scale, 1);
}
_ => panic!("expected Time"),
}
}
#[tokio::test]
async fn time_scale_2() {
let md = scale_metadata(TdsDataType::TimeN, 3, 2);
let mut buf = vec![3u8];
buf.extend_from_slice(&[1, 0, 0]);
let val = assert_decode_equivalence(buf, &md).await;
match val {
ColumnValues::Time(t) => {
assert_eq!(t.time_nanoseconds, 100_000);
assert_eq!(t.scale, 2);
}
_ => panic!("expected Time"),
}
}
#[tokio::test]
async fn time_scale_4() {
let md = scale_metadata(TdsDataType::TimeN, 4, 4);
let mut buf = vec![4u8];
buf.extend_from_slice(&1u32.to_le_bytes());
let val = assert_decode_equivalence(buf, &md).await;
match val {
ColumnValues::Time(t) => {
assert_eq!(t.time_nanoseconds, 1_000);
assert_eq!(t.scale, 4);
}
_ => panic!("expected Time"),
}
}
#[tokio::test]
async fn time_scale_5() {
let md = scale_metadata(TdsDataType::TimeN, 5, 5);
let mut buf = vec![5u8];
buf.extend_from_slice(&[1, 0, 0, 0, 0]); let val = assert_decode_equivalence(buf, &md).await;
match val {
ColumnValues::Time(t) => {
assert_eq!(t.time_nanoseconds, 100);
assert_eq!(t.scale, 5);
}
_ => panic!("expected Time"),
}
}
#[tokio::test]
async fn time_scale_6() {
let md = scale_metadata(TdsDataType::TimeN, 5, 6);
let mut buf = vec![5u8];
buf.extend_from_slice(&[1, 0, 0, 0, 0]);
let val = assert_decode_equivalence(buf, &md).await;
match val {
ColumnValues::Time(t) => {
assert_eq!(t.time_nanoseconds, 10);
assert_eq!(t.scale, 6);
}
_ => panic!("expected Time"),
}
}
#[tokio::test]
async fn datetime2_byte_len_underflow() {
let md = scale_metadata(TdsDataType::DateTime2N, 2, 7);
let buf = vec![2u8, 0, 0];
assert_decode_err(buf, &md).await;
}
#[tokio::test]
async fn datetimeoffset_byte_len_underflow() {
let md = scale_metadata(TdsDataType::DateTimeOffsetN, 1, 7);
let buf = vec![1u8, 0];
assert_decode_err(buf, &md).await;
}
#[tokio::test]
async fn intn_tinyint() {
let md = varlen_metadata(TdsDataType::IntN, 1);
let val = assert_decode_equivalence(vec![1u8, 42u8], &md).await;
assert_eq!(val, ColumnValues::TinyInt(42));
}
#[tokio::test]
async fn intn_smallint() {
let md = varlen_metadata(TdsDataType::IntN, 2);
let mut buf = vec![2u8];
buf.extend_from_slice(&(-100i16).to_le_bytes());
let val = assert_decode_equivalence(buf, &md).await;
assert_eq!(val, ColumnValues::SmallInt(-100));
}
#[tokio::test]
async fn intn_invalid_length() {
let md = varlen_metadata(TdsDataType::IntN, 4);
let buf = vec![3u8, 0, 0, 0]; assert_decode_err(buf, &md).await;
}
#[tokio::test]
async fn guid_invalid_length() {
let md = varlen_metadata(TdsDataType::Guid, 16);
let mut buf = vec![5u8];
buf.extend_from_slice(&[0; 5]);
assert_decode_err(buf, &md).await;
}
#[tokio::test]
async fn fltn_f64() {
let md = varlen_metadata(TdsDataType::FltN, 8);
let mut buf = vec![8u8];
buf.extend_from_slice(&99.5f64.to_le_bytes());
let val = assert_decode_equivalence(buf, &md).await;
assert_eq!(val, ColumnValues::Float(99.5));
}
#[tokio::test]
async fn moneyn_smallmoney() {
let md = varlen_metadata(TdsDataType::MoneyN, 8);
let mut buf = vec![4u8];
buf.extend_from_slice(&10000i32.to_le_bytes());
let val = assert_decode_equivalence(buf, &md).await;
assert!(matches!(val, ColumnValues::SmallMoney(_)));
}
#[tokio::test]
async fn moneyn_money() {
let md = varlen_metadata(TdsDataType::MoneyN, 8);
let mut buf = vec![8u8];
buf.extend_from_slice(&0i32.to_le_bytes()); buf.extend_from_slice(&10000i32.to_le_bytes()); let val = assert_decode_equivalence(buf, &md).await;
assert!(matches!(val, ColumnValues::Money(_)));
}
#[tokio::test]
async fn moneyn_invalid_length() {
let md = varlen_metadata(TdsDataType::MoneyN, 8);
let buf = vec![3u8, 0, 0, 0];
assert_decode_err(buf, &md).await;
}
#[tokio::test]
async fn datetimen_smalldatetime() {
let md = varlen_metadata(TdsDataType::DateTimeN, 8);
let mut buf = vec![4u8];
buf.extend_from_slice(&1000u16.to_le_bytes()); buf.extend_from_slice(&60u16.to_le_bytes()); let val = assert_decode_equivalence(buf, &md).await;
assert!(matches!(
val,
ColumnValues::SmallDateTime(SqlSmallDateTime {
days: 1000,
time: 60
})
));
}
#[tokio::test]
async fn datetimen_datetime() {
let md = varlen_metadata(TdsDataType::DateTimeN, 8);
let mut buf = vec![8u8];
buf.extend_from_slice(&43000i32.to_le_bytes()); buf.extend_from_slice(&100u32.to_le_bytes()); let val = assert_decode_equivalence(buf, &md).await;
assert!(matches!(
val,
ColumnValues::DateTime(SqlDateTime {
days: 43000,
time: 100
})
));
}
#[tokio::test]
async fn bigbinary_empty() {
let md = varlen_metadata(TdsDataType::BigBinary, 100);
let val = assert_decode_equivalence(vec![0x00, 0x00], &md).await;
assert_eq!(val, ColumnValues::Bytes(vec![]));
}
#[tokio::test]
async fn datetime2n_value() {
let md = scale_metadata(TdsDataType::DateTime2N, 8, 7);
let mut buf = vec![8u8];
buf.extend_from_slice(&[0, 0, 0, 0, 0]); buf.extend_from_slice(&[1, 0, 0]); let val = assert_decode_equivalence(buf, &md).await;
assert!(matches!(val, ColumnValues::DateTime2(_)));
}
#[tokio::test]
async fn datetimeoffsetn_value() {
let md = scale_metadata(TdsDataType::DateTimeOffsetN, 10, 7);
let mut buf = vec![10u8];
buf.extend_from_slice(&[0, 0, 0, 0, 0]); buf.extend_from_slice(&[1, 0, 0]); buf.extend_from_slice(&60i16.to_le_bytes()); let val = assert_decode_equivalence(buf, &md).await;
assert!(matches!(val, ColumnValues::DateTimeOffset(_)));
}
}
mod sink_destination_tests {
use std::mem::MaybeUninit;
use super::decode_into_tests::{ByteReader, plp_metadata, varlen_metadata};
use crate::datatypes::column_values::{
SqlDate, SqlDateTime, SqlDateTime2, SqlDateTimeOffset, SqlMoney, SqlSmallDateTime,
SqlSmallMoney, SqlTime, SqlXml,
};
use crate::datatypes::decoder::{DecimalParts, GenericDecoder};
use crate::datatypes::row_writer::{RowWriter, ValueKind};
use crate::datatypes::sql_json::SqlJson;
use crate::datatypes::sql_string::EncodingType;
use crate::datatypes::sql_vector::SqlVector;
use crate::datatypes::sqldatatypes::{PartialLengthType, TdsDataType};
use crate::query::metadata::ColumnMetadata;
use crate::token::tokens::SqlCollation;
use std::borrow::Cow;
use uuid::Uuid;
#[derive(Debug, PartialEq)]
enum Event {
Null,
Sunk {
bytes: Vec<u8>,
encoding: Option<EncodingType>,
complete: bool,
},
OwnedBytes(Vec<u8>),
OwnedString(Vec<u8>),
}
struct RecordingSink {
accept: bool,
destination_len: Option<usize>,
destination_requests: usize,
events: Vec<Event>,
pending: Option<(usize, Option<EncodingType>)>,
storage: Vec<u8>,
}
impl RecordingSink {
fn new(accept: bool) -> Self {
Self {
accept,
destination_len: None,
destination_requests: 0,
events: Vec::new(),
pending: None,
storage: Vec::new(),
}
}
fn with_destination_len(length: usize) -> Self {
Self {
destination_len: Some(length),
..Self::new(true)
}
}
}
impl RowWriter for RecordingSink {
fn value_destination<'a>(
&'a mut self,
_col: usize,
kind: ValueKind<'_>,
length: usize,
) -> Option<&'a mut [MaybeUninit<u8>]> {
self.destination_requests += 1;
if !self.accept {
return None;
}
let length = self.destination_len.unwrap_or(length);
let start = self.storage.len();
self.storage.resize(start + length, 0xAA);
let encoding = match kind {
ValueKind::Bytes => None,
ValueKind::String(encoding) => Some(*encoding),
};
self.pending = Some((start, encoding));
let storage = &mut self.storage[start..];
Some(unsafe {
std::slice::from_raw_parts_mut(
storage.as_mut_ptr().cast::<MaybeUninit<u8>>(),
storage.len(),
)
})
}
fn commit_value(&mut self, _col: usize, complete: bool) {
let (start, encoding) = self.pending.take().expect("commit without destination");
self.events.push(Event::Sunk {
bytes: self.storage[start..].to_vec(),
encoding,
complete,
});
self.storage.truncate(start);
}
fn write_null(&mut self, _col: usize) {
self.events.push(Event::Null);
}
fn write_bytes(&mut self, _col: usize, bytes: Cow<'_, [u8]>) {
self.events.push(Event::OwnedBytes(bytes.into_owned()));
}
fn write_string(
&mut self,
_col: usize,
bytes: Cow<'_, [u8]>,
_encoding_type: EncodingType,
) {
self.events.push(Event::OwnedString(bytes.into_owned()));
}
fn write_bool(&mut self, _col: usize, _val: bool) {}
fn write_u8(&mut self, _col: usize, _val: u8) {}
fn write_i16(&mut self, _col: usize, _val: i16) {}
fn write_i32(&mut self, _col: usize, _val: i32) {}
fn write_i64(&mut self, _col: usize, _val: i64) {}
fn write_f32(&mut self, _col: usize, _val: f32) {}
fn write_f64(&mut self, _col: usize, _val: f64) {}
fn write_decimal(&mut self, _col: usize, _val: DecimalParts) {}
fn write_numeric(&mut self, _col: usize, _val: DecimalParts) {}
fn write_date(&mut self, _col: usize, _val: SqlDate) {}
fn write_time(&mut self, _col: usize, _val: SqlTime) {}
fn write_datetime(&mut self, _col: usize, _val: SqlDateTime) {}
fn write_smalldatetime(&mut self, _col: usize, _val: SqlSmallDateTime) {}
fn write_datetime2(&mut self, _col: usize, _val: SqlDateTime2) {}
fn write_datetimeoffset(&mut self, _col: usize, _val: SqlDateTimeOffset) {}
fn write_money(&mut self, _col: usize, _val: SqlMoney) {}
fn write_smallmoney(&mut self, _col: usize, _val: SqlSmallMoney) {}
fn write_uuid(&mut self, _col: usize, _val: Uuid) {}
fn write_xml(&mut self, _col: usize, _val: SqlXml) {}
fn write_json(&mut self, _col: usize, _val: SqlJson) {}
fn write_vector(&mut self, _col: usize, _val: SqlVector) {}
fn end_row(&mut self) {}
}
fn sunk(bytes: &[u8], encoding: Option<EncodingType>) -> Event {
Event::Sunk {
bytes: bytes.to_vec(),
encoding,
complete: true,
}
}
fn plp_known(payload: &[u8], chunk: usize) -> Vec<u8> {
let mut buf = (payload.len() as u64).to_le_bytes().to_vec();
for part in payload.chunks(chunk.max(1)) {
buf.extend_from_slice(&(part.len() as u32).to_le_bytes());
buf.extend_from_slice(part);
}
buf.extend_from_slice(&0u32.to_le_bytes());
buf
}
fn plp_unknown(payload: &[u8]) -> Vec<u8> {
let mut buf = 0xFFFFFFFFFFFFFFFEu64.to_le_bytes().to_vec();
buf.extend_from_slice(&(payload.len() as u32).to_le_bytes());
buf.extend_from_slice(payload);
buf.extend_from_slice(&0u32.to_le_bytes());
buf
}
fn plp_null() -> Vec<u8> {
0xFFFFFFFFFFFFFFFFu64.to_le_bytes().to_vec()
}
async fn decode_generic(wire: Vec<u8>, md: &ColumnMetadata, accept: bool) -> Vec<Event> {
let mut reader = ByteReader::new(wire);
let mut writer = RecordingSink::new(accept);
GenericDecoder::default()
.decode_into(&mut reader, md, 0, &mut writer)
.await
.unwrap();
writer.events
}
async fn decode_string(wire: Vec<u8>, md: &ColumnMetadata, accept: bool) -> Vec<Event> {
decode_generic(wire, md, accept).await
}
#[tokio::test]
async fn bigvarbinary_non_plp_never_requests_a_destination() {
let md = varlen_metadata(TdsDataType::BigVarBinary, 8000);
let mut wire = 4u16.to_le_bytes().to_vec();
wire.extend_from_slice(&[1, 2, 3, 4]);
assert_eq!(
decode_generic(wire, &md, true).await,
vec![Event::OwnedBytes(vec![1, 2, 3, 4])]
);
}
#[tokio::test]
async fn bigbinary_never_requests_a_destination() {
let md = varlen_metadata(TdsDataType::BigBinary, 8000);
let mut wire = 3u16.to_le_bytes().to_vec();
wire.extend_from_slice(&[9, 8, 7]);
assert_eq!(
decode_generic(wire, &md, true).await,
vec![Event::OwnedBytes(vec![9, 8, 7])]
);
}
#[tokio::test]
async fn bigvarbinary_plp_known_length_uses_sink_across_chunks() {
let md = plp_metadata(
TdsDataType::BigVarBinary,
PartialLengthType::BigVarBinary,
None,
);
let payload: Vec<u8> = (0..=255u8).collect();
assert_eq!(
decode_generic(plp_known(&payload, 30), &md, true).await,
vec![sunk(&payload, None)]
);
}
#[tokio::test]
async fn declining_writer_falls_back_to_owned_path() {
let md = plp_metadata(
TdsDataType::BigVarBinary,
PartialLengthType::BigVarBinary,
None,
);
let payload: Vec<u8> = (0..64u8).collect();
assert_eq!(
decode_generic(plp_known(&payload, 7), &md, false).await,
vec![Event::OwnedBytes(payload)]
);
}
#[tokio::test]
async fn plp_unknown_length_never_requests_a_destination() {
let md = plp_metadata(
TdsDataType::BigVarBinary,
PartialLengthType::BigVarBinary,
None,
);
let payload: Vec<u8> = (0..40u8).collect();
assert_eq!(
decode_generic(plp_unknown(&payload), &md, true).await,
vec![Event::OwnedBytes(payload)]
);
}
#[tokio::test]
async fn plp_null_never_requests_a_destination() {
let md = plp_metadata(
TdsDataType::BigVarBinary,
PartialLengthType::BigVarBinary,
None,
);
assert_eq!(
decode_generic(plp_null(), &md, true).await,
vec![Event::Null]
);
}
#[tokio::test]
async fn non_plp_null_marker_never_requests_a_destination() {
let md = varlen_metadata(TdsDataType::BigVarBinary, 8000);
let wire = 0xFFFFu16.to_le_bytes().to_vec();
assert_eq!(decode_generic(wire, &md, true).await, vec![Event::Null]);
}
#[tokio::test]
async fn empty_plp_value_sinks_an_empty_slice() {
let md = plp_metadata(
TdsDataType::BigVarBinary,
PartialLengthType::BigVarBinary,
None,
);
assert_eq!(
decode_generic(plp_known(&[], 1), &md, true).await,
vec![sunk(&[], None)]
);
}
#[tokio::test]
async fn nvarchar_non_plp_never_requests_a_destination() {
let md = varlen_metadata(TdsDataType::NVarChar, 0xFF);
let utf16: Vec<u8> = "hé".encode_utf16().flat_map(u16::to_le_bytes).collect();
let mut wire = (utf16.len() as u16).to_le_bytes().to_vec();
wire.extend_from_slice(&utf16);
assert_eq!(
decode_string(wire, &md, true).await,
vec![Event::OwnedString(utf16)]
);
}
#[tokio::test]
async fn nvarchar_plp_known_length_sinks_raw_wire_bytes_across_chunks() {
let md = plp_metadata(TdsDataType::NVarChar, PartialLengthType::NVarChar, None);
let utf16: Vec<u8> = "the quick brown fox"
.encode_utf16()
.flat_map(u16::to_le_bytes)
.collect();
assert_eq!(
decode_string(plp_known(&utf16, 6), &md, true).await,
vec![sunk(&utf16, Some(EncodingType::Utf16))]
);
}
#[tokio::test]
async fn varchar_plp_known_length_reports_its_collation_encoding() {
let collation = SqlCollation::default();
let md = plp_metadata(
TdsDataType::BigVarChar,
PartialLengthType::BigVarChar,
Some(collation),
);
let payload = b"plain text";
assert_eq!(
decode_string(plp_known(payload, 3), &md, true).await,
vec![sunk(payload, Some(EncodingType::LcidBased(collation)))]
);
}
#[tokio::test]
async fn ntext_long_len_never_requests_a_destination() {
let md = varlen_metadata(TdsDataType::NText, 0x7FFFFFFF);
let payload: Vec<u8> = "hello world"
.encode_utf16()
.flat_map(u16::to_le_bytes)
.collect();
let mut wire = vec![16u8];
wire.extend_from_slice(&[0u8; 16]); wire.extend_from_slice(&[0u8; 8]); wire.extend_from_slice(&(payload.len() as u32).to_le_bytes());
wire.extend_from_slice(&payload);
assert_eq!(
decode_string(wire, &md, true).await,
vec![Event::OwnedString(payload)]
);
}
#[tokio::test]
async fn ntext_null_text_pointer_never_requests_a_destination() {
let md = varlen_metadata(TdsDataType::NText, 0x7FFFFFFF);
assert_eq!(decode_string(vec![0u8], &md, true).await, vec![Event::Null]);
}
#[tokio::test]
async fn string_sink_declined_falls_back_to_owned_path() {
let md = plp_metadata(TdsDataType::NVarChar, PartialLengthType::NVarChar, None);
let utf16: Vec<u8> = "abc".encode_utf16().flat_map(u16::to_le_bytes).collect();
assert_eq!(
decode_string(plp_known(&utf16, 2), &md, false).await,
vec![Event::OwnedString(utf16)]
);
}
#[tokio::test]
async fn truncated_payload_commits_the_value_as_incomplete() {
let md = plp_metadata(
TdsDataType::BigVarBinary,
PartialLengthType::BigVarBinary,
None,
);
let mut wire = 8u64.to_le_bytes().to_vec();
wire.extend_from_slice(&8u32.to_le_bytes());
wire.extend_from_slice(&[1, 2, 3]);
let mut reader = ByteReader::new(wire);
let mut writer = RecordingSink::new(true);
let err = GenericDecoder::default()
.decode_into(&mut reader, &md, 0, &mut writer)
.await;
assert!(err.is_err());
assert!(matches!(
writer.events.as_slice(),
[Event::Sunk {
complete: false,
..
}]
));
}
#[tokio::test]
async fn plp_shorter_than_declared_length_zero_fills_the_remainder() {
let md = plp_metadata(
TdsDataType::BigVarBinary,
PartialLengthType::BigVarBinary,
None,
);
let mut wire = 8u64.to_le_bytes().to_vec();
wire.extend_from_slice(&3u32.to_le_bytes());
wire.extend_from_slice(&[1, 2, 3]);
wire.extend_from_slice(&0u32.to_le_bytes());
let sink = decode_generic(wire.clone(), &md, true).await;
let owned = decode_generic(wire, &md, false).await;
assert_eq!(sink, vec![sunk(&[1, 2, 3, 0, 0, 0, 0, 0], None)]);
assert_eq!(owned, vec![Event::OwnedBytes(vec![1, 2, 3, 0, 0, 0, 0, 0])]);
}
#[tokio::test]
async fn plp_chunk_past_declared_length_is_rejected_on_both_paths() {
let md = plp_metadata(
TdsDataType::BigVarBinary,
PartialLengthType::BigVarBinary,
None,
);
let mut wire = 2u64.to_le_bytes().to_vec();
wire.extend_from_slice(&4u32.to_le_bytes());
wire.extend_from_slice(&[1, 2, 3, 4]);
wire.extend_from_slice(&0u32.to_le_bytes());
for accept in [true, false] {
let mut reader = ByteReader::new(wire.clone());
let mut writer = RecordingSink::new(accept);
let err = GenericDecoder::default()
.decode_into(&mut reader, &md, 0, &mut writer)
.await;
assert!(err.is_err(), "accept={accept} should have been rejected");
}
}
#[tokio::test]
#[should_panic(expected = "value_destination returned 1 bytes for a 2-byte value")]
async fn wrong_sized_destination_is_rejected() {
let md = plp_metadata(
TdsDataType::BigVarBinary,
PartialLengthType::BigVarBinary,
None,
);
let mut reader = ByteReader::new(plp_known(&[1, 2], 2));
let mut writer = RecordingSink::with_destination_len(1);
GenericDecoder::default()
.decode_into(&mut reader, &md, 0, &mut writer)
.await
.unwrap();
}
#[tokio::test]
async fn overlong_destination_does_not_relax_wire_declared_length() {
let mut wire = 4u32.to_le_bytes().to_vec();
wire.extend_from_slice(&[1, 2, 3, 4]);
wire.extend_from_slice(&0u32.to_le_bytes());
let mut reader = ByteReader::new(wire);
let mut destination = [MaybeUninit::uninit(); 4];
let err = GenericDecoder::read_plp_chunks_into_slice(&mut reader, &mut destination, 2)
.await
.unwrap_err();
assert!(err.to_string().contains("wire_declared_len=2"));
}
#[tokio::test]
async fn xml_json_udt_and_vector_never_request_destinations() {
for (data_type, partial_type) in [
(TdsDataType::Xml, PartialLengthType::Xml),
(TdsDataType::Json, PartialLengthType::Json),
(TdsDataType::Udt, PartialLengthType::Udt),
] {
let md = plp_metadata(data_type, partial_type, None);
let mut reader = ByteReader::new(plp_known(&[], 1));
let mut writer = RecordingSink::new(true);
let _ = GenericDecoder::default()
.decode_into(&mut reader, &md, 0, &mut writer)
.await;
assert_eq!(
writer.destination_requests, 0,
"{data_type:?} requested a destination"
);
}
let md = varlen_metadata(TdsDataType::Vector, 8);
let mut reader = ByteReader::new(Vec::new());
let mut writer = RecordingSink::new(true);
let _ = GenericDecoder::default()
.decode_into(&mut reader, &md, 0, &mut writer)
.await;
assert_eq!(writer.destination_requests, 0);
}
#[tokio::test]
async fn owned_path_matches_sink_path_byte_for_byte() {
let md = plp_metadata(
TdsDataType::BigVarBinary,
PartialLengthType::BigVarBinary,
None,
);
let payload: Vec<u8> = (0..200u8).map(|b| b.wrapping_mul(7)).collect();
let wire = plp_known(&payload, 17);
let Event::Sunk { bytes, .. } =
decode_generic(wire.clone(), &md, true).await.pop().unwrap()
else {
panic!("expected a sunk value");
};
let Event::OwnedBytes(owned) = decode_generic(wire, &md, false).await.pop().unwrap()
else {
panic!("expected an owned value");
};
assert_eq!(bytes, owned);
}
}
}