Skip to main content

sequel_mcp/policy/
model.rs

1//! Policy model: categories, actions, connection baselines, table rules.
2//!
3//! Mirrors the legacy TypeScript schema (defaults included) and adds the
4//! v2 table-rule layer. Validation bounds match `types.ts` exactly.
5
6use serde::{Deserialize, Serialize};
7use std::collections::BTreeMap;
8
9pub const SQL_CATEGORIES: [&str; 5] = ["read", "write", "ddl", "admin", "txCtrl"];
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
12#[serde(rename_all = "camelCase")]
13pub enum SqlCategory {
14    Read,
15    Write,
16    Ddl,
17    Admin,
18    TxCtrl,
19}
20
21impl std::fmt::Display for SqlCategory {
22    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
23        f.write_str(self.as_str())
24    }
25}
26
27impl SqlCategory {
28    pub fn as_str(&self) -> &'static str {
29        match self {
30            SqlCategory::Read => "read",
31            SqlCategory::Write => "write",
32            SqlCategory::Ddl => "ddl",
33            SqlCategory::Admin => "admin",
34            SqlCategory::TxCtrl => "txCtrl",
35        }
36    }
37
38    pub fn parse(s: &str) -> Option<Self> {
39        match s {
40            "read" => Some(SqlCategory::Read),
41            "write" => Some(SqlCategory::Write),
42            "ddl" => Some(SqlCategory::Ddl),
43            "admin" => Some(SqlCategory::Admin),
44            "txCtrl" => Some(SqlCategory::TxCtrl),
45            _ => None,
46        }
47    }
48}
49
50#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)]
51#[serde(rename_all = "camelCase")]
52pub enum PolicyAction {
53    Allow,
54    Confirm,
55    Deny,
56}
57
58impl PolicyAction {
59    /// Strictness ranking used by strictest-wins resolution.
60    pub fn rank(&self) -> u8 {
61        match self {
62            PolicyAction::Allow => 0,
63            PolicyAction::Confirm => 1,
64            PolicyAction::Deny => 2,
65        }
66    }
67
68    pub fn as_str(&self) -> &'static str {
69        match self {
70            PolicyAction::Allow => "allow",
71            PolicyAction::Confirm => "confirm",
72            PolicyAction::Deny => "deny",
73        }
74    }
75}
76
77#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)]
78#[serde(rename_all = "camelCase")]
79pub enum BackupOverflow {
80    Abort,
81    Truncate,
82}
83
84/// Complete connection baseline. Field defaults replicate the legacy zod
85/// schema so a v1 policy object with omitted fields parses identically.
86#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
87#[serde(rename_all = "camelCase", default)]
88pub struct Policy {
89    pub read: PolicyAction,
90    pub write: PolicyAction,
91    pub ddl: PolicyAction,
92    pub admin: PolicyAction,
93    pub tx_ctrl: PolicyAction,
94    pub row_cap: u32,
95    pub stmt_timeout_ms: u32,
96    pub require_touch_id: bool,
97    pub max_backup_rows: u32,
98    pub max_backup_bytes: u64,
99    pub on_backup_overflow: BackupOverflow,
100}
101
102impl Default for Policy {
103    fn default() -> Self {
104        Self {
105            read: PolicyAction::Allow,
106            write: PolicyAction::Confirm,
107            ddl: PolicyAction::Deny,
108            admin: PolicyAction::Deny,
109            tx_ctrl: PolicyAction::Allow,
110            row_cap: 1000,
111            stmt_timeout_ms: 10_000,
112            require_touch_id: false,
113            max_backup_rows: 10_000,
114            max_backup_bytes: 50 * 1024 * 1024,
115            on_backup_overflow: BackupOverflow::Abort,
116        }
117    }
118}
119
120impl Policy {
121    pub fn action_for(&self, category: SqlCategory) -> PolicyAction {
122        match category {
123            SqlCategory::Read => self.read,
124            SqlCategory::Write => self.write,
125            SqlCategory::Ddl => self.ddl,
126            SqlCategory::Admin => self.admin,
127            SqlCategory::TxCtrl => self.tx_ctrl,
128        }
129    }
130
131    /// Legacy semantics: a partial override merges field-by-field onto the
132    /// baseline. `None` fields fall through to `self`.
133    pub fn merged_with(&self, over: &PartialPolicy) -> Policy {
134        let mut next = self.clone();
135        if let Some(v) = over.read {
136            next.read = v;
137        }
138        if let Some(v) = over.write {
139            next.write = v;
140        }
141        if let Some(v) = over.ddl {
142            next.ddl = v;
143        }
144        if let Some(v) = over.admin {
145            next.admin = v;
146        }
147        if let Some(v) = over.tx_ctrl {
148            next.tx_ctrl = v;
149        }
150        if let Some(v) = over.row_cap {
151            next.row_cap = v;
152        }
153        if let Some(v) = over.stmt_timeout_ms {
154            next.stmt_timeout_ms = v;
155        }
156        if let Some(v) = over.require_touch_id {
157            next.require_touch_id = v;
158        }
159        if let Some(v) = over.max_backup_rows {
160            next.max_backup_rows = v;
161        }
162        if let Some(v) = over.max_backup_bytes {
163            next.max_backup_bytes = v;
164        }
165        if let Some(v) = over.on_backup_overflow {
166            next.on_backup_overflow = v;
167        }
168        next
169    }
170
171    /// Validate against the legacy bounds. Errors carry the offending field.
172    pub fn validate(&self) -> Result<(), &'static str> {
173        if self.row_cap == 0 || self.row_cap > 100_000 {
174            return Err("rowCap must be an integer in 1..=100000");
175        }
176        if self.stmt_timeout_ms == 0 || self.stmt_timeout_ms > 600_000 {
177            return Err("stmtTimeoutMs must be an integer in 1..=600000");
178        }
179        if self.max_backup_rows == 0 || self.max_backup_rows > 1_000_000 {
180            return Err("maxBackupRows must be an integer in 1..=1000000");
181        }
182        if self.max_backup_bytes == 0 {
183            return Err("maxBackupBytes must be positive");
184        }
185        Ok(())
186    }
187}
188
189/// Partial policy override (legacy `PartialPolicySchema` + table rules).
190#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize, schemars::JsonSchema)]
191#[serde(rename_all = "camelCase", default)]
192pub struct PartialPolicy {
193    pub read: Option<PolicyAction>,
194    pub write: Option<PolicyAction>,
195    pub ddl: Option<PolicyAction>,
196    pub admin: Option<PolicyAction>,
197    pub tx_ctrl: Option<PolicyAction>,
198    pub row_cap: Option<u32>,
199    pub stmt_timeout_ms: Option<u32>,
200    pub require_touch_id: Option<bool>,
201    pub max_backup_rows: Option<u32>,
202    pub max_backup_bytes: Option<u64>,
203    pub on_backup_overflow: Option<BackupOverflow>,
204}
205
206impl PartialPolicy {
207    /// Review finding: table rules previously bypassed ALL bounds
208    /// validation — a rule could set `rowCap: 0` or a 71-minute
209    /// `stmtTimeoutMs` that the baseline `Policy::validate` would never
210    /// allow. Every present field must sit inside the same bounds.
211    pub fn validate(&self) -> Result<(), &'static str> {
212        if let Some(v) = self.row_cap
213            && (v == 0 || v > 100_000)
214        {
215            return Err("rowCap must be an integer in 1..=100000");
216        }
217        if let Some(v) = self.stmt_timeout_ms
218            && (v == 0 || v > 600_000)
219        {
220            return Err("stmtTimeoutMs must be an integer in 1..=600000");
221        }
222        if let Some(v) = self.max_backup_rows
223            && (v == 0 || v > 1_000_000)
224        {
225            return Err("maxBackupRows must be an integer in 1..=1000000");
226        }
227        if let Some(v) = self.max_backup_bytes
228            && v == 0
229        {
230            return Err("maxBackupBytes must be positive");
231        }
232        Ok(())
233    }
234}
235
236impl PartialPolicy {
237    pub fn action_for(&self, category: SqlCategory) -> Option<PolicyAction> {
238        match category {
239            SqlCategory::Read => self.read,
240            SqlCategory::Write => self.write,
241            SqlCategory::Ddl => self.ddl,
242            SqlCategory::Admin => self.admin,
243            SqlCategory::TxCtrl => self.tx_ctrl,
244        }
245    }
246
247    pub fn is_empty(&self) -> bool {
248        *self == PartialPolicy::default()
249    }
250}
251
252/// Named presets. `development` is the new name for legacy `dev`; both parse.
253#[derive(Debug, Clone, Copy, PartialEq, Eq)]
254pub enum PolicyPresetName {
255    ReadOnly,
256    Development,
257    Administration,
258}
259
260impl PolicyPresetName {
261    pub fn parse(s: &str) -> Option<Self> {
262        match s {
263            "read-only" => Some(Self::ReadOnly),
264            "dev" | "development" => Some(Self::Development),
265            "admin" | "administration" => Some(Self::Administration),
266            _ => None,
267        }
268    }
269
270    pub fn as_str(&self) -> &'static str {
271        match self {
272            PolicyPresetName::ReadOnly => "read-only",
273            PolicyPresetName::Development => "development",
274            PolicyPresetName::Administration => "administration",
275        }
276    }
277}
278
279pub fn policy_from_preset(preset: PolicyPresetName) -> Policy {
280    match preset {
281        PolicyPresetName::ReadOnly => Policy {
282            read: PolicyAction::Allow,
283            write: PolicyAction::Deny,
284            ddl: PolicyAction::Deny,
285            admin: PolicyAction::Deny,
286            tx_ctrl: PolicyAction::Allow,
287            row_cap: 1000,
288            stmt_timeout_ms: 10_000,
289            require_touch_id: false,
290            max_backup_rows: 10_000,
291            max_backup_bytes: 50 * 1024 * 1024,
292            on_backup_overflow: BackupOverflow::Abort,
293        },
294        PolicyPresetName::Development => Policy {
295            read: PolicyAction::Allow,
296            write: PolicyAction::Confirm,
297            ddl: PolicyAction::Confirm,
298            admin: PolicyAction::Deny,
299            tx_ctrl: PolicyAction::Allow,
300            row_cap: 5000,
301            stmt_timeout_ms: 30_000,
302            require_touch_id: false,
303            max_backup_rows: 10_000,
304            max_backup_bytes: 50 * 1024 * 1024,
305            on_backup_overflow: BackupOverflow::Abort,
306        },
307        PolicyPresetName::Administration => Policy {
308            read: PolicyAction::Allow,
309            write: PolicyAction::Confirm,
310            ddl: PolicyAction::Confirm,
311            admin: PolicyAction::Confirm,
312            tx_ctrl: PolicyAction::Allow,
313            row_cap: 5000,
314            stmt_timeout_ms: 60_000,
315            require_touch_id: true,
316            max_backup_rows: 10_000,
317            max_backup_bytes: 50 * 1024 * 1024,
318            on_backup_overflow: BackupOverflow::Abort,
319        },
320    }
321}
322
323pub const POLICY_PRESET_NAMES: [&str; 3] = ["read-only", "development", "administration"];
324
325/// Retention defaults per category (legacy `DEFAULT_RETENTION_BY_CATEGORY`).
326pub fn default_retention_days(category: SqlCategory) -> u32 {
327    match category {
328        SqlCategory::Read => 7,
329        SqlCategory::Write => 30,
330        SqlCategory::Ddl => 90,
331        SqlCategory::Admin => 180,
332        SqlCategory::TxCtrl => 7,
333    }
334}
335
336#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
337#[serde(rename_all = "camelCase", default)]
338pub struct RetentionByCategory {
339    pub read: u32,
340    pub write: u32,
341    pub ddl: u32,
342    pub admin: u32,
343    pub tx_ctrl: u32,
344}
345
346impl RetentionByCategory {
347    pub fn get(&self, category: SqlCategory) -> u32 {
348        match category {
349            SqlCategory::Read => self.read,
350            SqlCategory::Write => self.write,
351            SqlCategory::Ddl => self.ddl,
352            SqlCategory::Admin => self.admin,
353            SqlCategory::TxCtrl => self.tx_ctrl,
354        }
355    }
356}
357
358impl Default for RetentionByCategory {
359    fn default() -> Self {
360        Self {
361            read: 7,
362            write: 30,
363            ddl: 90,
364            admin: 180,
365            tx_ctrl: 7,
366        }
367    }
368}
369
370#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
371#[serde(rename_all = "camelCase", default)]
372pub struct RetentionConfig {
373    pub retention_days_by_category: RetentionByCategory,
374    pub backup_days: u32,
375    pub audit_max_mb: u32,
376    pub backup_max_mb: u32,
377    pub auto_cleanup_hours: u32,
378    pub redact_sql_in_log: bool,
379    pub tamper_evident_chain: bool,
380}
381
382impl Default for RetentionConfig {
383    fn default() -> Self {
384        Self {
385            retention_days_by_category: RetentionByCategory::default(),
386            backup_days: 30,
387            audit_max_mb: 500,
388            backup_max_mb: 1000,
389            auto_cleanup_hours: 24,
390            redact_sql_in_log: false,
391            tamper_evident_chain: false,
392        }
393    }
394}
395
396/// A qualified table target: `<database>.<table>`.
397#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
398pub struct TableId {
399    pub database: String,
400    pub table: String,
401}
402
403impl TableId {
404    pub fn new(database: impl Into<String>, table: impl Into<String>) -> Self {
405        Self {
406            database: database.into(),
407            table: table.into(),
408        }
409    }
410}
411
412/// v2 layer-2 rule key. `Wildcard(db)` represents `db.*` (legacy per-database
413/// override); `Exact(TableId)` represents an exact table rule. Serializes
414/// as its rendered string so JSON config maps stay string-keyed.
415#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash)]
416pub enum TableRuleKey {
417    Exact(TableId),
418    Wildcard { database: String },
419}
420
421impl serde::Serialize for TableRuleKey {
422    fn serialize<S: serde::Serializer>(&self, s: S) -> Result<S::Ok, S::Error> {
423        s.serialize_str(&self.render())
424    }
425}
426
427impl<'de> serde::Deserialize<'de> for TableRuleKey {
428    fn deserialize<D: serde::Deserializer<'de>>(d: D) -> Result<Self, D::Error> {
429        let s = String::deserialize(d)?;
430        TableRuleKey::parse(&s)
431            .ok_or_else(|| serde::de::Error::custom(format!("invalid table rule key {s:?}")))
432    }
433}
434
435impl TableRuleKey {
436    pub fn parse(s: &str) -> Option<Self> {
437        let (db, table) = s.split_once('.')?;
438        if db.is_empty() {
439            return None;
440        }
441        if table == "*" {
442            Some(TableRuleKey::Wildcard {
443                database: db.to_string(),
444            })
445        } else if table.is_empty() {
446            None
447        } else {
448            Some(TableRuleKey::Exact(TableId::new(db, table)))
449        }
450    }
451
452    pub fn render(&self) -> String {
453        match self {
454            TableRuleKey::Exact(t) => format!("{}.{}", t.database, t.table),
455            TableRuleKey::Wildcard { database } => format!("{database}.*"),
456        }
457    }
458
459    /// Does this rule govern `table`? Exact rules match one table; wildcard
460    /// rules match every table in the database.
461    pub fn governs(&self, table: &TableId) -> bool {
462        match self {
463            TableRuleKey::Exact(t) => t == table,
464            TableRuleKey::Wildcard { database } => database == &table.database,
465        }
466    }
467}
468
469pub type TablePolicies = BTreeMap<TableRuleKey, PartialPolicy>;
470
471#[cfg(test)]
472mod tests {
473    use super::*;
474
475    #[test]
476    fn defaults_match_legacy_schema() {
477        let p = Policy::default();
478        assert_eq!(p.read, PolicyAction::Allow);
479        assert_eq!(p.write, PolicyAction::Confirm);
480        assert_eq!(p.ddl, PolicyAction::Deny);
481        assert_eq!(p.admin, PolicyAction::Deny);
482        assert_eq!(p.tx_ctrl, PolicyAction::Allow);
483        assert_eq!(p.row_cap, 1000);
484        assert_eq!(p.stmt_timeout_ms, 10_000);
485        assert_eq!(p.max_backup_rows, 10_000);
486        assert_eq!(p.max_backup_bytes, 52_428_800);
487    }
488
489    #[test]
490    fn presets_parse_legacy_and_new_names() {
491        assert_eq!(
492            PolicyPresetName::parse("dev"),
493            Some(PolicyPresetName::Development)
494        );
495        assert_eq!(
496            PolicyPresetName::parse("admin"),
497            Some(PolicyPresetName::Administration)
498        );
499        assert_eq!(
500            policy_from_preset(PolicyPresetName::ReadOnly).write,
501            PolicyAction::Deny
502        );
503        assert!(policy_from_preset(PolicyPresetName::Administration).require_touch_id);
504    }
505
506    #[test]
507    fn table_rule_keys_parse_and_match() {
508        let exact = TableRuleKey::parse("app.jobs").unwrap();
509        let wild = TableRuleKey::parse("app.*").unwrap();
510        let jobs = TableId::new("app", "jobs");
511        assert!(exact.governs(&jobs));
512        assert!(wild.governs(&jobs));
513        assert!(!wild.governs(&TableId::new("analytics", "jobs")));
514        assert_eq!(exact.render(), "app.jobs");
515        assert_eq!(wild.render(), "app.*");
516        assert!(TableRuleKey::parse("nodot").is_none());
517        assert!(TableRuleKey::parse(".x").is_none());
518    }
519
520    #[test]
521    fn merge_matches_legacy_spread() {
522        let base = Policy::default();
523        let over = PartialPolicy {
524            write: Some(PolicyAction::Deny),
525            row_cap: Some(42),
526            ..PartialPolicy::default()
527        };
528        let merged = base.merged_with(&over);
529        assert_eq!(merged.write, PolicyAction::Deny);
530        assert_eq!(merged.row_cap, 42);
531        assert_eq!(merged.read, PolicyAction::Allow);
532    }
533}