use crate::connection::tds_client::{ExecuteOptions, TdsClient};
use crate::core::{CancelHandle, TdsResult};
use crate::datatypes::bulk_copy_metadata::BulkCopyColumnMetadata;
use crate::error::Error;
use crate::sql_identifier::{
CATALOG_INDEX, SCHEMA_INDEX, TABLE_INDEX, build_multipart_name, escape_identifier,
escape_string_literal, parse_multipart_identifier,
};
use crate::token::tokens::ColMetadataToken;
use async_trait::async_trait;
use tracing::{debug, instrument, trace};
#[derive(Debug, Clone)]
pub(crate) struct TableMetadataResult {
pub col_metadata: ColMetadataToken,
pub collation_names: Vec<Option<String>>,
}
#[async_trait]
pub trait MetadataRetriever: Send {
async fn retrieve_metadata(
&mut self,
client: &mut TdsClient,
table_name: &str,
timeout_sec: u32,
) -> TdsResult<Vec<BulkCopyColumnMetadata>>;
}
#[derive(Debug, Default)]
pub struct FmtOnlyMetadataRetriever;
impl FmtOnlyMetadataRetriever {
pub fn new() -> Self {
Self
}
}
#[async_trait]
impl MetadataRetriever for FmtOnlyMetadataRetriever {
async fn retrieve_metadata(
&mut self,
client: &mut TdsClient,
table_name: &str,
timeout_sec: u32,
) -> TdsResult<Vec<BulkCopyColumnMetadata>> {
let timeout = if timeout_sec == 0 {
None
} else {
Some(timeout_sec)
};
let metadata_result = fetch_table_metadata(client, table_name, timeout, None).await?;
Vec::<BulkCopyColumnMetadata>::try_from(metadata_result)
}
}
impl TryFrom<TableMetadataResult> for Vec<BulkCopyColumnMetadata> {
type Error = Error;
fn try_from(result: TableMetadataResult) -> Result<Self, Self::Error> {
if result.col_metadata.columns.is_empty() {
return Err(Error::UsageError(
"Table not found or has no columns".to_string(),
));
}
let mut metadata: Vec<BulkCopyColumnMetadata> = result
.col_metadata
.columns
.iter()
.map(BulkCopyColumnMetadata::from)
.collect();
for (i, col_meta) in metadata.iter_mut().enumerate() {
if i < result.collation_names.len() {
col_meta.collation_name = result.collation_names[i].clone();
}
}
for (i, col_meta) in metadata.iter_mut().enumerate() {
let Some(crypto) = result
.col_metadata
.columns
.get(i)
.and_then(|c| c.crypto_metadata.as_ref())
else {
continue;
};
let cek_entry = result
.col_metadata
.cek_table
.get(crypto.cek_table_ordinal as usize)
.ok_or_else(|| {
Error::ProtocolError(format!(
"Encrypted column '{}' references CEK ordinal {} but the CEK table has \
{} entries",
col_meta.column_name,
crypto.cek_table_ordinal,
result.col_metadata.cek_table.len()
))
})?;
col_meta.encryption = Some(
crate::datatypes::bulk_copy_metadata::BulkCopyColumnEncryption {
crypto_metadata: crypto.clone(),
cek_entry: cek_entry.clone(),
},
);
}
Ok(metadata)
}
}
#[instrument(skip(client), level = "info")]
pub(crate) async fn fetch_table_metadata(
client: &mut TdsClient,
table_name: &str,
timeout_sec: Option<u32>,
cancel_handle: Option<&CancelHandle>,
) -> TdsResult<TableMetadataResult> {
let parts = parse_multipart_identifier(table_name, false)?;
let table_part = parts[TABLE_INDEX]
.as_ref()
.ok_or_else(|| Error::UsageError(format!("Invalid table name: {}", table_name)))?;
let is_temp_table = table_part.starts_with('#');
let catalog = if is_temp_table && parts[CATALOG_INDEX].is_none() {
"tempdb".to_string()
} else if let Some(cat) = &parts[CATALOG_INDEX] {
escape_identifier(cat)
} else {
String::new()
};
let full_name = build_multipart_name(&parts);
let escaped_full_name = escape_string_literal(&full_name);
let catalog_prefix = if !catalog.is_empty() {
format!("{}.", catalog)
} else {
String::new()
};
let catalog_for_sproc = if !catalog.is_empty() {
format!("{}..", catalog)
} else {
String::new()
};
let schema_name = parts[SCHEMA_INDEX]
.as_ref()
.map(|s| escape_identifier(&escape_string_literal(s)))
.unwrap_or_else(|| "dbo".to_string());
let table_name_escaped = escape_identifier(&escape_string_literal(table_part));
let query = format!(
r#"SELECT @@TRANCOUNT;
DECLARE @Column_Names NVARCHAR(MAX) = NULL;
DECLARE @object_id INT = OBJECT_ID('{escaped_full_name}');
DECLARE @sql NVARCHAR(MAX);
SET @sql = N'SELECT @CN = COALESCE(@CN + N'', '', N'''') + QUOTENAME([name]) FROM {catalog_prefix}sys.all_columns WHERE [object_id] = @ObjId';
IF EXISTS (SELECT TOP 1 * FROM sys.all_columns WHERE [object_id] = OBJECT_ID('sys.all_columns') AND [name] = 'graph_type')
SET @sql = @sql + N' AND COALESCE([graph_type], 0) NOT IN (1, 3, 4, 6, 7)';
SET @sql = @sql + N' ORDER BY [column_id] ASC';
EXEC sp_executesql @sql, N'@CN NVARCHAR(MAX) OUTPUT, @ObjId INT', @CN = @Column_Names OUTPUT, @ObjId = @object_id;
SELECT @Column_Names = COALESCE(@Column_Names, '*');
SET FMTONLY ON;
EXEC(N'SELECT ' + @Column_Names + N' FROM {escaped_full_name}');
SET FMTONLY OFF;
EXEC {catalog_for_sproc}sp_tablecollations_100 N'{schema_name}.{table_name_escaped}';"#
);
debug!("Fetching table metadata with FMTONLY and collations");
client
.execute(
query,
ExecuteOptions {
timeout: timeout_sec,
cancel: cancel_handle,
..Default::default()
},
)
.await?;
trace!("Consuming @@TRANCOUNT result set");
while let Some(_row) = client.get_next_row().await? {
}
trace!("Moving to FMTONLY metadata result set");
if !client.advance_to_rows().await? {
return Err(Error::UsageError(format!(
"Failed to move to FMTONLY metadata for table {}",
table_name
)));
}
let col_metadata = client
.get_current_metadata()
.ok_or_else(|| {
Error::UsageError(format!("Failed to fetch metadata for table {table_name}"))
})?
.clone();
debug!(
"Fetched {} columns from table metadata",
col_metadata.columns.len()
);
for (i, col) in col_metadata.columns.iter().enumerate() {
trace!(
"Column {}: name='{}', tds_type=0x{:02X}, nullable={}",
i,
col.column_name,
col.data_type as u8,
col.is_nullable()
);
}
trace!("Consuming FMTONLY result set rows (should be none)");
while let Some(_row) = client.get_next_row().await? {
}
trace!("Moving to sp_tablecollations_100 result set");
if !client.advance_to_rows().await? {
return Err(Error::UsageError(format!(
"Failed to move to collation metadata for table {}",
table_name
)));
}
let mut collation_names: Vec<Option<String>> = Vec::new();
trace!("Reading collation data from sp_tablecollations_100");
while let Some(row) = client.get_next_row().await? {
if row.len() > 3 {
let collation_value = &row[3];
let collation_name = match collation_value {
crate::datatypes::column_values::ColumnValues::String(s) => Some(s.to_string()),
crate::datatypes::column_values::ColumnValues::Null => None,
_ => {
trace!("Unexpected collation value type: {:?}", collation_value);
None
}
};
if let Some(name) = &collation_name {
trace!("Column {} collation: {}", collation_names.len(), name);
}
collation_names.push(collation_name);
} else {
trace!("Row has fewer than 4 columns, using None for collation");
collation_names.push(None);
}
}
debug!("Fetched {} collation names", collation_names.len());
while collation_names.len() < col_metadata.columns.len() {
collation_names.push(None);
}
client.close_query().await?;
Ok(TableMetadataResult {
col_metadata,
collation_names,
})
}
#[cfg(test)]
mod tests {
use super::*;
use crate::datatypes::sqldatatypes::{
FixedLengthTypes, TdsDataType, TypeInfo, TypeInfoVariant, VariableLengthTypes,
};
use crate::query::metadata::ColumnMetadata;
use crate::token::tokens::SqlCollation;
fn make_column(name: &str, flags: u16, type_info_variant: TypeInfoVariant) -> ColumnMetadata {
let tds_type = match &type_info_variant {
TypeInfoVariant::FixedLen(_) => TdsDataType::IntN,
TypeInfoVariant::VarLenString(_, _, _) => TdsDataType::NVarChar,
_ => TdsDataType::IntN,
};
ColumnMetadata {
user_type: 0,
flags,
type_info: TypeInfo {
tds_type,
length: 4,
type_info_variant,
},
data_type: tds_type,
column_name: name.to_string(),
multi_part_name: None,
crypto_metadata: None,
}
}
fn make_table_metadata(
columns: Vec<ColumnMetadata>,
collation_names: Vec<Option<String>>,
) -> TableMetadataResult {
TableMetadataResult {
col_metadata: ColMetadataToken {
column_count: columns.len() as u16,
columns,
cek_table: Vec::new(),
},
collation_names,
}
}
fn nvarchar_variant() -> TypeInfoVariant {
TypeInfoVariant::VarLenString(
VariableLengthTypes::NVarChar,
100,
Some(SqlCollation {
info: 0,
lcid_language_id: 0,
col_flags: 0,
sort_id: 0,
}),
)
}
#[test]
fn test_try_from_table_metadata_result_empty_columns() {
let result = make_table_metadata(vec![], vec![]);
let err = Vec::<BulkCopyColumnMetadata>::try_from(result).unwrap_err();
assert!(matches!(err, Error::UsageError(msg) if msg.contains("no columns")),);
}
#[test]
fn test_try_from_table_metadata_result_with_collations() {
let columns = vec![
make_column("col1", 0x00, nvarchar_variant()),
make_column("col2", 0x00, nvarchar_variant()),
];
let collations = vec![
Some("Latin1_General_CI_AS".to_string()),
Some("SQL_Latin1_General_CP1_CI_AS".to_string()),
];
let result = make_table_metadata(columns, collations);
let metadata = Vec::<BulkCopyColumnMetadata>::try_from(result).unwrap();
assert_eq!(metadata.len(), 2);
assert_eq!(
metadata[0].collation_name.as_deref(),
Some("Latin1_General_CI_AS")
);
assert_eq!(
metadata[1].collation_name.as_deref(),
Some("SQL_Latin1_General_CP1_CI_AS")
);
}
#[test]
fn test_try_from_table_metadata_result_collations_shorter_than_columns() {
let columns = vec![
make_column("col1", 0x00, nvarchar_variant()),
make_column("col2", 0x00, nvarchar_variant()),
make_column(
"col3",
0x00,
TypeInfoVariant::FixedLen(FixedLengthTypes::Int4),
),
];
let collations = vec![Some("Latin1_General_CI_AS".to_string())];
let result = make_table_metadata(columns, collations);
let metadata = Vec::<BulkCopyColumnMetadata>::try_from(result).unwrap();
assert_eq!(metadata.len(), 3);
assert_eq!(
metadata[0].collation_name.as_deref(),
Some("Latin1_General_CI_AS")
);
assert!(metadata[1].collation_name.is_none());
assert!(metadata[2].collation_name.is_none());
}
#[test]
fn test_try_from_table_metadata_result_no_collations() {
let columns = vec![
make_column("col1", 0x00, nvarchar_variant()),
make_column(
"col2",
0x00,
TypeInfoVariant::FixedLen(FixedLengthTypes::Int4),
),
];
let result = make_table_metadata(columns, vec![]);
let metadata = Vec::<BulkCopyColumnMetadata>::try_from(result).unwrap();
assert_eq!(metadata.len(), 2);
assert!(metadata[0].collation_name.is_none());
assert!(metadata[1].collation_name.is_none());
}
#[test]
fn test_try_from_table_metadata_result_preserves_identity_flag() {
let columns = vec![
make_column(
"id",
0x10,
TypeInfoVariant::FixedLen(FixedLengthTypes::Int4),
),
make_column("name", 0x00, nvarchar_variant()),
];
let result = make_table_metadata(columns, vec![]);
let metadata = Vec::<BulkCopyColumnMetadata>::try_from(result).unwrap();
assert!(metadata[0].is_identity);
assert!(!metadata[1].is_identity);
}
#[test]
fn test_try_from_table_metadata_captures_encryption_material() {
use crate::query::metadata::{CekTableEntry, CryptoMetadata};
let mut col = make_column(
"secret",
0x0800, TypeInfoVariant::FixedLen(FixedLengthTypes::Int4),
);
col.crypto_metadata = Some(CryptoMetadata {
cek_table_ordinal: 0,
base_data_type: TdsDataType::IntN,
base_type_info: TypeInfo {
tds_type: TdsDataType::IntN,
length: 4,
type_info_variant: TypeInfoVariant::FixedLen(FixedLengthTypes::Int4),
},
cipher_algorithm_id: 2,
cipher_algorithm_name: None,
encryption_type: 1,
normalization_rule_version: 1,
});
let cek_entry = CekTableEntry {
database_id: 7,
cek_id: 11,
cek_version: 1,
cek_md_version: [1, 2, 3, 4, 5, 6, 7, 8],
encrypted_cek_values: Vec::new(),
};
let mut table = make_table_metadata(vec![col], vec![None]);
table.col_metadata.cek_table = vec![cek_entry];
let metadata = Vec::<BulkCopyColumnMetadata>::try_from(table).unwrap();
assert_eq!(metadata.len(), 1);
assert!(metadata[0].is_encrypted);
let enc = metadata[0]
.encryption
.as_ref()
.expect("encryption material should be captured for an encrypted column");
assert_eq!(enc.crypto_metadata.encryption_type, 1);
assert_eq!(enc.crypto_metadata.cek_table_ordinal, 0);
assert_eq!(enc.cek_entry.database_id, 7);
assert_eq!(enc.cek_entry.cek_id, 11);
}
#[test]
fn test_try_from_table_metadata_rejects_dangling_cek_ordinal() {
use crate::query::metadata::CryptoMetadata;
let mut col = make_column(
"secret",
0x0800,
TypeInfoVariant::FixedLen(FixedLengthTypes::Int4),
);
col.crypto_metadata = Some(CryptoMetadata {
cek_table_ordinal: 3, base_data_type: TdsDataType::IntN,
base_type_info: TypeInfo {
tds_type: TdsDataType::IntN,
length: 4,
type_info_variant: TypeInfoVariant::FixedLen(FixedLengthTypes::Int4),
},
cipher_algorithm_id: 2,
cipher_algorithm_name: None,
encryption_type: 1,
normalization_rule_version: 1,
});
let table = make_table_metadata(vec![col], vec![None]);
let err = Vec::<BulkCopyColumnMetadata>::try_from(table).unwrap_err();
assert!(matches!(err, Error::ProtocolError(_)));
}
}