use std::convert::TryFrom;
use byteorder::{ByteOrder, LittleEndian};
use bytes::{BufMut, BytesMut};
use crate::{
error::Error,
sql_read_bytes::SqlReadBytes,
tds::{codec::guid, codec::Encode, Collation, Numeric},
ColumnData, FixedLenType, VarLenType,
};
async fn read_bytes<R>(src: &mut R, len: usize) -> crate::Result<Vec<u8>>
where
R: SqlReadBytes + Unpin,
{
let mut buf = Vec::with_capacity(len);
for _ in 0..len {
buf.push(src.read_u8().await?);
}
Ok(buf)
}
pub(crate) async fn decode<R>(src: &mut R) -> crate::Result<ColumnData<'static>>
where
R: SqlReadBytes + Unpin,
{
let total_len = src.read_u32_le().await? as usize;
if total_len == 0 {
return Ok(ColumnData::String(None));
}
if total_len < 2 {
return Err(Error::Protocol(
format!("sql_variant: invalid total length {}", total_len).into(),
));
}
let base_type = src.read_u8().await?;
let prop_bytes = src.read_u8().await? as usize;
if total_len < 2 + prop_bytes {
return Err(Error::Protocol(
format!(
"sql_variant: total length {} too small for {} property bytes",
total_len, prop_bytes
)
.into(),
));
}
let data_len = total_len - 2 - prop_bytes;
if data_len > MAX_VARIANT_PAYLOAD {
return Err(Error::Protocol(
format!("sql_variant: value length {data_len} exceeds the maximum").into(),
));
}
if let Ok(fixed) = FixedLenType::try_from(base_type) {
if prop_bytes != 0 {
return Err(Error::Protocol(
format!(
"sql_variant: fixed base type {:?} must not carry property bytes",
fixed
)
.into(),
));
}
return super::fixed_len::decode(src, &fixed).await;
}
let var = VarLenType::try_from(base_type).map_err(|_| {
Error::Protocol(format!("sql_variant: unknown base type 0x{:02x}", base_type).into())
})?;
let res = match var {
VarLenType::Guid => {
let bytes = read_bytes(src, 16).await?;
let mut data: [u8; 16] = bytes
.try_into()
.map_err(|_| Error::Protocol("sql_variant: short guid".into()))?;
guid::reorder_bytes(&mut data);
ColumnData::Guid(Some(uuid::Uuid::from_bytes(data)))
}
VarLenType::Decimaln | VarLenType::Numericn => {
let _precision = src.read_u8().await?;
let scale = src.read_u8().await?;
decode_numeric(src, data_len, scale).await?
}
VarLenType::BigChar | VarLenType::BigVarChar => {
let collation = read_collation(src).await?;
let _max_len = src.read_u16_le().await?;
let buf = read_bytes(src, data_len).await?;
let encoder = collation.encoding()?;
let s = encoder
.decode_without_bom_handling_and_without_replacement(buf.as_ref())
.ok_or_else(|| Error::Encoding("sql_variant: invalid sequence".into()))?
.to_string();
ColumnData::String(Some(s.into()))
}
VarLenType::NChar | VarLenType::NVarchar => {
let _collation = read_collation(src).await?;
let _max_len = src.read_u16_le().await?;
let buf = read_bytes(src, data_len).await?;
if buf.len() % 2 != 0 {
return Err(Error::Protocol("sql_variant: invalid nchar length".into()));
}
let buf: Vec<u16> = buf.chunks(2).map(LittleEndian::read_u16).collect();
ColumnData::String(Some(String::from_utf16(&buf)?.into()))
}
VarLenType::BigBinary | VarLenType::BigVarBin => {
let _max_len = src.read_u16_le().await?;
let buf = read_bytes(src, data_len).await?;
ColumnData::Binary(Some(buf.into()))
}
#[cfg(feature = "tds73")]
VarLenType::Daten => {
ColumnData::Date(Some(crate::tds::time::Date::decode(src).await?))
}
#[cfg(feature = "tds73")]
VarLenType::Timen => {
let scale = src.read_u8().await? as usize;
let time = crate::tds::time::Time::decode(src, scale, data_len).await?;
ColumnData::Time(Some(time))
}
#[cfg(feature = "tds73")]
VarLenType::Datetime2 => {
let scale = src.read_u8().await? as usize;
let time_len = data_len
.checked_sub(3)
.ok_or_else(|| Error::Protocol("sql_variant: datetime2 value too short".into()))?;
let dt = crate::tds::time::DateTime2::decode(src, scale, time_len).await?;
ColumnData::DateTime2(Some(dt))
}
#[cfg(feature = "tds73")]
VarLenType::DatetimeOffsetn => {
let scale = src.read_u8().await? as usize;
let time_len = data_len.checked_sub(5).ok_or_else(|| {
Error::Protocol("sql_variant: datetimeoffset value too short".into())
})?;
let time_len = u8::try_from(time_len).map_err(|_| {
Error::Protocol("sql_variant: datetimeoffset time length too large".into())
})?;
let dto = crate::tds::time::DateTimeOffset::decode(src, scale, time_len).await?;
ColumnData::DateTimeOffset(Some(dto))
}
other => {
return Err(Error::Protocol(
format!("sql_variant: unsupported base type {:?}", other).into(),
))
}
};
Ok(res)
}
async fn read_collation<R>(src: &mut R) -> crate::Result<Collation>
where
R: SqlReadBytes + Unpin,
{
let info = src.read_u32_le().await?;
let sort_id = src.read_u8().await?;
Ok(Collation::new(info, sort_id))
}
async fn decode_numeric<R>(
src: &mut R,
data_len: usize,
scale: u8,
) -> crate::Result<ColumnData<'static>>
where
R: SqlReadBytes + Unpin,
{
if data_len == 0 {
return Err(Error::Protocol("sql_variant: empty numeric value".into()));
}
if scale > 38 {
return Err(Error::Protocol(
format!("sql_variant: invalid numeric scale {scale}").into(),
));
}
let sign = match src.read_u8().await? {
0 => -1i128,
1 => 1i128,
_ => return Err(Error::Protocol("sql_variant: invalid numeric sign".into())),
};
let magnitude = read_bytes(src, data_len - 1).await?;
let value = match magnitude.len() {
4 => LittleEndian::read_u32(&magnitude) as i128,
8 => LittleEndian::read_u64(&magnitude) as i128,
12 => {
let low = LittleEndian::read_u64(&magnitude[0..8]) as i128;
let high = LittleEndian::read_u32(&magnitude[8..12]) as i128;
low + high * (1i128 << 64)
}
16 => {
let low = LittleEndian::read_u64(&magnitude[0..8]) as i128;
let high = LittleEndian::read_u64(&magnitude[8..16]) as i128;
low + high * (1i128 << 64)
}
n => {
return Err(Error::Protocol(
format!("sql_variant: invalid numeric magnitude length {}", n).into(),
))
}
};
Ok(ColumnData::Numeric(Some(Numeric::new_with_scale(
value * sign,
scale,
))))
}
const MAX_VARIANT_PAYLOAD: usize = 8000;
pub(crate) fn encode(dst: &mut BytesMut, data: ColumnData<'_>) -> crate::Result<()> {
let mut body = BytesMut::new();
let has_value = match data {
ColumnData::Bit(Some(val)) => {
body.put_u8(FixedLenType::Bit as u8);
body.put_u8(0);
body.put_u8(val as u8);
true
}
ColumnData::U8(Some(val)) => {
body.put_u8(FixedLenType::Int1 as u8);
body.put_u8(0);
body.put_u8(val);
true
}
ColumnData::I16(Some(val)) => {
body.put_u8(FixedLenType::Int2 as u8);
body.put_u8(0);
body.put_i16_le(val);
true
}
ColumnData::I32(Some(val)) => {
body.put_u8(FixedLenType::Int4 as u8);
body.put_u8(0);
body.put_i32_le(val);
true
}
ColumnData::I64(Some(val)) => {
body.put_u8(FixedLenType::Int8 as u8);
body.put_u8(0);
body.put_i64_le(val);
true
}
ColumnData::F32(Some(val)) => {
body.put_u8(FixedLenType::Float4 as u8);
body.put_u8(0);
body.put_f32_le(val);
true
}
ColumnData::F64(Some(val)) => {
body.put_u8(FixedLenType::Float8 as u8);
body.put_u8(0);
body.put_f64_le(val);
true
}
ColumnData::DateTime(Some(dt)) => {
body.put_u8(FixedLenType::Datetime as u8);
body.put_u8(0);
dt.encode(&mut body)?;
true
}
ColumnData::SmallDateTime(Some(dt)) => {
body.put_u8(FixedLenType::Datetime4 as u8);
body.put_u8(0);
dt.encode(&mut body)?;
true
}
ColumnData::Guid(Some(uuid)) => {
body.put_u8(VarLenType::Guid as u8);
body.put_u8(0);
let mut bytes = *uuid.as_bytes();
guid::reorder_bytes(&mut bytes);
body.extend_from_slice(&bytes);
true
}
ColumnData::Numeric(Some(num)) => {
body.put_u8(VarLenType::Numericn as u8);
body.put_u8(2);
body.put_u8(num.precision());
body.put_u8(num.scale());
let mut tmp = BytesMut::new();
num.encode(&mut tmp)?;
body.extend_from_slice(&tmp[1..]);
true
}
ColumnData::String(Some(ref s)) => {
let utf16: Vec<u8> = s.encode_utf16().flat_map(|c| c.to_le_bytes()).collect();
if utf16.len() > MAX_VARIANT_PAYLOAD {
return Err(Error::Conversion(
format!(
"sql_variant: string of {} bytes exceeds the {} byte limit",
utf16.len(),
MAX_VARIANT_PAYLOAD
)
.into(),
));
}
body.put_u8(VarLenType::NVarchar as u8);
body.put_u8(7);
body.extend_from_slice(&[0u8; 5]);
body.put_u16_le(MAX_VARIANT_PAYLOAD as u16);
body.extend_from_slice(&utf16);
true
}
ColumnData::Binary(Some(ref bytes)) => {
if bytes.len() > MAX_VARIANT_PAYLOAD {
return Err(Error::Conversion(
format!(
"sql_variant: binary of {} bytes exceeds the {} byte limit",
bytes.len(),
MAX_VARIANT_PAYLOAD
)
.into(),
));
}
body.put_u8(VarLenType::BigVarBin as u8);
body.put_u8(2);
body.put_u16_le(MAX_VARIANT_PAYLOAD as u16);
body.extend_from_slice(bytes);
true
}
#[cfg(feature = "tds73")]
ColumnData::Date(Some(date)) => {
body.put_u8(VarLenType::Daten as u8);
body.put_u8(0);
date.encode(&mut body)?;
true
}
#[cfg(feature = "tds73")]
ColumnData::Time(Some(time)) => {
body.put_u8(VarLenType::Timen as u8);
body.put_u8(1);
body.put_u8(time.scale());
time.encode(&mut body)?;
true
}
#[cfg(feature = "tds73")]
ColumnData::DateTime2(Some(dt)) => {
body.put_u8(VarLenType::Datetime2 as u8);
body.put_u8(1);
body.put_u8(dt.time().scale());
dt.encode(&mut body)?;
true
}
#[cfg(feature = "tds73")]
ColumnData::DateTimeOffset(Some(dto)) => {
body.put_u8(VarLenType::DatetimeOffsetn as u8);
body.put_u8(1);
body.put_u8(dto.datetime2().time().scale());
dto.encode(&mut body)?;
true
}
ColumnData::Xml(Some(_)) => {
return Err(Error::Conversion(
"sql_variant: xml is not a valid sql_variant base type".into(),
));
}
_ => false,
};
if has_value {
dst.put_u32_le(body.len() as u32);
dst.extend_from_slice(&body);
} else {
dst.put_u32_le(0);
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
use crate::sql_read_bytes::test_utils::IntoSqlReadBytes;
use bytes::{BufMut, BytesMut};
fn variant_reader(payload: &[u8]) -> impl SqlReadBytes + Unpin {
let mut buf = BytesMut::new();
buf.put_u32_le(payload.len() as u32);
buf.extend_from_slice(payload);
buf.into_sql_read_bytes()
}
#[tokio::test]
async fn decode_null() {
let mut buf = BytesMut::new();
buf.put_u32_le(0);
let data = decode(&mut buf.into_sql_read_bytes()).await.unwrap();
assert_eq!(data, ColumnData::String(None));
}
#[tokio::test]
async fn decode_int() {
let mut payload = vec![FixedLenType::Int4 as u8, 0];
payload.extend_from_slice(&42i32.to_le_bytes());
let data = decode(&mut variant_reader(&payload)).await.unwrap();
assert_eq!(data, ColumnData::I32(Some(42)));
}
#[tokio::test]
async fn decode_bigint() {
let mut payload = vec![FixedLenType::Int8 as u8, 0];
payload.extend_from_slice(&(-7i64).to_le_bytes());
let data = decode(&mut variant_reader(&payload)).await.unwrap();
assert_eq!(data, ColumnData::I64(Some(-7)));
}
#[tokio::test]
async fn decode_bit() {
let payload = vec![FixedLenType::Bit as u8, 0, 1];
let data = decode(&mut variant_reader(&payload)).await.unwrap();
assert_eq!(data, ColumnData::Bit(Some(true)));
}
#[tokio::test]
async fn decode_nvarchar() {
let text = "hi€";
let utf16: Vec<u8> = text.encode_utf16().flat_map(|c| c.to_le_bytes()).collect();
let mut payload = vec![VarLenType::NVarchar as u8, 7];
payload.extend_from_slice(&0u32.to_le_bytes()); payload.push(0); payload.extend_from_slice(&40u16.to_le_bytes()); payload.extend_from_slice(&utf16);
let data = decode(&mut variant_reader(&payload)).await.unwrap();
assert_eq!(data, ColumnData::String(Some(text.into())));
}
#[tokio::test]
async fn decode_varchar() {
let mut payload = vec![VarLenType::BigVarChar as u8, 7];
payload.extend_from_slice(&13632521u32.to_le_bytes());
payload.push(52);
payload.extend_from_slice(&40u16.to_le_bytes());
payload.extend_from_slice(b"abc");
let data = decode(&mut variant_reader(&payload)).await.unwrap();
assert_eq!(data, ColumnData::String(Some("abc".into())));
}
#[tokio::test]
async fn decode_binary() {
let mut payload = vec![VarLenType::BigVarBin as u8, 2];
payload.extend_from_slice(&40u16.to_le_bytes());
payload.extend_from_slice(&[1u8, 2, 3, 4]);
let data = decode(&mut variant_reader(&payload)).await.unwrap();
assert_eq!(data, ColumnData::Binary(Some(vec![1, 2, 3, 4].into())));
}
#[tokio::test]
async fn decode_guid() {
let uuid = uuid::Uuid::from_u128(0x0102030405060708090a0b0c0d0e0f10);
let mut wire = *uuid.as_bytes();
guid::reorder_bytes(&mut wire);
let mut payload = vec![VarLenType::Guid as u8, 0];
payload.extend_from_slice(&wire);
let data = decode(&mut variant_reader(&payload)).await.unwrap();
assert_eq!(data, ColumnData::Guid(Some(uuid)));
}
#[tokio::test]
async fn decode_numeric_value() {
let mut payload = vec![VarLenType::Numericn as u8, 2, 18, 2];
payload.push(1); payload.extend_from_slice(&123u32.to_le_bytes());
let data = decode(&mut variant_reader(&payload)).await.unwrap();
assert_eq!(
data,
ColumnData::Numeric(Some(Numeric::new_with_scale(123, 2)))
);
}
#[cfg(feature = "tds73")]
#[tokio::test]
async fn decode_date_value() {
use crate::tds::time::Date;
let mut payload = vec![VarLenType::Daten as u8, 0];
payload.extend_from_slice(&730119u32.to_le_bytes()[..3]);
let data = decode(&mut variant_reader(&payload)).await.unwrap();
assert_eq!(data, ColumnData::Date(Some(Date::new(730119))));
}
async fn round_trip(value: ColumnData<'static>) {
let mut buf = BytesMut::new();
encode(&mut buf, value.clone()).expect("encode must succeed");
let reader = &mut buf.into_sql_read_bytes();
let decoded = decode(reader).await.expect("decode must succeed");
assert_eq!(decoded, value);
reader
.read_u8()
.await
.expect_err("decode must consume the entire buffer");
}
#[tokio::test]
async fn round_trip_bit() {
round_trip(ColumnData::Bit(Some(true))).await;
round_trip(ColumnData::Bit(Some(false))).await;
}
#[tokio::test]
async fn round_trip_integers() {
round_trip(ColumnData::U8(Some(200))).await;
round_trip(ColumnData::I16(Some(-1234))).await;
round_trip(ColumnData::I32(Some(42))).await;
round_trip(ColumnData::I64(Some(-9_000_000_000))).await;
}
#[tokio::test]
async fn round_trip_floats() {
round_trip(ColumnData::F32(Some(1.5))).await;
round_trip(ColumnData::F64(Some(-2.5))).await;
}
#[tokio::test]
async fn round_trip_guid() {
let uuid = uuid::Uuid::from_u128(0x0102030405060708090a0b0c0d0e0f10);
round_trip(ColumnData::Guid(Some(uuid))).await;
}
#[tokio::test]
async fn round_trip_numeric() {
round_trip(ColumnData::Numeric(Some(Numeric::new_with_scale(123, 2)))).await;
round_trip(ColumnData::Numeric(Some(Numeric::new_with_scale(-4567, 4)))).await;
round_trip(ColumnData::Numeric(Some(Numeric::new_with_scale(
10i128.pow(30),
0,
))))
.await;
}
#[tokio::test]
async fn round_trip_string() {
round_trip(ColumnData::String(Some("hello€".into()))).await;
round_trip(ColumnData::String(Some("".into()))).await;
}
#[tokio::test]
async fn round_trip_binary() {
round_trip(ColumnData::Binary(Some(vec![1u8, 2, 3, 4, 5].into()))).await;
round_trip(ColumnData::Binary(Some(vec![].into()))).await;
}
#[tokio::test]
async fn round_trip_datetime() {
use crate::tds::time::{DateTime, SmallDateTime};
round_trip(ColumnData::DateTime(Some(DateTime::new(200, 3000)))).await;
round_trip(ColumnData::SmallDateTime(Some(SmallDateTime::new(
200, 3000,
))))
.await;
}
#[cfg(feature = "tds73")]
#[tokio::test]
async fn round_trip_temporal_tds73() {
use crate::tds::time::{Date, DateTime2, DateTimeOffset, Time};
round_trip(ColumnData::Date(Some(Date::new(730119)))).await;
round_trip(ColumnData::Time(Some(Time::new(222, 7)))).await;
round_trip(ColumnData::DateTime2(Some(DateTime2::new(
Date::new(55),
Time::new(222, 7),
))))
.await;
round_trip(ColumnData::DateTimeOffset(Some(DateTimeOffset::new(
DateTime2::new(Date::new(55), Time::new(222, 7)),
-8,
))))
.await;
}
#[tokio::test]
async fn round_trip_null() {
let mut buf = BytesMut::new();
encode(&mut buf, ColumnData::I32(None)).expect("encode must succeed");
let decoded = decode(&mut buf.into_sql_read_bytes()).await.unwrap();
assert_eq!(decoded, ColumnData::String(None));
}
#[tokio::test]
async fn xml_is_rejected() {
use crate::xml::XmlData;
use std::borrow::Cow;
let mut buf = BytesMut::new();
let err = encode(
&mut buf,
ColumnData::Xml(Some(Cow::Owned(XmlData::new("<a/>")))),
)
.expect_err("xml must not encode as a sql_variant");
assert!(matches!(err, Error::Conversion(_)));
}
}