use serde::{Deserialize, Serialize};
use std::collections::HashMap;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum Dialect {
MySql,
PostgreSql,
Sqlite,
Oracle,
Mssql,
}
impl Dialect {
pub fn as_str(&self) -> &'static str {
match self {
Dialect::MySql => "mysql",
Dialect::PostgreSql => "postgresql",
Dialect::Sqlite => "sqlite",
Dialect::Oracle => "oracle",
Dialect::Mssql => "mssql",
}
}
pub fn supports_tls(&self) -> bool {
!matches!(self, Dialect::Sqlite)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
pub enum CheckStatus {
Pass,
Fail,
Skipped,
NotApplicable,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DialectSecurityConfig {
pub dialect: Dialect,
pub tls_enabled: bool,
pub auth_configured: bool,
pub conn_str_masked: bool,
pub pool_params_valid: bool,
pub available: bool,
pub skip_reason: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DialectSecurityResult {
pub dialect: Dialect,
pub tls: CheckStatus,
pub auth: CheckStatus,
pub conn_str_masking: CheckStatus,
pub pool_params: CheckStatus,
pub evidence: Vec<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct DialectSecurityReport {
pub results: Vec<DialectSecurityResult>,
}
impl DialectSecurityReport {
pub fn all_pass(&self) -> bool {
self.results.iter().all(|r| {
r.tls != CheckStatus::Fail
&& r.auth != CheckStatus::Fail
&& r.conn_str_masking != CheckStatus::Fail
&& r.pool_params != CheckStatus::Fail
})
}
}
pub struct DialectSecurityVerifier {
configs: HashMap<Dialect, DialectSecurityConfig>,
}
impl DialectSecurityVerifier {
pub fn new(configs: Vec<DialectSecurityConfig>) -> Self {
let map = configs.into_iter().map(|c| (c.dialect, c)).collect();
Self { configs: map }
}
pub fn verify(&self) -> DialectSecurityReport {
let mut results = Vec::new();
for dialect in [
Dialect::MySql,
Dialect::PostgreSql,
Dialect::Sqlite,
Dialect::Oracle,
Dialect::Mssql,
] {
results.push(self.verify_one(dialect));
}
DialectSecurityReport { results }
}
fn verify_one(&self, dialect: Dialect) -> DialectSecurityResult {
let config = self.configs.get(&dialect);
let mut evidence = Vec::new();
if let Some(cfg) = config {
if !cfg.available {
return DialectSecurityResult {
dialect,
tls: CheckStatus::Skipped,
auth: CheckStatus::Skipped,
conn_str_masking: CheckStatus::Skipped,
pool_params: CheckStatus::Skipped,
evidence: vec![cfg
.skip_reason
.clone()
.unwrap_or_else(|| "not available".into())],
};
}
let tls = if dialect.supports_tls() {
if cfg.tls_enabled {
evidence.push(format!("{} TLS enabled", dialect.as_str()));
CheckStatus::Pass
} else {
evidence.push(format!("{} TLS not enabled", dialect.as_str()));
CheckStatus::Fail
}
} else {
evidence.push(format!("{} TLS N/A (file-based)", dialect.as_str()));
CheckStatus::NotApplicable
};
let auth = if cfg.auth_configured {
evidence.push(format!("{} auth configured", dialect.as_str()));
CheckStatus::Pass
} else {
CheckStatus::Fail
};
let conn_str_masking = if cfg.conn_str_masked {
evidence.push(format!("{} conn_str masked", dialect.as_str()));
CheckStatus::Pass
} else {
evidence.push(format!(
"{} conn_str has plaintext password",
dialect.as_str()
));
CheckStatus::Fail
};
let pool_params = if cfg.pool_params_valid {
evidence.push(format!("{} pool params valid", dialect.as_str()));
CheckStatus::Pass
} else {
CheckStatus::Fail
};
DialectSecurityResult {
dialect,
tls,
auth,
conn_str_masking,
pool_params,
evidence,
}
} else {
DialectSecurityResult {
dialect,
tls: CheckStatus::Skipped,
auth: CheckStatus::Skipped,
conn_str_masking: CheckStatus::Skipped,
pool_params: CheckStatus::Skipped,
evidence: vec!["no config provided".into()],
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_dialect_supports_tls() {
assert!(Dialect::MySql.supports_tls());
assert!(Dialect::PostgreSql.supports_tls());
assert!(!Dialect::Sqlite.supports_tls());
assert!(Dialect::Oracle.supports_tls());
assert!(Dialect::Mssql.supports_tls());
}
#[test]
fn test_verifier_all_pass() {
let configs = vec![
DialectSecurityConfig {
dialect: Dialect::MySql,
tls_enabled: true,
auth_configured: true,
conn_str_masked: true,
pool_params_valid: true,
available: true,
skip_reason: None,
},
DialectSecurityConfig {
dialect: Dialect::PostgreSql,
tls_enabled: true,
auth_configured: true,
conn_str_masked: true,
pool_params_valid: true,
available: true,
skip_reason: None,
},
DialectSecurityConfig {
dialect: Dialect::Sqlite,
tls_enabled: false,
auth_configured: true,
conn_str_masked: true,
pool_params_valid: true,
available: true,
skip_reason: None,
},
];
let verifier = DialectSecurityVerifier::new(configs);
let report = verifier.verify();
assert!(report.all_pass());
let sqlite = report
.results
.iter()
.find(|r| r.dialect == Dialect::Sqlite)
.unwrap();
assert_eq!(sqlite.tls, CheckStatus::NotApplicable);
}
#[test]
fn test_verifier_skipped_for_unavailable() {
let configs = vec![DialectSecurityConfig {
dialect: Dialect::Mssql,
tls_enabled: false,
auth_configured: false,
conn_str_masked: false,
pool_params_valid: false,
available: false,
skip_reason: Some("MSSQL not installed".into()),
}];
let verifier = DialectSecurityVerifier::new(configs);
let report = verifier.verify();
let mssql = report
.results
.iter()
.find(|r| r.dialect == Dialect::Mssql)
.unwrap();
assert_eq!(mssql.tls, CheckStatus::Skipped);
assert!(mssql.evidence[0].contains("MSSQL not installed"));
}
#[test]
fn test_verifier_fail_for_plaintext_password() {
let configs = vec![DialectSecurityConfig {
dialect: Dialect::MySql,
tls_enabled: true,
auth_configured: true,
conn_str_masked: false,
pool_params_valid: true,
available: true,
skip_reason: None,
}];
let verifier = DialectSecurityVerifier::new(configs);
let report = verifier.verify();
assert!(!report.all_pass());
let mysql = report
.results
.iter()
.find(|r| r.dialect == Dialect::MySql)
.unwrap();
assert_eq!(mysql.conn_str_masking, CheckStatus::Fail);
}
#[test]
fn test_verifier_no_config_skipped() {
let verifier = DialectSecurityVerifier::new(vec![]);
let report = verifier.verify();
assert_eq!(report.results.len(), 5);
for r in &report.results {
assert_eq!(r.tls, CheckStatus::Skipped);
}
}
}