diff --git a/mssql-tds/src/datatypes/decoder.rs b/mssql-tds/src/datatypes/decoder.rs index 94fd6d2e..b01214f6 100644 --- a/mssql-tds/src/datatypes/decoder.rs +++ b/mssql-tds/src/datatypes/decoder.rs @@ -28,7 +28,7 @@ use crate::{ }; use crate::{query::metadata::ColumnMetadata, token::tokens::SqlCollation}; -use super::row_writer::{RowWriter, write_column_value}; +use super::row_writer::{CaptureWriter, RowWriter, write_column_value}; /// Reads an encrypted column's cipher bytes from the wire and turns them back /// into a plaintext [`ColumnValues`]. @@ -846,26 +846,6 @@ impl GenericDecoder { Ok(ColumnValues::DateTimeOffset(datetime_offset)) } - async fn read_intn(&self, reader: &mut T, byte_len: u8) -> TdsResult - where - T: TdsPacketReader + Send + Sync, - { - let value: ColumnValues = match byte_len { - 1 => ColumnValues::TinyInt(reader.read_byte().await?), // Some(reader.read_byte().await? as i64), - 2 => ColumnValues::SmallInt(reader.read_int16().await?), // Some(reader.read_int16().await? as i64), - 4 => ColumnValues::Int(reader.read_int32().await?), - 8 => ColumnValues::BigInt(reader.read_int64().await?), - 0 => ColumnValues::Null, - _ => { - return Err(crate::error::Error::from(Error::new( - std::io::ErrorKind::InvalidData, - "Invalid IntN length", - ))); - } - }; - Ok(value) - } - async fn read_money4(&self, reader: &mut T) -> TdsResult where T: TdsPacketReader + Send + Sync, @@ -1123,9 +1103,17 @@ impl GenericDecoder { } /// Decodes a column value from the wire and writes it directly into a - /// [`RowWriter`], bypassing the intermediate `ColumnValues` enum for - /// common types. Rare types (XML, JSON, Vector, Image, UDT, SsVariant) - /// fall back to `decode()` + `write_column_value()`. + /// [`RowWriter`], bypassing the intermediate `ColumnValues` enum. + /// + /// This is the single per-type switch in the decoder: `SqlTypeDecode::decode` + /// is defined as this method plus a [`CaptureWriter`], so there is no second + /// switch that could drift out of sync with it. A few rare shapes (Vector, + /// SsVariant) still materialize a `ColumnValues` internally because their + /// sub-decoders produce one, but they are dispatched from here, not from a + /// parallel match. + /// + /// Every arm must either write exactly one value or return an error; + /// returning `Ok` without writing is reported by [`CaptureWriter::into_value`]. pub(crate) async fn decode_into( &self, reader: &mut T, @@ -1397,325 +1385,65 @@ impl GenericDecoder { } } - // === Fallback: rare types go through decode() → write_column_value() === - _ => { - let value = self.decode(reader, metadata).await?; - write_column_value(writer, col, value); - } - } - Ok(()) - } -} - -impl SqlTypeDecode for GenericDecoder { - async fn decode(&self, reader: &mut T, metadata: &ColumnMetadata) -> TdsResult - where - T: TdsPacketReader + Send + Sync, - { - let result = match metadata.data_type { - TdsDataType::Int1 => { - let value = reader.read_byte().await?; - ColumnValues::from(value) - } - TdsDataType::Int2 => { - let value = reader.read_int16().await?; - ColumnValues::SmallInt(value) - } - TdsDataType::Int4 => { - let value = reader.read_int32().await?; - ColumnValues::from(value) - } - TdsDataType::Int8 => { - let value = reader.read_int64().await?; - ColumnValues::BigInt(value) - } - TdsDataType::Flt4 => { - let value = reader.read_float32().await?; - ColumnValues::Real(value) - } - TdsDataType::Flt8 => { - let value = reader.read_float64().await?; - ColumnValues::Float(value) - } - TdsDataType::Money4 => ColumnValues::SmallMoney(self.read_money4(reader).await?), - TdsDataType::Money => ColumnValues::Money(self.read_money8(reader).await?), - TdsDataType::MoneyN => { - let byte_len = reader.read_byte().await?; - 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 = reader.read_byte().await?; - 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 = reader.read_byte().await?; - self.read_intn(reader, byte_len).await? - } - TdsDataType::BigBinary => { - let length = reader.read_uint16().await?; - // 0xFFFF is the USHORTLEN NULL marker (CHARBIN_NULL). - 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 = reader.read_uint16().await?; - // 0xFFFF is the USHORTLEN NULL marker (CHARBIN_NULL). - 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) - } - } - } + // === Rare shapes: boxed so their locals stay out of the hot path's + // future, but still handled here rather than in a second switch. === 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, - } + Box::pin(Self::read_plp_into( + reader, + metadata, + "XML", + col, + writer, + |bytes| ColumnValues::Xml(SqlXml { bytes }), + )) + .await?; } 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 = reader.read_byte().await?; - if byte_len > 0 { - let value = reader.read_byte().await?; - ColumnValues::Bit(value == 1) - } else { - ColumnValues::Null - } - } - TdsDataType::Guid => { - let length = reader.read_byte().await?; - Self::read_guid(reader, length).await? - } - TdsDataType::FltN => { - // This is variable length float, hence the length needs to be read first - let length = reader.read_byte().await?; - if length == 0 { - return Ok(ColumnValues::Null); - } - if length == 4 { - let value = reader.read_float32().await?; - ColumnValues::Real(value) - } else { - let value = reader.read_float64().await?; - ColumnValues::Float(value) - } - } - TdsDataType::DateTimeN => { - let length = reader.read_byte().await?; - // If length is 0, then it is NULL - if length == 0 { - return Ok(ColumnValues::Null); - } else if length == 4 { - // SmallDateTime - let smalldatetime = self.read_small_datetime(reader).await?; - return Ok(ColumnValues::SmallDateTime(smalldatetime)); - } else { - // DateTime - return Ok(ColumnValues::DateTime(self.read_datetime(reader).await?)); - } - } - TdsDataType::DateN => { - let length = reader.read_byte().await?; - return Self::read_daten(reader, length).await; + Box::pin(Self::read_plp_into( + reader, + metadata, + "JSON", + col, + writer, + |bytes| ColumnValues::Json(SqlJson::new(bytes)), + )) + .await?; } - TdsDataType::TimeN => { - let length = reader.read_byte().await?; - 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::Udt => { + Box::pin(Self::read_plp_into( + reader, + metadata, + "UDT", + col, + writer, + ColumnValues::Bytes, + )) + .await?; } - TdsDataType::DateTime2N => { - let length = reader.read_byte().await?; - 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 = reader.read_byte().await?; - 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 = reader.read_byte().await? as usize; - - let length = if text_ptr_len > 0 { - const TIMESTAMP_BYTE_COUNT: usize = 8; - reader.skip_bytes(text_ptr_len).await?; - reader.skip_bytes(TIMESTAMP_BYTE_COUNT).await?; - reader.read_uint32().await? as usize - } 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) - } + Box::pin(Self::read_image_into(reader, col, writer)).await?; } - 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 = reader.read_uint16().await?; - let timepart = reader.read_uint16().await?; - ColumnValues::SmallDateTime(SqlSmallDateTime { - days: daypart, - time: timepart, - }) + TdsDataType::Vector => { + let value = Box::pin(self.decode_vector(reader, metadata)).await?; + write_column_value(writer, col, value); } - 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::SsVariant => { + let value = Box::pin(self.read_sql_variant(reader)).await?; + write_column_value(writer, col, value); } - TdsDataType::Numeric => { + + // Fixed-length Decimal/Numeric have no reader; the server sends + // DecimalN/NumericN instead. + TdsDataType::Decimal | TdsDataType::Numeric => { return Err(crate::error::Error::UnimplementedFeature { - feature: "Fixed-length Numeric type".to_string(), + feature: format!("Fixed-length {:?} type", metadata.data_type), context: format!( - "Data type {:?} (0x{:02X}) is not implemented. Use NumericN instead.", - metadata.data_type, metadata.data_type as u8 + "Data type {:?} (0x{:02X}) is not implemented. Use {:?}N instead.", + metadata.data_type, metadata.data_type as u8, metadata.data_type ), }); } + _ => { return Err(crate::error::Error::UnimplementedFeature { feature: format!("Data type {:?}", metadata.data_type), @@ -1725,8 +1453,89 @@ impl SqlTypeDecode for GenericDecoder { ), }); } - }; - Ok(result) + } + Ok(()) + } + + /// Reads a PLP column that is only ever sent partially-length-prefixed. + async fn read_plp_into( + reader: &mut T, + metadata: &ColumnMetadata, + label: &str, + col: usize, + writer: &mut W, + wrap: impl FnOnce(Vec) -> ColumnValues, + ) -> TdsResult<()> + where + T: TdsPacketReader + Send + Sync, + W: RowWriter + ?Sized, + { + if !metadata.is_plp() { + return Err(crate::error::Error::ProtocolError(format!( + "{label} column metadata is not partially-length-prefixed" + ))); + } + match Self::read_plp_bytes(reader).await? { + Some(bytes) => write_column_value(writer, col, wrap(bytes)), + None => writer.write_null(col), + } + Ok(()) + } + + /// Reads a LONGLEN `image` payload. A zero-length payload is NULL here, + /// unlike `text`, which reports it as an empty string. + async fn read_image_into(reader: &mut T, col: usize, writer: &mut W) -> TdsResult<()> + where + T: TdsPacketReader + Send + Sync, + W: RowWriter + ?Sized, + { + match Self::read_long_len_bytes(reader).await? { + Some(bytes) if !bytes.is_empty() => writer.write_bytes(col, bytes), + _ => writer.write_null(col), + } + Ok(()) + } + + /// Reads the LONGLEN wire form shared by `text`, `ntext` and `image`: + /// a text-pointer length, the pointer, an 8-byte timestamp, then a `u32` + /// payload length. `None` means the text pointer was absent, i.e. NULL. + /// + /// Callers differ on what an empty payload means, so that decision is left + /// to them. + async fn read_long_len_bytes(reader: &mut T) -> TdsResult>> + where + T: TdsPacketReader + Send + Sync, + { + let text_ptr_len = reader.read_byte().await? as usize; + if text_ptr_len == 0 { + return Ok(None); + } + + const TIMESTAMP_BYTE_COUNT: usize = 8; + reader.skip_bytes(text_ptr_len).await?; + reader.skip_bytes(TIMESTAMP_BYTE_COUNT).await?; + let length = reader.read_uint32().await? as usize; + + if length > MAX_ALLOC_SIZE { + return Err(crate::error::Error::ProtocolError(format!( + "LOB data length {length} exceeds maximum allowed size of {MAX_ALLOC_SIZE} bytes" + ))); + } + + let mut buffer = vec![0u8; length]; + reader.read_bytes(&mut buffer).await?; + Ok(Some(buffer)) + } +} + +impl SqlTypeDecode for GenericDecoder { + async fn decode(&self, reader: &mut T, metadata: &ColumnMetadata) -> TdsResult + where + T: TdsPacketReader + Send + Sync, + { + let mut capture = CaptureWriter::default(); + self.decode_into(reader, metadata, 0, &mut capture).await?; + capture.into_value() } } @@ -1757,32 +1566,10 @@ impl StringDecoder { None => writer.write_null(col), } } else if Self::is_long_len_type(metadata.data_type) { - let text_ptr_len = reader.read_byte().await? as usize; - - 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 = reader.read_uint32().await? 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" - ))); + match GenericDecoder::read_long_len_bytes(reader).await? { + Some(bytes) => writer.write_string(col, SqlString::new(bytes, encoding_type)), + None => writer.write_null(col), } - - 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) - }; - writer.write_string(col, sql_string); } else { let length = reader.read_uint16().await? as usize; if length == 0xFFFF { @@ -1802,82 +1589,10 @@ impl SqlTypeDecode for StringDecoder { where T: TdsPacketReader + Send + Sync, { - let encoding_type = get_encoding_type(metadata); - - // If Plp Column. (BIGVARCHARTYPE, BIGVARBINARYTYPE, NVARCHARTYPE with md.length == ushort.max) - 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) { - // Legacy LOB types (TEXT/NTEXT/IMAGE) reading implementation - // - // WIRE FORMAT (from .NET TdsParser.cs:6517-6600): - // 1. textptr_len (1 byte): Length of text pointer - // - 0x00 = NULL value - // - 0x10 (16) = Valid pointer (typical) - // 2. textptr (textptr_len bytes): Text pointer (usually 16 bytes) - // - Server-managed pointer, client treats as opaque - // 3. timestamp (8 bytes): Row timestamp - // - Used for optimistic concurrency - // 4. data_length (4 bytes, uint32): Actual data length in bytes - // - For NTEXT: byte count (divide by 2 for char count) - // - For TEXT: byte count in the collation's encoding - // 5. data (data_length bytes): The actual string data - // - For NTEXT: UTF-16LE encoded - // - For TEXT: encoded per collation (LCID-based) - // - // CURRENT IMPLEMENTATION STATUS: - // Reads textptr_len (1 byte) - // Skips textptr (16 bytes) and timestamp (8 bytes) - // Reads data_length (4 bytes, uint32) - // Allocates buffer and reads data - // Creates SqlString with appropriate encoding type - // NULL handling works (textptr_len = 0) - // LCID-based decoding implemented (see sql_string.rs) - let text_ptr_len = reader.read_byte().await? as usize; - - let length = if text_ptr_len > 0 { - const TIMESTAMP_BYTE_COUNT: usize = 8; - reader.skip_bytes(text_ptr_len).await?; - reader.skip_bytes(TIMESTAMP_BYTE_COUNT).await?; - reader.read_uint32().await? as usize - } else { - // text_ptr_len == 0 means NULL value - return Ok(ColumnValues::Null); - }; - - // Empty string (length == 0 but textptr_len > 0) is valid - return empty string, not 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 { - // Create empty SqlString with appropriate encoding - 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 = reader.read_uint16().await? 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)) - } - } + let mut capture = CaptureWriter::default(); + self.decode_string_into(reader, metadata, 0, &mut capture) + .await?; + capture.into_value() } } @@ -4706,5 +4421,150 @@ mod test { let val = assert_decode_equivalence(buf, &md).await; assert!(matches!(val, ColumnValues::DateTimeOffset(_))); } + + // The arms below used to be reachable only through `decode`; `decode_into` + // delegated to it. Now that the delegation runs the other way, each is + // checked through both entry points. + + #[tokio::test] + async fn decode_into_xml_plp() { + let md = plp_metadata(TdsDataType::Xml, PartialLengthType::Xml, None); + let val = assert_decode_equivalence(plp_wire(b""), &md).await; + match val { + ColumnValues::Xml(x) => assert_eq!(x.bytes, b""), + other => panic!("unexpected: {other:?}"), + } + } + + #[tokio::test] + async fn decode_into_json_plp() { + let md = plp_metadata(TdsDataType::Json, PartialLengthType::Json, None); + let val = assert_decode_equivalence(plp_wire(b"{}"), &md).await; + assert!(matches!(val, ColumnValues::Json(_))); + } + + #[tokio::test] + async fn decode_into_udt_plp() { + let md = plp_metadata(TdsDataType::Udt, PartialLengthType::Udt, None); + let val = assert_decode_equivalence(plp_wire(&[1, 2, 3]), &md).await; + assert_eq!(val, ColumnValues::Bytes(vec![1, 2, 3])); + } + + #[tokio::test] + async fn decode_into_xml_plp_null() { + let md = plp_metadata(TdsDataType::Xml, PartialLengthType::Xml, None); + let buf = 0xFFFFFFFFFFFFFFFFu64.to_le_bytes().to_vec(); + let val = assert_decode_equivalence(buf, &md).await; + assert_eq!(val, ColumnValues::Null); + } + + #[tokio::test] + async fn decode_into_json_rejects_non_plp_metadata() { + let md = varlen_metadata(TdsDataType::Json, 100); + let decoder = GenericDecoder::default(); + let mut reader = ByteReader::new(plp_wire(b"{}")); + let mut writer = DefaultRowWriter::new(1); + let err = decoder + .decode_into(&mut reader, &md, 0, &mut writer) + .await + .unwrap_err(); + assert!( + err.to_string() + .contains("JSON column metadata is not partially-length-prefixed"), + "unexpected: {err}" + ); + } + + /// LONGLEN payload: text pointer length, pointer, timestamp, then u32 length. + fn long_len_wire(payload: &[u8]) -> Vec { + let mut buf = vec![16u8]; + buf.extend_from_slice(&[0u8; 16]); + buf.extend_from_slice(&[0u8; 8]); + buf.extend_from_slice(&(payload.len() as u32).to_le_bytes()); + buf.extend_from_slice(payload); + buf + } + + #[tokio::test] + async fn decode_into_image_value() { + let md = varlen_metadata(TdsDataType::Image, 0x7FFFFFFF); + let val = assert_decode_equivalence(long_len_wire(&[7, 8, 9]), &md).await; + assert_eq!(val, ColumnValues::Bytes(vec![7, 8, 9])); + } + + #[tokio::test] + async fn decode_into_image_null_text_pointer() { + let md = varlen_metadata(TdsDataType::Image, 0x7FFFFFFF); + let val = assert_decode_equivalence(vec![0u8], &md).await; + assert_eq!(val, ColumnValues::Null); + } + + /// `image` reports a zero-length payload as NULL, unlike `text`. + #[tokio::test] + async fn decode_into_image_empty_payload_is_null() { + let md = varlen_metadata(TdsDataType::Image, 0x7FFFFFFF); + let val = assert_decode_equivalence(long_len_wire(&[]), &md).await; + assert_eq!(val, ColumnValues::Null); + } + + #[tokio::test] + async fn decode_into_ssvariant() { + let md = varlen_metadata(TdsDataType::SsVariant, 8); + // total length 6, Int4 base type, 0 prop bytes, payload + let mut buf = 6u32.to_le_bytes().to_vec(); + buf.push(TdsDataType::Int4 as u8); + buf.push(0); + buf.extend_from_slice(&99i32.to_le_bytes()); + let val = assert_decode_equivalence(buf, &md).await; + assert_eq!(val, ColumnValues::Int(99)); + } + + #[tokio::test] + async fn decode_into_fixed_decimal_is_unimplemented() { + let md = varlen_metadata(TdsDataType::Decimal, 9); + let decoder = GenericDecoder::default(); + let mut reader = ByteReader::new(vec![0u8; 16]); + let mut writer = DefaultRowWriter::new(1); + let err = decoder + .decode_into(&mut reader, &md, 0, &mut writer) + .await + .unwrap_err(); + assert!( + err.to_string().contains("Use DecimalN instead"), + "unexpected: {err}" + ); + } + + #[tokio::test] + async fn decode_into_fixed_numeric_is_unimplemented() { + let md = varlen_metadata(TdsDataType::Numeric, 9); + let decoder = GenericDecoder::default(); + let mut reader = ByteReader::new(vec![0u8; 16]); + let mut writer = DefaultRowWriter::new(1); + let err = decoder + .decode_into(&mut reader, &md, 0, &mut writer) + .await + .unwrap_err(); + assert!( + err.to_string().contains("Use NumericN instead"), + "unexpected: {err}" + ); + } + + #[tokio::test] + async fn decode_into_unknown_type_is_unimplemented() { + let md = varlen_metadata(TdsDataType::Void, 0); + let decoder = GenericDecoder::default(); + let mut reader = ByteReader::new(vec![0u8; 8]); + let mut writer = DefaultRowWriter::new(1); + let err = decoder + .decode_into(&mut reader, &md, 0, &mut writer) + .await + .unwrap_err(); + assert!( + err.to_string().contains("not yet supported in the decoder"), + "unexpected: {err}" + ); + } } } diff --git a/mssql-tds/src/datatypes/row_writer.rs b/mssql-tds/src/datatypes/row_writer.rs index 53c146fd..a690d323 100644 --- a/mssql-tds/src/datatypes/row_writer.rs +++ b/mssql-tds/src/datatypes/row_writer.rs @@ -1,6 +1,7 @@ // Copyright (c) Microsoft Corporation. // Licensed under the MIT License. +use crate::core::TdsResult; use crate::datatypes::column_values::{ ColumnValues, SqlDate, SqlDateTime, SqlDateTime2, SqlDateTimeOffset, SqlMoney, SqlSmallDateTime, SqlSmallMoney, SqlTime, SqlXml, @@ -9,6 +10,7 @@ use crate::datatypes::decoder::DecimalParts; use crate::datatypes::sql_json::SqlJson; use crate::datatypes::sql_string::SqlString; use crate::datatypes::sql_vector::SqlVector; +use crate::error::Error; use uuid::Uuid; /// Pluggable decode sink for TDS row data. @@ -260,6 +262,123 @@ pub fn write_column_value(writer: &mut W, col: usize, val } } +/// Single-slot writer that captures the value decoded for one column. +/// +/// This is the inverse of [`write_column_value`], and it is what lets +/// `SqlTypeDecode::decode` reuse the `RowWriter` decode path instead of carrying +/// a second copy of the per-type switch that could silently drift from it. +#[derive(Default)] +pub(crate) struct CaptureWriter { + value: Option, +} + +impl CaptureWriter { + /// Records the decoded value for the single captured column. + /// + /// A second write means a `decode_into` arm produced two values for one + /// column, so it is caught in debug builds rather than silently keeping + /// the last one. + fn set(&mut self, value: ColumnValues) { + debug_assert!( + self.value.is_none(), + "decode_into wrote twice for one column: {:?} then {value:?}", + self.value + ); + self.value = Some(value); + } + + /// Takes the captured value. + /// + /// A SQL `NULL` is written explicitly via [`RowWriter::write_null`], so an + /// empty slot means a `decode_into` arm returned `Ok` without writing + /// anything. That is a decoder bug, and reporting it as `NULL` would be + /// exactly the silent wrong-data failure this convergence exists to + /// prevent, so it is an error instead. + pub(crate) fn into_value(self) -> TdsResult { + self.value.ok_or_else(|| { + Error::ImplementationError( + "decode_into returned without writing a value for the column".to_string(), + ) + }) + } +} + +impl RowWriter for CaptureWriter { + fn write_null(&mut self, _col: usize) { + self.set(ColumnValues::Null); + } + fn write_bool(&mut self, _col: usize, val: bool) { + self.set(ColumnValues::Bit(val)); + } + fn write_u8(&mut self, _col: usize, val: u8) { + self.set(ColumnValues::TinyInt(val)); + } + fn write_i16(&mut self, _col: usize, val: i16) { + self.set(ColumnValues::SmallInt(val)); + } + fn write_i32(&mut self, _col: usize, val: i32) { + self.set(ColumnValues::Int(val)); + } + fn write_i64(&mut self, _col: usize, val: i64) { + self.set(ColumnValues::BigInt(val)); + } + fn write_f32(&mut self, _col: usize, val: f32) { + self.set(ColumnValues::Real(val)); + } + fn write_f64(&mut self, _col: usize, val: f64) { + self.set(ColumnValues::Float(val)); + } + fn write_string(&mut self, _col: usize, val: SqlString) { + self.set(ColumnValues::String(val)); + } + fn write_bytes(&mut self, _col: usize, val: Vec) { + self.set(ColumnValues::Bytes(val)); + } + fn write_decimal(&mut self, _col: usize, val: DecimalParts) { + self.set(ColumnValues::Decimal(val)); + } + fn write_numeric(&mut self, _col: usize, val: DecimalParts) { + self.set(ColumnValues::Numeric(val)); + } + fn write_date(&mut self, _col: usize, val: SqlDate) { + self.set(ColumnValues::Date(val)); + } + fn write_time(&mut self, _col: usize, val: SqlTime) { + self.set(ColumnValues::Time(val)); + } + fn write_datetime(&mut self, _col: usize, val: SqlDateTime) { + self.set(ColumnValues::DateTime(val)); + } + fn write_smalldatetime(&mut self, _col: usize, val: SqlSmallDateTime) { + self.set(ColumnValues::SmallDateTime(val)); + } + fn write_datetime2(&mut self, _col: usize, val: SqlDateTime2) { + self.set(ColumnValues::DateTime2(val)); + } + fn write_datetimeoffset(&mut self, _col: usize, val: SqlDateTimeOffset) { + self.set(ColumnValues::DateTimeOffset(val)); + } + fn write_money(&mut self, _col: usize, val: SqlMoney) { + self.set(ColumnValues::Money(val)); + } + fn write_smallmoney(&mut self, _col: usize, val: SqlSmallMoney) { + self.set(ColumnValues::SmallMoney(val)); + } + fn write_uuid(&mut self, _col: usize, val: Uuid) { + self.set(ColumnValues::Uuid(val)); + } + fn write_xml(&mut self, _col: usize, val: SqlXml) { + self.set(ColumnValues::Xml(val)); + } + fn write_json(&mut self, _col: usize, val: SqlJson) { + self.set(ColumnValues::Json(val)); + } + fn write_vector(&mut self, _col: usize, val: SqlVector) { + self.set(ColumnValues::Vector(val)); + } + fn end_row(&mut self) {} +} + #[cfg(test)] mod tests { use super::*; @@ -400,4 +519,29 @@ mod tests { assert_eq!(row[6], ColumnValues::Bit(false)); assert_eq!(row[7], ColumnValues::Null); } + + #[test] + fn capture_writer_returns_written_value() { + let mut capture = CaptureWriter::default(); + capture.write_i32(0, 42); + assert_eq!(capture.into_value().unwrap(), ColumnValues::Int(42)); + } + + #[test] + fn capture_writer_distinguishes_written_null_from_no_write() { + let mut capture = CaptureWriter::default(); + capture.write_null(0); + assert_eq!(capture.into_value().unwrap(), ColumnValues::Null); + } + + #[test] + fn capture_writer_errors_when_nothing_was_written() { + let err = CaptureWriter::default() + .into_value() + .expect_err("an unwritten column must not silently decode as NULL"); + assert!( + matches!(err, Error::ImplementationError(ref msg) if msg.contains("without writing")), + "unexpected error: {err:?}" + ); + } }