1use crate::config::Connection;
5use crate::policy::classifier::{ClassifiedStatement, TableRef};
6use crate::policy::model::{
7 PartialPolicy, Policy, PolicyAction, SqlCategory, TableId, TableRuleKey,
8};
9use std::collections::BTreeMap;
10
11#[derive(Debug, Clone, PartialEq)]
13pub enum DenyReason {
14 FileIo,
15 UnresolvedTarget { table: String },
16 NoTargetsEstablished,
17 Policy,
18}
19
20#[derive(Debug, Clone, PartialEq)]
21pub struct TableContribution {
22 pub table: TableId,
23 pub kind: &'static str, pub category: SqlCategory,
25 pub action: PolicyAction,
26 pub baseline_action: PolicyAction,
27 pub rule: Option<String>, pub elevated: bool,
29}
30
31#[derive(Debug, Clone, PartialEq)]
32pub struct Resolution {
33 pub action: PolicyAction,
34 pub effective: Policy,
35 pub contributions: Vec<TableContribution>,
36 pub elevated: bool,
37 pub deny_reason: Option<DenyReason>,
38 pub contributing_databases: Vec<String>,
40}
41
42fn resolve_table_action(
43 baseline: PolicyAction,
44 rule: Option<PolicyAction>,
45) -> (PolicyAction, bool) {
46 match rule {
47 None => (baseline, false),
48 Some(a) => {
49 if a == PolicyAction::Allow && baseline != PolicyAction::Allow {
50 (PolicyAction::Confirm, true)
53 } else {
54 (a, false)
55 }
56 }
57 }
58}
59
60fn governing_rule<'a>(
62 table_policies: &'a BTreeMap<TableRuleKey, PartialPolicy>,
63 table: &TableId,
64) -> Option<(&'a PartialPolicy, String)> {
65 if let Some(p) = table_policies.get(&TableRuleKey::Exact(table.clone())) {
66 return Some((p, TableRuleKey::Exact(table.clone()).render()));
67 }
68 if let Some(p) = table_policies.get(&TableRuleKey::Wildcard {
69 database: table.database.clone(),
70 }) {
71 return Some((
72 p,
73 TableRuleKey::Wildcard {
74 database: table.database.clone(),
75 }
76 .render(),
77 ));
78 }
79 None
80}
81
82#[allow(clippy::too_many_arguments)]
83fn resolve_object(
84 _conn: &Connection,
85 baseline: &Policy,
86 table_policies: &BTreeMap<TableRuleKey, PartialPolicy>,
87 category: SqlCategory,
88 table: &TableRef,
89 fallback_db: Option<&str>,
90 kind: &'static str,
91 contributions: &mut Vec<TableContribution>,
92 strictest: &mut Option<(PolicyAction, Policy, TableContribution)>,
93) -> Result<(), DenyReason> {
94 let database = match &table.database {
95 Some(db) => db.clone(),
96 None => match fallback_db {
97 Some(db) => db.to_string(),
98 None => {
100 return Err(DenyReason::UnresolvedTarget {
101 table: table.table.clone(),
102 });
103 }
104 },
105 };
106 let id = TableId {
107 database,
108 table: table.table.clone(),
109 };
110 let base_action = baseline.action_for(category);
111 let (partial, rule_name) = match governing_rule(table_policies, &id) {
112 Some((p, name)) => (Some(p), Some(name)),
113 None => (None, None),
114 };
115 let rule_action = partial.and_then(|p| p.action_for(category));
116 let (action, elevated) = resolve_table_action(base_action, rule_action);
117 let effective = match partial {
120 Some(p) => baseline.merged_with(p),
121 None => baseline.clone(),
122 };
123 let contribution = TableContribution {
124 table: id,
125 kind,
126 category,
127 action,
128 baseline_action: base_action,
129 rule: rule_name,
130 elevated,
131 };
132 let better = match strictest {
133 None => true,
134 Some((a, _, _)) => action.rank() > a.rank(),
135 };
136 if better {
137 *strictest = Some((action, effective, contribution.clone()));
138 }
139 contributions.push(contribution);
140 Ok(())
141}
142
143pub fn resolve(
156 conn: &Connection,
157 classified: &ClassifiedStatement,
158 fallback_database: Option<&str>,
159) -> Resolution {
160 let baseline = conn.policy();
161 let table_policies = conn.table_policies();
162 let fallback = fallback_database.or_else(|| conn.database());
163
164 if classified.file_io {
165 return Resolution {
166 action: PolicyAction::Deny,
167 effective: baseline.clone(),
168 contributions: Vec::new(),
169 elevated: false,
170 deny_reason: Some(DenyReason::FileIo),
171 contributing_databases: Vec::new(),
172 };
173 }
174
175 let mut contributions = Vec::new();
176 let mut strictest: Option<(PolicyAction, Policy, TableContribution)> = None;
177 let mut elevated_any = false;
178 let mut dbs: Vec<String> = Vec::new();
179
180 let consider = |category: SqlCategory,
181 table: &TableRef,
182 kind: &'static str,
183 contributions: &mut Vec<TableContribution>,
184 strictest: &mut Option<(PolicyAction, Policy, TableContribution)>,
185 elevated_any: &mut bool,
186 dbs: &mut Vec<String>|
187 -> Result<(), DenyReason> {
188 let res = resolve_object(
189 conn,
190 baseline,
191 table_policies,
192 category,
193 table,
194 fallback,
195 kind,
196 contributions,
197 strictest,
198 );
199 if let Ok(c) = res.as_ref() {
200 let _ = c;
201 }
202 if res.is_ok()
203 && let Some(c) = contributions.last()
204 {
205 if c.elevated {
206 *elevated_any = true;
207 }
208 if !dbs.contains(&c.table.database) {
209 dbs.push(c.table.database.clone());
210 }
211 }
212 res
213 };
214
215 let mutation_category = classified.category;
216 let mut denied: Option<DenyReason> = None;
217
218 for t in &classified.read_tables {
219 if let Err(e) = consider(
220 SqlCategory::Read,
221 t,
222 "read",
223 &mut contributions,
224 &mut strictest,
225 &mut elevated_any,
226 &mut dbs,
227 ) {
228 denied = denied.or(Some(e));
229 }
230 if classified.locking_read
231 && let Err(e) = consider(
232 SqlCategory::Write,
233 t,
234 "locking-read",
235 &mut contributions,
236 &mut strictest,
237 &mut elevated_any,
238 &mut dbs,
239 )
240 {
241 denied = denied.or(Some(e));
242 }
243 }
244 for t in &classified.mutated_tables {
245 if mutation_category == SqlCategory::Read {
246 denied = denied.or(Some(DenyReason::Policy));
249 continue;
250 }
251 if let Err(e) = consider(
252 mutation_category,
253 t,
254 "mutated",
255 &mut contributions,
256 &mut strictest,
257 &mut elevated_any,
258 &mut dbs,
259 ) {
260 denied = denied.or(Some(e));
261 }
262 }
263
264 if let Some(reason) = denied {
265 return Resolution {
266 action: PolicyAction::Deny,
267 effective: baseline.clone(),
268 contributions,
269 elevated: elevated_any,
270 deny_reason: Some(reason),
271 contributing_databases: dbs,
272 };
273 }
274
275 match strictest {
276 Some((action, effective, _)) => Resolution {
277 action,
278 effective,
279 contributions,
280 elevated: elevated_any,
281 deny_reason: None,
282 contributing_databases: dbs,
283 },
284 None => {
285 let action = baseline.action_for(mutation_category);
288 let mut dbs = Vec::new();
289 if classified.target_databases.is_empty() {
290 if let Some(f) = fallback {
291 dbs.push(f.to_string());
292 }
293 } else {
294 dbs.extend(classified.target_databases.iter().cloned());
295 }
296 Resolution {
297 action,
298 effective: baseline.clone(),
299 contributions: Vec::new(),
300 elevated: false,
301 deny_reason: None,
302 contributing_databases: dbs,
303 }
304 }
305 }
306}
307
308pub fn resolve_databases_legacy(
312 conn: &Connection,
313 category: SqlCategory,
314 target_databases: &[String],
315 fallback_database: Option<&str>,
316) -> Resolution {
317 let baseline = conn.policy();
318 let table_policies = conn.table_policies();
319 let mut dbs: Vec<String> = if target_databases.is_empty() {
320 match fallback_database.or_else(|| conn.database()) {
321 Some(db) => vec![db.to_string()],
322 None => Vec::new(),
323 }
324 } else {
325 target_databases.to_vec()
326 };
327 dbs.sort();
328 dbs.dedup();
329
330 if dbs.is_empty() {
331 return Resolution {
332 action: baseline.action_for(category),
333 effective: baseline.clone(),
334 contributions: Vec::new(),
335 elevated: false,
336 deny_reason: None,
337 contributing_databases: Vec::new(),
338 };
339 }
340
341 let mut contributions = Vec::new();
342 let mut strictest: Option<(PolicyAction, Policy)> = None;
343 for db in &dbs {
344 let id = TableId {
345 database: db.clone(),
346 table: "*".to_string(),
347 };
348 let base_action = baseline.action_for(category);
349 let (partial, rule_name) = match governing_rule(table_policies, &id) {
350 Some((p, name)) => (Some(p), Some(name)),
351 None => (None, None),
352 };
353 let rule_action = partial.and_then(|p| p.action_for(category));
354 let (action, elevated) = resolve_table_action(base_action, rule_action);
355 let effective = match partial {
356 Some(p) => baseline.merged_with(p),
357 None => baseline.clone(),
358 };
359 contributions.push(TableContribution {
360 table: id,
361 kind: "database",
362 category,
363 action,
364 baseline_action: base_action,
365 rule: rule_name,
366 elevated,
367 });
368 let better = match &strictest {
369 None => true,
370 Some((a, _)) => action.rank() > a.rank(),
371 };
372 if better {
373 strictest = Some((action, effective));
374 }
375 }
376
377 let (action, effective) =
378 strictest.unwrap_or_else(|| (baseline.action_for(category), baseline.clone()));
379 let elevated = contributions.iter().any(|c| c.elevated);
380 Resolution {
381 action,
382 effective,
383 contributions,
384 elevated,
385 deny_reason: None,
386 contributing_databases: dbs,
387 }
388}
389
390#[cfg(test)]
391mod tests {
392 use super::*;
393 use crate::config::{MySqlConnection, SqliteConnection};
394 use crate::policy::classifier::classify_statement;
395 use crate::policy::model::{PolicyPresetName, policy_from_preset};
396
397 fn mysql_conn(policy: Policy, rules: Vec<(TableRuleKey, PartialPolicy)>) -> Connection {
398 let mut c = MySqlConnection {
399 name: "c1".into(),
400 host: "db.example.invalid".into(),
401 user: "u".into(),
402 ..MySqlConnection::default()
403 };
404 c.policy = policy;
405 for (k, v) in rules {
406 c.table_policies.insert(k, v);
407 }
408 Connection::Mysql(c)
409 }
410
411 fn sqlite_conn(rules: Vec<(TableRuleKey, PartialPolicy)>) -> Connection {
412 let mut c = SqliteConnection {
413 name: "s1".into(),
414 path: "/tmp/app.sqlite".into(),
415 ..SqliteConnection::default()
416 };
417 c.policy = policy_from_preset(PolicyPresetName::ReadOnly);
418 for (k, v) in rules {
419 c.table_policies.insert(k, v);
420 }
421 Connection::Sqlite(c)
422 }
423
424 fn write_allow() -> PartialPolicy {
425 PartialPolicy {
426 write: Some(PolicyAction::Allow),
427 ..PartialPolicy::default()
428 }
429 }
430
431 fn resolve_sql(conn: &Connection, sql: &str) -> Resolution {
432 let c = classify_statement(sql, crate::policy::classifier::Dialect::MySql).unwrap();
433 resolve(conn, &c, None)
434 }
435
436 #[test]
437 fn read_only_baseline_allows_reads_denies_writes() {
438 let conn = mysql_conn(policy_from_preset(PolicyPresetName::ReadOnly), vec![]);
439 assert_eq!(
440 resolve_sql(&conn, "SELECT * FROM app.users").action,
441 PolicyAction::Allow
442 );
443 let r = resolve_sql(&conn, "UPDATE app.jobs SET state = 'x' WHERE id = 1");
444 assert_eq!(r.action, PolicyAction::Deny);
445 }
446
447 #[test]
448 fn table_elevation_requires_confirmation() {
449 let conn = mysql_conn(
450 policy_from_preset(PolicyPresetName::ReadOnly),
451 vec![(TableRuleKey::parse("app.jobs").unwrap(), write_allow())],
452 );
453 let r = resolve_sql(&conn, "UPDATE app.jobs SET state = 'x' WHERE id = 1");
454 assert_eq!(r.action, PolicyAction::Confirm);
455 assert!(r.elevated);
456 assert_eq!(r.contributions.len(), 1);
457 assert_eq!(r.contributions[0].rule.as_deref(), Some("app.jobs"));
458 }
459
460 #[test]
461 fn cross_table_strictest_wins() {
462 let conn = mysql_conn(
463 policy_from_preset(PolicyPresetName::ReadOnly),
464 vec![
465 (TableRuleKey::parse("app.jobs").unwrap(), write_allow()),
466 (TableRuleKey::parse("app.users").unwrap(), {
467 PartialPolicy {
468 write: Some(PolicyAction::Deny),
469 ..PartialPolicy::default()
470 }
471 }),
472 ],
473 );
474 let r = resolve_sql(
475 &conn,
476 "UPDATE app.jobs JOIN app.users ON app.users.id = app.jobs.user_id SET app.jobs.state = 'ok'",
477 );
478 assert_eq!(r.action, PolicyAction::Deny);
479 assert_eq!(r.contributions.len(), 2);
480 }
481
482 #[test]
483 fn insert_select_authorizes_source_as_read() {
484 let conn = mysql_conn(
485 policy_from_preset(PolicyPresetName::ReadOnly),
486 vec![(
487 TableRuleKey::parse("app.jobs").unwrap(),
488 PartialPolicy {
489 write: Some(PolicyAction::Allow),
490 read: Some(PolicyAction::Deny),
491 ..PartialPolicy::default()
492 },
493 )],
494 );
495 let r = resolve_sql(&conn, "INSERT INTO app.jobs (id) SELECT id FROM app.jobs");
497 assert_eq!(r.action, PolicyAction::Deny);
498 }
499
500 #[test]
501 fn exact_rule_beats_wildcard() {
502 let conn = mysql_conn(
503 policy_from_preset(PolicyPresetName::ReadOnly),
504 vec![
505 (TableRuleKey::parse("app.*").unwrap(), write_allow()),
506 (
507 TableRuleKey::parse("app.audit").unwrap(),
508 PartialPolicy {
509 write: Some(PolicyAction::Deny),
510 ..PartialPolicy::default()
511 },
512 ),
513 ],
514 );
515 let elevated = resolve_sql(&conn, "UPDATE app.jobs SET x = 1");
516 assert_eq!(elevated.action, PolicyAction::Confirm);
517 let blocked = resolve_sql(&conn, "UPDATE app.audit SET x = 1");
518 assert_eq!(blocked.action, PolicyAction::Deny);
519 }
520
521 #[test]
522 fn unqualified_without_fallback_denies() {
523 let conn = mysql_conn(policy_from_preset(PolicyPresetName::ReadOnly), vec![]);
524 let r = resolve_sql(&conn, "UPDATE jobs SET x = 1");
525 assert_eq!(r.action, PolicyAction::Deny);
526 assert!(matches!(
527 r.deny_reason,
528 Some(DenyReason::UnresolvedTarget { .. })
529 ));
530 }
531
532 #[test]
533 fn unqualified_uses_connection_database() {
534 let mut c = MySqlConnection {
535 name: "c1".into(),
536 host: "db.example.invalid".into(),
537 user: "u".into(),
538 database: Some("app".into()),
539 ..MySqlConnection::default()
540 };
541 c.policy = policy_from_preset(PolicyPresetName::ReadOnly);
542 {
543 let (k, v) = (TableRuleKey::parse("app.jobs").unwrap(), write_allow());
544 c.table_policies.insert(k, v);
545 }
546 let conn = Connection::Mysql(c);
547 let r = resolve_sql(&conn, "UPDATE jobs SET x = 1");
548 assert_eq!(r.action, PolicyAction::Confirm);
549 }
550
551 #[test]
552 fn locking_reads_need_write_authorization() {
553 let conn = mysql_conn(policy_from_preset(PolicyPresetName::ReadOnly), vec![]);
554 let r = resolve_sql(&conn, "SELECT * FROM app.users FOR UPDATE");
555 assert_eq!(r.action, PolicyAction::Deny);
557 }
558
559 #[test]
560 fn file_io_always_denied() {
561 let mut conn = mysql_conn(Policy::default(), vec![]);
562 if let Connection::Mysql(ref mut m) = conn {
563 m.policy.admin = PolicyAction::Allow;
564 m.policy.write = PolicyAction::Allow;
565 }
566 let r = resolve_sql(&conn, "SELECT * FROM users INTO OUTFILE '/tmp/x'");
567 assert_eq!(r.action, PolicyAction::Deny);
568 assert_eq!(r.deny_reason, Some(DenyReason::FileIo));
569 }
570
571 #[test]
572 fn table_free_statements_use_baseline() {
573 let conn = mysql_conn(policy_from_preset(PolicyPresetName::ReadOnly), vec![]);
574 assert_eq!(resolve_sql(&conn, "BEGIN").action, PolicyAction::Allow);
575 let mut dev = policy_from_preset(PolicyPresetName::Development);
576 dev.admin = PolicyAction::Deny;
577 let conn2 = mysql_conn(dev, vec![]);
578 assert_eq!(resolve_sql(&conn2, "KILL 42").action, PolicyAction::Deny);
579 }
580
581 #[test]
582 fn sqlite_read_confirm_rule_prompts_reads() {
583 let conn = sqlite_conn(vec![(
584 TableRuleKey::parse("main.secrets").unwrap(),
585 PartialPolicy {
586 read: Some(PolicyAction::Confirm),
587 ..PartialPolicy::default()
588 },
589 )]);
590 let c = classify_statement(
591 "SELECT * FROM secrets",
592 crate::policy::classifier::Dialect::SQLite,
593 )
594 .unwrap();
595 let r = resolve(&conn, &c, None);
596 assert_eq!(r.action, PolicyAction::Confirm);
597 let c2 = classify_statement(
599 "SELECT * FROM main.users",
600 crate::policy::classifier::Dialect::SQLite,
601 )
602 .unwrap();
603 assert_eq!(resolve(&conn, &c2, None).action, PolicyAction::Allow);
604 }
605
606 #[test]
607 fn legacy_database_resolution_via_wildcards() {
608 let conn = mysql_conn(
609 policy_from_preset(PolicyPresetName::ReadOnly),
610 vec![(
611 TableRuleKey::parse("app.*").unwrap(),
612 PartialPolicy {
613 write: Some(PolicyAction::Confirm),
614 ..PartialPolicy::default()
615 },
616 )],
617 );
618 let r = resolve_databases_legacy(&conn, SqlCategory::Write, &["app".into()], None);
619 assert_eq!(r.action, PolicyAction::Confirm);
620 let r2 = resolve_databases_legacy(
622 &conn,
623 SqlCategory::Write,
624 &["app".into(), "analytics".into()],
625 None,
626 );
627 assert_eq!(r2.action, PolicyAction::Deny);
628 }
629
630 #[test]
633 fn matches_legacy_resolver_fixtures() {
634 let path = concat!(
635 env!("CARGO_MANIFEST_DIR"),
636 "/tests/fixtures/legacy/resolver.json"
637 );
638 let Ok(data) = std::fs::read_to_string(path) else {
641 eprintln!("skipping: legacy fixture corpus not present ({path})");
642 return;
643 };
644 let cases: Vec<serde_json::Value> = serde_json::from_str(&data).unwrap();
645 for case in cases {
646 let label = case["label"].as_str().unwrap();
647 let args = &case["args"];
648 let conn_json = &args["connection"];
649 let mut rules = Vec::new();
650 if let Some(dbs) = conn_json["databasePolicies"].as_object() {
651 for (db, partial) in dbs {
652 let p: PartialPolicy = serde_json::from_value(partial.clone()).unwrap();
653 rules.push((
654 TableRuleKey::Wildcard {
655 database: db.clone(),
656 },
657 p,
658 ));
659 }
660 }
661 let mut m = MySqlConnection {
662 name: "c1".into(),
663 host: "db.example.invalid".into(),
664 user: "u".into(),
665 database: conn_json["database"].as_str().map(str::to_string),
666 ..MySqlConnection::default()
667 };
668 m.policy = serde_json::from_value(conn_json["policy"].clone()).unwrap();
669 for (k, v) in rules {
670 m.table_policies.insert(k, v);
671 }
672 let conn = Connection::Mysql(m);
673 let category = SqlCategory::parse(args["category"].as_str().unwrap()).unwrap();
674 let targets: Vec<String> = args["targetDatabases"]
675 .as_array()
676 .unwrap()
677 .iter()
678 .map(|v| v.as_str().unwrap().to_string())
679 .collect();
680 let fallback = args["fallbackDatabase"].as_str();
681 let r = resolve_databases_legacy(&conn, category, &targets, fallback);
682 let want = case["result"]["action"].as_str().unwrap();
683 assert_eq!(r.action.as_str(), want, "case {label}");
684 }
685 }
686}