Skip to main content

sequel_mcp/policy/
resolver.rs

1//! Two-layer policy resolution (connection baseline + table rules) with
2//! strictest-wins and elevation-requires-confirmation semantics.
3
4use 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/// Why a resolution denied the statement (audit/explain surface).
12#[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, // "read" | "mutated" | "locking-read"
24    pub category: SqlCategory,
25    pub action: PolicyAction,
26    pub baseline_action: PolicyAction,
27    pub rule: Option<String>, // governing rule rendered
28    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    /// Legacy-compat: distinct target databases considered.
39    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                // A table rule elevates what the baseline denies (or gates):
51                // eligible, but never silent — confirmation required.
52                (PolicyAction::Confirm, true)
53            } else {
54                (a, false)
55            }
56        }
57    }
58}
59
60/// Find the governing partial rule for a table: exact beats wildcard.
61fn 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            // Unqualified table without an established database: fail closed.
99            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    // Effective limits come from the baseline merged with the governing
118    // rule's overrides (legacy merge semantics).
119    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
143/// Resolve a classified statement against a connection's two-layer policy.
144///
145/// Rules (normative):
146/// - every read table is authorized for `read`; every mutated table for the
147///   statement's mutation category (`write`/`ddl`/`admin`);
148/// - locking reads additionally require `write` authorization on their
149///   tables;
150/// - file I/O is denied regardless of grants;
151/// - strictest action across all contributions wins;
152/// - unresolved (unqualified, no database) targets deny;
153/// - table-rule elevation of a stricter baseline always downgrades to
154///   `confirm`.
155pub 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            // Defensive: mutated tables under a read classification cannot
247            // happen via the classifier; deny if it ever does.
248            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            // Table-free statements (txCtrl, admin fast paths, bare SELECT 1):
286            // fall back to the baseline, matching the legacy resolver.
287            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
308/// Legacy-compat helper used by the database-policy wrapper tools and
309/// explain surfaces: strictest action across explicit databases with
310/// wildcard-table-rule overrides standing in for per-database policies.
311pub 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        // Write elevated to confirm on app.jobs; read of app.jobs denied.
496        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        // Read allow + write deny ⇒ deny.
556        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        // Ordinary reads on other tables stay prompt-free.
598        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        // Strictest across two databases.
621        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    /// Differential check against the legacy resolver fixtures (database
631    /// granularity, via wildcard rules standing in for databasePolicies).
632    #[test]
633    fn matches_legacy_resolver_fixtures() {
634        let path = concat!(
635            env!("CARGO_MANIFEST_DIR"),
636            "/tests/fixtures/legacy/resolver.json"
637        );
638        // Untracked corpus generated from the legacy checkout; skip on
639        // fresh CI checkouts where it is absent.
640        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}