1use serde::{Deserialize, Serialize};
7use std::collections::HashMap;
8
9#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
11pub enum Dialect {
12 MySql,
14 PostgreSql,
16 Sqlite,
18 Oracle,
20 Mssql,
22}
23
24impl Dialect {
25 pub fn as_str(&self) -> &'static str {
27 match self {
28 Dialect::MySql => "mysql",
29 Dialect::PostgreSql => "postgresql",
30 Dialect::Sqlite => "sqlite",
31 Dialect::Oracle => "oracle",
32 Dialect::Mssql => "mssql",
33 }
34 }
35
36 pub fn supports_tls(&self) -> bool {
38 !matches!(self, Dialect::Sqlite)
39 }
40}
41
42#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
44pub enum CheckStatus {
45 Pass,
47 Fail,
49 Skipped,
51 NotApplicable,
53}
54
55#[derive(Debug, Clone, Serialize, Deserialize)]
57pub struct DialectSecurityConfig {
58 pub dialect: Dialect,
60 pub tls_enabled: bool,
62 pub auth_configured: bool,
64 pub conn_str_masked: bool,
66 pub pool_params_valid: bool,
68 pub available: bool,
70 pub skip_reason: Option<String>,
72}
73
74#[derive(Debug, Clone, Serialize, Deserialize)]
76pub struct DialectSecurityResult {
77 pub dialect: Dialect,
79 pub tls: CheckStatus,
81 pub auth: CheckStatus,
83 pub conn_str_masking: CheckStatus,
85 pub pool_params: CheckStatus,
87 pub evidence: Vec<String>,
89}
90
91#[derive(Debug, Clone, Serialize, Deserialize)]
93pub struct DialectSecurityReport {
94 pub results: Vec<DialectSecurityResult>,
96}
97
98impl DialectSecurityReport {
99 pub fn all_pass(&self) -> bool {
101 self.results.iter().all(|r| {
102 r.tls != CheckStatus::Fail
103 && r.auth != CheckStatus::Fail
104 && r.conn_str_masking != CheckStatus::Fail
105 && r.pool_params != CheckStatus::Fail
106 })
107 }
108}
109
110pub struct DialectSecurityVerifier {
112 configs: HashMap<Dialect, DialectSecurityConfig>,
113}
114
115impl DialectSecurityVerifier {
116 pub fn new(configs: Vec<DialectSecurityConfig>) -> Self {
118 let map = configs.into_iter().map(|c| (c.dialect, c)).collect();
119 Self { configs: map }
120 }
121
122 pub fn verify(&self) -> DialectSecurityReport {
124 let mut results = Vec::new();
125 for dialect in [
126 Dialect::MySql,
127 Dialect::PostgreSql,
128 Dialect::Sqlite,
129 Dialect::Oracle,
130 Dialect::Mssql,
131 ] {
132 results.push(self.verify_one(dialect));
133 }
134 DialectSecurityReport { results }
135 }
136
137 fn verify_one(&self, dialect: Dialect) -> DialectSecurityResult {
138 let config = self.configs.get(&dialect);
139 let mut evidence = Vec::new();
140
141 if let Some(cfg) = config {
142 if !cfg.available {
143 return DialectSecurityResult {
144 dialect,
145 tls: CheckStatus::Skipped,
146 auth: CheckStatus::Skipped,
147 conn_str_masking: CheckStatus::Skipped,
148 pool_params: CheckStatus::Skipped,
149 evidence: vec![cfg
150 .skip_reason
151 .clone()
152 .unwrap_or_else(|| "not available".into())],
153 };
154 }
155
156 let tls = if dialect.supports_tls() {
157 if cfg.tls_enabled {
158 evidence.push(format!("{} TLS enabled", dialect.as_str()));
159 CheckStatus::Pass
160 } else {
161 evidence.push(format!("{} TLS not enabled", dialect.as_str()));
162 CheckStatus::Fail
163 }
164 } else {
165 evidence.push(format!("{} TLS N/A (file-based)", dialect.as_str()));
166 CheckStatus::NotApplicable
167 };
168
169 let auth = if cfg.auth_configured {
170 evidence.push(format!("{} auth configured", dialect.as_str()));
171 CheckStatus::Pass
172 } else {
173 CheckStatus::Fail
174 };
175
176 let conn_str_masking = if cfg.conn_str_masked {
177 evidence.push(format!("{} conn_str masked", dialect.as_str()));
178 CheckStatus::Pass
179 } else {
180 evidence.push(format!(
181 "{} conn_str has plaintext password",
182 dialect.as_str()
183 ));
184 CheckStatus::Fail
185 };
186
187 let pool_params = if cfg.pool_params_valid {
188 evidence.push(format!("{} pool params valid", dialect.as_str()));
189 CheckStatus::Pass
190 } else {
191 CheckStatus::Fail
192 };
193
194 DialectSecurityResult {
195 dialect,
196 tls,
197 auth,
198 conn_str_masking,
199 pool_params,
200 evidence,
201 }
202 } else {
203 DialectSecurityResult {
204 dialect,
205 tls: CheckStatus::Skipped,
206 auth: CheckStatus::Skipped,
207 conn_str_masking: CheckStatus::Skipped,
208 pool_params: CheckStatus::Skipped,
209 evidence: vec!["no config provided".into()],
210 }
211 }
212 }
213}
214
215#[cfg(test)]
216mod tests {
217 use super::*;
218
219 #[test]
220 fn test_dialect_supports_tls() {
221 assert!(Dialect::MySql.supports_tls());
222 assert!(Dialect::PostgreSql.supports_tls());
223 assert!(!Dialect::Sqlite.supports_tls());
224 assert!(Dialect::Oracle.supports_tls());
225 assert!(Dialect::Mssql.supports_tls());
226 }
227
228 #[test]
229 fn test_verifier_all_pass() {
230 let configs = vec![
231 DialectSecurityConfig {
232 dialect: Dialect::MySql,
233 tls_enabled: true,
234 auth_configured: true,
235 conn_str_masked: true,
236 pool_params_valid: true,
237 available: true,
238 skip_reason: None,
239 },
240 DialectSecurityConfig {
241 dialect: Dialect::PostgreSql,
242 tls_enabled: true,
243 auth_configured: true,
244 conn_str_masked: true,
245 pool_params_valid: true,
246 available: true,
247 skip_reason: None,
248 },
249 DialectSecurityConfig {
250 dialect: Dialect::Sqlite,
251 tls_enabled: false,
252 auth_configured: true,
253 conn_str_masked: true,
254 pool_params_valid: true,
255 available: true,
256 skip_reason: None,
257 },
258 ];
259 let verifier = DialectSecurityVerifier::new(configs);
260 let report = verifier.verify();
261 assert!(report.all_pass());
262 let sqlite = report
263 .results
264 .iter()
265 .find(|r| r.dialect == Dialect::Sqlite)
266 .unwrap();
267 assert_eq!(sqlite.tls, CheckStatus::NotApplicable);
268 }
269
270 #[test]
271 fn test_verifier_skipped_for_unavailable() {
272 let configs = vec![DialectSecurityConfig {
273 dialect: Dialect::Mssql,
274 tls_enabled: false,
275 auth_configured: false,
276 conn_str_masked: false,
277 pool_params_valid: false,
278 available: false,
279 skip_reason: Some("MSSQL not installed".into()),
280 }];
281 let verifier = DialectSecurityVerifier::new(configs);
282 let report = verifier.verify();
283 let mssql = report
284 .results
285 .iter()
286 .find(|r| r.dialect == Dialect::Mssql)
287 .unwrap();
288 assert_eq!(mssql.tls, CheckStatus::Skipped);
289 assert!(mssql.evidence[0].contains("MSSQL not installed"));
290 }
291
292 #[test]
293 fn test_verifier_fail_for_plaintext_password() {
294 let configs = vec![DialectSecurityConfig {
295 dialect: Dialect::MySql,
296 tls_enabled: true,
297 auth_configured: true,
298 conn_str_masked: false,
299 pool_params_valid: true,
300 available: true,
301 skip_reason: None,
302 }];
303 let verifier = DialectSecurityVerifier::new(configs);
304 let report = verifier.verify();
305 assert!(!report.all_pass());
306 let mysql = report
307 .results
308 .iter()
309 .find(|r| r.dialect == Dialect::MySql)
310 .unwrap();
311 assert_eq!(mysql.conn_str_masking, CheckStatus::Fail);
312 }
313
314 #[test]
315 fn test_verifier_no_config_skipped() {
316 let verifier = DialectSecurityVerifier::new(vec![]);
317 let report = verifier.verify();
318 assert_eq!(report.results.len(), 5);
319 for r in &report.results {
320 assert_eq!(r.tls, CheckStatus::Skipped);
321 }
322 }
323}