use std::{io::Error, mem::size_of};
use async_trait::async_trait;
use super::super::tokens::Tokens;
use super::common::TokenParser;
use crate::{core::TdsResult, io::packet_reader::TdsPacketReader};
use crate::{
datatypes::sqldatatypes::{TdsDataType, read_type_info},
io::token_stream::ParserContext,
query::metadata::{
CekTableEntry, ColumnMetadata, CryptoMetadata, EncryptedCekValue, MultiPartName,
},
token::tokens::ColMetadataToken,
};
const FLAG_ENCRYPTED: u16 = 0x0800;
const COLUMN_PREALLOC_BYTES: usize = 160 * 1024;
const COLUMN_PREALLOC_CAP: usize = COLUMN_PREALLOC_BYTES / size_of::<ColumnMetadata>();
const CEK_TABLE_PREALLOC_BYTES: usize = 12 * 1024;
const CEK_TABLE_PREALLOC_CAP: usize = CEK_TABLE_PREALLOC_BYTES / size_of::<CekTableEntry>();
fn bounded_capacity(count: u16, cap: usize) -> usize {
(count as usize).min(cap)
}
fn new_column_vec(col_count: u16) -> Vec<ColumnMetadata> {
Vec::with_capacity(bounded_capacity(col_count, COLUMN_PREALLOC_CAP))
}
fn new_cek_table_vec(entry_count: u16) -> Vec<CekTableEntry> {
Vec::with_capacity(bounded_capacity(entry_count, CEK_TABLE_PREALLOC_CAP))
}
const CUSTOM_CIPHER_ALGORITHM_ID: u8 = 0x00;
#[derive(Default)]
pub(crate) struct ColMetadataTokenParser;
#[async_trait]
impl<T> TokenParser<T> for ColMetadataTokenParser
where
T: TdsPacketReader + Send + Sync,
{
async fn parse(&self, reader: &mut T, context: &ParserContext) -> TdsResult<Tokens> {
let is_column_encryption_supported = context.is_column_encryption_supported();
let col_count = reader.read_uint16().await?;
if col_count == 0xFFFF {
return Ok(Tokens::from(ColMetadataToken::default()));
}
let cek_table = if is_column_encryption_supported {
parse_cek_table(reader).await?
} else {
Vec::new()
};
let has_cek_table = is_column_encryption_supported;
let mut column_metadata = new_column_vec(col_count);
for _ in 0..col_count {
let user_type = reader.read_uint32().await?;
let flags = reader.read_uint16().await?;
let raw_data_type = reader.read_byte().await?;
let some_data_type = TdsDataType::try_from(raw_data_type);
if some_data_type.is_err() {
return Err(crate::error::Error::from(Error::new(
std::io::ErrorKind::InvalidData,
format!("Invalid data type: {raw_data_type}"),
)));
}
let data_type = some_data_type?;
let type_info = read_type_info(reader, data_type).await?;
let multi_part_name = match data_type {
TdsDataType::Text | TdsDataType::NText | TdsDataType::Image => {
let mut part_count = reader.read_byte().await?;
if part_count == 0 {
None
} else {
let mut mpt = MultiPartName::default();
while part_count > 0 {
let part_name = reader.read_varchar_u16_length().await?;
if part_count == 4 {
mpt.server_name = part_name;
} else if part_count == 3 {
mpt.catalog_name = part_name;
} else if part_count == 2 {
mpt.schema_name = part_name;
} else if part_count == 1 {
mpt.table_name = part_name.unwrap_or_default();
}
part_count -= 1;
}
Some(mpt)
}
}
_ => None,
};
let crypto_metadata = if is_column_encryption_supported && (flags & FLAG_ENCRYPTED) != 0
{
Some(parse_crypto_metadata(reader, has_cek_table).await?)
} else {
None
};
let col_name = reader.read_varchar_u8_length().await?;
let col_metadata = ColumnMetadata {
user_type,
flags,
data_type,
type_info,
column_name: col_name,
multi_part_name,
crypto_metadata,
};
column_metadata.push(col_metadata);
}
let metadata = ColMetadataToken {
column_count: col_count,
columns: column_metadata,
cek_table,
};
Ok(Tokens::from(metadata))
}
}
async fn parse_cek_table<T>(reader: &mut T) -> TdsResult<Vec<CekTableEntry>>
where
T: TdsPacketReader + Send + Sync,
{
let entry_count = reader.read_uint16().await?;
let mut entries = new_cek_table_vec(entry_count);
for _ in 0..entry_count {
entries.push(parse_cek_table_entry(reader).await?);
}
Ok(entries)
}
async fn parse_cek_table_entry<T>(reader: &mut T) -> TdsResult<CekTableEntry>
where
T: TdsPacketReader + Send + Sync,
{
let database_id = reader.read_int32().await?;
let cek_id = reader.read_int32().await?;
let cek_version = reader.read_int32().await?;
let mut cek_md_version = [0u8; 8];
reader.read_bytes(&mut cek_md_version).await?;
let value_count = reader.read_byte().await?;
let mut encrypted_cek_values: Vec<EncryptedCekValue> = Vec::with_capacity(value_count as usize);
for _ in 0..value_count {
let encrypted_len = reader.read_uint16().await? as usize;
let mut encrypted_key = vec![0u8; encrypted_len];
reader.read_bytes(&mut encrypted_key).await?;
let key_store_len = reader.read_byte().await? as usize;
let key_store_name = reader.read_unicode(key_store_len).await?;
let key_path_len = reader.read_uint16().await? as usize;
let key_path = reader.read_unicode(key_path_len).await?;
let algorithm_len = reader.read_byte().await? as usize;
let algorithm_name = reader.read_unicode(algorithm_len).await?;
encrypted_cek_values.push(EncryptedCekValue {
encrypted_key,
key_store_name,
key_path,
algorithm_name,
});
}
Ok(CekTableEntry {
database_id,
cek_id,
cek_version,
cek_md_version,
encrypted_cek_values,
})
}
pub(super) async fn parse_crypto_metadata<T>(
reader: &mut T,
has_cek_table: bool,
) -> TdsResult<CryptoMetadata>
where
T: TdsPacketReader + Send + Sync,
{
let cek_table_ordinal = if has_cek_table {
reader.read_uint16().await?
} else {
0
};
let _user_type = reader.read_uint32().await?;
let raw_base_type = reader.read_byte().await?;
let base_data_type = TdsDataType::try_from(raw_base_type).map_err(|_| {
crate::error::Error::from(Error::new(
std::io::ErrorKind::InvalidData,
format!("Invalid base data type in crypto metadata: {raw_base_type}"),
))
})?;
let base_type_info = read_type_info(reader, base_data_type).await?;
let cipher_algorithm_id = reader.read_byte().await?;
let cipher_algorithm_name = if cipher_algorithm_id == CUSTOM_CIPHER_ALGORITHM_ID {
let name_len = reader.read_byte().await? as usize;
Some(reader.read_unicode(name_len).await?)
} else {
None
};
let encryption_type = reader.read_byte().await?;
let normalization_rule_version = reader.read_byte().await?;
Ok(CryptoMetadata {
cek_table_ordinal,
base_data_type,
base_type_info,
cipher_algorithm_id,
cipher_algorithm_name,
encryption_type,
normalization_rule_version,
})
}
#[cfg(test)]
mod tests {
use super::super::common::test_utils::MockReader;
use super::*;
use crate::datatypes::sqldatatypes::TdsDataType;
use byteorder::{ByteOrder, LittleEndian};
fn build_colmetadata_bytes(col_count: u16, columns: Vec<ColumnData>) -> Vec<u8> {
let mut data = Vec::new();
let mut buf = [0u8; 2];
LittleEndian::write_u16(&mut buf, col_count);
data.extend_from_slice(&buf);
for col in columns {
let mut buf = [0u8; 4];
LittleEndian::write_u32(&mut buf, col.user_type);
data.extend_from_slice(&buf);
let mut buf = [0u8; 2];
LittleEndian::write_u16(&mut buf, col.flags);
data.extend_from_slice(&buf);
data.push(col.data_type_byte);
data.extend_from_slice(&col.type_info_bytes);
let name_bytes = MockReader::encode_utf16(&col.name);
data.push((name_bytes.len() / 2) as u8); data.extend_from_slice(&name_bytes);
}
data
}
#[derive(Clone)]
struct ColumnData {
user_type: u32,
flags: u16,
data_type_byte: u8,
type_info_bytes: Vec<u8>,
name: String,
}
#[tokio::test]
async fn test_parse_no_metadata() {
let data = vec![0xFF, 0xFF];
let mut reader = MockReader::new(data);
let parser = ColMetadataTokenParser;
let context = ParserContext::default();
let result = parser.parse(&mut reader, &context).await.unwrap();
match result {
Tokens::ColMetadata(token) => {
assert_eq!(token.column_count, 0);
assert_eq!(token.columns.len(), 0);
}
_ => panic!("Expected ColMetadata token"),
}
}
#[tokio::test]
async fn test_parse_single_int_column() {
let columns = vec![ColumnData {
user_type: 0,
flags: 0x00, data_type_byte: TdsDataType::Int4 as u8,
type_info_bytes: vec![], name: "id".to_string(),
}];
let data = build_colmetadata_bytes(1, columns);
let mut reader = MockReader::new(data);
let parser = ColMetadataTokenParser;
let context = ParserContext::default();
let result = parser.parse(&mut reader, &context).await.unwrap();
match result {
Tokens::ColMetadata(token) => {
assert_eq!(token.column_count, 1);
assert_eq!(token.columns.len(), 1);
assert_eq!(token.columns[0].user_type, 0);
assert_eq!(token.columns[0].flags, 0x00);
assert_eq!(token.columns[0].data_type, TdsDataType::Int4);
assert_eq!(token.columns[0].column_name, "id");
}
_ => panic!("Expected ColMetadata token"),
}
}
#[tokio::test]
async fn test_parse_nullable_column() {
let columns = vec![ColumnData {
user_type: 0,
flags: 0x01, data_type_byte: TdsDataType::IntN as u8,
type_info_bytes: vec![0x04], name: "age".to_string(),
}];
let data = build_colmetadata_bytes(1, columns);
let mut reader = MockReader::new(data);
let parser = ColMetadataTokenParser;
let context = ParserContext::default();
let result = parser.parse(&mut reader, &context).await.unwrap();
match result {
Tokens::ColMetadata(token) => {
assert_eq!(token.columns.len(), 1);
assert_eq!(token.columns[0].flags, 0x01);
assert_eq!(token.columns[0].data_type, TdsDataType::IntN);
assert!(token.columns[0].is_nullable());
}
_ => panic!("Expected ColMetadata token"),
}
}
#[tokio::test]
async fn test_parse_multiple_columns() {
let columns = vec![
ColumnData {
user_type: 0,
flags: 0x00,
data_type_byte: TdsDataType::Int4 as u8,
type_info_bytes: vec![],
name: "id".to_string(),
},
ColumnData {
user_type: 0,
flags: 0x01,
data_type_byte: TdsDataType::IntN as u8,
type_info_bytes: vec![0x04],
name: "age".to_string(),
},
ColumnData {
user_type: 0,
flags: 0x01,
data_type_byte: TdsDataType::BigVarChar as u8,
type_info_bytes: {
let mut bytes = vec![
0x32, 0x00, ];
bytes.extend_from_slice(&[0x09, 0x04, 0xD0, 0x00, 0x34]);
bytes
},
name: "name".to_string(),
},
];
let data = build_colmetadata_bytes(3, columns.clone());
let mut reader = MockReader::new(data);
let parser = ColMetadataTokenParser;
let context = ParserContext::default();
let result = parser.parse(&mut reader, &context).await.unwrap();
match result {
Tokens::ColMetadata(token) => {
assert_eq!(token.column_count, 3);
assert_eq!(token.columns.len(), 3);
assert_eq!(token.columns[0].column_name, "id");
assert_eq!(token.columns[0].data_type, TdsDataType::Int4);
assert_eq!(token.columns[1].column_name, "age");
assert_eq!(token.columns[1].data_type, TdsDataType::IntN);
assert_eq!(token.columns[2].column_name, "name");
assert_eq!(token.columns[2].data_type, TdsDataType::BigVarChar);
}
_ => panic!("Expected ColMetadata token"),
}
}
#[tokio::test]
async fn test_parse_bigint_column() {
let columns = vec![ColumnData {
user_type: 0,
flags: 0x00,
data_type_byte: TdsDataType::Int8 as u8,
type_info_bytes: vec![],
name: "bigid".to_string(),
}];
let data = build_colmetadata_bytes(1, columns);
let mut reader = MockReader::new(data);
let parser = ColMetadataTokenParser;
let context = ParserContext::default();
let result = parser.parse(&mut reader, &context).await.unwrap();
match result {
Tokens::ColMetadata(token) => {
assert_eq!(token.columns.len(), 1);
assert_eq!(token.columns[0].data_type, TdsDataType::Int8);
}
_ => panic!("Expected ColMetadata token"),
}
}
#[tokio::test]
async fn test_parse_identity_column() {
let columns = vec![ColumnData {
user_type: 0,
flags: 0x10, data_type_byte: TdsDataType::Int4 as u8,
type_info_bytes: vec![],
name: "id".to_string(),
}];
let data = build_colmetadata_bytes(1, columns);
let mut reader = MockReader::new(data);
let parser = ColMetadataTokenParser;
let context = ParserContext::default();
let result = parser.parse(&mut reader, &context).await.unwrap();
match result {
Tokens::ColMetadata(token) => {
assert_eq!(token.columns.len(), 1);
assert_eq!(token.columns[0].flags, 0x10);
assert!(token.columns[0].is_identity());
}
_ => panic!("Expected ColMetadata token"),
}
}
#[tokio::test]
async fn test_invalid_data_type() {
let mut data = vec![0x01, 0x00]; data.extend_from_slice(&[0x00, 0x00, 0x00, 0x00]); data.extend_from_slice(&[0x00, 0x00]); data.push(0xFF);
let mut reader = MockReader::new(data);
let parser = ColMetadataTokenParser;
let context = ParserContext::default();
let result = parser.parse(&mut reader, &context).await;
assert!(result.is_err());
}
#[test]
fn test_column_vec_capacity_is_bounded() {
assert_eq!(new_column_vec(0).capacity(), 0);
assert_eq!(new_column_vec(10).capacity(), 10);
assert_eq!(new_column_vec(5000).capacity(), COLUMN_PREALLOC_CAP);
assert_eq!(new_column_vec(0xFFFE).capacity(), COLUMN_PREALLOC_CAP);
}
#[test]
fn test_cek_table_vec_capacity_is_bounded() {
assert_eq!(new_cek_table_vec(0).capacity(), 0);
assert_eq!(new_cek_table_vec(10).capacity(), 10);
assert_eq!(
new_cek_table_vec(u16::MAX).capacity(),
CEK_TABLE_PREALLOC_CAP
);
}
#[tokio::test]
async fn test_large_column_count_without_data_returns_error() {
let data = vec![0x84, 0xB8]; let mut reader = MockReader::new(data);
let parser = ColMetadataTokenParser;
let context = ParserContext::default();
assert!(parser.parse(&mut reader, &context).await.is_err());
}
#[tokio::test]
async fn test_large_valid_column_count_without_data_returns_error() {
let data = vec![0x88, 0x13]; let mut reader = MockReader::new(data);
let parser = ColMetadataTokenParser;
let context = ParserContext::default();
assert!(parser.parse(&mut reader, &context).await.is_err());
}
#[tokio::test]
async fn test_column_encryption_empty_cek_table_no_encrypted_columns() {
let mut data = Vec::new();
data.extend_from_slice(&[0x01, 0x00]); data.extend_from_slice(&[0x00, 0x00]); data.extend_from_slice(&[0x00, 0x00, 0x00, 0x00]); data.extend_from_slice(&[0x00, 0x00]); data.push(TdsDataType::Int4 as u8); let name = MockReader::encode_utf16("id");
data.push((name.len() / 2) as u8);
data.extend_from_slice(&name);
let mut reader = MockReader::new(data);
let parser = ColMetadataTokenParser;
let context = ParserContext::ColumnEncryption(true);
let result = parser.parse(&mut reader, &context).await.unwrap();
match result {
Tokens::ColMetadata(token) => {
assert_eq!(token.column_count, 1);
assert!(token.cek_table.is_empty());
assert_eq!(token.columns.len(), 1);
assert!(token.columns[0].crypto_metadata.is_none());
assert!(!token.columns[0].is_encrypted());
}
_ => panic!("Expected ColMetadata token"),
}
}
#[tokio::test]
async fn test_column_encryption_encrypted_column_with_cek_table() {
let mut data = Vec::new();
data.extend_from_slice(&[0x01, 0x00]);
data.extend_from_slice(&[0x01, 0x00]); data.extend_from_slice(&5i32.to_le_bytes()); data.extend_from_slice(&7i32.to_le_bytes()); data.extend_from_slice(&1i32.to_le_bytes()); data.extend_from_slice(&[1, 2, 3, 4, 5, 6, 7, 8]); data.push(0x01); data.extend_from_slice(&[0x04, 0x00]); data.extend_from_slice(&[0xDE, 0xAD, 0xBE, 0xEF]); data.push(0x03); data.extend_from_slice(&MockReader::encode_utf16("AKV"));
data.extend_from_slice(&[0x03, 0x00]); data.extend_from_slice(&MockReader::encode_utf16("url"));
data.push(0x03); data.extend_from_slice(&MockReader::encode_utf16("RSA"));
data.extend_from_slice(&[0x00, 0x00, 0x00, 0x00]); data.extend_from_slice(&FLAG_ENCRYPTED.to_le_bytes()); data.push(TdsDataType::Int4 as u8); data.extend_from_slice(&[0x00, 0x00]); data.extend_from_slice(&[0x00, 0x00, 0x00, 0x00]); data.push(TdsDataType::Int4 as u8); data.push(0x02); data.push(0x01); data.push(0x01); let name = MockReader::encode_utf16("secret");
data.push((name.len() / 2) as u8);
data.extend_from_slice(&name);
let mut reader = MockReader::new(data);
let parser = ColMetadataTokenParser;
let context = ParserContext::ColumnEncryption(true);
let result = parser.parse(&mut reader, &context).await.unwrap();
match result {
Tokens::ColMetadata(token) => {
assert_eq!(token.column_count, 1);
assert_eq!(token.cek_table.len(), 1);
let entry = &token.cek_table[0];
assert_eq!(entry.database_id, 5);
assert_eq!(entry.cek_id, 7);
assert_eq!(entry.cek_version, 1);
assert_eq!(entry.cek_md_version, [1, 2, 3, 4, 5, 6, 7, 8]);
assert_eq!(entry.encrypted_cek_values.len(), 1);
let value = &entry.encrypted_cek_values[0];
assert_eq!(value.encrypted_key, vec![0xDE, 0xAD, 0xBE, 0xEF]);
assert_eq!(value.key_store_name, "AKV");
assert_eq!(value.key_path, "url");
assert_eq!(value.algorithm_name, "RSA");
assert_eq!(token.columns.len(), 1);
let col = &token.columns[0];
assert_eq!(col.column_name, "secret");
assert!(col.is_encrypted());
let crypto = col.crypto_metadata.as_ref().expect("crypto metadata");
assert_eq!(crypto.cek_table_ordinal, 0);
assert_eq!(crypto.base_data_type, TdsDataType::Int4);
assert_eq!(crypto.cipher_algorithm_id, 0x02);
assert!(crypto.cipher_algorithm_name.is_none());
assert_eq!(crypto.encryption_type, 1);
assert_eq!(crypto.normalization_rule_version, 1);
}
_ => panic!("Expected ColMetadata token"),
}
}
#[tokio::test]
async fn test_column_encryption_disabled_skips_cek_table() {
let columns = vec![ColumnData {
user_type: 0,
flags: 0x00,
data_type_byte: TdsDataType::Int4 as u8,
type_info_bytes: vec![],
name: "id".to_string(),
}];
let data = build_colmetadata_bytes(1, columns);
let mut reader = MockReader::new(data);
let parser = ColMetadataTokenParser;
let context = ParserContext::None(());
let result = parser.parse(&mut reader, &context).await.unwrap();
match result {
Tokens::ColMetadata(token) => {
assert_eq!(token.column_count, 1);
assert!(token.cek_table.is_empty());
assert!(token.columns[0].crypto_metadata.is_none());
}
_ => panic!("Expected ColMetadata token"),
}
}
#[tokio::test]
async fn parse_crypto_metadata_reads_custom_cipher_name() {
let mut data = Vec::new();
data.extend_from_slice(&0u32.to_le_bytes()); data.push(TdsDataType::Int4 as u8); data.push(CUSTOM_CIPHER_ALGORITHM_ID); data.push(3); data.extend_from_slice(&MockReader::encode_utf16("AES"));
data.push(1); data.push(1);
let mut reader = MockReader::new(data);
let md = parse_crypto_metadata(&mut reader, false).await.unwrap();
assert_eq!(md.cek_table_ordinal, 0);
assert_eq!(md.base_data_type, TdsDataType::Int4);
assert_eq!(md.cipher_algorithm_id, CUSTOM_CIPHER_ALGORITHM_ID);
assert_eq!(md.cipher_algorithm_name.as_deref(), Some("AES"));
assert_eq!(md.encryption_type, 1);
assert_eq!(md.normalization_rule_version, 1);
}
#[tokio::test]
async fn parse_crypto_metadata_rejects_invalid_base_type() {
let mut data = Vec::new();
data.extend_from_slice(&0u32.to_le_bytes()); data.push(0x01);
let mut reader = MockReader::new(data);
let result = parse_crypto_metadata(&mut reader, false).await;
assert!(result.is_err());
}
}