use sha2::{Digest, Sha256};
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum MaskStrategy {
Mask {
keep: usize,
},
Hash,
Truncate(usize),
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct MaskingRule {
pub column: String,
pub strategy: MaskStrategy,
}
#[derive(Debug, Clone, Default)]
pub struct MaskingEngine {
rules: Vec<MaskingRule>,
}
impl MaskingEngine {
pub fn new() -> Self {
Self { rules: Vec::new() }
}
pub fn rule(mut self, column: impl Into<String>, strategy: MaskStrategy) -> Self {
let column = column.into();
self.rules.retain(|r| r.column != column);
self.rules.push(MaskingRule { column, strategy });
self
}
pub fn is_empty(&self) -> bool {
self.rules.is_empty()
}
pub fn apply_value(
&self,
strategy: &MaskStrategy,
value: &serde_json::Value,
) -> serde_json::Value {
let s = match value {
serde_json::Value::String(s) => s.clone(),
serde_json::Value::Null => return serde_json::Value::Null,
other => other.to_string(),
};
let out = match strategy {
MaskStrategy::Mask { keep } => {
let chars: Vec<char> = s.chars().collect();
let total = chars.len();
let keep = (*keep).min(total / 2);
if total > keep * 2 {
let mut out2: String = chars[..keep].iter().collect();
for _ in 0..(total - keep * 2) {
out2.push('*');
}
out2.extend(chars[total - keep..].iter());
out2
} else {
"*".repeat(total)
}
}
MaskStrategy::Hash => {
let mut hasher = Sha256::new();
hasher.update(s.as_bytes());
let digest = hasher.finalize();
hex_encode(&digest)
}
MaskStrategy::Truncate(n) => {
let mut out2: String = s.chars().take(*n).collect();
if s.chars().count() > *n {
out2.push('…');
}
out2
}
};
serde_json::Value::String(out)
}
pub fn apply(&self, rows: &mut [serde_json::Value]) {
if self.rules.is_empty() {
return;
}
for row in rows.iter_mut() {
let Some(obj) = row.as_object_mut() else {
continue;
};
for rule in &self.rules {
if let Some(v) = obj.get_mut(&rule.column) {
let masked = self.apply_value(&rule.strategy, v);
*v = masked;
}
}
}
}
}
fn hex_encode(bytes: &[u8]) -> String {
let mut out = String::with_capacity(bytes.len() * 2);
for b in bytes {
out.push_str(&format!("{b:02x}"));
}
out
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct RlsPolicy {
pub table: String,
pub column: String,
pub value: String,
}
#[derive(Debug, Clone, Default)]
pub struct RlsEngine {
policies: Vec<RlsPolicy>,
}
impl RlsEngine {
pub fn new() -> Self {
Self {
policies: Vec::new(),
}
}
pub fn policy(
mut self,
table: impl Into<String>,
column: impl Into<String>,
value: impl Into<String>,
) -> Self {
self.policies.push(RlsPolicy {
table: table.into(),
column: column.into(),
value: value.into(),
});
self
}
pub fn is_empty(&self) -> bool {
self.policies.is_empty()
}
pub fn inject(&self, sql: &str, primary_table: Option<&str>) -> String {
if self.policies.is_empty() {
return sql.to_string();
}
let Some(table) = primary_table else {
return sql.to_string();
};
let Some(policy) = self.policies.iter().find(|p| p.table == table) else {
return sql.to_string();
};
let value = policy.value.replace('\'', "''");
let predicate = format!("{} = '{}'", policy.column, value);
let trimmed = sql.trim_end().trim_end_matches(';');
if contains_where(trimmed) {
format!("{trimmed} AND {predicate}")
} else {
format!("{trimmed} WHERE {predicate}")
}
}
}
fn contains_where(sql: &str) -> bool {
let lower = sql.to_ascii_lowercase();
lower.contains(" where ") || lower.starts_with("where ") || lower.contains(")where ")
}
use std::sync::Arc;
use tokio::sync::RwLock;
#[derive(Debug, Clone, Default)]
pub struct DataProtection {
pub masking: Option<Arc<MaskingEngine>>,
pub rls: Option<Arc<RlsEngine>>,
}
impl DataProtection {
pub fn load(arc: &Arc<RwLock<DataProtection>>) -> DataProtection {
arc.blocking_read().clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn test_masking_hash_strategy() {
let engine = MaskingEngine::new().rule("email", MaskStrategy::Hash);
let mut rows = vec![json!({"email": "alice@example.com"})];
engine.apply(&mut rows);
let masked = rows[0]["email"].as_str().unwrap();
assert_eq!(masked.len(), 64, "SHA-256 十六进制应为 64 字符");
assert!(!masked.contains("alice"));
assert!(!masked.contains('@'));
}
#[test]
fn test_masking_mask_and_truncate() {
let engine = MaskingEngine::new()
.rule("phone", MaskStrategy::Mask { keep: 3 })
.rule("bio", MaskStrategy::Truncate(5));
let mut rows = vec![json!({"phone": "13812345678", "bio": "很长的个人简介内容"})];
engine.apply(&mut rows);
let phone = rows[0]["phone"].as_str().unwrap();
assert!(phone.starts_with("138"), "应保留前 3 位");
assert!(phone.contains('*'), "中间应为掩码");
let bio = rows[0]["bio"].as_str().unwrap();
assert!(bio.chars().count() <= 6, "截断 5 字符 + 省略号");
assert!(bio.ends_with('…'));
}
#[test]
fn test_masking_untouched_columns() {
let engine = MaskingEngine::new().rule("secret", MaskStrategy::Hash);
let mut rows = vec![json!({"id": 1, "name": "keep"})];
engine.apply(&mut rows);
assert_eq!(rows[0]["name"], "keep");
assert_eq!(rows[0]["id"], 1);
}
#[test]
fn test_rls_injection() {
let rls = RlsEngine::new().policy("orders", "tenant_id", "t-100");
let out = rls.inject("SELECT * FROM orders", Some("orders"));
assert_eq!(out, "SELECT * FROM orders WHERE tenant_id = 't-100'");
let out2 = rls.inject("SELECT * FROM orders WHERE amount > 10", Some("orders"));
assert!(out2.ends_with("AND tenant_id = 't-100'"));
assert_eq!(
rls.inject("SELECT * FROM users", Some("users")),
"SELECT * FROM users"
);
assert_eq!(rls.inject("SELECT 1", None), "SELECT 1");
}
}