#![allow(dead_code)]
mod certificate;
pub use certificate::RsaKeyStoreProvider;
use std::collections::HashMap;
use std::sync::{Arc, Mutex};
use async_trait::async_trait;
use zeroize::Zeroizing;
use crate::core::TdsResult;
use crate::datatypes::column_values::ColumnValues;
use crate::error::Error;
use crate::query::metadata::{CekTableEntry, CryptoMetadata};
use crate::security::cell_decryptor::CellDecryptor;
use crate::security::encryption::decrypt_cell;
#[async_trait]
pub trait ColumnEncryptionKeyStoreProvider: Send + Sync {
async fn decrypt_column_encryption_key(
&self,
master_key_path: &str,
encryption_algorithm: &str,
encrypted_cek: &[u8],
) -> TdsResult<Vec<u8>>;
}
#[derive(Clone, Default)]
pub(crate) struct ColumnEncryptionKeyStoreProviderRegistry {
providers: HashMap<String, Arc<dyn ColumnEncryptionKeyStoreProvider>>,
}
impl ColumnEncryptionKeyStoreProviderRegistry {
pub(crate) fn new() -> Self {
Self::default()
}
pub(crate) fn register(
&mut self,
name: impl AsRef<str>,
provider: Arc<dyn ColumnEncryptionKeyStoreProvider>,
) {
self.providers
.insert(name.as_ref().to_ascii_uppercase(), provider);
}
pub(crate) fn get(&self, name: &str) -> Option<Arc<dyn ColumnEncryptionKeyStoreProvider>> {
self.providers.get(&name.to_ascii_uppercase()).cloned()
}
pub(crate) fn is_empty(&self) -> bool {
self.providers.is_empty()
}
}
#[derive(Clone, PartialEq, Eq, Hash)]
struct CekCacheKey {
provider_name: String,
master_key_path: String,
encrypted_cek: Vec<u8>,
}
#[derive(Default)]
pub(crate) struct CekCache {
entries: Mutex<HashMap<CekCacheKey, Arc<Zeroizing<Vec<u8>>>>>,
}
impl CekCache {
pub(crate) fn new() -> Self {
Self::default()
}
fn get(&self, key: &CekCacheKey) -> Option<Arc<Zeroizing<Vec<u8>>>> {
self.entries
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.get(key)
.cloned()
}
fn insert(&self, key: CekCacheKey, value: Arc<Zeroizing<Vec<u8>>>) {
self.entries
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.insert(key, value);
}
}
pub(crate) async fn decrypt_cek(
registry: &ColumnEncryptionKeyStoreProviderRegistry,
cache: &CekCache,
entry: &CekTableEntry,
trusted_key_paths: &[String],
) -> TdsResult<Arc<Zeroizing<Vec<u8>>>> {
if entry.encrypted_cek_values.is_empty() {
return Err(Error::ColumnEncryptionError(
"CEK table entry has no encrypted key values".to_string(),
));
}
let mut last_error: Option<Error> = None;
for value in &entry.encrypted_cek_values {
if !trusted_key_paths.is_empty()
&& !trusted_key_paths
.iter()
.any(|trusted| trusted.eq_ignore_ascii_case(&value.key_path))
{
last_error = Some(Error::ColumnEncryptionError(format!(
"The column master key path '{}' is not in the trusted master key paths list \
configured for this server; refusing to use it to unwrap a column encryption key.",
value.key_path
)));
continue;
}
let cache_key = CekCacheKey {
provider_name: value.key_store_name.to_ascii_uppercase(),
master_key_path: value.key_path.clone(),
encrypted_cek: value.encrypted_key.clone(),
};
if let Some(cached) = cache.get(&cache_key) {
return Ok(cached);
}
let Some(provider) = registry.get(&value.key_store_name) else {
last_error = Some(Error::ColumnEncryptionError(format!(
"No column master key store provider is registered with name '{}'",
value.key_store_name
)));
continue;
};
match provider
.decrypt_column_encryption_key(
&value.key_path,
&value.algorithm_name,
&value.encrypted_key,
)
.await
{
Ok(plaintext) => {
let plaintext = Arc::new(Zeroizing::new(plaintext));
cache.insert(cache_key, plaintext.clone());
return Ok(plaintext);
}
Err(error) => last_error = Some(error),
}
}
Err(last_error.unwrap_or_else(|| {
Error::ColumnEncryptionError(
"Unable to decrypt column encryption key with any registered provider".to_string(),
)
}))
}
pub(crate) struct ResolvedCekDecryptor {
ceks: Vec<Result<Arc<Zeroizing<Vec<u8>>>, String>>,
}
impl std::fmt::Debug for ResolvedCekDecryptor {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
struct RedactedCek<'a>(&'a Result<Arc<Zeroizing<Vec<u8>>>, String>);
impl std::fmt::Debug for RedactedCek<'_> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self.0 {
Ok(_) => write!(f, "<resolved>"),
Err(e) => write!(f, "<unresolved: {e}>"),
}
}
}
f.debug_struct("ResolvedCekDecryptor")
.field(
"ceks",
&self.ceks.iter().map(RedactedCek).collect::<Vec<_>>(),
)
.finish()
}
}
impl ResolvedCekDecryptor {
pub(crate) async fn resolve(
registry: &ColumnEncryptionKeyStoreProviderRegistry,
cache: &CekCache,
cek_table: &[CekTableEntry],
trusted_key_paths: &[String],
) -> Self {
let mut ceks = Vec::with_capacity(cek_table.len());
for entry in cek_table {
ceks.push(
decrypt_cek(registry, cache, entry, trusted_key_paths)
.await
.map_err(|error| error.to_string()),
);
}
Self { ceks }
}
}
impl CellDecryptor for ResolvedCekDecryptor {
fn decrypt(
&self,
crypto_metadata: &CryptoMetadata,
cipher_blob: &[u8],
) -> TdsResult<ColumnValues> {
let ordinal = crypto_metadata.cek_table_ordinal as usize;
let cek = self.ceks.get(ordinal).ok_or_else(|| {
Error::ColumnEncryptionError(format!(
"Encrypted column references CEK table ordinal {ordinal}, but the CEK table has \
{} entries",
self.ceks.len()
))
})?;
let cek = cek.as_ref().map_err(|message| {
Error::ColumnEncryptionError(format!(
"Column encryption key at ordinal {ordinal} could not be resolved: {message}"
))
})?;
decrypt_cell(crypto_metadata, cek, cipher_blob)
}
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use super::*;
use crate::query::metadata::EncryptedCekValue;
struct MockProvider {
plaintext: Vec<u8>,
calls: AtomicUsize,
}
impl MockProvider {
fn new(plaintext: Vec<u8>) -> Self {
Self {
plaintext,
calls: AtomicUsize::new(0),
}
}
}
#[async_trait]
impl ColumnEncryptionKeyStoreProvider for MockProvider {
async fn decrypt_column_encryption_key(
&self,
_master_key_path: &str,
_encryption_algorithm: &str,
_encrypted_cek: &[u8],
) -> TdsResult<Vec<u8>> {
self.calls.fetch_add(1, Ordering::SeqCst);
Ok(self.plaintext.clone())
}
}
struct FailingProvider;
#[async_trait]
impl ColumnEncryptionKeyStoreProvider for FailingProvider {
async fn decrypt_column_encryption_key(
&self,
_master_key_path: &str,
_encryption_algorithm: &str,
_encrypted_cek: &[u8],
) -> TdsResult<Vec<u8>> {
Err(Error::ColumnEncryptionError("provider failed".to_string()))
}
}
fn cek_value(store: &str, path: &str, key: &[u8]) -> EncryptedCekValue {
EncryptedCekValue {
encrypted_key: key.to_vec(),
key_store_name: store.to_string(),
key_path: path.to_string(),
algorithm_name: "RSA_OAEP".to_string(),
}
}
fn entry(values: Vec<EncryptedCekValue>) -> CekTableEntry {
CekTableEntry {
database_id: 1,
cek_id: 2,
cek_version: 1,
cek_md_version: [0u8; 8],
encrypted_cek_values: values,
}
}
#[test]
fn registry_lookup_is_case_insensitive() {
let mut registry = ColumnEncryptionKeyStoreProviderRegistry::new();
registry.register(
"MSSQL_CERTIFICATE_STORE",
Arc::new(MockProvider::new(vec![1u8; 32])),
);
assert!(registry.get("mssql_certificate_store").is_some());
assert!(registry.get("MSSQL_Certificate_Store").is_some());
assert!(registry.get("UNKNOWN").is_none());
assert!(!registry.is_empty());
}
#[tokio::test]
async fn decrypt_cek_resolves_and_caches() {
let provider = Arc::new(MockProvider::new(vec![7u8; 32]));
let mut registry = ColumnEncryptionKeyStoreProviderRegistry::new();
registry.register("PROVIDER", provider.clone());
let cache = CekCache::new();
let entry = entry(vec![cek_value("PROVIDER", "path", &[0xAB, 0xCD])]);
let first = decrypt_cek(®istry, &cache, &entry, &[]).await.unwrap();
assert_eq!(first.as_slice(), vec![7u8; 32].as_slice());
let second = decrypt_cek(®istry, &cache, &entry, &[]).await.unwrap();
assert_eq!(second.as_slice(), vec![7u8; 32].as_slice());
assert_eq!(provider.calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn decrypt_cek_falls_back_to_next_value() {
let mut registry = ColumnEncryptionKeyStoreProviderRegistry::new();
registry.register("FAILING", Arc::new(FailingProvider));
registry.register("GOOD", Arc::new(MockProvider::new(vec![9u8; 32])));
let cache = CekCache::new();
let entry = entry(vec![
cek_value("FAILING", "p1", &[1]),
cek_value("GOOD", "p2", &[2]),
]);
let key = decrypt_cek(®istry, &cache, &entry, &[]).await.unwrap();
assert_eq!(key.as_slice(), vec![9u8; 32].as_slice());
}
#[tokio::test]
async fn decrypt_cek_errors_when_no_provider_registered() {
let registry = ColumnEncryptionKeyStoreProviderRegistry::new();
let cache = CekCache::new();
let entry = entry(vec![cek_value("MISSING", "p", &[1])]);
let error = decrypt_cek(®istry, &cache, &entry, &[])
.await
.unwrap_err();
assert!(matches!(error, Error::ColumnEncryptionError(_)));
}
#[tokio::test]
async fn decrypt_cek_errors_when_no_values() {
let registry = ColumnEncryptionKeyStoreProviderRegistry::new();
let cache = CekCache::new();
let entry = entry(vec![]);
let error = decrypt_cek(®istry, &cache, &entry, &[])
.await
.unwrap_err();
assert!(matches!(error, Error::ColumnEncryptionError(_)));
}
#[tokio::test]
async fn decrypt_cek_rejects_untrusted_key_path() {
let provider = Arc::new(MockProvider::new(vec![7u8; 32]));
let mut registry = ColumnEncryptionKeyStoreProviderRegistry::new();
registry.register("PROVIDER", provider.clone());
let cache = CekCache::new();
let entry = entry(vec![cek_value(
"PROVIDER",
"https://vault/keys/attacker",
&[0xAB],
)]);
let trusted = vec!["https://vault/keys/trusted".to_string()];
let error = decrypt_cek(®istry, &cache, &entry, &trusted)
.await
.unwrap_err();
assert!(matches!(error, Error::ColumnEncryptionError(_)));
assert_eq!(provider.calls.load(Ordering::SeqCst), 0);
}
#[tokio::test]
async fn decrypt_cek_allows_trusted_key_path_case_insensitively() {
let provider = Arc::new(MockProvider::new(vec![7u8; 32]));
let mut registry = ColumnEncryptionKeyStoreProviderRegistry::new();
registry.register("PROVIDER", provider.clone());
let cache = CekCache::new();
let entry = entry(vec![cek_value(
"PROVIDER",
"https://Vault/Keys/Trusted",
&[0xAB],
)]);
let trusted = vec!["https://vault/keys/trusted".to_string()];
let key = decrypt_cek(®istry, &cache, &entry, &trusted)
.await
.unwrap();
assert_eq!(key.as_slice(), vec![7u8; 32].as_slice());
assert_eq!(provider.calls.load(Ordering::SeqCst), 1);
}
#[tokio::test]
async fn decrypt_cek_falls_back_from_untrusted_to_trusted_key_path() {
let provider = Arc::new(MockProvider::new(vec![7u8; 32]));
let mut registry = ColumnEncryptionKeyStoreProviderRegistry::new();
registry.register("PROVIDER", provider.clone());
let cache = CekCache::new();
let entry = entry(vec![
cek_value("PROVIDER", "https://vault/keys/attacker", &[0xAB]),
cek_value("PROVIDER", "https://vault/keys/trusted", &[0xCD]),
]);
let trusted = vec!["https://vault/keys/trusted".to_string()];
let key = decrypt_cek(®istry, &cache, &entry, &trusted)
.await
.unwrap();
assert_eq!(key.as_slice(), vec![7u8; 32].as_slice());
assert_eq!(provider.calls.load(Ordering::SeqCst), 1);
}
use crate::datatypes::column_values::ColumnValues;
use crate::datatypes::sqldatatypes::{
FixedLengthTypes, TdsDataType, TypeInfo, TypeInfoVariant,
};
use crate::query::metadata::CryptoMetadata;
use crate::security::encryption::{AeadAes256CbcHmacSha256, ColumnEncryptionType};
fn int_crypto_metadata(ordinal: u16) -> CryptoMetadata {
CryptoMetadata {
cek_table_ordinal: ordinal,
base_data_type: TdsDataType::Int4,
base_type_info: TypeInfo {
tds_type: TdsDataType::Int4,
length: 4,
type_info_variant: TypeInfoVariant::FixedLen(FixedLengthTypes::Int4),
},
cipher_algorithm_id: 0x02,
cipher_algorithm_name: None,
encryption_type: 1,
normalization_rule_version: 1,
}
}
#[tokio::test]
async fn resolved_decryptor_decrypts_cell() {
let cek = vec![3u8; 32];
let value: i64 = 0x1234_5678;
let cipher = AeadAes256CbcHmacSha256::new(&cek)
.unwrap()
.encrypt(&value.to_le_bytes(), ColumnEncryptionType::Deterministic)
.unwrap();
let mut registry = ColumnEncryptionKeyStoreProviderRegistry::new();
registry.register("PROVIDER", Arc::new(MockProvider::new(cek.clone())));
let cache = CekCache::new();
let cek_table = vec![entry(vec![cek_value("PROVIDER", "path", &[0xAB])])];
let decryptor = ResolvedCekDecryptor::resolve(®istry, &cache, &cek_table, &[]).await;
let crypto = int_crypto_metadata(0);
let decrypted = decryptor.decrypt(&crypto, &cipher).unwrap();
assert_eq!(decrypted, ColumnValues::Int(0x1234_5678));
}
#[tokio::test]
async fn resolved_decryptor_reports_unresolved_cek_on_use() {
let mut registry = ColumnEncryptionKeyStoreProviderRegistry::new();
registry.register("FAILING", Arc::new(FailingProvider));
let cache = CekCache::new();
let cek_table = vec![entry(vec![cek_value("FAILING", "p", &[1])])];
let decryptor = ResolvedCekDecryptor::resolve(®istry, &cache, &cek_table, &[]).await;
let error = decryptor
.decrypt(&int_crypto_metadata(0), &[0u8; 16])
.unwrap_err();
assert!(matches!(error, Error::ColumnEncryptionError(_)));
}
#[tokio::test]
async fn resolved_decryptor_errors_on_out_of_range_ordinal() {
let registry = ColumnEncryptionKeyStoreProviderRegistry::new();
let cache = CekCache::new();
let decryptor = ResolvedCekDecryptor::resolve(®istry, &cache, &[], &[]).await;
let error = decryptor.decrypt(&int_crypto_metadata(5), &[]).unwrap_err();
assert!(matches!(error, Error::ColumnEncryptionError(_)));
}
}