use std::io::Error;
use async_trait::async_trait;
use tracing::{debug, trace};
use super::super::tokens::{RowToken, Tokens};
use super::common::TokenParser;
use crate::{core::TdsResult, io::packet_reader::TdsPacketReader};
use crate::{
datatypes::{
column_values::ColumnValues,
decoder::{SqlTypeDecode, decrypt_encrypted_column},
},
io::token_stream::ParserContext,
};
pub(crate) struct RowTokenParser<T: SqlTypeDecode> {
decoder: T,
}
impl<T: SqlTypeDecode + Default> Default for RowTokenParser<T> {
fn default() -> Self {
Self {
decoder: T::default(),
}
}
}
#[async_trait]
impl<D: SqlTypeDecode + Default + Send + Sync, P: TdsPacketReader + Send + Sync> TokenParser<P>
for RowTokenParser<D>
{
async fn parse(&self, reader: &mut P, context: &ParserContext) -> TdsResult<Tokens> {
let (column_metadata_token, decryptor) = match context {
ParserContext::ColumnMetadata(metadata, decryptor) => {
trace!("Metadata during Row Parsing: {:?}", metadata);
(metadata, decryptor.as_ref())
}
_ => {
return Err(crate::error::Error::from(Error::new(
std::io::ErrorKind::InvalidData,
"Expected ColumnMetadata in context",
)));
}
};
let all_metadata = &column_metadata_token.columns;
let mut all_values: Vec<ColumnValues> =
Vec::with_capacity(column_metadata_token.column_count as usize);
for metadata in all_metadata {
trace!("Metadata: {:?}", metadata);
let column_value = match (metadata.crypto_metadata.is_some(), decryptor) {
(true, Some(dec)) => {
decrypt_encrypted_column(&self.decoder, reader, metadata, dec).await?
}
(true, None) => {
debug!(
column = %metadata.column_name,
"Encrypted column has no column-encryption decryptor available \
(Always Encrypted disabled for this command, or no key-store \
provider registered); returning the raw ciphertext varbinary"
);
self.decoder.decode(reader, metadata).await?
}
(false, _) => self.decoder.decode(reader, metadata).await?,
};
all_values.push(column_value);
}
Ok(Tokens::from(RowToken::new(all_values)))
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use super::*;
use crate::datatypes::sqldatatypes::{
FixedLengthTypes, TdsDataType, TypeInfo, TypeInfoVariant,
};
use crate::io::token_stream::ParserContext;
use crate::query::metadata::ColumnMetadata;
use crate::token::parsers::common::test_utils::MockReader;
use crate::token::tokens::ColMetadataToken;
#[derive(Default)]
struct MockDecoder;
impl SqlTypeDecode for MockDecoder {
async fn decode<T>(
&self,
_reader: &mut T,
_metadata: &ColumnMetadata,
) -> TdsResult<ColumnValues>
where
T: TdsPacketReader + Send + Sync,
{
Ok(ColumnValues::Int(42))
}
}
fn make_int_column(name: &str) -> ColumnMetadata {
ColumnMetadata {
user_type: 0,
flags: 0,
type_info: TypeInfo {
tds_type: TdsDataType::Int4,
length: 4,
type_info_variant: TypeInfoVariant::FixedLen(FixedLengthTypes::Int4),
},
data_type: TdsDataType::Int4,
column_name: name.to_string(),
multi_part_name: None,
crypto_metadata: None,
}
}
fn make_context(columns: Vec<ColumnMetadata>) -> ParserContext {
ParserContext::ColumnMetadata(
Arc::new(ColMetadataToken {
column_count: columns.len() as u16,
columns,
cek_table: Vec::new(),
}),
None,
)
}
#[tokio::test]
async fn test_parse_no_metadata_context() {
let parser = RowTokenParser::<MockDecoder>::default();
let mut reader = MockReader::new(vec![]);
let context = ParserContext::None(());
let result = parser.parse(&mut reader, &context).await;
assert!(result.is_err());
assert!(
result
.unwrap_err()
.to_string()
.contains("Expected ColumnMetadata in context")
);
}
#[tokio::test]
async fn test_parse_single_column() {
let parser = RowTokenParser::<MockDecoder>::default();
let mut reader = MockReader::new(vec![]);
let context = make_context(vec![make_int_column("id")]);
let result = parser.parse(&mut reader, &context).await.unwrap();
match result {
Tokens::Row(row) => {
assert_eq!(row.all_values.len(), 1);
assert_eq!(row.all_values[0], ColumnValues::Int(42));
}
_ => panic!("Expected Row token"),
}
}
#[tokio::test]
async fn test_parse_multiple_columns() {
let parser = RowTokenParser::<MockDecoder>::default();
let mut reader = MockReader::new(vec![]);
let context = make_context(vec![
make_int_column("a"),
make_int_column("b"),
make_int_column("c"),
]);
let result = parser.parse(&mut reader, &context).await.unwrap();
match result {
Tokens::Row(row) => assert_eq!(row.all_values.len(), 3),
_ => panic!("Expected Row token"),
}
}
#[tokio::test]
async fn test_parse_zero_columns() {
let parser = RowTokenParser::<MockDecoder>::default();
let mut reader = MockReader::new(vec![]);
let context = make_context(vec![]);
let result = parser.parse(&mut reader, &context).await.unwrap();
match result {
Tokens::Row(row) => assert!(row.all_values.is_empty()),
_ => panic!("Expected Row token"),
}
}
#[derive(Default)]
struct FailingDecoder;
impl SqlTypeDecode for FailingDecoder {
async fn decode<T>(
&self,
_reader: &mut T,
_metadata: &ColumnMetadata,
) -> TdsResult<ColumnValues>
where
T: TdsPacketReader + Send + Sync,
{
Err(crate::error::Error::ProtocolError(
"decode failure".to_string(),
))
}
}
#[tokio::test]
async fn test_parse_decoder_error_propagates() {
let parser = RowTokenParser::<FailingDecoder>::default();
let mut reader = MockReader::new(vec![]);
let context = make_context(vec![make_int_column("a")]);
let result = parser.parse(&mut reader, &context).await;
assert!(result.is_err());
assert!(result.unwrap_err().to_string().contains("decode failure"));
}
}