use std::collections::HashMap;
use crate::dek_buffer::EncryptionAlgo;
#[derive(Debug, Clone)]
pub struct ColumnCryptoConfig {
pub table: String,
pub column: String,
pub algorithm: EncryptionAlgo,
pub key_version: u32,
}
impl ColumnCryptoConfig {
pub fn new(table: impl Into<String>, column: impl Into<String>) -> Self {
Self {
table: table.into(),
column: column.into(),
algorithm: EncryptionAlgo::default(),
key_version: 1,
}
}
pub fn with_algorithm(mut self, algo: EncryptionAlgo) -> Self {
self.algorithm = algo;
self
}
pub fn with_key_version(mut self, version: u32) -> Self {
self.key_version = version;
self
}
pub fn key(&self) -> String {
format!("{}.{}", self.table, self.column)
}
}
pub struct ColumnEncryptionPolicy {
configs: HashMap<String, ColumnCryptoConfig>,
}
impl Default for ColumnEncryptionPolicy {
fn default() -> Self {
Self::new()
}
}
impl ColumnEncryptionPolicy {
pub fn new() -> Self {
Self {
configs: HashMap::new(),
}
}
pub fn from_configs(configs: Vec<ColumnCryptoConfig>) -> Self {
let mut policy = Self::new();
for config in configs {
let _ = policy.add_column(config);
}
policy
}
pub fn add_column(&mut self, config: ColumnCryptoConfig) -> Result<(), String> {
if config.table.is_empty() || config.column.is_empty() {
return Err("table 和 column 不能为空".to_string());
}
let key = config.key();
self.configs.insert(key, config);
Ok(())
}
pub fn find(&self, table: &str, column: &str) -> Option<&ColumnCryptoConfig> {
self.configs.get(&format!("{}.{}", table, column))
}
pub fn remove(&mut self, table: &str, column: &str) -> Option<ColumnCryptoConfig> {
self.configs.remove(&format!("{}.{}", table, column))
}
pub fn len(&self) -> usize {
self.configs.len()
}
pub fn is_empty(&self) -> bool {
self.configs.is_empty()
}
pub fn is_encrypted(&self, table: &str, column: &str) -> bool {
self.find(table, column).is_some()
}
pub fn reload(&mut self, configs: Vec<ColumnCryptoConfig>) {
self.configs.clear();
for config in configs {
let _ = self.add_column(config);
}
}
pub fn all_configs(&self) -> Vec<&ColumnCryptoConfig> {
self.configs.values().collect()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn policy_add_and_find() {
let mut policy = ColumnEncryptionPolicy::new();
policy
.add_column(ColumnCryptoConfig::new("users", "ssn").with_key_version(2))
.unwrap();
let config = policy.find("users", "ssn").unwrap();
assert_eq!(config.table, "users");
assert_eq!(config.column, "ssn");
assert_eq!(config.key_version, 2);
assert_eq!(config.algorithm, EncryptionAlgo::Aes256Gcm);
}
#[test]
fn policy_not_found() {
let policy = ColumnEncryptionPolicy::new();
assert!(policy.find("users", "ssn").is_none());
assert!(!policy.is_encrypted("users", "ssn"));
}
#[test]
fn policy_reject_empty() {
let mut policy = ColumnEncryptionPolicy::new();
assert!(policy
.add_column(ColumnCryptoConfig::new("", "ssn"))
.is_err());
assert!(policy
.add_column(ColumnCryptoConfig::new("users", ""))
.is_err());
}
#[test]
fn policy_is_encrypted() {
let mut policy = ColumnEncryptionPolicy::new();
policy
.add_column(ColumnCryptoConfig::new("users", "ssn"))
.unwrap();
assert!(policy.is_encrypted("users", "ssn"));
assert!(!policy.is_encrypted("users", "name"));
}
#[test]
fn policy_remove() {
let mut policy = ColumnEncryptionPolicy::new();
policy
.add_column(ColumnCryptoConfig::new("users", "ssn"))
.unwrap();
assert_eq!(policy.len(), 1);
policy.remove("users", "ssn");
assert_eq!(policy.len(), 0);
}
#[test]
fn policy_reload() {
let mut policy = ColumnEncryptionPolicy::new();
policy
.add_column(ColumnCryptoConfig::new("users", "ssn"))
.unwrap();
assert_eq!(policy.len(), 1);
policy.reload(vec![
ColumnCryptoConfig::new("users", "email"),
ColumnCryptoConfig::new("orders", "credit_card"),
]);
assert_eq!(policy.len(), 2);
assert!(policy.is_encrypted("users", "email"));
assert!(policy.is_encrypted("orders", "credit_card"));
assert!(!policy.is_encrypted("users", "ssn"));
}
#[test]
fn policy_from_configs() {
let policy = ColumnEncryptionPolicy::from_configs(vec![
ColumnCryptoConfig::new("users", "ssn"),
ColumnCryptoConfig::new("orders", "card"),
]);
assert_eq!(policy.len(), 2);
assert!(policy.is_encrypted("users", "ssn"));
assert!(policy.is_encrypted("orders", "card"));
}
#[test]
fn policy_all_configs() {
let policy = ColumnEncryptionPolicy::from_configs(vec![
ColumnCryptoConfig::new("users", "ssn"),
ColumnCryptoConfig::new("orders", "card"),
]);
assert_eq!(policy.all_configs().len(), 2);
}
}